已合并
在 experimental 中新增 Ascend950 FP4/FP8 量化矩阵乘 #791
Chen_HaoWen创建于 7月1日
在 experimental 中新增 Ascend950 FP4/FP8 量化矩阵乘 #791
已合并
Chen_HaoWen创建于 7月1日
共 38 个文件变更+4496-34
@@ -0,0 +1,13 @@
1+# ----------------------------------------------------------------------------
2+# This program is free software, you can redistribute it and/or modify.
longjihui
longjihuilongjihui8月8日

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

likedislike
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+ 
longjihui
longjihuilongjihui8月8日

如果这个样例的功能只是相对于54_ascend950_fp4_mx_matmul多了反量化乘的话,这个命名改成xx quant matmul

likedislike
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+#ifndef K_MAX_SHAPE_DIM
15+#define K_MAX_SHAPE_DIM 0
16+#endif
17+ 
18+#include <iostream>
19+#include <vector>
20+#include <cstdlib>
21+#include <cstring>
22+ 
23+#include "catlass/catlass.hpp"
24+#include "catlass/arch/arch.hpp"
25+#include "catlass/gemm/block/block_mmad.hpp"
26+#include "catlass/gemm/block/block_swizzle.hpp"
27+#include "catlass/gemm/dispatch_policy.hpp"
28+#include "catlass/gemm/gemm_type.hpp"
29+#include "catlass/gemm/kernel/mx_matmul_pertoken_perchannel_tla.hpp"
30+#include "catlass/layout/layout.hpp"
31+#include "catlass/status.hpp"
32+#include "catlass/gemm/device/device_gemm.hpp"
33+#include "tla/layout.hpp"
34+#include "helper.hpp"
35+#include "golden.hpp"
36+#include "catlass/epilogue/block/block_epilogue.hpp"
37+#include "catlass/epilogue/dispatch_policy.hpp"
38+#include "catlass/epilogue/tile/tile_broadcast_mul.hpp"
39+#include "catlass/epilogue/tile/tile_broadcast_one_blk.hpp"
40+#include "catlass/epilogue/tile/tile_swizzle.hpp"
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+@only_on_3510
38+@pytest.mark.parametrize("shape", QUANT_MATMUL_SHAPES)
39+@pytest.mark.parametrize("seed", range(10))
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+@only_on_3510
57+@pytest.mark.parametrize("input_index", range(6))
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+@only_on_3510
66+@pytest.mark.parametrize("axis", range(3))
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+ 
longjihui
longjihuilongjihui8月8日

这个样例只支持e4m3的话,建议参考29样例命名

likedislike
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+#include <iostream>
34+#include <vector>
35+#include <cstdlib>
36+#include <cstring>
37+ 
38+#include "catlass/arch/arch.hpp"
39+#include "catlass/catlass.hpp"
40+#include "catlass/gemm/block/block_mmad.hpp"
41+#include "catlass/gemm/block/block_swizzle.hpp"
42+#include "catlass/gemm/device/device_gemm.hpp"
43+#include "catlass/gemm/dispatch_policy.hpp"
44+#include "catlass/gemm/kernel/matmul_per_token_per_channel_epilogue_tla.hpp"
45+#include "catlass/gemm/tile/tile_copy.hpp"
46+#include "catlass/epilogue/block/block_epilogue.hpp"
47+#include "catlass/epilogue/dispatch_policy.hpp"
48+#include "catlass/epilogue/tile/tile_broadcast_mul.hpp"
49+#include "catlass/epilogue/tile/tile_broadcast_one_blk.hpp"
50+#include "catlass/epilogue/tile/tile_swizzle.hpp"
51+#include "catlass/layout/layout.hpp"
52+#include "catlass/matrix_coord.hpp"
53+#include "catlass/status.hpp"
54+#include "tla/layout.hpp"
55+ 
56+#include "golden.hpp"
57+#include "helper.hpp"
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)
@@ -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+@only_on_3510
20+@pytest.mark.parametrize("shape", QUANT_MATMUL_SHAPES + [(17, 33, 65)])
21+@pytest.mark.parametrize("seed", range(10))
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+@only_on_3510
37+@pytest.mark.parametrize("input_index", range(4))
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+@only_on_3510
46+@pytest.mark.parametrize("axis", range(3))
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#include "catlass/epilogue/block/block_epilogue_dequant.hpp"56#include "catlass/epilogue/block/block_epilogue_dequant.hpp"
57#include "catlass/epilogue/block/block_epilogue_dual_level_quant_mx.hpp"57#include "catlass/epilogue/block/block_epilogue_dual_level_quant_mx.hpp"
58#include "catlass/epilogue/block/block_epilogue_per_block_quant_tla.hpp"58#include "catlass/epilogue/block/block_epilogue_per_block_quant_tla.hpp"
59+#include "catlass/epilogue/block/block_epilogue_per_token_dequant_fp8_regbase.hpp"
60+#include "catlass/epilogue/block/block_epilogue_per_token_dequant_fp8_regbase_l0c2ub.hpp"
59#include "catlass/epilogue/block/block_epilogue_visitor.hpp"61#include "catlass/epilogue/block/block_epilogue_visitor.hpp"
60#include "catlass/epilogue/block/block_epilogue_swiglu_mx_quant.hpp"62#include "catlass/epilogue/block/block_epilogue_swiglu_mx_quant.hpp"
61#include "catlass/epilogue/block/block_epilogue_finalize_routing.hpp"63#include "catlass/epilogue/block/block_epilogue_finalize_routing.hpp"
@@ -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+#ifndef CATLASS_EPILOGUE_BLOCK_EPILOGUE_PER_TOKEN_DEQUANT_FP8_REGBASE_HPP
13+#define CATLASS_EPILOGUE_BLOCK_EPILOGUE_PER_TOKEN_DEQUANT_FP8_REGBASE_HPP
14+ 
15+#include "catlass/catlass.hpp"
16+#include "catlass/arch/resource.hpp"
17+#include "catlass/epilogue/dispatch_policy.hpp"
18+#include "catlass/gemm_coord.hpp"
19+#include "catlass/gemm/gemm_type.hpp"
20+#include "catlass/matrix_coord.hpp"
21+#include "catlass/layout/layout.hpp"
22+#include "catlass/detail/callback.hpp"
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+#endif // CATLASS_EPILOGUE_BLOCK_EPILOGUE_PER_TOKEN_DEQUANT_FP8_REGBASE_HPP
@@ -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+#ifndef CATLASS_EPILOGUE_BLOCK_EPILOGUE_PER_TOKEN_DEQUANT_FP8_REGBASE_L0C2UB_HPP
13+#define CATLASS_EPILOGUE_BLOCK_EPILOGUE_PER_TOKEN_DEQUANT_FP8_REGBASE_L0C2UB_HPP
14+ 
15+#include "catlass/catlass.hpp"
16+#include "catlass/arch/resource.hpp"
17+#include "catlass/epilogue/dispatch_policy.hpp"
18+#include "catlass/gemm_coord.hpp"
19+#include "catlass/gemm/gemm_type.hpp"
20+#include "catlass/matrix_coord.hpp"
21+#include "catlass/layout/layout.hpp"
22+#include "catlass/detail/callback.hpp"
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+#endif
@@ -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 of3+ * 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#ifndef CATLASS_EPILOGUE_DISPATCH_POLICY_HPP16#ifndef CATLASS_EPILOGUE_DISPATCH_POLICY_HPP
12#define CATLASS_EPILOGUE_DISPATCH_POLICY_HPP17#define CATLASS_EPILOGUE_DISPATCH_POLICY_HPP
@@ -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 dequant191// For Ascend950, perGroup + perBlock dequant
180struct BlockEpiloguePertile {192struct 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+ 
276struct EpilogueElemWiseNoSourceFromUB {294struct EpilogueElemWiseNoSourceFromUB {
277 using ArchTag = Arch::Ascend950;295 using ArchTag = Arch::Ascend950;
278 // Number of operands. Including Src, Dst 2 operands296 // 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 of3+ * 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#ifndef CATLASS_EPILOGUE_TILE_TILE_COPY_GM_TO_UB_HPP16#ifndef CATLASS_EPILOGUE_TILE_TILE_COPY_GM_TO_UB_HPP
12#define CATLASS_EPILOGUE_TILE_TILE_COPY_GM_TO_UB_HPP17#define CATLASS_EPILOGUE_TILE_TILE_COPY_GM_TO_UB_HPP
@@ -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 of3+ * 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#ifndef CATLASS_EPILOGUE_TILE_TILE_COPY_UB_TO_GM_HPP16#ifndef CATLASS_EPILOGUE_TILE_TILE_COPY_UB_TO_GM_HPP
12#define CATLASS_EPILOGUE_TILE_TILE_COPY_UB_TO_GM_HPP17#define CATLASS_EPILOGUE_TILE_TILE_COPY_UB_TO_GM_HPP
@@ -167,6 +172,7 @@ struct CopyUb2Gm<Arch::Ascend950, Gemm::GemmType<Element, layout::VectorLayout>>
167 }172 }
168 };173 };
169};174};
175+ 
170#endif // CATLASS_ARCH == 3510 || __NPU_ARCH__ == 3510176#endif // CATLASS_ARCH == 3510 || __NPU_ARCH__ == 3510
171 177 
172} // namespace Catlass::Epilogue::Tile178} // 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 of3+ * 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#ifndef CATLASS_EPILOGUE_TILE_TILE_COPY_HPP16#ifndef CATLASS_EPILOGUE_TILE_TILE_COPY_HPP
12#define CATLASS_EPILOGUE_TILE_TILE_COPY_HPP17#define CATLASS_EPILOGUE_TILE_TILE_COPY_HPP
@@ -144,6 +144,7 @@ struct BlockPrologue {
144#include "catlass/gemm/block/block_mmad_pingpong_dequant_tla.hpp"144#include "catlass/gemm/block/block_mmad_pingpong_dequant_tla.hpp"
145#include "catlass/gemm/block/block_mmad_pingpong_tla_v2.hpp"145#include "catlass/gemm/block/block_mmad_pingpong_tla_v2.hpp"
146#include "catlass/gemm/block/block_mmad_preload_tla.hpp"146#include "catlass/gemm/block/block_mmad_preload_tla.hpp"
147+#include "catlass/gemm/block/block_mmad_pingpong_preload_tla.hpp"
147#include "catlass/gemm/block/block_mmad_preload_async_with_callback_tla.hpp"148#include "catlass/gemm/block/block_mmad_preload_async_with_callback_tla.hpp"
148#include "catlass/gemm/block/block_mmad_fai_pv_tla.hpp"149#include "catlass/gemm/block/block_mmad_fai_pv_tla.hpp"
149#include "catlass/gemm/block/block_mmad_fai_qk_tla.hpp"150#include "catlass/gemm/block/block_mmad_fai_qk_tla.hpp"
@@ -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+#ifndef CATLASS_GEMM_BLOCK_BLOCK_MMAD_PINGPONG_PRELOAD_TLA_HPP
13+#define CATLASS_GEMM_BLOCK_BLOCK_MMAD_PINGPONG_PRELOAD_TLA_HPP
14+ 
15+#include "catlass/catlass.hpp"
16+#include "catlass/arch/resource.hpp"
17+#include "catlass/coord.hpp"
18+#include "catlass/detail/callback.hpp"
19+#include "catlass/gemm_coord.hpp"
20+#include "catlass/gemm/dispatch_policy.hpp"
21+#include "catlass/gemm/helper.hpp"
22+#include "tla/layout.hpp"
23+#include "tla/tensor.hpp"
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+#endif // CATLASS_GEMM_BLOCK_BLOCK_MMAD_PINGPONG_TLA_HPP
@@ -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+ 
365template <class ArchTag_, bool ENABLE_UNIT_FLAG_ = false>380template <class ArchTag_, bool ENABLE_UNIT_FLAG_ = false>
366struct MmadPingpongSymmLeft : public MmadBase<ArchTag_, false> {381struct MmadPingpongSymmLeft : public MmadBase<ArchTag_, false> {
367 static constexpr uint32_t STAGES = 2;382 static constexpr uint32_t STAGES = 2;