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

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.10.06 02:42:01