如何基于新旧值映射转换PyTorch二维张量的元素?
基于映射张量转换元素值
给定二维张量old:
import torch old = torch.Tensor([ [1, 2, 12, 12], [0, 1, 12, 12], [3, 5, 12, 12], [7, 8, 12, 12], [6, 7, 12, 12], [9, 11, 12, 12]])
以及映射张量mapping(第一列为原元素值,第二列为对应转换后的值):
mapping = torch.Tensor([ [0, 0], [1, 6], [2, 1], [3, 6], [4, 2], [5, 6], [6, 3], [7, 6], [8, 4], [9, 6], [10, 5], [11, 6], [12, 6]])
期望得到转换后的张量:
new_or_desired = torch.Tensor([ [6, 1, 6, 6], [0, 6, 6, 6], [6, 6, 6, 6], [6, 4, 6, 6], [3, 6, 6, 6], [6, 6, 6, 6]])
原方法的问题
你尝试的old[old == mapping[:, 0]] = mapping[:, 1]会报错,原因是old == mapping[:,0]会触发广播机制,生成一个(6,4,13)的布尔张量,通过它索引出的元素数量和mapping[:,1]的长度不匹配,且布尔索引返回的是一维张量,无法直接赋值回原二维形状,导致形状不匹配错误。
解决方案
方法1:构建索引映射表(最直观高效)
利用mapping的结构,先创建一个以原元素值为索引、目标值为对应内容的映射表,再直接用old的整数索引获取转换后的值:
# 获取原元素的最大值,确定映射表长度 max_original_val = int(mapping[:, 0].max().item()) # 初始化映射表 map_table = torch.zeros(max_original_val + 1, dtype=torch.float32) # 将映射关系填充到表中:原数值位置放入对应目标值 map_table[mapping[:, 0].long()] = mapping[:, 1] # 直接用old的整数索引取映射值,自动匹配原形状 new_tensor = map_table[old.long()]
运行后new_tensor完全符合期望输出。
方法2:使用scatter_构建映射表
如果你想用scatter_,可以用它来填充映射表,原理和方法1一致:
max_original_val = int(mapping[:, 0].max().item()) map_table = torch.zeros(max_original_val + 1, dtype=torch.float32) # scatter_参数:dim=0(沿第0维填充),index为原数值的整数索引,src为目标值 map_table.scatter_(0, mapping[:, 0].long(), mapping[:, 1]) new_tensor = map_table[old.long()]
scatter_在这里的作用是把mapping[:,1]的值,根据mapping[:,0].long()的索引位置,填充到map_table中,最终得到和方法1相同的映射表,再通过索引得到结果。
内容的提问来源于stack exchange,提问作者KDecker
相关产品推荐
相关产品推荐

