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

如何在PyTorch中无需循环从1D张量按起止索引提取切片

如何在PyTorch中从1D张量批量生成等长切片的2D张量

给定1D数据张量和批量切片的起始索引(切片长度固定):

import torch
data = torch.arange(10)
starts = torch.tensor([0, 3, 4, 1])
# 切片长度固定为2,ends = starts + 2

目标是不通过循环,直接生成如下2D张量:

tensor([[0, 1],
        [3, 4],
        [4, 5],
        [1, 2]])

问题原因

直接使用data[starts:ends]会报错,因为PyTorch的切片语法仅支持单个整数或单元素张量作为切片的起始/结束位置,不支持批量的起始/结束索引张量。

解决方案:利用广播生成批量索引

因为所有切片长度固定,我们可以通过广播机制生成所有切片的索引矩阵,再直接索引原张量:

slice_length = 2
# 生成每个切片内的偏移量:[0, 1],扩展维度以支持广播
offsets = torch.arange(slice_length).unsqueeze(0)
# 将starts扩展维度后与偏移量广播,得到(4,2)的索引矩阵
indices = starts.unsqueeze(1) + offsets
# 索引原张量得到结果
dataSlices = data[indices]

运行结果:

>>> dataSlices
tensor([[0, 1],
        [3, 4],
        [4, 5],
        [1, 2]])

另一种方法:使用torch.gather

如果需要更灵活的索引场景,也可以用torch.gather实现:

indices = starts.unsqueeze(1) + torch.arange(slice_length)
dataSlices = torch.gather(data.unsqueeze(0).repeat(len(starts), 1), 1, indices)

不过这种方法需要先重复原张量,效率略低于第一种广播索引的方式。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.31 17:50:34