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

如何在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]])

原理说明

  1. 偏移量row_offsets和col_offsets定义了3x3补丁相对于中心位置的偏移;
  2. 通过[:, None, None]扩展a、b、c的维度,让它们能和偏移量进行广播,生成每个样本对应的3x3网格索引;
  3. 直接用广播后的索引从data中取值,自动拼接成(k, 3, 3)的结果。

补充说明

  • 该方案基于PyTorch原生操作,效率远高于循环;
  • 已假设输入的a、b、c均在有效索引范围内(无需处理边界越界);
  • 若无需梯度,可在操作后添加.detach()或转换为NumPy数组,但上述PyTorch方案已足够高效。

内容的提问来源于stack exchange,提问作者Nagabhushan S N

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.27 07:43:17