第 3 章 配置体系 工业级项目的第一块基石,是把超参数管好。本章讲解如何用 Python 标准库的 ,把模型架构、蒸馏超参、训练流程三类配置统一管理,并支持命令行灵活覆盖。 3.1 为什么配置如此重要 深度学习实验的特点是:超参数多、要反复调、必须可复现。如果超参数散落在代码各处(硬编码),会带来三个灾难: 难复现:三个月后你想重跑某次实验,却记不清当时的学习率是多少。 难调参:想试一组新参数,得翻遍代码逐个改,极易漏改。 难分享:把代码发给同事,对方不知道该改哪些值。 解决方案是「配置即代码」:把所有超参数集中到一个数据类里,默认值清晰、类型明确、可序列化、可被命令行覆盖。本项目正是这么做的。 3.
工业级项目的第一块基石,是把超参数管好。本章讲解如何用 Python 标准库的
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__。print(config) 一次性看清所有配置。config.__dict__ 可直接存 JSON,便于复现。本项目的配置类把超参分成三组,对应蒸馏的三个层面。先看整体结构:
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 # ... 还有日志、保存、设备等字段
下面逐组讲解关键字段。
这一组定义「教师是谁、学生长什么样」。
| 字段 | 默认 | 说明 |
|---|---|---|
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,压缩比可观又不会太小到无法学。你可以按需放大或缩小。
这一组控制「怎么蒸馏」,是本项目的灵魂配置。
| 字段 | 默认 | 说明 |
|---|---|---|
temperature |
2.0 |
蒸馏温度,软化教师分布(原理见第 1 章) |
alpha |
0.5 |
硬标签权重,0~1 之间 |
use_teacher_cache |
False |
是否预跑教师并缓存 logits(见第 10 章) |
温度 T 与权重 α 的调参是蒸馏效果的关键,详细实验在《第 10 章 进阶实战》。
这一组定义「训练怎么跑」,和普通训练项目类似。
| 字段 | 默认 | 说明 |
|---|---|---|
batch_size |
32 |
批大小 |
learning_rate |
3e-4 |
AdamW 峰值学习率 |
weight_decay |
0.1 |
权重衰减 |
max_iters |
5000 |
总训练步数 |
warmup_iters |
100 |
学习率预热步数 |
grad_clip |
1.0 |
梯度裁剪阈值 |
学生模型用的是 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 等)时,只需写一个类似的映射方法即可。
光有默认配置不够,实验时总要调参。本项目通过 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 能正确取到值。
把这些拼起来,就是训练入口的标准开头:
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
为了保证实验可复现,每次存 checkpoint 时,当前配置也会一起存进去:
torch.save({ "step": step, "student_state_dict": student.state_dict(), "distill_config": cfg.__dict__, # ← 关键:把配置也存了 # ... }, path)
这样,未来加载任意一个 checkpoint,都能从里面恢复出当时的完整配置,从而精确重建模型架构。这个能力在《第 8 章 评估》和《第 9 章 推理》中都会用到——加载学生权重时,会优先从 checkpoint 读取当时的学生架构配置。
| 原则 | 体现 |
|---|---|
| 集中 | 所有超参在一个数据类里 |
| 有默认 | 每个字段都有合理默认值,开箱即用 |
| 可覆盖 | 命令行能覆盖任意字段 |
| 可复现 | 配置随 checkpoint 保存 |
| 分组清晰 | 模型 / 蒸馏 / 训练三组分明 |
| 类型注解 | 字段都有类型,IDE 友好 |
dataclass 把蒸馏项目的全部超参集中管理。to_student_gpt2_kwargs() 把项目配置映射为 HuggingFace 配置。None,实现「只覆盖用户明确指定的字段」。动手实验:尝试自定义一组配置并打印。新建一个脚本,实例化 DistillConfig,修改几个字段,打印出来,确认你能控制每个超参。
下一站:配置体系搭好了,接下来看数据从哪里来。在《第 4 章 数据处理流水线》中,我们将讲解文本如何变成模型能吃的张量。