三亩地 三亩地SAN MU DI · CODE DIARY
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_tensors="pt").input_ids.cuda(), do_sample=True, temperature=0.1, max_new_tokens=20, )

报错或异常行为:

RuntimeError: probability tensor contains either `inf`, `nan` or element < 0

或者没有抛错,但生成结果全是重复 token、半角符号、乱码——本质是 softmax 之后某一项变成了inf,softmax 被inf污染,采样退化为 argmax 或随机乱跳。

最迷惑的地方在于:同一个模型,把temperature去掉(纯 greedy /do_sample=False)就完全正常;把temperature调到1.0也基本正常;只有temperature取很小的值(比如0.10.01)时高频出现。这种「条件触发」让很多同学以为是数据问题,浪费大量时间。

二、背景

temperature在采样里的标准做法是:先把 logits 除以温度,再做 softmax,然后 multinomial 采样。

# 标准公式 probs = torch.softmax(logits / temperature, dim=-1) next_token = torch.multinomial(probs, num_samples=1)

温度越小,logits 被放得越大。当temperature=0.1时,相当于 logits 放大 10 倍。如果原始 logits 里某些通道数值偏大(尤其在 fp16 下,logits 本身就是半精度,动态范围窄),放大之后很容易超过 fp16 的表示上限 65504,变成+inf+inf进了 softmax:

softmax([inf, x, y]) -> [1, 0, 0] # 看起来还行?

但如果不止一个+inf(比如多个通道都溢出),softmax 变成inf - infnan,然后 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 与采样。

具体三处:

  1. 缩放顺序错误:代码在 fp16 张量上直接logits / temperature,放大发生在「减最大值」之前(或之后但用了放大后的值),稳定项失效。
  2. dtype 不匹配:logits 是 fp16,温度缩放与 softmax 全在 fp16 算,动态范围不够。正确做法是把缩放挪到 fp32 下做,再回 fp16/交给采样。
  3. 无穷未兜底: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_samples=1) # fp16 下,构造几个很大的 logit(模拟模型输出极值) logits_fp16 = torch.tensor([[30000.0, -5.0, 20.0, 100.0]], dtype=torch.float16, device="cuda") for t in [1.0, 0.5, 0.1, 0.01]: try: tok = naive_sample(logits_fp16, t) print(f"temperature={t}: 采样成功, token={tok.item()}") except Exception as e: print(f"temperature={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())

跑出来你会看到:temperature=1.0还可能正常,一旦到0.10.01scaled30000/0.01 = 3_000_000远超 fp16 上限 →infsoftmax([inf, ...])在多无穷情况下出nanmultinomial抛错。这就精确复现了线上现象。

五、解决方案(第一层:最小直接修复)

最小修复:在 fp32 下做温度缩放与 softmax,并对无穷做兜底。这是一个可直接替换进采样流程的稳定函数:

import torch def stable_sample(logits, temperature=1.0, top_k=0, top_p=1.0, generator=None): # 1) 统一升到 fp32,避免半精度溢出 scores = logits.to(torch.float32) # 2) 先做数值稳定(减最大值),再做温度缩放 scores = scores - scores.max(dim=-1, keepdim=True).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, descending=True) 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_samples=1, generator=generator) # 用第四节的溢出 logits 验证 bad = torch.tensor([[30000.0, -5.0, 20.0, 100.0]], dtype=torch.float16, device="cuda") for t in [1.0, 0.5, 0.1, 0.01]: tok = stable_sample(bad, temperature=t) print(f"temperature={t}: 稳定采样成功, token={tok.item()}")

关键改动:

  • 缩放前先升 fp32,再减最大值,温度缩放作用于「已稳定」的 scores,不会再溢出。
  • torch.where(isfinite, x, -1e4)把任何inf/nan变成「极小但有限」的值。这样多个溢出通道不会凑出nan
  • 采样前最后再 check 一次,彻底杜绝multinomial抛错。

这一步单独就能让带温度的 fp16 采样稳定运行。

六、解决方案(第二层:结构性改进)

第一层是「在采样函数里修一处」。但生成入口很多(model.generateTextGenerationPipeline、各种LogitsProcessor、训练时 teacher forcing 的采样),最好在框架层放一个统一的「温度缩放 + 稳定化」策略,让所有入口共用。

下面用 dataclass 作为单一事实来源:

from dataclasses import dataclass, field from typing import Optional import torch @dataclass class TemperatureScaler: """统一的温度缩放与数值稳定策略。""" # 是否在缩放前升 fp32(fp16/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, keepdim=True).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, temperature=0.1) probs = scaler.safe_softmax(scaled)

结构上的收益:

  • 统一入口generatepipeline、训练采样器全部调用同一个TemperatureScaler,不会某个入口忘做稳定化。
  • 配置化upcast_to_float32finite_floormin_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]], dtype=dtype, device="cuda" 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, device=probs.device), atol=1e-4) def test_multinomial_runs_on_overflow(): scaler = TemperatureScaler() logits = torch.tensor([[30000.0, 30001.0, -5.0]], dtype=torch.float16) scaled = scaler(logits, 0.01) probs = scaler.safe_softmax(scaled) # 不抛 RuntimeError sampled = torch.multinomial(probs, num_samples=1) assert sampled.shape == (1, 1) def test_temperature_zero_guard(): scaler = TemperatureScaler(min_temperature=1e-3) logits = torch.randn(1, 10) # 即使传 0,也会被钳到 min_temperature,不除零 scaled = scaler(logits, 0.0) assert torch.isfinite(scaled).all()

CI 常驻跑这三个测试后,任何「把缩放挪回 fp16」「去掉兜底」的改动都会立刻失败。

八、排查清单

采样出现inf/nan时按顺序排查:

  1. 先去掉temperature试 greedy:正常说明问题在「缩放 + 精度」,不在模型或数据。
  2. 确认logits.dtype:如果是 fp16/bf16,立刻怀疑溢出。把缩放改到 fp32 再试。
  3. 确认缩放顺序:必须「先减最大值(稳定)再除以温度」,而不是反过来。
  4. 检查有没有torch.where(cond, -inf, x)这类 mask:-inf会和+inf共存导致nan,需统一兜底。
  5. 确认temperature不会被传0:除以 0 直接inf,必须钳最小值。
  6. 若用了top_k/top_p,确认过滤用的是「替换成有限极小值」而非「乘 0」——乘 0 后-inf*0 = nan
  7. 多卡/AMP 下,确认 logits 进入采样前没有在半精度下经历额外的大数运算(如重复缩放)。

九、小结

「带 temperature 采样出现 inf」不是玄学,而是半精度下「先做温度缩放、后做稳定化」顺序颠倒,叠加无穷未兜底导致的数值溢出。修复三层次:第一层在 fp32 下「减最大值 → 除温度 → 兜底无穷」,让单次采样稳住;第二层用TemperatureScalerdataclass 把策略收敛为框架统一入口,所有采样路径共用;第三层用 pytest 守「任意温度任意精度都返回有限概率」「溢出 logits 不崩 multinomial」。

工程启示:任何涉及「除以一个可能很小/很大系数」的半精度计算,都要把缩放挪到高位精度、缩放后做稳定、并对无穷显式兜底。采样、对比学习温度系数、对比损失里的tau、知识蒸馏的T都是同一个坑,照此处理即可。

← 返回列表