如何在PyTorch中基于其他张量值无循环提取多维张量的3x3补丁
PyTorch无循环提取指定位置的3x3补丁张量
需求说明
现有4个PyTorch张量:
data:形状(l, m, n)a:形状(k,),long类型,指定第0维度的选取索引b:形状(k,),long类型,指定第1维度的中心位置c:形状(k,),long类型,指定第2维度的中心位置
需要无循环实现:从data中按a选取第0维度元素,在第1、2维度围绕b、c提取3x3补丁,最终得到形状(k, 3, 3)的张量。
错误示例
直接使用切片会触发类型错误:
import torch data = torch.arange(200).reshape((2, 10, 10)) a = torch.Tensor([1, 0, 1, 1, 0]).long() b = torch.Tensor([5, 6, 3, 4, 7]).long() c = torch.Tensor([4, 3, 7, 6, 5]).long() data1 = data[a, b-1:b+1, c-1:c+1] # 报错 # TypeError: only integer tensors of a single element can be converted to an index
预期输出
data1[0] = [[143,144,145],[153,154,155],[163,164,165]] data1[1] = [[52,53,54],[62,63,64],[72,73,74]] data1[2] = [[126,127,128],[136,137,138],[146,147,148]] # 后续元素以此类推
解决方案
利用PyTorch的广播机制生成所有补丁的索引,无需循环:
import torch data = torch.arange(200).reshape((2, 10, 10)) a = torch.Tensor([1, 0, 1, 1, 0]).long() b = torch.Tensor([5, 6, 3, 4, 7]).long() c = torch.Tensor([4, 3, 7, 6, 5]).long() # 生成3x3补丁的行/列偏移量 row_offsets = torch.arange(-1, 2) # [-1, 0, 1] col_offsets = torch.arange(-1, 2) # 扩展维度实现广播,生成每个样本的3x3行/列索引 rows = b[:, None, None] + row_offsets[None, :, None] cols = c[:, None, None] + col_offsets[None, None, :] # 提取目标张量,自动广播索引维度 result = data[a[:, None, None], rows, cols] print(result.shape) # torch.Size([5, 3, 3]) print(result[0]) # tensor([[143, 144, 145], # [153, 154, 155], # [163, 164, 165]])
原理说明
- 偏移量
row_offsets和col_offsets定义了3x3补丁相对于中心位置的偏移; - 通过
[:, None, None]扩展a、b、c的维度,让它们能和偏移量进行广播,生成每个样本对应的3x3网格索引; - 直接用广播后的索引从
data中取值,自动拼接成(k, 3, 3)的结果。
补充说明
- 该方案基于PyTorch原生操作,效率远高于循环;
- 已假设输入的
a、b、c均在有效索引范围内(无需处理边界越界); - 若无需梯度,可在操作后添加
.detach()或转换为NumPy数组,但上述PyTorch方案已足够高效。
内容的提问来源于stack exchange,提问作者Nagabhushan S N
相关产品推荐
相关产品推荐

