[ PROMPT_NODE_22340 ]
minillm
[ SKILL_DOCUMENTATION ]
# MiniLLM:用于 LLM 蒸馏的反向 KL 散度
基于 arXiv 2306.08543 (2024) - MiniLLM: 大语言模型的知识蒸馏
## 概述
**来源**: https://arxiv.org/abs/2306.08543
**GitHub**: https://github.com/microsoft/LMOps/tree/main/minillm
MiniLLM 使用反向 KLD 替代前向 KLD 进行知识蒸馏,在生成式语言模型上取得了更好的性能。
## 标准 KLD 的问题
### 前向 KL 散度 (标准)
**公式**: `KL(Student || Teacher)`
**最小化行为**: 模式寻求 (Mode-seeking)
学生模型试图匹配教师模型的平均行为
→ 学生模型专注于教师模型概率最高的区域
→ 学生模型忽略了低概率但有效的生成结果
**生成模型的问题**: 限制了多样性,学生模型生成的输出安全但枯燥。
### 为什么前向 KL 不适用于生成任务
python
# 教师分布 (多样化)
teacher_probs = [0.3, 0.3, 0.2, 0.1, 0.1] # 多个有效选项
# 前向 KL 最小化
# 学生学习到: [0.6, 0.3, 0.1, 0.0, 0.0]
# 问题: 完全忽略了选项 4-5 (模式寻求)
## MiniLLM 解决方案:反向 KLD
### 反向 KL 散度
**公式**: `KL(Teacher || Student)`
**最小化行为**: 模式覆盖 (Mode-covering)
学生模型试图覆盖教师模型的所有模式
→ 学生模型学习到多样化的生成
→ 学生模型不会忽略任何有效的教师输出
### 数学公式
**前向 KL** (标准蒸馏):
L_forward = Σ p_student(x) log(p_student(x) / p_teacher(x))
= E_{x~student} [log p_student(x) - log p_teacher(x)]
**反向 KL** (MiniLLM):
L_reverse = Σ p_teacher(x) log(p_teacher(x) / p_student(x))
= E_{x~teacher} [log p_teacher(x) - log p_student(x)]
**关键区别**: 期望值是基于教师分布还是学生分布。
## 实现
### 反向 KLD 损失
python
import torch
import torch.nn.functional as F
def reverse_kl_loss(student_logits, teacher_logits, temperature=1.0):
"""
反向 KL 散度: KL(Teacher || Student)。
参数:
student_logits: 模型预测 (batch, seq_len, vocab_size)
teacher_logits: 教师预测 (batch, seq_len, vocab_size)
temperature: 平滑参数
返回:
反向 KL 散度损失
"""
# 教师分布 (目标,已分离)
p_teacher = F.softmax(teacher_logits / temperature, dim=-1)
p_teacher = p_teacher.detach() # 不通过教师模型进行反向传播
# Studen