如何用PyTorch高效将嵌套列表转换为指定形状的one-hot张量?
高效实现方案
原方法依赖Python循环逐个处理子列表,再堆叠结果,数据量较大时会因循环开销和多次小张量操作拖慢速度。推荐用向量化索引赋值的方式实现,完全利用PyTorch的张量操作优化,效率提升明显:
实现代码
import torch foo = [[1], [2, 3], [4]] num_classes = 100 # 生成行索引和列索引 row_indices = [] col_indices = [] for idx, items in enumerate(foo): row_indices.extend([idx] * len(items)) col_indices.extend(items) # 转换为张量 row_indices = torch.tensor(row_indices) col_indices = torch.tensor(col_indices) # 创建全0目标张量,批量赋值 bar = torch.zeros(len(foo), num_classes, dtype=torch.int32) bar[row_indices, col_indices] = 1
原理说明
- 先把嵌套列表展开,生成每个元素对应的行索引(子列表在
foo中的位置)和列索引(元素本身),整理出所有需要设为1的位置索引对。 - 初始化全0目标张量后,通过高级索引一次性将所有目标位置设为1,全程是张量层面的向量化操作,没有Python循环的额外开销,GPU运行时还能利用并行计算进一步提速。
对比优势
- 避免了原方法中对每个子列表单独调用
F.one_hot和sum的重复操作,减少了张量创建与销毁的开销。 - 数据量越大,向量化操作的性能优势越显著,完全适配PyTorch的底层优化逻辑。
内容的提问来源于stack exchange,提问作者Loua
相关产品推荐
相关产品推荐

