PyTorch GPU环境下低秩块对角矩阵与向量高效乘法方案求解
PyTorch无循环高效实现方案
你要的运算可以通过分组聚合加广播的向量化操作实现,完全不需要写循环,所有运算都是PyTorch原生CUDA支持的操作,GPU上运行效率极高,时间复杂度为O(n),和k的大小无关。
核心逻辑拆解
你需要的运算可以拆解为两步:
- 对每个分组
x == i_j,计算该分组下v和w的点积:dot_j = sum(v[x==i_j] * w[x==i_j]) - 把每个分组的点积结果,乘以该分组下所有位置的
w元素,得到对应位置的y值
完整实现代码
import torch def group_dot_mul(v: torch.Tensor, w: torch.Tensor, x: torch.Tensor) -> torch.Tensor: # 对x做重映射,得到连续的分组索引,避免x取值稀疏导致内存浪费 unique_x, inverse_idx = torch.unique(x, return_inverse=True) k = unique_x.numel() # 计算v和w的逐元素乘积 elem_prod = v * w # 按分组求和,得到每个组的点积结果 group_dot = torch.zeros(k, dtype=v.dtype, device=v.device).scatter_reduce_( dim=0, index=inverse_idx, src=elem_prod, reduce="sum", include_self=False ) # 分组点积广播到每个位置后乘以w,得到最终结果 return group_dot[inverse_idx] * w
性能说明
- 所有操作均为PyTorch原生实现,支持CUDA加速,没有Python层循环开销
- 无需对输入做排序和逆排列操作,省去了排序的O(n log n)时间开销
- 对k的大小不敏感,无论k极小还是接近n,都能保持稳定的高性能
内容的提问来源于stack exchange,提问作者mkcohen
相关产品推荐
相关产品推荐

