You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

PyTorch中obj_ids[:, None, None]与mask比较运算的运行原理疑问

代码运行逻辑解析

这行代码的核心是借助PyTorch的广播机制和维度扩展操作,一次性批量生成多个目标类别的分割掩码,具体运行逻辑拆解如下:

1. 初始变量维度

  • 示例中的mask是3×3的二维单通道分割标签,形状为(H, W)(H为高度,W为宽度,这里H=3、W=3)
  • obj_ids是需要提取的目标类别id列表,示例值为[1,2],是长度为2的一维张量,初始形状为(N,)(N为目标类别数量,这里N=2)

2. [:, None, None]的维度扩展作用

PyTorch中None等价于torch.unsqueeze,作用是在指定位置插入长度为1的新维度:
obj_ids[:, None, None]会将原本形状为(2,)的一维张量,依次在第1、第2个维度位置插入新维度,最终得到形状为(2, 1, 1)的三维张量,示例扩展后的取值为:

[[[1]], [[2]]]

3. 相等判断的广播运算过程

形状为(2, 1, 1)的扩展后obj_ids,和形状为(3, 3)的mask做==相等判断时,会自动触发广播机制:

  1. 首先自动将mask的维度补全为3维,形状变为(1, 3, 3)
  2. 两个张量都会自动把维度为1的位置复制扩展到和对方匹配的尺寸:(2, 1, 1)扩展为(2, 3, 3),(1, 3, 3)也扩展为(2, 3, 3)
  3. 逐元素对比两个同尺寸张量的对应位置值是否相等,最终输出布尔类型的结果张量,形状正好是(2, 3, 3),也就是你观察到的2×原矩阵尺寸的结果

4. 示例结果验证

代入你给出的示例值计算,最终得到的masks包含两个3×3的布尔掩码:

  • 第0个通道对应类别1:所有mask取值为1的位置为True,其余为False,结果为[[False, False, False], [True, False, False], [True, False, True]]
  • 第1个通道对应类别2:所有mask取值为2的位置为True,其余为False,结果为[[False, False, False], [False, False, False], [False, True, False]]

内容的提问来源于stack exchange,提问作者Han

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.09.27 22:15:08