尧图建网站 尧图建网站 YAOTU WEB BUILD 免费咨询
ARTICLE DETAIL

资讯详情

深耕网站建设与建站编程的一线实战洞察。

【Bug已解决】INF encountered when using sampling with temperature. 解决方案

【Bug已解决】INF encountered when using sampling with temperature. 解决方案 【Bug已解决】INF encountered when using sampling with temperature. 解决方案一、现象长什么样用 Transformers 做采样生成时一旦带上temperature偶尔会整批输出inf甚至直接崩在torch.multinomial上from transformers import AutoModelForCausalLM, AutoTokenizer tok AutoTokenizer.from_pretrained(gpt2) model AutoModelForCausalLM.from_pretrained(gpt2).cuda().half() # fp16 out model.generate( tok(Hello, return_tensorspt).input_ids.cuda(), do_sampleTrue, temperature0.1, max_new_tokens20, )报错或异常行为RuntimeError: probability tensor contains either inf, nan or element 0或者没有抛错但生成结果全是重复 token、半角符号、乱码——本质是 softmax 之后某一项变成了infsoftmax 被inf污染采样退化为 argmax 或随机乱跳。最迷惑的地方在于同一个模型把temperature去掉纯 greedy /do_sampleFalse就完全正常把temperature调到1.0也基本正常只有temperature取很小的值比如0.1、0.01时高频出现。这种「条件触发」让很多同学以为是数据问题浪费大量时间。二、背景temperature在采样里的标准做法是先把 logits 除以温度再做 softmax然后 multinomial 采样。# 标准公式 probs torch.softmax(logits / temperature, dim-1) next_token torch.multinomial(probs, num_samples1)温度越小logits 被放得越大。当temperature0.1时相当于 logits 放大 10 倍。如果原始 logits 里某些通道数值偏大尤其在 fp16 下logits 本身就是半精度动态范围窄放大之后很容易超过 fp16 的表示上限 65504变成inf。inf进了 softmaxsoftmax([inf, x, y]) - [1, 0, 0] # 看起来还行但如果不止一个inf比如多个通道都溢出softmax 变成inf - inf的nan然后 multinomial 直接抛上面的RuntimeError。更隐蔽的是另一类成因减最大值numerical stability这一步在 fp16 下被「反向放大」了。常规 softmax 会先logits - logits.max()防溢出但这是针对「不缩放」的情况。一旦先除以温度再减最大值或者减最大值用的是放大后的数值稳定项本身也被放大等于没稳定。还有第三个成因在LogitsProcessor里有的实现用torch.where(condition, -inf, scores)做 mask然后在 fp16 下-inf / temperature仍是-inf但反过来的inf通道没被处理于是正负无穷并存softmax 出nan。三、根因根因归纳为一句话温度缩放把数值放大后既没有在缩放前做稳定化也没有对溢出/无穷做兜底导致 fp16 下的 logits 溢出成inf/nan污染了 softmax 与采样。具体三处缩放顺序错误代码在 fp16 张量上直接logits / temperature放大发生在「减最大值」之前或之后但用了放大后的值稳定项失效。dtype 不匹配logits 是 fp16温度缩放与 softmax 全在 fp16 算动态范围不够。正确做法是把缩放挪到 fp32 下做再回 fp16/交给采样。无穷未兜底mask 产生的-inf、溢出产生的inf没有被统一检测与替换softmax 在「多无穷」时产生nan。这不是模型的问题也不是数据的锅而是采样前置处理temperature 缩放 softmax在半精度下的数值稳定性缺失。四、最小可运行复现下面这段不依赖真实大模型手动构造会溢出的 logits把问题放大给你看import torch def naive_sample(logits, temperature): # 模拟 transformers 里「直接在 logits 上除温度」的朴素实现 scaled logits / temperature probs torch.softmax(scaled, dim-1) return torch.multinomial(probs, num_samples1) # fp16 下构造几个很大的 logit模拟模型输出极值 logits_fp16 torch.tensor([[30000.0, -5.0, 20.0, 100.0]], dtypetorch.float16, devicecuda) for t in [1.0, 0.5, 0.1, 0.01]: try: tok naive_sample(logits_fp16, t) print(ftemperature{t}: 采样成功, token{tok.item()}) except Exception as e: print(ftemperature{t}: 崩溃 - {type(e).__name__}: {e}) # 验证溢出直接看缩放后的值 scaled logits_fp16 / 0.01 print(scaled 是否含 inf:, torch.isinf(scaled).any().item()) print(scaled 是否含 nan:, torch.isnan(scaled).any().item())跑出来你会看到temperature1.0还可能正常一旦到0.1、0.01scaled里30000/0.01 3_000_000远超 fp16 上限 →infsoftmax([inf, ...])在多无穷情况下出nanmultinomial抛错。这就精确复现了线上现象。五、解决方案第一层最小直接修复最小修复在 fp32 下做温度缩放与 softmax并对无穷做兜底。这是一个可直接替换进采样流程的稳定函数import torch def stable_sample(logits, temperature1.0, top_k0, top_p1.0, generatorNone): # 1) 统一升到 fp32避免半精度溢出 scores logits.to(torch.float32) # 2) 先做数值稳定减最大值再做温度缩放 scores scores - scores.max(dim-1, keepdimTrue).values if temperature ! 1.0 and temperature 0: scores scores / temperature # 3) 兜底任何残留的 inf/nan 都替换成极端但不致命的值 scores torch.where(torch.isfinite(scores), scores, torch.full_like(scores, -1e4)) # 4) top_k / top_p 过滤可选但建议保留 if top_k and top_k 0: kth torch.topk(scores, top_k).values[..., -1, None] scores torch.where(scores kth, scores, torch.full_like(scores, -1e4)) if top_p 1.0: sorted_logits, sorted_idx torch.sort(scores, descendingTrue) cum torch.cumsum(torch.softmax(sorted_logits, -1), -1) remove cum top_p remove[..., 1:] remove[..., :-1].clone() remove[..., 0] False mask remove.scatter(-1, sorted_idx, remove) scores torch.where(mask, torch.full_like(scores, -1e4), scores) probs torch.softmax(scores, dim-1) # 5) 采样前再确认没有 inf/nan if torch.isnan(probs).any() or torch.isinf(probs).any(): probs torch.ones_like(probs) / probs.shape[-1] return torch.multinomial(probs, num_samples1, generatorgenerator) # 用第四节的溢出 logits 验证 bad torch.tensor([[30000.0, -5.0, 20.0, 100.0]], dtypetorch.float16, devicecuda) for t in [1.0, 0.5, 0.1, 0.01]: tok stable_sample(bad, temperaturet) print(ftemperature{t}: 稳定采样成功, token{tok.item()})关键改动缩放前先升 fp32再减最大值温度缩放作用于「已稳定」的 scores不会再溢出。用torch.where(isfinite, x, -1e4)把任何inf/nan变成「极小但有限」的值。这样多个溢出通道不会凑出nan。采样前最后再 check 一次彻底杜绝multinomial抛错。这一步单独就能让带温度的 fp16 采样稳定运行。六、解决方案第二层结构性改进第一层是「在采样函数里修一处」。但生成入口很多model.generate、TextGenerationPipeline、各种LogitsProcessor、训练时 teacher forcing 的采样最好在框架层放一个统一的「温度缩放 稳定化」策略让所有入口共用。下面用 dataclass 作为单一事实来源from dataclasses import dataclass, field from typing import Optional import torch dataclass class TemperatureScaler: 统一的温度缩放与数值稳定策略。 # 是否在缩放前升 fp32fp16/bf16 下强烈建议 True upcast_to_float32: bool True # 缩放后兜底替换值有限避免 nan finite_floor: float -1e4 # 允许的最小温度避免 0 导致除零 min_temperature: float 1e-3 # 缩放前是否减去最大值做稳定 subtract_max: bool True def __call__(self, logits: torch.Tensor, temperature: float) - torch.Tensor: if temperature is None or temperature 1.0: return logits temp max(temperature, self.min_temperature) work logits if self.upcast_to_float32: work work.float() if self.subtract_max: work work - work.max(dim-1, keepdimTrue).values work work / temp work torch.where( torch.isfinite(work), work, torch.full_like(work, self.finite_floor), ) return work def safe_softmax(self, scaled: torch.Tensor) - torch.Tensor: probs torch.softmax(scaled, dim-1) bad torch.isnan(probs) | torch.isinf(probs) if bad.any(): # 退化到均匀分布保证采样永远可跑 probs torch.where(bad, torch.full_like(probs, 1.0 / probs.shape[-1]), probs) return probs # 用法任何采样入口都先过它 scaler TemperatureScaler() scaled scaler(fp16_logits, temperature0.1) probs scaler.safe_softmax(scaled)结构上的收益统一入口generate、pipeline、训练采样器全部调用同一个TemperatureScaler不会某个入口忘做稳定化。配置化upcast_to_float32、finite_floor、min_temperature都可按硬件/精度调不用改逻辑。兜底确定性safe_softmax保证「永远返回合法的有限概率分布」下游multinomial再也不会因inf/nan崩。七、解决方案第三层断言 / CI 守护写一组 pytest守两条铁律(1) 任意温度任意精度下采样都返回有限概率(2) 溢出 logits 不会让采样崩。import torch import pytest from your_lib import TemperatureScaler pytest.mark.parametrize(temperature, [1.0, 0.5, 0.1, 0.01, 1e-4]) pytest.mark.parametrize(dtype, [torch.float32, torch.float16, torch.bfloat16]) def test_temperature_never_produces_inf_or_nan(temperature, dtype): if dtype torch.float16 and not torch.cuda.is_available(): pytest.skip(fp16 需 cuda) scaler TemperatureScaler() # 构造会溢出的极端 logits logits torch.tensor( [[30000.0, -5.0, 20.0, 100.0, -30000.0]], dtypedtype, devicecuda if torch.cuda.is_available() else cpu, ) scaled scaler(logits, temperature) probs scaler.safe_softmax(scaled) assert torch.isfinite(probs).all(), f{dtype} t{temperature} 出现非有限概率 assert probs.shape (1, 5) # 概率和应为 1 assert torch.allclose(probs.sum(-1), torch.ones(1, deviceprobs.device), atol1e-4) def test_multinomial_runs_on_overflow(): scaler TemperatureScaler() logits torch.tensor([[30000.0, 30001.0, -5.0]], dtypetorch.float16) scaled scaler(logits, 0.01) probs scaler.safe_softmax(scaled) # 不抛 RuntimeError sampled torch.multinomial(probs, num_samples1) assert sampled.shape (1, 1) def test_temperature_zero_guard(): scaler TemperatureScaler(min_temperature1e-3) logits torch.randn(1, 10) # 即使传 0也会被钳到 min_temperature不除零 scaled scaler(logits, 0.0) assert torch.isfinite(scaled).all()CI 常驻跑这三个测试后任何「把缩放挪回 fp16」「去掉兜底」的改动都会立刻失败。八、排查清单采样出现inf/nan时按顺序排查先去掉temperature试 greedy正常说明问题在「缩放 精度」不在模型或数据。确认logits.dtype如果是 fp16/bf16立刻怀疑溢出。把缩放改到 fp32 再试。确认缩放顺序必须「先减最大值稳定再除以温度」而不是反过来。检查有没有torch.where(cond, -inf, x)这类 mask-inf会和inf共存导致nan需统一兜底。确认temperature不会被传0除以 0 直接inf必须钳最小值。若用了top_k/top_p确认过滤用的是「替换成有限极小值」而非「乘 0」——乘 0 后-inf*0 nan。多卡/AMP 下确认 logits 进入采样前没有在半精度下经历额外的大数运算如重复缩放。九、小结「带 temperature 采样出现 inf」不是玄学而是半精度下「先做温度缩放、后做稳定化」顺序颠倒叠加无穷未兜底导致的数值溢出。修复三层次第一层在 fp32 下「减最大值 → 除温度 → 兜底无穷」让单次采样稳住第二层用TemperatureScalerdataclass 把策略收敛为框架统一入口所有采样路径共用第三层用 pytest 守「任意温度任意精度都返回有限概率」「溢出 logits 不崩 multinomial」。工程启示任何涉及「除以一个可能很小/很大系数」的半精度计算都要把缩放挪到高位精度、缩放后做稳定、并对无穷显式兜底。采样、对比学习温度系数、对比损失里的tau、知识蒸馏的T都是同一个坑照此处理即可。
返回列表