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

如何用PyTorch高效批量实现矩阵指定边的替换操作?

基于PyTorch高效批量替换距离矩阵边值的实现方法

问题说明

现有两个形状为(batch_size, pomo_size, problem_size, problem_size)的距离张量up和down,设定参数:

  • batch_size=1
  • pomo_size=2
  • problem_size=5

同时有一个形状为(batch_size, pomo_size, selected_length)的节点序列selected_node_list,selected_length=8。

需求:将down张量中,selected_node_list序列所表示的边(包含序列首尾节点构成的边)的值,替换为up张量对应位置的值,要求用PyTorch的批量操作(如gather、scatter_等)实现,避免多层for循环带来的性能损耗。


给定张量

up张量

import torch

up = torch.tensor([[[[8, 5, 6, 8, 2],
                     [2, 9, 7, 7, 0],
                     [5, 9, 1, 7, 7],
                     [0, 5, 8, 6, 3],
                     [5, 9, 1, 0, 2]],
                    
                    [[2, 8, 1, 4, 7],
                     [0, 0, 2, 2, 7],
                     [4, 7, 7, 9, 4],
                     [6, 6, 7, 1, 3],
                     [3, 9, 9, 7, 2]]]])

down张量

down = torch.tensor([[[[6, 1, 7, 9, 1],
                       [7, 6, 2, 7, 9],
                       [5, 3, 8, 6, 3],
                       [0, 8, 6, 3, 3],
                       [8, 9, 4, 8, 1]],
                      
                      [[7, 2, 0, 6, 0],
                       [6, 7, 5, 3, 9],
                       [4, 8, 6, 6, 1],
                       [9, 3, 2, 2, 5],
                       [2, 5, 2, 1, 8]]]])

selected_node_list张量

selected_node_list = torch.tensor([[[0, 1, 2, 3, 4, 1, 3, 2],
                                    [1, 0, 3, 4, 0, 2, 4, 3]]])

替换示例与目标结果

单个pomo替换示例

第一个batch的第一个pomo矩阵替换后结果:

[[6, 5, 7, 9, 1],
 [7, 6, 7, 7, 9],
 [5, 3, 8, 7, 3],
 [0, 8, 8, 3, 3],
 [8, 9, 4, 8, 1]]

完整替换后的down张量

tensor([[[[6, 5, 7, 9, 1],
          [7, 6, 7, 7, 9],
          [5, 3, 8, 7, 3],
          [0, 8, 8, 3, 3],
          [8, 9, 4, 8, 1]],
         
         [[7, 2, 1, 4, 0],
          [0, 7, 5, 3, 9],
          [4, 8, 6, 6, 4],
          [9, 6, 2, 2, 3],
          [3, 5, 2, 7, 8]]]])

高效批量实现代码

# 获取维度参数
batch_size, pomo_size, selected_length = selected_node_list.shape
problem_size = up.shape[-1]

# 生成边的起点和终点:序列中连续节点对 + 首尾节点对
start_nodes = selected_node_list
end_nodes = torch.cat([selected_node_list[..., 1:], selected_node_list[..., :1]], dim=-1)

# 生成batch和pomo的索引,用于定位每个边对应的维度
batch_idx = torch.arange(batch_size).view(-1, 1, 1).repeat(1, pomo_size, selected_length)
pomo_idx = torch.arange(pomo_size).view(1, -1, 1).repeat(batch_size, 1, selected_length)

# 从up中提取需要替换的值
values_to_replace = up[batch_idx, pomo_idx, start_nodes, end_nodes]

# 若需保留原down张量,先创建副本;否则直接修改原张量
down_copy = down.clone()
down_copy.scatter_(dim=-1, 
                   index=end_nodes.unsqueeze(-1).repeat(1,1,1,problem_size), 
                   src=values_to_replace.unsqueeze(-1).repeat(1,1,1,problem_size))

# 验证结果
print(down_copy)

代码解释

  1. 生成边对:通过拼接selected_node_list与自身偏移一位的张量,得到所有需要替换的边的起点(start_nodes)和终点(end_nodes),自动包含首尾节点构成的边。
  2. 维度索引生成:创建batch_idx和pomo_idx,确保每个边都能精准定位到对应的batch和pomo维度。
  3. 提取替换值:利用张量索引直接从up中批量取出所有需要替换的边的值。
  4. 批量替换:使用scatter_函数,在down的对应位置批量写入up中的值,全程无循环,实现高效批量操作。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.19 03:50:09