第 0 章 项目导览与学习路线


文档摘要

第 0 章 项目导览与学习路线 在动任何代码之前,先建立全局心智模型。本章回答三个问题:这是什么?为什么学它?怎么学? 0.1 知识蒸馏解决什么问题 大模型(LLM)能力强大,但代价昂贵:参数动辄数十亿至上千亿,推理时占用大量显存与算力,难以部署到手机、边缘设备等资源受限的环境。 知识蒸馏(Knowledge Distillation) 正是为这个问题而生。它的核心思想可以用一句话概括: 让一个「小而快」的学生模型,去模仿一个「大而强」的教师模型的输出行为,从而在参数量大幅缩减的同时,尽可能保住教师的能力。 用一个生活化的比喻:教师是经验丰富的老教授,学生是年轻助教。学生不必从零摸索(只啃课本/硬标签),而是通过观察老教授如何判断问题、如何给出选项的概率(软标签),更快地积累「直觉」。

第 0 章 项目导览与学习路线

在动任何代码之前,先建立全局心智模型。本章回答三个问题:这是什么?为什么学它?怎么学?

0.1 知识蒸馏解决什么问题

大模型(LLM)能力强大,但代价昂贵:参数动辄数十亿至上千亿,推理时占用大量显存与算力,难以部署到手机、边缘设备等资源受限的环境。

知识蒸馏(Knowledge Distillation) 正是为这个问题而生。它的核心思想可以用一句话概括:

让一个「小而快」的学生模型,去模仿一个「大而强」的教师模型的输出行为,从而在参数量大幅缩减的同时,尽可能保住教师的能力。

用一个生活化的比喻:教师是经验丰富的老教授,学生是年轻助教。学生不必从零摸索(只啃课本/硬标签),而是通过观察老教授如何判断问题、如何给出选项的概率(软标签),更快地积累「直觉」。

本教程实现的,正是把一个标准预训练 GPT2(教师)的能力,蒸馏到一个只有其约 1/12 参数量的迷你 GPT2(学生)中。

0.2 技术栈一览

类别 选型 为什么选它
深度学习框架 PyTorch ≥ 2.0 工业标准,动态图调试友好
模型实现 HuggingFace transformersGPT2LMHeadModel 工业级稳定、规模可配、生态完善
分词器 tiktoken(p50k_base OpenAI 出品,Rust 实现,性能强
数据集 tiny_shakespeare(Karpathy 经典) 体量小(约 1MB),单机几分钟可训完
优化器 AdamW + 自定义余弦退火 GPT-2/3 标准配方
蒸馏方法 经典 Logits 蒸馏(Hinton 2015) 原理清晰、效果稳定、最适合入门

0.3 全局数据流

理解一个项目,最重要的是抓住「数据流」。本项目的数据流如下:

几个关键点要记住:

  • 教师是「顾问」:它被完全冻结(requires_grad=False),只在前向时提供「软标签」,不参与梯度更新。
  • 学生是「主角」:它从零初始化,是唯一被训练的对象。
  • 损失是「组合拳」:硬标签(真实下一个 token)的交叉熵 + 软标签(教师分布)的 KL 散度,二者加权求和。

0.4 你能学到什么

读完这套教程并跑通代码,你将掌握:

能力 对应章节
理解知识蒸馏的数学原理(温度、KL、暗知识) 第 1 章
搭建可复现的蒸馏实验环境 第 2 章
dataclass 管理蒸馏/训练全超参 第 3 章
构造带索引的自回归样本,支持教师缓存 第 4 章
区分并构建「预训练教师」与「从零学生」 第 5 章
手写蒸馏损失:KL + CE + shift 对齐 第 6 章
写一个带损失分解日志的训练循环 第 7 章
用困惑度、Top-1、分布 KL 评估蒸馏质量 第 8 章
实现带温度/top-k 的自回归推理 第 9 章
掌握教师缓存、调参等进阶技巧 第 10 章

0.5 逻辑模块地图

本项目按职责拆分为若干逻辑模块,每个模块负责流水线中的一个环节:

逻辑模块 职责 对应章节
核心配置模块 用数据类统管模型/蒸馏/训练三类超参 第 3 章
数据处理模块 下载文本、分词编码、切窗、构造样本 第 4 章
模型构建模块 加载预训练教师、从零初始化学生 第 5 章
蒸馏损失模块 计算软硬标签组合损失 第 6 章
训练主程序 整合所有模块,跑训练循环 第 7 章
评估对比模块 衡量学生与教师的性能差距 第 8 章
推理生成模块 用学生模型做文本生成 第 9 章

各模块的依赖关系是单向的:数据 → 模型 → 损失 → 训练 → 评估/推理。建议按这个顺序阅读。

0.6 推荐学习路线

根据你的背景,可以选择不同路线:

路线 A:零基础系统学习(推荐)

按章节顺序 00 → 01 → 02 → ... → 10 逐章精读,每章跑通代码再进入下一章。预计耗时 1-2 周。

路线 B:有深度学习基础,只学蒸馏新知识

直接跳到第 1 章(原理)和第 6 章(损失函数),这是本项目的灵魂;再按需查阅第 7、8、10 章。

路线 C:只想快速跑起来

读第 2 章(环境准备),按命令跑通训练;遇到问题再回查对应章节或附录 C。

0.7 本项目的局限与延伸

为保持入门友好,本项目做了一些简化,同时也为进阶留了接口:

  • 蒸馏类型:只实现最经典的 Logits 蒸馏。进阶可扩展到中间层蒸馏(隐藏状态对齐、注意力对齐),这部分在《第 10 章》讨论。
  • 教师来源:教师固定为预训练 GPT2。若教师是黑盒 API(无法取 logits),则需改用「序列级蒸馏」(教师生成伪数据),见第 10 章。
  • 任务类型:聚焦文本生成(LM)。分类任务的蒸馏思想类似,但损失实现略有不同。

0.8 约定与前置知识

阅读本教程需要以下前置知识:

  • Python:熟悉函数、类、类型注解。
  • PyTorch:了解 tensorautogradModuleDataLoader
  • Transformer 基础:知道自注意力、自回归生成的大致概念即可(本项目不手写 Transformer,复用现成实现)。
  • 概率论基础:了解概率分布、交叉熵即可,KL 散度会在第 1、6 章详细推导。

代码示例中的注释默认为中文,与项目源码风格保持一致。

本章小结

  • 知识蒸馏 = 小模型模仿大模型的输出行为,实现压缩与加速。
  • 本项目把预训练 GPT2(教师)蒸馏到迷你 GPT2(学生),压缩比约 12:1。
  • 数据流:文本 → 分词 → 样本 → 教师/学生前向 → 组合损失 → 训练学生。
  • 模块按职责拆分,单向依赖。

下一站:在《第 1 章 知识蒸馏原理详解》中,我们将深入 Hinton 的奠基论文,彻底搞懂温度软化、暗知识、KL 散度这些核心概念背后的数学。


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