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

手写数值稳定的 Softmax 与交叉熵

手写数值稳定的 Softmax 与交叉熵的可运行 PyTorch 模板、复杂度、数值边界、变体与对拍方法。

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

本文目录

    Compute graph and automatic differentiation。辅助检查本题的数据流与张量关系。

    图:辅助检查本题的数据流与张量关系。来源:Compute graph and automatic differentiation,作者 D2L.ai authors,许可 CC BY-SA 4.0。

    识别信号

    看到“不能调用 softmax/cross_entropy”“解释减最大值”“推导梯度”时,先确认输入 logits 形状 [B,C]、标签是类别下标还是分布,以及 reduction 方式。

    核心思路

    利用 softmax 的平移不变性,每行减去最大 logits,避免 exp 上溢。交叉熵不要先算概率再取 log,而用 logsumexp:loss_i=logΣexp(z_ij)-z_i,y。其梯度对 logits 为 (p-one_hot)/B。

    实现模板

    import torch
    
    def stable_cross_entropy(logits, target):
        # logits: [B, C], target: [B]
        row_max = logits.max(dim=-1, keepdim=True).values
        shifted = logits - row_max
        logsumexp = torch.log(torch.exp(shifted).sum(dim=-1))
        chosen = shifted.gather(1, target[:, None]).squeeze(1)
        return (logsumexp - chosen).mean()
    
    def softmax(logits):
        shifted = logits - logits.max(dim=-1, keepdim=True).values
        exp = torch.exp(shifted)
        return exp / exp.sum(dim=-1, keepdim=True)

    面试时边写边报 shape、不变量和数值稳定策略。先完成可验证基线,再讨论融合 kernel 或分布式优化。

    复杂度

    时间 O(BC),若显式保存概率则空间 O(BC)。融合 kernel 可避免中间张量和额外 HBM 往返,但不改变渐进复杂度。

    边界条件

    • logits 含极大正负值
    • batch=1 或 class=1
    • target 越界
    • FP16/BF16 的指数和累加精度
    • ignore_index、label smoothing 与 soft labels

    常见变体

    继续追问 log-softmax、KL divergence、label smoothing,以及为何训练框架会把 log-softmax 与 NLLLoss 融合。

    测试用例

    x=torch.tensor([ [1000.,999.,-1000.],[-1000.,-999.,-998.]],requires_grad=True)
    y=torch.tensor([0,2])
    ref=torch.nn.functional.cross_entropy(x,y)
    out=stable_cross_entropy(x,y)
    assert torch.allclose(out,ref,atol=1e-6)
    out.backward()
    assert torch.isfinite(x.grad).all()

    正确性验证与对拍

    “手写数值稳定的 Softmax 与交叉熵”先与 PyTorch 官方实现对拍前向结果和梯度。小规模测试使用 float64,并覆盖随机 shape、极值和非法输入;优化版还要与朴素实现比较误差、峰值显存和耗时。

    复习自测

    • 能否在十分钟内写出“手写数值稳定的 Softmax 与交叉熵”的核心实现,并解释每一维 shape?
    • 哪三类输入最容易产生 NaN、越界或错误广播?
    • 数据规模扩大十倍后,瓶颈更可能落在算术、内存还是 kernel 启动?

    来源与关联知识

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

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