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

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.30 00:02:32