当前位置: 首页 > news >正文

【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都是同一个坑,照此处理即可。

http://www.jsqmd.com/news/1355189/

相关文章:

  • 10种精选配色方案:GitHub ReadME Terminal让你的主页脱颖而出
  • 如何快速搭建专业级视频监控平台:wvp-GB28181-pro零代码部署实战指南
  • RetrofitCache实战案例:构建离线优先的Android应用
  • Distill-Any-Depth模型优化技巧:如何在保持精度的同时减小模型体积?
  • 终极指南:如何快速找回遗忘的7z/Zip/Rar压缩包密码?完整解决方案
  • 终极指南:使用Swagger UI Express快速构建API文档
  • 深入理解Ookii.Dialogs.WinForms的设计模式与架构
  • 深度解析DevDocs存储架构:从资源管理到性能优化实战指南
  • TencentDB Agent Memory源码解析:核心模块与关键算法实现
  • 原来重庆这些校园广播系统公司这么靠谱,究竟是哪些呢?
  • UsbDk高级技巧:批量传输与等时传输的优化实现
  • 3大创新架构:如何构建250+格式的零信任本地化文件转换引擎
  • Ookii.Dialogs.WinForms跨版本支持:从.NET Framework到.NET 6的终极指南
  • 终极指南:如何用开源工具biliTickerBuy轻松搞定B站会员购抢票难题
  • 构建跨云AI代理:Agent Governance Toolkit多云部署策略
  • 【Bug已解决】[serge] integration failure triage - 2026-07-05 解决方案
  • word转pdf软件有哪些?七款PDF格式转换工具实测盘点
  • CloudWalker Platform源代码解析:Go语言实现的高性能检测引擎
  • 开发者必看:Apify MCP Server核心组件与架构详解
  • python神经网络编程入门(二十七)——RNN IMBD搭建情感分类器与基础训练
  • 终极指南:如何在macOS上使用BlackHole实现零延迟音频环回
  • js-stellar-sdk错误处理完全手册:解决90%的Stellar开发问题
  • word转图片怎么转?7款PDF格式转换工具实测盘点,免费与官方方法一次说清
  • MongoKit索引优化指南:提升MongoDB查询性能的完整方案
  • 如何在5分钟内为Tailwind项目添加Apple式平滑圆角?Corner Smoothing插件快速上手
  • 如何构建企业级语义层:Cube Core实战架构指南与性能优化策略
  • Ookii.Dialogs.WinForms高级技巧:如何实现Vista风格文件对话框
  • Grapple.nvim项目作用域详解:Git仓库、LSP与自定义作用域配置教程
  • Klock实战教程:如何在Android与iOS项目中集成日期时间功能
  • 解决mechabar常见问题:从依赖安装到主题适配的完整解决方案