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

如何将PyTorch张量构成的字典元组转换为张量字典?

解决方案:将字典元组转换为张量字典(PyTorch)

最优实现方式

直接使用字典推导式结合torch.stack(),这是GPU环境下最高效的方法——torch.stack()是PyTorch底层优化的张量操作,能避免Python层面的循环开销,充分利用GPU并行计算能力:

import torch

tuple_of_dicts_of_tensors = (
    {'key_1': torch.tensor([1,1,1]), 'key_2': torch.tensor([4,4,4])},
    {'key_1': torch.tensor([2,2,2]), 'key_2': torch.tensor([5,5,5])},
    {'key_1': torch.tensor([3,3,3]), 'key_2': torch.tensor([6,6,6])}
)

# 核心转换代码
dict_of_tensors = {
    key: torch.stack([d[key] for d in tuple_of_dicts_of_tensors])
    for key in tuple_of_dicts_of_tensors[0].keys()
}

# 验证结果
print(dict_of_tensors['key_1'])
# 输出:
# tensor([[1, 1, 1],
#         [2, 2, 2],
#         [3, 3, 3]])

效率说明

  • 仅对字典的key进行循环(通常key数量远小于元组长度),张量堆叠操作由PyTorch的C++后端执行,完全避开Python循环带来的性能损耗。
  • GPU上torch.stack()直接在设备内存中完成张量拼接,无需频繁数据传输或Python层面计算,最大化利用硬件性能。

边界情况处理(可选)

如果元组中存在key不一致的字典,可先提取所有字典的公共key再处理:

# 获取所有字典的公共key
common_keys = set.intersection(*(set(d.keys()) for d in tuple_of_dicts_of_tensors))

dict_of_tensors = {
    key: torch.stack([d[key] for d in tuple_of_dicts_of_tensors])
    for key in common_keys
}

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.12 14:28:23