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做==相等判断时,会自动触发广播机制:
- 首先自动将
mask的维度补全为3维,形状变为(1, 3, 3) - 两个张量都会自动把维度为1的位置复制扩展到和对方匹配的尺寸:
(2, 1, 1)扩展为(2, 3, 3),(1, 3, 3)也扩展为(2, 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
相关产品推荐
相关产品推荐

