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

PyTorch中用3D张量索引选取4D张量对应值的技术问题

PyTorch中用3D坐标张量索引4D张量的解决方案

问题回顾

给定4D张量 possible_values(形状 torch.Size([2, 5, 5, 4])),维度定义为:

  • dim 0: batch
  • dim 1: x_axis
  • dim 2: y_axis
  • dim 3: 坐标(x_i,y_j)对应的特征值

同时有3D坐标张量 coordinates(形状 torch.Size([2, 5, 2])),维度定义为:

  • dim 0: batch
  • dim 1: (x,y)坐标序列
  • dim 2: 单个坐标的(x,y)值

需要从每个batch中,选取coordinates指定坐标对应的特征值(即dim3的4个值)。

关键注意点

示例中的坐标是1-based(比如[1,5]),但PyTorch张量采用0-based索引,必须先将坐标转换为0-based,否则会触发索引越界错误。

解决方案代码

import torch

# 初始化示例张量
possible_values = torch.randn(2, 5, 5, 4)  # [batch, x_axis, y_axis, feature]
coordinates = torch.tensor([
    [[1,5], [3,3], [2,4], [1,3], [2,3]],
    [[1,5], [4,3], [2,1], [5,3], [5,3]]
])

# 1. 转换为0-based索引
coords_0based = coordinates - 1

# 2. 拆分x、y坐标分量
x_idx = coords_0based[..., 0]  # 形状 [2,5],对应每个batch的x坐标
y_idx = coords_0based[..., 1]  # 形状 [2,5],对应每个batch的y坐标

# 3. 生成batch维度的索引,确保坐标与所属batch对应
batch_idx = torch.arange(possible_values.size(0))[:, None].repeat(1, x_idx.size(1))

# 4. 执行高级索引取值
selected_features = possible_values[batch_idx, x_idx, y_idx]

# 验证结果形状:预期为 [2,5,4]
print(selected_features.size())  # 输出 torch.Size([2, 5, 4])

原理说明

PyTorch的高级索引支持同时使用多个同形状的张量对不同维度进行索引:

  • batch_idx 对应batch维度,确保每个坐标从对应的batch中取值
  • x_idx 和 y_idx 分别对应x_axis和y_axis维度,定位具体坐标位置
  • 索引后自动保留feature维度(dim3),最终得到每个坐标对应的4维特征值

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.08 05:10:32