已开启
feat(examples): 新增 AIN + Blaze AllGatherMatmul 流水示例 #226
izanaami创建于 12 天前
feat(examples): 新增 AIN + Blaze AllGatherMatmul 流水示例 #226
已开启
共 14 个文件变更+1490-1
| @@ -9,6 +9,7 @@ | |||
| 9 | | 样例 | 说明 | 支持产品 | | 9 | | 样例 | 说明 | 支持产品 | |
| 10 | | --- | --- | --- | | 10 | | --- | --- | --- | |
| 11 | | [aicore/ain/01_basic_ring](./aicore/ain/01_basic_ring/README.md) | 演示多卡环形场景下AIV Kernel通过AIN接口调用`Put`/`Get`进行点对点单边通信,并通过`AinBarrierSession`完成同步与结果校验。 | Ascend 950PR/Ascend 950DT | | 11 | | [aicore/ain/01_basic_ring](./aicore/ain/01_basic_ring/README.md) | 演示多卡环形场景下AIV Kernel通过AIN接口调用`Put`/`Get`进行点对点单边通信,并通过`AinBarrierSession`完成同步与结果校验。 | Ascend 950PR/Ascend 950DT | |
| 12 | +| [aicore/ain/02_all_gather_matmul](./aicore/ain/02_all_gather_matmul/README.md) | 使用 AIN Get 与 Blaze BlockMmad 实现 FP16/BF16 AllGatherMatmul,提供融合流水、串行/HCCL 对照及数值校验。 | Ascend 950PR/Ascend 950DT | | ||
| 12 | | [aicore/hcomm/01_hcomm_write_read_nbi](./aicore/hcomm/01_hcomm_write_read_nbi/README.md) | 演示多卡场景下AIV Kernel通过URMA路径调用`Hcomm::WriteNbi`和`Hcomm::ReadNbi`,并校验通信结果。 | Ascend 950PR/Ascend 950DT | | 13 | | [aicore/hcomm/01_hcomm_write_read_nbi](./aicore/hcomm/01_hcomm_write_read_nbi/README.md) | 演示多卡场景下AIV Kernel通过URMA路径调用`Hcomm::WriteNbi`和`Hcomm::ReadNbi`,并校验通信结果。 | Ascend 950PR/Ascend 950DT | |
| 13 | | [aicore/hcomm/02_hcomm_batch_write](./aicore/hcomm/02_hcomm_batch_write/README.md) | 演示多个URMA Channel共享Jetty,并通过`MakeBatchHandle`、`GetHandleRef`、`BatchCommit`和`Drain`批量提交跨peer写任务。 | Ascend 950PR/Ascend 950DT | | 14 | | [aicore/hcomm/02_hcomm_batch_write](./aicore/hcomm/02_hcomm_batch_write/README.md) | 演示多个URMA Channel共享Jetty,并通过`MakeBatchHandle`、`GetHandleRef`、`BatchCommit`和`Drain`批量提交跨peer写任务。 | Ascend 950PR/Ascend 950DT | |
| 14 | | [aicore/hcomm/03_one_multi_path](./aicore/hcomm/03_one_multi_path/README.md) | 查询`UB_MEM`链路,为每个peer创建one path/multi path Channel,逐个处理peer,并通过双Stream并发搬运和校验该peer的远端数据。 | Ascend 950PR/Ascend 950DT | | 15 | | [aicore/hcomm/03_one_multi_path](./aicore/hcomm/03_one_multi_path/README.md) | 查询`UB_MEM`链路,为每个peer创建one path/multi path Channel,逐个处理peer,并通过双Stream并发搬运和校验该peer的远端数据。 | Ascend 950PR/Ascend 950DT | |
| @@ -9,6 +9,7 @@ This directory contains usage samples for asc-comm APIs. | |||
| 9 | | Sample | Description | Supported Products | | 9 | | Sample | Description | Supported Products | |
| 10 | | --- | --- | --- | | 10 | | --- | --- | --- | |
| 11 | | [aicore/ain/01_basic_ring](./aicore/ain/01_basic_ring/README_en.md) | Demonstrates AIV Kernel invoking `Put`/`Get` over the AIN interface for point-to-point one-sided communication in a multi-card ring scenario, with synchronization and result verification through `AinBarrierSession`. | Ascend 950PR / Ascend 950DT | | 11 | | [aicore/ain/01_basic_ring](./aicore/ain/01_basic_ring/README_en.md) | Demonstrates AIV Kernel invoking `Put`/`Get` over the AIN interface for point-to-point one-sided communication in a multi-card ring scenario, with synchronization and result verification through `AinBarrierSession`. | Ascend 950PR / Ascend 950DT | |
| 12 | +| [aicore/ain/02_all_gather_matmul](./aicore/ain/02_all_gather_matmul/README_en.md) | Combines AIN Get and Blaze BlockMmad for FP16/BF16 AllGatherMatmul, with an overlapped pipeline, serial/HCCL baselines and numerical verification. | Ascend 950PR / Ascend 950DT | | ||
| 12 | | [aicore/hcomm/01_hcomm_write_read_nbi](./aicore/hcomm/01_hcomm_write_read_nbi/README_en.md) | Demonstrates AIV Kernel invoking `Hcomm::WriteNbi` and `Hcomm::ReadNbi` over the URMA path in a multi-card scenario, with communication result verification. | Ascend 950PR / Ascend 950DT | | 13 | | [aicore/hcomm/01_hcomm_write_read_nbi](./aicore/hcomm/01_hcomm_write_read_nbi/README_en.md) | Demonstrates AIV Kernel invoking `Hcomm::WriteNbi` and `Hcomm::ReadNbi` over the URMA path in a multi-card scenario, with communication result verification. | Ascend 950PR / Ascend 950DT | |
| 13 | | [aicore/hcomm/02_hcomm_batch_write](./aicore/hcomm/02_hcomm_batch_write/README_en.md) | Demonstrates multiple URMA channels sharing a Jetty and submitting cross-peer writes through `MakeBatchHandle`, `GetHandleRef`, `BatchCommit`, and `Drain`. | Ascend 950PR / Ascend 950DT | | 14 | | [aicore/hcomm/02_hcomm_batch_write](./aicore/hcomm/02_hcomm_batch_write/README_en.md) | Demonstrates multiple URMA channels sharing a Jetty and submitting cross-peer writes through `MakeBatchHandle`, `GetHandleRef`, `BatchCommit`, and `Drain`. | Ascend 950PR / Ascend 950DT | |
| 14 | | [aicore/hcomm/03_one_multi_path](./aicore/hcomm/03_one_multi_path/README_en.md) | Queries a `UB_MEM` link, creates one path/multi path channels for every peer, processes peers serially, and concurrently moves and verifies each peer's remote data over two streams. | Ascend 950PR / Ascend 950DT | | 15 | | [aicore/hcomm/03_one_multi_path](./aicore/hcomm/03_one_multi_path/README_en.md) | Queries a `UB_MEM` link, creates one path/multi path channels for every peer, processes peers serially, and concurrently moves and verifies each peer's remote data over two streams. | Ascend 950PR / Ascend 950DT | |
| @@ -0,0 +1,49 @@ | |||
| 1 | +# Copyright (c) 2026 Huawei Technologies Co., Ltd. | ||
| 2 | +# Licensed under the CANN Open Software License Agreement Version 2.0. | ||
| 3 | +cmake_minimum_required(VERSION 3.16) | ||
| 4 | +option(AGMM_HOST_TESTS_ONLY "Build portable host checks without CANN" OFF) | ||
| 5 | + | ||
| 6 | +if(AGMM_HOST_TESTS_ONLY) | ||
| 7 | + project(ain_agmm_host_tests LANGUAGES CXX) | ||
| 8 | + enable_testing() | ||
| 9 | + add_executable(agmm_host_tests tests/host_tests.cpp) | ||
| 10 | + target_compile_features(agmm_host_tests PRIVATE cxx_std_17) | ||
| 11 | + target_include_directories(agmm_host_tests PRIVATE . ../..) | ||
| 12 | + target_compile_options(agmm_host_tests PRIVATE -Wall -Wextra -Werror) | ||
| 13 | + add_test(NAME agmm_host_tests COMMAND agmm_host_tests) | ||
| 14 | + return() | ||
| 15 | +endif() | ||
| 16 | + | ||
| 17 | +set(CMAKE_ASC_ARCHITECTURES "dav-3510" CACHE STRING "Ascend 950 architecture") | ||
| 18 | +set(OPS_TENSOR_ROOT "" CACHE PATH "ops-tensor source with initialized include/tensor_api submodule") | ||
| 19 | +foreach(dep include/blaze/gemm/block/block_mmad_matmul_basic.h | ||
| 20 | + include/tensor_api/include/tensor_api/tensor.h examples/common/cce_pipe_stub.h) | ||
| 21 | + if(NOT EXISTS "${OPS_TENSOR_ROOT}/${dep}") | ||
| 22 | + message(FATAL_ERROR "Missing ${OPS_TENSOR_ROOT}/${dep}; set OPS_TENSOR_ROOT and initialize its Tensor API submodule") | ||
| 23 | + endif() | ||
| 24 | +endforeach() | ||
| 25 | +find_package(ASC REQUIRED) | ||
| 26 | +project(ain_all_gather_matmul LANGUAGES ASC CXX) | ||
| 27 | +set(CANN_ROOT "$ENV{ASCEND_HOME_PATH}") | ||
| 28 | +set(CANN_ARCH "${CMAKE_SYSTEM_PROCESSOR}-linux") | ||
| 29 | +add_executable(ain_all_gather_matmul main.asc) | ||
| 30 | +set_target_properties(ain_all_gather_matmul PROPERTIES LINKER_LANGUAGE ASC CXX_STANDARD 17) | ||
| 31 | +target_include_directories(ain_all_gather_matmul PRIVATE | ||
| 32 | + . ../.. | ||
| 33 | + "${OPS_TENSOR_ROOT}/include" "${OPS_TENSOR_ROOT}/include/blaze" | ||
| 34 | + "${OPS_TENSOR_ROOT}/include/tensor_api/include" "${OPS_TENSOR_ROOT}/include/tensor_api" | ||
| 35 | + "${CANN_ROOT}/include" "${CANN_ROOT}/${CANN_ARCH}/include" | ||
| 36 | + "${CANN_ROOT}/asc/include/comm_api/aicore" "${CANN_ROOT}/asc/include/comm_api/hcomm" | ||
| 37 | + "${CANN_ROOT}/${CANN_ARCH}/asc" "${CANN_ROOT}/${CANN_ARCH}/asc/include" | ||
| 38 | + "${CANN_ROOT}/${CANN_ARCH}/asc/include/tiling" | ||
| 39 | + "${CANN_ROOT}/include/op_common" "${CANN_ROOT}/pkg_inc/op_common" | ||
| 40 | + "${CANN_ROOT}/pkg_inc/base" "${CANN_ROOT}/pkg_inc" | ||
| 41 | + "${CANN_ROOT}/asc/impl/c_api" "${CANN_ROOT}/asc/impl/utils" | ||
| 42 | + "${CANN_ROOT}/compiler/ascendc/include/highlevel_api") | ||
| 43 | +target_link_directories(ain_all_gather_matmul PRIVATE "${CANN_ROOT}/lib64" "${CANN_ROOT}/${CANN_ARCH}/lib64") | ||
| 44 | +target_link_libraries(ain_all_gather_matmul PRIVATE hcomm hccl ascendcl ascendc_runtime runtime platform tiling_api dl m) | ||
| 45 | +target_compile_options(ain_all_gather_matmul PRIVATE | ||
| 46 | + $<$<COMPILE_LANGUAGE:ASC>:--npu-arch=${CMAKE_ASC_ARCHITECTURES}> | ||
| 47 | + $<$<COMPILE_LANGUAGE:ASC>:-iquote${OPS_TENSOR_ROOT}/include/tensor_api/include> | ||
| 48 | + $<$<COMPILE_LANGUAGE:ASC>:-iquote${OPS_TENSOR_ROOT}/include/tensor_api> | ||
| 49 | + $<$<COMPILE_LANGUAGE:ASC>:-include${OPS_TENSOR_ROOT}/examples/common/cce_pipe_stub.h>) | ||
| @@ -0,0 +1,149 @@ | |||
| 1 | +# AIN + Blaze AllGatherMatmul | ||
| 2 | + | ||
| 3 | +简体中文 | [English](README_en.md) | ||
| 4 | + | ||
| 5 | +本示例在一个 MIX AIC:AIV=1:1 kernel 中组合 AIN `Get` 和 Blaze `BlockMmad`,演示通信与计算重叠。 | ||
| 6 | +Host 使用 HCCL 建域和对称 window/team 资源接口,数据面由 AIV 提交 AIN 通信,AIC 执行矩阵乘。 | ||
| 7 | + | ||
| 8 | +## 语义与边界 | ||
| 9 | + | ||
| 10 | +每个 rank `r` 提供 `A_r[M,K]` 和 `B_r[K,N]`,得到: | ||
| 11 | + | ||
| 12 | +```text | ||
| 13 | +C_r[P*M,N] = Concat(A_0, ..., A_(P-1)) @ B_r | ||
| 14 | +``` | ||
| 15 | + | ||
| 16 | +`B_r` 可以在各 rank 上不同,输出按源 rank 排序。本版支持: | ||
| 17 | + | ||
| 18 | +- Ascend 950PR/950DT、`dav-3510`、单机同一互联组内的 2~64 rank(且不超过可用设备数)。 | ||
| 19 | +- FP16/BF16 输入和同类型输出、FP32 累加、ND 连续内存、无转置、无 bias。 | ||
| 20 | +- `M>0`,N/K 为正的 16 元素倍数;M 可以不对齐,支持通信和 Cube tile 尾块。 | ||
| 21 | +- 独立示例接口,不提供 ACLNN 注册、量化、变长 rank 输入或公开的 gatherOut 输出。 | ||
| 22 | + | ||
| 23 | +**必须按实际互联拓扑选择设备。** 例如 A5_4p 的 0~3 和 4~7 分属两个互联组,四卡应选择 | ||
| 24 | +`--devices 0,1,2,3` 或 `--devices 4,5,6,7`,不能使用 `1,2,3,4`。编号是当前 ACL 环境中的 | ||
| 25 | +逻辑 device ID;不要在同时设置可见设备映射后继续使用原物理编号。 | ||
| 26 | + | ||
| 27 | +## 数据流与同步 | ||
| 28 | + | ||
| 29 | +1. 为各 rank 的 A 注册发送 window,为 `[P,M,K]` 接收区注册接收 window,创建 AIV/UB_CTP team。 | ||
| 30 | +2. 将远端 peer 分配给 AIV,每个本端 peer channel 只有一个提交者和一个完成等待者。 | ||
| 31 | +3. AIC 先直接计算本地 `A_r @ B_r`;AIV 同时按 M 分块向所有远端发起 Get。 | ||
| 32 | +4. 每轮 AIV 使用 `FlushAsync + Wait` 等待自己的 peer channel,再进行本卡 AIV barrier。 | ||
| 33 | + 每个 AIV 通知配对 AIC,本轮远端数据已就绪;AIC 从该轮所有 peer 的 M/N tiles 中分配任务。 | ||
| 34 | +5. 下一轮 Get 与本轮远端计算重叠。两个 ready/free 槽限制预取深度;最后排空所有通知。 | ||
| 35 | + | ||
| 36 | +接收 GM 使用完整 rank-major 空间;数据不会在同一 kernel 内被循环覆盖。本 rank 的接收槽未使用, | ||
| 37 | +本地计算直接读原始 A,远端计算直接读接收 GM,没有额外 pack/unpack。每卡接收空间为 `2*P*M*K` 字节。 | ||
| 38 | + | ||
| 39 | +AIN 完成等待、本卡核间通知、AIC MTE2 加载顺序分别处理: | ||
| 40 | + | ||
| 41 | +- 每个 Get 立即提交,并在下一轮复用同一 channel 前等待完成;单次不超过 256 MiB。 | ||
| 42 | +- mode 0 / flag 0 用于 AIV 全核同步;mode 2 / flags 2、3 用于 ready,flags 4、5 用于 free。 | ||
| 43 | + flag ID 不随分块数增长。空闲计算核也参与通知消费和归还。 | ||
| 44 | +- 融合 kernel 使用 `__schedmode__(1)`;AIC 等待阻塞 MTE2,接收 A 的 Tensor 关闭 L2 cache hint。 | ||
| 45 | +- Host 在每次调用前 barrier、stream 完成后汇集各 rank 时间,保护远端输入生命周期。期间 A 不变。 | ||
| 46 | + TCP 同步调用位于 device event 之外;公共控制通道启用 `TCP_NODELAY`,避免延迟 ACK 造成 | ||
| 47 | + 毫秒级启动偏斜。剩余启动偏斜仍可能进入 collective 的等待时间。没有逐轮跨卡 barrier。 | ||
| 48 | + | ||
| 49 | +计算部分复用 Blaze Basic `BlockMmad`。默认 Cube M/N tile 为 256,K L0/L1 为 64/128、L1 双缓冲。 | ||
| 50 | +通信 tile 自动选择为能提供至少一轮 AIC 工作量的 M 分块,并受 M 和 Get 长度限制;也可用 | ||
| 51 | +`--tile-m` 固定。非最后一个通信块的行数必须为 256 的倍数。自动值是初始启发式,不是性能最优保证。 | ||
| 52 | + | ||
| 53 | +## 构建 | ||
| 54 | + | ||
| 55 | +需要同时包含 AIN API、HCCL team/window API 和相应实现的 CANN 包。实际验证版本见 | ||
| 56 | +[VALIDATION.md](VALIDATION.md),不能只凭版本号推断包中一定带有这些接口。 | ||
| 57 | + | ||
| 58 | +Blaze 为外部头文件依赖,不复制到 asc-comm,也不修改 ops-transformer。配套依赖版本: | ||
| 59 | + | ||
| 60 | +```text | ||
| 61 | +ops-tensor: e6cd0c20c716813ac4c76f2dba7152f63a4514ab | ||
| 62 | +Tensor API submodule: fd51475d8ceb26eb21befffe4391528bc16726a3 | ||
| 63 | +``` | ||
| 64 | + | ||
| 65 | +```bash | ||
| 66 | +source /path/to/cann/set_env.sh | ||
| 67 | +git clone https://gitcode.com/cann/ops-tensor.git /path/to/ops-tensor | ||
| 68 | +git -C /path/to/ops-tensor checkout e6cd0c20c716813ac4c76f2dba7152f63a4514ab | ||
| 69 | +git -C /path/to/ops-tensor submodule update --init --recursive include/tensor_api | ||
| 70 | + | ||
| 71 | +# 在本示例目录执行 | ||
| 72 | +cmake -S . -B build -DOPS_TENSOR_ROOT=/path/to/ops-tensor | ||
| 73 | +cmake --build build -j | ||
| 74 | +``` | ||
| 75 | + | ||
| 76 | +CMake 检查 Blaze/Tensor API 依赖,通过 CANN ASC 编译,并沿用 Blaze 示例的 host `PIPE_FIX` 兼容头。 | ||
| 77 | +ASC 默认目标为 `dav-3510`。本示例无需安装自定义 OPP 包。 | ||
| 78 | + | ||
| 79 | +## 运行与对照模式 | ||
| 80 | + | ||
| 81 | +```bash | ||
| 82 | +# 两卡,小 shape 全量随机 golden,串行与融合对照 | ||
| 83 | +./build/ain_all_gather_matmul tcp://127.0.0.1:29640 2 257 80 272 \ | ||
| 84 | + --devices 1,2 --tile-m 256 --random --warmup 1 --iters 5 | ||
| 85 | + | ||
| 86 | +# 四卡稳态测量;按本机实际互联组修改 devices | ||
| 87 | +./build/ain_all_gather_matmul tcp://127.0.0.1:29640 4 1024 4096 4096 \ | ||
| 88 | + --devices 0,1,2,3 --tile-m 256 --dtype bf16 --trace | ||
| 89 | +``` | ||
| 90 | + | ||
| 91 | +位置参数为 `endpoint ranks M K N`。进程管理器在本机 fork 全部 rank,无需 mpirun。 | ||
| 92 | + | ||
| 93 | +| 参数 | 默认值 | 含义 | | ||
| 94 | +| --- | --- | --- | | ||
| 95 | +| `--devices` | `0,...,P-1` | 每个 rank 对应的 ACL device ID,必须唯一且拓扑可达 | | ||
| 96 | +| `--dtype` | `fp16` | `fp16` / `bf16` | | ||
| 97 | +| `--mode` | `both` | `both` 依次运行 serial、fused | | ||
| 98 | +| `--tile-m` | 自动 | 通信 M 分块,末块按真实行数处理 | | ||
| 99 | +| `--cores` | 设备 AIC 数 | 启动核数;可用 1 验证每核多个 peer | | ||
| 100 | +| `--warmup` / `--iters` | 10 / 100 | 预热 / 稳态次数 | | ||
| 101 | +| `--verify-iters` | 2 | 改变输入并验证的次数;测量后再验证最终结果 | | ||
| 102 | +| `--random` | 关闭 | 全 K 随机输入与 O(PMKN) CPU golden,只建议小 shape | | ||
| 103 | +| `--trace` | 关闭 | 为 fused 额外运行一次带系统 cycle 记录的诊断 kernel | | ||
| 104 | +| `--timeout` | 300 秒 | 整个进程组的超时,包含初始化、校验与测量 | | ||
| 105 | +| `--delay-rank` / `--delay-ms` | 不延迟 / 0 | 输入就绪后延迟指定 rank 的启动,检查输入生命周期 | | ||
| 106 | + | ||
| 107 | +| mode | 行为 | | ||
| 108 | +| --- | --- | | ||
| 109 | +| `fused` | 一个 MIX kernel 内执行 AIN Get 和 Blaze 计算 | | ||
| 110 | +| `serial` | AIN Get kernel 完成后,启动相同 Blaze 计算 kernel | | ||
| 111 | +| `hccl` | HCCL AllGather 后启动相同 Blaze 计算 kernel;HCCL 使用环境默认算法 | | ||
| 112 | +| `comm` | 独立 AIN Get,校验远端接收数据 | | ||
| 113 | +| `matmul` | Host 填充完整输入,独立验证 Blaze 计算 | | ||
| 114 | + | ||
| 115 | +默认数据在 K 上以 16 为周期,rank/行/列/迭代参与数据生成,可在大 shape 下快速校验**所有输出元素**。 | ||
| 116 | +`--random` 使用全部 K 索引独立生成数据。两种输入均为可精确表示的 1/16 倍数。CPU golden 使用 | ||
| 117 | +FP32 累加并转成输出类型;FP16/BF16 比较阈值分别为 `0.001 + 0.001*abs(expected)` 和 | ||
| 118 | +`0.001 + 0.008*abs(expected)`,非有限值判失败。远端接收 A 逐元素按位比较。 | ||
| 119 | + | ||
| 120 | +`PERF` 记录每次所有 rank device event 耗时的最大值,再统计 median/P95。资源创建、输入复制、 | ||
| 121 | +CPU 校验和 TCP 同步均在计时区间外。serial/hccl 的时间包含两个 kernel 间隔。 | ||
| 122 | +这是稳态设备执行时间,不是包括 Host 建域的端到端请求延迟,也不是跨设备时钟校准后的全局包络。 | ||
| 123 | +`host_submit_skew_p95_us` 记录各 rank 调用起始 event 前的 Host 单调时钟偏斜 P95,用于检查 | ||
| 124 | +Host 发射一致性;它不是设备 kernel 启动偏斜,也不从 event 耗时中直接扣除。 | ||
| 125 | + | ||
| 126 | +`TRACE` 是独立诊断 kernel 中配对 AIV0/AIC0 的系统 cycle 区间;950 上 1000 cycles = 1 us。 | ||
| 127 | +计算阶段记录前后同步流水,ready 等待不计入对应远端计算阶段。诊断会改变时序,**不混入 PERF 样本**; | ||
| 128 | +阶段区间相交表示阶段并行,不能解释为网络链路/Cube 单元在每个 cycle 都同时繁忙。不同 rank 的 | ||
| 129 | +绝对 cycle 不用于跨卡比较。 | ||
| 130 | + | ||
| 131 | +## 验证 | ||
| 132 | + | ||
| 133 | +```bash | ||
| 134 | +# 可在无 CANN 的 macOS/Linux 上运行;不代表 NPU 覆盖 | ||
| 135 | +cmake -S . -B /tmp/ain-agmm-host -DAGMM_HOST_TESTS_ONLY=ON | ||
| 136 | +cmake --build /tmp/ain-agmm-host | ||
| 137 | +ctest --test-dir /tmp/ain-agmm-host --output-on-failure | ||
| 138 | + | ||
| 139 | +# 目标环境:按实际互联组选择设备,输出目录必须不存在 | ||
| 140 | +python3 tests/run_device_tests.py --binary build/ain_all_gather_matmul \ | ||
| 141 | + --devices 0,1,2,3 --output-dir device-results | ||
| 142 | +``` | ||
| 143 | + | ||
| 144 | +设备测试覆盖 FP16/BF16、M/N/K tile 尾块、18 个通信分块、单核多 peer、延迟 rank、独立通信和 | ||
| 145 | +独立计算、HCCL 对照,以及资源复用时改变输入。每条命令和日志保存在输出目录,`results.json` | ||
| 146 | +保留退出码与判定。只有全部 rank 校验通过且进程组成功,才打印最终 `Status=PASS`。 | ||
| 147 | + | ||
| 148 | +任何特定 shape 的收益都不能推广为全量优于传统 MC2;传统 MC2 算子还需用匹配的 shape、dtype、 | ||
| 149 | +布局和通信引擎另行对照。本示例没有把 HCCL cascade 标成传统 MC2 性能。 | ||
| @@ -0,0 +1,126 @@ | |||
| 1 | +# AIN + Blaze AllGatherMatmul | ||
| 2 | + | ||
| 3 | +[简体中文](README.md) | English | ||
| 4 | + | ||
| 5 | +This example composes AIN `Get` on AIV with Blaze `BlockMmad` on AIC in a single 1:1 MIX kernel: | ||
| 6 | + | ||
| 7 | +```text | ||
| 8 | +C_r[P*M,N] = Concat(A_0[M,K], ..., A_(P-1)[M,K]) @ B_r[K,N] | ||
| 9 | +``` | ||
| 10 | + | ||
| 11 | +Each rank may have different B. C is ordered by source rank. The initial scope is Ascend 950PR/950DT, | ||
| 12 | +single-node ranks in one connected fabric group, FP16/BF16 input and output with FP32 accumulation, | ||
| 13 | +contiguous ND, no transpose or bias. M may be unaligned; N/K must be positive multiples of 16. | ||
| 14 | +This is a standalone example, not an ACLNN operator or a quantized MC2 replacement. | ||
| 15 | + | ||
| 16 | +## Dependencies and build | ||
| 17 | + | ||
| 18 | +Use a CANN installation containing the AIN, HCCL team/window APIs and their implementations. | ||
| 19 | +See [VALIDATION.md](VALIDATION.md) for the actual tested environment. | ||
| 20 | +Blaze remains an external header dependency; ops-transformer is not needed. | ||
| 21 | + | ||
| 22 | +```bash | ||
| 23 | +source /path/to/cann/set_env.sh | ||
| 24 | +git clone https://gitcode.com/cann/ops-tensor.git /path/to/ops-tensor | ||
| 25 | +git -C /path/to/ops-tensor checkout e6cd0c20c716813ac4c76f2dba7152f63a4514ab | ||
| 26 | +git -C /path/to/ops-tensor submodule update --init --recursive include/tensor_api | ||
| 27 | +# Tensor API revision: fd51475d8ceb26eb21befffe4391528bc16726a3 | ||
| 28 | +cmake -S . -B build -DOPS_TENSOR_ROOT=/path/to/ops-tensor | ||
| 29 | +cmake --build build -j | ||
| 30 | +``` | ||
| 31 | + | ||
| 32 | +The CMake target uses ASC/dav-3510 and the host PIPE_FIX compatibility header supplied by Blaze. | ||
| 33 | +No custom OPP package needs installation. | ||
| 34 | + | ||
| 35 | +## Scheduling and memory | ||
| 36 | + | ||
| 37 | +AIC first computes its local A, directly into its rank's output rows. Meanwhile AIVs pull M chunks | ||
| 38 | +from peer windows. Each local peer channel has exactly one owner. That owner submits all assigned | ||
| 39 | +Gets and waits per channel using `FlushAsync`/`Wait`. All local AIVs then synchronize and notify | ||
| 40 | +their paired AIC. Remote tiles across peers are distributed together to the compute cores. | ||
| 41 | + | ||
| 42 | +The next communication round overlaps with the previous computation round. Two ready/free slots | ||
| 43 | +bound the notification count; all cores, including idle cores, participate. There is no growing | ||
| 44 | +per-round flag ID. AIV mode-0 flag 0 is the local barrier, mode-2 flags 2/3 are ready and 4/5 are free. | ||
| 45 | +The fused kernel uses batch scheduling. AIC waits on MTE2 before reading received A, whose Tensor | ||
| 46 | +disables the L2 cache hint. All final acknowledgements are consumed. | ||
| 47 | + | ||
| 48 | +The receive window is a full rank-major allocation of `2*P*M*K` bytes, without in-kernel reuse. | ||
| 49 | +The self slot is intentionally unused. No packing or output reordering is needed. Host collectives | ||
| 50 | +before launch and after stream completion protect remote A lifetime. Inputs stay immutable during | ||
| 51 | +each invocation. Host collectives are outside event timing. The shared TCP control channel enables | ||
| 52 | +`TCP_NODELAY` to avoid delayed-ACK launch skew. Residual launch skew can still enter collective | ||
| 53 | +wait time; there is no inter-rank barrier per chunk. | ||
| 54 | + | ||
| 55 | +Blaze tiles start at M/N=256, K L0/L1=64/128, with two L1 stages. Communication tile M is chosen | ||
| 56 | +to supply at least one compute wave, capped by M and the 256 MiB Get limit. Non-final communication | ||
| 57 | +tiles must have a multiple of 256 rows. This heuristic is a starting point, not a performance guarantee. | ||
| 58 | + | ||
| 59 | +## Run | ||
| 60 | + | ||
| 61 | +```bash | ||
| 62 | +./build/ain_all_gather_matmul tcp://127.0.0.1:29640 2 257 80 272 \ | ||
| 63 | + --devices 1,2 --tile-m 256 --random --warmup 1 --iters 5 | ||
| 64 | +``` | ||
| 65 | + | ||
| 66 | +Positional arguments are `endpoint ranks M K N`. The executable forks local ranks without MPI. | ||
| 67 | +Select devices in the **same physical connectivity group**. On A5_4p, groups are 0–3 and 4–7: | ||
| 68 | +use `--devices 0,1,2,3` for four ranks, not `1,2,3,4`. IDs refer to the current ACL logical device | ||
| 69 | +namespace; account for any externally configured visibility mapping. | ||
| 70 | + | ||
| 71 | +| Option | Default | Meaning | | ||
| 72 | +| --- | --- | --- | | ||
| 73 | +| `--devices` | `0,...,P-1` | One unique ACL device ID per rank | | ||
| 74 | +| `--dtype` | fp16 | fp16 or bf16 | | ||
| 75 | +| `--mode` | both | both runs serial then fused | | ||
| 76 | +| `--tile-m` | auto | Communication chunk rows | | ||
| 77 | +| `--cores` | device AIC count | Launch core count | | ||
| 78 | +| `--warmup`, `--iters` | 10, 100 | Warmup and measurement counts | | ||
| 79 | +| `--verify-iters` | 2 | Reuse resources with changed inputs and validate | | ||
| 80 | +| `--random` | off | Full-K random data and O(PMKN) CPU reference, for small shapes | | ||
| 81 | +| `--trace` | off | One additional instrumented fused invocation, outside measurements | | ||
| 82 | +| `--timeout` | 300 seconds | Whole process-group timeout including CPU checks | | ||
| 83 | +| `--delay-rank`, `--delay-ms` | none, 0 | Delay one rank after input readiness | | ||
| 84 | + | ||
| 85 | +Modes: `fused` is one MIX kernel; `serial` is AIN Get then the same Blaze compute kernel; | ||
| 86 | +`hccl` is HCCL AllGather then that compute kernel; `comm` checks AIN only; `matmul` checks Blaze | ||
| 87 | +with host-filled inputs. HCCL uses the environment's default algorithm. | ||
| 88 | + | ||
| 89 | +Default fixtures have a K period of 16, with rank/row/column/iteration-dependent values, allowing | ||
| 90 | +full-output golden checks in O(PMN*16). `--random` uses independent full-K inputs. Values are | ||
| 91 | +exactly representable multiples of 1/16. Received A is checked bitwise. C is checked against an | ||
| 92 | +FP32 reference rounded to output dtype, with FP16 tolerance `0.001 + 0.001*abs(expected)` and | ||
| 93 | +BF16 tolerance `0.001 + 0.008*abs(expected)`. Non-finite values fail. Inputs are changed during | ||
| 94 | +validation; a final output check follows measurement. | ||
| 95 | + | ||
| 96 | +`PERF` reports median/P95 of each iteration's maximum rank device-event duration. Initialization, | ||
| 97 | +copies, CPU checks and TCP collectives are excluded. Serial/cascade durations include the gap | ||
| 98 | +between kernels. These are steady-state device intervals, not globally clock-aligned envelopes | ||
| 99 | +or end-to-end host request latency. | ||
| 100 | +`host_submit_skew_p95_us` reports the P95 spread of host monotonic timestamps immediately before | ||
| 101 | +start-event submission. It is not device kernel-start skew and is not subtracted from event duration. | ||
| 102 | + | ||
| 103 | +`TRACE` records paired AIV0/AIC0 phase intervals using the same-device system counter | ||
| 104 | +(1000 cycles/us on 950). Compute intervals synchronize pipelines; remote readiness waits are | ||
| 105 | +outside compute intervals. Instrumentation changes timing and is excluded from `PERF` samples. | ||
| 106 | +Intersecting intervals prove overlapping phases, not that network/Cube hardware is busy every | ||
| 107 | +cycle. Do not compare absolute counter values across devices. | ||
| 108 | + | ||
| 109 | +## Checks | ||
| 110 | + | ||
| 111 | +```bash | ||
| 112 | +# Portable checks without CANN: NOT device execution | ||
| 113 | +cmake -S . -B /tmp/ain-agmm-host -DAGMM_HOST_TESTS_ONLY=ON | ||
| 114 | +cmake --build /tmp/ain-agmm-host | ||
| 115 | +ctest --test-dir /tmp/ain-agmm-host --output-on-failure | ||
| 116 | + | ||
| 117 | +# NPU regression; output directory must not already exist | ||
| 118 | +python3 tests/run_device_tests.py --binary build/ain_all_gather_matmul \ | ||
| 119 | + --devices 0,1,2,3 --output-dir device-results | ||
| 120 | +``` | ||
| 121 | + | ||
| 122 | +The suite checks both dtypes, tails, 18 communication rounds, one core owning multiple peers, | ||
| 123 | +delayed rank entry, independent communication/computation and the HCCL cascade. Logs and | ||
| 124 | +`results.json` preserve commands and exit status. The process group reports PASS only if all | ||
| 125 | +ranks succeed. A faster result for a particular shape is not a general claim against traditional | ||
| 126 | +MC2; compare that operator separately with matching layout, dtype and communication engine. | ||
| @@ -0,0 +1,134 @@ | |||
| 1 | +# AIN + Blaze AllGatherMatmul 验证记录 | ||
| 2 | + | ||
| 3 | +验证日期:2026-09-24。结论针对本示例当前工作区版本,不代表传统 MC2 算子或所有 shape。 | ||
| 4 | + | ||
| 5 | +## 环境和版本 | ||
| 6 | + | ||
| 7 | +| 项目 | 实际值 | | ||
| 8 | +| --- | --- | | ||
| 9 | +| asc-comm 基线 | `master`,`79de7e489db8db78f8fe3b616682099a888a204d`,叠加本次示例和 TCP 控制通道修改 | | ||
| 10 | +| ops-tensor | `e6cd0c20c716813ac4c76f2dba7152f63a4514ab` | | ||
| 11 | +| Tensor API submodule | `fd51475d8ceb26eb21befffe4391528bc16726a3` | | ||
| 12 | +| 目标 | SSH `A5_4p`,Linux x86_64,Ascend950PR,28 AIC | | ||
| 13 | +| Device | 两卡 `0,1`;四卡 `0,1,2,3` | | ||
| 14 | +| CANN 环境入口 | `/mnt/docker/jxy/cann/cann/set_env.sh` | | ||
| 15 | +| 实际 ASCEND_HOME_PATH | `/mnt/docker/jxy/cann/cann-9.3.0` | | ||
| 16 | +| 安装清单版本 | runtime / asc-devkit / hccl / hcomm 均为 `Version=9.2.0`,`timestamp=20260923_000324205` | | ||
| 17 | +| 驱动工具 | `npu-smi 25.7.rc1.6` | | ||
| 18 | +| 构建 / 启动 | CANN ASC,`dav-3510`;示例本机 fork,无 MPI、无自定义 OPP 安装 | | ||
| 19 | + | ||
| 20 | +目录名与安装清单版本不一致,不能把本次验证简称为“CANN 9.3 验证”。`ldd` 确认 HCOMM、HCCL、 | ||
| 21 | +runtime 等库来自上述实际 CANN 路径。AIN 头文件和运行库使用该安装包;本次没有重新构建整套 asc-comm。 | ||
| 22 | +Blaze 编译保留一条外部 `dispatch_policy.h:67` 的 `cce_global attribute ignored` 警告,最终构建成功。 | ||
| 23 | + | ||
| 24 | +远端独立工作目录为 `/mnt/docker/jxy/ain-blaze-agmm.GilOAt`。源码位于 `src`,构建产物位于 `build`。 | ||
| 25 | +执行后 `npu-smi` 显示 0~3 卡无运行进程;5~7 卡的其他用户进程未改动。测试并非整机独占性能验收。 | ||
| 26 | + | ||
| 27 | +## 构建和正确性 | ||
| 28 | + | ||
| 29 | +- macOS arm64:便携 Host CMake/CTest 通过,覆盖参数、溢出、尾块地址、数值转换和 golden;不代表设备执行。 | ||
| 30 | +- Linux x86_64 + ASC:最终版本编译成功,包含普通和独立 cycle 采样 kernel。 | ||
| 31 | +- NPU:最终版本 **11/11 用例通过**,所有参与 rank 校验成功、进程退出码 0。 | ||
| 32 | + | ||
| 33 | +| 用例 | Rank | M / K / N | 重点 | | ||
| 34 | +| --- | --- | --- | --- | | ||
| 35 | +| matmul_random | 2 | 17 / 80 / 272 | 独立 Blaze,完整 K 随机参考 | | ||
| 36 | +| get_random | 2 | 257 / 80 / 272 | 独立 Get,逐元素按位校验 | | ||
| 37 | +| fused_fp16_tail | 2 | 257 / 80 / 272 | FP16,M/N/K Cube 尾块 | | ||
| 38 | +| fused_bf16_tail | 2 | 513 / 144 / 272 | BF16,完整 K 随机参考 | | ||
| 39 | +| many_rounds | 2 | 4353 / 80 / 272 | 18 轮通信,循环复用 ready/free flag | | ||
| 40 | +| one_core | 2 | 257 / 80 / 272 | 单个 AIC/AIV 配对 | | ||
| 41 | +| delayed_peer | 2 | 257 / 80 / 272 | rank 1 输入就绪后延迟 20 ms 启动 | | ||
| 42 | +| hccl_cascade | 2 | 257 / 80 / 272 | HCCL AllGather + 相同 Blaze 计算 | | ||
| 43 | +| four_rank_fp16 | 4 | 513 / 144 / 528 | FP16 四卡多 peer | | ||
| 44 | +| four_rank_bf16 | 4 | 257 / 80 / 272 | BF16 四卡、完整 K 随机参考 | | ||
| 45 | +| multiple_peers_per_core | 4 | 257 / 80 / 272 | 单个 AIV 提交并等待多个 peer channel | | ||
| 46 | + | ||
| 47 | +每项默认改变输入验证两次,再预热 1 次、测量 5 次并校验最终输出。除独立模式和 HCCL 用例外, | ||
| 48 | +同时覆盖 AIN 串行与融合版本。远端接收 A 按位检查;C 全元素检查,阈值见 README。 | ||
| 49 | + | ||
| 50 | +复现命令(在远端工作目录执行): | ||
| 51 | + | ||
| 52 | +```bash | ||
| 53 | +source /mnt/docker/jxy/cann/cann/set_env.sh | ||
| 54 | +cmake -S src/examples/aicore/ain/02_all_gather_matmul -B build \ | ||
| 55 | + -DOPS_TENSOR_ROOT="$PWD/ops-tensor" | ||
| 56 | +cmake --build build -j4 | ||
| 57 | +python3 src/examples/aicore/ain/02_all_gather_matmul/tests/run_device_tests.py \ | ||
| 58 | + --binary build/ain_all_gather_matmul --devices 0,1,2,3 \ | ||
| 59 | + --output-dir device-tests-rerun --port 29800 | ||
| 60 | +``` | ||
| 61 | + | ||
| 62 | +此前误用 `1,2,3,4` 的四卡测试在资源初始化失败:`HcclTeamCreate` 返回 6/1、 | ||
| 63 | +`HcclCommSymWinRegister` 返回 9,尚未进入计算 kernel。按用户提供的 0~3 / 4~7 拓扑改用 | ||
| 64 | +`0,1,2,3` 后通过。失败日志保留于 `device-tests/four_rank_fp16.log`,不计入验收或性能统计。 | ||
| 65 | + | ||
| 66 | +## 稳态性能 | ||
| 67 | + | ||
| 68 | +四卡 `0,1,2,3`,每 rank `M=1024, K=4096, N=1024`,`tileM=256`,28 核,预热 10 次、 | ||
| 69 | +测量 100 次。每次取所有 rank 的 device event 耗时最大值,再统计 median/P95;没有混入诊断采样。 | ||
| 70 | + | ||
| 71 | +| dtype | 模式 | median / us | P95 / us | Host 提交偏斜 P95 / us | | ||
| 72 | +| --- | --- | ---: | ---: | ---: | | ||
| 73 | +| FP16 | AIN serial | 408.699 | 410.609 | 6.129 | | ||
| 74 | +| FP16 | AIN fused | 257.151 | 260.299 | 6.039 | | ||
| 75 | +| BF16 | AIN serial | 409.139 | 412.402 | 10.416 | | ||
| 76 | +| BF16 | AIN fused | 257.781 | 262.249 | 6.119 | | ||
| 77 | +| FP16 | HCCL cascade | 491.656 | 505.657 | 6.049 | | ||
| 78 | + | ||
| 79 | +该 shape 的融合版本相对 AIN 串行版本,中位耗时分别降低 **37.081% / 36.994%**, | ||
| 80 | +加速比分别为 **1.589x / 1.587x**。两者采用相同 Blaze 计算逻辑;结果含 kernel 间隔与实际调度代价。 | ||
| 81 | +HCCL cascade 使用环境默认通信算法,不等同于 ops-transformer 的传统 MC2 融合算子。 | ||
| 82 | + | ||
| 83 | +```bash | ||
| 84 | +# FP16;BF16 使用端口 29821 并增加 --dtype bf16 | ||
| 85 | +./build/ain_all_gather_matmul tcp://127.0.0.1:29820 4 1024 4096 1024 \ | ||
| 86 | + --devices 0,1,2,3 --tile-m 256 --mode both --trace --timeout 600 | ||
| 87 | +# HCCL cascade | ||
| 88 | +./build/ain_all_gather_matmul tcp://127.0.0.1:29822 4 1024 4096 1024 \ | ||
| 89 | + --devices 0,1,2,3 --tile-m 256 --mode hccl --timeout 600 | ||
| 90 | +``` | ||
| 91 | + | ||
| 92 | +### TCP 同步造成的旧基线偏差 | ||
| 93 | + | ||
| 94 | +修复前,公共 TCP 控制通道对消息头、载荷分开发送,未禁用 Nagle。四进程独立 probe 发现 | ||
| 95 | +连续 barrier 的 rank 退出偏斜约 40,980~41,000 us;旧 HCCL cascade 中位耗时为 42,627.992 us。 | ||
| 96 | +这组数据存在明显启动偏斜,不能用于加速比结论。 | ||
| 97 | + | ||
| 98 | +给已连接 socket 设置 `TCP_NODELAY` 后,相同 probe 的 10 次偏斜为 11.177~47.361 us; | ||
| 99 | +最终算子测量另行记录了上表的 Host 提交偏斜。HCCL cascade 随之恢复至 491.656 us。 | ||
| 100 | +保留 `sync_probe.cpp`、`sync-before.log`、`sync-after.log` 和旧 `perf-*.log` 供复核。 | ||
| 101 | + | ||
| 102 | +TCP 调用位于 event 外部,并不意味着启动偏斜不会进入 collective 等待。Host 提交偏斜也不等于 | ||
| 103 | +设备 kernel 启动偏斜;这里没有直接扣除偏斜或声称已取得跨设备校时后的全局包络。 | ||
| 104 | + | ||
| 105 | +## 通信计算重叠证据 | ||
| 106 | + | ||
| 107 | +每种 dtype 额外执行一次带 cycle 记录的 kernel,采集每个 rank 的配对 AIV0/AIC0, | ||
| 108 | +共 72 个完整阶段区间。950 上按 1000 cycles/us 转换,各 rank 单独计算区间交集: | ||
| 109 | + | ||
| 110 | +- local matmul 与 Get 第 0 轮重叠。 | ||
| 111 | +- remote matmul 第 j 轮与 Get 第 j+1 轮重叠,j=0、1、2。 | ||
| 112 | +- FP16 的 16 个交集为 **41.677~42.947 us**,BF16 为 **41.813~44.413 us**。 | ||
| 113 | + | ||
| 114 | +例如 FP16 rank 0:本地计算与 Get 第 0 轮交集 41.840 us;远端计算第 0 轮与 Get 第 1 轮 | ||
| 115 | +交集 42.947 us。区间直接来自 `perf-final-fp16.log`,未使用跨 rank 的绝对 cycle 差。 | ||
| 116 | + | ||
| 117 | +这是阶段并行证据;采样 kernel 会改变流水时序,且只记录 core 0,不能据此声称所有核、网络链路 | ||
| 118 | +和 Cube 在整个交集内均忙碌。本次未采集完整 msprof 任务时间线。 | ||
| 119 | + | ||
| 120 | +## 证据位置与限制 | ||
| 121 | + | ||
| 122 | +本地原始证据目录: | ||
| 123 | +`/Users/izanaami/Documents/MC2_AICPU/logs/ain_blaze_agmm_20260924_GilOAt`。 | ||
| 124 | + | ||
| 125 | +- `device-tests-final/results.json`:最终 11 项完整命令、退出码及日志路径。 | ||
| 126 | +- `build-final.log`、`host-tests.log`:目标构建与便携测试分别记录。 | ||
| 127 | +- `perf-final-{fp16,bf16,hccl}.log`:最终性能、每 rank 校验和原始 cycle。 | ||
| 128 | +- `audit_results.py`、`audit-summary.json`:校验 11 项结果、5 行性能和 32 个正交集。 | ||
| 129 | +- `environment-final.log`、`package-versions.log`、`npu-after.log`:包路径、依赖和设备状态。 | ||
| 130 | +- `source-sha256.json`、`binary.sha256`:参与构建的示例/公共工具源码及二进制摘要;本地与远端源码比对一致。 | ||
| 131 | + | ||
| 132 | +已验证 2/4 rank,未验证跨互联组、8/16 rank、950DT 或其他 shape 的性能范围。参数允许的 rank | ||
| 133 | +上限不是设备验收范围。未进行 Get/Put A/B、完整传统 MC2 对照、转置/量化/ACLNN 接入和大规模 | ||
| 134 | +随机压力测试;当前 tile 参数也未做全面调优。本记录对应完成上板验证时的源码,交付状态以关联 PR 为准。 | ||
| @@ -0,0 +1,239 @@ | |||
| 1 | +/* | ||
| 2 | + * Copyright (c) 2026 Huawei Technologies Co., Ltd. | ||
| 3 | + * Licensed under the CANN Open Software License Agreement Version 2.0. | ||
| 4 | + * See LICENSE in the root of the software repository. | ||
| 5 | + */ | ||
| 6 | + | ||
| 7 | + | ||
| 8 | + | ||
| 9 | + | ||
| 10 | + | ||
| 11 | + | ||
| 12 | + | ||
| 13 | + | ||
| 14 | + | ||
| 15 | +namespace ain_agmm { | ||
| 16 | + | ||
| 17 | +// 同步分两层:mode 0 汇合本卡所有 AIV,mode 2 连接同 blockIdx 的 AIV/AIC。 | ||
| 18 | +// ready[0/1] 表示通信数据就绪,free[0/1] 表示计算已消费该轮通知。 | ||
| 19 | +// flag ID 固定复用,不能直接用 round 作 ID;也不能与 SyncAll 的保留标志混用。 | ||
| 20 | +constexpr uint16_t AIV_BARRIER = 0; | ||
| 21 | +constexpr uint16_t READY_BASE = 2; | ||
| 22 | +constexpr uint16_t FREE_BASE = 4; | ||
| 23 | + | ||
| 24 | +__aicore__ inline uint64_t DivUp(uint64_t n, uint64_t d) { return (n + d - 1) / d; } | ||
| 25 | +__aicore__ inline uint64_t Smaller(uint64_t a, uint64_t b) { return a < b ? a : b; } | ||
| 26 | + | ||
| 27 | +// 仅供诊断:记录当前阶段边界,只由 AIV0/AIC0 写各自区域。 | ||
| 28 | +// 先同步流水,再取本设备系统 cycle,并将标量写入刷出 cache;两类核的区域分开对齐。 | ||
| 29 | +// 区间包含阶段内等待,不等同于网络忙碌时间,也不能直接跨设备比较绝对 cycle。 | ||
| 30 | +__aicore__ inline void RecordCycle(GM_ADDR trace, uint64_t index) | ||
| 31 | +{ | ||
| 32 | + if (AscendC::GetBlockIdx() == 0) { | ||
| 33 | + AscendC::PipeBarrier<PIPE_ALL>(); | ||
| 34 | + const uint64_t cycle = static_cast<uint64_t>(AscendC::GetSystemCycle()); | ||
| 35 | + AscendC::GlobalTensor<uint64_t> ticks; | ||
| 36 | + ticks.SetGlobalBuffer(reinterpret_cast<__gm__ uint64_t*>(trace)); | ||
| 37 | + ticks.SetValue(index, cycle); | ||
| 38 | + AscendC::DataCacheCleanAndInvalid<uint64_t, AscendC::CacheLine::SINGLE_CACHE_LINE, | ||
| 39 | + AscendC::DcciDst::CACHELINE_OUT>(ticks[index]); | ||
| 40 | + } | ||
| 41 | +} | ||
| 42 | + | ||
| 43 | +// 拉取一轮 M 分块:从所有远端 A_peer[row:row+rows, :] 写入本卡 gathered[peer, row:row+rows, :]。 | ||
| 44 | +// Get 的 offset/length 单位均为字节,FP16/BF16 每元素占 2 字节;尾轮按实际 rows 搬运。 | ||
| 45 | +__aicore__ inline void GatherRound(AscendC::Ain<>& ain, AscendC::AinDescriptorUbuf& ub, | ||
| 46 | + GM_ADDR team, GM_ADDR srcWin, GM_ADDR dstWin, | ||
| 47 | + const Tiling& t, uint32_t rank, uint32_t round) | ||
| 48 | +{ | ||
| 49 | + const uint64_t row = static_cast<uint64_t>(round) * t.tileM; | ||
| 50 | + const uint64_t rows = Smaller(t.tileM, t.m - row); | ||
| 51 | + const uint64_t srcOffset = row * t.k * 2; | ||
| 52 | + // i 是剔除自身后的 peer 序号,例如 rank=1 时,i=0/1/2 对应 peer=0/2/3。 | ||
| 53 | + // 按 blockIdx 跨步分配,确保同一 peer channel 只有一个 AIV 提交和等待。 | ||
| 54 | + for (uint32_t i = AscendC::GetBlockIdx(); i < t.ranks - 1; i += t.cores) { | ||
| 55 | + const uint32_t peer = i < rank ? i : i + 1; | ||
| 56 | + const uint64_t dstOffset = (static_cast<uint64_t>(peer) * t.m + row) * t.k * 2; | ||
| 57 | + ain.Get(team, peer, dstWin, dstOffset, srcWin, srcOffset, rows * t.k * 2, ub); | ||
| 58 | + } | ||
| 59 | + // 先提交本核所有 peer 的 Get,再逐 channel 等待,避免每提交一个 peer 就立刻串行等待。 | ||
| 60 | + // FlushAsync 在此取得 channel,真正的完成等待由 Wait 承担;返回后本核负责的数据才可使用。 | ||
| 61 | + for (uint32_t i = AscendC::GetBlockIdx(); i < t.ranks - 1; i += t.cores) { | ||
| 62 | + const uint32_t peer = i < rank ? i : i + 1; | ||
| 63 | + AscendC::ChannelHandle channel; | ||
| 64 | + ain.FlushAsync(team, peer, &channel); | ||
| 65 | + ain.Wait(channel); | ||
| 66 | + } | ||
| 67 | +} | ||
| 68 | + | ||
| 69 | +// AIV 侧通信流水:Fused=true 时与 AIC 交换 ready/free;独立通信 kernel 不执行配对通知。 | ||
| 70 | +template <bool Fused, bool Trace = false> | ||
| 71 | +__aicore__ inline void Gather(GM_ADDR team, GM_ADDR srcWin, GM_ADDR dstWin, const Tiling& t, uint32_t rank, | ||
| 72 | + GM_ADDR trace = nullptr) | ||
| 73 | +{ | ||
| 74 | + // 每个 AIV 独占一块 UB,供 AIN 构造通信描述符;实际矩阵数据位于对称 GM window。 | ||
| 75 | + AscendC::TPipe pipe; | ||
| 76 | + AscendC::TBuf<AscendC::TPosition::VECOUT> buffer; | ||
| 77 | + pipe.InitBuffer(buffer, AscendC::HCOMM_URMA_TMP_BUF_SIZE); | ||
| 78 | + auto local = buffer.Get<uint8_t>(); | ||
| 79 | + AscendC::AinDescriptorUbuf ub{reinterpret_cast<__ubuf__ uint8_t*>(local.GetPhyAddr()), | ||
| 80 | + AscendC::HCOMM_URMA_TMP_BUF_SIZE, 0}; | ||
| 81 | + AscendC::Ain<> ain(0); | ||
| 82 | + for (uint32_t round = 0; round < t.rounds; ++round) { | ||
| 83 | + if constexpr (Fused) { | ||
| 84 | + // 前两轮可直接使用两个通知槽;第 r 轮复用槽前,消费第 r-2 轮的 free。 | ||
| 85 | + // 双槽限制通知的超前量,不是 GM 双缓冲:各轮数据写入不同 GM 行,不会覆盖。 | ||
| 86 | + if (round >= 2) { | ||
| 87 | + AscendC::CrossCoreWaitFlag<2, PIPE_S>(FREE_BASE + (round & 1)); | ||
| 88 | + } | ||
| 89 | + } | ||
| 90 | + if constexpr (Trace) { RecordCycle(trace, static_cast<uint64_t>(round) * 2); } | ||
| 91 | + GatherRound(ain, ub, team, srcWin, dstWin, t, rank, round); | ||
| 92 | + if constexpr (Fused) { | ||
| 93 | + // 先等本卡所有 AIV 的 Get 完成,再通知 AIC;AIC 会跨所有 peer 分配计算任务。 | ||
| 94 | + // 即使本 AIV 没分到 peer,也必须参加 barrier 并通知配对 AIC。 | ||
| 95 | + AscendC::CrossCoreSetFlag<0, PIPE_MTE3>(AIV_BARRIER); | ||
| 96 | + AscendC::CrossCoreWaitFlag<0, PIPE_S>(AIV_BARRIER); | ||
| 97 | + if constexpr (Trace) { RecordCycle(trace, static_cast<uint64_t>(round) * 2 + 1); } | ||
| 98 | + AscendC::CrossCoreSetFlag<2, PIPE_MTE3>(READY_BASE + (round & 1)); | ||
| 99 | + } | ||
| 100 | + } | ||
| 101 | + if constexpr (Fused) { | ||
| 102 | + // 循环中只消费了复用槽时的 free;退出前排空最后一到两轮的确认,避免遗留通知。 | ||
| 103 | + const uint32_t begin = t.rounds > 2 ? t.rounds - 2 : 0; | ||
| 104 | + for (uint32_t round = begin; round < t.rounds; ++round) { | ||
| 105 | + AscendC::CrossCoreWaitFlag<2, PIPE_S>(FREE_BASE + (round & 1)); | ||
| 106 | + } | ||
| 107 | + } | ||
| 108 | +} | ||
| 109 | + | ||
| 110 | +// 封装单个 AIC 的 Blaze 块矩阵乘;外层负责 peer/轮次/核分工,Blaze 负责块内 K 循环和搬运。 | ||
| 111 | +template <typename T> | ||
| 112 | +class Compute { | ||
| 113 | + using Layout = asc::te::nd_ext_layout_ptn; | ||
| 114 | + using Mmad = Blaze::Gemm::Block::BlockMmad<Blaze::Gemm::MatmulMultiBlockBasic<>, | ||
| 115 | + T, Layout, T, Layout, T, Layout, T, Layout>; | ||
| 116 | + using MakeLayout = asc::te::frame_layout_format<Layout>; | ||
| 117 | + Mmad mm_; | ||
| 118 | +public: | ||
| 119 | + // 设置 Cube tile:M/N=256,K 在 L1/L0 分别为 128/64,L1 使用两级缓冲。 | ||
| 120 | + __aicore__ inline void Init() | ||
| 121 | + { | ||
| 122 | + typename Mmad::Params p; | ||
| 123 | + p.mL1 = CUBE_M; | ||
| 124 | + p.nL1 = CUBE_N; | ||
| 125 | + p.kL1 = CUBE_K * 2; | ||
| 126 | + p.mL0 = CUBE_M; | ||
| 127 | + p.nL0 = CUBE_N; | ||
| 128 | + p.kL0 = CUBE_K; | ||
| 129 | + p.l1Stages = 2; | ||
| 130 | + p.l0cStages = 1; | ||
| 131 | + mm_.Init(p); | ||
| 132 | + } | ||
| 133 | + | ||
| 134 | + // 计算 [row, row+rows) 行范围。local=true 只读本地 A,否则读所有远端 peer 的接收区。 | ||
| 135 | + // B 始终是当前 rank 的 B;C 直接写入源 rank 对应的行段,后续无需重排。 | ||
| 136 | + __aicore__ inline void Run(GM_ADDR a, GM_ADDR gathered, GM_ADDR b, GM_ADDR c, | ||
| 137 | + const Tiling& t, uint32_t rank, uint64_t row, uint64_t rows, bool local) | ||
| 138 | + { | ||
| 139 | + const uint64_t mTiles = DivUp(rows, CUBE_M); | ||
| 140 | + const uint64_t nTiles = DivUp(t.n, CUBE_N); | ||
| 141 | + const uint64_t tilesPerPeer = mTiles * nTiles; | ||
| 142 | + const uint64_t tiles = tilesPerPeer * (local ? 1 : t.ranks - 1); | ||
| 143 | + auto gmB = asc::te::make_tensor(asc::te::make_mem_ptr<asc::te::location::gm>( | ||
| 144 | + reinterpret_cast<__gm__ T*>(b)), MakeLayout{}(static_cast<int64_t>(t.k), static_cast<int64_t>(t.n))); | ||
| 145 | + // 将 peer × M tile × N tile 展平成任务序号,按核号跨步领取,避免按 peer 固定绑定计算核。 | ||
| 146 | + for (uint64_t task = AscendC::GetBlockIdx(); task < tiles; task += t.cores) { | ||
| 147 | + const uint32_t index = static_cast<uint32_t>(task / tilesPerPeer); | ||
| 148 | + const uint32_t peer = local ? rank : (index < rank ? index : index + 1); | ||
| 149 | + const uint64_t tile = task % tilesPerPeer; | ||
| 150 | + const uint64_t m = (tile / nTiles) * CUBE_M; | ||
| 151 | + const uint64_t n = (tile % nTiles) * CUBE_N; | ||
| 152 | + const int64_t actualM = static_cast<int64_t>(Smaller(CUBE_M, rows - m)); | ||
| 153 | + const int64_t actualN = static_cast<int64_t>(Smaller(CUBE_N, t.n - n)); | ||
| 154 | + // gathered 看作 [P*M,K],C 看作 [P*M,N];peer*M 定位源 rank,row+m 定位块内行。 | ||
| 155 | + const uint64_t globalRow = static_cast<uint64_t>(peer) * t.m + row + m; | ||
| 156 | + auto aPtr = reinterpret_cast<__gm__ T*>(local ? a : gathered) + | ||
| 157 | + (local ? row + m : globalRow) * t.k; | ||
| 158 | + auto cPtr = reinterpret_cast<__gm__ T*>(c) + globalRow * t.n; | ||
| 159 | + auto gmA = asc::te::make_tensor(asc::te::make_mem_ptr<asc::te::location::gm>(aPtr), | ||
| 160 | + MakeLayout{}(actualM, static_cast<int64_t>(t.k))); | ||
| 161 | + // window 跨调用复用,接收地址可能曾被缓存;关闭 A 的 L2 hint,配合 ready 等待读取新数据。 | ||
| 162 | + gmA.set_l2_cache_hint(asc::te::cache_mode::disable); | ||
| 163 | + auto gmC = asc::te::make_tensor(asc::te::make_mem_ptr<asc::te::location::gm>(cPtr), | ||
| 164 | + MakeLayout{}(actualM, static_cast<int64_t>(t.n))); | ||
| 165 | + auto blockB = gmB.slice(asc::te::make_coord(0L, static_cast<int64_t>(n)), | ||
| 166 | + asc::te::make_shape(static_cast<int64_t>(t.k), actualN)); | ||
| 167 | + auto blockC = gmC.slice(asc::te::make_coord(0L, static_cast<int64_t>(n)), | ||
| 168 | + asc::te::make_shape(actualM, actualN)); | ||
| 169 | + // 本策略关闭 bias;接口仍需传入合法 Tensor 占位,不传空 Tensor 指针。 | ||
| 170 | + auto bias = blockC; | ||
| 171 | + typename Mmad::TupleShape shape{actualM, actualN, static_cast<int64_t>(t.k), 1L}; | ||
| 172 | + mm_(gmA, blockB, bias, blockC, shape); | ||
| 173 | + } | ||
| 174 | + } | ||
| 175 | +}; | ||
| 176 | + | ||
| 177 | +// AIC 侧计算流水:先算不依赖通信的本地 A,再逐轮消费远端块,与下一轮 Get 并行。 | ||
| 178 | +template <typename T, bool Fused, bool Trace = false> | ||
| 179 | +__aicore__ inline void Matmul(GM_ADDR a, GM_ADDR gathered, GM_ADDR b, GM_ADDR c, | ||
| 180 | + const Tiling& t, uint32_t rank, GM_ADDR trace = nullptr) | ||
| 181 | +{ | ||
| 182 | + Compute<T> compute; | ||
| 183 | + compute.Init(); | ||
| 184 | + // 前 2*rounds 个 uint64 保存 Get 起止时间;AIC 区域按 32 个 uint64 对齐,避免与 AIV 共用脏行。 | ||
| 185 | + const uint64_t traceBase = DivUp(static_cast<uint64_t>(t.rounds) * 2, 32) * 32; | ||
| 186 | + if constexpr (Trace) { RecordCycle(trace, traceBase); } | ||
| 187 | + compute.Run(a, gathered, b, c, t, rank, 0, t.m, true); | ||
| 188 | + if constexpr (Trace) { RecordCycle(trace, traceBase + 1); } | ||
| 189 | + for (uint32_t round = 0; round < t.rounds; ++round) { | ||
| 190 | + if constexpr (Fused) { | ||
| 191 | + // 诊断版阻塞标量流水,使后续时间戳不包含 ready 等待;普通版在 MTE2 上阻止提前加载 A。 | ||
| 192 | + if constexpr (Trace) { | ||
| 193 | + AscendC::CrossCoreWaitFlag<2, PIPE_S>(READY_BASE + (round & 1)); | ||
| 194 | + } else { | ||
| 195 | + AscendC::CrossCoreWaitFlag<2, PIPE_MTE2>(READY_BASE + (round & 1)); | ||
| 196 | + } | ||
| 197 | + } | ||
| 198 | + if constexpr (Trace) { RecordCycle(trace, traceBase + 2 + static_cast<uint64_t>(round) * 2); } | ||
| 199 | + const uint64_t row = static_cast<uint64_t>(round) * t.tileM; | ||
| 200 | + compute.Run(a, gathered, b, c, t, rank, row, Smaller(t.tileM, t.m - row), false); | ||
| 201 | + if constexpr (Trace) { RecordCycle(trace, traceBase + 3 + static_cast<uint64_t>(round) * 2); } | ||
| 202 | + if constexpr (Fused) { | ||
| 203 | + // 在 FIX 流水上归还槽,保持与结果写回的顺序;无计算任务的核也要归还,避免 AIV 死等。 | ||
| 204 | + AscendC::CrossCoreSetFlag<2, PIPE_FIX>(FREE_BASE + (round & 1)); | ||
| 205 | + } | ||
| 206 | + } | ||
| 207 | + AscendC::PipeBarrier<PIPE_ALL>(); | ||
| 208 | +} | ||
| 209 | + | ||
| 210 | +// 融合入口:同一 kernel 启动配对 AIC/AIV,分别进入计算/通信分支;调度属性配合跨核同步。 | ||
| 211 | +template <typename T, bool Trace = false> | ||
| 212 | +__global__ __aicore__ __schedmode__(1) void FusedKernel( | ||
| 213 | + GM_ADDR team, GM_ADDR srcWin, GM_ADDR dstWin, GM_ADDR a, GM_ADDR gathered, GM_ADDR b, GM_ADDR c, | ||
| 214 | + Tiling tiling, uint32_t rank, GM_ADDR trace) | ||
| 215 | +{ | ||
| 216 | + KERNEL_TASK_TYPE_DEFAULT(KERNEL_TYPE_MIX_AIC_1_1); | ||
| 217 | + AscendC::InitSocState(); | ||
| 218 | + if ASCEND_IS_AIV { Gather<true, Trace>(team, srcWin, dstWin, tiling, rank, trace); } | ||
| 219 | + if ASCEND_IS_AIC { Matmul<T, true, Trace>(a, gathered, b, c, tiling, rank, trace); } | ||
| 220 | +} | ||
| 221 | + | ||
| 222 | +// 串行基线的通信入口:只启动 AIV;Host 在同一 stream 后续排入 MatmulKernel。 | ||
| 223 | +__global__ __aicore__ void GatherKernel(GM_ADDR team, GM_ADDR srcWin, GM_ADDR dstWin, Tiling tiling, uint32_t rank) | ||
| 224 | +{ | ||
| 225 | + KERNEL_TASK_TYPE_DEFAULT(KERNEL_TYPE_AIV_ONLY); | ||
| 226 | + AscendC::InitSocState(); | ||
| 227 | + Gather<false>(team, srcWin, dstWin, tiling, rank); | ||
| 228 | +} | ||
| 229 | + | ||
| 230 | +// 独立计算入口:远端数据已由通信 kernel、HCCL 或 Host 填好,因此不等待 ready/free。 | ||
| 231 | +template <typename T> | ||
| 232 | +__global__ __aicore__ void MatmulKernel(GM_ADDR a, GM_ADDR gathered, GM_ADDR b, GM_ADDR c, Tiling tiling, uint32_t rank) | ||
| 233 | +{ | ||
| 234 | + KERNEL_TASK_TYPE_DEFAULT(KERNEL_TYPE_AIC_ONLY); | ||
| 235 | + AscendC::InitSocState(); | ||
| 236 | + Matmul<T, false>(a, gathered, b, c, tiling, rank); | ||
| 237 | +} | ||
| 238 | + | ||
| 239 | +} // namespace ain_agmm | ||
| @@ -0,0 +1,369 @@ | |||
| 1 | +/* | ||
| 2 | + * Copyright (c) 2026 Huawei Technologies Co., Ltd. | ||
| 3 | + * Licensed under the CANN Open Software License Agreement Version 2.0. | ||
| 4 | + * See LICENSE in the root of the software repository. | ||
| 5 | + */ | ||
| 6 | +#include <algorithm> | ||
| 7 | +#include <chrono> | ||
| 8 | +#include <cstdio> | ||
| 9 | +#include <stdexcept> | ||
| 10 | +#include <thread> | ||
| 11 | +#include <vector> | ||
| 12 | +#include "acl/acl.h" | ||
| 13 | +#include "hccl/hccl.h" | ||
| 14 | +#include "hccl/hccl_comm.h" | ||
| 15 | +#include "hccl/hccl_team.h" | ||
| 16 | +#include "hcomm/hcomm_res.h" | ||
| 17 | +#include "platform/platform_ascendc.h" | ||
| 18 | +#include "utils/hccl_comm_init.h" | ||
| 19 | +#include "utils/process_manager.h" | ||
| 20 | +#include "utils/rank_sync.h" | ||
| 21 | +#include "options.h" | ||
| 22 | +#include "reference.h" | ||
| 23 | +#include "kernel.h" | ||
| 24 | + | ||
| 25 | +namespace ain_agmm { | ||
| 26 | + | ||
| 27 | +// 将 ACL/HCCL/HCOMM 返回码统一转为异常,并保留失败调用名称;宏实参只求值一次。 | ||
| 28 | +template <typename Status> | ||
| 29 | +void Check(Status status, const char* expression) | ||
| 30 | +{ | ||
| 31 | + if (static_cast<int64_t>(status) != 0) { | ||
| 32 | + throw std::runtime_error(std::string(expression) + " returned " + std::to_string(status)); | ||
| 33 | + } | ||
| 34 | +} | ||
| 35 | +#define CHECK_CALL(expr) Check((expr), #expr) | ||
| 36 | + | ||
| 37 | +// 检查 Host 协议和校验条件,失败后由进程管理器汇总为整个用例失败。 | ||
| 38 | +void Require(bool pass, const char* message) | ||
| 39 | +{ | ||
| 40 | + if (!pass) { throw std::runtime_error(message); } | ||
| 41 | +} | ||
| 42 | + | ||
| 43 | +// 单个 rank 独占的资源。A/gathered 注册为通信 window,B/C 是本地计算内存。 | ||
| 44 | +// 按已成功创建的成员清理,支持初始化途中失败,也允许显式 Close 后再析构。 | ||
| 45 | +struct Resources { | ||
| 46 | + uint32_t device; | ||
| 47 | + bool acl = false, deviceSet = false; | ||
| 48 | + aclrtStream stream = nullptr; | ||
| 49 | + aclrtEvent start = nullptr, end = nullptr; | ||
| 50 | + HcclComm comm = nullptr; | ||
| 51 | + HcommTeamHandle team = nullptr; | ||
| 52 | + HcclCommSymWindow srcWin = nullptr, dstWin = nullptr; | ||
| 53 | + void *a = nullptr, *gathered = nullptr, *b = nullptr, *c = nullptr, *trace = nullptr; | ||
| 54 | + bool failedCleanup = false; | ||
| 55 | + | ||
| 56 | + // 清理阶段不抛异常,记录失败并继续释放其他资源;正常返回路径据此给出失败退出码。 | ||
| 57 | + void Record(int64_t status) | ||
| 58 | + { | ||
| 59 | + if (status != 0) { | ||
| 60 | + failedCleanup = true; | ||
| 61 | + std::fprintf(stderr, "cleanup failed: %lld\n", static_cast<long long>(status)); | ||
| 62 | + } | ||
| 63 | + } | ||
| 64 | + // 先等待设备工作结束,再拆 team/window,最后释放内存、通信域和 ACL 环境。 | ||
| 65 | + void Close() | ||
| 66 | + { | ||
| 67 | + if (stream) { Record(aclrtSynchronizeStream(stream)); } | ||
| 68 | + if (team) { Record(HcclTeamDestroy(team)); team = nullptr; } | ||
| 69 | + if (dstWin) { Record(HcclCommSymWinDeregister(dstWin)); dstWin = nullptr; } | ||
| 70 | + if (srcWin) { Record(HcclCommSymWinDeregister(srcWin)); srcWin = nullptr; } | ||
| 71 | + if (a) { Record(HcommMemFree(a)); a = nullptr; } | ||
| 72 | + if (gathered) { Record(HcommMemFree(gathered)); gathered = nullptr; } | ||
| 73 | + if (b) { Record(aclrtFree(b)); b = nullptr; } | ||
| 74 | + if (c) { Record(aclrtFree(c)); c = nullptr; } | ||
| 75 | + if (trace) { Record(aclrtFree(trace)); trace = nullptr; } | ||
| 76 | + if (comm) { Record(HcclCommDestroy(comm)); comm = nullptr; } | ||
| 77 | + if (start) { Record(aclrtDestroyEvent(start)); start = nullptr; } | ||
| 78 | + if (end) { Record(aclrtDestroyEvent(end)); end = nullptr; } | ||
| 79 | + if (stream) { Record(aclrtDestroyStream(stream)); stream = nullptr; } | ||
| 80 | + if (deviceSet) { Record(aclrtResetDevice(device)); deviceSet = false; } | ||
| 81 | + if (acl) { Record(aclFinalize()); acl = false; } | ||
| 82 | + } | ||
| 83 | + ~Resources() { Close(); } | ||
| 84 | +}; | ||
| 85 | + | ||
| 86 | +// 绑定 rank 对应的 ACL device,建立 HCCL 域,再创建供 AIV/UB_CTP 使用的对称 window 和 team。 | ||
| 87 | +void Initialize(Resources& r, examples::RankSyncContext& sync, const Options& o, uint32_t rank) | ||
| 88 | +{ | ||
| 89 | + CHECK_CALL(aclInit(nullptr)); | ||
| 90 | + r.acl = true; | ||
| 91 | + uint32_t count = 0; | ||
| 92 | + CHECK_CALL(aclrtGetDeviceCount(&count)); | ||
| 93 | + for (auto id : o.devices) { Require(id < count, "requested ACL device ID is unavailable"); } | ||
| 94 | + CHECK_CALL(aclrtSetDevice(r.device)); | ||
| 95 | + r.deviceSet = true; | ||
| 96 | + CHECK_CALL(aclrtCreateStream(&r.stream)); | ||
| 97 | + CHECK_CALL(aclrtCreateEvent(&r.start)); | ||
| 98 | + CHECK_CALL(aclrtCreateEvent(&r.end)); | ||
| 99 | + CHECK_CALL(examples::InitCommByRootInfo(sync, &r.comm)); | ||
| 100 | + // A:[M,K],gathered:[P,M,K],B:[K,N],C:[P*M,N];FP16/BF16 均按 2 字节分配。 | ||
| 101 | + const uint64_t aBytes = static_cast<uint64_t>(o.m) * o.k * 2; | ||
| 102 | + CHECK_CALL(HcommMemAlloc(&r.a, aBytes)); | ||
| 103 | + CHECK_CALL(HcommMemAlloc(&r.gathered, aBytes * o.ranks)); | ||
| 104 | + CHECK_CALL(aclrtMalloc(&r.b, static_cast<uint64_t>(o.k) * o.n * 2, ACL_MEM_MALLOC_HUGE_FIRST)); | ||
| 105 | + CHECK_CALL(aclrtMalloc(&r.c, static_cast<uint64_t>(o.ranks) * o.m * o.n * 2, ACL_MEM_MALLOC_HUGE_FIRST)); | ||
| 106 | + // 各 rank 以相同顺序、相同大小注册 window,AIN 通过 window 和 peer 定位远端内存。 | ||
| 107 | + CHECK_CALL(HcclCommSymWinRegister(r.comm, r.a, aBytes, &r.srcWin, 1)); | ||
| 108 | + CHECK_CALL(HcclCommSymWinRegister(r.comm, r.gathered, aBytes * o.ranks, &r.dstWin, 1)); | ||
| 109 | + HcclTeamCreateDesc desc; | ||
| 110 | + CHECK_CALL(HcclTeamCreateDescInit(&desc)); | ||
| 111 | + std::vector<uint32_t> ids(o.ranks); | ||
| 112 | + for (uint32_t i = 0; i < o.ranks; ++i) { ids[i] = i; } | ||
| 113 | + desc.rankIds = ids.data(); | ||
| 114 | + desc.rankNum = o.ranks; | ||
| 115 | + desc.selfRankId = rank; | ||
| 116 | + desc.netLayer = 0; | ||
| 117 | + desc.requirement.barrierCount = 1; | ||
| 118 | + desc.engine = COMM_ENGINE_AIV; | ||
| 119 | + desc.protocol = COMM_PROTOCOL_UB_CTP; | ||
| 120 | + // 每个 peer 使用一条 channel;kernel 按 peer 分配唯一 AIV owner,避免多核竞争同一 channel。 | ||
| 121 | + desc.channelCnt = 1; | ||
| 122 | + CHECK_CALL(HcclTeamCreate(r.comm, &desc, &r.team)); | ||
| 123 | +} | ||
| 124 | + | ||
| 125 | +// 根据 rank/坐标/seed 生成可复现输入;各 rank 的 B 可以不同。 | ||
| 126 | +// 默认 K 周期为 16,便于大 shape 全量校验;random 模式逐 K 生成,仅适合较小测试。 | ||
| 127 | +void FillInputs(Resources& r, const Options& o, uint32_t rank, uint32_t seed, bool fillGathered) | ||
| 128 | +{ | ||
| 129 | + std::vector<uint16_t> a(static_cast<uint64_t>(o.m) * o.k); | ||
| 130 | + std::vector<uint16_t> b(static_cast<uint64_t>(o.k) * o.n); | ||
| 131 | + auto fillA = [&](uint32_t peer) { | ||
| 132 | + for (uint32_t m = 0; m < o.m; ++m) { | ||
| 133 | + for (uint32_t k = 0; k < o.k; ++k) { | ||
| 134 | + a[static_cast<uint64_t>(m) * o.k + k] = | ||
| 135 | + Encode(InputValue(peer, m, o.random ? k : k % 16, seed, false), o.bf16); | ||
| 136 | + } | ||
| 137 | + } | ||
| 138 | + }; | ||
| 139 | + fillA(rank); | ||
| 140 | + CHECK_CALL(aclrtMemcpy(r.a, a.size() * 2, a.data(), a.size() * 2, ACL_MEMCPY_HOST_TO_DEVICE)); | ||
| 141 | + for (uint32_t k = 0; k < o.k; ++k) { | ||
| 142 | + for (uint32_t n = 0; n < o.n; ++n) { | ||
| 143 | + b[static_cast<uint64_t>(k) * o.n + n] = | ||
| 144 | + Encode(InputValue(rank, o.random ? k : k % 16, n, seed, true), o.bf16); | ||
| 145 | + } | ||
| 146 | + } | ||
| 147 | + CHECK_CALL(aclrtMemcpy(r.b, b.size() * 2, b.data(), b.size() * 2, ACL_MEMCPY_HOST_TO_DEVICE)); | ||
| 148 | + // 独立 matmul 模式没有通信,由 Host 预填所有 peer 的 A,隔离计算侧问题。 | ||
| 149 | + if (fillGathered) { | ||
| 150 | + for (uint32_t peer = 0; peer < o.ranks; ++peer) { | ||
| 151 | + fillA(peer); | ||
| 152 | + auto dst = static_cast<uint8_t*>(r.gathered) + peer * a.size() * 2; | ||
| 153 | + CHECK_CALL(aclrtMemcpy(dst, a.size() * 2, a.data(), a.size() * 2, ACL_MEMCPY_HOST_TO_DEVICE)); | ||
| 154 | + } | ||
| 155 | + } | ||
| 156 | +} | ||
| 157 | + | ||
| 158 | +// 按模式向同一 stream 提交工作;serial/hccl 的通信排在计算前,融合模式只启动一个 MIX kernel。 | ||
| 159 | +template <typename T> | ||
| 160 | +void Launch(Resources& r, const Options& o, const Tiling& t, uint32_t rank, const std::string& mode, bool trace) | ||
| 161 | +{ | ||
| 162 | + auto a = static_cast<uint8_t*>(r.a); | ||
| 163 | + auto g = static_cast<uint8_t*>(r.gathered); | ||
| 164 | + auto b = static_cast<uint8_t*>(r.b); | ||
| 165 | + auto c = static_cast<uint8_t*>(r.c); | ||
| 166 | + auto team = reinterpret_cast<uint8_t*>(r.team); | ||
| 167 | + auto srcWin = reinterpret_cast<uint8_t*>(r.srcWin); | ||
| 168 | + auto dstWin = reinterpret_cast<uint8_t*>(r.dstWin); | ||
| 169 | + if (mode == "fused") { | ||
| 170 | + if (trace) { | ||
| 171 | + FusedKernel<T, true><<<t.cores, nullptr, r.stream>>>(team, srcWin, dstWin, a, g, b, c, t, rank, | ||
| 172 | + static_cast<uint8_t*>(r.trace)); | ||
| 173 | + } else { | ||
| 174 | + FusedKernel<T><<<t.cores, nullptr, r.stream>>>(team, srcWin, dstWin, a, g, b, c, t, rank, nullptr); | ||
| 175 | + } | ||
| 176 | + return; | ||
| 177 | + } | ||
| 178 | + if (mode == "serial" || mode == "comm") { | ||
| 179 | + GatherKernel<<<t.cores, nullptr, r.stream>>>(team, srcWin, dstWin, t, rank); | ||
| 180 | + } else if (mode == "hccl") { | ||
| 181 | + CHECK_CALL(HcclAllGather(r.a, r.gathered, static_cast<uint64_t>(o.m) * o.k, | ||
| 182 | + o.bf16 ? HCCL_DATA_TYPE_BFP16 : HCCL_DATA_TYPE_FP16, r.comm, r.stream)); | ||
| 183 | + } | ||
| 184 | + if (mode != "comm") { | ||
| 185 | + MatmulKernel<T><<<t.cores, nullptr, r.stream>>>(a, g, b, c, t, rank); | ||
| 186 | + } | ||
| 187 | +} | ||
| 188 | + | ||
| 189 | +// 通信结果按位比对,矩阵输出与 CPU golden 全元素比较;golden 先舍入到输出 dtype。 | ||
| 190 | +// comm 只验接收 A,matmul 只验 C,其余模式两者都验;非有限值直接失败。 | ||
| 191 | +void Verify(Resources& r, const Options& o, uint32_t rank, uint32_t seed, const std::string& mode) | ||
| 192 | +{ | ||
| 193 | + uint64_t failures = 0; | ||
| 194 | + if (mode != "matmul") { | ||
| 195 | + std::vector<uint16_t> a(static_cast<uint64_t>(o.m) * o.k); | ||
| 196 | + for (uint32_t peer = 0; peer < o.ranks; ++peer) { | ||
| 197 | + if (peer == rank) { continue; } // Get 不填自身槽,本地 matmul 直接读原始 A。 | ||
| 198 | + auto src = static_cast<uint8_t*>(r.gathered) + peer * a.size() * 2; | ||
| 199 | + CHECK_CALL(aclrtMemcpy(a.data(), a.size() * 2, src, a.size() * 2, ACL_MEMCPY_DEVICE_TO_HOST)); | ||
| 200 | + for (uint32_t m = 0; m < o.m; ++m) { | ||
| 201 | + for (uint32_t k = 0; k < o.k; ++k) { | ||
| 202 | + const auto expected = Encode(InputValue(peer, m, o.random ? k : k % 16, seed, false), o.bf16); | ||
| 203 | + if (a[static_cast<uint64_t>(m) * o.k + k] != expected) { ++failures; } | ||
| 204 | + } | ||
| 205 | + } | ||
| 206 | + } | ||
| 207 | + } | ||
| 208 | + if (mode != "comm") { | ||
| 209 | + std::vector<uint16_t> c(static_cast<uint64_t>(o.ranks) * o.m * o.n); | ||
| 210 | + CHECK_CALL(aclrtMemcpy(c.data(), c.size() * 2, r.c, c.size() * 2, ACL_MEMCPY_DEVICE_TO_HOST)); | ||
| 211 | + for (uint32_t peer = 0; peer < o.ranks; ++peer) { | ||
| 212 | + for (uint32_t m = 0; m < o.m; ++m) { | ||
| 213 | + for (uint32_t n = 0; n < o.n; ++n) { | ||
| 214 | + const uint64_t offset = (static_cast<uint64_t>(peer) * o.m + m) * o.n + n; | ||
| 215 | + const float expected = Decode(Encode(Golden(peer, rank, m, n, o.k, seed, o.random), o.bf16), o.bf16); | ||
| 216 | + const float actual = Decode(c[offset], o.bf16); | ||
| 217 | + const float tolerance = (o.bf16 ? 0.008F : 0.001F) * std::fabs(expected) + 0.001F; | ||
| 218 | + if (!std::isfinite(actual) || !std::isfinite(expected) || std::fabs(actual - expected) > tolerance) { | ||
| 219 | + if (failures < 4) { | ||
| 220 | + std::fprintf(stderr, "rank=%u source=%u row=%u col=%u got=%g expected=%g\n", | ||
| 221 | + rank, peer, m, n, actual, expected); | ||
| 222 | + } | ||
| 223 | + ++failures; | ||
| 224 | + } | ||
| 225 | + } | ||
| 226 | + } | ||
| 227 | + } | ||
| 228 | + } | ||
| 229 | + std::printf("CHECK rank=%u mode=%s seed=%u errors=%llu\n", rank, mode.c_str(), seed, | ||
| 230 | + static_cast<unsigned long long>(failures)); | ||
| 231 | + Require(failures == 0, "numerical or gathered-data verification failed"); | ||
| 232 | +} | ||
| 233 | + | ||
| 234 | +// 执行一次并返回所有 rank 中最大的 device event 耗时(us);输入准备/CPU 校验在调用外。 | ||
| 235 | +float RunOnce(Resources& r, examples::RankSyncContext& sync, const Options& o, | ||
| 236 | + const Tiling& t, uint32_t rank, const std::string& mode, bool trace = false, float* launchSkew = nullptr) | ||
| 237 | +{ | ||
| 238 | + // 各 rank 先完成输入准备再启动;结束后的 Allgather 确认所有 rank 已完成读取。 | ||
| 239 | + // 这两个同步共同保护远端 A 的生命周期,防止快 rank 提前更新输入或释放 window。 | ||
| 240 | + Require(sync.Barrier(examples::kTagUserBase), "pre-launch barrier failed"); | ||
| 241 | + if (rank == o.delayRank) { std::this_thread::sleep_for(std::chrono::milliseconds(o.delayMs)); } | ||
| 242 | + struct Sample { int64_t hostNs; float ms; }; | ||
| 243 | + Sample sample{}; | ||
| 244 | + sample.hostNs = std::chrono::duration_cast<std::chrono::nanoseconds>( | ||
| 245 | + std::chrono::steady_clock::now().time_since_epoch()).count(); | ||
| 246 | + // event 包住通信与计算提交;TCP 同步调用本身在外,但 rank 启动偏斜仍可能进入 collective 等待。 | ||
| 247 | + CHECK_CALL(aclrtRecordEvent(r.start, r.stream)); | ||
| 248 | + if (o.bf16) { Launch<bfloat16_t>(r, o, t, rank, mode, trace); } | ||
| 249 | + else { Launch<half>(r, o, t, rank, mode, trace); } | ||
| 250 | + CHECK_CALL(aclrtRecordEvent(r.end, r.stream)); | ||
| 251 | + CHECK_CALL(aclrtSynchronizeStream(r.stream)); | ||
| 252 | + CHECK_CALL(aclrtEventElapsedTime(&sample.ms, r.start, r.end)); | ||
| 253 | + Require(std::isfinite(sample.ms) && sample.ms > 0, "invalid event duration"); | ||
| 254 | + std::vector<Sample> all(o.ranks); | ||
| 255 | + Require(sync.Allgather(examples::kTagUserBase + 1, &sample, sizeof(sample), all.data()), "timing allgather failed"); | ||
| 256 | + float maxMs = 0; | ||
| 257 | + int64_t first = sample.hostNs, last = sample.hostNs; | ||
| 258 | + for (const auto& s : all) { | ||
| 259 | + maxMs = std::max(maxMs, s.ms); | ||
| 260 | + first = std::min(first, s.hostNs); | ||
| 261 | + last = std::max(last, s.hostNs); | ||
| 262 | + } | ||
| 263 | + // 本示例所有 rank 在同一 Host 上 fork,共享单调时钟;这里只量 Host 提交偏斜, | ||
| 264 | + // 不代表设备 kernel 启动偏斜,也不从 event 耗时中直接扣除。 | ||
| 265 | + if (launchSkew) { *launchSkew = static_cast<float>(last - first) / 1000; } | ||
| 266 | + return maxMs * 1000; | ||
| 267 | +} | ||
| 268 | + | ||
| 269 | +// 额外运行一次采样版本,拷回并打印 AIV0/AIC0 的阶段区间;该次耗时不进入稳态统计。 | ||
| 270 | +void TraceOnce(Resources& r, examples::RankSyncContext& sync, const Options& o, const Tiling& t, uint32_t rank) | ||
| 271 | +{ | ||
| 272 | + // 布局与 kernel 的 traceBase 一致:Get 区、对齐填充、本地计算起止、各轮远端计算起止。 | ||
| 273 | + const uint64_t base = CeilDiv(static_cast<uint64_t>(t.rounds) * 2, 32) * 32; | ||
| 274 | + std::vector<uint64_t> ticks(base + 2 + static_cast<uint64_t>(t.rounds) * 2); | ||
| 275 | + CHECK_CALL(aclrtMalloc(&r.trace, ticks.size() * sizeof(uint64_t), ACL_MEM_MALLOC_HUGE_FIRST)); | ||
| 276 | + CHECK_CALL(aclrtMemset(r.trace, ticks.size() * sizeof(uint64_t), 0, ticks.size() * sizeof(uint64_t))); | ||
| 277 | + (void)RunOnce(r, sync, o, t, rank, "fused", true); | ||
| 278 | + CHECK_CALL(aclrtMemcpy(ticks.data(), ticks.size() * sizeof(uint64_t), r.trace, | ||
| 279 | + ticks.size() * sizeof(uint64_t), ACL_MEMCPY_DEVICE_TO_HOST)); | ||
| 280 | + auto print = [&](const char* stage, uint32_t round, uint64_t index) { | ||
| 281 | + Require(ticks[index] != 0 && ticks[index + 1] >= ticks[index], "invalid diagnostic timestamps"); | ||
| 282 | + std::printf("TRACE rank=%u core=0 stage=%s round=%u begin_cycle=%llu end_cycle=%llu\n", rank, stage, round, | ||
| 283 | + static_cast<unsigned long long>(ticks[index]), static_cast<unsigned long long>(ticks[index + 1])); | ||
| 284 | + }; | ||
| 285 | + print("local_matmul", 0, base); | ||
| 286 | + for (uint32_t j = 0; j < t.rounds; ++j) { | ||
| 287 | + print("get", j, static_cast<uint64_t>(j) * 2); | ||
| 288 | + print("remote_matmul", j, base + 2 + static_cast<uint64_t>(j) * 2); | ||
| 289 | + } | ||
| 290 | + Verify(r, o, rank, o.verifyIterations, "fused"); | ||
| 291 | + CHECK_CALL(aclrtFree(r.trace)); | ||
| 292 | + r.trace = nullptr; | ||
| 293 | +} | ||
| 294 | + | ||
| 295 | +// 单个子进程的完整流程:控制连接 -> 建域/分配 -> 校验 -> 可选采样 -> 预热/测量 -> 释放。 | ||
| 296 | +int RunRank(uint32_t rank, Options o) | ||
| 297 | +{ | ||
| 298 | + // rank 是通信域内编号,devices[rank] 是 ACL device ID;两者不能混用。 | ||
| 299 | + Resources r{o.devices[rank]}; | ||
| 300 | + auto sync = rank == 0 ? examples::RankSyncContext::Listen(o.endpoint, o.ranks) : | ||
| 301 | + examples::RankSyncContext::Connect(rank, o.ranks, o.endpoint); | ||
| 302 | + Require(sync.Ok(), "control channel initialization failed"); | ||
| 303 | + Initialize(r, sync, o, rank); | ||
| 304 | + auto* platform = platform_ascendc::PlatformAscendCManager::GetInstance(); | ||
| 305 | + Require(platform != nullptr, "platform query failed"); | ||
| 306 | + const uint32_t availableCores = platform->GetCoreNumAic(); | ||
| 307 | + Require(o.cores == 0 || o.cores <= availableCores, "cores exceeds device AIC count"); | ||
| 308 | + const Tiling t = MakeTiling(o.m, o.n, o.k, o.ranks, o.cores ? o.cores : availableCores, o.tileM); | ||
| 309 | + std::printf("CONFIG rank=%u device=%u M=%u K=%u N=%u dtype=%s cores=%u tileM=%u rounds=%u\n", | ||
| 310 | + rank, r.device, t.m, t.k, t.n, o.bf16 ? "bf16" : "fp16", t.cores, t.tileM, t.rounds); | ||
| 311 | + const std::vector<std::string> modes = o.mode == "both" ? std::vector<std::string>{"serial", "fused"} : | ||
| 312 | + std::vector<std::string>{o.mode}; | ||
| 313 | + for (const auto& mode : modes) { | ||
| 314 | + // 保持同一套 window/team,改变 seed 后再次调用,不重建 channel 或额外重置 flag。 | ||
| 315 | + // 用不同输入检查陈旧数据/通知;输出先填为非有限值,便于发现漏写。 | ||
| 316 | + for (uint32_t v = 0; v < o.verifyIterations; ++v) { | ||
| 317 | + FillInputs(r, o, rank, v + 1, mode == "matmul"); | ||
| 318 | + const uint64_t cBytes = static_cast<uint64_t>(o.ranks) * o.m * o.n * 2; | ||
| 319 | + CHECK_CALL(aclrtMemset(r.c, cBytes, 0xff, cBytes)); | ||
| 320 | + (void)RunOnce(r, sync, o, t, rank, mode); | ||
| 321 | + Verify(r, o, rank, v + 1, mode); | ||
| 322 | + } | ||
| 323 | + if (o.trace && mode == "fused") { TraceOnce(r, sync, o, t, rank); } | ||
| 324 | + for (uint32_t i = 0; i < o.warmup; ++i) { (void)RunOnce(r, sync, o, t, rank, mode); } | ||
| 325 | + std::vector<float> times; | ||
| 326 | + std::vector<float> skews; | ||
| 327 | + for (uint32_t i = 0; i < o.iterations; ++i) { | ||
| 328 | + float skew = 0; | ||
| 329 | + times.push_back(RunOnce(r, sync, o, t, rank, mode, false, &skew)); | ||
| 330 | + skews.push_back(skew); | ||
| 331 | + } | ||
| 332 | + Verify(r, o, rank, o.verifyIterations, mode); | ||
| 333 | + // 先逐轮取最大 rank 耗时,再统计 median/P95,不能先算各 rank 均值后取最大。 | ||
| 334 | + if (rank == 0) { | ||
| 335 | + std::sort(times.begin(), times.end()); | ||
| 336 | + std::sort(skews.begin(), skews.end()); | ||
| 337 | + const float median = (times[(times.size() - 1) / 2] + times[times.size() / 2]) / 2; | ||
| 338 | + const size_t p95 = static_cast<size_t>(std::ceil(times.size() * 0.95)) - 1; | ||
| 339 | + std::printf("PERF mode=%s ranks=%u M=%u K=%u N=%u dtype=%s tileM=%u samples=%u " | ||
| 340 | + "max_rank_median_us=%.3f max_rank_p95_us=%.3f host_submit_skew_p95_us=%.3f\n", | ||
| 341 | + mode.c_str(), o.ranks, t.m, t.k, t.n, | ||
| 342 | + o.bf16 ? "bf16" : "fp16", t.tileM, o.iterations, median, times[p95], skews[p95]); | ||
| 343 | + } | ||
| 344 | + Require(sync.Barrier(examples::kTagUserBase + 2), "post-verification barrier failed"); | ||
| 345 | + } | ||
| 346 | + r.Close(); | ||
| 347 | + return r.failedCleanup ? 1 : 0; | ||
| 348 | +} | ||
| 349 | +} // namespace ain_agmm | ||
| 350 | + | ||
| 351 | +// 父进程解析参数并 fork 各 rank,统一管理超时和退出状态;任一 rank 失败则整组失败。 | ||
| 352 | +int main(int argc, char** argv) | ||
| 353 | +{ | ||
| 354 | + try { | ||
| 355 | + const auto options = ain_agmm::ParseOptions(argc, argv); | ||
| 356 | + std::vector<examples::ArgsOf<decltype(ain_agmm::RunRank)>> args; | ||
| 357 | + for (uint32_t rank = 0; rank < options.ranks; ++rank) { args.emplace_back(rank, options); } | ||
| 358 | + examples::ProcessGroupOptions groupOptions; | ||
| 359 | + groupOptions.timeoutSeconds = options.timeout; | ||
| 360 | + examples::ProcessGroup group(groupOptions); | ||
| 361 | + if (!group.Launch(options.ranks, ain_agmm::RunRank, std::move(args))) { return 1; } | ||
| 362 | + const bool pass = group.WaitAll(); | ||
| 363 | + std::printf("RESULT | Example=ain_all_gather_matmul | Status=%s\n", pass ? "PASS" : "FAIL"); | ||
| 364 | + return pass ? 0 : 1; | ||
| 365 | + } catch (const std::exception& e) { | ||
| 366 | + std::fprintf(stderr, "%s\n", e.what()); | ||
| 367 | + return 1; | ||
| 368 | + } | ||
| 369 | +} | ||
| @@ -0,0 +1,92 @@ | |||
| 1 | +/* | ||
| 2 | + * Copyright (c) 2026 Huawei Technologies Co., Ltd. | ||
| 3 | + * Licensed under the CANN Open Software License Agreement Version 2.0. | ||
| 4 | + * See LICENSE in the root of the software repository. | ||
| 5 | + */ | ||
| 6 | + | ||
| 7 | + | ||
| 8 | + | ||
| 9 | + | ||
| 10 | + | ||
| 11 | + | ||
| 12 | + | ||
| 13 | +namespace ain_agmm { | ||
| 14 | +// 示例的 Host 配置:tileM/cores 为 0 表示自动选择,devices 保存 rank 到 ACL device 的映射。 | ||
| 15 | +struct Options { | ||
| 16 | + std::string endpoint; | ||
| 17 | + std::string mode = "both"; | ||
| 18 | + uint32_t ranks = 0, m = 0, n = 0, k = 0; | ||
| 19 | + uint32_t tileM = 0, cores = 0, warmup = 10, iterations = 100, verifyIterations = 2; | ||
| 20 | + uint32_t timeout = 300; | ||
| 21 | + uint32_t delayRank = UINT32_MAX, delayMs = 0; | ||
| 22 | + bool bf16 = false, random = false, trace = false; | ||
| 23 | + std::vector<uint32_t> devices; | ||
| 24 | +}; | ||
| 25 | + | ||
| 26 | +// 统一解析无符号参数并检查最小值,避免负数或超范围值被静默截断。 | ||
| 27 | +inline uint32_t Number(const char* text, uint32_t minimum = 0) | ||
| 28 | +{ | ||
| 29 | + uint32_t result; | ||
| 30 | + if (!examples::ParseUint32(text, result, minimum)) { | ||
| 31 | + throw std::invalid_argument(std::string("invalid integer: ") + text); | ||
| 32 | + } | ||
| 33 | + return result; | ||
| 34 | +} | ||
| 35 | + | ||
| 36 | +// 解析 endpoint ranks M K N 及可选项,在 fork/设备分配前拒绝非法配置。 | ||
| 37 | +inline Options ParseOptions(int argc, char** argv) | ||
| 38 | +{ | ||
| 39 | + if (argc < 6) { | ||
| 40 | + throw std::invalid_argument("Usage: ain_all_gather_matmul tcp://ip:port ranks M K N " | ||
| 41 | + "[--dtype fp16|bf16] [--mode both|fused|serial|hccl|comm|matmul] " | ||
| 42 | + "[--devices 1,2] [--tile-m M] [--cores C] [--warmup 10] [--iters 100] " | ||
| 43 | + "[--verify-iters 2] [--random] [--trace] [--timeout 300] [--delay-rank R --delay-ms MS]"); | ||
| 44 | + } | ||
| 45 | + Options o; | ||
| 46 | + o.endpoint = argv[1]; | ||
| 47 | + o.ranks = Number(argv[2], 2); | ||
| 48 | + o.m = Number(argv[3], 1); | ||
| 49 | + o.k = Number(argv[4], 1); | ||
| 50 | + o.n = Number(argv[5], 1); | ||
| 51 | + for (int i = 6; i < argc; ++i) { | ||
| 52 | + const std::string key = argv[i]; | ||
| 53 | + if (key == "--random") { o.random = true; continue; } | ||
| 54 | + if (key == "--trace") { o.trace = true; continue; } | ||
| 55 | + if (++i == argc) { throw std::invalid_argument("missing value for " + key); } | ||
| 56 | + const std::string value = argv[i]; | ||
| 57 | + if (key == "--dtype") { | ||
| 58 | + if (value != "fp16" && value != "bf16") { throw std::invalid_argument("dtype must be fp16 or bf16"); } | ||
| 59 | + o.bf16 = value == "bf16"; | ||
| 60 | + } else if (key == "--mode") { o.mode = value; } | ||
| 61 | + else if (key == "--tile-m") { o.tileM = Number(argv[i], 1); } | ||
| 62 | + else if (key == "--cores") { o.cores = Number(argv[i], 1); } | ||
| 63 | + else if (key == "--warmup") { o.warmup = Number(argv[i]); } | ||
| 64 | + else if (key == "--iters") { o.iterations = Number(argv[i], 1); } | ||
| 65 | + else if (key == "--verify-iters") { o.verifyIterations = Number(argv[i], 1); } | ||
| 66 | + else if (key == "--timeout") { o.timeout = Number(argv[i], 1); } | ||
| 67 | + else if (key == "--delay-rank") { o.delayRank = Number(argv[i]); } | ||
| 68 | + else if (key == "--delay-ms") { o.delayMs = Number(argv[i]); } | ||
| 69 | + else if (key == "--devices") { | ||
| 70 | + std::istringstream input(value); | ||
| 71 | + std::string part; | ||
| 72 | + if (value.empty() || value.back() == ',') { throw std::invalid_argument("invalid devices list"); } | ||
| 73 | + while (std::getline(input, part, ',')) { o.devices.push_back(Number(part.c_str())); } | ||
| 74 | + } else { throw std::invalid_argument("unknown option " + key); } | ||
| 75 | + } | ||
| 76 | + const std::vector<std::string> modes{"both", "fused", "serial", "hccl", "comm", "matmul"}; | ||
| 77 | + if (std::find(modes.begin(), modes.end(), o.mode) == modes.end()) { throw std::invalid_argument("invalid mode"); } | ||
| 78 | + // 尚未查询硬件核数,先用占位核数检查 shape;RunRank 会按真实核数重新生成 tiling。 | ||
| 79 | + (void)MakeTiling(o.m, o.n, o.k, o.ranks, o.cores ? o.cores : 1, o.tileM); | ||
| 80 | + if (o.devices.empty()) { | ||
| 81 | + for (uint32_t i = 0; i < o.ranks; ++i) { o.devices.push_back(i); } | ||
| 82 | + } | ||
| 83 | + // 只排序副本检查重复,不改变用户指定的 rank->device 顺序;互联可达性需由运行者选择。 | ||
| 84 | + auto sorted = o.devices; | ||
| 85 | + std::sort(sorted.begin(), sorted.end()); | ||
| 86 | + if (sorted.size() != o.ranks || std::adjacent_find(sorted.begin(), sorted.end()) != sorted.end()) { | ||
| 87 | + throw std::invalid_argument("devices must contain one unique ACL device ID per rank"); | ||
| 88 | + } | ||
| 89 | + if (o.delayMs != 0 && o.delayRank >= o.ranks) { throw std::invalid_argument("delay-rank must be in team"); } | ||
| 90 | + return o; | ||
| 91 | +} | ||
| 92 | +} // namespace ain_agmm | ||
| @@ -0,0 +1,92 @@ | |||
| 1 | +/* | ||
| 2 | + * Copyright (c) 2026 Huawei Technologies Co., Ltd. | ||
| 3 | + * Licensed under the CANN Open Software License Agreement Version 2.0. | ||
| 4 | + * See LICENSE in the root of the software repository. | ||
| 5 | + */ | ||
| 6 | + | ||
| 7 | + | ||
| 8 | + | ||
| 9 | + | ||
| 10 | + | ||
| 11 | + | ||
| 12 | +namespace ain_agmm { | ||
| 13 | + | ||
| 14 | +// 用 rank、行列坐标和 seed 的整数哈希产生可复现输入;weight 区分 A/B 数据。 | ||
| 15 | +// 数值为 [-4,4]/16,在 FP16/BF16 中均可精确表示,便于区分搬运错误和输出舍入。 | ||
| 16 | +inline float InputValue(uint32_t rank, uint32_t row, uint32_t col, uint32_t seed, bool weight) | ||
| 17 | +{ | ||
| 18 | + uint32_t x = rank * 0x9e3779b9U + row * 0x85ebca6bU + col * 0xc2b2ae35U + seed * 97U; | ||
| 19 | + x ^= weight ? 0xa511e9b3U : 0x63d83595U; | ||
| 20 | + x ^= x >> 16; | ||
| 21 | + x *= 0x7feb352dU; | ||
| 22 | + x ^= x >> 15; | ||
| 23 | + return static_cast<float>(static_cast<int32_t>(x % 9U) - 4) / 16.0F; | ||
| 24 | +} | ||
| 25 | + | ||
| 26 | +// Host 侧将 float 转成 FP16/BF16 的原始 16 位存储,不依赖 Host 是否支持半精度类型。 | ||
| 27 | +// 使用最近偶数舍入,并单独处理 NaN、无穷大及 FP16 次正规数。 | ||
| 28 | +inline uint16_t Encode(float value, bool bf16) | ||
| 29 | +{ | ||
| 30 | + uint32_t u = 0; | ||
| 31 | + std::memcpy(&u, &value, sizeof(u)); | ||
| 32 | + if (bf16) { | ||
| 33 | + if ((u & 0x7fffffffU) > 0x7f800000U) { | ||
| 34 | + return static_cast<uint16_t>((u >> 16) | 0x40U); | ||
| 35 | + } | ||
| 36 | + // 保留高 16 位:低位加上舍入偏置,恰好一半时由保留位的奇偶决定进位。 | ||
| 37 | + return static_cast<uint16_t>((u + 0x7fffU + ((u >> 16) & 1U)) >> 16); | ||
| 38 | + } | ||
| 39 | + const uint32_t sign = (u >> 16) & 0x8000U; | ||
| 40 | + const uint32_t exp = (u >> 23) & 0xffU; | ||
| 41 | + const uint32_t mantissa = u & 0x7fffffU; | ||
| 42 | + if (exp == 255) { | ||
| 43 | + return static_cast<uint16_t>(sign | 0x7c00U | (mantissa ? 0x200U : 0U)); | ||
| 44 | + } | ||
| 45 | + int32_t halfExp = static_cast<int32_t>(exp) - 127 + 15; | ||
| 46 | + if (halfExp >= 31) { | ||
| 47 | + return static_cast<uint16_t>(sign | 0x7c00U); | ||
| 48 | + } | ||
| 49 | + if (halfExp < -10) { | ||
| 50 | + return static_cast<uint16_t>(sign); | ||
| 51 | + } | ||
| 52 | + // 次正规数没有隐含的首位 1,需连同该位一起右移,再按最近偶数舍入。 | ||
| 53 | + if (halfExp <= 0) { | ||
| 54 | + const uint32_t shift = static_cast<uint32_t>(14 - halfExp); | ||
| 55 | + const uint32_t m = mantissa | 0x800000U; | ||
| 56 | + return static_cast<uint16_t>(sign | ((m + ((1U << (shift - 1)) - 1) + ((m >> shift) & 1)) >> shift)); | ||
| 57 | + } | ||
| 58 | + const uint32_t rounded = mantissa + 0xfffU + ((mantissa >> 13) & 1U); | ||
| 59 | + return static_cast<uint16_t>(sign | ((static_cast<uint32_t>(halfExp) << 10) + (rounded >> 13))); | ||
| 60 | +} | ||
| 61 | + | ||
| 62 | +// 将设备输出的 16 位存储还原成 float,供 CPU 比较及错误打印使用。 | ||
| 63 | +inline float Decode(uint16_t value, bool bf16) | ||
| 64 | +{ | ||
| 65 | + if (bf16) { | ||
| 66 | + uint32_t u = static_cast<uint32_t>(value) << 16; | ||
| 67 | + float result; | ||
| 68 | + std::memcpy(&result, &u, sizeof(result)); | ||
| 69 | + return result; | ||
| 70 | + } | ||
| 71 | + const uint32_t exp = (value >> 10) & 31U; | ||
| 72 | + const uint32_t frac = value & 1023U; | ||
| 73 | + float result = exp == 0 ? std::ldexp(static_cast<float>(frac), -24) : | ||
| 74 | + exp == 31 ? (frac ? NAN : INFINITY) : | ||
| 75 | + std::ldexp(static_cast<float>(1024 + frac), static_cast<int>(exp) - 25); | ||
| 76 | + return value & 0x8000U ? -result : result; | ||
| 77 | +} | ||
| 78 | + | ||
| 79 | +// 计算单个 C 元素:A 来自 sourceRank,B 来自当前输出所属的 weightRank。 | ||
| 80 | +// 周期输入只累加 16 项再乘 K/16(调用前已校验 K 为 16 的倍数);random 模式累加全部 K。 | ||
| 81 | +inline float Golden(uint32_t sourceRank, uint32_t weightRank, uint32_t row, uint32_t col, | ||
| 82 | + uint32_t k, uint32_t seed, bool random) | ||
| 83 | +{ | ||
| 84 | + float sum = 0; | ||
| 85 | + const uint32_t count = random ? k : 16; | ||
| 86 | + for (uint32_t i = 0; i < count; ++i) { | ||
| 87 | + sum += InputValue(sourceRank, row, i, seed, false) * InputValue(weightRank, i, col, seed, true); | ||
| 88 | + } | ||
| 89 | + return random ? sum : sum * static_cast<float>(k / 16); | ||
| 90 | +} | ||
| 91 | + | ||
| 92 | +} // namespace ain_agmm | ||
| @@ -0,0 +1,89 @@ | |||
| 1 | +/* | ||
| 2 | + * Copyright (c) 2026 Huawei Technologies Co., Ltd. | ||
| 3 | + * Licensed under the CANN Open Software License Agreement Version 2.0. | ||
| 4 | + * See LICENSE in the root of the software repository. | ||
| 5 | + */ | ||
| 6 | + | ||
| 7 | + | ||
| 8 | + | ||
| 9 | + | ||
| 10 | + | ||
| 11 | + | ||
| 12 | + | ||
| 13 | +using namespace ain_agmm; | ||
| 14 | +void Expect(bool pass) { if (!pass) { throw std::runtime_error("host check failed"); } } | ||
| 15 | +template <typename F> void Reject(F f) | ||
| 16 | +{ | ||
| 17 | + bool rejected = false; | ||
| 18 | + try { f(); } catch (const std::invalid_argument&) { rejected = true; } | ||
| 19 | + Expect(rejected); | ||
| 20 | +} | ||
| 21 | + | ||
| 22 | +Options Parse(std::vector<std::string> args) | ||
| 23 | +{ | ||
| 24 | + std::vector<char*> argv; | ||
| 25 | + for (auto& arg : args) { argv.push_back(arg.data()); } | ||
| 26 | + return ParseOptions(static_cast<int>(argv.size()), argv.data()); | ||
| 27 | +} | ||
| 28 | + | ||
| 29 | +int main() | ||
| 30 | +{ | ||
| 31 | + Reject([] { MakeTiling(0, 16, 16, 2, 1, 0); }); | ||
| 32 | + Reject([] { MakeTiling(128, 17, 16, 2, 1, 0); }); | ||
| 33 | + Reject([] { MakeTiling(128, 16, 17, 2, 1, 0); }); | ||
| 34 | + Reject([] { MakeTiling(512, 256, 256, 2, 8, 257); }); | ||
| 35 | + Reject([] { MakeTiling(UINT32_MAX, 0xfffffff0U, 0xfffffff0U, 64, 48, 0); }); | ||
| 36 | + Reject([] { MakeTiling(512, 16, 1048592, 2, 1, 256); }); | ||
| 37 | + Reject([] { CheckedProduct(UINT64_MAX, 2); }); | ||
| 38 | + Reject([] { Parse({"demo", "tcp://127.0.0.1:1", "2", "64", "16", "16", "--devices", "1,1"}); }); | ||
| 39 | + Reject([] { Parse({"demo", "tcp://127.0.0.1:1", "2", "64", "16", "16", "--iters", "0"}); }); | ||
| 40 | + Reject([] { Parse({"demo", "tcp://127.0.0.1:1", "2", "64", "16", "16", "--dtype", "fp32"}); }); | ||
| 41 | + const auto o = Parse({"demo", "tcp://127.0.0.1:1", "2", "64", "16", "16", "--devices", "1,3", "--random"}); | ||
| 42 | + Expect(o.devices == std::vector<uint32_t>({1, 3}) && o.random); | ||
| 43 | + | ||
| 44 | + // Communication address coverage: odd tails, more than 16 rounds, and >4 GiB | ||
| 45 | + // offsets. Every row must occur once and each operation fits its window. | ||
| 46 | + for (uint32_t m : {1U, 17U, 255U, 256U, 257U, 4097U, 65537U}) { | ||
| 47 | + for (uint32_t ranks : {2U, 4U, 8U}) { | ||
| 48 | + const auto t = MakeTiling(m, 272, 16384, ranks, 8, m > 256 ? 256 : m); | ||
| 49 | + uint64_t next = 0; | ||
| 50 | + for (uint32_t j = 0; j < t.rounds; ++j) { | ||
| 51 | + const uint64_t row = static_cast<uint64_t>(j) * t.tileM; | ||
| 52 | + const uint64_t rows = std::min<uint64_t>(t.tileM, t.m - row); | ||
| 53 | + Expect(row == next && rows > 0 && rows * t.k * 2 <= MAX_GET_BYTES); | ||
| 54 | + next += rows; | ||
| 55 | + for (uint32_t p = 0; p < ranks; ++p) { | ||
| 56 | + const uint64_t end = (static_cast<uint64_t>(p) * m + row + rows) * t.k * 2; | ||
| 57 | + Expect(end <= static_cast<uint64_t>(ranks) * m * t.k * 2); | ||
| 58 | + } | ||
| 59 | + } | ||
| 60 | + Expect(next == m); | ||
| 61 | + } | ||
| 62 | + } | ||
| 63 | + for (bool bf16 : {false, true}) { | ||
| 64 | + for (int i = -128; i <= 128; ++i) { | ||
| 65 | + const float v = i / 16.0F; | ||
| 66 | + Expect(Decode(Encode(v, bf16), bf16) == v); | ||
| 67 | + } | ||
| 68 | + Expect(std::isnan(Decode(Encode(NAN, bf16), bf16))); | ||
| 69 | + Expect(std::isinf(Decode(Encode(INFINITY, bf16), bf16))); | ||
| 70 | + } | ||
| 71 | + Expect(Encode(1.0F, false) == 0x3c00 && Encode(1.0F, true) == 0x3f80); | ||
| 72 | + Expect(Encode(std::ldexp(1.0F, -24), false) == 1); | ||
| 73 | + Expect(Encode(1.00048828125F, false) == 0x3c00); // Tie to even. | ||
| 74 | + // Independent dense reference must agree with the fast periodic golden. | ||
| 75 | + for (uint32_t seed : {1U, 2U}) { | ||
| 76 | + for (uint32_t p = 0; p < 4; ++p) { | ||
| 77 | + for (uint32_t m = 0; m < 17; ++m) { | ||
| 78 | + for (uint32_t n = 0; n < 19; ++n) { | ||
| 79 | + float sum = 0; | ||
| 80 | + for (uint32_t k = 0; k < 80; ++k) { | ||
| 81 | + sum += InputValue(p, m, k % 16, seed, false) * InputValue(3, k % 16, n, seed, true); | ||
| 82 | + } | ||
| 83 | + Expect(sum == Golden(p, 3, m, n, 80, seed, false)); | ||
| 84 | + } | ||
| 85 | + } | ||
| 86 | + } | ||
| 87 | + } | ||
| 88 | + std::cout << "PASS: host tiling, argument validation and reference checks (not device execution)\n"; | ||
| 89 | +} | ||
| @@ -0,0 +1,64 @@ | |||
| 1 | +#!/usr/bin/env python3 | ||
| 2 | +# Copyright (c) 2026 Huawei Technologies Co., Ltd. | ||
| 3 | +# Licensed under the CANN Open Software License Agreement Version 2.0. | ||
| 4 | +"""Run bounded, full-output device checks; preserve each command and its log.""" | ||
| 5 | +import argparse | ||
| 6 | +import json | ||
| 7 | +from pathlib import Path | ||
| 8 | +import subprocess | ||
| 9 | +import sys | ||
| 10 | + | ||
| 11 | + | ||
| 12 | +def main(): | ||
| 13 | + parser = argparse.ArgumentParser(description=__doc__) | ||
| 14 | + parser.add_argument("--binary", type=Path, required=True) | ||
| 15 | + parser.add_argument("--devices", required=True, help="at least two ACL device IDs in the same interconnect group") | ||
| 16 | + parser.add_argument("--output-dir", type=Path, required=True) | ||
| 17 | + parser.add_argument("--port", type=int, default=29650) | ||
| 18 | + args = parser.parse_args() | ||
| 19 | + devices = [int(x) for x in args.devices.split(",")] | ||
| 20 | + if len(devices) < 2 or len(set(devices)) != len(devices) or min(devices) < 0: | ||
| 21 | + parser.error("provide at least two unique nonnegative devices") | ||
| 22 | + if not args.binary.is_file(): | ||
| 23 | + parser.error("binary does not exist") | ||
| 24 | + # Refuse to overwrite a previous verification report. | ||
| 25 | + args.output_dir.mkdir(parents=True, exist_ok=False) | ||
| 26 | + cases = [ | ||
| 27 | + ("matmul_random", 2, 17, 80, 272, ["--mode", "matmul", "--random"]), | ||
| 28 | + ("get_random", 2, 257, 80, 272, ["--mode", "comm", "--tile-m", "256", "--random"]), | ||
| 29 | + ("fused_fp16_tail", 2, 257, 80, 272, ["--tile-m", "256", "--random"]), | ||
| 30 | + ("fused_bf16_tail", 2, 513, 144, 272, ["--dtype", "bf16", "--tile-m", "256", "--random"]), | ||
| 31 | + ("many_rounds", 2, 4353, 80, 272, ["--tile-m", "256"]), | ||
| 32 | + ("one_core", 2, 257, 80, 272, ["--cores", "1", "--tile-m", "256"]), | ||
| 33 | + ("delayed_peer", 2, 257, 80, 272, ["--tile-m", "256", "--delay-rank", "1", "--delay-ms", "20"]), | ||
| 34 | + ("hccl_cascade", 2, 257, 80, 272, ["--mode", "hccl", "--random"]), | ||
| 35 | + ] | ||
| 36 | + if len(devices) >= 4: | ||
| 37 | + cases.extend([ | ||
| 38 | + ("four_rank_fp16", 4, 513, 144, 528, ["--tile-m", "256"]), | ||
| 39 | + ("four_rank_bf16", 4, 257, 80, 272, ["--dtype", "bf16", "--tile-m", "256", "--random"]), | ||
| 40 | + ("multiple_peers_per_core", 4, 257, 80, 272, ["--cores", "1", "--tile-m", "256"]), | ||
| 41 | + ]) | ||
| 42 | + results = [] | ||
| 43 | + for number, (name, ranks, m, k, n, extra) in enumerate(cases): | ||
| 44 | + command = [str(args.binary.resolve()), f"tcp://127.0.0.1:{args.port + number}", | ||
| 45 | + str(ranks), str(m), str(k), str(n), "--devices", ",".join(map(str, devices[:ranks])), | ||
| 46 | + "--warmup", "1", "--iters", "5", "--timeout", "180", *extra] | ||
| 47 | + print(name, " ".join(command), flush=True) | ||
| 48 | + log_path = args.output_dir / f"{name}.log" | ||
| 49 | + with log_path.open("w") as log: | ||
| 50 | + result = subprocess.run(command, stdout=log, stderr=subprocess.STDOUT, check=False) | ||
| 51 | + log_text = log_path.read_text() | ||
| 52 | + passed = result.returncode == 0 and "RESULT | Example=ain_all_gather_matmul | Status=PASS" in log_text | ||
| 53 | + results.append({"case": name, "command": command, "returncode": result.returncode, | ||
| 54 | + "pass": passed, "log": str(log_path)}) | ||
| 55 | + (args.output_dir / "results.json").write_text(json.dumps(results, indent=2) + "\n") | ||
| 56 | + print(f"{name}: {'PASS' if passed else 'FAIL'}", flush=True) | ||
| 57 | + if not passed: | ||
| 58 | + print(log_text[-8000:], file=sys.stderr) | ||
| 59 | + return 1 | ||
| 60 | + return 0 | ||
| 61 | + | ||
| 62 | + | ||
| 63 | +if __name__ == "__main__": | ||
| 64 | + sys.exit(main()) | ||
| @@ -0,0 +1,79 @@ | |||
| 1 | +/* | ||
| 2 | + * Copyright (c) 2026 Huawei Technologies Co., Ltd. | ||
| 3 | + * Licensed under the CANN Open Software License Agreement Version 2.0. | ||
| 4 | + * See LICENSE in the root of the software repository. | ||
| 5 | + */ | ||
| 6 | + | ||
| 7 | + | ||
| 8 | + | ||
| 9 | + | ||
| 10 | + | ||
| 11 | + | ||
| 12 | + | ||
| 13 | + | ||
| 14 | +namespace ain_agmm { | ||
| 15 | + | ||
| 16 | +constexpr uint32_t CUBE_M = 256; | ||
| 17 | +constexpr uint32_t CUBE_N = 256; | ||
| 18 | +constexpr uint32_t CUBE_K = 64; | ||
| 19 | +constexpr uint64_t MAX_GET_BYTES = 256ULL * 1024 * 1024; | ||
| 20 | + | ||
| 21 | +// Host 生成后按值传入 kernel;M 是每 rank 的输入行数,输出总行数为 ranks*M。 | ||
| 22 | +// 设备侧地址乘积使用 uint64_t,避免大矩阵跨过 4 GiB 时发生 32 位截断。 | ||
| 23 | +struct Tiling { | ||
| 24 | + uint32_t m; | ||
| 25 | + uint32_t n; | ||
| 26 | + uint32_t k; | ||
| 27 | + uint32_t ranks; | ||
| 28 | + uint32_t cores; // 启动的 AIC 数;融合模式按 1:1 配对相同数量的 AIV。 | ||
| 29 | + uint32_t tileM; // 一轮通信处理的 M 行数,不等于单个 Cube tile 的行数。 | ||
| 30 | + uint32_t rounds; // ceil(M/tileM),最后一轮可能不足 tileM。 | ||
| 31 | +}; | ||
| 32 | + | ||
| 33 | +// 向上取整除法,用商和余数避免 value+divisor-1 的中间加法溢出。 | ||
| 34 | +inline uint64_t CeilDiv(uint64_t value, uint64_t divisor) | ||
| 35 | +{ | ||
| 36 | + return value / divisor + (value % divisor != 0); | ||
| 37 | +} | ||
| 38 | + | ||
| 39 | +// 在申请内存前检查尺寸乘积,确保不超出指针差值可表示范围。 | ||
| 40 | +inline uint64_t CheckedProduct(uint64_t a, uint64_t b) | ||
| 41 | +{ | ||
| 42 | + const uint64_t limit = static_cast<uint64_t>(std::numeric_limits<std::ptrdiff_t>::max()); | ||
| 43 | + if (b != 0 && a > limit / b) { | ||
| 44 | + throw std::invalid_argument("tensor size exceeds addressable memory"); | ||
| 45 | + } | ||
| 46 | + return a * b; | ||
| 47 | +} | ||
| 48 | + | ||
| 49 | +// 校验 shape/内存范围,并确定通信 tileM;这里做 Host 侧校验,kernel 依赖这些约束。 | ||
| 50 | +inline Tiling MakeTiling(uint32_t m, uint32_t n, uint32_t k, uint32_t ranks, uint32_t cores, uint32_t tileM) | ||
| 51 | +{ | ||
| 52 | + if (m == 0 || n == 0 || k == 0 || ranks < 2 || ranks > 64 || cores == 0 || n % 16 || k % 16) { | ||
| 53 | + throw std::invalid_argument("require M>0, N/K positive multiples of 16, 2<=ranks<=64 and cores>0"); | ||
| 54 | + } | ||
| 55 | + // 依次检查完整接收区、输出 C 和 B 的字节数,FP16/BF16 每元素均为 2 字节。 | ||
| 56 | + CheckedProduct(CheckedProduct(CheckedProduct(m, k), ranks), 2); | ||
| 57 | + CheckedProduct(CheckedProduct(CheckedProduct(m, n), ranks), 2); | ||
| 58 | + CheckedProduct(CheckedProduct(k, n), 2); | ||
| 59 | + const uint64_t maxRows = MAX_GET_BYTES / (static_cast<uint64_t>(k) * 2); | ||
| 60 | + if (maxRows == 0) { | ||
| 61 | + throw std::invalid_argument("one A row exceeds the 256 MiB AIN Get limit"); | ||
| 62 | + } | ||
| 63 | + if (tileM == 0) { | ||
| 64 | + const uint64_t nTiles = CeilDiv(n, CUBE_N); | ||
| 65 | + // 每个 M tile 可提供 (P-1)*nTiles 个远端计算任务;优先凑够一轮核数, | ||
| 66 | + // 再受 M 和单次 Get 的 256 MiB 上限约束。这只是初始启发式,不保证最优。 | ||
| 67 | + const uint64_t rowTiles = CeilDiv(cores, (ranks - 1) * nTiles); | ||
| 68 | + const uint64_t alignedMaxRows = maxRows / CUBE_M * CUBE_M; | ||
| 69 | + const uint64_t capacity = alignedMaxRows == 0 ? maxRows : alignedMaxRows; | ||
| 70 | + tileM = static_cast<uint32_t>(std::min<uint64_t>(m, std::min(rowTiles * CUBE_M, capacity))); | ||
| 71 | + } | ||
| 72 | + // 中间通信块按 Cube M 对齐;整段不足一块或最后的尾块允许非对齐,kernel 按实际行数处理。 | ||
| 73 | + if (tileM > m || tileM > maxRows || (tileM < m && tileM % CUBE_M != 0)) { | ||
| 74 | + throw std::invalid_argument("tile-m must fit M and 256 MiB; non-final tiles must be multiples of 256"); | ||
| 75 | + } | ||
| 76 | + return {m, n, k, ranks, cores, tileM, static_cast<uint32_t>(CeilDiv(m, tileM))}; | ||
| 77 | +} | ||
| 78 | + | ||
| 79 | +} // namespace ain_agmm | ||
| @@ -35,6 +35,7 @@ | |||
| 35 | 35 | ||
| 36 | 36 | ||
| 37 | 37 | ||
| 38 | + | ||
| 38 | 39 | ||
| 39 | 40 | ||
| 40 | 41 | ||
| @@ -213,7 +214,11 @@ private: | |||
| 213 | { | 214 | { |
| 214 | struct timeval tv = {}; | 215 | struct timeval tv = {}; |
| 215 | tv.tv_sec = kSockTimeoutSec; | 216 | tv.tv_sec = kSockTimeoutSec; |
| 216 | - return setsockopt(fd, SOL_SOCKET, SO_RCVTIMEO, &tv, sizeof(tv)) == 0 && | 217 | + // Headers and payloads are sent separately. Disable Nagle so delayed ACK |
| 218 | + // does not hold a barrier reply and skew rank launch times by tens of ms. | ||
| 219 | + const int noDelay = 1; | ||
| 220 | + return setsockopt(fd, IPPROTO_TCP, TCP_NODELAY, &noDelay, sizeof(noDelay)) == 0 && | ||
| 221 | + setsockopt(fd, SOL_SOCKET, SO_RCVTIMEO, &tv, sizeof(tv)) == 0 && | ||
| 217 | setsockopt(fd, SOL_SOCKET, SO_SNDTIMEO, &tv, sizeof(tv)) == 0; | 222 | setsockopt(fd, SOL_SOCKET, SO_SNDTIMEO, &tv, sizeof(tv)) == 0; |
| 218 | } | 223 | } |
| 219 | 224 | ||