PyTorch GPU性能优化:移除分组组合计数逻辑中的Python循环
PyTorch 文档共享特征统计性能优化方案
现有实现瓶颈分析
现有代码的性能损耗主要来自两点:
- Python侧的列表推导式循环会打断GPU计算流,每次
torch.combinations调用都需要触发CPU-GPU同步,空耗大量等待时间 torch.split产出的每个子张量长度不定,torch.combinations每次都要做单独的显存分配和回收,无法利用GPU批量并行的特性
优化方案
方案1:矩阵乘法直接计算(优先推荐)
你的需求本质是计算所有文档对的特征共现次数,这个逻辑可以直接通过矩阵乘法实现,完全不需要循环和组合生成,GPU并行效率最高:
# 假设X为n_docs x m_features的二值矩阵,1表示对应文档包含该特征,已移到GPU显存 X = X.to('cuda') # 矩阵乘法直接得到所有文档对的共享特征数,shape为n_docs x n_docs shared_counts = X @ X.T # 如果只需要无重复的文档对结果,添加上三角掩码过滤即可 mask = torch.triu(torch.ones_like(shared_counts, dtype=torch.bool), diagonal=1) unique_pair_counts = shared_counts[mask]
如果你的X是稀疏特征矩阵,可以先转成torch.sparse格式再做乘法,显存占用和计算速度会比稠密矩阵更优。
方案2:按特征对应文档数分组批量处理
如果你的文档数量极大,矩阵乘法显存不足,可以采用你提到的分组思路实现,把Python循环次数从原来的feature数量级降到k的不同取值数量级(通常k的取值最多几十种):
from collections import defaultdict import torch # 预处理:过滤掉仅出现在1个文档中的特征,得到每个有效feature对应的doc_id列表 valid_feat_doc_ids = [doc_ids for doc_ids in torch.split(X, C) if len(doc_ids) >= 2] # 按每个feature对应的文档数量k分组 k_groups = defaultdict(list) for doc_ids in valid_feat_doc_ids: k = len(doc_ids) k_groups[k].append(doc_ids) all_pairs = [] for k, batch_docs in k_groups.items(): # 同k的所有doc列表拼成批量张量,shape为[当前k对应的feature数量, k] batch = torch.stack(batch_docs, dim=0) # 一次性生成该k值对应的所有组合索引,无需循环 comb_idx = torch.combinations(torch.arange(k, device=batch.device), r=2) # 批量生成所有文档对,shape为[当前k对应的总pair数量, 2] batch_pairs = batch[:, comb_idx].reshape(-1, 2) all_pairs.append(batch_pairs) # 合并所有pair并统计频次 all_pairs = torch.cat(all_pairs, dim=0) unique_pairs, shared_counts = torch.unique(all_pairs, dim=0, return_counts=True)
方案选择建议
- 当文档数n_docs ≤ 1e4时,优先选择矩阵乘法方案,计算速度最快,代码最简洁
- 当文档数n_docs > 1e4且特征非常稀疏时,选择分组批量处理方案,显存占用更低
内容的提问来源于stack exchange,提问作者chase
相关产品推荐
相关产品推荐

