已合并
在 experimental 中新增 Ascend950 FP4/FP8 量化矩阵乘 #791
Chen_HaoWen创建于 7月1日
在 experimental 中新增 Ascend950 FP4/FP8 量化矩阵乘 #791
已合并
共 38 个文件变更+4496-34
| @@ -0,0 +1,13 @@ | |||
| 1 | +# ---------------------------------------------------------------------------- | ||
| 2 | +# This program is free software, you can redistribute it and/or modify. | ||
| 3 | +# Copyright (c) 2026 Huawei Technologies Co., Ltd. | ||
| 4 | +# This file is a part of the CANN Open Software. | ||
| 5 | +# Licensed under CANN Open Software License Agreement Version 2.0 (the "License"). | ||
| 6 | +# Please refer to the License for details. You may not use this file except in compliance with the License. | ||
| 7 | +# THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED, INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE. | ||
| 8 | +# See LICENSE in the root of the software repository for the full text of the License. | ||
| 9 | +# ---------------------------------------------------------------------------- | ||
| 10 | + | ||
| 11 | +set_source_files_properties(fp4_mx_quant_matmul.cpp PROPERTIES LANGUAGE ASC) | ||
| 12 | +catlass_example_add_executable(ascend950_fp4_mx_quant_matmul mix fp4_mx_quant_matmul.cpp) | ||
| 13 | +target_compile_definitions(ascend950_fp4_mx_quant_matmul PRIVATE L2_CACHE_HINT) | ||
| @@ -0,0 +1,53 @@ | |||
| 1 | +# Ascend950 FP4 MX Quant Matmul Example Readme | ||
| 2 | + | ||
| 3 | +> **注意**:本样例位于 `experimental/` 目录下,如需编译运行,请先将样例目录拷贝至 `examples/` 下,并在 `examples/CMakeLists.txt` 中添加样例名称 `ascend950_fp4_mx_quant_matmul`。 | ||
| 4 | + | ||
如果这个样例的功能只是相对于54_ascend950_fp4_mx_matmul多了反量化乘的话,这个命名改成xx quant matmul ![]() ![]() | |||
| 5 | +## 功能介绍 | ||
| 6 | + | ||
| 7 | +- 演示 Ascend950 上的 **MX FP4 矩阵乘 + per-token/per-channel 量化 epilogue**。 | ||
| 8 | +- 计算:`D = perTokenScale * (MxScaleA * A @ MxScaleB * B) * perChannelScale`,本示例中 MxScale 固定为 1.0。 | ||
| 9 | +- A、B 元素类型为 `float4_e2m1x2_t`,per-token/per-channel scale 为 `float8_e4m3_t`,输出为 FP32。 | ||
| 10 | +- 默认布局为 A `RowMajor`、B `RowMajor`,与 `gen_data.py` 生成的数据一致。 | ||
| 11 | +- torch 接口的 MX scale 每个覆盖 K 轴 32 个元素,相邻两个 scale 成对存储。A scale 的连续存储形状为 `(M, ceil(K/64), 2)`,B scale 为 `(ceil(K/64), N, 2)`;K 的尾分组须补齐至一对。 | ||
| 12 | + | ||
| 13 | +## 代码组织 | ||
| 14 | +```text | ||
| 15 | +experimental | ||
| 16 | +└── matmul | ||
| 17 | + └── ascend950_fp4_mx_quant_matmul | ||
| 18 | + ├── CMakeLists.txt | ||
| 19 | + ├── README.md | ||
| 20 | + ├── gen_data.py | ||
| 21 | + ├── fp4_mx_quant_matmul.cpp | ||
| 22 | + └── test_85_ascend950_fp4_mx_quant_matmul.py | ||
| 23 | +``` | ||
| 24 | + | ||
| 25 | +## 使用示例 | ||
| 26 | +- 获取代码之后编译相应的算子可执行文件,可参考 [quickstart](../../../docs/zh/1_Practice/01_quick_start.md#编译执行)。本用例为 Ascend950(3510)算子,编译时需加 `-DCATLASS_ARCH=3510`。 | ||
| 27 | +- 执行算子 | ||
| 28 | +``` | ||
| 29 | +# 编译指定用例 | ||
| 30 | +bash scripts/build.sh ascend950_fp4_mx_quant_matmul -DCATLASS_ARCH=3510 | ||
| 31 | +# 生成测试样例(在 examples/ascend950_fp4_mx_quant_matmul/data 下生成 input/ 与 golden/) | ||
| 32 | +python3 examples/ascend950_fp4_mx_quant_matmul/gen_data.py 256 256 128 | ||
| 33 | +# 可选:--data-root <DIR> 指定在 DIR/data/ 下生成(默认在脚本所在目录下生成) | ||
| 34 | +# 输入参数分别对应 m, n, k;当前 n、k 需为偶数 | ||
| 35 | +# 在 output/bin 中执行,以匹配示例读取数据的相对路径 | ||
| 36 | +cd output/bin | ||
| 37 | +./ascend950_fp4_mx_quant_matmul 256 256 128 0 | ||
| 38 | +# 可执行文件名 |矩阵m轴|n轴|k轴|Device ID | ||
| 39 | +# Device ID 可选,默认为 0 | ||
| 40 | +``` | ||
| 41 | +执行结果如下,说明精度比对成功。 | ||
| 42 | +``` | ||
| 43 | +Compare success. | ||
| 44 | +``` | ||
| 45 | + | ||
| 46 | +## optest 测试 | ||
| 47 | + | ||
| 48 | +先按 [optest 说明](../../../tests/optest/README.md) 构建并安装 `torch_catlass`,然后在仓库根目录执行迁移后的测试件: | ||
| 49 | + | ||
| 50 | +```bash | ||
| 51 | +PYTHONPATH="$PWD/tests/optest/tests${PYTHONPATH:+:$PYTHONPATH}" \ | ||
| 52 | +python3 -m pytest experimental/matmul/ascend950_fp4_mx_quant_matmul/test_85_ascend950_fp4_mx_quant_matmul.py -v | ||
| 53 | +``` | ||
| @@ -0,0 +1,334 @@ | |||
| 1 | +/** | ||
| 2 | + * This program is free software, you can redistribute it and/or modify. | ||
| 3 | + * Copyright (c) 2026 Huawei Technologies Co., Ltd. | ||
| 4 | + * This file is a part of the CANN Open Software. | ||
| 5 | + * Licensed under CANN Open Software License Agreement Version 2.0 (the "License"). | ||
| 6 | + * Please refer to the License for details. You may not use this file except in compliance with the License. | ||
| 7 | + * THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED, INCLUDING | ||
| 8 | + * BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE. See LICENSE in the root of | ||
| 9 | + * the software repository for the full text of the License. | ||
| 10 | + */ | ||
| 11 | + | ||
| 12 | +// By setting the K_MAX_SHAPE_DIM macro, the dimension of the AscendC Tensor's ShapeInfo is configured to 0, | ||
| 13 | +// optimizing stack space. If you need to use the ShapeInfo of the AscendC Tensor, please undefine this macro. | ||
| 14 | + | ||
| 15 | + | ||
| 16 | + | ||
| 17 | + | ||
| 18 | + | ||
| 19 | + | ||
| 20 | + | ||
| 21 | + | ||
| 22 | + | ||
| 23 | + | ||
| 24 | + | ||
| 25 | + | ||
| 26 | + | ||
| 27 | + | ||
| 28 | + | ||
| 29 | + | ||
| 30 | + | ||
| 31 | + | ||
| 32 | + | ||
| 33 | + | ||
| 34 | + | ||
| 35 | + | ||
| 36 | + | ||
| 37 | + | ||
| 38 | + | ||
| 39 | + | ||
| 40 | + | ||
| 41 | + | ||
| 42 | +using namespace Catlass; | ||
| 43 | +using namespace tla; | ||
| 44 | + | ||
| 45 | +template <class Dtype> | ||
| 46 | +bool MatmulKernelRun( | ||
| 47 | + GM_ADDR deviceA, GM_ADDR deviceB, GM_ADDR deviceMxScaleA, GM_ADDR deviceMxScaleB, GM_ADDR deviceScale, | ||
| 48 | + GM_ADDR devicePerTokenScale, GM_ADDR deviceD, uint8_t*& deviceWorkspace, uint32_t m, uint32_t n, uint32_t k, | ||
| 49 | + aclrtStream stream) | ||
| 50 | +{ | ||
| 51 | + auto aicCoreNum = platform_ascendc::PlatformAscendCManager::GetInstance()->GetCoreNumAic(); | ||
| 52 | + | ||
| 53 | + constexpr uint32_t workspaceStages = 2; | ||
| 54 | + size_t sizeWorkspace = 0; | ||
| 55 | + uint32_t mxScaleK = CeilDiv<MX_SCALE_GROUP_NUM>(k); | ||
| 56 | + uint32_t mxScaleAlignedK = RoundUp<2>(mxScaleK); | ||
| 57 | + | ||
| 58 | + using ElementA = float4_e2m1x2_t; | ||
| 59 | + using ElementB = float4_e2m1x2_t; | ||
| 60 | + using ElementC = float; | ||
| 61 | + using ElementMxScale = float8_e8m0_t; | ||
| 62 | + using ElementScale = float8_e4m3_t; // per-channel scale | ||
| 63 | + using ElementPerTokenScale = float8_e4m3_t; // per-token scale | ||
| 64 | + using ElementD = Dtype; | ||
| 65 | + | ||
| 66 | + using LayoutTagA = layout::RowMajor; | ||
| 67 | + using LayoutTagB = layout::RowMajor; | ||
| 68 | + using LayoutTagC = layout::RowMajor; | ||
| 69 | + using LayoutTagD = layout::RowMajor; | ||
| 70 | + using LayoutTagScale = layout::VectorLayout; | ||
| 71 | + using LayoutTagPerTokenScale = layout::VectorLayout; | ||
| 72 | + | ||
| 73 | + auto layoutA = tla::MakeLayout<ElementA, LayoutTagA>(m, k); | ||
| 74 | + auto layoutB = tla::MakeLayout<ElementB, LayoutTagB>(k, n); | ||
| 75 | + auto layoutMxScaleA = tla::MakeMxScaleLayout<ElementMxScale, LayoutTagA, false>(m, mxScaleK); | ||
| 76 | + auto layoutMxScaleB = tla::MakeMxScaleLayout<ElementMxScale, LayoutTagB, true>(mxScaleK, n); | ||
| 77 | + | ||
| 78 | + LayoutTagD layoutD{m, n}; | ||
| 79 | + LayoutTagScale layoutScale{n}; | ||
| 80 | + LayoutTagPerTokenScale layoutPerTokenScale{m}; | ||
| 81 | + | ||
| 82 | + using ArchTag = Arch::Ascend950; | ||
| 83 | + constexpr bool enableUnitFlag = true; | ||
| 84 | + using DispatchPolicy = Gemm::MmadMx<ArchTag, enableUnitFlag>; | ||
| 85 | + | ||
| 86 | + using L1TileShape = Shape<Int<256>, Int<256>, Int<512>>; | ||
| 87 | + using L0TileShape = Shape<Int<256>, Int<256>, Int<256>>; | ||
| 88 | + | ||
| 89 | + using TileCopy = Gemm::Tile::PackedMxTileCopyTla< | ||
| 90 | + ArchTag, ElementA, LayoutTagA, ElementB, LayoutTagB, ElementMxScale, decltype(layoutMxScaleA), ElementMxScale, | ||
| 91 | + decltype(layoutMxScaleB), ElementC, LayoutTagC, void>; | ||
| 92 | + | ||
| 93 | + using BlockMmad = Gemm::Block::BlockMmadTla< | ||
| 94 | + DispatchPolicy, L1TileShape, L0TileShape, ElementA, ElementB, ElementC, void, TileCopy>; | ||
| 95 | + | ||
| 96 | + using CType = Gemm::GemmType<ElementC, layout::RowMajor>; | ||
| 97 | + | ||
| 98 | + constexpr uint32_t ubStages = 2; | ||
| 99 | + using EpilogueDispatchPolicy = Epilogue::EpilogueAscend950PerTokenPerChannelQuant<ubStages>; | ||
| 100 | + using ScaleType = Gemm::GemmType<ElementScale, layout::VectorLayout>; | ||
| 101 | + using PerTokenScaleType = Gemm::GemmType<ElementPerTokenScale, layout::VectorLayout>; | ||
| 102 | + using DType = Gemm::GemmType<ElementD, layout::RowMajor>; | ||
| 103 | + | ||
| 104 | + using RowBroadcastMulType = Gemm::GemmType<float, layout::RowMajor>; | ||
| 105 | + using BroadcastOneBlkType = Gemm::GemmType<float, layout::RowMajor>; | ||
| 106 | + using OneBlkColumnBroadcastMulType = Gemm::GemmType<float, layout::RowMajor>; | ||
| 107 | + | ||
| 108 | + using EpilogueTileShape = MatrixShape<32, 256>; | ||
| 109 | + using TileRowBroadcastMul = Epilogue::Tile::TileRowBroadcastMul<ArchTag, RowBroadcastMulType, EpilogueTileShape>; | ||
| 110 | + using TileBroadcastOneBlk = | ||
| 111 | + Epilogue::Tile::TileBroadcastOneBlk<ArchTag, BroadcastOneBlkType, EpilogueTileShape::ROW>; | ||
| 112 | + using TileOneBlkColumnBroadcastMul = | ||
| 113 | + Epilogue::Tile::TileOneBlkColumnBroadcastMul<ArchTag, OneBlkColumnBroadcastMulType, EpilogueTileShape>; | ||
| 114 | + using TileCopyEpilogue = Epilogue::Tile::TileCopy<ArchTag, CType, ScaleType, PerTokenScaleType, DType>; | ||
| 115 | + using TileScheduler = Epilogue::Tile::EpilogueHorizontalTileSwizzle; | ||
| 116 | + | ||
| 117 | + using BlockEpilogue = Epilogue::Block::BlockEpilogue< | ||
| 118 | + EpilogueDispatchPolicy, CType, ScaleType, PerTokenScaleType, DType, TileRowBroadcastMul, TileBroadcastOneBlk, | ||
| 119 | + TileOneBlkColumnBroadcastMul, TileCopyEpilogue, TileScheduler>; | ||
| 120 | + | ||
| 121 | + using BlockScheduler = typename Gemm::Block::GemmIdentityBlockSwizzle<3, 0>; | ||
| 122 | + | ||
| 123 | + using MatmulKernel = | ||
| 124 | + Gemm::Kernel::MxMatmulPerTokenPerChannelTla<BlockMmad, BlockEpilogue, BlockScheduler, workspaceStages>; | ||
| 125 | + | ||
| 126 | + using MatmulAdapter = Gemm::Device::DeviceGemm<MatmulKernel>; | ||
| 127 | + | ||
| 128 | + GemmCoord problemShape{m, n, k}; | ||
| 129 | + typename MatmulKernel::Arguments arguments{ | ||
| 130 | + problemShape, aicCoreNum, deviceA, layoutA, deviceB, layoutB, | ||
| 131 | + deviceMxScaleA, layoutMxScaleA, deviceMxScaleB, layoutMxScaleB, deviceScale, layoutScale, | ||
| 132 | + devicePerTokenScale, layoutPerTokenScale, deviceD, layoutD}; | ||
| 133 | + | ||
| 134 | + MatmulAdapter matmulOp; | ||
| 135 | + if (matmulOp.CanImplement(arguments) != Status::kSuccess) { | ||
| 136 | + std::cerr << "Matmul shape cannot be implemented." << std::endl; | ||
| 137 | + return false; | ||
| 138 | + } | ||
| 139 | + sizeWorkspace = matmulOp.GetWorkspaceSize(arguments); | ||
| 140 | + if (sizeWorkspace > 0) { | ||
| 141 | + ACL_CHECK(aclrtMalloc(reinterpret_cast<void**>(&deviceWorkspace), sizeWorkspace, ACL_MEM_MALLOC_HUGE_FIRST)); | ||
| 142 | + } | ||
| 143 | + matmulOp.Initialize(arguments, deviceWorkspace); | ||
| 144 | + matmulOp(stream, aicCoreNum); | ||
| 145 | + return true; | ||
| 146 | +} | ||
| 147 | + | ||
| 148 | +using Options = GemmOptions; | ||
| 149 | + | ||
| 150 | +static const std::string kDataRoot = "../../examples/ascend950_fp4_mx_quant_matmul/data"; | ||
| 151 | + | ||
| 152 | +int Run(const Options& options) | ||
| 153 | +{ | ||
| 154 | + if (options.problemShape.m() == 0 || options.problemShape.n() == 0 || options.problemShape.k() == 0) { | ||
| 155 | + std::cerr << "M, N and K must be positive." << std::endl; | ||
| 156 | + return EXIT_FAILURE; | ||
| 157 | + } | ||
| 158 | + if (options.problemShape.n() % 2 != 0 || options.problemShape.k() % 2 != 0) { | ||
| 159 | + std::cerr << "N and K must be even for packed FP4." << std::endl; | ||
| 160 | + return EXIT_FAILURE; | ||
| 161 | + } | ||
| 162 | + aclrtStream stream{nullptr}; | ||
| 163 | + | ||
| 164 | + ACL_CHECK(aclInit(nullptr)); | ||
| 165 | + ACL_CHECK(aclrtSetDevice(options.deviceId)); | ||
| 166 | + ACL_CHECK(aclrtCreateStream(&stream)); | ||
| 167 | + | ||
| 168 | + uint32_t m = options.problemShape.m(); | ||
| 169 | + uint32_t n = options.problemShape.n(); | ||
| 170 | + uint32_t k = options.problemShape.k(); | ||
| 171 | + uint32_t mxScaleK = CeilDiv<MX_SCALE_GROUP_NUM>(k); | ||
| 172 | + uint32_t mxScaleAlignedK = RoundUp<2>(mxScaleK); | ||
| 173 | + | ||
| 174 | + using ElementA = float4_e2m1x2_t; | ||
| 175 | + using ElementB = float4_e2m1x2_t; | ||
| 176 | + using ElementMxScale = float8_e8m0_t; | ||
| 177 | + using ElementScale = float8_e4m3_t; | ||
| 178 | + using ElementPerTokenScale = float8_e4m3_t; | ||
| 179 | + | ||
| 180 | + using LayoutTagA = layout::RowMajor; | ||
| 181 | + using LayoutTagB = layout::ColumnMajor; | ||
| 182 | + | ||
| 183 | + LayoutTagA tagA = LayoutTagA::MakeLayout<ElementA>(m, k); | ||
| 184 | + LayoutTagB tagB = LayoutTagB::MakeLayout<ElementB>(k, n); | ||
| 185 | + | ||
| 186 | + size_t lenA = tagA.Capacity(); | ||
| 187 | + size_t lenB = tagB.Capacity(); | ||
| 188 | + uint32_t lenMxScaleA = m * mxScaleAlignedK; | ||
| 189 | + uint32_t lenMxScaleB = mxScaleAlignedK * n; | ||
| 190 | + size_t lenD = static_cast<size_t>(m) * n; | ||
| 191 | + size_t lenScale = static_cast<size_t>(n); | ||
| 192 | + size_t lenPerTokenScale = static_cast<size_t>(m); | ||
| 193 | + | ||
| 194 | + size_t sizeA = lenA / 2; | ||
| 195 | + size_t sizeB = lenB / 2; | ||
| 196 | + size_t sizeMxScaleA = lenMxScaleA * sizeof(ElementMxScale); | ||
| 197 | + size_t sizeMxScaleB = lenMxScaleB * sizeof(ElementMxScale); | ||
| 198 | + size_t sizeD = lenD * sizeof(float); // default float output | ||
| 199 | + size_t sizeScale = lenScale * sizeof(ElementScale); | ||
| 200 | + size_t sizePerTokenScale = lenPerTokenScale * sizeof(ElementPerTokenScale); | ||
| 201 | + | ||
| 202 | + std::vector<int8_t> hostA(sizeA); | ||
| 203 | + std::vector<int8_t> hostB(sizeB); | ||
| 204 | + // MxScale 固定为 1: 用 fp8_e8m0 编码的 1.0,即指数 = 127,编码 = 127 | ||
| 205 | + // 但直接填 0x7F (127) 即 2^(127-128) = 2^(-1) = 0.5,不对 | ||
| 206 | + // fp8_e8m0: value = 2^(exp - 128),exp=128 => value=1.0,但 exp 是 8 位无符号,范围 0~255 | ||
| 207 | + // exp byte = 128 => 0x80 => value = 2^(128-128) = 2^0 = 1.0 | ||
| 208 | + std::vector<uint8_t> hostMxScaleA(lenMxScaleA, 0x7F); // all 1.0 | ||
| 209 | + std::vector<uint8_t> hostMxScaleB(lenMxScaleB, 0x7F); // all 1.0 | ||
| 210 | + std::vector<int8_t> hostScale(lenScale); | ||
| 211 | + std::vector<int8_t> hostPerTokenScale(lenPerTokenScale); | ||
| 212 | + | ||
| 213 | + const auto releaseAclEarly = [&]() { | ||
| 214 | + ACL_CHECK(aclrtDestroyStream(stream)); | ||
| 215 | + ACL_CHECK(aclrtResetDevice(options.deviceId)); | ||
| 216 | + ACL_CHECK(aclFinalize()); | ||
| 217 | + }; | ||
| 218 | + if (!ReadFile(kDataRoot + "/input/a_4.bin", hostA.data(), sizeA)) { | ||
| 219 | + releaseAclEarly(); | ||
| 220 | + return EXIT_FAILURE; | ||
| 221 | + } | ||
| 222 | + if (!ReadFile(kDataRoot + "/input/b_4.bin", hostB.data(), sizeB)) { | ||
| 223 | + releaseAclEarly(); | ||
| 224 | + return EXIT_FAILURE; | ||
| 225 | + } | ||
| 226 | + if (!ReadFile(kDataRoot + "/input/a_scale4.bin", hostPerTokenScale.data(), sizePerTokenScale)) { | ||
| 227 | + releaseAclEarly(); | ||
| 228 | + return EXIT_FAILURE; | ||
| 229 | + } | ||
| 230 | + if (!ReadFile(kDataRoot + "/input/b_scale4.bin", hostScale.data(), sizeScale)) { | ||
| 231 | + releaseAclEarly(); | ||
| 232 | + return EXIT_FAILURE; | ||
| 233 | + } | ||
| 234 | + | ||
| 235 | + uint8_t* deviceA{nullptr}; | ||
| 236 | + ACL_CHECK(aclrtMalloc(reinterpret_cast<void**>(&deviceA), sizeA, ACL_MEM_MALLOC_HUGE_FIRST)); | ||
| 237 | + ACL_CHECK(aclrtMemcpy(deviceA, sizeA, hostA.data(), sizeA, ACL_MEMCPY_HOST_TO_DEVICE)); | ||
| 238 | + | ||
| 239 | + uint8_t* deviceB{nullptr}; | ||
| 240 | + ACL_CHECK(aclrtMalloc(reinterpret_cast<void**>(&deviceB), sizeB, ACL_MEM_MALLOC_HUGE_FIRST)); | ||
| 241 | + ACL_CHECK(aclrtMemcpy(deviceB, sizeB, hostB.data(), sizeB, ACL_MEMCPY_HOST_TO_DEVICE)); | ||
| 242 | + | ||
| 243 | + uint8_t* deviceMxScaleA{nullptr}; | ||
| 244 | + ACL_CHECK(aclrtMalloc(reinterpret_cast<void**>(&deviceMxScaleA), sizeMxScaleA, ACL_MEM_MALLOC_HUGE_FIRST)); | ||
| 245 | + ACL_CHECK(aclrtMemcpy(deviceMxScaleA, sizeMxScaleA, hostMxScaleA.data(), sizeMxScaleA, ACL_MEMCPY_HOST_TO_DEVICE)); | ||
| 246 | + | ||
| 247 | + uint8_t* deviceMxScaleB{nullptr}; | ||
| 248 | + ACL_CHECK(aclrtMalloc(reinterpret_cast<void**>(&deviceMxScaleB), sizeMxScaleB, ACL_MEM_MALLOC_HUGE_FIRST)); | ||
| 249 | + ACL_CHECK(aclrtMemcpy(deviceMxScaleB, sizeMxScaleB, hostMxScaleB.data(), sizeMxScaleB, ACL_MEMCPY_HOST_TO_DEVICE)); | ||
| 250 | + | ||
| 251 | + uint8_t* deviceScale{nullptr}; | ||
| 252 | + ACL_CHECK(aclrtMalloc(reinterpret_cast<void**>(&deviceScale), sizeScale, ACL_MEM_MALLOC_HUGE_FIRST)); | ||
| 253 | + ACL_CHECK(aclrtMemcpy(deviceScale, sizeScale, hostScale.data(), sizeScale, ACL_MEMCPY_HOST_TO_DEVICE)); | ||
| 254 | + | ||
| 255 | + uint8_t* devicePerTokenScale{nullptr}; | ||
| 256 | + ACL_CHECK( | ||
| 257 | + aclrtMalloc(reinterpret_cast<void**>(&devicePerTokenScale), sizePerTokenScale, ACL_MEM_MALLOC_HUGE_FIRST)); | ||
| 258 | + ACL_CHECK(aclrtMemcpy( | ||
| 259 | + devicePerTokenScale, sizePerTokenScale, hostPerTokenScale.data(), sizePerTokenScale, | ||
| 260 | + ACL_MEMCPY_HOST_TO_DEVICE)); | ||
| 261 | + | ||
| 262 | + uint8_t* deviceD{nullptr}; | ||
| 263 | + ACL_CHECK(aclrtMalloc(reinterpret_cast<void**>(&deviceD), sizeD, ACL_MEM_MALLOC_HUGE_FIRST)); | ||
| 264 | + | ||
| 265 | + uint8_t* deviceWorkspace{nullptr}; | ||
| 266 | + | ||
| 267 | + bool launched = MatmulKernelRun<float>( | ||
| 268 | + deviceA, deviceB, deviceMxScaleA, deviceMxScaleB, deviceScale, devicePerTokenScale, deviceD, deviceWorkspace, m, | ||
| 269 | + n, k, stream); | ||
| 270 | + | ||
| 271 | + if (!launched) { | ||
| 272 | + ACL_CHECK(aclrtFree(deviceA)); | ||
| 273 | + ACL_CHECK(aclrtFree(deviceB)); | ||
| 274 | + ACL_CHECK(aclrtFree(deviceMxScaleA)); | ||
| 275 | + ACL_CHECK(aclrtFree(deviceMxScaleB)); | ||
| 276 | + ACL_CHECK(aclrtFree(deviceScale)); | ||
| 277 | + ACL_CHECK(aclrtFree(devicePerTokenScale)); | ||
| 278 | + ACL_CHECK(aclrtFree(deviceD)); | ||
| 279 | + releaseAclEarly(); | ||
| 280 | + return EXIT_FAILURE; | ||
| 281 | + } | ||
| 282 | + | ||
| 283 | + ACL_CHECK(aclrtSynchronizeStream(stream)); | ||
| 284 | + | ||
| 285 | + std::vector<float> hostD(lenD); | ||
| 286 | + ACL_CHECK(aclrtMemcpy(hostD.data(), sizeD, deviceD, sizeD, ACL_MEMCPY_DEVICE_TO_HOST)); | ||
| 287 | + | ||
| 288 | + std::vector<float> hostGolden(lenD); | ||
| 289 | + if (!ReadFile(kDataRoot + "/golden/expected_data.bin", hostGolden.data(), sizeD)) { | ||
| 290 | + ACL_CHECK(aclrtFree(deviceA)); | ||
| 291 | + ACL_CHECK(aclrtFree(deviceB)); | ||
| 292 | + ACL_CHECK(aclrtFree(deviceMxScaleA)); | ||
| 293 | + ACL_CHECK(aclrtFree(deviceMxScaleB)); | ||
| 294 | + ACL_CHECK(aclrtFree(deviceScale)); | ||
| 295 | + ACL_CHECK(aclrtFree(devicePerTokenScale)); | ||
| 296 | + ACL_CHECK(aclrtFree(deviceD)); | ||
| 297 | + if (deviceWorkspace != nullptr) { | ||
| 298 | + ACL_CHECK(aclrtFree(deviceWorkspace)); | ||
| 299 | + } | ||
| 300 | + releaseAclEarly(); | ||
| 301 | + return EXIT_FAILURE; | ||
| 302 | + } | ||
| 303 | + | ||
| 304 | + std::vector<uint64_t> errorIndices = golden::CompareData(hostD, hostGolden, k); | ||
| 305 | + if (errorIndices.empty()) { | ||
| 306 | + std::cout << "Compare success." << std::endl; | ||
| 307 | + } else { | ||
| 308 | + std::cerr << "Compare failed. Error count: " << errorIndices.size() << std::endl; | ||
| 309 | + } | ||
| 310 | + | ||
| 311 | + ACL_CHECK(aclrtFree(deviceA)); | ||
| 312 | + ACL_CHECK(aclrtFree(deviceB)); | ||
| 313 | + ACL_CHECK(aclrtFree(deviceMxScaleA)); | ||
| 314 | + ACL_CHECK(aclrtFree(deviceMxScaleB)); | ||
| 315 | + ACL_CHECK(aclrtFree(deviceScale)); | ||
| 316 | + ACL_CHECK(aclrtFree(devicePerTokenScale)); | ||
| 317 | + ACL_CHECK(aclrtFree(deviceD)); | ||
| 318 | + if (deviceWorkspace != nullptr) { | ||
| 319 | + ACL_CHECK(aclrtFree(deviceWorkspace)); | ||
| 320 | + } | ||
| 321 | + ACL_CHECK(aclrtDestroyStream(stream)); | ||
| 322 | + ACL_CHECK(aclrtResetDevice(options.deviceId)); | ||
| 323 | + ACL_CHECK(aclFinalize()); | ||
| 324 | + return errorIndices.empty() ? EXIT_SUCCESS : EXIT_FAILURE; | ||
| 325 | +} | ||
| 326 | + | ||
| 327 | +int main(int argc, const char** argv) | ||
| 328 | +{ | ||
| 329 | + Options options; | ||
| 330 | + if (options.Parse(argc, argv) != 0) { | ||
| 331 | + return -1; | ||
| 332 | + } | ||
| 333 | + return Run(options); | ||
| 334 | +} | ||
| @@ -0,0 +1,114 @@ | |||
| 1 | +#!/usr/bin/env python3 | ||
| 2 | +# -*- coding: utf-8 -*- | ||
| 3 | +# ---------------------------------------------------------------------------- | ||
| 4 | +# This program is free software, you can redistribute it and/or modify. | ||
| 5 | +# Copyright (c) 2026 Huawei Technologies Co., Ltd. | ||
| 6 | +# This file is a part of the CANN Open Software. | ||
| 7 | +# Licensed under CANN Open Software License Agreement Version 2.0 (the "License"). | ||
| 8 | +# Please refer to the License for details. You may not use this file except in compliance with the License. | ||
| 9 | +# THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED, INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE. | ||
| 10 | +# See LICENSE in the root of the software repository for the full text of the License. | ||
| 11 | +# ---------------------------------------------------------------------------- | ||
| 12 | + | ||
| 13 | +import argparse | ||
| 14 | +import os | ||
| 15 | +from typing import Optional | ||
| 16 | + | ||
| 17 | +import numpy as np | ||
| 18 | +import torch | ||
| 19 | +from ml_dtypes import float4_e2m1fn, float8_e4m3fn | ||
| 20 | + | ||
| 21 | +_SCRIPT_DIR = os.path.dirname(os.path.abspath(__file__)) | ||
| 22 | +_FP4_MAX = 6.0 | ||
| 23 | + | ||
| 24 | + | ||
| 25 | +def _resolve_workspace(data_root_cli: Optional[str]) -> str: | ||
| 26 | + if data_root_cli is not None: | ||
| 27 | + root = data_root_cli.strip() | ||
| 28 | + if root: | ||
| 29 | + return os.path.abspath(os.path.expanduser(root)) | ||
| 30 | + return _SCRIPT_DIR | ||
| 31 | + | ||
| 32 | + | ||
| 33 | +def _pack_fp4_to_int8(fp4_array: np.ndarray) -> np.ndarray: | ||
| 34 | + flat = fp4_array.flatten() | ||
| 35 | + if len(flat) % 2 != 0: | ||
| 36 | + flat = np.concatenate([flat, np.array([0], dtype=flat.dtype)]) | ||
| 37 | + low = flat[0::2].view(np.uint8) & 0x0F | ||
| 38 | + high = flat[1::2].view(np.uint8) & 0x0F | ||
| 39 | + packed = (high << 4) | low | ||
| 40 | + return packed.astype(np.int8) | ||
| 41 | + | ||
| 42 | + | ||
| 43 | +def _unpack_fp4_from_int8(packed: np.ndarray, num_elements: int) -> np.ndarray: | ||
| 44 | + flat = packed.astype(np.uint8) | ||
| 45 | + low = flat & 0x0F | ||
| 46 | + high = (flat >> 4) & 0x0F | ||
| 47 | + unpacked = np.empty(flat.size * 2, dtype=np.uint8) | ||
| 48 | + unpacked[0::2] = low | ||
| 49 | + unpacked[1::2] = high | ||
| 50 | + return unpacked[:num_elements].view(float4_e2m1fn).astype(np.float32) | ||
| 51 | + | ||
| 52 | + | ||
| 53 | +def gen_data(m: int, n: int, k: int, workspace: str) -> None: | ||
| 54 | + if k % 2 != 0: | ||
| 55 | + raise ValueError(f"K={k} must be even for fp4x2 format") | ||
| 56 | + if n % 2 != 0: | ||
| 57 | + raise ValueError(f"N={n} must be even for fp4x2 format") | ||
| 58 | + | ||
| 59 | + data_dir = os.path.join(workspace, "data") | ||
| 60 | + input_dir = os.path.join(data_dir, "input") | ||
| 61 | + golden_dir = os.path.join(data_dir, "golden") | ||
| 62 | + os.makedirs(input_dir, exist_ok=True) | ||
| 63 | + os.makedirs(golden_dir, exist_ok=True) | ||
| 64 | + | ||
| 65 | + torch.manual_seed(0) | ||
| 66 | + a_fp32 = torch.randn((m, k), dtype=torch.float32) * 2 | ||
| 67 | + b_fp32 = torch.randn((k, n), dtype=torch.float32) * 2 | ||
| 68 | + | ||
| 69 | + a_max = a_fp32.abs().max(dim=1).values | ||
| 70 | + b_max = b_fp32.abs().max(dim=0).values | ||
| 71 | + a_max[a_max == 0] = 1.0 | ||
| 72 | + b_max[b_max == 0] = 1.0 | ||
| 73 | + | ||
| 74 | + a_scale_fp32 = a_max / _FP4_MAX | ||
| 75 | + b_scale_fp32 = b_max / _FP4_MAX | ||
| 76 | + | ||
| 77 | + a_fp4 = (a_fp32 / a_scale_fp32.view(-1, 1)).numpy().astype(float4_e2m1fn) | ||
| 78 | + b_fp4 = (b_fp32 / b_scale_fp32.view(1, -1)).numpy().astype(float4_e2m1fn) | ||
| 79 | + a_scale_fp8 = a_scale_fp32.numpy().astype(float8_e4m3fn) | ||
| 80 | + b_scale_fp8 = b_scale_fp32.numpy().astype(float8_e4m3fn) | ||
| 81 | + | ||
| 82 | + _pack_fp4_to_int8(a_fp4).tofile(os.path.join(input_dir, "a_4.bin")) | ||
| 83 | + _pack_fp4_to_int8(b_fp4).tofile(os.path.join(input_dir, "b_4.bin")) | ||
| 84 | + a_scale_fp8.view(np.int8).tofile(os.path.join(input_dir, "a_scale4.bin")) | ||
| 85 | + b_scale_fp8.view(np.int8).tofile(os.path.join(input_dir, "b_scale4.bin")) | ||
| 86 | + | ||
| 87 | + a_val = _unpack_fp4_from_int8(np.fromfile(os.path.join(input_dir, "a_4.bin"), dtype=np.int8), m * k).reshape(m, k) | ||
| 88 | + b_val = _unpack_fp4_from_int8(np.fromfile(os.path.join(input_dir, "b_4.bin"), dtype=np.int8), k * n).reshape(k, n) | ||
| 89 | + per_token = a_scale_fp8.astype(np.float32) | ||
| 90 | + per_channel = b_scale_fp8.astype(np.float32) | ||
| 91 | + | ||
| 92 | + c_fp32 = a_val @ b_val | ||
| 93 | + golden = c_fp32 * per_token[:, np.newaxis] * per_channel[np.newaxis, :] | ||
| 94 | + golden.astype(np.float32).tofile(os.path.join(golden_dir, "expected_data.bin")) | ||
| 95 | + | ||
| 96 | + | ||
| 97 | +if __name__ == "__main__": | ||
| 98 | + parser = argparse.ArgumentParser( | ||
| 99 | + description="Generate MX-FP4 per-token/per-channel inputs and FP32 golden under " | ||
| 100 | + "<data-root>/data/.", | ||
| 101 | + ) | ||
| 102 | + parser.add_argument( | ||
| 103 | + "--data-root", | ||
| 104 | + default=None, | ||
| 105 | + metavar="DIR", | ||
| 106 | + help="Directory under which data/input and data/golden are created. " | ||
| 107 | + "Default: this script's directory.", | ||
| 108 | + ) | ||
| 109 | + parser.add_argument("m", type=int) | ||
| 110 | + parser.add_argument("n", type=int) | ||
| 111 | + parser.add_argument("k", type=int) | ||
| 112 | + args = parser.parse_args() | ||
| 113 | + workspace = _resolve_workspace(args.data_root) | ||
| 114 | + gen_data(args.m, args.n, args.k, workspace) | ||
| @@ -0,0 +1,82 @@ | |||
| 1 | +# This program is free software, you can redistribute it and/or modify. | ||
| 2 | +# Copyright (c) 2026 Huawei Technologies Co., Ltd. | ||
| 3 | +# This file is a part of the CANN Open Software. | ||
| 4 | +# Licensed under CANN Open Software License Agreement Version 2.0 (the "License"). | ||
| 5 | +# Please refer to the License for details. You may not use this file except in compliance with the License. | ||
| 6 | +# THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED, INCLUDING | ||
| 7 | +# BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE. See LICENSE in the root of | ||
| 8 | +# the software repository for the full text of the License. | ||
| 9 | + | ||
| 10 | +import pytest | ||
| 11 | +import torch | ||
| 12 | +import torch_catlass | ||
| 13 | + | ||
| 14 | +from common import only_on_3510 | ||
| 15 | +from mx_golden import prepare_fp4_mx_quant_matmul_inputs | ||
| 16 | +from quant_matmul_test_utils import QUANT_MATMUL_SHAPES, check_quant_result, signed_scales, with_storage_offset | ||
| 17 | + | ||
| 18 | + | ||
| 19 | +def _fp4_dtype_supported() -> bool: | ||
| 20 | + if not hasattr(torch, "float4_e2m1fn_x2"): | ||
| 21 | + return False | ||
| 22 | + try: | ||
| 23 | + packed = torch.zeros(4, dtype=torch.uint8) | ||
| 24 | + out = torch.empty((2, 2), dtype=torch.float4_e2m1fn_x2) | ||
| 25 | + out.view(torch.uint8).flatten()[: packed.numel()].copy_(packed) | ||
| 26 | + return out.shape == (2, 2) | ||
| 27 | + except (RuntimeError, TypeError): | ||
| 28 | + return False | ||
| 29 | + | ||
| 30 | + | ||
| 31 | +pytestmark = pytest.mark.skipif( | ||
| 32 | + not _fp4_dtype_supported(), | ||
| 33 | + reason="torch.float4_e2m1fn_x2 tensor construction is unavailable in this PyTorch build", | ||
| 34 | +) | ||
| 35 | + | ||
| 36 | + | ||
| 37 | + | ||
| 38 | + | ||
| 39 | + | ||
| 40 | +def test_ascend950_fp4_mx_quant_matmul(shape, seed, record_property): | ||
| 41 | + """Compare example 85 against the quantized CPU reference.""" | ||
| 42 | + m, n, k = shape | ||
| 43 | + a, b, mx_scale_a, mx_scale_b, per_token, per_channel, expected = ( | ||
| 44 | + prepare_fp4_mx_quant_matmul_inputs(m, n, k, device="npu", seed=seed, vary_mx_scales=bool(seed % 2)) | ||
| 45 | + ) | ||
| 46 | + if seed % 2: | ||
| 47 | + per_token, per_channel, expected = signed_scales(per_token, per_channel, expected) | ||
| 48 | + | ||
| 49 | + result = torch_catlass.ascend950_fp4_mx_quant_matmul( | ||
| 50 | + a, b, mx_scale_a, mx_scale_b, per_token, per_channel | ||
| 51 | + ) | ||
| 52 | + | ||
| 53 | + check_quant_result(result, expected, record_property) | ||
| 54 | + | ||
| 55 | + | ||
| 56 | + | ||
| 57 | + | ||
| 58 | +def test_fp4_quant_storage_offset(input_index, record_property): | ||
| 59 | + *inputs, expected = prepare_fp4_mx_quant_matmul_inputs(127, 130, 66, vary_mx_scales=True) | ||
| 60 | + inputs[input_index] = with_storage_offset(inputs[input_index]) | ||
| 61 | + result = torch_catlass.ascend950_fp4_mx_quant_matmul(*inputs) | ||
| 62 | + check_quant_result(result, expected, record_property) | ||
| 63 | + | ||
| 64 | + | ||
| 65 | + | ||
| 66 | + | ||
| 67 | +def test_fp4_quant_rejects_empty_shape(axis): | ||
| 68 | + shape = [16, 16, 16] | ||
| 69 | + shape[axis] = 0 | ||
| 70 | + m, n, k = shape | ||
| 71 | + a = torch.empty((m, k), dtype=torch.float4_e2m1fn_x2, device="npu") | ||
| 72 | + b = torch.empty((k, n), dtype=torch.float4_e2m1fn_x2, device="npu") | ||
| 73 | + mx_a = torch.empty((m, 2), dtype=torch.float8_e8m0fnu, device="npu") | ||
| 74 | + mx_b = torch.empty((n, 2), dtype=torch.float8_e8m0fnu, device="npu") | ||
| 75 | + token = torch.empty(m, dtype=torch.float8_e4m3fn, device="npu") | ||
| 76 | + channel = torch.empty(n, dtype=torch.float8_e4m3fn, device="npu") | ||
| 77 | + with pytest.raises(RuntimeError, match="must be positive"): | ||
| 78 | + torch_catlass.ascend950_fp4_mx_quant_matmul(a, b, mx_a, mx_b, token, channel) | ||
| 79 | + | ||
| 80 | + | ||
| 81 | +if __name__ == "__main__": | ||
| 82 | + pytest.main([__file__, "-v", "-s"]) | ||
| @@ -0,0 +1,13 @@ | |||
| 1 | +# ---------------------------------------------------------------------------- | ||
| 2 | +# This program is free software, you can redistribute it and/or modify. | ||
| 3 | +# Copyright (c) 2026 Huawei Technologies Co., Ltd. | ||
| 4 | +# This file is a part of the CANN Open Software. | ||
| 5 | +# Licensed under CANN Open Software License Agreement Version 2.0 (the "License"). | ||
| 6 | +# Please refer to the License for details. You may not use this file except in compliance with the License. | ||
| 7 | +# THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED, INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE. | ||
| 8 | +# See LICENSE in the root of the software repository for the full text of the License. | ||
| 9 | +# ---------------------------------------------------------------------------- | ||
| 10 | + | ||
| 11 | +set_source_files_properties(fp8_e4m3_quant_matmul.cpp PROPERTIES LANGUAGE ASC) | ||
| 12 | +catlass_example_add_executable(ascend950_fp8_e4m3_quant_matmul mix fp8_e4m3_quant_matmul.cpp) | ||
| 13 | +target_compile_definitions(ascend950_fp8_e4m3_quant_matmul PRIVATE L2_CACHE_HINT) | ||
| @@ -0,0 +1,52 @@ | |||
| 1 | +# Ascend950 FP8 E4M3 Quant Matmul Example Readme | ||
| 2 | + | ||
| 3 | +> **注意**:本样例位于 `experimental/` 目录下,如需编译运行,请先将样例目录拷贝至 `examples/` 下,并在 `examples/CMakeLists.txt` 中添加样例名称 `ascend950_fp8_e4m3_quant_matmul`。 | ||
| 4 | + | ||
这个样例只支持e4m3的话,建议参考29样例命名 ![]() ![]() | |||
| 5 | +## 功能介绍 | ||
| 6 | + | ||
| 7 | +- 演示 Ascend950 上的 **FP8 矩阵乘 + epilogue per-token/per-channel 量化**。 | ||
| 8 | +- 计算:`out = (A @ B) * perTokenScale * perChannelScale`。 | ||
| 9 | +- A、B 元素类型为 `float8_e4m3_t`,scale 为 `float8_e4m3_t`,输出为 FP32。 | ||
| 10 | +- 默认布局为 A `RowMajor`、B `RowMajor`,与 `gen_data.py` 生成的数据一致。 | ||
| 11 | + | ||
| 12 | +## 代码组织 | ||
| 13 | +```text | ||
| 14 | +experimental | ||
| 15 | +└── matmul | ||
| 16 | + └── ascend950_fp8_e4m3_quant_matmul | ||
| 17 | + ├── CMakeLists.txt | ||
| 18 | + ├── README.md | ||
| 19 | + ├── gen_data.py | ||
| 20 | + ├── fp8_e4m3_quant_matmul.cpp | ||
| 21 | + └── test_86_ascend950_fp8_e4m3_quant_matmul.py | ||
| 22 | +``` | ||
| 23 | + | ||
| 24 | +## 使用示例 | ||
| 25 | +- 获取代码之后编译相应的算子可执行文件,可参考 [quickstart](../../../docs/zh/1_Practice/01_quick_start.md#编译执行)。本用例为 Ascend950(3510)算子,编译时需加 `-DCATLASS_ARCH=3510`。 | ||
| 26 | +- 执行算子 | ||
| 27 | +``` | ||
| 28 | +# 编译指定用例 | ||
| 29 | +bash scripts/build.sh ascend950_fp8_e4m3_quant_matmul -DCATLASS_ARCH=3510 | ||
| 30 | +# 生成测试样例(在 examples/ascend950_fp8_e4m3_quant_matmul/data 下生成 input/ 与 golden/) | ||
| 31 | +python3 examples/ascend950_fp8_e4m3_quant_matmul/gen_data.py 256 256 128 | ||
| 32 | +# 可选:--data-root <DIR> 指定在 DIR/data/ 下生成(默认在脚本所在目录下生成) | ||
| 33 | +# 输入参数分别对应 m, n, k | ||
| 34 | +# 在 output/bin 中执行,以匹配示例读取数据的相对路径 | ||
| 35 | +cd output/bin | ||
| 36 | +./ascend950_fp8_e4m3_quant_matmul 256 256 128 0 | ||
| 37 | +# 可执行文件名 |矩阵m轴|n轴|k轴|Device ID | ||
| 38 | +# Device ID 可选,默认为 0 | ||
| 39 | +``` | ||
| 40 | +执行结果如下,说明精度比对成功。 | ||
| 41 | +``` | ||
| 42 | +Compare success. | ||
| 43 | +``` | ||
| 44 | + | ||
| 45 | +## optest 测试 | ||
| 46 | + | ||
| 47 | +先按 [optest 说明](../../../tests/optest/README.md) 构建并安装 `torch_catlass`,然后在仓库根目录执行迁移后的测试件: | ||
| 48 | + | ||
| 49 | +```bash | ||
| 50 | +PYTHONPATH="$PWD/tests/optest/tests${PYTHONPATH:+:$PYTHONPATH}" \ | ||
| 51 | +python3 -m pytest experimental/matmul/ascend950_fp8_e4m3_quant_matmul/test_86_ascend950_fp8_e4m3_quant_matmul.py -v | ||
| 52 | +``` | ||
| @@ -0,0 +1,389 @@ | |||
| 1 | +/** | ||
| 2 | + * This program is free software, you can redistribute it and/or modify. | ||
| 3 | + * Copyright (c) 2026 Huawei Technologies Co., Ltd. | ||
| 4 | + * This file is a part of the CANN Open Software. | ||
| 5 | + * Licensed under CANN Open Software License Agreement Version 2.0 (the "License"). | ||
| 6 | + * Please refer to the License for details. You may not use this file except in compliance with the License. | ||
| 7 | + * THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED, INCLUDING | ||
| 8 | + * BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE. See LICENSE in the root of | ||
| 9 | + * the software repository for the full text of the License. | ||
| 10 | + */ | ||
| 11 | + | ||
| 12 | +/////////////////////////////////////////////////////////////////////////////// | ||
| 13 | +/// Example 81: Ascend950 FP8 Per-Token Per-Channel Quantized Matrix Multiplication | ||
| 14 | +/// with Epilogue-based Quantization | ||
| 15 | +/// | ||
| 16 | +/// Computes: out = (A @ B) * perTokenScale * perChannelScale | ||
| 17 | +/// | ||
| 18 | +/// Where: | ||
| 19 | +/// - A: float8_e4m3_t, shape (M, K) | ||
| 20 | +/// - B: float8_e4m3_t, shape (K, N) | ||
| 21 | +/// - perTokenScale: float8_e4m3_t, shape (M,) | ||
| 22 | +/// - perChannelScale: float8_e4m3_t, shape (N,) | ||
| 23 | +/// - out: float, shape (M, N) | ||
| 24 | +/// | ||
| 25 | +/// Architecture: Ascend950, epilogue-based quantization | ||
| 26 | +/// AIC core: reads A and B from GM, performs matmul, writes intermediate C to shared UB | ||
| 27 | +/// AIV core: reads C from shared UB, applies per-token and per-channel scales, writes output to GM | ||
| 28 | +/// | ||
| 29 | +/// This example uses raw element types and layout tags instead of GemmType wrappers, | ||
| 30 | +/// following the TLA pattern. | ||
| 31 | +/////////////////////////////////////////////////////////////////////////////// | ||
| 32 | + | ||
| 33 | + | ||
| 34 | + | ||
| 35 | + | ||
| 36 | + | ||
| 37 | + | ||
| 38 | + | ||
| 39 | + | ||
| 40 | + | ||
| 41 | + | ||
| 42 | + | ||
| 43 | + | ||
| 44 | + | ||
| 45 | + | ||
| 46 | + | ||
| 47 | + | ||
| 48 | + | ||
| 49 | + | ||
| 50 | + | ||
| 51 | + | ||
| 52 | + | ||
| 53 | + | ||
| 54 | + | ||
| 55 | + | ||
| 56 | + | ||
| 57 | + | ||
| 58 | + | ||
| 59 | +using namespace Catlass; | ||
| 60 | +using namespace tla; | ||
| 61 | +using Options = GemmOptions; | ||
| 62 | + | ||
| 63 | +static const std::string kDataRoot = "../../examples/ascend950_fp8_e4m3_quant_matmul/data"; | ||
| 64 | + | ||
| 65 | +/////////////////////////////////////////////////////////////////////////////// | ||
| 66 | +/// Run: Main execution function | ||
| 67 | +/////////////////////////////////////////////////////////////////////////////// | ||
| 68 | +static int Run(const Options& options) | ||
| 69 | +{ | ||
| 70 | + if (options.problemShape.m() == 0 || options.problemShape.n() == 0 || options.problemShape.k() == 0) { | ||
| 71 | + std::cerr << "M, N and K must be positive." << std::endl; | ||
| 72 | + return EXIT_FAILURE; | ||
| 73 | + } | ||
| 74 | + aclrtStream stream{nullptr}; | ||
| 75 | + | ||
| 76 | + ACL_CHECK(aclInit(nullptr)); | ||
| 77 | + ACL_CHECK(aclrtSetDevice(options.deviceId)); | ||
| 78 | + ACL_CHECK(aclrtCreateStream(&stream)); | ||
| 79 | + | ||
| 80 | + uint32_t m = options.problemShape.m(); | ||
| 81 | + uint32_t n = options.problemShape.n(); | ||
| 82 | + uint32_t k = options.problemShape.k(); | ||
| 83 | + | ||
| 84 | + // ======================================== | ||
| 85 | + // Step 1: Type definitions | ||
| 86 | + // ======================================== | ||
| 87 | + | ||
| 88 | + using ArchTag = Arch::Ascend950; | ||
| 89 | + | ||
| 90 | + // Input data types | ||
| 91 | + using ElementA = float8_e4m3_t; | ||
| 92 | + using ElementB = float8_e4m3_t; | ||
| 93 | + using ElementC = float; // Intermediate result type (matmul accumulator) | ||
| 94 | + using ElementD = float; // Output type | ||
| 95 | + | ||
| 96 | + // Scale types | ||
| 97 | + using ElementPerTokenScale = float8_e4m3_t; | ||
| 98 | + using ElementPerChannelScale = float8_e4m3_t; | ||
| 99 | + | ||
| 100 | + // Layout tags (TLA layout types for Ascend950) | ||
| 101 | + using LayoutTagA = layout::RowMajor; | ||
| 102 | + using LayoutTagB = layout::RowMajor; | ||
| 103 | + using LayoutTagC = layout::RowMajor; | ||
| 104 | + using LayoutTagD = layout::RowMajor; | ||
| 105 | + using LayoutTagScale = layout::VectorLayout; | ||
| 106 | + using LayoutTagPerTokenScale = layout::VectorLayout; | ||
| 107 | + | ||
| 108 | + LayoutTagA tagA{m, k}; | ||
| 109 | + LayoutTagB tagB{k, n}; | ||
| 110 | + LayoutTagC tagC{m, n}; | ||
| 111 | + | ||
| 112 | + LayoutTagD tagD{m, n}; | ||
| 113 | + auto layoutA = MakeLayoutFromTag(tagA); | ||
| 114 | + auto layoutB = MakeLayoutFromTag(tagB); | ||
| 115 | + | ||
| 116 | + LayoutTagD layoutD{m, n}; | ||
| 117 | + LayoutTagScale layoutPerChannelScale{n}; | ||
| 118 | + LayoutTagPerTokenScale layoutPerTokenScale{m}; | ||
| 119 | + | ||
| 120 | + // ======================================== | ||
| 121 | + // Step 2: Tile shapes | ||
| 122 | + // ======================================== | ||
| 123 | + | ||
| 124 | + // L1/L0 tile shapes optimized for Ascend950 | ||
| 125 | + using L1TileShape = Shape<Int<256>, Int<256>, Int<512>>; | ||
| 126 | + using L0TileShape = Shape<Int<256>, Int<256>, Int<64>>; | ||
| 127 | + | ||
| 128 | + // Epilogue tile shape for UB operations | ||
| 129 | + using EpilogueTileShape = MatrixShape<128, 256>; | ||
| 130 | + | ||
| 131 | + // ======================================== | ||
| 132 | + // Step 3: Tile compute operations for epilogue | ||
| 133 | + // ======================================== | ||
| 134 | + | ||
| 135 | + // Compute type for epilogue | ||
| 136 | + using ComputeType = Gemm::GemmType<float, LayoutTagC>; | ||
| 137 | + | ||
| 138 | + // Broadcast ops for per-token scale (broadcast along columns) | ||
| 139 | + using TileBroadcastOneBlkOp = Epilogue::Tile::TileBroadcastOneBlk<ArchTag, ComputeType, EpilogueTileShape::ROW>; | ||
| 140 | + using TileOneBlkColumnBroadcastMulOp = | ||
| 141 | + Epilogue::Tile::TileOneBlkColumnBroadcastMul<ArchTag, ComputeType, EpilogueTileShape>; | ||
| 142 | + | ||
| 143 | + // Broadcast ops for per-channel scale (broadcast along rows) | ||
| 144 | + using TileRowBroadcastMulOp = Epilogue::Tile::TileRowBroadcastMul<ArchTag, ComputeType, EpilogueTileShape>; | ||
| 145 | + | ||
| 146 | + // ======================================== | ||
| 147 | + // Step 4: Tile copy operations for epilogue | ||
| 148 | + // ======================================== | ||
| 149 | + | ||
| 150 | + // GemmType wrappers for copy operations | ||
| 151 | + using CType = Gemm::GemmType<ElementC, LayoutTagC>; | ||
| 152 | + constexpr uint32_t ubStages = 1; | ||
| 153 | + | ||
| 154 | + using EpilogueDispatchPolicy = Epilogue::EpilogueAscend950Fp8PerTokenPerChannelDequant<ubStages>; | ||
| 155 | + using ScaleType = Gemm::GemmType<ElementPerChannelScale, layout::VectorLayout>; | ||
| 156 | + using PerTokenScaleType = Gemm::GemmType<ElementPerTokenScale, layout::VectorLayout>; | ||
| 157 | + using DType = Gemm::GemmType<ElementD, layout::RowMajor>; | ||
| 158 | + | ||
| 159 | + using RowBroadcastMulType = Gemm::GemmType<float, layout::RowMajor>; | ||
| 160 | + using BroadcastOneBlkType = Gemm::GemmType<float, layout::RowMajor>; | ||
| 161 | + using OneBlkColumnBroadcastMulType = Gemm::GemmType<float, layout::RowMajor>; | ||
| 162 | + | ||
| 163 | + using TileRowBroadcastMul = Epilogue::Tile::TileRowBroadcastMul<ArchTag, RowBroadcastMulType, EpilogueTileShape>; | ||
| 164 | + using TileBroadcastOneBlk = | ||
| 165 | + Epilogue::Tile::TileBroadcastOneBlk<ArchTag, BroadcastOneBlkType, EpilogueTileShape::ROW>; | ||
| 166 | + using TileOneBlkColumnBroadcastMul = | ||
| 167 | + Epilogue::Tile::TileOneBlkColumnBroadcastMul<ArchTag, OneBlkColumnBroadcastMulType, EpilogueTileShape>; | ||
| 168 | + using TileCopyEpilogue = Epilogue::Tile::TileCopy<ArchTag, CType, ScaleType, PerTokenScaleType, DType>; | ||
| 169 | + using TileScheduler = Epilogue::Tile::EpilogueHorizontalTileSwizzle; | ||
| 170 | + | ||
| 171 | + using BlockEpilogue = Epilogue::Block::BlockEpilogue< | ||
| 172 | + EpilogueDispatchPolicy, CType, ScaleType, PerTokenScaleType, DType, TileRowBroadcastMul, TileBroadcastOneBlk, | ||
| 173 | + TileOneBlkColumnBroadcastMul, TileCopyEpilogue, TileScheduler>; | ||
| 174 | + | ||
| 175 | + // ======================================== | ||
| 176 | + // Step 6: Block MMAD (AIC core) | ||
| 177 | + // ======================================== | ||
| 178 | + | ||
| 179 | + // Tile copy for matmul | ||
| 180 | + using TileCopy = Gemm::Tile::PackedTileCopyTlaToUB< | ||
| 181 | + ArchTag, ElementA, LayoutTagA, ElementB, LayoutTagB, ElementC, LayoutTagC, void, | ||
| 182 | + Gemm::Tile::CopyL0CToUBMode::SPLIT_M>; | ||
| 183 | + | ||
| 184 | + constexpr bool enableUnitFlag = true; | ||
| 185 | + constexpr bool useHF32 = false; | ||
| 186 | + using DispatchPolicy = Gemm::MmadPingpongPreLoad<ArchTag, enableUnitFlag, useHF32>; | ||
| 187 | + | ||
| 188 | + // Block MMAD | ||
| 189 | + using BlockMmadTla = Gemm::Block::BlockMmadTla< | ||
| 190 | + DispatchPolicy, L1TileShape, L0TileShape, ElementA, ElementB, ElementC, void, TileCopy>; | ||
| 191 | + | ||
| 192 | + // ======================================== | ||
| 193 | + // Step 7: Kernel | ||
| 194 | + // ======================================== | ||
| 195 | + | ||
| 196 | + using BlockScheduler = Gemm::Block::GemmIdentityBlockSwizzle<3, 0>; | ||
| 197 | + using GemmKernel = | ||
| 198 | + Gemm::Kernel::MatmulPerTokenPerChannelEpilogueTla<BlockMmadTla, BlockEpilogue, BlockScheduler, ubStages>; | ||
| 199 | + using GemmAdapter = Gemm::Device::DeviceGemm<GemmKernel>; | ||
| 200 | + | ||
| 201 | + // ======================================== | ||
| 202 | + // Step 8: Memory allocation | ||
| 203 | + // ======================================== | ||
| 204 | + | ||
| 205 | + size_t lenInputA = static_cast<size_t>(m) * k; | ||
| 206 | + size_t lenInputB = static_cast<size_t>(k) * n; | ||
| 207 | + size_t lenPerTokenScale = static_cast<size_t>(m); | ||
| 208 | + size_t lenPerChannelScale = static_cast<size_t>(n); | ||
| 209 | + size_t lenD = static_cast<size_t>(m) * n; | ||
| 210 | + | ||
| 211 | + size_t sizeInputA = lenInputA * sizeof(ElementA); | ||
| 212 | + size_t sizeInputB = lenInputB * sizeof(ElementB); | ||
| 213 | + size_t sizePerTokenScale = lenPerTokenScale * sizeof(ElementPerTokenScale); | ||
| 214 | + size_t sizePerChannelScale = lenPerChannelScale * sizeof(ElementPerChannelScale); | ||
| 215 | + size_t sizeD = lenD * sizeof(ElementD); | ||
| 216 | + | ||
| 217 | + // Host memory | ||
| 218 | + std::vector<uint8_t> hostA(lenInputA * sizeof(ElementA)); | ||
| 219 | + std::vector<uint8_t> hostB(lenInputB * sizeof(ElementB)); | ||
| 220 | + std::vector<uint8_t> hostPerTokenScale(lenPerTokenScale * sizeof(ElementPerTokenScale)); | ||
| 221 | + std::vector<uint8_t> hostPerChannelScale(lenPerChannelScale * sizeof(ElementPerChannelScale)); | ||
| 222 | + | ||
| 223 | + // Read input data from files | ||
| 224 | + const auto releaseAclEarly = [&]() { | ||
| 225 | + ACL_CHECK(aclrtDestroyStream(stream)); | ||
| 226 | + ACL_CHECK(aclrtResetDevice(options.deviceId)); | ||
| 227 | + ACL_CHECK(aclFinalize()); | ||
| 228 | + }; | ||
| 229 | + if (!ReadFile(kDataRoot + "/input/a_8.bin", hostA.data(), sizeInputA)) { | ||
| 230 | + releaseAclEarly(); | ||
| 231 | + return EXIT_FAILURE; | ||
| 232 | + } | ||
| 233 | + if (!ReadFile(kDataRoot + "/input/b_8.bin", hostB.data(), sizeInputB)) { | ||
| 234 | + releaseAclEarly(); | ||
| 235 | + return EXIT_FAILURE; | ||
| 236 | + } | ||
| 237 | + if (!ReadFile(kDataRoot + "/input/a_scale.bin", hostPerTokenScale.data(), sizePerTokenScale)) { | ||
| 238 | + releaseAclEarly(); | ||
| 239 | + return EXIT_FAILURE; | ||
| 240 | + } | ||
| 241 | + if (!ReadFile(kDataRoot + "/input/b_scale.bin", hostPerChannelScale.data(), sizePerChannelScale)) { | ||
| 242 | + releaseAclEarly(); | ||
| 243 | + return EXIT_FAILURE; | ||
| 244 | + } | ||
| 245 | + | ||
| 246 | + // Device memory | ||
| 247 | + uint8_t* deviceA{nullptr}; | ||
| 248 | + ACL_CHECK(aclrtMalloc(reinterpret_cast<void**>(&deviceA), sizeInputA, ACL_MEM_MALLOC_HUGE_FIRST)); | ||
| 249 | + ACL_CHECK(aclrtMemcpy(deviceA, sizeInputA, hostA.data(), sizeInputA, ACL_MEMCPY_HOST_TO_DEVICE)); | ||
| 250 | + | ||
| 251 | + uint8_t* deviceB{nullptr}; | ||
| 252 | + ACL_CHECK(aclrtMalloc(reinterpret_cast<void**>(&deviceB), sizeInputB, ACL_MEM_MALLOC_HUGE_FIRST)); | ||
| 253 | + ACL_CHECK(aclrtMemcpy(deviceB, sizeInputB, hostB.data(), sizeInputB, ACL_MEMCPY_HOST_TO_DEVICE)); | ||
| 254 | + | ||
| 255 | + uint8_t* devicePerTokenScale{nullptr}; | ||
| 256 | + ACL_CHECK( | ||
| 257 | + aclrtMalloc(reinterpret_cast<void**>(&devicePerTokenScale), sizePerTokenScale, ACL_MEM_MALLOC_HUGE_FIRST)); | ||
| 258 | + ACL_CHECK(aclrtMemcpy( | ||
| 259 | + devicePerTokenScale, sizePerTokenScale, hostPerTokenScale.data(), sizePerTokenScale, | ||
| 260 | + ACL_MEMCPY_HOST_TO_DEVICE)); | ||
| 261 | + | ||
| 262 | + uint8_t* devicePerChannelScale{nullptr}; | ||
| 263 | + ACL_CHECK( | ||
| 264 | + aclrtMalloc(reinterpret_cast<void**>(&devicePerChannelScale), sizePerChannelScale, ACL_MEM_MALLOC_HUGE_FIRST)); | ||
| 265 | + ACL_CHECK(aclrtMemcpy( | ||
| 266 | + devicePerChannelScale, sizePerChannelScale, hostPerChannelScale.data(), sizePerChannelScale, | ||
| 267 | + ACL_MEMCPY_HOST_TO_DEVICE)); | ||
| 268 | + | ||
| 269 | + uint8_t* deviceD{nullptr}; | ||
| 270 | + ACL_CHECK(aclrtMalloc(reinterpret_cast<void**>(&deviceD), sizeD, ACL_MEM_MALLOC_HUGE_FIRST)); | ||
| 271 | + | ||
| 272 | + // ======================================== | ||
| 273 | + // Step 9: Kernel configuration and execution | ||
| 274 | + // ======================================== | ||
| 275 | + | ||
| 276 | + auto aicCoreNum = platform_ascendc::PlatformAscendCManager::GetInstance()->GetCoreNumAic(); | ||
| 277 | + | ||
| 278 | + uint32_t taskNum = CeilDiv(m, static_cast<uint32_t>(tla::get<0>(L1TileShape{}))) * | ||
| 279 | + CeilDiv(n, static_cast<uint32_t>(tla::get<1>(L1TileShape{}))); | ||
| 280 | + | ||
| 281 | + std::cout << "taskNum: " << taskNum << std::endl; | ||
| 282 | + uint32_t aicCoreUsed = min(aicCoreNum, taskNum); | ||
| 283 | + | ||
| 284 | + // Set up kernel arguments (using TLA layouts defined earlier) | ||
| 285 | + typename GemmKernel::Arguments arguments{ | ||
| 286 | + options.problemShape, | ||
| 287 | + deviceA, | ||
| 288 | + layoutA, | ||
| 289 | + deviceB, | ||
| 290 | + layoutB, | ||
| 291 | + devicePerTokenScale, | ||
| 292 | + layoutPerTokenScale, | ||
| 293 | + devicePerChannelScale, | ||
| 294 | + layoutPerChannelScale, | ||
| 295 | + deviceD, | ||
| 296 | + layoutD, | ||
| 297 | + aicCoreUsed}; | ||
| 298 | + | ||
| 299 | + GemmAdapter gemmOp; | ||
| 300 | + | ||
| 301 | + // Check if kernel can be implemented | ||
| 302 | + if (gemmOp.CanImplement(arguments) == Status::kInvalid) { | ||
| 303 | + std::cerr << "Gemm op cannot be implemented. Please check shape requirements." << std::endl; | ||
| 304 | + ACL_CHECK(aclrtFree(deviceA)); | ||
| 305 | + ACL_CHECK(aclrtFree(deviceB)); | ||
| 306 | + ACL_CHECK(aclrtFree(devicePerTokenScale)); | ||
| 307 | + ACL_CHECK(aclrtFree(devicePerChannelScale)); | ||
| 308 | + ACL_CHECK(aclrtFree(deviceD)); | ||
| 309 | + releaseAclEarly(); | ||
| 310 | + return EXIT_FAILURE; | ||
| 311 | + } | ||
| 312 | + | ||
| 313 | + // Allocate workspace for intermediate C | ||
| 314 | + size_t sizeWorkspace = gemmOp.GetWorkspaceSize(arguments); | ||
| 315 | + uint8_t* deviceWorkspace{nullptr}; | ||
| 316 | + if (sizeWorkspace > 0) { | ||
| 317 | + ACL_CHECK(aclrtMalloc(reinterpret_cast<void**>(&deviceWorkspace), sizeWorkspace, ACL_MEM_MALLOC_HUGE_FIRST)); | ||
| 318 | + } | ||
| 319 | + | ||
| 320 | + // Initialize and run | ||
| 321 | + gemmOp.Initialize(arguments, deviceWorkspace); | ||
| 322 | + gemmOp(stream, aicCoreUsed); | ||
| 323 | + | ||
| 324 | + ACL_CHECK(aclrtSynchronizeStream(stream)); | ||
| 325 | + | ||
| 326 | + // ======================================== | ||
| 327 | + // Step 10: Result readback and save | ||
| 328 | + // ======================================== | ||
| 329 | + | ||
| 330 | + std::vector<float> hostD(lenD); | ||
| 331 | + ACL_CHECK(aclrtMemcpy(hostD.data(), sizeD, deviceD, sizeD, ACL_MEMCPY_DEVICE_TO_HOST)); | ||
| 332 | + | ||
| 333 | + std::vector<float> hostGolden(lenD); | ||
| 334 | + if (!ReadFile(kDataRoot + "/golden/expected_data.bin", hostGolden.data(), sizeD)) { | ||
| 335 | + ACL_CHECK(aclrtFree(deviceA)); | ||
| 336 | + ACL_CHECK(aclrtFree(deviceB)); | ||
| 337 | + ACL_CHECK(aclrtFree(devicePerTokenScale)); | ||
| 338 | + ACL_CHECK(aclrtFree(devicePerChannelScale)); | ||
| 339 | + ACL_CHECK(aclrtFree(deviceD)); | ||
| 340 | + if (deviceWorkspace != nullptr) { | ||
| 341 | + ACL_CHECK(aclrtFree(deviceWorkspace)); | ||
| 342 | + } | ||
| 343 | + releaseAclEarly(); | ||
| 344 | + return EXIT_FAILURE; | ||
| 345 | + } | ||
| 346 | + | ||
| 347 | + std::vector<uint64_t> errorIndices = golden::CompareData(hostD, hostGolden, k); | ||
| 348 | + if (errorIndices.empty()) { | ||
| 349 | + std::cout << "Compare success." << std::endl; | ||
| 350 | + } else { | ||
| 351 | + std::cerr << "Compare failed. Error count: " << errorIndices.size() << std::endl; | ||
| 352 | + } | ||
| 353 | + | ||
| 354 | + std::cout << "Ascend950 FP8 Per-Token Per-Channel Quantized Matmul (Epilogue) completed." << std::endl; | ||
| 355 | + std::cout << "Problem shape: M=" << m << " N=" << n << " K=" << k << std::endl; | ||
| 356 | + std::cout << "Workspace size: " << sizeWorkspace << " bytes" << std::endl; | ||
| 357 | + | ||
| 358 | + // ======================================== | ||
| 359 | + // Step 11: Cleanup | ||
| 360 | + // ======================================== | ||
| 361 | + | ||
| 362 | + ACL_CHECK(aclrtFree(deviceA)); | ||
| 363 | + ACL_CHECK(aclrtFree(deviceB)); | ||
| 364 | + ACL_CHECK(aclrtFree(devicePerTokenScale)); | ||
| 365 | + ACL_CHECK(aclrtFree(devicePerChannelScale)); | ||
| 366 | + ACL_CHECK(aclrtFree(deviceD)); | ||
| 367 | + if (deviceWorkspace != nullptr) { | ||
| 368 | + ACL_CHECK(aclrtFree(deviceWorkspace)); | ||
| 369 | + } | ||
| 370 | + | ||
| 371 | + ACL_CHECK(aclrtDestroyStream(stream)); | ||
| 372 | + ACL_CHECK(aclrtResetDevice(options.deviceId)); | ||
| 373 | + ACL_CHECK(aclFinalize()); | ||
| 374 | + return errorIndices.empty() ? EXIT_SUCCESS : EXIT_FAILURE; | ||
| 375 | +} | ||
| 376 | + | ||
| 377 | +/////////////////////////////////////////////////////////////////////////////// | ||
| 378 | +/// Main | ||
| 379 | +/////////////////////////////////////////////////////////////////////////////// | ||
| 380 | +int main(int argc, const char** argv) | ||
| 381 | +{ | ||
| 382 | + Options options; | ||
| 383 | + if (options.Parse(argc, argv) != 0) { | ||
| 384 | + return -1; | ||
| 385 | + } | ||
| 386 | + | ||
| 387 | + std::cout << "Running Ascend950 FP8 Per-Token Per-Channel Quantized Matmul (Epilogue)..." << std::endl; | ||
| 388 | + return Run(options); | ||
| 389 | +} | ||
| @@ -0,0 +1,89 @@ | |||
| 1 | +#!/usr/bin/env python3 | ||
| 2 | +# -*- coding: utf-8 -*- | ||
| 3 | +# ---------------------------------------------------------------------------- | ||
| 4 | +# This program is free software, you can redistribute it and/or modify. | ||
| 5 | +# Copyright (c) 2026 Huawei Technologies Co., Ltd. | ||
| 6 | +# This file is a part of the CANN Open Software. | ||
| 7 | +# Licensed under CANN Open Software License Agreement Version 2.0 (the "License"). | ||
| 8 | +# Please refer to the License for details. You may not use this file except in compliance with the License. | ||
| 9 | +# THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED, INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE. | ||
| 10 | +# See LICENSE in the root of the software repository for the full text of the License. | ||
| 11 | +# ---------------------------------------------------------------------------- | ||
| 12 | + | ||
| 13 | +import argparse | ||
| 14 | +import os | ||
| 15 | +from typing import Optional | ||
| 16 | + | ||
| 17 | +import numpy as np | ||
| 18 | +import torch | ||
| 19 | +from ml_dtypes import float8_e4m3fn | ||
| 20 | + | ||
| 21 | +_SCRIPT_DIR = os.path.dirname(os.path.abspath(__file__)) | ||
| 22 | +_FP8_MAX = 448.0 | ||
| 23 | + | ||
| 24 | + | ||
| 25 | +def _resolve_workspace(data_root_cli: Optional[str]) -> str: | ||
| 26 | + if data_root_cli is not None: | ||
| 27 | + root = data_root_cli.strip() | ||
| 28 | + if root: | ||
| 29 | + return os.path.abspath(os.path.expanduser(root)) | ||
| 30 | + return _SCRIPT_DIR | ||
| 31 | + | ||
| 32 | + | ||
| 33 | +def gen_data(m: int, n: int, k: int, workspace: str) -> None: | ||
| 34 | + data_dir = os.path.join(workspace, "data") | ||
| 35 | + input_dir = os.path.join(data_dir, "input") | ||
| 36 | + golden_dir = os.path.join(data_dir, "golden") | ||
| 37 | + os.makedirs(input_dir, exist_ok=True) | ||
| 38 | + os.makedirs(golden_dir, exist_ok=True) | ||
| 39 | + | ||
| 40 | + torch.manual_seed(0) | ||
| 41 | + a_fp32 = torch.randn((m, k), dtype=torch.float32) * 2 | ||
| 42 | + b_fp32 = torch.randn((k, n), dtype=torch.float32) * 2 | ||
| 43 | + | ||
| 44 | + a_max = a_fp32.abs().max(dim=1).values | ||
| 45 | + b_max = b_fp32.abs().max(dim=0).values | ||
| 46 | + a_max[a_max == 0] = 1.0 | ||
| 47 | + b_max[b_max == 0] = 1.0 | ||
| 48 | + | ||
| 49 | + a_scale_fp32 = a_max / _FP8_MAX | ||
| 50 | + b_scale_fp32 = b_max / _FP8_MAX | ||
| 51 | + | ||
| 52 | + a_fp8 = (a_fp32 / a_scale_fp32.view(-1, 1)).numpy().astype(float8_e4m3fn) | ||
| 53 | + b_fp8 = (b_fp32 / b_scale_fp32.view(1, -1)).numpy().astype(float8_e4m3fn) | ||
| 54 | + a_scale_fp8 = a_scale_fp32.numpy().astype(float8_e4m3fn) | ||
| 55 | + b_scale_fp8 = b_scale_fp32.numpy().astype(float8_e4m3fn) | ||
| 56 | + | ||
| 57 | + a_fp8.view(np.int8).tofile(os.path.join(input_dir, "a_8.bin")) | ||
| 58 | + b_fp8.view(np.int8).tofile(os.path.join(input_dir, "b_8.bin")) | ||
| 59 | + a_scale_fp8.view(np.int8).tofile(os.path.join(input_dir, "a_scale.bin")) | ||
| 60 | + b_scale_fp8.view(np.int8).tofile(os.path.join(input_dir, "b_scale.bin")) | ||
| 61 | + | ||
| 62 | + a_val = a_fp8.astype(np.float32) | ||
| 63 | + b_val = b_fp8.astype(np.float32) | ||
| 64 | + per_token = a_scale_fp8.astype(np.float32) | ||
| 65 | + per_channel = b_scale_fp8.astype(np.float32) | ||
| 66 | + | ||
| 67 | + c_fp32 = a_val @ b_val | ||
| 68 | + golden = c_fp32 * per_token[:, np.newaxis] * per_channel[np.newaxis, :] | ||
| 69 | + golden.astype(np.float32).tofile(os.path.join(golden_dir, "expected_data.bin")) | ||
| 70 | + | ||
| 71 | + | ||
| 72 | +if __name__ == "__main__": | ||
| 73 | + parser = argparse.ArgumentParser( | ||
| 74 | + description="Generate FP8 per-token/per-channel inputs and FP32 golden under " | ||
| 75 | + "<data-root>/data/.", | ||
| 76 | + ) | ||
| 77 | + parser.add_argument( | ||
| 78 | + "--data-root", | ||
| 79 | + default=None, | ||
| 80 | + metavar="DIR", | ||
| 81 | + help="Directory under which data/input and data/golden are created. " | ||
| 82 | + "Default: this script's directory.", | ||
| 83 | + ) | ||
| 84 | + parser.add_argument("m", type=int) | ||
| 85 | + parser.add_argument("n", type=int) | ||
| 86 | + parser.add_argument("k", type=int) | ||
| 87 | + args = parser.parse_args() | ||
| 88 | + workspace = _resolve_workspace(args.data_root) | ||
| 89 | + gen_data(args.m, args.n, args.k, workspace) | ||
Aexperimental/matmul/ascend950_fp8_e4m3_quant_matmul/test_86_ascend950_fp8_e4m3_quant_matmul.py+57-0
| @@ -0,0 +1,57 @@ | |||
| 1 | +# This program is free software, you can redistribute it and/or modify. | ||
| 2 | +# Copyright (c) 2026 Huawei Technologies Co., Ltd. | ||
| 3 | +# This file is a part of the CANN Open Software. | ||
| 4 | +# Licensed under CANN Open Software License Agreement Version 2.0 (the "License"). | ||
| 5 | +# Please refer to the License for details. You may not use this file except in compliance with the License. | ||
| 6 | +# THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED, INCLUDING | ||
| 7 | +# BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE. See LICENSE in the root of | ||
| 8 | +# the software repository for the full text of the License. | ||
| 9 | + | ||
| 10 | +import pytest | ||
| 11 | +import torch | ||
| 12 | +import torch_catlass | ||
| 13 | + | ||
| 14 | +from common import only_on_3510 | ||
| 15 | +from mx_golden import prepare_fp8_e4m3_quant_matmul_inputs | ||
| 16 | +from quant_matmul_test_utils import QUANT_MATMUL_SHAPES, check_quant_result, signed_scales, with_storage_offset | ||
| 17 | + | ||
| 18 | + | ||
| 19 | + | ||
| 20 | + | ||
| 21 | + | ||
| 22 | +def test_ascend950_fp8_e4m3_quant_matmul(shape, seed, record_property): | ||
| 23 | + """Compare example 86 against the quantized CPU reference.""" | ||
| 24 | + m, n, k = shape | ||
| 25 | + a, b, per_token, per_channel, expected = prepare_fp8_e4m3_quant_matmul_inputs( | ||
| 26 | + m, n, k, device="npu", seed=seed | ||
| 27 | + ) | ||
| 28 | + if seed % 2: | ||
| 29 | + per_token, per_channel, expected = signed_scales(per_token, per_channel, expected) | ||
| 30 | + | ||
| 31 | + result = torch_catlass.ascend950_fp8_e4m3_quant_matmul(a, b, per_token, per_channel) | ||
| 32 | + | ||
| 33 | + check_quant_result(result, expected, record_property) | ||
| 34 | + | ||
| 35 | + | ||
| 36 | + | ||
| 37 | + | ||
| 38 | +def test_fp8_quant_storage_offset(input_index, record_property): | ||
| 39 | + *inputs, expected = prepare_fp8_e4m3_quant_matmul_inputs(127, 130, 66) | ||
| 40 | + inputs[input_index] = with_storage_offset(inputs[input_index]) | ||
| 41 | + result = torch_catlass.ascend950_fp8_e4m3_quant_matmul(*inputs) | ||
| 42 | + check_quant_result(result, expected, record_property) | ||
| 43 | + | ||
| 44 | + | ||
| 45 | + | ||
| 46 | + | ||
| 47 | +def test_fp8_quant_rejects_empty_shape(axis): | ||
| 48 | + shape = [16, 16, 16] | ||
| 49 | + shape[axis] = 0 | ||
| 50 | + m, n, k = shape | ||
| 51 | + inputs = [torch.empty(s, dtype=torch.float8_e4m3fn, device="npu") for s in [(m, k), (k, n), (m,), (n,)]] | ||
| 52 | + with pytest.raises(RuntimeError, match="must be positive"): | ||
| 53 | + torch_catlass.ascend950_fp8_e4m3_quant_matmul(*inputs) | ||
| 54 | + | ||
| 55 | + | ||
| 56 | +if __name__ == "__main__": | ||
| 57 | + pytest.main([__file__, "-v", "-s"]) | ||
| @@ -56,6 +56,8 @@ class BlockEpilogue { | |||
| 56 | 56 | ||
| 57 | 57 | ||
| 58 | 58 | ||
| 59 | + | ||
| 60 | + | ||
| 59 | 61 | ||
| 60 | 62 | ||
| 61 | 63 | ||
| @@ -0,0 +1,544 @@ | |||
| 1 | +/** | ||
| 2 | + * This program is free software, you can redistribute it and/or modify. | ||
| 3 | + * Copyright (c) 2026 Huawei Technologies Co., Ltd. | ||
| 4 | + * This file is a part of the CANN Open Software. | ||
| 5 | + * Licensed under CANN Open Software License Agreement Version 2.0 (the "License"). | ||
| 6 | + * Please refer to the License for details. You may not use this file except in compliance with the License. | ||
| 7 | + * THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED, INCLUDING | ||
| 8 | + * BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE. See LICENSE in the root of | ||
| 9 | + * the software repository for the full text of the License. | ||
| 10 | + */ | ||
| 11 | + | ||
| 12 | + | ||
| 13 | + | ||
| 14 | + | ||
| 15 | + | ||
| 16 | + | ||
| 17 | + | ||
| 18 | + | ||
| 19 | + | ||
| 20 | + | ||
| 21 | + | ||
| 22 | + | ||
| 23 | + | ||
| 24 | +namespace Catlass::Epilogue::Block { | ||
| 25 | + | ||
| 26 | +template < | ||
| 27 | + uint32_t UB_STAGES_, class CType_, class ScaleType_, class PerTokenScaleType_, class DType_, | ||
| 28 | + class TileRowBroadcastMul_, class TileBroadcastOneBlk_, class TileOneBlkColumnBroadcastMul_, class TileCopy_, | ||
| 29 | + class EpilogueTileSwizzle_> | ||
| 30 | +class BlockEpilogue< | ||
| 31 | + EpilogueAscend950Fp8PerTokenPerChannelDequant<UB_STAGES_>, CType_, ScaleType_, PerTokenScaleType_, DType_, | ||
| 32 | + TileRowBroadcastMul_, TileBroadcastOneBlk_, TileOneBlkColumnBroadcastMul_, TileCopy_, EpilogueTileSwizzle_> { | ||
| 33 | +public: | ||
| 34 | + using DispatchPolicy = EpilogueAscend950Fp8PerTokenPerChannelDequant<UB_STAGES_>; | ||
| 35 | + using ArchTag = typename DispatchPolicy::ArchTag; | ||
| 36 | + static constexpr uint32_t UB_STAGES = UB_STAGES_; | ||
| 37 | + | ||
| 38 | + // Data infos | ||
| 39 | + using ElementC = typename CType_::Element; // 现在float | ||
| 40 | + using LayoutC = typename CType_::Layout; | ||
| 41 | + using ElementScale = typename ScaleType_::Element; // 现在fp8_e5m2 | ||
| 42 | + using LayoutScale = typename ScaleType_::Layout; | ||
| 43 | + using ElementPerTokenScale = typename PerTokenScaleType_::Element; // 现在fp8_e5m2 | ||
| 44 | + using LayoutPerTokenScale = typename PerTokenScaleType_::Layout; | ||
| 45 | + using ElementD = typename DType_::Element; // 现在根据输入确定是half or float | ||
| 46 | + using LayoutD = typename DType_::Layout; | ||
| 47 | + | ||
| 48 | + static_assert( | ||
| 49 | + std::is_same_v<LayoutC, layout::RowMajor> && std::is_same_v<LayoutScale, layout::VectorLayout> && | ||
| 50 | + std::is_same_v<LayoutPerTokenScale, layout::VectorLayout> && std::is_same_v<LayoutD, layout::RowMajor>, | ||
| 51 | + "The layout template parameters of BlockEpilogue are wrong"); | ||
| 52 | + | ||
| 53 | + // Tile compute ops | ||
| 54 | + using TileRowBroadcastMul = TileRowBroadcastMul_; | ||
| 55 | + using TileBroadcastOneBlk = TileBroadcastOneBlk_; | ||
| 56 | + using TileOneBlkColumnBroadcastMul = TileOneBlkColumnBroadcastMul_; | ||
| 57 | + | ||
| 58 | + // Tile copy | ||
| 59 | + using CopyGmToUbScale = typename TileCopy_::CopyGmToUbX; | ||
| 60 | + using CopyGmToUbPerTokenScale = typename TileCopy_::CopyGmToUbY; | ||
| 61 | + using CopyUbToGmD = typename TileCopy_::CopyUbToGmD; | ||
| 62 | + | ||
| 63 | + using EpilogueTileSwizzle = EpilogueTileSwizzle_; | ||
| 64 | + | ||
| 65 | + using TileShape = typename TileRowBroadcastMul::TileShape; | ||
| 66 | + | ||
| 67 | + static_assert( | ||
| 68 | + TileShape::ROW == TileBroadcastOneBlk::COMPUTE_LENGTH && | ||
| 69 | + std::is_same_v<TileShape, typename TileOneBlkColumnBroadcastMul::TileShape>, | ||
| 70 | + "TileShape must be consistent for all tile compute ops"); | ||
| 71 | + | ||
| 72 | + static_assert( | ||
| 73 | + (UB_STAGES * (TileShape::COUNT * sizeof(ElementC) + TileShape::COUNT * sizeof(half) + | ||
| 74 | + TileShape::COLUMN * sizeof(ElementScale) + TileShape::ROW * sizeof(ElementPerTokenScale)) + | ||
| 75 | + UB_STAGES * (TileShape::COLUMN + TileShape::ROW) * sizeof(float) + | ||
| 76 | + UB_STAGES * TileShape::ROW * BYTE_PER_BLK) <= ArchTag::UB_SIZE, | ||
| 77 | + "TileShape is too large to fit in UB"); | ||
| 78 | + | ||
| 79 | + struct Params { | ||
| 80 | + __gm__ ElementScale* ptrScale{nullptr}; | ||
| 81 | + LayoutScale layoutScale{}; | ||
| 82 | + __gm__ ElementPerTokenScale* ptrPerTokenScale{nullptr}; | ||
| 83 | + LayoutPerTokenScale layoutPerTokenScale{}; | ||
| 84 | + __gm__ ElementD* ptrD{nullptr}; | ||
| 85 | + LayoutD layoutD{}; | ||
| 86 | + | ||
| 87 | + CATLASS_DEVICE | ||
| 88 | + Params() {}; | ||
| 89 | + | ||
| 90 | + CATLASS_DEVICE | ||
| 91 | + Params( | ||
| 92 | + __gm__ ElementScale* ptrScale_, LayoutScale const& layoutScale_, | ||
| 93 | + __gm__ ElementPerTokenScale* ptrPerTokenScale_, LayoutPerTokenScale const& layoutPerTokenScale_, | ||
| 94 | + __gm__ ElementD* ptrD_, LayoutD const& layoutD_) | ||
| 95 | + : ptrScale(ptrScale_), | ||
| 96 | + layoutScale(layoutScale_), | ||
| 97 | + ptrPerTokenScale(ptrPerTokenScale_), | ||
| 98 | + layoutPerTokenScale(layoutPerTokenScale_), | ||
| 99 | + ptrD(ptrD_), | ||
| 100 | + layoutD(layoutD_) | ||
| 101 | + {} | ||
| 102 | + }; | ||
| 103 | + | ||
| 104 | + CATLASS_DEVICE | ||
| 105 | + BlockEpilogue(Arch::Resource<ArchTag> const& resource, Params const& params = Params{}) : params(params) | ||
| 106 | + { | ||
| 107 | + size_t ubOffset = 0; | ||
| 108 | + for (uint32_t i = 0; i < UB_STAGES; ++i) { | ||
| 109 | + ubCList[i] = resource.ubBuf.template GetBufferByByte<ElementC>(ubOffset); | ||
| 110 | + ubOffset += TileShape::COUNT * sizeof(ElementC); | ||
| 111 | + } | ||
| 112 | + | ||
| 113 | + for (uint32_t i = 0; i < UB_STAGES; ++i) { | ||
| 114 | + ubScaleList[i] = resource.ubBuf.template GetBufferByByte<ElementScale>(ubOffset); | ||
| 115 | + ubOffset += TileShape::COLUMN * sizeof(ElementScale); | ||
| 116 | + ubPerTokenScaleList[i] = resource.ubBuf.template GetBufferByByte<ElementPerTokenScale>(ubOffset); | ||
| 117 | + ubOffset += TileShape::ROW * sizeof(ElementPerTokenScale); | ||
| 118 | + if constexpr (!std::is_same_v<ElementD, float>) { | ||
| 119 | + ubDList[i] = resource.ubBuf.template GetBufferByByte<ElementD>(ubOffset); | ||
| 120 | + ubOffset += TileShape::COUNT * sizeof(ElementD); | ||
| 121 | + } | ||
| 122 | + ubPerTokenScaleFp32List[i] = resource.ubBuf.template GetBufferByByte<float>(ubOffset); | ||
| 123 | + ubOffset += TileShape::ROW * sizeof(float); | ||
| 124 | + ubPerTokenMulList[i] = ubCList[i]; | ||
| 125 | + } | ||
| 126 | + } | ||
| 127 | + | ||
| 128 | + CATLASS_DEVICE | ||
| 129 | + ~BlockEpilogue() | ||
| 130 | + {} | ||
| 131 | + | ||
| 132 | + CATLASS_DEVICE | ||
| 133 | + void UpdateParams(Params const& params_) | ||
| 134 | + { | ||
| 135 | + params = params_; | ||
| 136 | + } | ||
| 137 | + | ||
| 138 | + CATLASS_DEVICE | ||
| 139 | + void operator()( | ||
| 140 | + GemmCoord const& blockShapeMNK, GemmCoord const& blockCoordMNK, GemmCoord const& actualBlockShapeMNK, | ||
| 141 | + AscendC::GlobalTensor<ElementC> const& gmBlockC, LayoutC const& layoutBlockC, uint32_t const& stageId, | ||
| 142 | + Arch::CrossCoreFlag preCross, Arch::CrossCoreFlag afterCross, Callback&& callback = Callback{}) | ||
| 143 | + { | ||
| 144 | + if (actualBlockShapeMNK.k() == 0) { | ||
| 145 | + return; | ||
| 146 | + } | ||
| 147 | + callback(); | ||
| 148 | + | ||
| 149 | + ubListId = stageId; | ||
| 150 | + | ||
| 151 | + // Calculate the offset of the current block | ||
| 152 | + MatrixCoord blockShape = blockShapeMNK.GetCoordMN(); | ||
| 153 | + MatrixCoord blockCoord = blockCoordMNK.GetCoordMN(); | ||
| 154 | + MatrixCoord actualBlockShape = actualBlockShapeMNK.GetCoordMN(); | ||
| 155 | + MatrixCoord blockOffset = blockCoord * blockShape; | ||
| 156 | + | ||
| 157 | + AscendC::GlobalTensor<ElementScale> gmScale; | ||
| 158 | + gmScale.SetGlobalBuffer(params.ptrScale); | ||
| 159 | + AscendC::GlobalTensor<ElementPerTokenScale> gmPerTokenScale; | ||
| 160 | + gmPerTokenScale.SetGlobalBuffer(params.ptrPerTokenScale); | ||
| 161 | + AscendC::GlobalTensor<ElementD> gmD; | ||
| 162 | + gmD.SetGlobalBuffer(params.ptrD); | ||
| 163 | + | ||
| 164 | + auto ubTileStride = MakeCoord(static_cast<int64_t>(TileShape::COLUMN), 1L); | ||
| 165 | + auto tileShape = MakeCoord((actualBlockShape.row() + 1) / 2, TileShape::COLUMN); | ||
| 166 | + EpilogueTileSwizzle epilogueTileSwizzle(actualBlockShape, tileShape); | ||
| 167 | + uint32_t tileLoops = epilogueTileSwizzle.GetLoops(); | ||
| 168 | + uint32_t subblockIdx = AscendC::GetSubBlockIdx(); | ||
| 169 | + uint32_t subblockNum = AscendC::GetSubBlockNum(); | ||
| 170 | + | ||
| 171 | + AscendC::SetFlag<AscendC::HardEvent::V_MTE2>(EVENT_ID0); | ||
| 172 | + AscendC::SetFlag<AscendC::HardEvent::V_MTE2>(EVENT_ID1); | ||
| 173 | + // An idle AIV must still consume the ready flag and acknowledge the stage. | ||
| 174 | + // Otherwise the AIC waits forever when an M tail has only one row. | ||
| 175 | + if (subblockIdx >= tileLoops) { | ||
| 176 | + Arch::CrossCoreWaitFlag(preCross); | ||
| 177 | + Arch::CrossCoreSetFlag<0x2, PIPE_MTE3>(afterCross); | ||
| 178 | + } | ||
| 179 | + for (uint32_t loopIdx = subblockIdx; loopIdx < tileLoops; loopIdx += subblockNum) { | ||
| 180 | + auto tileCoord = epilogueTileSwizzle.GetTileCoord(loopIdx); | ||
| 181 | + auto actualTileShape = epilogueTileSwizzle.GetActualTileShape(tileCoord); | ||
| 182 | + auto tileOffsetInBlock = tileCoord * tileShape; | ||
| 183 | + auto tileOffset = blockOffset + tileOffsetInBlock; | ||
| 184 | + | ||
| 185 | + auto gmTileC = gmBlockC[layoutBlockC.GetOffset(tileOffsetInBlock)]; | ||
| 186 | + auto layoutGmTileC = layoutBlockC.GetTileLayout(actualTileShape); | ||
| 187 | + | ||
| 188 | + auto& ubC = ubCList[ubListId]; | ||
| 189 | + LayoutC layoutUbC{actualTileShape, ubTileStride}; | ||
| 190 | + | ||
| 191 | + auto eventId = ubListId ? EVENT_ID0 : EVENT_ID1; | ||
| 192 | + | ||
| 193 | + AscendC::WaitFlag<AscendC::HardEvent::V_MTE2>(eventId); | ||
| 194 | + | ||
| 195 | + auto scaleTileOffset = tileOffset.template GetCoordByAxis<1>(); | ||
| 196 | + auto scaleTileShape = actualTileShape.template GetCoordByAxis<1>(); | ||
| 197 | + | ||
| 198 | + auto gmTileScale = gmScale[params.layoutScale.GetOffset(scaleTileOffset)]; | ||
| 199 | + auto layoutGmTileScale = params.layoutScale.GetTileLayout(scaleTileShape); | ||
| 200 | + | ||
| 201 | + auto& ubScale = ubScaleList[ubListId]; | ||
| 202 | + auto layoutUbScale = LayoutScale::template MakeLayoutInUb<ElementScale>(scaleTileShape); | ||
| 203 | + | ||
| 204 | + copyGmToUbScale(ubScale, gmTileScale, layoutUbScale, layoutGmTileScale); | ||
| 205 | + | ||
| 206 | + auto perTokenScaleTileOffset = tileOffset.template GetCoordByAxis<0>(); | ||
| 207 | + auto perTokenScaleTileShape = actualTileShape.template GetCoordByAxis<0>(); | ||
| 208 | + | ||
| 209 | + auto gmTilePerTokenScale = gmPerTokenScale[params.layoutPerTokenScale.GetOffset(perTokenScaleTileOffset)]; | ||
| 210 | + auto layoutGmTilePerTokenScale = params.layoutPerTokenScale.GetTileLayout(perTokenScaleTileShape); | ||
| 211 | + | ||
| 212 | + auto& ubPerTokenScale = ubPerTokenScaleList[ubListId]; | ||
| 213 | + auto layoutUbPerTokenScale = | ||
| 214 | + LayoutScale::template MakeLayoutInUb<ElementPerTokenScale>(perTokenScaleTileShape); | ||
| 215 | + | ||
| 216 | + auto& ubD = ubDList[ubListId]; | ||
| 217 | + LayoutD layoutUbD{actualTileShape, ubTileStride}; | ||
| 218 | + | ||
| 219 | + copyGmToUbPerTokenScale( | ||
| 220 | + ubPerTokenScale, gmTilePerTokenScale, layoutUbPerTokenScale, layoutGmTilePerTokenScale); | ||
| 221 | + | ||
| 222 | + AscendC::SetFlag<AscendC::HardEvent::MTE2_V>(eventId); | ||
| 223 | + AscendC::WaitFlag<AscendC::HardEvent::MTE2_V>(eventId); | ||
| 224 | + | ||
| 225 | + Arch::CrossCoreWaitFlag(preCross); | ||
| 226 | + | ||
| 227 | + if constexpr (TileShape::COLUMN == 256 && TileShape::ROW == 128) { | ||
| 228 | + if constexpr (std::is_same<ElementPerTokenScale, float8_e8m0_t>::value) { | ||
| 229 | + // High-level API doesn't support float8_e8m0_t -> float, use MicroAPI instead | ||
| 230 | + __ubuf__ float8_e8m0_t* srcAddr = (__ubuf__ float8_e8m0_t*)ubPerTokenScale.GetPhyAddr(); | ||
| 231 | + __ubuf__ float* dstAddr = (__ubuf__ float*)ubPerTokenScaleFp32List[ubListId].GetPhyAddr(); | ||
| 232 | + CastFp8E8m0ToFp32(dstAddr, srcAddr, TileShape::ROW); | ||
| 233 | + | ||
| 234 | + } else if constexpr (!std::is_same<ElementPerTokenScale, float>::value) { | ||
| 235 | + AscendC::Cast( | ||
| 236 | + ubPerTokenScaleFp32List[ubListId], ubPerTokenScale, AscendC::RoundMode::CAST_NONE, | ||
| 237 | + TileShape::ROW); | ||
| 238 | + } | ||
| 239 | + AscendC::PipeBarrier<PIPE_V>(); | ||
| 240 | + MulCompute( | ||
| 241 | + (__ubuf__ float*)ubPerTokenMulList[ubListId].GetPhyAddr(), (__ubuf__ float*)ubC.GetPhyAddr(), | ||
| 242 | + (__ubuf__ ElementScale*)ubScale.GetPhyAddr(), | ||
| 243 | + (__ubuf__ float*)ubPerTokenScaleFp32List[ubListId].GetPhyAddr()); | ||
| 244 | + if constexpr (!std::is_same_v<ElementD, float>) { | ||
| 245 | + AscendC::Cast( | ||
| 246 | + ubD, ubPerTokenMulList[ubListId], AscendC::RoundMode::CAST_RINT, | ||
| 247 | + TileShape::ROW * TileShape::COLUMN); | ||
| 248 | + } | ||
| 249 | + } else { | ||
| 250 | + if constexpr (std::is_same<ElementPerTokenScale, float8_e8m0_t>::value) { | ||
| 251 | + __ubuf__ float8_e8m0_t* srcAddr = (__ubuf__ float8_e8m0_t*)ubPerTokenScale.GetPhyAddr(); | ||
| 252 | + __ubuf__ float* dstAddr = (__ubuf__ float*)ubPerTokenScaleFp32List[ubListId].GetPhyAddr(); | ||
| 253 | + CastFp8E8m0ToFp32(dstAddr, srcAddr, TileShape::ROW); | ||
| 254 | + } else if constexpr (!std::is_same<ElementPerTokenScale, float>::value) { | ||
| 255 | + AscendC::Cast( | ||
| 256 | + ubPerTokenScaleFp32List[ubListId], ubPerTokenScale, AscendC::RoundMode::CAST_NONE, | ||
| 257 | + TileShape::ROW); | ||
| 258 | + } | ||
| 259 | + AscendC::PipeBarrier<PIPE_V>(); | ||
| 260 | + MulComputeGeneric( | ||
| 261 | + (__ubuf__ float*)ubPerTokenMulList[ubListId].GetPhyAddr(), (__ubuf__ float*)ubC.GetPhyAddr(), | ||
| 262 | + (__ubuf__ ElementScale*)ubScale.GetPhyAddr(), | ||
| 263 | + (__ubuf__ float*)ubPerTokenScaleFp32List[ubListId].GetPhyAddr()); | ||
| 264 | + if constexpr (!std::is_same_v<ElementD, float>) { | ||
| 265 | + AscendC::Cast( | ||
| 266 | + ubD, ubPerTokenMulList[ubListId], AscendC::RoundMode::CAST_RINT, | ||
| 267 | + TileShape::ROW * TileShape::COLUMN); | ||
| 268 | + } | ||
| 269 | + } | ||
| 270 | + AscendC::SetFlag<AscendC::HardEvent::V_MTE3>(eventId); | ||
| 271 | + AscendC::WaitFlag<AscendC::HardEvent::V_MTE3>(eventId); | ||
| 272 | + | ||
| 273 | + auto gmTileD = gmD[params.layoutD.GetOffset(tileOffset)]; | ||
| 274 | + auto layoutGmTileD = params.layoutD.GetTileLayout(actualTileShape); | ||
| 275 | + | ||
| 276 | + if constexpr (!std::is_same_v<ElementD, float>) { | ||
| 277 | + copyUbToGmD(gmTileD, ubD, layoutGmTileD, layoutUbD); | ||
| 278 | + } else { | ||
| 279 | + copyUbToGmD(gmTileD, ubPerTokenMulList[ubListId], layoutGmTileD, layoutUbD); | ||
| 280 | + } | ||
| 281 | + Arch::CrossCoreSetFlag<0x2, PIPE_MTE3>(afterCross); | ||
| 282 | + AscendC::SetFlag<AscendC::HardEvent::V_MTE2>(eventId); | ||
| 283 | + | ||
| 284 | + ubListId = (ubListId + 1 < UB_STAGES) ? (ubListId + 1) : 0; | ||
| 285 | + } | ||
| 286 | + AscendC::WaitFlag<AscendC::HardEvent::V_MTE2>(EVENT_ID0); | ||
| 287 | + AscendC::WaitFlag<AscendC::HardEvent::V_MTE2>(EVENT_ID1); | ||
| 288 | + } | ||
| 289 | + | ||
| 290 | +private: | ||
| 291 | + Params params; | ||
| 292 | + | ||
| 293 | + AscendC::LocalTensor<ElementC> ubCList[UB_STAGES]; | ||
| 294 | + AscendC::LocalTensor<ElementScale> ubScaleList[UB_STAGES]; | ||
| 295 | + AscendC::LocalTensor<ElementPerTokenScale> ubPerTokenScaleList[UB_STAGES]; | ||
| 296 | + AscendC::LocalTensor<ElementD> ubDList[UB_STAGES]; | ||
| 297 | + AscendC::LocalTensor<float> ubPerTokenScaleFp32List[UB_STAGES]; | ||
| 298 | + AscendC::LocalTensor<float> ubPerTokenMulList[UB_STAGES]; | ||
| 299 | + | ||
| 300 | + uint32_t ubListId{0}; | ||
| 301 | + | ||
| 302 | + CopyGmToUbScale copyGmToUbScale; | ||
| 303 | + CopyGmToUbPerTokenScale copyGmToUbPerTokenScale; | ||
| 304 | + CopyUbToGmD copyUbToGmD; | ||
| 305 | + | ||
| 306 | + /// Helper function to cast scale to fp32 | ||
| 307 | + template <typename T> | ||
| 308 | + CATLASS_DEVICE void CastScaleToFp32(AscendC::LocalTensor<float>& dst, AscendC::LocalTensor<T>& src, uint32_t count) | ||
| 309 | + { | ||
| 310 | + if constexpr (std::is_same_v<T, float8_e8m0_t>) { | ||
| 311 | + CastFp8E8m0ToFp32(dst, src, count); | ||
| 312 | + } else if constexpr (std::is_same_v<T, float8_e4m3_t> || std::is_same_v<T, float8_e5m2_t>) { | ||
| 313 | + AscendC::Cast(dst, src, AscendC::RoundMode::CAST_NONE, count); | ||
| 314 | + } else if constexpr (std::is_same_v<T, float>) { | ||
| 315 | + // Already fp32, copy | ||
| 316 | + AscendC::DataCopy(dst, src, count); | ||
| 317 | + } else { | ||
| 318 | + AscendC::Cast(dst, src, AscendC::RoundMode::CAST_NONE, count); | ||
| 319 | + } | ||
| 320 | + } | ||
| 321 | + | ||
| 322 | + /// Cast float8_e8m0_t to fp32 using MicroAPI | ||
| 323 | + __simd_vf__ inline void CastFp8E8m0ToFp32( | ||
| 324 | + __ubuf__ float* dstPtr, __ubuf__ ElementPerTokenScale* srcPtr, uint32_t count) | ||
| 325 | + { | ||
| 326 | + namespace MicroAPI = AscendC::MicroAPI; | ||
| 327 | + | ||
| 328 | + static constexpr MicroAPI::CastTrait castTraitFp8ToBf16 = { | ||
| 329 | + MicroAPI::RegLayout::ZERO, MicroAPI::SatMode::UNKNOWN, MicroAPI::MaskMergeMode::ZEROING, | ||
| 330 | + AscendC::RoundMode::CAST_RINT}; | ||
| 331 | + | ||
| 332 | + static constexpr MicroAPI::CastTrait castTraitBf16ToFp32 = { | ||
| 333 | + MicroAPI::RegLayout::ZERO, MicroAPI::SatMode::UNKNOWN, MicroAPI::MaskMergeMode::ZEROING, | ||
| 334 | + AscendC::RoundMode::CAST_NONE}; | ||
| 335 | + | ||
| 336 | + MicroAPI::RegTensor<ElementPerTokenScale> vRegFp8; | ||
| 337 | + MicroAPI::RegTensor<bfloat16_t> vRegBf16; | ||
| 338 | + MicroAPI::RegTensor<float> vRegFp32; | ||
| 339 | + MicroAPI::MaskReg maskAll; | ||
| 340 | + | ||
| 341 | + constexpr uint32_t ELE_NUM_PER_REPEAT = 64; | ||
| 342 | + uint16_t repeatTimes = static_cast<uint16_t>((count + ELE_NUM_PER_REPEAT - 1) / ELE_NUM_PER_REPEAT); | ||
| 343 | + | ||
| 344 | + for (uint16_t i = 0; i < repeatTimes; ++i) { | ||
| 345 | + maskAll = MicroAPI::UpdateMask<float>(count); | ||
| 346 | + MicroAPI::DataCopy<ElementPerTokenScale, MicroAPI::LoadDist::DIST_UNPACK4_B8>( | ||
| 347 | + vRegFp8, srcPtr + i * ELE_NUM_PER_REPEAT); | ||
| 348 | + MicroAPI::Cast<bfloat16_t, ElementPerTokenScale, castTraitFp8ToBf16>(vRegBf16, vRegFp8, maskAll); | ||
| 349 | + MicroAPI::Cast<float, bfloat16_t, castTraitBf16ToFp32>(vRegFp32, vRegBf16, maskAll); | ||
| 350 | + MicroAPI::StoreAlign<float, MicroAPI::StoreDist::DIST_NORM_B32>( | ||
| 351 | + dstPtr + i * ELE_NUM_PER_REPEAT * 4, vRegFp32, maskAll); | ||
| 352 | + } | ||
| 353 | + } | ||
| 354 | + | ||
| 355 | + /// Cast float8_e8m0_t to fp32 using MicroAPI | ||
| 356 | + __simd_vf__ inline void MulCompute( | ||
| 357 | + __ubuf__ float* dstPtr, __ubuf__ float* srcPtr, __ubuf__ ElementPerTokenScale* src1Ptr, __ubuf__ float* src2Ptr) | ||
| 358 | + { | ||
| 359 | + namespace MicroAPI = AscendC::MicroAPI; | ||
| 360 | + | ||
| 361 | + static constexpr MicroAPI::CastTrait castTraitFp8ToBf16 = { | ||
| 362 | + MicroAPI::RegLayout::ZERO, MicroAPI::SatMode::UNKNOWN, MicroAPI::MaskMergeMode::ZEROING, | ||
| 363 | + AscendC::RoundMode::CAST_RINT}; | ||
| 364 | + | ||
| 365 | + static constexpr MicroAPI::CastTrait castTraitBf16ToFp32 = { | ||
| 366 | + MicroAPI::RegLayout::ZERO, MicroAPI::SatMode::UNKNOWN, MicroAPI::MaskMergeMode::ZEROING, | ||
| 367 | + AscendC::RoundMode::CAST_NONE}; | ||
| 368 | + | ||
| 369 | + static constexpr MicroAPI::CastTrait castTraitFp8ToFp32 = { | ||
| 370 | + MicroAPI::RegLayout::ZERO, MicroAPI::SatMode::UNKNOWN, MicroAPI::MaskMergeMode::ZEROING, | ||
| 371 | + AscendC::RoundMode::UNKNOWN}; | ||
| 372 | + | ||
| 373 | + static constexpr MicroAPI::CastTrait castTraitFp16ToFp32 = { | ||
| 374 | + MicroAPI::RegLayout::ZERO, MicroAPI::SatMode::UNKNOWN, MicroAPI::MaskMergeMode::ZEROING, | ||
| 375 | + AscendC::RoundMode::UNKNOWN}; | ||
| 376 | + | ||
| 377 | + MicroAPI::RegTensor<ElementPerTokenScale> vRegFp8_1; | ||
| 378 | + MicroAPI::RegTensor<ElementPerTokenScale> vRegFp8_2; | ||
| 379 | + MicroAPI::RegTensor<ElementPerTokenScale> vRegFp8_3; | ||
| 380 | + MicroAPI::RegTensor<ElementPerTokenScale> vRegFp8_4; | ||
| 381 | + MicroAPI::RegTensor<bfloat16_t> vRegBf16_1; | ||
| 382 | + MicroAPI::RegTensor<bfloat16_t> vRegBf16_2; | ||
| 383 | + MicroAPI::RegTensor<bfloat16_t> vRegBf16_3; | ||
| 384 | + MicroAPI::RegTensor<bfloat16_t> vRegBf16_4; | ||
| 385 | + | ||
| 386 | + MicroAPI::MaskReg maskAll = MicroAPI::CreateMask<float, MicroAPI::MaskPattern::ALL>(); | ||
| 387 | + | ||
| 388 | + MicroAPI::RegTensor<float> vRegFp32_1; | ||
| 389 | + MicroAPI::RegTensor<float> vRegFp32_2; | ||
| 390 | + MicroAPI::RegTensor<float> vRegFp32_3; | ||
| 391 | + MicroAPI::RegTensor<float> vRegFp32_4; | ||
| 392 | + MicroAPI::RegTensor<float> vRegMuls1_1; | ||
| 393 | + MicroAPI::RegTensor<float> vRegMuls1_2; | ||
| 394 | + MicroAPI::RegTensor<float> vRegMuls1_3; | ||
| 395 | + MicroAPI::RegTensor<float> vRegMuls1_4; | ||
| 396 | + MicroAPI::RegTensor<float> vRegMuls2; | ||
| 397 | + | ||
| 398 | + uint32_t row = TileShape::ROW; | ||
| 399 | + uint32_t col = TileShape::COLUMN; | ||
| 400 | + | ||
| 401 | + constexpr uint32_t ELE_NUM_PER_REPEAT = 64; | ||
| 402 | + | ||
| 403 | + if constexpr (std::is_same_v<ElementPerTokenScale, float8_e8m0_t>) { | ||
| 404 | + MicroAPI::DataCopy<ElementPerTokenScale, MicroAPI::LoadDist::DIST_UNPACK4_B8>(vRegFp8_1, src1Ptr); | ||
| 405 | + MicroAPI::Cast<bfloat16_t, ElementPerTokenScale, castTraitFp8ToBf16>(vRegBf16_1, vRegFp8_1, maskAll); | ||
| 406 | + MicroAPI::Cast<float, bfloat16_t, castTraitBf16ToFp32>(vRegMuls1_1, vRegBf16_1, maskAll); | ||
| 407 | + MicroAPI::DataCopy<ElementPerTokenScale, MicroAPI::LoadDist::DIST_UNPACK4_B8>( | ||
| 408 | + vRegFp8_2, src1Ptr + ELE_NUM_PER_REPEAT); | ||
| 409 | + MicroAPI::Cast<bfloat16_t, ElementPerTokenScale, castTraitFp8ToBf16>(vRegBf16_2, vRegFp8_2, maskAll); | ||
| 410 | + MicroAPI::Cast<float, bfloat16_t, castTraitBf16ToFp32>(vRegMuls1_2, vRegBf16_2, maskAll); | ||
| 411 | + MicroAPI::DataCopy<ElementPerTokenScale, MicroAPI::LoadDist::DIST_UNPACK4_B8>( | ||
| 412 | + vRegFp8_3, src1Ptr + ELE_NUM_PER_REPEAT * 2); | ||
| 413 | + MicroAPI::Cast<bfloat16_t, ElementPerTokenScale, castTraitFp8ToBf16>(vRegBf16_3, vRegFp8_3, maskAll); | ||
| 414 | + MicroAPI::Cast<float, bfloat16_t, castTraitBf16ToFp32>(vRegMuls1_3, vRegBf16_3, maskAll); | ||
| 415 | + MicroAPI::DataCopy<ElementPerTokenScale, MicroAPI::LoadDist::DIST_UNPACK4_B8>( | ||
| 416 | + vRegFp8_4, src1Ptr + ELE_NUM_PER_REPEAT * 3); | ||
| 417 | + MicroAPI::Cast<bfloat16_t, ElementPerTokenScale, castTraitFp8ToBf16>(vRegBf16_4, vRegFp8_4, maskAll); | ||
| 418 | + MicroAPI::Cast<float, bfloat16_t, castTraitBf16ToFp32>(vRegMuls1_4, vRegBf16_4, maskAll); | ||
| 419 | + } else if constexpr ( | ||
| 420 | + std::is_same_v<ElementPerTokenScale, float8_e4m3_t> || | ||
| 421 | + std::is_same_v<ElementPerTokenScale, float8_e5m2_t>) { | ||
| 422 | + MicroAPI::DataCopy<ElementPerTokenScale, MicroAPI::LoadDist::DIST_UNPACK4_B8>(vRegFp8_1, src1Ptr); | ||
| 423 | + MicroAPI::Cast<float, ElementPerTokenScale, castTraitFp8ToFp32>(vRegMuls1_1, vRegFp8_1, maskAll); | ||
| 424 | + MicroAPI::DataCopy<ElementPerTokenScale, MicroAPI::LoadDist::DIST_UNPACK4_B8>( | ||
| 425 | + vRegFp8_2, src1Ptr + ELE_NUM_PER_REPEAT); | ||
| 426 | + MicroAPI::Cast<float, ElementPerTokenScale, castTraitFp8ToFp32>(vRegMuls1_2, vRegFp8_2, maskAll); | ||
| 427 | + MicroAPI::DataCopy<ElementPerTokenScale, MicroAPI::LoadDist::DIST_UNPACK4_B8>( | ||
| 428 | + vRegFp8_3, src1Ptr + ELE_NUM_PER_REPEAT * 2); | ||
| 429 | + MicroAPI::Cast<float, ElementPerTokenScale, castTraitFp8ToFp32>(vRegMuls1_3, vRegFp8_3, maskAll); | ||
| 430 | + MicroAPI::DataCopy<ElementPerTokenScale, MicroAPI::LoadDist::DIST_UNPACK4_B8>( | ||
| 431 | + vRegFp8_4, src1Ptr + ELE_NUM_PER_REPEAT * 3); | ||
| 432 | + MicroAPI::Cast<float, ElementPerTokenScale, castTraitFp8ToFp32>(vRegMuls1_4, vRegFp8_4, maskAll); | ||
| 433 | + } else if constexpr (std::is_same_v<ElementPerTokenScale, float>) { | ||
| 434 | + // Already fp32, copy | ||
| 435 | + MicroAPI::DataCopy<ElementPerTokenScale, MicroAPI::LoadDist::DIST_NORM>(vRegMuls1_1, src1Ptr); | ||
| 436 | + MicroAPI::DataCopy<ElementPerTokenScale, MicroAPI::LoadDist::DIST_NORM>( | ||
| 437 | + vRegMuls1_2, src1Ptr + ELE_NUM_PER_REPEAT); | ||
| 438 | + MicroAPI::DataCopy<ElementPerTokenScale, MicroAPI::LoadDist::DIST_NORM>( | ||
| 439 | + vRegMuls1_3, src1Ptr + ELE_NUM_PER_REPEAT * 2); | ||
| 440 | + MicroAPI::DataCopy<ElementPerTokenScale, MicroAPI::LoadDist::DIST_NORM>( | ||
| 441 | + vRegMuls1_4, src1Ptr + ELE_NUM_PER_REPEAT * 3); | ||
| 442 | + } else { | ||
| 443 | + MicroAPI::DataCopy<ElementPerTokenScale, MicroAPI::LoadDist::DIST_UNPACK_B16>(vRegFp8_1, src1Ptr); | ||
| 444 | + MicroAPI::Cast<float, ElementPerTokenScale, castTraitFp16ToFp32>(vRegMuls1_1, vRegBf16_1, maskAll); | ||
| 445 | + MicroAPI::DataCopy<ElementPerTokenScale, MicroAPI::LoadDist::DIST_UNPACK_B16>( | ||
| 446 | + vRegFp8_2, src1Ptr + ELE_NUM_PER_REPEAT); | ||
| 447 | + MicroAPI::Cast<float, ElementPerTokenScale, castTraitFp16ToFp32>(vRegMuls1_2, vRegBf16_2, maskAll); | ||
| 448 | + MicroAPI::DataCopy<ElementPerTokenScale, MicroAPI::LoadDist::DIST_UNPACK_B16>( | ||
| 449 | + vRegFp8_3, src1Ptr + ELE_NUM_PER_REPEAT * 2); | ||
| 450 | + MicroAPI::Cast<float, ElementPerTokenScale, castTraitFp16ToFp32>(vRegMuls1_3, vRegBf16_3, maskAll); | ||
| 451 | + MicroAPI::DataCopy<ElementPerTokenScale, MicroAPI::LoadDist::DIST_UNPACK_B16>( | ||
| 452 | + vRegFp8_4, src1Ptr + ELE_NUM_PER_REPEAT * 3); | ||
| 453 | + MicroAPI::Cast<float, ElementPerTokenScale, castTraitFp16ToFp32>(vRegMuls1_4, vRegBf16_4, maskAll); | ||
| 454 | + } | ||
| 455 | + | ||
| 456 | + for (uint16_t i = 0; i < row; ++i) { | ||
| 457 | + MicroAPI::DataCopy<float, MicroAPI::LoadDist::DIST_NORM>(vRegFp32_1, srcPtr + i * col); | ||
| 458 | + MicroAPI::DataCopy<float, MicroAPI::LoadDist::DIST_NORM>(vRegFp32_2, srcPtr + i * col + ELE_NUM_PER_REPEAT); | ||
| 459 | + MicroAPI::DataCopy<float, MicroAPI::LoadDist::DIST_NORM>( | ||
| 460 | + vRegFp32_3, srcPtr + i * col + ELE_NUM_PER_REPEAT * 2); | ||
| 461 | + MicroAPI::DataCopy<float, MicroAPI::LoadDist::DIST_NORM>( | ||
| 462 | + vRegFp32_4, srcPtr + i * col + ELE_NUM_PER_REPEAT * 3); | ||
| 463 | + MicroAPI::DataCopy<float, MicroAPI::LoadDist::DIST_BRC_B32>(vRegMuls2, src2Ptr + i); | ||
| 464 | + MicroAPI::Mul(vRegFp32_1, vRegFp32_1, vRegMuls1_1, maskAll); | ||
| 465 | + MicroAPI::Mul(vRegFp32_2, vRegFp32_2, vRegMuls1_2, maskAll); | ||
| 466 | + MicroAPI::Mul(vRegFp32_3, vRegFp32_3, vRegMuls1_3, maskAll); | ||
| 467 | + MicroAPI::Mul(vRegFp32_4, vRegFp32_4, vRegMuls1_4, maskAll); | ||
| 468 | + MicroAPI::Mul(vRegFp32_1, vRegFp32_1, vRegMuls2, maskAll); | ||
| 469 | + MicroAPI::Mul(vRegFp32_2, vRegFp32_2, vRegMuls2, maskAll); | ||
| 470 | + MicroAPI::Mul(vRegFp32_3, vRegFp32_3, vRegMuls2, maskAll); | ||
| 471 | + MicroAPI::Mul(vRegFp32_4, vRegFp32_4, vRegMuls2, maskAll); | ||
| 472 | + MicroAPI::StoreAlign<float, MicroAPI::StoreDist::DIST_NORM_B32>(dstPtr + i * col, vRegFp32_1, maskAll); | ||
| 473 | + MicroAPI::StoreAlign<float, MicroAPI::StoreDist::DIST_NORM_B32>( | ||
| 474 | + dstPtr + i * col + ELE_NUM_PER_REPEAT, vRegFp32_2, maskAll); | ||
| 475 | + MicroAPI::StoreAlign<float, MicroAPI::StoreDist::DIST_NORM_B32>( | ||
| 476 | + dstPtr + i * col + ELE_NUM_PER_REPEAT * 2, vRegFp32_3, maskAll); | ||
| 477 | + MicroAPI::StoreAlign<float, MicroAPI::StoreDist::DIST_NORM_B32>( | ||
| 478 | + dstPtr + i * col + ELE_NUM_PER_REPEAT * 3, vRegFp32_4, maskAll); | ||
| 479 | + } | ||
| 480 | + } | ||
| 481 | + | ||
| 482 | + __simd_vf__ inline void MulComputeGeneric( | ||
| 483 | + __ubuf__ float* dstPtr, __ubuf__ float* srcPtr, __ubuf__ ElementScale* src1Ptr, __ubuf__ float* src2Ptr) | ||
| 484 | + { | ||
| 485 | + namespace MicroAPI = AscendC::MicroAPI; | ||
| 486 | + | ||
| 487 | + static constexpr MicroAPI::CastTrait castTraitFp8ToBf16 = { | ||
| 488 | + MicroAPI::RegLayout::ZERO, MicroAPI::SatMode::UNKNOWN, MicroAPI::MaskMergeMode::ZEROING, | ||
| 489 | + AscendC::RoundMode::CAST_RINT}; | ||
| 490 | + | ||
| 491 | + static constexpr MicroAPI::CastTrait castTraitBf16ToFp32 = { | ||
| 492 | + MicroAPI::RegLayout::ZERO, MicroAPI::SatMode::UNKNOWN, MicroAPI::MaskMergeMode::ZEROING, | ||
| 493 | + AscendC::RoundMode::CAST_NONE}; | ||
| 494 | + | ||
| 495 | + static constexpr MicroAPI::CastTrait castTraitFp8ToFp32 = { | ||
| 496 | + MicroAPI::RegLayout::ZERO, MicroAPI::SatMode::UNKNOWN, MicroAPI::MaskMergeMode::ZEROING, | ||
| 497 | + AscendC::RoundMode::UNKNOWN}; | ||
| 498 | + | ||
| 499 | + static constexpr MicroAPI::CastTrait castTraitFp16ToFp32 = { | ||
| 500 | + MicroAPI::RegLayout::ZERO, MicroAPI::SatMode::UNKNOWN, MicroAPI::MaskMergeMode::ZEROING, | ||
| 501 | + AscendC::RoundMode::UNKNOWN}; | ||
| 502 | + | ||
| 503 | + uint32_t row = TileShape::ROW; | ||
| 504 | + uint32_t col = TileShape::COLUMN; | ||
| 505 | + constexpr uint32_t ELE_NUM_PER_REPEAT = 64; | ||
| 506 | + | ||
| 507 | + MicroAPI::RegTensor<float> vRegC; | ||
| 508 | + MicroAPI::RegTensor<float> vRegScale; | ||
| 509 | + MicroAPI::RegTensor<float> vRegPerTokenScale; | ||
| 510 | + MicroAPI::RegTensor<ElementScale> vRegScaleRaw; | ||
| 511 | + MicroAPI::RegTensor<bfloat16_t> vRegBf16; | ||
| 512 | + MicroAPI::MaskReg maskAll = MicroAPI::CreateMask<float, MicroAPI::MaskPattern::ALL>(); | ||
| 513 | + | ||
| 514 | + for (uint32_t i = 0; i < row; ++i) { | ||
| 515 | + MicroAPI::DataCopy<float, MicroAPI::LoadDist::DIST_NORM>(vRegPerTokenScale, src2Ptr + i); | ||
| 516 | + | ||
| 517 | + for (uint32_t j = 0; j < col; j += ELE_NUM_PER_REPEAT) { | ||
| 518 | + MicroAPI::DataCopy<float, MicroAPI::LoadDist::DIST_NORM>(vRegC, srcPtr + i * col + j); | ||
| 519 | + | ||
| 520 | + if constexpr (std::is_same_v<ElementScale, float8_e8m0_t>) { | ||
| 521 | + MicroAPI::DataCopy<ElementScale, MicroAPI::LoadDist::DIST_UNPACK4_B8>(vRegScaleRaw, src1Ptr + j); | ||
| 522 | + MicroAPI::Cast<bfloat16_t, ElementScale, castTraitFp8ToBf16>(vRegBf16, vRegScaleRaw, maskAll); | ||
| 523 | + MicroAPI::Cast<float, bfloat16_t, castTraitBf16ToFp32>(vRegScale, vRegBf16, maskAll); | ||
| 524 | + } else if constexpr ( | ||
| 525 | + std::is_same_v<ElementScale, float8_e4m3_t> || std::is_same_v<ElementScale, float8_e5m2_t>) { | ||
| 526 | + MicroAPI::DataCopy<ElementScale, MicroAPI::LoadDist::DIST_UNPACK4_B8>(vRegScaleRaw, src1Ptr + j); | ||
| 527 | + MicroAPI::Cast<float, ElementScale, castTraitFp8ToFp32>(vRegScale, vRegScaleRaw, maskAll); | ||
| 528 | + } else if constexpr (std::is_same_v<ElementScale, float>) { | ||
| 529 | + MicroAPI::DataCopy<float, MicroAPI::LoadDist::DIST_NORM>(vRegScale, src1Ptr + j); | ||
| 530 | + } else { | ||
| 531 | + MicroAPI::DataCopy<ElementScale, MicroAPI::LoadDist::DIST_UNPACK_B16>(vRegScaleRaw, src1Ptr + j); | ||
| 532 | + MicroAPI::Cast<float, ElementScale, castTraitFp16ToFp32>(vRegScale, vRegScaleRaw, maskAll); | ||
| 533 | + } | ||
| 534 | + | ||
| 535 | + MicroAPI::Mul(vRegC, vRegC, vRegScale, maskAll); | ||
| 536 | + MicroAPI::Mul(vRegC, vRegC, vRegPerTokenScale, maskAll); | ||
| 537 | + MicroAPI::StoreAlign<float, MicroAPI::StoreDist::DIST_NORM_B32>(dstPtr + i * col + j, vRegC, maskAll); | ||
| 538 | + } | ||
| 539 | + } | ||
| 540 | + } | ||
| 541 | +}; | ||
| 542 | +} // namespace Catlass::Epilogue::Block | ||
| 543 | + | ||
| 544 | + | ||
| @@ -0,0 +1,363 @@ | |||
| 1 | +/** | ||
| 2 | + * This program is free software, you can redistribute it and/or modify. | ||
| 3 | + * Copyright (c) 2026 Huawei Technologies Co., Ltd. | ||
| 4 | + * This file is a part of the CANN Open Software. | ||
| 5 | + * Licensed under CANN Open Software License Agreement Version 2.0 (the "License"). | ||
| 6 | + * Please refer to the License for details. You may not use this file except in compliance with the License. | ||
| 7 | + * THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED, INCLUDING | ||
| 8 | + * BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE. See LICENSE in the root of | ||
| 9 | + * the software repository for the full text of the License. | ||
| 10 | + */ | ||
| 11 | + | ||
| 12 | + | ||
| 13 | + | ||
| 14 | + | ||
| 15 | + | ||
| 16 | + | ||
| 17 | + | ||
| 18 | + | ||
| 19 | + | ||
| 20 | + | ||
| 21 | + | ||
| 22 | + | ||
| 23 | + | ||
| 24 | +namespace Catlass::Epilogue::Block { | ||
| 25 | + | ||
| 26 | +template < | ||
| 27 | + uint32_t UB_STAGES_, class CType_, class ScaleType_, class PerTokenScaleType_, class DType_, | ||
| 28 | + class TileRowBroadcastMul_, class TileBroadcastOneBlk_, class TileOneBlkColumnBroadcastMul_, class TileCopy_, | ||
| 29 | + class EpilogueTileSwizzle_> | ||
| 30 | +class BlockEpilogue< | ||
| 31 | + EpilogueAscend950PerTokenPerChannelQuant<UB_STAGES_>, CType_, ScaleType_, PerTokenScaleType_, DType_, | ||
| 32 | + TileRowBroadcastMul_, TileBroadcastOneBlk_, TileOneBlkColumnBroadcastMul_, TileCopy_, EpilogueTileSwizzle_> { | ||
| 33 | +public: | ||
| 34 | + using DispatchPolicy = EpilogueAscend950PerTokenPerChannelQuant<UB_STAGES_>; | ||
| 35 | + using ArchTag = typename DispatchPolicy::ArchTag; | ||
| 36 | + static constexpr uint32_t UB_STAGES = UB_STAGES_; | ||
| 37 | + | ||
| 38 | + // Data infos | ||
| 39 | + using ElementC = typename CType_::Element; // 现在float | ||
| 40 | + using LayoutC = typename CType_::Layout; | ||
| 41 | + using ElementScale = typename ScaleType_::Element; // 现在fp8_e5m2 | ||
| 42 | + using LayoutScale = typename ScaleType_::Layout; | ||
| 43 | + using ElementPerTokenScale = typename PerTokenScaleType_::Element; // 现在fp8_e5m2 | ||
| 44 | + using LayoutPerTokenScale = typename PerTokenScaleType_::Layout; | ||
| 45 | + using ElementD = typename DType_::Element; // 现在根据输入确定是half or float | ||
| 46 | + using LayoutD = typename DType_::Layout; | ||
| 47 | + | ||
| 48 | + static_assert( | ||
| 49 | + std::is_same_v<LayoutC, layout::RowMajor> && std::is_same_v<LayoutScale, layout::VectorLayout> && | ||
| 50 | + std::is_same_v<LayoutPerTokenScale, layout::VectorLayout> && std::is_same_v<LayoutD, layout::RowMajor>, | ||
| 51 | + "The layout template parameters of BlockEpilogue are wrong"); | ||
| 52 | + | ||
| 53 | + // Tile compute ops | ||
| 54 | + using TileRowBroadcastMul = TileRowBroadcastMul_; | ||
| 55 | + using TileBroadcastOneBlk = TileBroadcastOneBlk_; | ||
| 56 | + using TileOneBlkColumnBroadcastMul = TileOneBlkColumnBroadcastMul_; | ||
| 57 | + | ||
| 58 | + // Tile copy | ||
| 59 | + using CopyGmToUbC = typename TileCopy_::CopyGmToUbC; | ||
| 60 | + using CopyGmToUbScale = typename TileCopy_::CopyGmToUbX; | ||
| 61 | + using CopyGmToUbPerTokenScale = typename TileCopy_::CopyGmToUbY; | ||
| 62 | + using CopyUbToGmD = typename TileCopy_::CopyUbToGmD; | ||
| 63 | + | ||
| 64 | + using EpilogueTileSwizzle = EpilogueTileSwizzle_; | ||
| 65 | + | ||
| 66 | + using TileShape = typename TileRowBroadcastMul::TileShape; | ||
| 67 | + | ||
| 68 | + static_assert( | ||
| 69 | + TileShape::ROW == TileBroadcastOneBlk::COMPUTE_LENGTH && | ||
| 70 | + std::is_same_v<TileShape, typename TileOneBlkColumnBroadcastMul::TileShape>, | ||
| 71 | + "TileShape must be consistent for all tile compute ops"); | ||
| 72 | + | ||
| 73 | + static_assert( | ||
| 74 | + (UB_STAGES * (TileShape::COUNT * sizeof(ElementC) + TileShape::COLUMN * sizeof(ElementScale) + | ||
| 75 | + TileShape::ROW * sizeof(ElementPerTokenScale) + TileShape::COUNT * sizeof(ElementD)) + | ||
| 76 | + (TileShape::COUNT + TileShape::COLUMN + TileShape::ROW) * sizeof(float) + TileShape::ROW * BYTE_PER_BLK) <= | ||
| 77 | + ArchTag::UB_SIZE, | ||
| 78 | + "TileShape is too large to fit in UB"); | ||
| 79 | + | ||
| 80 | + struct Params { | ||
| 81 | + __gm__ ElementScale* ptrScale{nullptr}; | ||
| 82 | + LayoutScale layoutScale{}; | ||
| 83 | + __gm__ ElementPerTokenScale* ptrPerTokenScale{nullptr}; | ||
| 84 | + LayoutPerTokenScale layoutPerTokenScale{}; | ||
| 85 | + __gm__ ElementD* ptrD{nullptr}; | ||
| 86 | + LayoutD layoutD{}; | ||
| 87 | + | ||
| 88 | + CATLASS_DEVICE | ||
| 89 | + Params() {}; | ||
| 90 | + | ||
| 91 | + CATLASS_DEVICE | ||
| 92 | + Params( | ||
| 93 | + __gm__ ElementScale* ptrScale_, LayoutScale const& layoutScale_, | ||
| 94 | + __gm__ ElementPerTokenScale* ptrPerTokenScale_, LayoutPerTokenScale const& layoutPerTokenScale_, | ||
| 95 | + __gm__ ElementD* ptrD_, LayoutD const& layoutD_) | ||
| 96 | + : ptrScale(ptrScale_), | ||
| 97 | + layoutScale(layoutScale_), | ||
| 98 | + ptrPerTokenScale(ptrPerTokenScale_), | ||
| 99 | + layoutPerTokenScale(layoutPerTokenScale_), | ||
| 100 | + ptrD(ptrD_), | ||
| 101 | + layoutD(layoutD_) | ||
| 102 | + {} | ||
| 103 | + }; | ||
| 104 | + | ||
| 105 | + CATLASS_DEVICE | ||
| 106 | + BlockEpilogue(Arch::Resource<ArchTag> const& resource, Params const& params = Params{}) : params(params) | ||
| 107 | + { | ||
| 108 | + size_t ubOffset = 0; | ||
| 109 | + int32_t eventVMTE2 = 0; | ||
| 110 | + int32_t eventMTE2V = 0; | ||
| 111 | + int32_t eventMTE3V = 0; | ||
| 112 | + int32_t eventVMTE3 = 0; | ||
| 113 | + | ||
| 114 | + int32_t eventid = 0; | ||
| 115 | + for (uint32_t i = 0; i < UB_STAGES; ++i) { | ||
| 116 | + ubCList[i] = resource.ubBuf.template GetBufferByByte<ElementC>(ubOffset); | ||
| 117 | + ubOffset += TileShape::COUNT * sizeof(ElementC); | ||
| 118 | + ubScaleList[i] = resource.ubBuf.template GetBufferByByte<ElementScale>(ubOffset); | ||
| 119 | + ubOffset += TileShape::COLUMN * sizeof(ElementScale); | ||
| 120 | + ubPerTokenScaleList[i] = resource.ubBuf.template GetBufferByByte<ElementPerTokenScale>(ubOffset); | ||
| 121 | + ubOffset += TileShape::ROW * sizeof(ElementPerTokenScale); | ||
| 122 | + ubDList[i] = resource.ubBuf.template GetBufferByByte<ElementD>(ubOffset); | ||
| 123 | + ubOffset += TileShape::COUNT * sizeof(ElementD); | ||
| 124 | + } | ||
| 125 | + ubScaleFp32 = resource.ubBuf.template GetBufferByByte<float>(ubOffset); | ||
| 126 | + ubOffset += TileShape::COLUMN * sizeof(float); | ||
| 127 | + ubMul = resource.ubBuf.template GetBufferByByte<float>(ubOffset); | ||
| 128 | + ubOffset += TileShape::COUNT * sizeof(float); | ||
| 129 | + ubPerTokenScaleFp32 = resource.ubBuf.template GetBufferByByte<float>(ubOffset); | ||
| 130 | + ubOffset += TileShape::ROW * sizeof(float); | ||
| 131 | + ubPerTokenScaleFp32Brcb = resource.ubBuf.template GetBufferByByte<float>(ubOffset); | ||
| 132 | + ubOffset += TileShape::ROW * BYTE_PER_BLK; | ||
| 133 | + ubPerTokenMul = ubMul; | ||
| 134 | + } | ||
| 135 | + | ||
| 136 | + CATLASS_DEVICE | ||
| 137 | + ~BlockEpilogue() | ||
| 138 | + {} | ||
| 139 | + | ||
| 140 | + CATLASS_DEVICE | ||
| 141 | + void UpdateParams(Params const& params_) | ||
| 142 | + { | ||
| 143 | + params = params_; | ||
| 144 | + } | ||
| 145 | + | ||
| 146 | + CATLASS_DEVICE | ||
| 147 | + void operator()( | ||
| 148 | + GemmCoord const& blockShapeMNK, GemmCoord const& blockCoordMNK, GemmCoord const& actualBlockShapeMNK, | ||
| 149 | + AscendC::GlobalTensor<ElementC> const& gmBlockC, LayoutC const& layoutBlockC, Callback&& callback = Callback{}) | ||
| 150 | + { | ||
| 151 | + if (actualBlockShapeMNK.k() == 0) { | ||
| 152 | + return; | ||
| 153 | + } | ||
| 154 | + callback(); | ||
| 155 | + | ||
| 156 | + // Calculate the offset of the current block | ||
| 157 | + MatrixCoord blockShape = blockShapeMNK.GetCoordMN(); | ||
| 158 | + MatrixCoord blockCoord = blockCoordMNK.GetCoordMN(); | ||
| 159 | + MatrixCoord actualBlockShape = actualBlockShapeMNK.GetCoordMN(); | ||
| 160 | + MatrixCoord blockOffset = blockCoord * blockShape; | ||
| 161 | + | ||
| 162 | + AscendC::GlobalTensor<ElementScale> gmScale; | ||
| 163 | + gmScale.SetGlobalBuffer(params.ptrScale); | ||
| 164 | + AscendC::GlobalTensor<ElementPerTokenScale> gmPerTokenScale; | ||
| 165 | + gmPerTokenScale.SetGlobalBuffer(params.ptrPerTokenScale); | ||
| 166 | + AscendC::GlobalTensor<ElementD> gmD; | ||
| 167 | + gmD.SetGlobalBuffer(params.ptrD); | ||
| 168 | + | ||
| 169 | + auto ubTileStride = MakeCoord(static_cast<int64_t>(TileShape::COLUMN), 1L); | ||
| 170 | + auto tileShape = TileShape::ToCoord(); | ||
| 171 | + EpilogueTileSwizzle epilogueTileSwizzle(actualBlockShape, tileShape); | ||
| 172 | + uint32_t tileLoops = epilogueTileSwizzle.GetLoops(); | ||
| 173 | + uint32_t subblockIdx = AscendC::GetSubBlockIdx(); | ||
| 174 | + uint32_t subblockNum = AscendC::GetSubBlockNum(); | ||
| 175 | + | ||
| 176 | + AscendC::SetFlag<AscendC::HardEvent::MTE3_MTE2>(0); | ||
| 177 | + AscendC::SetFlag<AscendC::HardEvent::MTE3_MTE2>(1); | ||
| 178 | + for (uint32_t loopIdx = subblockIdx; loopIdx < tileLoops; loopIdx += subblockNum) { | ||
| 179 | + auto tileCoord = epilogueTileSwizzle.GetTileCoord(loopIdx); | ||
| 180 | + auto actualTileShape = epilogueTileSwizzle.GetActualTileShape(tileCoord); | ||
| 181 | + auto tileOffsetInBlock = tileCoord * tileShape; | ||
| 182 | + auto tileOffset = blockOffset + tileOffsetInBlock; | ||
| 183 | + | ||
| 184 | + auto gmTileC = gmBlockC[layoutBlockC.GetOffset(tileOffsetInBlock)]; | ||
| 185 | + auto layoutGmTileC = layoutBlockC.GetTileLayout(actualTileShape); | ||
| 186 | + | ||
| 187 | + auto& ubC = ubCList[ubListId]; | ||
| 188 | + LayoutC layoutUbC{actualTileShape, ubTileStride}; | ||
| 189 | + | ||
| 190 | + auto eventId = ubListId ? EVENT_ID0 : EVENT_ID1; | ||
| 191 | + AscendC::PipeBarrier<PIPE_ALL>(); | ||
| 192 | + | ||
| 193 | + AscendC::WaitFlag<AscendC::HardEvent::MTE3_MTE2>(eventId); | ||
| 194 | + | ||
| 195 | + copyGmToUbC(ubC, gmTileC, layoutUbC, layoutGmTileC); | ||
| 196 | + | ||
| 197 | + auto scaleTileOffset = tileOffset.template GetCoordByAxis<1>(); | ||
| 198 | + auto scaleTileShape = actualTileShape.template GetCoordByAxis<1>(); | ||
| 199 | + | ||
| 200 | + auto gmTileScale = gmScale[params.layoutScale.GetOffset(scaleTileOffset)]; | ||
| 201 | + auto layoutGmTileScale = params.layoutScale.GetTileLayout(scaleTileShape); | ||
| 202 | + | ||
| 203 | + auto& ubScale = ubScaleList[ubListId]; | ||
| 204 | + auto layoutUbScale = LayoutScale::template MakeLayoutInUb<ElementScale>(scaleTileShape); | ||
| 205 | + | ||
| 206 | + copyGmToUbScale(ubScale, gmTileScale, layoutUbScale, layoutGmTileScale); | ||
| 207 | + | ||
| 208 | + auto perTokenScaleTileOffset = tileOffset.template GetCoordByAxis<0>(); | ||
| 209 | + auto perTokenScaleTileShape = actualTileShape.template GetCoordByAxis<0>(); | ||
| 210 | + | ||
| 211 | + auto gmTilePerTokenScale = gmPerTokenScale[params.layoutPerTokenScale.GetOffset(perTokenScaleTileOffset)]; | ||
| 212 | + auto layoutGmTilePerTokenScale = params.layoutPerTokenScale.GetTileLayout(perTokenScaleTileShape); | ||
| 213 | + | ||
| 214 | + auto& ubPerTokenScale = ubPerTokenScaleList[ubListId]; | ||
| 215 | + auto layoutUbPerTokenScale = | ||
| 216 | + LayoutScale::template MakeLayoutInUb<ElementPerTokenScale>(perTokenScaleTileShape); | ||
| 217 | + | ||
| 218 | + copyGmToUbPerTokenScale( | ||
| 219 | + ubPerTokenScale, gmTilePerTokenScale, layoutUbPerTokenScale, layoutGmTilePerTokenScale); | ||
| 220 | + | ||
| 221 | + AscendC::SetFlag<AscendC::HardEvent::MTE2_V>(eventId); | ||
| 222 | + AscendC::WaitFlag<AscendC::HardEvent::MTE2_V>(eventId); | ||
| 223 | + if constexpr (std::is_same<ElementScale, float8_e8m0_t>::value) { | ||
| 224 | + // High-level API doesn't support float8_e8m0_t -> float, use MicroAPI instead | ||
| 225 | + __ubuf__ float8_e8m0_t* srcAddr = (__ubuf__ float8_e8m0_t*)ubScale.GetPhyAddr(); | ||
| 226 | + __ubuf__ float* dstAddr = (__ubuf__ float*)ubScaleFp32.GetPhyAddr(); | ||
| 227 | + CastFp8E8m0ToFp32(dstAddr, srcAddr, TileShape::COLUMN); | ||
| 228 | + | ||
| 229 | + } else if (!std::is_same<ElementScale, float>::value) { | ||
| 230 | + AscendC::Cast(ubScaleFp32, ubScale, AscendC::RoundMode::CAST_NONE, TileShape::COLUMN); | ||
| 231 | + } | ||
| 232 | + | ||
| 233 | + if constexpr (std::is_same<ElementPerTokenScale, float8_e8m0_t>::value) { | ||
| 234 | + // High-level API doesn't support float8_e8m0_t -> float, use MicroAPI instead | ||
| 235 | + __ubuf__ float8_e8m0_t* srcAddr = (__ubuf__ float8_e8m0_t*)ubPerTokenScale.GetPhyAddr(); | ||
| 236 | + __ubuf__ float* dstAddr = (__ubuf__ float*)ubPerTokenScaleFp32.GetPhyAddr(); | ||
| 237 | + CastFp8E8m0ToFp32(dstAddr, srcAddr, TileShape::ROW); | ||
| 238 | + | ||
| 239 | + } else if (!std::is_same<ElementPerTokenScale, float>::value) { | ||
| 240 | + AscendC::Cast(ubPerTokenScaleFp32, ubPerTokenScale, AscendC::RoundMode::CAST_NONE, TileShape::ROW); | ||
| 241 | + } | ||
| 242 | + | ||
| 243 | + AscendC::PipeBarrier<PIPE_V>(); | ||
| 244 | + tileRowBroadcastMul(ubMul, ubC, ubScaleFp32); | ||
| 245 | + AscendC::PipeBarrier<PIPE_V>(); | ||
| 246 | + tileBroadcastOneBlk(ubPerTokenScaleFp32Brcb, ubPerTokenScaleFp32); | ||
| 247 | + AscendC::PipeBarrier<PIPE_V>(); | ||
| 248 | + tileOneBlkColumnBroadcastMul(ubPerTokenMul, ubMul, ubPerTokenScaleFp32Brcb); | ||
| 249 | + AscendC::PipeBarrier<PIPE_V>(); | ||
| 250 | + | ||
| 251 | + auto& ubD = ubDList[ubListId]; | ||
| 252 | + LayoutD layoutUbD{actualTileShape, ubTileStride}; | ||
| 253 | + | ||
| 254 | + AscendC::PipeBarrier<PIPE_ALL>(); | ||
| 255 | + | ||
| 256 | + if constexpr (std::is_same_v<ElementD, half>) { | ||
| 257 | + AscendC::Cast(ubD, ubPerTokenMul, AscendC::RoundMode::CAST_RINT, TileShape::COUNT); | ||
| 258 | + } | ||
| 259 | + AscendC::PipeBarrier<PIPE_ALL>(); | ||
| 260 | + | ||
| 261 | + auto gmTileD = gmD[params.layoutD.GetOffset(tileOffset)]; | ||
| 262 | + auto layoutGmTileD = params.layoutD.GetTileLayout(actualTileShape); | ||
| 263 | + | ||
| 264 | + if constexpr (std::is_same_v<ElementD, half>) { | ||
| 265 | + copyUbToGmD(gmTileD, ubD, layoutGmTileD, layoutUbD); | ||
| 266 | + } else { | ||
| 267 | + copyUbToGmD(gmTileD, ubPerTokenMul, layoutGmTileD, layoutUbD); | ||
| 268 | + } | ||
| 269 | + | ||
| 270 | + AscendC::SetFlag<AscendC::HardEvent::MTE3_MTE2>(eventId); | ||
| 271 | + | ||
| 272 | + ubListId = (ubListId + 1 < UB_STAGES) ? (ubListId + 1) : 0; | ||
| 273 | + } | ||
| 274 | + AscendC::WaitFlag<AscendC::HardEvent::MTE3_MTE2>(0); | ||
| 275 | + AscendC::WaitFlag<AscendC::HardEvent::MTE3_MTE2>(1); | ||
| 276 | + } | ||
| 277 | + | ||
| 278 | +private: | ||
| 279 | + Params params; | ||
| 280 | + | ||
| 281 | + AscendC::LocalTensor<ElementC> ubCList[UB_STAGES]; | ||
| 282 | + AscendC::LocalTensor<ElementScale> ubScaleList[UB_STAGES]; | ||
| 283 | + AscendC::LocalTensor<ElementPerTokenScale> ubPerTokenScaleList[UB_STAGES]; | ||
| 284 | + AscendC::LocalTensor<ElementD> ubDList[UB_STAGES]; | ||
| 285 | + | ||
| 286 | + int32_t eventUbCVMTE2List[UB_STAGES]; | ||
| 287 | + int32_t eventUbCMTE2VList[UB_STAGES]; | ||
| 288 | + int32_t eventUbScaleVMTE2List[UB_STAGES]; | ||
| 289 | + int32_t eventUbScaleMTE2VList[UB_STAGES]; | ||
| 290 | + int32_t eventUbPerTokenScaleVMTE2List[UB_STAGES]; | ||
| 291 | + int32_t eventUbPerTokenScaleMTE2VList[UB_STAGES]; | ||
| 292 | + int32_t eventUbDMTE3VList[UB_STAGES]; | ||
| 293 | + int32_t eventUbDVMTE3List[UB_STAGES]; | ||
| 294 | + int32_t eventList[UB_STAGES]; | ||
| 295 | + | ||
| 296 | + uint32_t ubListId{0}; | ||
| 297 | + | ||
| 298 | + AscendC::LocalTensor<float> ubScaleFp32; | ||
| 299 | + AscendC::LocalTensor<float> ubMul; | ||
| 300 | + AscendC::LocalTensor<float> ubPerTokenScaleFp32; | ||
| 301 | + AscendC::LocalTensor<float> ubPerTokenScaleFp32Brcb; | ||
| 302 | + AscendC::LocalTensor<float> ubPerTokenMul; | ||
| 303 | + | ||
| 304 | + TileRowBroadcastMul tileRowBroadcastMul; | ||
| 305 | + TileBroadcastOneBlk tileBroadcastOneBlk; | ||
| 306 | + TileOneBlkColumnBroadcastMul tileOneBlkColumnBroadcastMul; | ||
| 307 | + | ||
| 308 | + CopyGmToUbC copyGmToUbC; | ||
| 309 | + CopyGmToUbScale copyGmToUbScale; | ||
| 310 | + CopyGmToUbPerTokenScale copyGmToUbPerTokenScale; | ||
| 311 | + CopyUbToGmD copyUbToGmD; | ||
| 312 | + | ||
| 313 | + /// Helper function to cast scale to fp32 | ||
| 314 | + template <typename T> | ||
| 315 | + CATLASS_DEVICE void CastScaleToFp32(AscendC::LocalTensor<float>& dst, AscendC::LocalTensor<T>& src, uint32_t count) | ||
| 316 | + { | ||
| 317 | + if constexpr (std::is_same_v<T, float8_e8m0_t>) { | ||
| 318 | + CastFp8E8m0ToFp32(dst, src, count); | ||
| 319 | + } else if constexpr (std::is_same_v<T, float8_e4m3_t> || std::is_same_v<T, float8_e5m2_t>) { | ||
| 320 | + AscendC::Cast(dst, src, AscendC::RoundMode::CAST_NONE, count); | ||
| 321 | + } else if constexpr (std::is_same_v<T, float>) { | ||
| 322 | + // Already fp32, copy | ||
| 323 | + AscendC::DataCopy(dst, src, count); | ||
| 324 | + } else { | ||
| 325 | + AscendC::Cast(dst, src, AscendC::RoundMode::CAST_NONE, count); | ||
| 326 | + } | ||
| 327 | + } | ||
| 328 | + | ||
| 329 | + /// Cast float8_e8m0_t to fp32 using MicroAPI | ||
| 330 | + __simd_vf__ inline void CastFp8E8m0ToFp32( | ||
| 331 | + __ubuf__ float* dstPtr, __ubuf__ ElementPerTokenScale* srcPtr, uint32_t count) | ||
| 332 | + { | ||
| 333 | + namespace MicroAPI = AscendC::MicroAPI; | ||
| 334 | + | ||
| 335 | + static constexpr MicroAPI::CastTrait castTraitFp8ToBf16 = { | ||
| 336 | + MicroAPI::RegLayout::ZERO, MicroAPI::SatMode::UNKNOWN, MicroAPI::MaskMergeMode::ZEROING, | ||
| 337 | + AscendC::RoundMode::CAST_RINT}; | ||
| 338 | + | ||
| 339 | + static constexpr MicroAPI::CastTrait castTraitBf16ToFp32 = { | ||
| 340 | + MicroAPI::RegLayout::ZERO, MicroAPI::SatMode::UNKNOWN, MicroAPI::MaskMergeMode::ZEROING, | ||
| 341 | + AscendC::RoundMode::CAST_NONE}; | ||
| 342 | + | ||
| 343 | + MicroAPI::RegTensor<ElementPerTokenScale> vRegFp8; | ||
| 344 | + MicroAPI::RegTensor<bfloat16_t> vRegBf16; | ||
| 345 | + MicroAPI::RegTensor<float> vRegFp32; | ||
| 346 | + MicroAPI::MaskReg maskAll; | ||
| 347 | + | ||
| 348 | + constexpr uint32_t ELE_NUM_PER_REPEAT = 64; | ||
| 349 | + uint16_t repeatTimes = static_cast<uint16_t>((count + ELE_NUM_PER_REPEAT - 1) / ELE_NUM_PER_REPEAT); | ||
| 350 | + | ||
| 351 | + for (uint16_t i = 0; i < repeatTimes; ++i) { | ||
| 352 | + maskAll = MicroAPI::UpdateMask<float>(count); | ||
| 353 | + MicroAPI::DataCopy<ElementPerTokenScale, MicroAPI::LoadDist::DIST_UNPACK4_B8>( | ||
| 354 | + vRegFp8, srcPtr + i * ELE_NUM_PER_REPEAT); | ||
| 355 | + MicroAPI::Cast<bfloat16_t, ElementPerTokenScale, castTraitFp8ToBf16>(vRegBf16, vRegFp8, maskAll); | ||
| 356 | + MicroAPI::Cast<float, bfloat16_t, castTraitBf16ToFp32>(vRegFp32, vRegBf16, maskAll); | ||
| 357 | + MicroAPI::StoreAlign<float, MicroAPI::StoreDist::DIST_NORM_B32>( | ||
| 358 | + dstPtr + i * ELE_NUM_PER_REPEAT * 4, vRegFp32, maskAll); | ||
| 359 | + } | ||
| 360 | + } | ||
| 361 | +}; | ||
| 362 | +} // namespace Catlass::Epilogue::Block | ||
| 363 | + | ||
| @@ -1,12 +1,17 @@ | |||
| 1 | /** | 1 | /** |
| 2 | - * Copyright (c) 2025 Huawei Technologies Co., Ltd. | 2 | + * Copyright (c) 2025-2026 Huawei Technologies Co., Ltd. |
| 3 | - * This program is free software, you can redistribute it and/or modify it under the terms and conditions of | 3 | + * This program is free software, you can redistribute it |
| 4 | - * CANN Open Software License Agreement Version 2.0 (the "License"). | 4 | + * and/or modify it under the terms and conditions of |
| 5 | - * Please refer to the License for details. You may not use this file except in compliance with the License. | 5 | + * CANN Open Software License Agreement Version 2.0 (the |
| 6 | - * THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED, | 6 | + * "License"). |
| 7 | - * INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE. | 7 | + * Please refer to the License for details. You may not use this file except in compliance with the |
| 8 | - * See LICENSE in the root of the software repository for the full text of the License. | 8 | + * License. |
| 9 | - */ | 9 | + * THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR |
| 10 | + * IMPLIED, | ||
| 11 | + * INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE. | ||
| 12 | + * | ||
| 13 | + * See LICENSE in the root of the software repository for the full text of the License. | ||
| 14 | + */ | ||
| 10 | 15 | ||
| 11 | 16 | ||
| 12 | 17 | ||
| @@ -176,6 +181,13 @@ struct EpilogueAscend950PerTokenDequantTla { | |||
| 176 | static constexpr uint32_t UB_STAGES = UB_STAGES_; | 181 | static constexpr uint32_t UB_STAGES = UB_STAGES_; |
| 177 | }; | 182 | }; |
| 178 | 183 | ||
| 184 | +// For Ascend950, FP8 per-token/per-channel dequant (epilogue-based, AIC/AIV pipeline) | ||
| 185 | +template <uint32_t UB_STAGES_> | ||
| 186 | +struct EpilogueAscend950Fp8PerTokenPerChannelDequant { | ||
| 187 | + using ArchTag = Arch::Ascend950; | ||
| 188 | + static constexpr uint32_t UB_STAGES = UB_STAGES_; | ||
| 189 | +}; | ||
| 190 | + | ||
| 179 | // For Ascend950, perGroup + perBlock dequant | 191 | // For Ascend950, perGroup + perBlock dequant |
| 180 | struct BlockEpiloguePertile { | 192 | struct BlockEpiloguePertile { |
| 181 | using ArchTag = Arch::Ascend950; | 193 | using ArchTag = Arch::Ascend950; |
| @@ -273,6 +285,12 @@ struct EpilogueFARescaleO { | |||
| 273 | using ArchTag = Arch::Ascend950; | 285 | using ArchTag = Arch::Ascend950; |
| 274 | }; | 286 | }; |
| 275 | 287 | ||
| 288 | +template <uint32_t UB_STAGES_> | ||
| 289 | +struct EpilogueAscend950PerTokenPerChannelQuant { | ||
| 290 | + using ArchTag = Arch::Ascend950; | ||
| 291 | + static constexpr uint32_t UB_STAGES = UB_STAGES_; | ||
| 292 | +}; | ||
| 293 | + | ||
| 276 | struct EpilogueElemWiseNoSourceFromUB { | 294 | struct EpilogueElemWiseNoSourceFromUB { |
| 277 | using ArchTag = Arch::Ascend950; | 295 | using ArchTag = Arch::Ascend950; |
| 278 | // Number of operands. Including Src, Dst 2 operands | 296 | // Number of operands. Including Src, Dst 2 operands |
| @@ -1,12 +1,17 @@ | |||
| 1 | /** | 1 | /** |
| 2 | - * Copyright (c) 2025 Huawei Technologies Co., Ltd. | 2 | + * Copyright (c) 2025-2026 Huawei Technologies Co., Ltd. |
| 3 | - * This program is free software, you can redistribute it and/or modify it under the terms and conditions of | 3 | + * This program is free software, you can redistribute it |
| 4 | - * CANN Open Software License Agreement Version 2.0 (the "License"). | 4 | + * and/or modify it under the terms and conditions of |
| 5 | - * Please refer to the License for details. You may not use this file except in compliance with the License. | 5 | + * CANN Open Software License Agreement Version 2.0 (the |
| 6 | - * THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED, | 6 | + * "License"). |
| 7 | - * INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE. | 7 | + * Please refer to the License for details. You may not use this file except in compliance with the |
| 8 | - * See LICENSE in the root of the software repository for the full text of the License. | 8 | + * License. |
| 9 | - */ | 9 | + * THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR |
| 10 | + * IMPLIED, | ||
| 11 | + * INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE. | ||
| 12 | + * | ||
| 13 | + * See LICENSE in the root of the software repository for the full text of the License. | ||
| 14 | + */ | ||
| 10 | 15 | ||
| 11 | 16 | ||
| 12 | 17 | ||
| @@ -198,7 +203,7 @@ struct CopyGm2Ub<Arch::Ascend950, Gemm::GemmType<Element, layout::VectorLayout>> | |||
| 198 | { | 203 | { |
| 199 | AscendC::DataCopyExtParams dataCopyParams(1, layoutSrc.shape(0) * sizeof(Element), 0, 0, 0); | 204 | AscendC::DataCopyExtParams dataCopyParams(1, layoutSrc.shape(0) * sizeof(Element), 0, 0, 0); |
| 200 | if constexpr (AscendC::Std::is_one_of_v< | 205 | if constexpr (AscendC::Std::is_one_of_v< |
| 201 | - Element, float8_e4m3_t, float8_e5m2_t, float4_e2m1x2_t, float4_e1m2x2_t>) { | 206 | + Element, float8_e4m3_t, float8_e5m2_t, float8_e8m0_t, float4_e2m1x2_t, float4_e1m2x2_t>) { |
| 202 | AscendC::DataCopyPadExtParams<uint8_t> padParams(false, 0, 0, 0); | 207 | AscendC::DataCopyPadExtParams<uint8_t> padParams(false, 0, 0, 0); |
| 203 | AscendC::DataCopyPad( | 208 | AscendC::DataCopyPad( |
| 204 | dstTensor.template ReinterpretCast<uint8_t>(), srcTensor.template ReinterpretCast<uint8_t>(), | 209 | dstTensor.template ReinterpretCast<uint8_t>(), srcTensor.template ReinterpretCast<uint8_t>(), |
| @@ -1,12 +1,17 @@ | |||
| 1 | /** | 1 | /** |
| 2 | - * Copyright (c) 2025 Huawei Technologies Co., Ltd. | 2 | + * Copyright (c) 2025-2026 Huawei Technologies Co., Ltd. |
| 3 | - * This program is free software, you can redistribute it and/or modify it under the terms and conditions of | 3 | + * This program is free software, you can redistribute it |
| 4 | - * CANN Open Software License Agreement Version 2.0 (the "License"). | 4 | + * and/or modify it under the terms and conditions of |
| 5 | - * Please refer to the License for details. You may not use this file except in compliance with the License. | 5 | + * CANN Open Software License Agreement Version 2.0 (the |
| 6 | - * THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED, | 6 | + * "License"). |
| 7 | - * INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE. | 7 | + * Please refer to the License for details. You may not use this file except in compliance with the |
| 8 | - * See LICENSE in the root of the software repository for the full text of the License. | 8 | + * License. |
| 9 | - */ | 9 | + * THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR |
| 10 | + * IMPLIED, | ||
| 11 | + * INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE. | ||
| 12 | + * | ||
| 13 | + * See LICENSE in the root of the software repository for the full text of the License. | ||
| 14 | + */ | ||
| 10 | 15 | ||
| 11 | 16 | ||
| 12 | 17 | ||
| @@ -167,6 +172,7 @@ struct CopyUb2Gm<Arch::Ascend950, Gemm::GemmType<Element, layout::VectorLayout>> | |||
| 167 | } | 172 | } |
| 168 | }; | 173 | }; |
| 169 | }; | 174 | }; |
| 175 | + | ||
| 170 | 176 | ||
| 171 | 177 | ||
| 172 | } // namespace Catlass::Epilogue::Tile | 178 | } // namespace Catlass::Epilogue::Tile |
| @@ -1,12 +1,17 @@ | |||
| 1 | /** | 1 | /** |
| 2 | - * Copyright (c) 2025 Huawei Technologies Co., Ltd. | 2 | + * Copyright (c) 2025-2026 Huawei Technologies Co., Ltd. |
| 3 | - * This program is free software, you can redistribute it and/or modify it under the terms and conditions of | 3 | + * This program is free software, you can redistribute it |
| 4 | - * CANN Open Software License Agreement Version 2.0 (the "License"). | 4 | + * and/or modify it under the terms and conditions of |
| 5 | - * Please refer to the License for details. You may not use this file except in compliance with the License. | 5 | + * CANN Open Software License Agreement Version 2.0 (the |
| 6 | - * THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED, | 6 | + * "License"). |
| 7 | - * INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE. | 7 | + * Please refer to the License for details. You may not use this file except in compliance with the |
| 8 | - * See LICENSE in the root of the software repository for the full text of the License. | 8 | + * License. |
| 9 | - */ | 9 | + * THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR |
| 10 | + * IMPLIED, | ||
| 11 | + * INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE. | ||
| 12 | + * | ||
| 13 | + * See LICENSE in the root of the software repository for the full text of the License. | ||
| 14 | + */ | ||
| 10 | 15 | ||
| 11 | 16 | ||
| 12 | 17 | ||
| @@ -144,6 +144,7 @@ struct BlockPrologue { | |||
| 144 | 144 | ||
| 145 | 145 | ||
| 146 | 146 | ||
| 147 | + | ||
| 147 | 148 | ||
| 148 | 149 | ||
| 149 | 150 | ||
| @@ -0,0 +1,623 @@ | |||
| 1 | +/** | ||
| 2 | + * This program is free software, you can redistribute it and/or modify. | ||
| 3 | + * Copyright (c) 2026 Huawei Technologies Co., Ltd. | ||
| 4 | + * This file is a part of the CANN Open Software. | ||
| 5 | + * Licensed under CANN Open Software License Agreement Version 2.0 (the "License"). | ||
| 6 | + * Please refer to the License for details. You may not use this file except in compliance with the License. | ||
| 7 | + * THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED, INCLUDING | ||
| 8 | + * BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE. See LICENSE in the root of | ||
| 9 | + * the software repository for the full text of the License. | ||
| 10 | + */ | ||
| 11 | + | ||
| 12 | + | ||
| 13 | + | ||
| 14 | + | ||
| 15 | + | ||
| 16 | + | ||
| 17 | + | ||
| 18 | + | ||
| 19 | + | ||
| 20 | + | ||
| 21 | + | ||
| 22 | + | ||
| 23 | + | ||
| 24 | + | ||
| 25 | +namespace Catlass::Gemm::Block { | ||
| 26 | + | ||
| 27 | +template < | ||
| 28 | + class ArchTag_, bool ENABLE_UNIT_FLAG_, bool USE_HF32_MODE_, uint32_t L0C_STAGES_, bool ENABLE_L1_RESIDENT_, | ||
| 29 | + uint32_t L1A_STAGES_, uint32_t L1B_STAGES_, uint32_t L0A_STAGES_, uint32_t L0B_STAGES_, class L1TileShape_, | ||
| 30 | + class L0TileShape_, class ElementA_, class ElementB_, class ElementC_, class ElementBias_, class TileCopy_, | ||
| 31 | + class TileMmad_> | ||
| 32 | +struct BlockMmadTla< | ||
| 33 | + MmadPingpongPreLoad< | ||
| 34 | + ArchTag_, ENABLE_UNIT_FLAG_, USE_HF32_MODE_, L0C_STAGES_, ENABLE_L1_RESIDENT_, L1A_STAGES_, L1B_STAGES_, | ||
| 35 | + L0A_STAGES_, L0B_STAGES_>, | ||
| 36 | + L1TileShape_, L0TileShape_, ElementA_, ElementB_, ElementC_, ElementBias_, TileCopy_, TileMmad_> { | ||
| 37 | +public: | ||
| 38 | + // Type Aliases | ||
| 39 | + using DispatchPolicy = MmadPingpongPreLoad< | ||
| 40 | + ArchTag_, ENABLE_UNIT_FLAG_, USE_HF32_MODE_, L0C_STAGES_, ENABLE_L1_RESIDENT_, L1A_STAGES_, L1B_STAGES_, | ||
| 41 | + L0A_STAGES_, L0B_STAGES_>; | ||
| 42 | + using ArchTag = typename DispatchPolicy::ArchTag; | ||
| 43 | + using TileCopy = TileCopy_; | ||
| 44 | + using L1TileShape = L1TileShape_; | ||
| 45 | + using L0TileShape = L0TileShape_; | ||
| 46 | + using ElementA = ElementA_; | ||
| 47 | + using LayoutA = typename TileCopy::LayoutA; | ||
| 48 | + using ElementB = ElementB_; | ||
| 49 | + using LayoutB = typename TileCopy::LayoutB; | ||
| 50 | + using ElementC = ElementC_; | ||
| 51 | + using LayoutC = typename TileCopy::LayoutC; | ||
| 52 | + using ElementBias = ElementBias_; | ||
| 53 | + | ||
| 54 | + using TileMmad = TileMmad_; | ||
| 55 | + | ||
| 56 | + using CopyL1ToL0A = typename TileCopy::CopyL1ToL0A; | ||
| 57 | + using CopyL1ToL0B = typename TileCopy::CopyL1ToL0B; | ||
| 58 | + using CopyL1ToBT = typename TileCopy::CopyL1ToBT; | ||
| 59 | + | ||
| 60 | + using ElementAccumulator = typename TileCopy::ElementAccumulator; | ||
| 61 | + | ||
| 62 | + static constexpr bool HAS_BIAS = TileCopy::HAS_BIAS; | ||
| 63 | + | ||
| 64 | + using LayoutTagL1A = typename TileCopy::LayoutTagL1A; | ||
| 65 | + using LayoutTagL1B = typename TileCopy::LayoutTagL1B; | ||
| 66 | + using LayoutTagL0A = typename TileCopy::LayoutTagL0A; | ||
| 67 | + using LayoutTagL0B = typename TileCopy::LayoutTagL0B; | ||
| 68 | + | ||
| 69 | + static_assert( | ||
| 70 | + tla::is_tuple<L1TileShape>::value && tla::is_static<L1TileShape>::value, | ||
| 71 | + "L1TileShape must be tla::tuple and static!"); | ||
| 72 | + static_assert( | ||
| 73 | + tla::is_tuple<L0TileShape>::value && tla::is_static<L0TileShape>::value, | ||
| 74 | + "L0TileShape must be tla::tuple and static!"); | ||
| 75 | + | ||
| 76 | + static constexpr bool ENABLE_UNIT_FLAG = DispatchPolicy::ENABLE_UNIT_FLAG; | ||
| 77 | + static constexpr bool USE_HF32_MODE = DispatchPolicy::USE_HF32_MODE; | ||
| 78 | + static constexpr bool ENABLE_L1_RESIDENT = DispatchPolicy::ENABLE_L1_RESIDENT; | ||
| 79 | + static constexpr uint32_t L1A_STAGES = DispatchPolicy::L1A_STAGES; | ||
| 80 | + static constexpr uint32_t L1B_STAGES = DispatchPolicy::L1B_STAGES; | ||
| 81 | + static constexpr uint32_t L0A_STAGES = DispatchPolicy::L0A_STAGES; | ||
| 82 | + static constexpr uint32_t L0B_STAGES = DispatchPolicy::L0B_STAGES; | ||
| 83 | + static constexpr uint32_t L0C_STAGES = DispatchPolicy::L0C_STAGES; | ||
| 84 | + static constexpr uint32_t L1_TILE_M = tla::get<0>(L1TileShape{}); | ||
| 85 | + static constexpr uint32_t L1_TILE_N = tla::get<1>(L1TileShape{}); | ||
| 86 | + static constexpr uint32_t L1_TILE_K = tla::get<2>(L1TileShape{}); | ||
| 87 | + static constexpr uint32_t L0_TILE_M = tla::get<0>(L0TileShape{}); | ||
| 88 | + static constexpr uint32_t L0_TILE_N = tla::get<1>(L0TileShape{}); | ||
| 89 | + static constexpr uint32_t L0_TILE_K = tla::get<2>(L0TileShape{}); | ||
| 90 | + | ||
| 91 | + // L1 tile size | ||
| 92 | + static constexpr uint32_t L1A_TILE_SIZE = L1_TILE_M * L1_TILE_K * sizeof(ElementA); | ||
| 93 | + static constexpr uint32_t L1B_TILE_SIZE = L1_TILE_N * L1_TILE_K * sizeof(ElementB); | ||
| 94 | + // L0 tile size | ||
| 95 | + static constexpr uint32_t L0A_TILE_SIZE = L0_TILE_M * L0_TILE_K * sizeof(ElementA); | ||
| 96 | + static constexpr uint32_t L0B_TILE_SIZE = L0_TILE_K * L0_TILE_N * sizeof(ElementB); | ||
| 97 | + static constexpr uint32_t L0C_TILE_SIZE = L1_TILE_M * L1_TILE_N * sizeof(ElementAccumulator); | ||
| 98 | + | ||
| 99 | + // Check HF32_MODE | ||
| 100 | + static_assert( | ||
| 101 | + !USE_HF32_MODE || (USE_HF32_MODE && std::is_same_v<ElementA, float> && std::is_same_v<ElementB, float>), | ||
| 102 | + "HF32 MODE only supports in float!"); | ||
| 103 | + | ||
| 104 | + // Check L0C_STAGES | ||
| 105 | + static_assert(!(ENABLE_UNIT_FLAG && L0C_STAGES != 1), "L0C_STAGES must be 1 when UnitFlag is true!"); | ||
| 106 | + | ||
| 107 | + // Check LayoutC | ||
| 108 | + static_assert( | ||
| 109 | + tla::detail::isRowMajor<LayoutC>::value || | ||
| 110 | + ((std::is_same_v<ElementC, half> || std::is_same_v<ElementC, bfloat16_t> || | ||
| 111 | + std::is_same_v<ElementC, float>) && | ||
| 112 | + tla::detail::iszN<ElementC, LayoutC>::value), | ||
| 113 | + "LayoutC only supports zN in half or bfloat16 or float, RowMajor in all dtype yet!"); | ||
| 114 | + | ||
| 115 | + // Check L1TileShape | ||
| 116 | + static_assert( | ||
| 117 | + L1A_TILE_SIZE * L1A_STAGES + L1B_TILE_SIZE * L1B_STAGES <= ArchTag::L1_SIZE, | ||
| 118 | + "L1TileShape exceeding the L1 space!"); | ||
| 119 | + | ||
| 120 | + // Check L0TileShape | ||
| 121 | + static_assert(L0A_TILE_SIZE * L0A_STAGES <= ArchTag::L0A_SIZE, "L0TileShape exceeding the L0A space!"); | ||
| 122 | + static_assert(L0B_TILE_SIZE * L0B_STAGES <= ArchTag::L0B_SIZE, "L0TileShape exceeding the L0B space!"); | ||
| 123 | + static_assert(L0C_TILE_SIZE * L0C_STAGES <= ArchTag::L0C_SIZE, "L0TileShape exceeding the L0C space!"); | ||
| 124 | + | ||
| 125 | + static_assert( | ||
| 126 | + L1_TILE_M == L0_TILE_M && L1_TILE_N == L0_TILE_N, | ||
| 127 | + "The situation where the basic blocks of L1 and L0 differ on the m and n axes is not supported yet"); | ||
| 128 | + static_assert(L0_TILE_K <= L1_TILE_K, "L0TileShape::K cannot exceed L1TileShape::K"); | ||
| 129 | + | ||
| 130 | + static_assert( | ||
| 131 | + (!HAS_BIAS && (L1A_STAGES + L1B_STAGES) <= 8) || (HAS_BIAS && (L1A_STAGES + L1B_STAGES) <= 7), | ||
| 132 | + "L1 Buffer overflow: Exceeds the supported range of EVENT(0~7)"); | ||
| 133 | + | ||
| 134 | + static_assert( | ||
| 135 | + (!HAS_BIAS && (L0A_STAGES + L0B_STAGES) <= 8) || (HAS_BIAS && (L0A_STAGES + L0B_STAGES) <= 7), | ||
| 136 | + "L0 Buffer overflow: Exceeds the supported range of EVENT_ID(0~7)"); | ||
| 137 | + | ||
| 138 | + static constexpr auto L1A_LAYOUT = | ||
| 139 | + tla::MakeLayout<ElementA, LayoutTagL1A>(tla::Int<L1_TILE_M>{}, tla::Int<L1_TILE_K>{}); | ||
| 140 | + static constexpr auto L1B_LAYOUT = | ||
| 141 | + tla::MakeLayout<ElementB, LayoutTagL1B>(tla::Int<L1_TILE_K>{}, tla::Int<L1_TILE_N>{}); | ||
| 142 | + static constexpr auto L1BIAS_LAYOUT = tla::MakeLayout(tla::Int<L1_TILE_N>{}); | ||
| 143 | + static constexpr auto L0BIAS_LAYOUT = tla::MakeLayout(tla::Int<L0_TILE_N>{}); | ||
| 144 | + | ||
| 145 | + // When enableing L1 resident mode, restore the pointer and coordinates that record the last state | ||
| 146 | + // to the initial state. if tow blockmmad instances need to be consecutively invoked at the kernel layer, | ||
| 147 | + // RestoreStatus() must be inserted between them. | ||
| 148 | + CATLASS_DEVICE | ||
| 149 | + void RestoreStatus() | ||
| 150 | + { | ||
| 151 | + for (int i = 0; i < L1A_STAGES; ++i) { | ||
| 152 | + lastAddrA[i] = nullptr; | ||
| 153 | + lastCoordA[i] = MatrixCoord{0U, 0U}; | ||
| 154 | + } | ||
| 155 | + for (int i = 0; i < L1B_STAGES; ++i) { | ||
| 156 | + lastAddrB[i] = nullptr; | ||
| 157 | + lastCoordB[i] = MatrixCoord{0U, 0U}; | ||
| 158 | + } | ||
| 159 | + } | ||
| 160 | + | ||
| 161 | + /// Construct | ||
| 162 | + CATLASS_DEVICE | ||
| 163 | + BlockMmadTla(Arch::Resource<ArchTag>& resource, uint32_t l1BufAddrStart = 0) | ||
| 164 | + { | ||
| 165 | + if ASCEND_IS_AIC { | ||
| 166 | + // use HF32 when USE_HF32_MODE is true | ||
| 167 | + if constexpr (USE_HF32_MODE) { | ||
| 168 | + AscendC::SetHF32Mode(true); | ||
| 169 | + } else { | ||
| 170 | + AscendC::SetHF32Mode(false); | ||
| 171 | + } | ||
| 172 | + if constexpr (ENABLE_UNIT_FLAG && tla::detail::isRowMajor<LayoutC>::value) { | ||
| 173 | + AscendC::SetMMLayoutTransform(true); | ||
| 174 | + } | ||
| 175 | + uint32_t l1AOffset = l1BufAddrStart; | ||
| 176 | + uint32_t l1BOffset = l1BufAddrStart + L1A_TILE_SIZE * L1A_STAGES; | ||
| 177 | + // Init buffers | ||
| 178 | + for (uint32_t i = 0; i < L1A_STAGES; i++) { | ||
| 179 | + // Assign L1/L0A/L0B space for each stages | ||
| 180 | + l1ATensorList[i] = resource.l1Buf.template GetBufferByByte<ElementA>(l1AOffset + L1A_TILE_SIZE * i); | ||
| 181 | + // Assign event ID for each stages | ||
| 182 | + l1AEventList[i] = i; | ||
| 183 | + // The event id that needs to be set before the loop | ||
| 184 | + AscendC::SetFlag<AscendC::HardEvent::MTE1_MTE2>(l1AEventList[i]); | ||
| 185 | + } | ||
| 186 | + for (uint32_t i = 0; i < L1B_STAGES; i++) { | ||
| 187 | + // Assign L1/L0A/L0B space for each stages | ||
| 188 | + l1BTensorList[i] = resource.l1Buf.template GetBufferByByte<ElementB>(l1BOffset + L1B_TILE_SIZE * i); | ||
| 189 | + // Assign event ID for each stages | ||
| 190 | + l1BEventList[i] = i + L1A_STAGES; | ||
| 191 | + // The event id that needs to be set before the loop | ||
| 192 | + AscendC::SetFlag<AscendC::HardEvent::MTE1_MTE2>(l1BEventList[i]); | ||
| 193 | + } | ||
| 194 | + for (uint32_t i = 0; i < L0A_STAGES; i++) { | ||
| 195 | + // Assign L1/L0A/L0B space for each stages | ||
| 196 | + l0ATensorList[i] = resource.l0ABuf.template GetBufferByByte<ElementA>(L0A_TILE_SIZE * i); | ||
| 197 | + // Assign event ID for each stages | ||
| 198 | + l0AEventList[i] = i; | ||
| 199 | + // The event id that needs to be set before the loop | ||
| 200 | + AscendC::SetFlag<AscendC::HardEvent::M_MTE1>(l0AEventList[i]); | ||
| 201 | + } | ||
| 202 | + for (uint32_t i = 0; i < L0B_STAGES; i++) { | ||
| 203 | + // Assign L1/L0A/L0B space for each stages | ||
| 204 | + l0BTensorList[i] = resource.l0BBuf.template GetBufferByByte<ElementB>(L0B_TILE_SIZE * i); | ||
| 205 | + // Assign event ID for each stages | ||
| 206 | + l0BEventList[i] = i + L0A_STAGES; | ||
| 207 | + // The event id that needs to be set before the loop | ||
| 208 | + AscendC::SetFlag<AscendC::HardEvent::M_MTE1>(l0BEventList[i]); | ||
| 209 | + } | ||
| 210 | + if constexpr (!ENABLE_UNIT_FLAG) { | ||
| 211 | + for (uint32_t i = 0; i < L0C_STAGES; i++) { | ||
| 212 | + l0CTensorList[i] = resource.l0CBuf.template GetBufferByByte<ElementAccumulator>(L0C_TILE_SIZE * i); | ||
| 213 | + l0CEventList[i] = i; | ||
| 214 | + AscendC::SetFlag<AscendC::HardEvent::FIX_M>(l0CEventList[i]); | ||
| 215 | + } | ||
| 216 | + } else { | ||
| 217 | + l0CTensorList[0] = resource.l0CBuf.template GetBufferByByte<ElementAccumulator>(0); | ||
| 218 | + } | ||
| 219 | + if constexpr (HAS_BIAS) { | ||
| 220 | + uint32_t l1BiasOffset = l1BOffset + L1B_TILE_SIZE * L1B_STAGES; | ||
| 221 | + l1BiasTensor = resource.l1Buf.template GetBufferByByte<uint8_t>(l1BiasOffset); | ||
| 222 | + l0BiasTensor = resource.btBuf.template GetBufferByByte<ElementAccumulator>(0); | ||
| 223 | + AscendC::SetFlag<AscendC::HardEvent::MTE1_MTE2>(L1A_STAGES + L1B_STAGES); | ||
| 224 | + AscendC::SetFlag<AscendC::HardEvent::M_MTE1>(L0A_STAGES + L0B_STAGES); | ||
| 225 | + } | ||
| 226 | + | ||
| 227 | + if constexpr (ENABLE_L1_RESIDENT) { | ||
| 228 | + RestoreStatus(); | ||
| 229 | + } | ||
| 230 | + } | ||
| 231 | + } | ||
| 232 | + | ||
| 233 | + /// Destructor | ||
| 234 | + CATLASS_DEVICE | ||
| 235 | + ~BlockMmadTla() | ||
| 236 | + { | ||
| 237 | + if ASCEND_IS_AIC { | ||
| 238 | + if constexpr (USE_HF32_MODE) { | ||
| 239 | + AscendC::SetHF32Mode(false); | ||
| 240 | + } | ||
| 241 | + if constexpr (ENABLE_UNIT_FLAG && tla::detail::isRowMajor<LayoutC>::value) { | ||
| 242 | + AscendC::SetMMLayoutTransform(false); | ||
| 243 | + } | ||
| 244 | + for (uint32_t i = 0; i < L1A_STAGES; i++) { | ||
| 245 | + AscendC::WaitFlag<AscendC::HardEvent::MTE1_MTE2>(l1AEventList[i]); | ||
| 246 | + } | ||
| 247 | + for (uint32_t i = 0; i < L1B_STAGES; i++) { | ||
| 248 | + AscendC::WaitFlag<AscendC::HardEvent::MTE1_MTE2>(l1BEventList[i]); | ||
| 249 | + } | ||
| 250 | + for (uint32_t i = 0; i < L0A_STAGES; i++) { | ||
| 251 | + AscendC::WaitFlag<AscendC::HardEvent::M_MTE1>(l0AEventList[i]); | ||
| 252 | + } | ||
| 253 | + for (uint32_t i = 0; i < L0B_STAGES; i++) { | ||
| 254 | + AscendC::WaitFlag<AscendC::HardEvent::M_MTE1>(l0BEventList[i]); | ||
| 255 | + } | ||
| 256 | + if constexpr (!ENABLE_UNIT_FLAG) { | ||
| 257 | + for (uint32_t i = 0; i < L0C_STAGES; i++) { | ||
| 258 | + AscendC::WaitFlag<AscendC::HardEvent::FIX_M>(l0CEventList[i]); | ||
| 259 | + } | ||
| 260 | + } | ||
| 261 | + if constexpr (HAS_BIAS) { | ||
| 262 | + AscendC::WaitFlag<AscendC::HardEvent::MTE1_MTE2>(L1A_STAGES + L1B_STAGES); | ||
| 263 | + AscendC::WaitFlag<AscendC::HardEvent::M_MTE1>(L0A_STAGES + L0B_STAGES); | ||
| 264 | + } | ||
| 265 | + } | ||
| 266 | + } | ||
| 267 | + | ||
| 268 | + /// Perform a block-scoped matrix multiply-accumulate | ||
| 269 | + template <class TensorA, class TensorB, class TensorC, class TensorBias = EmptyClass> | ||
| 270 | + CATLASS_DEVICE void operator()( | ||
| 271 | + TensorA& tensorA, TensorB& tensorB, TensorC& tensorC, GemmCoord const& actualShape, | ||
| 272 | + TensorBias const& tensorBias = {}, Callback const& callbackBeforeFixpipe = {}, | ||
| 273 | + Callback const& callbackAfterFixpipe = {}) | ||
| 274 | + { | ||
| 275 | + // Check L1TileShape | ||
| 276 | + if constexpr (HAS_BIAS) { | ||
| 277 | + static constexpr uint32_t BIAS_BUF_SIZE = L0_TILE_N * sizeof(ElementAccumulator); | ||
| 278 | + static constexpr uint32_t L1BIAS_SIZE = L1_TILE_N * sizeof(ElementBias); | ||
| 279 | + static_assert( | ||
| 280 | + BIAS_BUF_SIZE <= ArchTag::BIAS_SIZE, "BIAS_BUF_SIZE exceeding the BT space! Reduce L0_TILE_N"); | ||
| 281 | + static_assert( | ||
| 282 | + L1A_TILE_SIZE * L1A_STAGES + L1B_TILE_SIZE * L1B_STAGES + L1BIAS_SIZE <= ArchTag::L1_SIZE, | ||
| 283 | + "L1TileShape exceeding the L1 space!"); | ||
| 284 | + } | ||
| 285 | + | ||
| 286 | + using CopyGmToL1A = typename TileCopy_::template CopyGmToL1A<TensorA>; | ||
| 287 | + using CopyGmToL1B = typename TileCopy_::template CopyGmToL1B<TensorB>; | ||
| 288 | + using CopyL0CToDst = typename TileCopy_::template CopyL0CToDst<TensorC>; | ||
| 289 | + CopyGmToL1A copyGmToL1A; | ||
| 290 | + CopyGmToL1B copyGmToL1B; | ||
| 291 | + CopyL0CToDst copyL0CToDst; | ||
| 292 | + | ||
| 293 | + uint32_t mBlockActual = actualShape.m(); | ||
| 294 | + uint32_t kBlockActual = actualShape.k(); | ||
| 295 | + uint32_t nBlockActual = actualShape.n(); | ||
| 296 | + | ||
| 297 | + uint32_t mL1Actual = mBlockActual; | ||
| 298 | + if constexpr (std::is_same_v<ArchTag, Arch::AtlasA2>) { | ||
| 299 | + // Avoid using the gemv mode in mmad | ||
| 300 | + if (mL1Actual == 1) { | ||
| 301 | + mL1Actual = 16; | ||
| 302 | + } | ||
| 303 | + } | ||
| 304 | + uint32_t nL1Actual = nBlockActual; | ||
| 305 | + | ||
| 306 | + auto layoutInL0C = tla::MakeLayoutL0C(mL1Actual, nL1Actual); | ||
| 307 | + auto tensorL0C = tla::MakeTensor(l0CTensorList[l0CListId], layoutInL0C, Arch::PositionL0C{}); | ||
| 308 | + auto tensorL0Bias = tla::MakeTensor(l0BiasTensor, L0BIAS_LAYOUT, Arch::PositionBias{}); | ||
| 309 | + | ||
| 310 | + uint32_t kL1Actual = min(kBlockActual, L1_TILE_K); | ||
| 311 | + // load first matrix A tile from GM to L1 | ||
| 312 | + AscendC::WaitFlag<AscendC::HardEvent::MTE1_MTE2>(l1AEventList[l1AListId]); | ||
| 313 | + auto tensorL1A = tla::MakeTensor(l1ATensorList[l1AListId], L1A_LAYOUT, Arch::PositionL1{}); | ||
| 314 | + auto tensorTileA = GetTileA(tensorA, 0, 0, mBlockActual, kL1Actual); | ||
| 315 | + if constexpr (ENABLE_L1_RESIDENT) { | ||
| 316 | + // If the currently loaded GM pointer and block coordinates are the same as the last loaded ones, | ||
| 317 | + // skip this loadding. | ||
| 318 | + if (lastAddrA[l1AListId] != tensorTileA.data().GetPhyAddr() || | ||
| 319 | + tla::get<0>(tensorTileA.coord()) != lastCoordA[l1AListId].row() || | ||
| 320 | + tla::get<1>(tensorTileA.coord()) != lastCoordA[l1AListId].column()) { | ||
| 321 | + copyGmToL1A(tensorL1A, tensorTileA); | ||
| 322 | + lastCoordA[l1AListId] = MatrixCoord{tla::get<0>(tensorTileA.coord()), tla::get<1>(tensorTileA.coord())}; | ||
| 323 | + lastAddrA[l1AListId] = const_cast<__gm__ typename AscendC::GlobalTensor<ElementA>::PrimType*>( | ||
| 324 | + tensorTileA.data().GetPhyAddr()); | ||
| 325 | + } | ||
| 326 | + } else { | ||
| 327 | + copyGmToL1A(tensorL1A, tensorTileA); | ||
| 328 | + } | ||
| 329 | + AscendC::SetFlag<AscendC::HardEvent::MTE2_MTE1>(l1AEventList[l1AListId]); | ||
| 330 | + | ||
| 331 | + // load first matrix B tile from GM to L1 | ||
| 332 | + AscendC::WaitFlag<AscendC::HardEvent::MTE1_MTE2>(l1BEventList[l1BListId]); | ||
| 333 | + auto tensorL1B = tla::MakeTensor(l1BTensorList[l1BListId], L1B_LAYOUT, Arch::PositionL1{}); | ||
| 334 | + auto tensorTileB = GetTile(tensorB, tla::MakeCoord(0, 0), tla::MakeShape(kL1Actual, nBlockActual)); | ||
| 335 | + if constexpr (ENABLE_L1_RESIDENT) { | ||
| 336 | + if (lastAddrB[l1BListId] != tensorTileB.data().GetPhyAddr() || | ||
| 337 | + tla::get<0>(tensorTileB.coord()) != lastCoordB[l1BListId].row() || | ||
| 338 | + tla::get<1>(tensorTileB.coord()) != lastCoordB[l1BListId].column()) { | ||
| 339 | + copyGmToL1B(tensorL1B, tensorTileB); | ||
| 340 | + lastCoordB[l1BListId] = MatrixCoord{tla::get<0>(tensorTileB.coord()), tla::get<1>(tensorTileB.coord())}; | ||
| 341 | + lastAddrB[l1BListId] = const_cast<__gm__ typename AscendC::GlobalTensor<ElementB>::PrimType*>( | ||
| 342 | + tensorTileB.data().GetPhyAddr()); | ||
| 343 | + } | ||
| 344 | + } else { | ||
| 345 | + copyGmToL1B(tensorL1B, tensorTileB); | ||
| 346 | + } | ||
| 347 | + AscendC::SetFlag<AscendC::HardEvent::MTE2_MTE1>(l1BEventList[l1BListId]); | ||
| 348 | + | ||
| 349 | + if constexpr (HAS_BIAS && !std::is_same_v<TensorBias, EmptyClass>) { | ||
| 350 | + using CopyGmToL1Bias = typename TileCopy::template CopyGmToL1Bias<TensorBias>; | ||
| 351 | + CopyGmToL1Bias copyGmToL1Bias; | ||
| 352 | + AscendC::WaitFlag<AscendC::HardEvent::MTE1_MTE2>(L1A_STAGES + L1B_STAGES); | ||
| 353 | + auto l1Bias = l1BiasTensor.template ReinterpretCast<ElementBias>(); | ||
| 354 | + auto tensorL1Bias = tla::MakeTensor(l1Bias, L1BIAS_LAYOUT, Arch::PositionL1{}); | ||
| 355 | + copyGmToL1Bias(tensorL1Bias, tensorBias); | ||
| 356 | + AscendC::SetFlag<AscendC::HardEvent::MTE2_MTE1>(L1A_STAGES + L1B_STAGES); | ||
| 357 | + } | ||
| 358 | + | ||
| 359 | + if constexpr (!ENABLE_UNIT_FLAG) { | ||
| 360 | + AscendC::WaitFlag<AscendC::HardEvent::FIX_M>(l0CEventList[l0CListId]); | ||
| 361 | + } | ||
| 362 | + | ||
| 363 | + uint32_t mL0Loop = CeilDiv<L0_TILE_M>(mL1Actual); | ||
| 364 | + uint32_t nL0Loop = CeilDiv<L0_TILE_N>(nL1Actual); | ||
| 365 | + | ||
| 366 | + // main loop | ||
| 367 | + uint32_t kL1Loop = CeilDiv<L1_TILE_K>(kBlockActual); | ||
| 368 | + for (uint32_t kL1Idx = 0; kL1Idx < kL1Loop; kL1Idx++) { | ||
| 369 | + uint32_t l1AListIdNext = (l1AListId + 1 < L1A_STAGES) ? (l1AListId + 1) : 0; | ||
| 370 | + uint32_t l1BListIdNext = (l1BListId + 1 < L1B_STAGES) ? (l1BListId + 1) : 0; | ||
| 371 | + uint32_t kL1ActualNext{0}; | ||
| 372 | + // preload next tile from GM to L1 | ||
| 373 | + if (kL1Idx < kL1Loop - 1) { | ||
| 374 | + uint32_t kL1IdxNext = kL1Idx + 1; | ||
| 375 | + kL1ActualNext = (kL1IdxNext < kL1Loop - 1) ? L1_TILE_K : (kBlockActual - kL1IdxNext * L1_TILE_K); | ||
| 376 | + | ||
| 377 | + // Get L1 tensor for next stage | ||
| 378 | + auto l1ATensor = l1ATensorList[l1AListIdNext]; | ||
| 379 | + auto l1BTensor = l1BTensorList[l1BListIdNext]; | ||
| 380 | + auto tensorL1A = tla::MakeTensor(l1ATensor, L1A_LAYOUT, Arch::PositionL1{}); | ||
| 381 | + auto tensorL1B = tla::MakeTensor(l1BTensor, L1B_LAYOUT, Arch::PositionL1{}); | ||
| 382 | + // Get GM tile for next stage | ||
| 383 | + auto tensorTileA = GetTileA(tensorA, 0, kL1IdxNext * L1_TILE_K, mBlockActual, kL1ActualNext); | ||
| 384 | + auto tensorTileB = GetTile( | ||
| 385 | + tensorB, tla::MakeCoord(kL1IdxNext * L1_TILE_K, 0), tla::MakeShape(kL1ActualNext, nBlockActual)); | ||
| 386 | + | ||
| 387 | + // load next matrix A tile from GM to L1 | ||
| 388 | + AscendC::WaitFlag<AscendC::HardEvent::MTE1_MTE2>(l1AEventList[l1AListIdNext]); | ||
| 389 | + if constexpr (ENABLE_L1_RESIDENT) { | ||
| 390 | + if (lastAddrA[l1AListIdNext] != tensorTileA.data().GetPhyAddr() || | ||
| 391 | + tla::get<0>(tensorTileA.coord()) != lastCoordA[l1AListIdNext].row() || | ||
| 392 | + tla::get<1>(tensorTileA.coord()) != lastCoordA[l1AListIdNext].column()) { | ||
| 393 | + copyGmToL1A(tensorL1A, tensorTileA); | ||
| 394 | + lastCoordA[l1AListIdNext] = | ||
| 395 | + MatrixCoord{tla::get<0>(tensorTileA.coord()), tla::get<1>(tensorTileA.coord())}; | ||
| 396 | + lastAddrA[l1AListIdNext] = | ||
| 397 | + const_cast<__gm__ typename AscendC::GlobalTensor<ElementA>::PrimType*>( | ||
| 398 | + tensorTileA.data().GetPhyAddr()); | ||
| 399 | + } | ||
| 400 | + } else { | ||
| 401 | + copyGmToL1A(tensorL1A, tensorTileA); | ||
| 402 | + } | ||
| 403 | + AscendC::SetFlag<AscendC::HardEvent::MTE2_MTE1>(l1AEventList[l1AListIdNext]); | ||
| 404 | + | ||
| 405 | + // load next matrix B tile from GM to L1 | ||
| 406 | + AscendC::WaitFlag<AscendC::HardEvent::MTE1_MTE2>(l1BEventList[l1BListIdNext]); | ||
| 407 | + if constexpr (ENABLE_L1_RESIDENT) { | ||
| 408 | + if (lastAddrB[l1BListIdNext] != tensorTileB.data().GetPhyAddr() || | ||
| 409 | + tla::get<0>(tensorTileB.coord()) != lastCoordB[l1BListIdNext].row() || | ||
| 410 | + tla::get<1>(tensorTileB.coord()) != lastCoordB[l1BListIdNext].column()) { | ||
| 411 | + copyGmToL1B(tensorL1B, tensorTileB); | ||
| 412 | + lastCoordB[l1BListIdNext] = | ||
| 413 | + MatrixCoord{tla::get<0>(tensorTileB.coord()), tla::get<1>(tensorTileB.coord())}; | ||
| 414 | + lastAddrB[l1BListIdNext] = | ||
| 415 | + const_cast<__gm__ typename AscendC::GlobalTensor<ElementB>::PrimType*>( | ||
| 416 | + tensorTileB.data().GetPhyAddr()); | ||
| 417 | + } | ||
| 418 | + } else { | ||
| 419 | + copyGmToL1B(tensorL1B, tensorTileB); | ||
| 420 | + } | ||
| 421 | + AscendC::SetFlag<AscendC::HardEvent::MTE2_MTE1>(l1BEventList[l1BListIdNext]); | ||
| 422 | + } | ||
| 423 | + | ||
| 424 | + // Get L1 tensor for current stage | ||
| 425 | + auto l1ATensor = l1ATensorList[l1AListId]; | ||
| 426 | + auto l1BTensor = l1BTensorList[l1BListId]; | ||
| 427 | + tensorL1A = tla::MakeTensor(l1ATensor, L1A_LAYOUT, Arch::PositionL1{}); | ||
| 428 | + tensorL1B = tla::MakeTensor(l1BTensor, L1B_LAYOUT, Arch::PositionL1{}); | ||
| 429 | + // Get the loop nums on L0 | ||
| 430 | + uint32_t kL0Loop = CeilDiv<L0_TILE_K>(kL1Actual); | ||
| 431 | + | ||
| 432 | + for (int mL0Idx = 0; mL0Idx < mL0Loop; mL0Idx++) { | ||
| 433 | + uint32_t mL0Actual = (mL0Idx < mL0Loop - 1) ? L0_TILE_M : (mL1Actual - mL0Idx * L0_TILE_M); | ||
| 434 | + | ||
| 435 | + for (int kL0Idx = 0; kL0Idx < kL0Loop; kL0Idx++) { | ||
| 436 | + uint32_t kL0Actual = (kL0Idx < kL0Loop - 1) ? L0_TILE_K : (kL1Actual - kL0Idx * L0_TILE_K); | ||
| 437 | + | ||
| 438 | + // Locate the current tile on L0A | ||
| 439 | + auto l0ATile = l0ATensorList[l0AListId]; | ||
| 440 | + auto layoutAInL0 = tla::MakeLayout<ElementA, LayoutTagL0A>(mL0Actual, kL0Actual); | ||
| 441 | + auto tensorL0A = tla::MakeTensor(l0ATile, layoutAInL0, Arch::PositionL0A{}); | ||
| 442 | + // Locate the current tile of matrix A on L1 | ||
| 443 | + auto tensorTileL1A = | ||
| 444 | + GetTileA(tensorL1A, mL0Idx * L0_TILE_M, kL0Idx * L0_TILE_K, mL0Actual, kL0Actual); | ||
| 445 | + | ||
| 446 | + AscendC::WaitFlag<AscendC::HardEvent::M_MTE1>(l0AEventList[l0AListId]); | ||
| 447 | + if ((mL0Idx == 0) && (kL0Idx == 0)) { | ||
| 448 | + AscendC::WaitFlag<AscendC::HardEvent::MTE2_MTE1>(l1AEventList[l1AListId]); | ||
| 449 | + } | ||
| 450 | + | ||
| 451 | + // Load current tile from L1 to L0A | ||
| 452 | + copyL1ToL0A(tensorL0A, tensorTileL1A); | ||
| 453 | + | ||
| 454 | + if ((mL0Idx == mL0Loop - 1) && (kL0Idx == kL0Loop - 1)) { | ||
| 455 | + AscendC::SetFlag<AscendC::HardEvent::MTE1_MTE2>(l1AEventList[l1AListId]); | ||
| 456 | + } | ||
| 457 | + | ||
| 458 | + bool initC = ((kL1Idx == 0) && (kL0Idx == 0)); | ||
| 459 | + for (int nL0Idx = 0; nL0Idx < nL0Loop; nL0Idx++) { | ||
| 460 | + uint32_t nL0Actual = (nL0Idx < nL0Loop - 1) ? L0_TILE_N : (nL1Actual - nL0Idx * L0_TILE_N); | ||
| 461 | + | ||
| 462 | + // Locate the current tile on L0B | ||
| 463 | + auto l0BTile = l0BTensorList[l0BListId]; | ||
| 464 | + auto layoutBInL0 = tla::MakeLayout<ElementB, LayoutTagL0B>(kL0Actual, nL0Actual); | ||
| 465 | + auto tensorL0B = tla::MakeTensor(l0BTile, layoutBInL0, Arch::PositionL0B{}); | ||
| 466 | + // Locate the current tile of matrix B on L1 | ||
| 467 | + auto tensorTileL1B = GetTile( | ||
| 468 | + tensorL1B, tla::MakeCoord(kL0Idx * L0_TILE_K, nL0Idx * L0_TILE_N), | ||
| 469 | + tla::MakeShape(kL0Actual, nL0Actual)); | ||
| 470 | + | ||
| 471 | + // Wait for mmad finished | ||
| 472 | + AscendC::WaitFlag<AscendC::HardEvent::M_MTE1>(l0BEventList[l0BListId]); | ||
| 473 | + // If the current tile is the first one on the k&n axis, wait for loading matrix B from GM to L1 | ||
| 474 | + if ((mL0Idx == 0) && (kL0Idx == 0) && (nL0Idx == 0)) { | ||
| 475 | + AscendC::WaitFlag<AscendC::HardEvent::MTE2_MTE1>(l1BEventList[l1BListId]); | ||
| 476 | + } | ||
| 477 | + | ||
| 478 | + // Load current tile from L1 to L0B | ||
| 479 | + copyL1ToL0B(tensorL0B, tensorTileL1B); | ||
| 480 | + | ||
| 481 | + // If the current tile is the last one on the k&n axis, notify to load matrix B from GM to L1 | ||
| 482 | + if ((mL0Idx == mL0Loop - 1) && (kL0Idx == kL0Loop - 1) && (nL0Idx == nL0Loop - 1)) { | ||
| 483 | + AscendC::SetFlag<AscendC::HardEvent::MTE1_MTE2>(l1BEventList[l1BListId]); | ||
| 484 | + } | ||
| 485 | + | ||
| 486 | + if constexpr (HAS_BIAS && !std::is_same_v<TensorBias, EmptyClass>) { | ||
| 487 | + if (initC) { | ||
| 488 | + if (nL0Idx == 0) { | ||
| 489 | + AscendC::WaitFlag<AscendC::HardEvent::MTE2_MTE1>(L1A_STAGES + L1B_STAGES); | ||
| 490 | + } | ||
| 491 | + AscendC::WaitFlag<AscendC::HardEvent::M_MTE1>(L0A_STAGES + L0B_STAGES); | ||
| 492 | + auto l1Bias = l1BiasTensor.template ReinterpretCast<ElementBias>(); | ||
| 493 | + auto tensorL1Bias = tla::MakeTensor(l1Bias, L1BIAS_LAYOUT, Arch::PositionL1{}); | ||
| 494 | + auto tensorTileL1Bias = GetTile( | ||
| 495 | + tensorL1Bias, tla::MakeCoord(nL0Idx * L0_TILE_N), tla::MakeShape(nL0Actual)); | ||
| 496 | + // Load bias to l0 biastable | ||
| 497 | + copyL1ToBT(tensorL0Bias, tensorTileL1Bias); | ||
| 498 | + if (nL0Idx == nL0Loop - 1) { | ||
| 499 | + AscendC::SetFlag<AscendC::HardEvent::MTE1_MTE2>(L1A_STAGES + L1B_STAGES); | ||
| 500 | + } | ||
| 501 | + } | ||
| 502 | + } | ||
| 503 | + | ||
| 504 | + // Notify to do mmad | ||
| 505 | + AscendC::SetFlag<AscendC::HardEvent::MTE1_M>(l0CEventList[l0CListId]); | ||
| 506 | + | ||
| 507 | + // Locate the current tile on L0C | ||
| 508 | + auto tensorTileL0C = GetTile( | ||
| 509 | + tensorL0C, tla::MakeCoord(mL0Idx * L0_TILE_M, nL0Idx * L0_TILE_N), | ||
| 510 | + tla::MakeShape(mL0Actual, nL0Actual)); | ||
| 511 | + | ||
| 512 | + // Compute the matrix multiplication on L0A and L0B and write the result to the accumulator | ||
| 513 | + // Wait for loading L0B | ||
| 514 | + AscendC::WaitFlag<AscendC::HardEvent::MTE1_M>(l0CEventList[l0CListId]); | ||
| 515 | + | ||
| 516 | + // If the unit flag is enabled, the unit flag is set according to the calculation progress | ||
| 517 | + uint8_t unitFlag = 0b00; | ||
| 518 | + if constexpr (ENABLE_UNIT_FLAG) { | ||
| 519 | + if ((kL1Idx == kL1Loop - 1) && (mL0Idx == mL0Loop - 1) && (kL0Idx == kL0Loop - 1) && | ||
| 520 | + (nL0Idx == nL0Loop - 1)) { | ||
| 521 | + unitFlag = 0b11; | ||
| 522 | + } else { | ||
| 523 | + unitFlag = 0b10; | ||
| 524 | + } | ||
| 525 | + } | ||
| 526 | + | ||
| 527 | + if constexpr (HAS_BIAS && !std::is_same_v<TensorBias, EmptyClass>) { | ||
| 528 | + if (initC) { | ||
| 529 | + tileMmad( | ||
| 530 | + tensorTileL0C, tensorL0A, tensorL0B, tensorL0Bias, mL0Actual, nL0Actual, kL0Actual, | ||
| 531 | + initC, unitFlag); | ||
| 532 | + AscendC::SetFlag<AscendC::HardEvent::M_MTE1>(L0A_STAGES + L0B_STAGES); | ||
| 533 | + } else { | ||
| 534 | + tileMmad( | ||
| 535 | + tensorTileL0C, tensorL0A, tensorL0B, mL0Actual, nL0Actual, kL0Actual, initC, | ||
| 536 | + unitFlag); | ||
| 537 | + } | ||
| 538 | + } else { | ||
| 539 | + tileMmad( | ||
| 540 | + tensorTileL0C, tensorL0A, tensorL0B, mL0Actual, nL0Actual, kL0Actual, initC, unitFlag); | ||
| 541 | + } | ||
| 542 | + | ||
| 543 | + // Notify to move the next L0B tile | ||
| 544 | + AscendC::SetFlag<AscendC::HardEvent::M_MTE1>(l0BEventList[l0BListId]); | ||
| 545 | + l0BListId = (l0BListId + 1 < L0B_STAGES) ? (l0BListId + 1) : 0; | ||
| 546 | + } | ||
| 547 | + AscendC::SetFlag<AscendC::HardEvent::M_MTE1>(l0AEventList[l0AListId]); | ||
| 548 | + l0AListId = (l0AListId + 1 < L0A_STAGES) ? (l0AListId + 1) : 0; | ||
| 549 | + } | ||
| 550 | + } | ||
| 551 | + l1AListId = l1AListIdNext; | ||
| 552 | + l1BListId = l1BListIdNext; | ||
| 553 | + kL1Actual = kL1ActualNext; | ||
| 554 | + } | ||
| 555 | + | ||
| 556 | + // copy block out | ||
| 557 | + if constexpr (!ENABLE_UNIT_FLAG) { | ||
| 558 | + AscendC::SetFlag<AscendC::HardEvent::M_FIX>(l0CEventList[l0CListId]); | ||
| 559 | + AscendC::WaitFlag<AscendC::HardEvent::M_FIX>(l0CEventList[l0CListId]); | ||
| 560 | + callbackBeforeFixpipe(); | ||
| 561 | + copyL0CToDst(tensorC, tensorL0C); | ||
| 562 | + callbackAfterFixpipe(); | ||
| 563 | + AscendC::SetFlag<AscendC::HardEvent::FIX_M>(l0CEventList[l0CListId]); | ||
| 564 | + l0CListId = (l0CListId + 1 < L0C_STAGES) ? (l0CListId + 1) : 0; | ||
| 565 | + } else { | ||
| 566 | + callbackBeforeFixpipe(); | ||
| 567 | + copyL0CToDst(tensorC, tensorL0C, 0b11); | ||
| 568 | + callbackAfterFixpipe(); | ||
| 569 | + } | ||
| 570 | + } | ||
| 571 | + | ||
| 572 | + template <class TensorC> | ||
| 573 | + CATLASS_DEVICE void SynchronizeBlock() | ||
| 574 | + {} | ||
| 575 | + | ||
| 576 | +protected: | ||
| 577 | + template <class TensorA> | ||
| 578 | + CATLASS_DEVICE auto GetTileA(TensorA& tensorA, uint32_t mIndex, uint32_t kIndex, uint32_t mSize, uint32_t kSize) | ||
| 579 | + { | ||
| 580 | + if constexpr (tla::detail::isVector<LayoutA>::value) { | ||
| 581 | + return GetTile(tensorA, tla::MakeCoord(kIndex), tla::MakeShape(kSize)); | ||
| 582 | + } else { | ||
| 583 | + return GetTile(tensorA, tla::MakeCoord(mIndex, kIndex), tla::MakeShape(mSize, kSize)); | ||
| 584 | + } | ||
| 585 | + } | ||
| 586 | + | ||
| 587 | + // Multi-stage tensors list | ||
| 588 | + AscendC::LocalTensor<ElementA> l1ATensorList[L1A_STAGES]; | ||
| 589 | + AscendC::LocalTensor<ElementB> l1BTensorList[L1B_STAGES]; | ||
| 590 | + AscendC::LocalTensor<ElementA> l0ATensorList[L0A_STAGES]; | ||
| 591 | + AscendC::LocalTensor<ElementB> l0BTensorList[L0B_STAGES]; | ||
| 592 | + AscendC::LocalTensor<ElementAccumulator> l0CTensorList[L0C_STAGES]; | ||
| 593 | + AscendC::LocalTensor<uint8_t> l1BiasTensor; | ||
| 594 | + AscendC::LocalTensor<ElementAccumulator> l0BiasTensor; | ||
| 595 | + | ||
| 596 | + // Multi-stage event id list | ||
| 597 | + int32_t l1AEventList[L1A_STAGES]; | ||
| 598 | + int32_t l1BEventList[L1B_STAGES]; | ||
| 599 | + int32_t l0AEventList[L0A_STAGES]; | ||
| 600 | + int32_t l0BEventList[L0B_STAGES]; | ||
| 601 | + int32_t l0CEventList[L0C_STAGES]; | ||
| 602 | + | ||
| 603 | + __gm__ typename AscendC::GlobalTensor<ElementA>::PrimType* lastAddrA[L1A_STAGES]; | ||
| 604 | + __gm__ typename AscendC::GlobalTensor<ElementB>::PrimType* lastAddrB[L1B_STAGES]; | ||
| 605 | + MatrixCoord lastCoordA[L1A_STAGES]; | ||
| 606 | + MatrixCoord lastCoordB[L1B_STAGES]; | ||
| 607 | + | ||
| 608 | + // The id of current stage | ||
| 609 | + uint32_t l1AListId{0}; | ||
| 610 | + uint32_t l1BListId{0}; | ||
| 611 | + uint32_t l0AListId{0}; | ||
| 612 | + uint32_t l0BListId{0}; | ||
| 613 | + uint32_t l0CListId{0}; | ||
| 614 | + | ||
| 615 | + TileMmad tileMmad; | ||
| 616 | + CopyL1ToL0A copyL1ToL0A; | ||
| 617 | + CopyL1ToL0B copyL1ToL0B; | ||
| 618 | + CopyL1ToBT copyL1ToBT; | ||
| 619 | +}; | ||
| 620 | + | ||
| 621 | +} // namespace Catlass::Gemm::Block | ||
| 622 | + | ||
| 623 | + | ||
| @@ -362,6 +362,21 @@ struct MmadPingpongMutex : public MmadBase<ArchTag_, false> { | |||
| 362 | static constexpr bool ENABLE_L1_RESIDENT = ENABLE_L1_RESIDENT_; | 362 | static constexpr bool ENABLE_L1_RESIDENT = ENABLE_L1_RESIDENT_; |
| 363 | }; | 363 | }; |
| 364 | 364 | ||
| 365 | +template < | ||
| 366 | + class ArchTag_, bool ENABLE_UNIT_FLAG_ = false, bool USE_HF32_MODE_ = false, uint32_t L0C_STAGES_ = 1, | ||
| 367 | + bool ENABLE_L1_RESIDENT_ = false, uint32_t L1A_STAGES_ = 2, uint32_t L1B_STAGES_ = 2, uint32_t L0A_STAGES_ = 2, | ||
| 368 | + uint32_t L0B_STAGES_ = 2> | ||
| 369 | +struct MmadPingpongPreLoad : public MmadBase<ArchTag_, true> { | ||
| 370 | + static constexpr uint32_t L1A_STAGES = L1A_STAGES_; | ||
| 371 | + static constexpr uint32_t L1B_STAGES = L1B_STAGES_; | ||
| 372 | + static constexpr uint32_t L0A_STAGES = L0A_STAGES_; | ||
| 373 | + static constexpr uint32_t L0B_STAGES = L0B_STAGES_; | ||
| 374 | + static constexpr uint32_t L0C_STAGES = L0C_STAGES_; | ||
| 375 | + static constexpr bool ENABLE_UNIT_FLAG = ENABLE_UNIT_FLAG_; | ||
| 376 | + static constexpr bool USE_HF32_MODE = USE_HF32_MODE_; | ||
| 377 | + static constexpr bool ENABLE_L1_RESIDENT = ENABLE_L1_RESIDENT_; | ||
| 378 | +}; | ||
| 379 | + | ||
| 365 | template <class ArchTag_, bool ENABLE_UNIT_FLAG_ = false> | 380 | template <class ArchTag_, bool ENABLE_UNIT_FLAG_ = false> |
| 366 | struct MmadPingpongSymmLeft : public MmadBase<ArchTag_, false> { | 381 | struct MmadPingpongSymmLeft : public MmadBase<ArchTag_, false> { |
| 367 | static constexpr uint32_t STAGES = 2; | 382 | static constexpr uint32_t STAGES = 2; |


copyright年份修改为2026,其他文件也检查下