通用算法与机器学习手写题
算法手写 高级 重要度 4/5 知识骨架 理解阶段
手写 Temperature、Top-k 与 Top-p 采样
手写 Temperature、Top-k 与 Top-p 采样的可运行 PyTorch 模板、复杂度、数值边界、变体与对拍方法。
发布于 2026-08-31 · 更新于 2026-09-07
本文目录
图:辅助检查本题的数据流与张量关系。来源:Beam search,作者 D2L.ai authors,许可 CC BY-SA 4.0。
识别信号
看到解码策略题先区分训练 logits 与推理采样;明确 temperature 的位置、过滤是在概率归一化前还是后,以及随机数种子。
核心思路
temperature 用 logits/T 调节相对差距;Top-k 保留最大的 k 个 logits;Top-p 先按概率降序,保留累计概率首次达到 p 的最小前缀。过滤后重新归一并 multinomial。通常先 temperature,再 top-k/top-p。
实现模板
import torch
def sample(logits, temperature=1.0, top_k=None, top_p=None):
if temperature <= 0: return logits.argmax(dim=-1)
x = logits / temperature
if top_k is not None:
threshold = x.topk(min(top_k, x.size(-1)), dim=-1).values[..., -1, None]
x = x.masked_fill(x < threshold, float('-inf'))
if top_p is not None:
sorted_x, idx = x.sort(dim=-1, descending=True)
probs = sorted_x.softmax(dim=-1)
remove = probs.cumsum(dim=-1) - probs > top_p
sorted_x = sorted_x.masked_fill(remove, float('-inf'))
x = torch.full_like(x, float('-inf')).scatter(-1, idx, sorted_x)
return torch.multinomial(x.softmax(dim=-1), 1).squeeze(-1)面试时边写边报 shape、不变量和数值稳定策略。先完成可验证基线,再讨论融合 kernel 或分布式优化。
复杂度
Top-k 可近似 O(V log k),朴素 Top-p 排序 O(V log V),空间 O(V)。词表很大或 batch 很高时,排序与采样 kernel 也会进入 Decode 开销。
边界条件
- temperature=0 定义为 greedy
- k 大于词表
- p≤0 或 p≥1
- 并列 logits
- 过滤后至少保留一个 token
- 分布式词表并行下的全局 top-k
常见变体
追问 repetition penalty、min-p、typical sampling、beam search,以及为何低 temperature 仍不等于确定性。安全场景还要说明 stop token 与最大长度。
测试用例
x=torch.tensor([ [5.,4.,1.,-2.]])
assert sample(x,temperature=0).item()==0
for _ in range(20):
assert sample(x,top_k=1).item()==0
torch.manual_seed(7)
a=sample(x,top_p=.8)
torch.manual_seed(7)
b=sample(x,top_p=.8)
assert a.item()==b.item()正确性验证与对拍
“手写 Temperature、Top-k 与 Top-p 采样”先与 PyTorch 官方实现对拍前向结果和梯度。小规模测试使用 float64,并覆盖随机 shape、极值和非法输入;优化版还要与朴素实现比较误差、峰值显存和耗时。
复习自测
- 能否在十分钟内写出“手写 Temperature、Top-k 与 Top-p 采样”的核心实现,并解释每一维 shape?
- 哪三类输入最容易产生 NaN、越界或错误广播?
- 数据规模扩大十倍后,瓶颈更可能落在算术、内存还是 kernel 启动?