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
相关产品推荐
相关产品推荐

