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

PyTorch中按行索引模运算规则提取张量数据的实现方法

PyTorch张量按模规则提取行通用实现

你可以用基于索引掩码的通用实现,兼容动态张量尺寸、自动适配CPU/GPU设备,无需硬编码每个i的取值:

基础实现(输出与原张量维度数一致,行按余数顺序排列)

import torch

def extract_rows_by_mod(X, n, dim=2):
    # dim参数为行所在的维度索引,你示例中的行维度对应索引为2
    row_count = X.size(dim)
    # 生成与输入张量同设备的行索引序列
    row_indices = torch.arange(row_count, device=X.device)
    # 按索引模n的结果排序,输出顺序为i=0对应所有行、i=1对应所有行...i=n-1对应所有行
    sorted_indices = row_indices[torch.argsort(row_indices % n)]
    # 按排序后的索引提取行
    return X.index_select(dim=dim, index=sorted_indices)

示例测试

针对你给出的形状为(1, 3, 16, 9)的输入张量,n=4时调用方式如下:

X = torch.randn(1, 3, 16, 9)
Y = extract_rows_by_mod(X, n=4)
print(Y.shape) # 输出:torch.Size([1, 3, 16, 9])

分组实现(输出新增维度存放不同i的行组)

如果你需要将不同i对应的行单独分组,额外新增一个维度存放i维度,可以使用以下实现:

def extract_rows_by_mod_grouped(X, n, dim=2):
    row_count = X.size(dim)
    row_indices = torch.arange(row_count, device=X.device)
    # 按模n结果分组提取索引,再展开索引进行提取,最后拆分出分组维度
    grouped_indices = torch.stack([row_indices[row_indices % n == i] for i in range(n)])
    return X.index_select(dim=dim, index=grouped_indices.flatten()).unflatten(dim, (n, row_count//n))

示例测试

Y_grouped = extract_rows_by_mod_grouped(X, n=4)
print(Y_grouped.shape) # 输出:torch.Size([1, 3, 4, 4, 9]),其中第三个维度为i从0到3的分组,第四个维度为每个分组对应的4行

方案优势

  • 无需硬编码i的取值,自动适配任意n的设置(需保证n可以整除对应行维度的长度,若不整除可自行补充截断或补全逻辑)
  • 自动适配输入张量的设备,CPU/GPU场景都可直接使用
  • 可通过修改dim参数适配任意维度的索引提取,不局限于你示例中的维度位置

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.09.26 22:36:07