PyTorch中如何实现支持起始索引张量的类torch.narrow切片操作?
向量化实现方案
你需要的逐行不同起始点的固定长度切片,完全可以用PyTorch内置的向量化算子实现,不需要写Python层for循环,性能可以达到底层并行计算的最优水平。
核心思路
对形状为(M, N)的输入张量,要取每行长为length的切片,本质是对第i行,取列位置为start_indices[i] + k的元素,其中k从0到length-1。只要构造出形状为(M, length)的列索引矩阵,就可以直接用内置算子批量取数,整个过程没有串行循环。
最优实现(基于torch.gather)
torch.gather是专门用于沿指定维度按索引批量取数的算子,不需要额外构造行索引,开销最低,代码如下:
import torch def variable_start_narrow(input: torch.Tensor, dim: int, start: torch.Tensor, length: int): assert dim == 1, "当前实现仅支持dim=1的二维张量场景" # 构造偏移量,广播后得到每个切片对应的列索引 offsets = torch.arange(length, device=input.device, dtype=start.dtype) col_indices = start.unsqueeze(-1) + offsets.unsqueeze(0) # 沿指定维度按索引取数 return torch.gather(input, dim=dim, index=col_indices) # 测试示例 if __name__ == "__main__": start_indices = torch.tensor([3, 1, 2]) dataset = torch.tensor([ [3, 5, 3, 4, 8, 0, 1], [3, 9, 7, 2, 7, 3, 7], [6, 0, 2, 3, 0, 2, 5] ]) res = variable_start_narrow(dataset, dim=1, start=start_indices, length=4) print(res) # 输出: # tensor([[4, 8, 0, 1], # [9, 7, 2, 7], # [2, 3, 0, 2]])
性能与边界说明
- 索引矩阵的大小仅为
(M, length),按你给出的场景M~1e3、length=4计算,索引总元素数仅为4000,内存/显存开销可以忽略,哪怕N>1e6也不会有额外内存压力。 - 所有计算都走PyTorch底层C++/CUDA并行实现,比Python层for循环快2个数量级以上,完全可以满足性能要求。
- 越界行为和原生
torch.narrow一致:当某行的start[i] + length > N时,会直接抛出索引越界错误,符合你的预期。
其他可选实现(高级索引)
如果不想用torch.gather,也可以通过构造行、列双索引用原生高级索引实现,逻辑更直观但性能略低于gather版本:
offsets = torch.arange(length, device=input.device) col_indices = start.unsqueeze(-1) + offsets.unsqueeze(0) row_indices = torch.arange(input.shape[0], device=input.device).unsqueeze(-1).expand_as(col_indices) result = input[row_indices, col_indices]
注:你之前尝试的
torch.index_select仅支持传入一维索引,无法处理逐行不同起始位置的切片需求,torch.gather就是对应这种可变索引取数场景的官方实现。
内容的提问来源于stack exchange,提问作者hmomin
相关产品推荐
相关产品推荐

