已关闭
【社区任务】MatmulGatherScatter算子贡献 #391
Delicate02创建于  8月3日关闭于  5 天前
Delicate02
Delicate02
8月3日 创建

Thanks for sending an requirement! Please fill in the following template to help quickly solve your problem.

Backgroud / 背景信息

需求背景

需求来源

7 月社区任务 MatmulGatherScatter 算子开发任务书,要求在 CATLASS 仓中实现 Ascend 950 上的 Gather、Matmul、Scatter 融合算子,并补充 optest 测试交付件和 README。

任务书:https://gitcode.com/cann/cann-ops-competitions/blob/master/04_tasks/01_community-task-2026/docs/202607/MatmulGatherScatter_task_doc.md

背景介绍

MatmulGatherScatter 算子实现优化

推荐系统和稀疏特征计算通常只需要更新稠密矩阵中的部分行。任务基线由 torch.zeros 初始化、Gather、torch.mm 和 Scatter 小算子拼接完成,需要多次 kernel launch,并会将 Gather 结果 (J,K) 和 Matmul 结果 (J,N) 写入 GM 后再被后续算子读取。

本需求基于 CATLASS 和 Ascend C 实现单次 MIX kernel 融合,在片上完成 D 清零、按行 Gather、Cube Matmul 和按行 Scatter,减少中间 GM 读写和多算子下发开销。

MatmulGatherScatter 算子实现现状分析

算子输入输出如下:

参数 参数含义 数据类型 格式 约束 形状
A Matmul 左矩阵 FP16 ND RowMajor 按 indices 读取行 (M,K)
B Matmul 右矩阵 FP16 ND RowMajor 连续二维矩阵 (K,N)
indices 行索引向量 INT32 ND 合法、无重复,由调用方保证 (J,)
D 输出矩阵 FP16 ND RowMajor 未索引行保持为零 (M,N)

MatmulGatherScatter 算子功能分析

先从 A 中按 indices 取出逻辑左矩阵:

Ag[i, k] = A[indices[i], k]

再计算矩阵乘:

C[i, n] = sum(Ag[i, k] * B[k, n]), k in [0, K)

最后使用同一 indices 写回 D:

D[indices[i], n] = C[i, n]

D 中不在 indices 内的其他行保持为零。等价向量化形式为:

D[indices, :] = A[indices, :] @ B

需求分析

需求描述

使用 CATLASS/Ascend C 实现 Ascend 950 上的 FP16 RowMajor Gather A + Matmul + Scatter D 融合算子。算子需在单次 kernel launch 内完成输出清零、Gather、Matmul 和 Scatter,输出与 PyTorch 参考实现保持一致,并在 Ascend 950PR 上达到任务测试集小算子拼接 baseline 的整体 1.2 倍性能目标。

需求拆解

  1. 支持输入 A(M,K)、B(K,N) 和 indices(J),数据类型分别为 FP16、FP16 和 INT32。
  2. 支持输出 D(M,N),数据类型为 FP16,布局为 RowMajor。
  3. Gather A 和 Scatter D 使用同一索引向量。
  4. indices 的范围与唯一性由调用方保证,kernel 不做范围检查或去重。
  5. 单次 kernel launch 完成 D 清零、Gather、Matmul 和 Scatter,不创建 (J,K) 或 (J,N) GM 中间张量。
  6. Cube 使用 FP32 累加,最终输出转换为 FP16。
  7. 未被 indices 选中的 D 行严格为零。
  8. 基于 CATLASS optest 补充 JIT kernel、C++ adapter、Python wrapper 和 pytest,覆盖任务测试集全部 180 个 shape。
  9. 正式性能使用 Ascend 950PR msprof op 逐项采集,并与任务测试集提供的 baseline 对比。

详细设计

算子分析

数学公式

Ag[i, k] = A[indices[i], k]
C[i, n] = sum(Ag[i, k] * B[k, n]), k in [0, K)

输出定义为:

D[r, n] = C[i, n], r = indices[i]
D[r, n] = 0,       r not in indices

逻辑 Matmul 的问题规模为 (J,N,K),A 和 D 的物理行数为 M。Gather 源行地址使用物理 K 作为行跨度,Scatter 目标行地址使用物理 N 作为行跨度。

支持数据类型

参数 数据类型 格式 说明
A FP16 RowMajor ND 左矩阵,shape (M,K)
B FP16 RowMajor ND 右矩阵,shape (K,N)
indices INT32 ND 行索引,shape (J,)
D FP16 RowMajor ND 输出矩阵,shape (M,N)

支持形状

维度 约束
M M > 0
J 0 < J <= M
N N > 0
K K > 0,当前实现要求为 16 的倍数
indices 元素位于 [0,M-1] 且互不重复

算子实现

Host 侧设计

Host/JIT 侧完成输入约束检查、输出申请、硬件同步地址获取、调度参数选择和 MIX kernel launch:

  1. 解析 M/J/N/K,检查输入 dtype、设备、连续 RowMajor 布局和形状关系。
  2. 根据 K、N、J、TileShape、L1 容量和输出规模选择 Gather 路径、TileM/TileN/L1K、B stage 数和 BlockScheduler。
  3. 调用 aclrtGetHardwareSyncAddr 获取全核同步地址,并通过 kernel 参数传入。
  4. 固定启动 28 个 AIC block,以 MIX (1 AIC, 2 AIV) 使用全部 28 个 Cube Core 和 56 个 Vector Core。
  5. 单次 launch 完成全部融合流程,workspace 大小为 0。

新增 numbered example 79_ascend950_matmul_gather_scatter,公开 Python 接口为:

torch_catlass.ascend950_matmul_gather_scatter(a, b, indices) -> d

Kernel 侧设计

Kernel 采用 AIC/AIV 协同执行:

  1. 56 个 AIV 对 M*N 输出区域分片清零;满足容量和规模条件时,只清零未被选中的连续行段。
  2. Gather 根据 shape 在 AIV MTE2 Gather 和 AIC 按 indices 直接 GM 到 L1 Gather 间选择;Ascend 950 SIMT Gather/Scatter 组件保留用于动态 shape 和实验探针。
  3. Gather 后的逻辑 A tile (TileM,K) 沿 K 维完整驻留 L1,并在多个 N tile 间复用。
  4. B 复用 CATLASS 的 GM/L2 到 L1 双缓冲或多阶段流水,L1 数据进入 L0A/L0B 后由 Cube 执行 FP32 累加。
  5. Fixpipe 将 L0C 结果送入单/双 AIV UB,AIV 完成 FP16 转换并按 D[indices[row], nOffset] Scatter 写回。
  6. 通用全量清零路径在首次 Scatter 前执行一次全核同步;后续 tile 通过 A-ready、A-reload-safe、C-ready 和 buffer-free 标志完成双缓冲流水同步。

Gather 与调度策略

当前设计按 shape 特征选择稳健路径,而不是为每个测试 case 写逐项特判:

场景 主要 Gather 策略 设计考虑
K<=512 AIV MTE2 Gather 小 K 场景降低 Gather 启动与轮询开销
K=768 AIV MTE/AIC Gather 实测分流 结合 TileShape 和 N/J 规模选择
K>=2048 AIC 直接 GM 到 L1 Gather 大 K 复用 L1 resident-A,减少 AIV Gather 开销
L1 容量边界或动态 shape 已验证的 AIV MTE/AIC 回退 保证容量安全和设备稳定性

当 J>TileM 时优先使用 FullLoadA 调度,使同一 J tile 的多个 N tile 尽量在同一 AIC 连续执行并复用 L1 A;小 J 场景展开 N tile 到多个 AIC,避免并行度不足。

性能优化策略

  1. 融合 D 清零、Gather、Matmul 和 Scatter,减少额外 kernel launch 与 GM 中间张量读写。
  2. Gather 后的 A tile 沿 K 维完整驻留 L1,并在多个 N tile 间复用。
  3. B 使用双缓冲或多阶段 L1 环形缓冲,重叠 GM/L2 搬运与 Cube 计算。
  4. 输出清零使用多 AIV 大块连续搬运;满足条件时只清零未索引行,并消除首次全核屏障。
  5. Gather、输出、Cast 和清零 scratch 按生命周期复用 UB,使用编译期断言保证区域不重叠。
  6. 通过代表 shape 对比 AIV MTE、AIC Gather、SIMT、TileShape 和 B stage,只固化能通过完整任务集回归的策略。

支持硬件

支持的芯片版本 涉及勾选
Ascend 950 √
Ascend 950PR √

算子约束限制

  1. A、B、D 当前仅支持 FP16,indices 仅支持 INT32。
  2. A、B、D 当前仅支持连续 ND RowMajor。
  3. A 的列数必须等于 B 的行数,indices 长度为 J,且 J <= M。
  4. indices 必须合法且无重复,由调用方保证。
  5. 输入输出内存不得非法重叠。

可维可测分析

精度标准与性能标准

验收标准 描述 标准来源
精度标准 输出与 zeros + index_select + matmul + index_copy 参考实现满足生态算子开源精度标准,未索引行严格为零 任务书、生态算子开源精度标准
性能标准 Ascend 950PR 上 fused kernel 整体性能达到任务测试集 torch.zeros + gather + torch.mm + scatter baseline 的 1.2 倍 任务书

性能采集使用 msprof op,任务测试集全部 180 个 case 逐项记录 fused 时间、baseline、speedup 和实际调度配置;涉及 Gather 方案、TileShape 或 stage 调整时备注实验条件和选择结论。

兼容性分析

该需求以独立 numbered example、CATLASS kernel/tile helper 和 optest 接口接入,不改变已有公开算子语义。新增注册项仅用于 ascend950_matmul_gather_scatter JIT/optest 路径;已有 CATLASS 示例和 optest 接口不受影响。

likedislike
Delicate02Delicate02
8月3日 关联了pull request:[feat]补充77_ascend950_matmul_gather_scatter样例
Delicate02Delicate02
8月3日 修改了issue 的描述
Ssunhao_hw成员
8月6日 将 Delicate02 设为负责人
longjihuilongjihui成员
8月12日 添加了label:requirement
longjihui
longjihui成员
5 天前 评论:

需求已闭环,详见关联PR

likedislike
longjihuilongjihui成员
5 天前 issue状态由 进行中 改变为 已完成
longjihuilongjihui成员
5 天前 关闭了 issue
CANN-robotCANN-robot成员
5 天前 添加了label:resolved