
本文介绍一种无需循环、不依赖 cpu 的优雅方法:将映射字典转为索引张量,利用 pytorch 的高级索引机制实现 o(1) 时间复杂度的批量映射。
本文介绍一种无需循环、不依赖 cpu 的优雅方法:将映射字典转为索引张量,利用 pytorch 的高级索引机制实现 o(1) 时间复杂度的批量映射。
在深度学习与科学计算中,常需将离散标签张量(如类别索引)批量映射为对应权重、嵌入向量或概率值。若使用 Python 循环或 torch.apply_(已弃用且仅支持 CPU),不仅性能低下,还破坏了 GPU 加速和自动微分的连续性。
最推荐的解决方案是将字典转换为张量,并利用张量索引机制完成向量化映射。核心思想是:字典的键作为新张量的索引位置,值作为该位置的元素。只要原始张量 t 中的元素均为非负整数且落在字典键范围内,即可直接用 t 作为索引访问映射张量。
例如:
import torch
t = torch.tensor([1, 0, 0, 1], dtype=torch.long) # 注意:索引需为 long 类型
weights = {0: 0.1, 1: 0.9}
# 将字典转换为张量:按键 0, 1, 2, ... 顺序排列值
weight_tensor = torch.tensor([weights[i] for i in sorted(weights.keys())])
# 向量化索引(支持 CPU/GPU/autograd)
new_t = weight_tensor[t]
print(new_t)
# tensor([0.9000, 0.1000, 0.1000, 0.9000])
✅ 优势说明:
- 完全向量化,无显式循环,速度极快;
- 原生支持 GPU 张量(只需确保
weight_tensor和t在同一设备); - 保留梯度流(若
weight_tensor是可训练参数,反向传播自动生效); - 内存友好,时间复杂度为 O(n),空间复杂度为 O(k)(k 为字典键数量)。
⚠️ 注意事项:
-
t必须是torch.long(或torch.int64)类型,否则索引会报错; - 字典键应为连续或至少覆盖
t中所有可能出现的值,否则索引越界; - 若键不连续(如
{0: 0.1, 2: 0.9}),需补全中间键(如设为默认值),或改用torch.nn.Embedding(更适用于大规模稀疏映射); - 对于超大字典(如万级类别),建议直接使用
nn.Embedding(num_classes, 1)并初始化权重,语义更清晰且支持优化器管理。
总结:将映射关系张量化 + 高级索引,是 PyTorch 中处理“离散值→连续值”批量转换的黄金范式——简洁、高效、可微、设备无关。











