:cross-attention分数提取与KL散度蒸馏损失的实现细节)
FiD源码深潜下cross-attention分数提取与KL散度蒸馏损失的实现细节【免费下载链接】FiDFusion-in-Decoder项目地址: https://gitcode.com/gh_mirrors/fi/FiD在开源问答项目FiDFusion-in-Decoder中cross-attention 分数提取与KL 散度蒸馏损失是把阅读器reader里蕴含的段落相关性知识蒸馏给轻量检索器retriever的两块基石。这篇源码深潜文章带你读懂这两个核心机制的实现细节分数从哪里来、如何聚合成一个标量、蒸馏损失又是怎么算的——面向新手少代码、讲原理。一、整体思路为什么用 cross-attention 分数当软标签传统检索器训练依赖人工标注的问题-段落正负样本代价高且信号稀疏。FiD 的巧思是FiD 阅读器在生成答案时decoder 的每个 token 都会通过cross-attention去看输入的 100 个段落——注意力落在哪个段落上就隐式表达了该段落与问题的相关程度只需取第一个解码步首个生成 token的注意力分数对注意力头、网络层、段落内有效 token 三者取平均就能把每个段落压缩成一个相关性标量这组标量经 softmax 后构成教师分布用KL 散度驱动双塔检索器去模仿它——这就是从 reader 蒸馏知识到 retriever的全部魔法。整个蒸馏流程只有 4 步对应 4 个脚本步骤脚本产出① 提取 cross-attention 分数test_reader.pydataset_wscores.json每段一个 score② KL 散度蒸馏训练检索器train_retriever.py双塔 Retriever 模型③ 索引知识库段落generate_passage_embeddings.py段落向量库④ 按问题检索段落passage_retrieval.py检索结果二、cross-attention 分数提取的实现细节分数提取的代码集中在 src/model.py配合入口脚本 test_reader.py。1. 猴子补丁给 cross-attention 装一个记录器FiD 的FiDT5类提供两个辅助方法src/model.py 第 87-127 行reset_score_storage()把每个 decoder block 中EncDecAttention的score_storage置空为下一个样本重新记录做准备overwrite_forward_crossattention()用types.MethodType把每个 decoder block 的 cross-attention 前向函数替换成自定义的cross_attention_forward。被替换后的前向函数第 194-254 行里最关键只有 5 行scores torch.einsum(bnqd,bnkd-bnqk, q, k) # pre-softmax 注意力分数 scores mask # padding 掩码 scores position_bias # T5 相对位置偏置 if self.score_storage is None: # 只在第一个解码步记录 self.score_storage scores两个设计要点值得新手注意记录的是 softmax 之前的原始分数含 mask 和位置偏置这才是反映相关性的连续量只在第一个解码步记录generate是逐 token 自回归的score_storage首次写入后就不再覆盖。由于首个解码步的 query 长度只有 1所以这一步几乎没有额外计算开销后续解码步完全不受影响。2. 聚合从 5 维张量到一个标量get_crossattention_scoressrc/model.py 第 95-118 行负责把每层的分数块拼起来、按段落归约scores torch.cat(scores, dim2) # (bsz, 头数, 层数, klen) scores scores.view(bsz, n_heads, n_layers, n_passages, -1) scores scores.masked_fill(~context_mask[:, None, None], 0.) # 屏蔽 padding scores scores.sum(dim[1, 2, 4]) # 对头、层、token 求和 ntokens context_mask.sum(dim[2]) * n_layers * n_heads scores scores / ntokens # 除以有效token数 → (bsz, n_passages)这里的masked_fill保证 padding 位置不参与平均除以ntokens得到的本质是一个三重平均平均所有注意力头、所有 decoder 层、段落内所有有效 token 的 cross-attention 分数。最终每个问题得到一个(n_passages,)的分数向量——这就是软标签的原始形态。3. 落盘test_reader.py 的主循环test_reader.py 第 27-62 行把上述能力串了起来开启--write_crossattention_scores后先执行猴子补丁并清空存储每个 batch 先reset_score_storage()再调用model.generate(max_length50)正常生成答案生成结束后调用get_crossattention_scores把每个段落的分数逐条写回数据example[ctxs][j][score] ...最后由save_distributed_datasetsrc/util.py 第 187-209 行把多张 GPU 各自的分片合并成一份dataset_wscores.json。可以看到提取分数不改变任何推理行为只是在生成答案的同时顺手记录了一下注意力。三、KL 散度蒸馏损失Retriever 的训练信号1. 双塔结构问题塔与段落塔Retrieversrc/model.py 第 276-331 行是一个基于 BERT 的双塔dual-encoder模型问题与每个段落各自独立过 BERT经embed_text第 333-353 行做掩码平均池化得到定长向量相似度用缩放点积计算第 320-325 行einsum(bd,bid-bi)再除以sqrt(d)——与 self-attention 中的缩放方式一致一个容易忽略的细节train_retriever.py 第 206-208 行把投影层硬编码为 768→256 维 LayerNorm最终入索引的向量只有 256 维兼顾精度与检索速度。2. KL 散度损失的三行核心代码损失函数只有三行src/model.py 第 355-358 行def kldivloss(self, score, gold_score): gold_score torch.softmax(gold_score, dim-1) # 教师分布 score torch.nn.functional.log_softmax(score, dim-1) # 学生 log 概率 return self.loss_fct(score, gold_score) # KLDivLoss新手需要理解的三个关键点教师分布 cross-attention 分数原始实数做 softmax学生分布 检索器打分的 log_softmax。KL 散度衡量两者分布差异训练目标就是让学生分布逼近教师分布PyTorch 的KLDivLoss约定 input 必须是 log 概率、target 是概率代码正好满足这个约定默认reductionmean没有温度系数T1这是最直接的模仿蒸馏实现简单效果已被验证。3. gold_score 从哪来数据侧的两个细节数据加载src/data.py 第 137-139 行若某段落没有score字段默认填1.0 / (k1)即用排名越靠前分数越高的位置先验兜底——这意味着没有 reader 打分时也能退化成用初始检索排序做蒸馏分数在数据侧始终保持原始实数RetrieverCollator第 170-171 行直接 stacksoftmax 统一放在损失函数内部完成避免归一化逻辑散落各处。训练主循环train_retriever.py 第 54-63 行则非常简单model(..., gold_score...)的第 4 个返回值就是 KL 散度损失直接backward()。四、训练中怎么看效果三个评估指标每 500 步训练会打印一组检索质量指标实现见 src/evaluation.py 第 148-175 行指标含义好/坏inversions预测排序相对 gold 排序的逆序对数越小越好avg top-k预测 top-k 中有多少比例落在 gold top-k 内越大越好idx top-k要覆盖全部 gold top-k 段落需要取到预测结果的第几位越小越好这三项指标本质上都在回答同一个问题检索器学到的相似度是否逼近了 reader 的 cross-attention 相关性判断。推荐超参数参考 README 官方配置--lr 1e-4 --optim adamw --scheduler linear --total_steps 20000 --scheduler_steps 30000配合--n_context 100。五、快速上手4 条命令跑通蒸馏全流程克隆仓库并准备数据后数据与预训练模型可用仓库自带的get-data.sh、get-model.sh下载按顺序执行git clone https://gitcode.com/gh_mirrors/fi/FiD # ① 提取 cross-attention 分数 → dataset_wscores.json python test_reader.py --model_path reader模型 --eval_data data.json \ --per_gpu_batch_size 4 --n_context 100 --name my_test \ --checkpoint_dir checkpoint --write_crossattention_scores # ② KL 散度蒸馏训练检索器 python train_retriever.py --lr 1e-4 --optim adamw --scheduler linear \ --train_data train_data.json --eval_data eval_data.json \ --n_context 100 --total_steps 20000 --scheduler_steps 30000 # ③ 为知识库如维基百科生成段落向量 python generate_passage_embeddings.py --model_path retriever目录 \ --passages passages.tsv --output_path wikipedia_embeddings \ --per_gpu_batch_size 500 # ④ 给定问题高效检索段落 python passage_retrieval.py --model_path retriever目录 \ --passages psgs_w100.tsv --data_path data.json \ --passages_embeddings wikipedia_embeddings/wiki_* \ --output_path retrieved_data.json --n-docs 100README 中还提到一个实用技巧把上面 4 步迭代多轮用新检索器检出的段落重新打分、再训练效果会进一步提升。官方报告的蒸馏检索器成绩NaturalQuestions R5 达 73.8、R20 达 84.3TriviaQA R5 达 77.0。六、小结回顾一下这次深潜的收获分数提取靠猴子补丁在 decoder 的 cross-attention 里偷拍第一个解码步的 pre-attention 分数再对头、层、token 做掩码平均零推理开销地拿到每个段落的相关性标量KL 蒸馏教师分数 softmax 后与学生打分的 log_softmax 一起喂给KLDivLoss三行代码完成分布对齐让轻量双塔检索器继承 reader 的判断力工程细节缺失分数用位置先验兜底、投影层硬编码 256 维、逆序对与 top-k 指标监控——这些正是开源实现里值得学习的实践。至此FiD 从阅读器到检索器的知识蒸馏链路就完整了。想继续深挖EncoderWrapper的decoder 内融合机制与显存优化技巧可阅读本系列上篇或 src/model.py 全文。【免费下载链接】FiDFusion-in-Decoder项目地址: https://gitcode.com/gh_mirrors/fi/FiD创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考