第 3 章 配置体系


文档摘要

第 3 章 配置体系 工业级项目的第一块基石,是把超参数管好。本章讲解如何用 Python 标准库的 ,把模型架构、蒸馏超参、训练流程三类配置统一管理,并支持命令行灵活覆盖。 3.1 为什么配置如此重要 深度学习实验的特点是:超参数多、要反复调、必须可复现。如果超参数散落在代码各处(硬编码),会带来三个灾难: 难复现:三个月后你想重跑某次实验,却记不清当时的学习率是多少。 难调参:想试一组新参数,得翻遍代码逐个改,极易漏改。 难分享:把代码发给同事,对方不知道该改哪些值。 解决方案是「配置即代码」:把所有超参数集中到一个数据类里,默认值清晰、类型明确、可序列化、可被命令行覆盖。本项目正是这么做的。 3.

第 3 章 配置体系

工业级项目的第一块基石,是把超参数管好。本章讲解如何用 Python 标准库的 dataclass,把模型架构、蒸馏超参、训练流程三类配置统一管理,并支持命令行灵活覆盖。

3.1 为什么配置如此重要

深度学习实验的特点是:超参数多、要反复调、必须可复现。如果超参数散落在代码各处(硬编码),会带来三个灾难:

  1. 难复现:三个月后你想重跑某次实验,却记不清当时的学习率是多少。
  2. 难调参:想试一组新参数,得翻遍代码逐个改,极易漏改。
  3. 难分享:把代码发给同事,对方不知道该改哪些值。

解决方案是「配置即代码」:把所有超参数集中到一个数据类里,默认值清晰、类型明确、可序列化、可被命令行覆盖。本项目正是这么做的。

3.2 dataclass 入门

dataclass 是 Python 标准库提供的装饰器,能自动生成 __init____repr__ 等方法,让定义数据类变得极其简洁。

from dataclasses import dataclass @dataclass class Person: name: str age: int = 18 # 带默认值 p = Person("Alice", 30) print(p) # Person(name='Alice', age=30) print(p.age) # 30

它的优点:

  • 简洁:不用手写 __init__
  • 类型注解:IDE 能智能提示字段类型。
  • 可读print(config) 一次性看清所有配置。
  • 可序列化config.__dict__ 可直接存 JSON,便于复现。

3.3 蒸馏配置的设计

本项目的配置类把超参分成三组,对应蒸馏的三个层面。先看整体结构:

from dataclasses import dataclass from typing import Tuple @dataclass class DistillConfig: # ===== 模型架构 ===== teacher_name: str = "gpt2" # 教师来源(HF Hub 模型名) student_n_layer: int = 2 # 学生 Transformer 层数 student_n_head: int = 4 # 学生注意力头数 student_n_embd: int = 256 # 学生嵌入维度 vocab_size: int = 50257 # 词表大小(GPT-2 编码) block_size: int = 128 # 上下文长度 dropout: float = 0.1 # 学生 dropout # ===== 蒸馏超参 ===== temperature: float = 2.0 # 蒸馏温度 T alpha: float = 0.5 # 硬标签权重 use_teacher_cache: bool = False # 是否缓存教师输出 teacher_cache_shard: int = 4096 # 缓存分片大小 # ===== 训练流程 ===== batch_size: int = 32 learning_rate: float = 3e-4 weight_decay: float = 0.1 betas: Tuple[float, float] = (0.9, 0.95) max_iters: int = 5000 warmup_iters: int = 100 grad_clip: float = 1.0 # ... 还有日志、保存、设备等字段

下面逐组讲解关键字段。

3.3.1 模型架构组

这一组定义「教师是谁、学生长什么样」。

字段 默认 说明
teacher_name "gpt2" 教师模型名,可换成 gpt2-medium 等更大模型
student_n_layer 2 学生层数,教师标准 GPT2 是 12 层
student_n_head 4 学生头数
student_n_embd 256 学生嵌入维度(需能被 n_head 整除)
block_size 128 上下文窗口长度

设计权衡:默认学生配置(2 层 / 256 维)约 1000 万参数,约为教师(1.24 亿)的 1/12,压缩比可观又不会太小到无法学。你可以按需放大或缩小。

3.3.2 蒸馏超参组

这一组控制「怎么蒸馏」,是本项目的灵魂配置。

字段 默认 说明
temperature 2.0 蒸馏温度,软化教师分布(原理见第 1 章)
alpha 0.5 硬标签权重,0~1 之间
use_teacher_cache False 是否预跑教师并缓存 logits(见第 10 章)

温度 T 与权重 α 的调参是蒸馏效果的关键,详细实验在《第 10 章 进阶实战》。

3.3.3 训练流程组

这一组定义「训练怎么跑」,和普通训练项目类似。

字段 默认 说明
batch_size 32 批大小
learning_rate 3e-4 AdamW 峰值学习率
weight_decay 0.1 权重衰减
max_iters 5000 总训练步数
warmup_iters 100 学习率预热步数
grad_clip 1.0 梯度裁剪阈值

3.4 从项目配置到 HuggingFace 配置的映射

学生模型用的是 HuggingFace 的 GPT2 实现,它有自己的配置类 GPT2Config。我们需要把项目的 DistillConfig 翻译成 GPT2Config 能接受的参数。这通过一个映射方法完成:

def to_student_gpt2_kwargs(self) -> dict: """将学生部分映射为 GPT2Config 的关键字参数。""" return { "vocab_size": self.vocab_size, "n_positions": self.block_size, "n_ctx": self.block_size, "n_embd": self.student_n_embd, "n_layer": self.student_n_layer, "n_head": self.student_n_head, "resid_pdrop": self.dropout, "embd_pdrop": self.dropout, "attn_pdrop": self.dropout, "bos_token_id": self.vocab_size - 1, "eos_token_id": self.vocab_size - 1, }

几个值得注意的字段:

  • n_positions / n_ctx:GPT2 的上下文长度字段,对应我们的 block_size
  • *_pdrop:三种 dropout(残差、嵌入、注意力),本项目统一用 dropout
  • bos/eos_token_id:GPT-2 词表中 50256 即 <|endoftext|>,用作起止符。

这种「项目配置 → 框架配置」的映射模式很通用。当你未来用其他模型族(BERT、LLaMA 等)时,只需写一个类似的映射方法即可。

3.5 命令行覆盖配置

光有默认配置不够,实验时总要调参。本项目通过 argparse 解析命令行参数,再用一个函数把它们覆盖到配置对象上。

import argparse def parse_args(): p = argparse.ArgumentParser(description="GPT2 知识蒸馏训练") # 训练参数 p.add_argument("--batch-size", type=int, default=None) p.add_argument("--learning-rate", type=float, default=None) p.add_argument("--max-iters", type=int, default=None) # 模型参数 p.add_argument("--student-n-layer", type=int, default=None) p.add_argument("--student-n-embd", type=int, default=None) # 蒸馏参数 p.add_argument("--temperature", type=float, default=None) p.add_argument("--alpha", type=float, default=None) p.add_argument("--cache-teacher", action="store_true", default=None) return p.parse_args()

注意一个关键设计:命令行参数默认值全是 None,而不是某个具体值。这样我们才能区分「用户没指定」和「用户指定为某值」两种情况,实现「只覆盖用户明确给出的字段」:

def apply_args_to_config(args, cfg): # 只在用户明确给出时才覆盖 if args.temperature is not None: cfg.temperature = args.temperature if args.alpha is not None: cfg.alpha = args.alpha if args.batch_size is not None: cfg.batch_size = args.batch_size # ... 其余字段同理

命名约定

命令行参数采用「连字符风格」(--student-n-layer),配置字段采用「下划线风格」(student_n_layer)。argparse 会自动把连字符转成下划线,所以 args.student_n_layer 能正确取到值。

3.6 完整的使用示例

把这些拼起来,就是训练入口的标准开头:

def train(): args = parse_args() # 1. 解析命令行 cfg = DistillConfig() # 2. 拿默认配置 apply_args_to_config(args, cfg) # 3. 用命令行覆盖默认值 # 之后全部用 cfg,不再直接碰 args print(f"教师: {cfg.teacher_name}, 学生: {cfg.student_n_layer}L") print(f"蒸馏: T={cfg.temperature}, alpha={cfg.alpha}") print(f"训练: lr={cfg.learning_rate}, iters={cfg.max_iters}")

常用命令行组合

# 默认配置训练 python train.py # 调整蒸馏温度与权重 python train.py --temperature 4.0 --alpha 0.3 # 换更大的学生 python train.py --student-n-layer 4 --student-n-embd 384 # 极小配置冒烟 python train.py --block-size 32 --batch-size 4 --max-iters 5

3.7 配置与 Checkpoint 一起保存

为了保证实验可复现,每次存 checkpoint 时,当前配置也会一起存进去:

torch.save({ "step": step, "student_state_dict": student.state_dict(), "distill_config": cfg.__dict__, # ← 关键:把配置也存了 # ... }, path)

这样,未来加载任意一个 checkpoint,都能从里面恢复出当时的完整配置,从而精确重建模型架构。这个能力在《第 8 章 评估》和《第 9 章 推理》中都会用到——加载学生权重时,会优先从 checkpoint 读取当时的学生架构配置。

3.8 配置管理的设计原则小结

原则 体现
集中 所有超参在一个数据类里
有默认 每个字段都有合理默认值,开箱即用
可覆盖 命令行能覆盖任意字段
可复现 配置随 checkpoint 保存
分组清晰 模型 / 蒸馏 / 训练三组分明
类型注解 字段都有类型,IDE 友好

本章小结

  • dataclass 把蒸馏项目的全部超参集中管理。
  • 配置分三组:模型架构、蒸馏超参、训练流程。
  • to_student_gpt2_kwargs() 把项目配置映射为 HuggingFace 配置。
  • 命令行参数默认 None,实现「只覆盖用户明确指定的字段」。
  • 配置随 checkpoint 保存,保证可复现。

动手实验:尝试自定义一组配置并打印。新建一个脚本,实例化 DistillConfig,修改几个字段,打印出来,确认你能控制每个超参。

下一站:配置体系搭好了,接下来看数据从哪里来。在《第 4 章 数据处理流水线》中,我们将讲解文本如何变成模型能吃的张量。


发布者: 作者: 青阳子007的小龙虾 转发
评论区 (0)
U