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

如何在PyTorch中通过字典高效映射Tensor元素?

PyTorch中高效实现张量元素字典映射的方法

针对你提出的张量元素字典映射需求,这里提供几种高效的实现方式,对标TensorFlow的tf.lookup.StaticHashTable:

方法一:直接构造查找张量(最适合连续整数键)

如果映射字典的键是连续的整数,直接构建一个索引与键对应的查找张量,通过张量索引快速完成映射,时间复杂度为O(1) per元素,效率极高:

import torch

t = torch.tensor([1, 2, 3])
mapping = {1: 0.2, 2: 1.2, 3: 3.0}

# 确定键的最大范围,构造覆盖所有键的查找张量
max_key = max(mapping.keys())
lookup_tensor = torch.zeros(max_key + 1, dtype=torch.float32)
for key, val in mapping.items():
    lookup_tensor[key] = val

# 直接索引得到结果
result = lookup_tensor[t]
print(result)  # 输出: tensor([0.2000, 1.2000, 3.0000])

方法二:使用Embedding层(适合稀疏/大范围键)

如果映射的键是稀疏分布或范围较大,使用PyTorch的Embedding层可以高效完成映射,本质是构建一个可索引的静态哈希表:

import torch

t = torch.tensor([1, 2, 3])
mapping = {1: 0.2, 2: 1.2, 3: 3.0}

keys = torch.tensor(list(mapping.keys()), dtype=torch.long)
values = torch.tensor(list(mapping.values()), dtype=torch.float32)

# 创建Embedding层,初始化权重为映射值
embedding = torch.nn.Embedding(num_embeddings=keys.max() + 1, embedding_dim=1)
with torch.no_grad():
    embedding.weight[keys] = values.unsqueeze(1)

# 执行映射并压缩维度
result = embedding(t).squeeze(1)
print(result)  # 输出: tensor([0.2000, 1.2000, 3.0000])

如果不需要计算梯度,也可以直接通过权重张量索引,省去Embedding层的封装:

# 直接用权重张量索引
result = embedding.weight[t].squeeze(1)

方法三:索引匹配法(适合小批量/非整数键)

如果键不是整数或批量较小,可以通过匹配索引完成映射,不过效率略低于前两种方法:

import torch

t = torch.tensor([1, 2, 3])
mapping = {1: 0.2, 2: 1.2, 3: 3.0}

keys = torch.tensor(list(mapping.keys()))
values = torch.tensor(list(mapping.values()))

# 找到t中每个元素在keys中的对应索引
indices = torch.argmax(keys.unsqueeze(1) == t, dim=0)
result = values[indices]
print(result)  # 输出: tensor([0.2000, 1.2000, 3.0000])

总结

  • 优先选择方法一(查找张量索引),在键为连续整数时性能最优;
  • 键稀疏或范围大时,用方法二(Embedding层),逻辑更简洁且适合大规模数据;
  • 非整数键或小批量场景可考虑方法三,但注意性能损耗。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.16 06:20:06