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

PyTorch GPU环境下低秩块对角矩阵与向量高效乘法方案求解

PyTorch无循环高效实现方案

你要的运算可以通过分组聚合加广播的向量化操作实现,完全不需要写循环,所有运算都是PyTorch原生CUDA支持的操作,GPU上运行效率极高,时间复杂度为O(n),和k的大小无关。

核心逻辑拆解

你需要的运算可以拆解为两步:

  1. 对每个分组x == i_j,计算该分组下v和w的点积:dot_j = sum(v[x==i_j] * w[x==i_j])
  2. 把每个分组的点积结果,乘以该分组下所有位置的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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.09.28 05:57:00