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
相关产品推荐
相关产品推荐

