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

PyTorch中高效合并同键张量字典的方法及现有实现优化问询

PyTorch中合并同键张量字典列表的高效方法

现有实现分析

你当前的实现是先用defaultdict(list)收集每个键对应的张量列表,最后通过torch.stack完成拼接,这是常规可行的方案,但确实存在优化空间。

更高效的实现方案

方案1:预分配张量内存(大幅减少动态扩容开销)

如果能提前知道每个键对应的张量总数(比如len(mention_indices))以及单个张量的维度,可以直接预分配大张量,再逐个填充,避免列表动态扩容的额外开销:

# 先获取第一个样本的结果,确定张量维度与设备信息
first_mention_input, _ = ...  # 需保证mention_indices非空
total_count = len(mention_indices)
mention_inputs = {}

# 为每个键预分配对应形状的张量
for key, value in first_mention_input.items():
    target_shape = (total_count,) + value.shape
    mention_inputs[key] = torch.empty(target_shape, device=value.device, dtype=value.dtype)

# 遍历填充张量
for idx in range(total_count):
    mention_input, _ = ...  # 对应mention_indices[idx]的计算结果
    for key, value in mention_input.items():
        mention_inputs[key][idx] = value

这种方式跳过了列表append和torch.stack的中间步骤,直接操作张量内存,在数据量较大时效率提升明显,尤其是高维度张量场景。

方案2:用torch.cat逐步拼接(小数据量场景适用)

如果无法提前确定总数,也可以直接用torch.cat逐步拼接,但注意每次拼接都会创建新张量,数据量大时内存拷贝开销较高:

mention_inputs = {}
for idx in mention_indices:
    mention_input, _ = ...
    for key, value in mention_input.items():
        if key not in mention_inputs:
            mention_inputs[key] = value.unsqueeze(0)
        else:
            mention_inputs[key] = torch.cat([mention_inputs[key], value.unsqueeze(0)], dim=0)

这种方式无需维护额外列表,直接操作张量,但仅适合小批量数据场景。

方案3:列表推导式简化代码(效率接近原实现)

如果追求代码简洁且数据量不大,可以先收集所有字典,再按键批量拼接:

# 先批量收集所有mention_input字典
all_mention_inputs = [get_mention_input(idx)[0] for idx in mention_indices]  # 替换成实际获取mention_input的逻辑
mention_inputs = {key: torch.stack([d[key] for d in all_mention_inputs]) for key in all_mention_inputs[0].keys()}

代码更紧凑,但本质逻辑和原实现一致,效率差别不大。

总结

  • 数据量大且能提前确定参数:优先用预分配张量的方案,效率最高
  • 数据量小或无法提前确定参数:可选择torch.cat逐步拼接或保留原实现
  • 追求代码简洁:用列表推导式的写法

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.28 23:30:06