ARTICLE DETAIL

资讯详情

深耕网站建设与运营推广的一线实战洞察。

torch.argmax 的 dim=1 含义与 one-hot 转整数标签实战

torch.argmax 的 dim=1 含义与 one-hot 转整数标签实战 torch.argmax 这个函数几乎是每个用 PyTorch 做分类任务的人最早接触到的 API 之一。你一定会遇到这样的场景模型输出一个形状为[batch_size, num_classes]的分数矩阵你想知道每个样本到底被预测成了哪个类别于是敲下torch.argmax(output, dim1)得到一串整数标签。而另一边如果你的训练数据是从 one-hot 格式存的比如[0, 0, 1, 0]这种向量你想把它还原成整数2最直接的办法同样是用torch.argmax(one_hot_vector, dim1)。看起来都是“沿着第 1 维找最大值索引”这一件事但真到代码里维度选错、形状搞混、one-hot 转不出来整数标签的问题我在实际项目里见过太多次了。这篇文章就把dim1的语义、它与 one-hot 转整数标签之间的对应关系、以及实战中常见的坑一次性讲透适合刚接触 PyTorch 的初学者也适合写了多年模型但偶尔被维度问题绊一脚的工程师。1. 先搞懂 dim1张量的“轴”就是业务语义的开瓶器很多人对argmax的理解停留在“取最大值的下标”但一碰到多维张量就开始懵根源是把“张量的维度”和“业务里的维度”搞混了。1.1 分类模型的默认形状约定[批量, 类别数]深度学习里处理分类任务时模型最后一层输出的 logits 一般都约定成[batch_size, num_classes]。比如批量大小为 8类别数为 10输出就是[8, 10]。在这个约定下第 0 维是“样本”维度代表这个批次里有几张图、几句话、几段序列。第 1 维是“类别”维度代表每个样本分别在 10 个类别上的得分。这里的第 1 维就是dim1对应的那个轴。torch.argmax(output, dim1)的意思就是对每个样本固定第 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, dim1)的结果是[2, 0]。怎么来的第一行里最大的是索引 2 对应的 3.0第二行里最大的是索引 0 对应的 4.0。你注意看dim1说的是“消掉第 1 维”也就是把[2, 3]变成[2]操作方向是横着对每一行内部做比较。如果你反过来用dim0那就是纵向比较对第 0 个类别取两个样本的最大值索引对第 1 个类别取两个样本的最大值索引对第 2 个类别取两个样本的最大值索引结果变成[1, 0, 0]这样的长度 3 的张量。这个结果在分类任务里几乎没有任何业务意义因为纵向比较的是“不同样本在同一个类别上的得分”这不是我们想要的东西。所以理解dim的核心心法只有一个dim 指定的是你要“消掉”哪个轴。剩下的轴就是你结果里保留的维度。1.3 dim1、dim-1 和 dim0 的惯用法实际项目里dim1和dim-1在二维张量上是同一个意思因为-1代表最后一个轴。但如果你的张量是三维的情况就不一样了。比如图像语义分割模型的输出形状是[B, C, H, W]类别维度 C 在中间这时你依然应该用dim1而不是dim-1。dim-1在这个场景下对应的是 W 维度会直接让你得到一堆坐标值而不是类别索引。我见过的另一种混乱来自“通道最后排”的格式比如某些模型输出[B, H, W, C]这时候才需要用dim-1。判断标准永远不是“别人说用哪个”而是“你的张量布局里类别到底在第几维”。要么在模型代码里统一约定输出格式要么在写argmax之前用一行注释写明形状# [B, C, H, W], C dimension num_classes这条注释能省掉很多排查时间。2. one-hot 与整数标签一对可以互相翻译的兄弟搞定了dim1的语义再来盘一盘 one-hot。标题里把它和argmax绑在一起是因为“one-hot 转整数标签”这件事数学上跟argmax(dim1)几乎是同一个操作。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_classes5) # 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, dim1)由于 one-hot 每一行只有一个 1其余全是 0最大值必然出现在那个 1 的位置上argmax(dim1)返回的索引就是原来的整数标签。这是完全无损的逆变换也是我在所有项目里推荐用、也唯一推荐用的逆变换方式。为什么不推荐写循环遍历因为向量化操作又快又简洁还能保持梯度链虽然 one-hot 本来也不需要梯度。为什么不推荐用torch.nonzero你当然可以用torch.nonzero(one_hot)[:, 1]拿到每个样本的索引但写法绕、易读性差、要处理空行和维度变化不如argmax一行清爽。2.3 三种常见 one-hot 生成方式与统一出口生成 one-hot 的方法不止F.one_hot一种我在不同代码库里看到过各种写法生成方式代码示例适用场景F.one_hotF.one_hot(labels, num_classesC)最标准推荐优先使用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, axis1)NumPy 里的np.argmax(one_hot, axis1)都一个道理。3. 实操把 dim1 用到完整的分类流程里理论说完了接下来看几个真实项目里会碰到的代码场景。dim1和 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, dim1) confidence, pred torch.max(probs, dim1)写法二pred torch.argmax(logits, dim1)它们其实不冲突。softmax是单调递增的归一化函数所以对同一个logits张量argmax(softmax(logits), dim1)和argmax(logits, dim1)的结果完全一致。区别在于你有没有想要置信度。如果你只想要类别编号直接argmax就够了省一次 softmax 的指数运算如果你还想要这个预测有多大概率那才需要先softmax再取max或者直接取对应位置的概率。对于单标签分类我推荐写成with torch.no_grad(): logits model(x) probs torch.softmax(logits, dim1) pred torch.argmax(probs, dim1) max_conf, _ torch.max(probs, dim1)其实这里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, dim1)转回整数再比较。我遇到过一个典型的写法错误有人把 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]这时dim1同样是类别维度。逐像素的分类预测就变成了pred_mask torch.argmax(logits, dim1) # 结果是 [B, H, W]每个位置上都是该像素的类别索引。这个操作在代码里很常见但很多人在这里会手滑写成dim-1因为注意力机制里dim-1用多了形成肌肉记忆。还有一个容易踩的点当你需要根据argmax结果生成一个 one-hot 掩码时顺序是反着的。先从整数掩码生成 one-hot 可以用F.one_hot(pred_mask, num_classesC)这时候pred_mask的形状是[B, H, W]生成的one_hot_mask会变成[B, H, W, C]类别通道跑到了最后。如果你在后面的代码里用了dim1的argmax维度就对不上了。要么提前permute(0, 3, 1, 2)把类别通道挪回中间要么后续直接用dim-1。4. 避坑实录argmax 翻车的几种典型姿势这一节聊得都是我或者身边同事真的在代码里踩过的坑每一个都不是“理论上的坑”而是“上线前差点漏掉的坑”。4.1 返回类型是 long不是 floattorch.argmax永远返回torch.int64也就是long张量而不是浮点。很多人忽略了这个细节后面直接把这个结果拿去和浮点张量做运算或者写进某个要求float32的数据容器遇到类型不匹配报错一时半会反应不过来。如果你需要把预测索引拼接成字符串、写入 CSV、或者传给某些只接受float的 API记得.long()已经满足要求必要时再.cpu().numpy().tolist()转成 Python 列表。另外注意argmax返回的索引本身是没有梯度可言的。它是在前向过程中对离散下标的选择torch.argmax这个操作是不可导的。因此它只能出现在推理、评估和 loss 计算的“下游”绝不能出现在需要回传梯度的网络结构中否则梯度直接断掉。4.2 keepdim 什么时候必须保留torch.argmax(input, dim1, keepdimTrue)会让输出的形状从[B]变成[B, 1]。这个参数平时可加可不加但有两个场景必须考虑。第一个场景是做行方向的广播运算。比如你要用预测的类别索引去 gather 每个样本对应的 logits 值selected_logits torch.gather(logits, dim1, indexpred.unsqueeze(1))这里需要pred是[B, 1]所以你或者用unsqueeze(1)或者一开始就keepdimTrue。第二个场景是混合掩码。分割任务里你算出了pred_mask你不仅需要它用来算指标还需要它按类别生成加权 maskkeepdim能帮你省掉一次手工扩维代码更连贯。但反过来keepdimTrue会让输出不再是一维如果你习惯默认输出是一维后面的代码又是按一维写的加上keepdim反而会造成隐形 bug。我的习惯是默认不写keepdim等到gather或其他需要广播形状的算子出现时显式用unsqueeze这样每一行的意图更清楚。4.3 多标签分类里别用 argmax多标签分类的场景里一个样本可以同时属于多个类别比如一张图既包含“天空”又包含“人”。模型输出通常经过sigmoid变成每个类别的独立概率然后通过阈值比如 0.5判断哪些类别激活preds (torch.sigmoid(logits) 0.5).long() # [B, C]这时候如果你用torch.argmax(logits, dim1)你只保留了一个概率最大的类别而且是唯一的一个。这直接丢失了所有“次高但同样有效”的预测明显不符合多标签任务的定义。正确的还原方式是把整数标签转回 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, dim1)[:4])第三步对 one-hot 标签做反向转换检查torch.argmax(one_hot_labels, dim1)是否还原出原始整数标签。assert torch.all(torch.argmax(F.one_hot(integer_labels, num_classesC), dim1) integer_labels)第三步这个断言建议在代码里留作单元测试。one-hot 到整数的转换稳定可靠但正因为太稳定所以出问题的时候往往是其他地方错了而它是个很好的“对照基准”。5.2 和 max、topk 的选择问题很多时候argmax并不是唯一选择。如果你只需要最大值的索引argmax是最轻量的。如果你想同时拿到最大概率值和索引就用torch.max(probs, dim1)。如果你还想要第二、第三大torch.topk(probs, k5, dim1)是标准姿势。社区里有个常见的误解是“先 softmax 再 argmax 会更准”。我郑重说明一下这不会更准也不会有任何数值上的差别因为 softmax 是严格单调的最大值在哪一个位置softmax 前后一模一样。唯一的区别是softmax 之后你能顺便拿到每个类别的概率方便你算置信度、画 PR 曲线。不要在“准不准”这个维度上纠结逻辑上说不通。5.3 从经验里总结的几个小习惯写代码时我会默认遵守一套关于argmax的小习惯逐个分享一下。第一凡是argmax之后参与指标计算的我都会让代码注释里明确写下形状变化# [B, C] - [B]。一行注释省得后面 review 的人反复推敲。第二凡是 one-hot 和整数标签互相转换我都在项目里抽成独立函数不希望散落在各个文件里。这样一旦 debug全项目只需要改一个地方。第三凡是新接手一个模型第一步是跑一个最小样例验证输入输出形状是否符合预期再把argmax接上去。这个流程特别笨但特别有效。你只要跑假数据确认argmax(dim1)的 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(..., dim1)成对出现。配合单元测试这个组合会让你在做分类任务时省掉至少 80% 的标签转换排查时间。
返回列表