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

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.30 14:45:36