通用算法与机器学习手写题
算法手写 高级 重要度 4/5 知识骨架 理解阶段

手写 Temperature、Top-k 与 Top-p 采样

手写 Temperature、Top-k 与 Top-p 采样的可运行 PyTorch 模板、复杂度、数值边界、变体与对拍方法。

发布于 2026-08-31 · 更新于 2026-09-07

本文目录

    Beam search。辅助检查本题的数据流与张量关系。

    图:辅助检查本题的数据流与张量关系。来源: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 启动?

    来源与关联知识

    输入关键词,查找全部技术文章。

      搜索范围:正文、标题、分类和标签