PyTorch中如何在指定位置为2D张量添加零行?
在PyTorch中给2D张量指定位置插入零行
你可以通过以下方式实现指定位置插入零行,这里提供两种实用方案:
方案一:通过索引映射填充
先创建全零的目标张量,再计算原张量每行在目标张量中的对应位置,将原数据填充到正确位置:
import torch x = torch.tensor([[1,1,1], [2,2,2], [3,3,3], [4,4,4]]) # 指定要插入零行的目标索引 insert_indices = [1, 3] # 构建目标形状 target_shape = (x.shape[0] + len(insert_indices), x.shape[1]) # 初始化全零张量 X = torch.zeros(target_shape, dtype=x.dtype) # 计算原张量行对应的目标行索引 source_indices = torch.arange(x.shape[0]) # 统计每个原索引前插入的零行数量 insert_counts = torch.tensor([sum(idx <= i for idx in insert_indices) for i in source_indices]) target_indices = source_indices + insert_counts # 填充原张量数据 X[target_indices] = x print(X)
方案二:布尔掩码法(更简洁高效)
通过布尔掩码标记需要填充原数据的位置,直接批量填充,适合插入数量较多的场景:
import torch def insert_zero_rows(x, insert_indices): insert_indices = sorted(insert_indices) # 确保索引有序 n, m = x.shape k = len(insert_indices) # 初始化全零目标张量 X = torch.zeros(n + k, m, dtype=x.dtype) # 构建掩码:True表示需要填充原数据的位置,False表示零行 mask = torch.ones(n + k, dtype=torch.bool) mask[insert_indices] = False # 批量填充原张量数据 X[mask] = x.flatten() return X # 测试示例 x = torch.tensor([[1,1,1], [2,2,2], [3,3,3], [4,4,4]]) X = insert_zero_rows(x, [1,3]) print(X)
输出结果
两种方案都会得到你想要的张量:
tensor([[1, 1, 1], [0, 0, 0], [2, 2, 2], [0, 0, 0], [3, 3, 3], [4, 4, 4]])
内容的提问来源于stack exchange,提问作者ChrisNick92
相关产品推荐
相关产品推荐

