如何用PyTorch高效批量实现矩阵指定边的替换操作?
基于PyTorch高效批量替换距离矩阵边值的实现方法
问题说明
现有两个形状为(batch_size, pomo_size, problem_size, problem_size)的距离张量up和down,设定参数:
batch_size=1pomo_size=2problem_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)
代码解释
- 生成边对:通过拼接
selected_node_list与自身偏移一位的张量,得到所有需要替换的边的起点(start_nodes)和终点(end_nodes),自动包含首尾节点构成的边。 - 维度索引生成:创建
batch_idx和pomo_idx,确保每个边都能精准定位到对应的batch和pomo维度。 - 提取替换值:利用张量索引直接从
up中批量取出所有需要替换的边的值。 - 批量替换:使用
scatter_函数,在down的对应位置批量写入up中的值,全程无循环,实现高效批量操作。
内容的提问来源于stack exchange,提问作者chihiro
相关产品推荐
相关产品推荐

