torch.argmax 这个函数,几乎是每个用 PyTorch 做分类任务的人最早接触到的 API 之一。你一定会遇到这样的场景:模型输出一个形状为[batch_size, num_classes]的分数矩阵,你想知道每个样本到底被预测成了哪个类别,于是敲下torch.argmax(output, dim=1),得到一串整数标签。而另一边,如果你的训练数据是从 one-hot 格式存的,比如[0, 0, 1, 0]这种向量,你想把它还原成整数2,最直接的办法同样是用torch.argmax(one_hot_vector, dim=1)。看起来都是“沿着第 1 维找最大值索引”这一件事,但真到代码里,维度选错、形状搞混、one-hot 转不出来整数标签的问题,我在实际项目里见过太多次了。这篇文章就把dim=1的语义、它与 one-hot 转整数标签之间的对应关系、以及实战中常见的坑一次性讲透,适合刚接触 PyTorch 的初学者,也适合写了多年模型但偶尔被维度问题绊一脚的工程师。
1. 先搞懂 dim=1:张量的“轴”就是业务语义的开瓶器
很多人对argmax的理解停留在“取最大值的下标”,但一碰到多维张量就开始懵,根源是把“张量的维度”和“业务里的维度”搞混了。
1.1 分类模型的默认形状约定:[批量, 类别数]
深度学习里处理分类任务时,模型最后一层输出的 logits 一般都约定成[batch_size, num_classes]。比如批量大小为 8,类别数为 10,输出就是[8, 10]。在这个约定下:
- 第 0 维是“样本”维度,代表这个批次里有几张图、几句话、几段序列。
- 第 1 维是“类别”维度,代表每个样本分别在 10 个类别上的得分。
这里的第 1 维,就是dim=1对应的那个轴。torch.argmax(output, dim=1)的意思就是:对每个样本(固定第 0 维的某个具体值),沿着类别方向扫一遍,找出得分最高的那个位置的编号。结果形状变成[8],每个元素是一个0~9之间的整数索引。
这个编号,就是模型认为该样本所属的类别。举个例子,某个样本在第 5 个类别的得分最高,那argmax的结果就是4(索引从 0 开始)。这一步正在做的事情,本质上就是把模型输出的“分数分布”翻译成“整数标签”。
1.2 argmax 在轴上到底是怎么扫的
拿一个最小例子来说。设logits = tensor([[1.0, 2.0, 3.0], [4.0, 1.0, 2.0]]),形状是[2, 3]。
torch.argmax(logits, dim=1)的结果是[2, 0]。怎么来的?第一行里最大的是索引 2 对应的 3.0,第二行里最大的是索引 0 对应的 4.0。你注意看,dim=1说的是“消掉第 1 维”,也就是把[2, 3]变成[2],操作方向是横着对每一行内部做比较。
如果你反过来用dim=0,那就是纵向比较:对第 0 个类别取两个样本的最大值索引,对第 1 个类别取两个样本的最大值索引,对第 2 个类别取两个样本的最大值索引,结果变成[1, 0, 0]这样的长度 3 的张量。这个结果在分类任务里几乎没有任何业务意义,因为纵向比较的是“不同样本在同一个类别上的得分”,这不是我们想要的东西。
所以,理解dim的核心心法只有一个:dim 指定的是你要“消掉”哪个轴。剩下的轴就是你结果里保留的维度。
1.3 dim=1、dim=-1 和 dim=0 的惯用法
实际项目里,dim=1和dim=-1在二维张量上是同一个意思,因为-1代表最后一个轴。但如果你的张量是三维的,情况就不一样了。比如图像语义分割模型的输出形状是[B, C, H, W],类别维度 C 在中间,这时你依然应该用dim=1,而不是dim=-1。dim=-1在这个场景下对应的是 W 维度,会直接让你得到一堆坐标值而不是类别索引。
我见过的另一种混乱来自“通道最后排”的格式,比如某些模型输出[B, H, W, C],这时候才需要用dim=-1。判断标准永远不是“别人说用哪个”,而是“你的张量布局里,类别到底在第几维”。要么在模型代码里统一约定输出格式,要么在写argmax之前用一行注释写明形状:# [B, C, H, W], C dimension = num_classes,这条注释能省掉很多排查时间。
2. one-hot 与整数标签:一对可以互相翻译的兄弟
搞定了dim=1的语义,再来盘一盘 one-hot。标题里把它和argmax绑在一起,是因为“one-hot 转整数标签”这件事,数学上跟argmax(dim=1)几乎是同一个操作。
2.1 one-hot 是标签的“展开形态”
one-hot 向量长这样:如果总共有 5 个类别,整数标签3对应的 one-hot 是[0, 0, 0, 1, 0]。它的含义很直白:在第 3 个(下标 3)位置上是 1,其余全是 0。这种表示把“类别编号”变成一个等长的向量,好处是类别之间没有大小关系,可以直接作为神经网络输出的监督目标。
在 PyTorch 里,one-hot 张量通常有两种来源。一种是数据预处理阶段直接存的,比如某些数据集给你的是经过 one-hot 编码的标签;另一种是代码里临时用函数生成的,比如:
import torch import torch.nn.functional as F y = torch.tensor([2, 0, 4]) one_hot = F.one_hot(y, num_classes=5) # one_hot 的 shape 是 [3, 5] # tensor([[0, 0, 1, 0, 0], # [1, 0, 0, 0, 0], # [0, 0, 0, 0, 1]])注意这里 an output 的最后一个维度就是类别数,形状是[样本数, 类别数]。你会发现它和模型输出的 logits 形状天然对齐,都是[B, C]。这正是后面一切转换能成立的基础。
2.2 还原整数标签:argmax 就是解码器
现在问题来了:如果你手里只有 one-hot 张量,怎么拿到整数标签?
做法很简单:
integer_labels = torch.argmax(one_hot, dim=1)由于 one-hot 每一行只有一个 1,其余全是 0,最大值必然出现在那个 1 的位置上,argmax(dim=1)返回的索引就是原来的整数标签。这是完全无损的逆变换,也是我在所有项目里推荐用、也唯一推荐用的逆变换方式。
为什么不推荐写循环遍历?因为向量化操作又快又简洁,还能保持梯度链(虽然 one-hot 本来也不需要梯度)。为什么不推荐用torch.nonzero?你当然可以用torch.nonzero(one_hot)[:, 1]拿到每个样本的索引,但写法绕、易读性差、要处理空行和维度变化,不如argmax一行清爽。
2.3 三种常见 one-hot 生成方式与统一出口
生成 one-hot 的方法不止F.one_hot一种,我在不同代码库里看到过各种写法:
| 生成方式 | 代码示例 | 适用场景 |
|---|---|---|
F.one_hot | F.one_hot(labels, num_classes=C) | 最标准,推荐优先使用 |
scatter_ | torch.zeros(B, C).scatter_(1, labels.unsqueeze(1), 1) | 老代码和自定义 loss 中常见 |
eye索引 | torch.eye(C)[labels] | 原理清晰,适合教学演示 |
无论你是用哪种方式生成的 one-hot,还原整数标签的出口统一都是argmax。这背后其实就是线性代数里的“标准基向量表示”:one-hot 是标准基,argmax告诉你基向量的序号。你在多个框架里都会看到同样的约定,TensorFlow 里的tf.argmax(one_hot, axis=1),NumPy 里的np.argmax(one_hot, axis=1),都一个道理。
3. 实操:把 dim=1 用到完整的分类流程里
理论说完了,接下来看几个真实项目里会碰到的代码场景。dim=1和 one-hot 转整数标签可不只是“写一行代码”这么简单,它在训练、推理、评估三个阶段分别扮演不同角色。
3.1 训练阶段:交叉熵损失为什么不需要你手动转 one-hot
很多新手一开始不理解:模型输出的是 logits,标签如果存成了 one-hot,那是不是要先手动把标签转成某种格式才能喂给损失函数?
答案是:如果用nn.CrossEntropyLoss,你根本不用转。
loss_fn = nn.CrossEntropyLoss() loss = loss_fn(logits, integer_labels) # integer_labels 是 [B] 的整数张量nn.CrossEntropyLoss内部会自己完成 softmax 和对数损失的计算,标签参数接收的就是整数索引,不需要 one-hot。这是 PyTorch 设计上的一个卡点:它不想让你在内存里多存一份[B, C]的稀疏大矩阵。
但如果你用的是nn.BCEWithLogitsLoss,那就是另一套语义了,它对应的是多标签分类,期望的是 0/1 标签而不是索引。所以,先看清楚你的损失函数是“单标签分类”还是“多标签分类”,再决定标签的形态。这比纠结 one-hot 转不转重要得多。
3.2 推理阶段:logits -> softmax -> argmax 的正确打开方式
推理时你通常想把模型输出变成“类别预测”和“置信度”两个东西。我见过两种常见写法:
写法一:
probs = torch.softmax(logits, dim=1) confidence, pred = torch.max(probs, dim=1)写法二:
pred = torch.argmax(logits, dim=1)它们其实不冲突。softmax是单调递增的归一化函数,所以对同一个logits张量,argmax(softmax(logits), dim=1)和argmax(logits, dim=1)的结果完全一致。区别在于你有没有想要置信度。如果你只想要类别编号,直接argmax就够了,省一次 softmax 的指数运算;如果你还想要这个预测有多大概率,那才需要先softmax再取max或者直接取对应位置的概率。
对于单标签分类,我推荐写成:
with torch.no_grad(): logits = model(x) probs = torch.softmax(logits, dim=1) pred = torch.argmax(probs, dim=1) max_conf, _ = torch.max(probs, dim=1)其实这里pred用logits算也可以,但都走probs会让代码语义更统一:你得到的pred和max_conf是从同一份概率分布里取出来的。
3.3 评估阶段:pred 和 label 的整数对齐是准确率计算的命门
算准确率时,标准的写法是:
correct = (pred == integer_labels).float().sum().item()注意这里pred是torch.argmax之后得到的整数索引,integer_labels也是整数索引。即使你的标签原本是 one-hot 格式,也要先通过torch.argmax(one_hot_labels, dim=1)转回整数再比较。
我遇到过一个典型的写法错误:有人把 one-hot 标签直接拿出来跟pred做比较,pred.shape是[B],one_hot_labels.shape是[B, C],广播机制默不作声地让比较变成“每个预测整数与每个 one-hot 元素比较”,最后算出来的准确率既不是 0 也不是真实值,而是一个莫名其妙的分数。这种情况你必须先打印两个张量的shape,打印完了基本一眼就能发现。
如果要做混淆矩阵,torch.argmax的结果同样是对齐利器。sklearn.metrics.confusion_matrix(integer_labels.cpu().numpy(), pred.cpu().numpy())这一步里,两个输入都必须是一维整数数组。模型输出和 one-hot 标签如果不先做argmax,根本没法喂给 sklearn。
3.4 进阶:分割任务和掩码任务里的维度处理
图像分割任务的输出形状是[B, C, H, W],这时dim=1同样是类别维度。逐像素的分类预测就变成了:
pred_mask = torch.argmax(logits, dim=1) # 结果是 [B, H, W]每个位置上都是该像素的类别索引。这个操作在代码里很常见,但很多人在这里会手滑写成dim=-1,因为注意力机制里dim=-1用多了,形成肌肉记忆。
还有一个容易踩的点:当你需要根据argmax结果生成一个 one-hot 掩码时,顺序是反着的。先从整数掩码生成 one-hot 可以用F.one_hot(pred_mask, num_classes=C),这时候pred_mask的形状是[B, H, W],生成的one_hot_mask会变成[B, H, W, C],类别通道跑到了最后。如果你在后面的代码里用了dim=1的argmax,维度就对不上了。要么提前permute(0, 3, 1, 2)把类别通道挪回中间,要么后续直接用dim=-1。
4. 避坑实录:argmax 翻车的几种典型姿势
这一节聊得都是我或者身边同事真的在代码里踩过的坑,每一个都不是“理论上的坑”,而是“上线前差点漏掉的坑”。
4.1 返回类型是 long,不是 float
torch.argmax永远返回torch.int64(也就是long)张量,而不是浮点。很多人忽略了这个细节,后面直接把这个结果拿去和浮点张量做运算,或者写进某个要求float32的数据容器,遇到类型不匹配报错一时半会反应不过来。
如果你需要把预测索引拼接成字符串、写入 CSV、或者传给某些只接受float的 API,记得.long()已经满足要求,必要时再.cpu().numpy().tolist()转成 Python 列表。
另外注意:argmax返回的索引本身是没有梯度可言的。它是在前向过程中对离散下标的选择,torch.argmax这个操作是不可导的。因此它只能出现在推理、评估和 loss 计算的“下游”,绝不能出现在需要回传梯度的网络结构中,否则梯度直接断掉。
4.2 keepdim 什么时候必须保留
torch.argmax(input, dim=1, keepdim=True)会让输出的形状从[B]变成[B, 1]。这个参数平时可加可不加,但有两个场景必须考虑。
第一个场景是做行方向的广播运算。比如你要用预测的类别索引去 gather 每个样本对应的 logits 值:
selected_logits = torch.gather(logits, dim=1, index=pred.unsqueeze(1))这里需要pred是[B, 1],所以你或者用unsqueeze(1),或者一开始就keepdim=True。
第二个场景是混合掩码。分割任务里你算出了pred_mask,你不仅需要它用来算指标,还需要它按类别生成加权 mask,keepdim能帮你省掉一次手工扩维,代码更连贯。
但反过来,keepdim=True会让输出不再是一维,如果你习惯默认输出是一维,后面的代码又是按一维写的,加上keepdim反而会造成隐形 bug。我的习惯是:默认不写keepdim,等到gather或其他需要广播形状的算子出现时,显式用unsqueeze,这样每一行的意图更清楚。
4.3 多标签分类里别用 argmax
多标签分类的场景里,一个样本可以同时属于多个类别,比如一张图既包含“天空”又包含“人”。模型输出通常经过sigmoid变成每个类别的独立概率,然后通过阈值(比如 0.5)判断哪些类别激活:
preds = (torch.sigmoid(logits) > 0.5).long() # [B, C]这时候如果你用torch.argmax(logits, dim=1),你只保留了一个概率最大的类别,而且是唯一的一个。这直接丢失了所有“次高但同样有效”的预测,明显不符合多标签任务的定义。正确的还原方式是把整数标签转回 one-hot 比较,或者用torch.where、nonzero等操作。
一句话总结:argmax 适用于“互斥的单标签问题”,不适用于“可并发的多标签问题”。别拿一把锤子去拧所有的螺丝。
4.4 NaN 和 -inf 的隐藏陷阱
logits 里如果混进了NaN,那么argmax的行为在不同版本里可能不一致,但大概率不会给你想要的结果——NaN的比较语义是未定义的,你拿到的索引常常是那个NaN本身所在的位置。你要是发现预测结果突然变成固定某几个值,先检查两点:
- 网络中间层有没有除零或溢出,导致 logits 变成了
NaN。 - 注意力 mask 里有没有把某些位置设为
-inf,如果有,argmax一般能正确避开,因为有限值永远比-inf大。但如果你用float('-inf')和float('inf')混用,麻烦就大了。
我在大规模分类模型里踩过一次:因为某个数据缺失,batch 里唯一的特征行全为 0,网络输出的一整行 logits 都是-inf,这行样本的argmax直接随机给了个 0。后续排查花了大半天,最后靠打印 logits 统计才发现数据源的问题。从此以后我在推理代码里都会加一个断言:
assert torch.isfinite(logits).all(), "logits contains non-finite values"这个习惯帮我拦下过至少三次潜在事故,代价几乎为零。
5. 常见问题速查表
下面这些是我在知乎私信、GitHub issue 和线下带实习生时被问过最多的问题,整理成自查表形式,遇到类似报错可以直接按这个顺序查。
| 问题现象 | 排查方向 |
|---|---|
| 预测结果全是同一个数(比如全是 0) | 首先打印logits.shape,确认类别维在第几维;再确认argmax的dim和类别维一致 |
报错说index out of range | 检查argmax结果的数值范围是否超出你后面gather、索引的维度上限 |
| 准确率计算结果异常高或异常低 | 检查标签是否做过 one-hot 但没有argmax回去就参与比较 |
torch.argmax和torch.max结果对不上 | 确认你是否用了同一个张量;max返回两个值,你拿的可能是 values 而不是 indices |
| 梯度回传时网络参数不更新 | 检查网络中是否调用了argmax、sort、nonzero等不可导操作 |
| ONNX 导出后推理结果和 PyTorch 不同 | ONNX 里对应的是ArgMax算子,axis参数可能映射成dim,检查导出时的 opset 版本和轴映射 |
5.1 实战中调通 argmax 的最快调试路径
当你怀疑argmax结果有问题,最快的定位方式不是瞎改dim,而是分三步走:
第一步,打印logits.shape和标签张量的shape并排对比。
print(logits.shape, labels.shape)第二步,打印一个 batch 里前 4 个样本的原始 logits 行,手工判断最大值应该在哪个索引,再和argmax输出对比。
print(logits[:4]) print(torch.argmax(logits, dim=1)[:4])第三步,对 one-hot 标签做反向转换,检查torch.argmax(one_hot_labels, dim=1)是否还原出原始整数标签。
assert torch.all(torch.argmax(F.one_hot(integer_labels, num_classes=C), dim=1) == integer_labels)第三步这个断言,建议在代码里留作单元测试。one-hot 到整数的转换稳定可靠,但正因为太稳定,所以出问题的时候往往是其他地方错了,而它是个很好的“对照基准”。
5.2 和 max、topk 的选择问题
很多时候argmax并不是唯一选择。如果你只需要最大值的索引,argmax是最轻量的。如果你想同时拿到最大概率值和索引,就用torch.max(probs, dim=1)。如果你还想要第二、第三大,torch.topk(probs, k=5, dim=1)是标准姿势。
社区里有个常见的误解是“先 softmax 再 argmax 会更准”。我郑重说明一下,这不会更准,也不会有任何数值上的差别,因为 softmax 是严格单调的,最大值在哪一个位置,softmax 前后一模一样。唯一的区别是,softmax 之后你能顺便拿到每个类别的概率,方便你算置信度、画 PR 曲线。不要在“准不准”这个维度上纠结,逻辑上说不通。
5.3 从经验里总结的几个小习惯
写代码时我会默认遵守一套关于argmax的小习惯,逐个分享一下。
第一,凡是argmax之后参与指标计算的,我都会让代码注释里明确写下形状变化:# [B, C] -> [B]。一行注释,省得后面 review 的人反复推敲。
第二,凡是 one-hot 和整数标签互相转换,我都在项目里抽成独立函数,不希望散落在各个文件里。这样一旦 debug,全项目只需要改一个地方。
第三,凡是新接手一个模型,第一步是跑一个最小样例,验证输入输出形状是否符合预期,再把argmax接上去。这个流程特别笨,但特别有效。你只要跑假数据确认argmax(dim=1)的 shape 是[B],你就已经排除了大部分维度问题。
第四,留意 batch 维度为 1 的情况。[1, C]和[C]长得像,但argmax返回的 shape 一个是一维[1],一个是标量[],返回值类型差了十万八千里,后续代码如果直接操作标量,很容易炸。碰到这种边界情况,稳妥做法是torch.argmax(logits, dim=-1).squeeze(0),把多余维度去掉。
我个人在实际项目里最深的体会是:torch.argmax本身只是一个几十行的 API,它的学习成本低到几乎可以忽略,但所有真正的复杂度都藏在“张量的形状约定”和“标签的表示形式”上。你只要习惯性地从[B, C]的视角去看待模型输出、从“one-hot 每一行只有一个 1”的角度去理解标签,那么这个函数和 one-hot 与整数标签之间的那点关系,就再也不会给你制造惊喜。最后再分享一个建议:在你项目的utils.py里放一个one_hot_to_label和label_to_one_hot的对偶函数,一个是F.one_hot,一个是torch.argmax(..., dim=1),成对出现。配合单元测试,这个组合会让你在做分类任务时省掉至少 80% 的标签转换排查时间。