
本文介绍如何仅用一次张量操作,从三维张量中按行分别提取不同起始位置、相同长度的切片(如第0行取索引5–9、第1行取10–14),避免分步索引+拼接,提升性能与代码简洁性。
本文介绍如何仅用一次张量操作,从三维张量中按行分别提取不同起始位置、相同长度的切片(如第0行取索引5–9、第1行取10–14),避免分步索引+拼接,提升性能与代码简洁性。
在 PyTorch 中,当需要从高维张量中按样本(即沿某维度)差异化地提取固定长度的连续子序列时(例如:对 batch 中每个样本独立指定起始索引,截取相同长度的 slice),若采用 arr[indices] 后再逐个切片并 torch.cat,不仅代码冗余,还引入中间张量和显式循环逻辑,不利于向量化与 GPU 利用率优化。
核心思路是:将“每行不同起始点 + 固定长度”的索引需求,转化为一个二维索引网格(shape: [N, L]),再通过 torch.gather 沿指定维度完成批量收集。
以原始问题为例:
import torch arr = torch.randint(0, 9, (100, 50, 3)) # shape: [100, 50, 3] indices = torch.tensor([5, 55]) # 选第5和第55个样本 → shape: [2] start_offsets = torch.tensor([5, 10]) # 第0个样本从 idx=5 开始,第1个从 idx=10 开始 length = 5 # 每个样本取5个连续元素(即 5:10 和 10:15)
✅ 正确的单步实现如下:
# Step 1: 构建相对偏移量(0 到 length-1) rel_idx = torch.arange(length, device=arr.device) # [0, 1, 2, 3, 4] # Step 2: 广播生成绝对索引矩阵:shape [N, length] # start_offsets[:, None] → [[5], [10]], rel_idx[None, :] → [[0,1,2,3,4]] abs_idx = start_offsets[:, None] + rel_idx[None, :] # shape: [2, 5] # Step 3: 扩展至目标维度(适配最后一维 channel=3) # gather 要求 index 与 input 在非 gather 维度上 shape 一致 # 这里沿 dim=1 gather,故 abs_idx 需扩展为 [2, 5, 1],再广播为 [2, 5, 3] abs_idx_exp = abs_idx.unsqueeze(-1).expand(-1, -1, arr.size(-1)) # [2, 5, 3] # Step 4: 先取样本子集,再 gather partial_arr = arr[indices] # shape: [2, 50, 3] final_result = torch.gather(partial_arr, dim=1, index=abs_idx_exp) # shape: [2, 5, 3] print(final_result.shape) # torch.Size([2, 5, 3])
? 关键说明:
- torch.gather(input, dim, index) 要求 index 的 shape 与 input 在除 dim 外的所有维度上严格匹配;index[i][j][k] 表示在 input 的 dim 维上取第 index[i][j][k] 个元素。
- 此处 dim=1(即 50 维),因此 abs_idx_exp 必须为 [2, 5, 3] —— 前两维对应 partial_arr 的 [2, 50],最后一维 3 由 expand 自动广播,确保每个通道都使用相同的行内索引。
- 若需最终合并为 [10, 3] 形状(即取消 batch 维),可追加 final_result.view(-1, 3)。
⚠️ 注意事项:
- 该方法要求所有行提取长度一致(即 length 固定),否则无法构造统一 shape 的 index 张量;
- 索引值必须在有效范围内(0 ≤ abs_idx
- 对于不规则长度需求(如每行长度不同),应改用 torch.nn.utils.rnn.pad_packed_sequence 或自定义 torch.vmap(PyTorch 2.0+)方案。
此技巧广泛适用于序列标注、动态窗口采样、Batch 内局部注意力掩码构建等场景,在保持计算图完整性和 GPU 友好性的同时,显著提升索引操作的表达力与执行效率。











