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

如何加速计算两个3D GPU数组对应多重集的交集大小?

MATLAB GPU数组多重集交集计算优化方案

问题背景

我们有两个uint16类型的3D GPU数组A和B,第2、3维度尺寸一致:

  • size(A,1)=300000,size(B,1)=2000
  • 第3维度规模极大(约1000000),需分块处理
    需求是对每个索引i,j,d,计算A(i,:,d)与B(j,:,d)这两个等长多重集的交集大小(即最大公共子多重集的元素数量),且已知B的行已排序。原代码因第3维度循环次数过多,内层计算逻辑存在性能瓶颈,需优化。

原代码瓶颈分析

原实现的核心低效点:

  1. 基于unique和逐元素相等比较的频次统计,会生成超大中间数组,占用GPU显存且计算冗余
  2. 内层for i=1:n循环遍历B的每一行,未利用GPU的大规模并行计算能力,算力浪费严重

优化方案

利用固定取值范围的频次统计+GPU广播并行运算,完全消除内层循环,充分发挥GPU算力优势:

优化核心思路

已知元素取值上限为N,直接统计每个行向量中所有可能值的出现次数,再通过维度广播将A和B的频次矩阵扩展为可批量计算的三维数组,最后逐元素取最小值并求和得到交集大小。

完整优化代码

s = 300000; % A的第1维度尺寸
n = 2000; % B的第1维度尺寸
c = 10; % A和B的第2维度尺寸
depth = 10; % 第3维度批量处理大小
N = 100; % 元素取值上限

A = randi(N,s,c,depth,'uint16','gpuArray');
B = randi(N,n,c,depth,'uint16','gpuArray');

Sizes_of_multiset_intersections = zeros(s,n,depth,'uint8');

for d=1:depth
    A_slice = A(:,:,d); % 当前批次的A切片:s×c
    B_slice = B(:,:,d); % 当前批次的B切片:n×c
    
    % --- 1. 批量统计A每行的元素频次 ---
    [s_dim, c_dim] = size(A_slice);
    row_idx_A = repmat((1:s_dim)', 1, c_dim); % 生成行索引矩阵:s×c
    % 用accumarray一次性统计所有行的频次,得到s×N的uint8数组
    A_counts = accumarray([row_idx_A(:), A_slice(:)], 1, [s_dim, N], @sum, 0, 'uint8');
    A_counts = gpuArray(A_counts);
    
    % --- 2. 批量统计B每行的元素频次 ---
    [n_dim, c_dim] = size(B_slice);
    row_idx_B = repmat((1:n_dim)', 1, c_dim); % 生成行索引矩阵:n×c
    B_counts = accumarray([row_idx_B(:), B_slice(:)], 1, [n_dim, N], @sum, 0, 'uint8');
    B_counts = gpuArray(B_counts);
    
    % --- 3. 广播计算逐对交集大小 ---
    % 扩展维度实现广播:A_counts(s×N) → s×1×N;B_counts(n×N) → 1×n×N
    A_exp = reshape(A_counts, s_dim, 1, N);
    B_exp = reshape(B_counts, 1, n_dim, N);
    
    % 逐元素取最小值,再沿N维度求和,得到s×n的交集大小矩阵
    Sizes_of_multiset_intersections_tmp = sum(min(A_exp, B_exp), 3, 'native');
    
    % 将结果回传到CPU内存
    Sizes_of_multiset_intersections(:,:,d) = gather(Sizes_of_multiset_intersections_tmp);
end

关键优化细节

  • 频次统计优化:用accumarray一次性完成所有行的频次统计,替代原代码中去重+逐元素比较的冗余逻辑,大幅减少中间数组占用
  • GPU并行利用:通过维度广播将逐对行的计算转化为GPU擅长的三维批量运算,彻底消除内层循环,最大化并行效率
  • 类型安全:全程使用uint8类型和'native'选项,保证运算不溢出且保持最高计算效率
  • 显存控制:分批次处理第3维度,避免一次性加载超大数组,适配GPU显存限制

内容的提问来源于stack exchange,提问作者M.G.

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.23 15:57:20