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

PyTorch中如何用形状[b,k]的张量索引[b,m,n]形状的张量

在PyTorch中实现按batch维度索引张量

你可以通过以下几种方法实现需求,核心是让索引张量的形状适配PyTorch的索引规则:

方法1:使用torch.take_along_dim(推荐,PyTorch 1.10+)

这个函数专门用于沿指定维度选取元素,语义更直观,不需要复杂的形状扩展:

import torch

# 构造示例张量
b, m, n, k = 2, 5, 3, 3
A = torch.randn(b, m, n)  # shape: [b, m, n]
B = torch.randint(0, m, (b, k))  # shape: [b, k],索引值范围0~m-1

# 将索引张量扩展为[b, k, 1],自动广播匹配n维度
indices = B.unsqueeze(-1)
result = torch.take_along_dim(A, indices, dim=1)
# result形状为[b, k, n],符合要求

方法2:使用torch.gather

需要将索引张量扩展为与输入张量A的形状匹配(仅在索引维度dim=1上长度不同):

# 将B扩展为[b, k, n],每个n维度位置复用相同的索引值
indices = B.unsqueeze(-1).expand(-1, -1, n)
result = torch.gather(A, dim=1, index=indices)
# result形状为[b, k, n]

方法3:使用高级索引

通过构造batch维度的索引,直接进行多维索引:

# 构造batch索引,形状[b, k],每个batch对应自身的索引
batch_idx = torch.arange(b).unsqueeze(1).expand(-1, k)
result = A[batch_idx, B, :]
# result形状为[b, k, n]

说明

  • torch.index_select/torch.take仅支持全局一维索引,无法直接处理batch维度的独立索引,所以不适用。
  • 之前使用torch.gather失败是因为未将索引张量扩展到与输入张量的所有非索引维度匹配,扩展后即可正常工作。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.25 16:43:24