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

如何在PyTorch中高效按特定值分割(batch, seq_len)形状的张量?

PyTorch高效按1的列索引分割张量方案

问题描述

给定形状为(batch, seq_len)的张量:

X = [[0, 0, 0, 1, 0, 0, 0, 0],
     [0, 0, 1, 0, 0, 0, 0, 0],
     [0, 0, 1, 0, 0, 1, 0, 0]]

需要按每行中值为1的列索引分割张量,得到如下输出:

X^ = [[0, 0, 0, 1],
      [0, 0, 0, 0],
      [0, 0, 1],
      [0, 0, 0, 0, 0],
      [0, 0, 1],
      [0, 0, 1],
      [0, 0]] 

手动逐行循环处理效率低下,尤其在GPU环境中,需要PyTorch下的最优实现方法。

最优实现思路

利用PyTorch的向量化张量操作替代Python循环,充分利用GPU并行计算能力。核心步骤是:定位所有1的位置→生成分割段的起止边界→提取片段并统一填充长度。

代码实现(基础版本)

import torch

# 输入张量
X = torch.tensor([
    [0, 0, 0, 1, 0, 0, 0, 0],
    [0, 0, 1, 0, 0, 0, 0, 0],
    [0, 0, 1, 0, 0, 1, 0, 0]
])
batch_size, seq_len = X.shape
device = X.device

# 1. 获取每行中1的列索引,末尾添加seq_len作为该行最后一段的结束
ones_indices = torch.nonzero(X == 1, as_tuple=False).split(batch_size, dim=0)
split_points = [torch.cat([idx[:, 1], torch.tensor([seq_len], device=device)]) for idx in ones_indices]

# 2. 生成所有分割段的起止索引
segments = []
for points in split_points:
    starts = torch.cat([torch.tensor([0], device=device), points[:-1] + 1])
    ends = points
    segments.extend(torch.stack([starts, ends], dim=1))

# 3. 提取片段并填充到统一长度
max_segment_len = max((end - start).item() for start, end in segments)
result = []
for row_idx, (start, end) in enumerate(segments):
    # 确定当前片段所属的原行
    original_row = row_idx // len(split_points[row_idx // len(split_points)])
    segment = X[original_row, start:end]
    # 填充0至最大长度
    padded_segment = torch.cat([segment, torch.zeros(max_segment_len - len(segment), dtype=X.dtype, device=device)])
    result.append(padded_segment)

# 转换为最终张量
result_tensor = torch.stack(result)
print(result_tensor)

完全向量化优化版本(适合大规模GPU场景)

如果batch规模很大,可进一步消除所有Python循环,全程用张量操作实现:

import torch

X = torch.tensor([
    [0, 0, 0, 1, 0, 0, 0, 0],
    [0, 0, 1, 0, 0, 0, 0, 0],
    [0, 0, 1, 0, 0, 1, 0, 0]
])
batch_size, seq_len = X.shape
device = X.device

# 定位所有1的位置(行索引、列索引)
row_ids, col_ids = torch.nonzero(X == 1, as_tuple=True)

# 添加每行的结束标记(seq_len位置)
row_ids_ext = torch.cat([row_ids, torch.arange(batch_size, device=device)])
col_ids_ext = torch.cat([col_ids, torch.full((batch_size,), seq_len, device=device)])

# 按行排序,确保分割段顺序正确
sort_keys = row_ids_ext * seq_len + col_ids_ext
sorted_indices = torch.argsort(sort_keys)
sorted_rows = row_ids_ext[sorted_indices]
sorted_cols = col_ids_ext[sorted_indices]

# 生成所有分割段的起止位置(扁平化索引)
starts_flat = torch.cat([torch.zeros(1, dtype=torch.long, device=device), sorted_cols[:-1] + 1])
starts_flat = sorted_rows * seq_len + starts_flat
ends_flat = sorted_rows * seq_len + sorted_cols

# 计算所有段的长度,确定最大长度用于填充
segment_lens = ends_flat - starts_flat
max_len = segment_lens.max().item()

# 提取所有片段并填充
flattened_X = X.flatten()
result = torch.zeros(len(starts_flat), max_len, dtype=X.dtype, device=device)
# 利用高级索引填充
mask = torch.arange(max_len, device=device)[None, :] < segment_lens[:, None]
result[mask] = flattened_X[torch.repeat_interleave(starts_flat, segment_lens) + torch.cat([torch.arange(l) for l in segment_lens])]

print(result)

方案说明

  • 基础版本保留少量循环,但循环次数仅为分割段总数(远小于原张量元素数),GPU上效率依然很高。
  • 完全向量化版本通过张量操作完成所有步骤,无Python循环,适合超大规模batch的场景,能最大化GPU并行利用率。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.26 16:02:50