Temperature:先保护除数,再保持 logits 的形状与索引
🎯 先猜一猜
两个 logits 是 4.0 和 3.1,temperature=0.5。
缩放后的分差是多少?
先补的知识
- logits 是 [batch, vocab_size] 的浮点张量,每一列仍代表固定的 token_id;Temperature 不排序也不删除列。
- Softmax 前除以 T:T 小于 1 会放大差距,T 大于 1 会缩小差距,T=1 不改变 logits。
- temperature 是 Python float,作为除数必须设置极小正数下限,避免 0 导致无定义结果。
- 逐元素除法不改变 logits 的 shape、dtype 或 device。
图解原理
Temperature 像调节分数表的对比度:低温把分差拉大,高温把分差压小。它只缩放数值,不负责把分数变成概率;Softmax 在完整解码流水线最后统一执行。
T < 1
除以小数,分差放大,采样更集中。
T = 1
数值不变,保留原始相对差距。
T > 1
除以大数,分差缩小,采样更分散。
语法热身:用考试分数练安全除法和差值变化
import torch
ratings = torch.tensor([[8.0, 5.0, 2.0]])
temperature = 0.5
safe_temperature = max(temperature, 1e-6)
adjusted = ratings / safe_temperature
print(adjusted.shape) # [1,3]
print(adjusted[0,0] - adjusted[0,1]) # 6.0从语法例子迁移到 TODO 1
ratings:对应 logits,每一列的位置含义不能改变safe_temperature:对应 temp,使用 Python max 设置正数下限adjusted:对应返回值 logits / temp,不在这里调用 Softmax差值 6.0:原差值 3.0 除以 0.5 后翻倍,对应可见测试的检查方式
巩固一下
TODO 1 为什么不直接 return F.softmax(logits / temp, dim=-1)?
学完这一段,试着做
用一个小动作确认自己理解了;最后再进入官方题目。