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

如何基于新旧值映射转换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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.03 12:40:42