计算困惑度的三大挑战:我的应对方案是否合理?
BART模型困惑度计算的三类问题解决方案
问题1:BART用中心窗口计算困惑度是否合理?
BART的核心结构是双向编码器+单向解码器,需分场景判断:
- 如果目标是衡量模型对文本的双向语义理解能力(比如筛选适配特定文本的模型),用掩码token为中心的窗口完全合理——编码器本身就是双向注意力结构,能同时利用掩码位的前后上下文,更贴合模型的双向特性。
- 如果是针对生成类任务(比如摘要、续写)的困惑度评估,还是要遵循解码器的单向逻辑:用掩码token之前的上下文预测该token,因为生成任务中模型是基于前文输出后续内容的。
问题2:大窗口计算困惑度时概率下溢的解决方案
直接计算概率乘积会因数值过小导致下溢,核心解决思路是切换到对数空间计算,具体操作:
- 困惑度的数学表达式为:
PPL = exp( (-1/N) * sum(log(p_i)) ),其中p_i是第i个token的预测概率,N是总token数。 - 实现时避免直接乘概率,转而累加每个token的负对数概率:
import torch # probs: 模型输出的每个token预测概率,shape为[batch_size, seq_len, vocab_size] # target_ids: 真实token索引,shape为[batch_size, seq_len] log_probs = torch.log(torch.gather(probs, dim=2, index=target_ids.unsqueeze(-1)).squeeze(-1)) total_neg_log_prob = -log_probs.sum() avg_neg_log_prob = total_neg_log_prob / target_ids.numel() ppl = torch.exp(avg_neg_log_prob) - 这种方式全程在对数空间运算,不会出现下溢问题,完全支持1024甚至更长的窗口。
问题3:超长文本采样片段的最优聚合方式
不要直接平均各片段的困惑度,优先选择基于总负对数概率的聚合方式:
- 直接平均困惑度的问题:困惑度是指数形式,短片段的困惑度波动会被放大,长片段的权重被低估,导致结果失真。
- 正确步骤:
- 对每个采样片段,计算其负对数概率之和(不要转成困惑度)。
- 将所有片段的负对数概率之和相加,除以所有片段的总token数。
- 最后对这个平均值取指数,得到整体困惑度。
- 采样建议:随机均匀选取文本的不同位置(开头、中间、结尾都要覆盖),避免采样偏差;如果文本有明显结构划分(比如章节),可按结构分层采样,提升结果代表性。
内容的提问来源于stack exchange,提问作者Agnes
相关产品推荐
相关产品推荐

