通过率 0% · 提交 0 · 通过 0
分类模型输出了 r 行、每行 c 个类别分数(logits),以及每行的真实类别。对每一行:先减去该行最大值,再计算 softmax 概率;随后计算全部行的平均交叉熵(自然对数)。注意:当分数跨度很大时,真实类别的概率可能下溢为 0,交叉熵必须用 log-sum-exp 的方式计算,即 −ln p = ln(Σexp(z−max)) − (z_label−max),不能对已下溢的概率取对数。
这类题属于算法机考高频题型中「华为 AI 岗 / Softmax」方向的高频题型,通常考察对「华为 AI 岗 / Softmax」的建模能力与边界条件处理。掌握本题的解题思路后,可举一反三应对同类真题方向,稳步提升机考通过率。
第一行输入 r c,分别为行数和类别数。接下来 r 行,每行 c 个实数 logits。最后一行输入 r 个整数,表示每行的真实类别(0 开始编号)。
输出 r 行,每行 c 个保留四位小数的概率,用空格分隔。最后一行输出保留六位小数的平均交叉熵。绝对值小于对应精度一半时输出 0.0000(或 0.000000)。
示例 1
输入示例
1 3 1 2 3 2
输出示例
0.0900 0.2447 0.6652 0.407606
单行三分类,交叉熵为 −ln(0.6652…)。
示例 2
输入示例
1 3 1000 999 0 0
输出示例
0.7311 0.2689 0.0000 0.313262
大数行必须先减最大值再取指数,否则溢出。
时间限制 2000 ms · 内存限制 256 MB
本平台为独立第三方培训机构,与华为技术有限公司无任何关联;课程的服务内容与权益以购买协议为准,学习效果因个人情况而异。「华为 OD」「华为可信」等仅为对岗位与考试方向的客观描述,相关商标归各自权利人所有。
这些是真正决定能不能 AC、但通用题解里常被略过的点。
稳定 Softmax 与平均交叉熵,两个都是面试手撕的数值稳定名场面。概率那一半减 max 就够;交叉熵那一半藏着一个更深的坑——下溢。
每行取 m = max(z),概率 w_j = exp(z_j − m) / Σ exp(z_t − m)。logits 绝对值可以到 1000,exp(1000) 溢出,减 max 后指数参数最大为 0,安全且数学等价。
平均交叉熵 = (1/r)·Σ −ln p_label。看似把上一步算好的 p_label 拿来 log 一下就完事,但当真实类别的分数远低于最大分数时(比如差 800),p_label ≈ exp(−800) 在 double 里下溢为 0,log(0) 得到 −inf,全盘皆输。
正确姿势是走 log-sum-exp,从头到尾不落地成概率:
−ln p_label = ln(Σ exp(z_t − m)) − (z_label − m)
右边两项都在安全范围内:求和至少为 1(含 exp(0) 那一项),对数不会炸;z_label − m 是普通减法。概率照常输出(下溢打出 0.0000 没问题),交叉熵用这条公式单独算。
时间 O(r·c),空间 O(c)。
-0.0000 未钳制。样例 1:一行 logits 1 2 3,真实类别 2。减 max=3 后是 −2、−1、0,指数 0.1353、0.3679、1,和为 1.5032:
这行数字温和,两条路径殊途同归;把 logits 换成 1 2 1000 那类用例,概率路径下溢为 0、log 炸成 −inf,log-sum-exp 路径面不改色——两条公式的分水岭就在这里。
数值稳定的 softmax 与交叉熵在面试里几乎是「写过训练代码吗」的试金石。标准讲法:exp 上溢用减 max 解决,这是恒等变换;交叉熵下溢用 log-sum-exp 解决,本质是把 log 和 exp 在公式层面抵消掉、绝不让概率落地成 0 再取对数。追问「框架里为什么把 softmax 和交叉熵合成一个算子」:分开算两次数值风险,合起来公式里的 exp 与 log 相消,又稳又省——PyTorch 的 CrossEntropyLoss 接收 logits 而非概率,原因就在这里。答到这一层,这道题就从格式题变成了理解题。
概率行对、交叉熵行错,是最有诊断价值的组合——它几乎锁定「对下溢概率取了对数」或对数底用错;两行都错则先查减 max。输出 inf/nan 直接指向没减 max 或 log(0)。最后一行保留六位、概率保留四位,格式混用也是一档独立的错误来源。
# 每行减 max 再 exp 求 softmax;交叉熵不对概率取对数,用 log-sum-exp:−ln p = ln Σexp(z−max) − (z_label−max)。
# 自然对数;概率四位、平均交叉熵六位,注意负零钳制。
import math
import sys
def fmt(value: float, digits: int, eps: float) -> str:
if abs(value) < eps:
value = 0.0
return f"{value:.{digits}f}"
def solve() -> None:
data = sys.stdin.buffer.read().split()
if not data:
return
it = iter(data)
r = int(next(it))
c = int(next(it))
rows = [[float(next(it)) for _ in range(c)] for _ in range(r)]
labels = [int(next(it)) for _ in range(r)]
out = []
ce_sum = 0.0
for row, label in zip(rows, labels):
mx = row[0]
for v in row:
if v > mx:
mx = v
denom = 0.0
exps = []
for v in row:
e = math.exp(v - mx)
denom += e
exps.append(e)
out.append(" ".join(fmt(e / denom, 4, 0.00005) for e in exps))
ce_sum += -((row[label] - mx) - math.log(denom))
out.append(fmt(ce_sum / r, 6, 0.0000005))
print("\n".join(out))
if __name__ == "__main__":
solve()
登录后可查看你在本题的历史提交,以及每次的各用例通过情况。