PyTorch中如何用列表列表高效批量设置张量元素值?
解决方案:PyTorch高效批量设置张量指定位置为-1
问题原因
你之前的写法报错,是因为torch.arange(len(list_of_indices))是一维张量(比如形状[9]),而list_of_indices转成张量后是变长二维结构(比如形状[9, 0]或各行长度不一),PyTorch的高级索引要求两个索引张量的形状必须完全匹配,所以触发形状不匹配的错误。
高效实现方式
核心思路是把行索引和列索引都展平成一维张量,让每个行索引对应一个要修改的列索引,再通过高级索引批量赋值。
方法1:简洁推导式写法
# 生成匹配的行、列索引张量 row_indices = torch.tensor([i for i, cols in enumerate(list_of_indices) for _ in cols], device=x.device) col_indices = torch.tensor([col for cols in list_of_indices for col in cols], device=x.device) # 批量赋值 x[row_indices, col_indices] = -1
方法2:循环扩展写法(逻辑更直观)
row_indices = [] col_indices = [] for row_idx, target_cols in enumerate(list_of_indices): # 该行有多少个要修改的列,就把行索引重复多少次 row_indices.extend([row_idx] * len(target_cols)) # 把该行的列索引全部加入列表 col_indices.extend(target_cols) # 转成同设备的张量 row_indices = torch.tensor(row_indices, device=x.device) col_indices = torch.tensor(col_indices, device=x.device) # 批量修改 x[row_indices, col_indices] = -1
说明
- 这两种写法都会自动跳过空的子列表(即没有要修改列的行),不会对这些行做任何操作。
- 如果你的张量
x在GPU上运行,一定要把row_indices和col_indices移到同一个设备上(通过device=x.device参数),避免设备不匹配错误。 - 这种方式的时间复杂度和for循环一致,但利用了PyTorch的批量操作优化,比手动逐行循环效率更高,尤其当张量规模较大时优势明显。
内容的提问来源于stack exchange,提问作者leleogere
相关产品推荐
相关产品推荐

