已合并
Upgrade the infrastructure to support future extensions #157
Maksim Vlasov创建于 8月11日
Upgrade the infrastructure to support future extensions #157
已合并
Maksim Vlasov创建于 8月11日
共 162 个文件变更+5279-1086
@@ -8,13 +8,12 @@ __pycache__/
8 8 
9# Distribution / packaging9# Distribution / packaging
10.Python10.Python
11-build/11+build*/
12develop-eggs/12develop-eggs/
13dist/13dist/
14downloads/14downloads/
15eggs/15eggs/
16.eggs/16.eggs/
17-lib/
18lib64/17lib64/
19parts/18parts/
20sdist/19sdist/
@@ -197,9 +196,9 @@ cython_debug/
197.abstra/196.abstra/
198 197 
199# Visual Studio Code198# Visual Studio Code
200-# Visual Studio Code specific template is maintained in a separate VisualStudioCode.gitignore 199+# Visual Studio Code specific template is maintained in a separate VisualStudioCode.gitignore
201# that can be found at https://github.com/github/gitignore/blob/main/Global/VisualStudioCode.gitignore200# that can be found at https://github.com/github/gitignore/blob/main/Global/VisualStudioCode.gitignore
202-# and can be added to the global gitignore or merged into this file. However, if you prefer, 201+# and can be added to the global gitignore or merged into this file. However, if you prefer,
203# you could uncomment the following to ignore the entire vscode folder202# you could uncomment the following to ignore the entire vscode folder
204# .vscode/203# .vscode/
205# Temporary file for partial code execution204# Temporary file for partial code execution
@@ -227,4 +226,4 @@ docs/python-api/rst/**/generated/
227.spec-workflow/226.spec-workflow/
228docs/superpowers/227docs/superpowers/
229AGENTS.md228AGENTS.md
230-CLAUDE.md229+CLAUDE.md
@@ -1,6 +1,18 @@
1repos:1repos:
2+ - repo: https://github.com/google/yapf
3+ rev: v0.43.0
4+ hooks:
5+ - id: yapf
6+ args: ["-p", "-i"]
7+ 
8+ - repo: https://github.com/astral-sh/ruff-pre-commit
9+ rev: v0.11.5
10+ hooks:
11+ - id: ruff
12+ files: ^python/asc/
13+ 
2 - repo: https://github.com/pre-commit/mirrors-clang-format14 - repo: https://github.com/pre-commit/mirrors-clang-format
3- rev: v16.0.015+ rev: v20.1.8
4 hooks:16 hooks:
5 - id: clang-format17 - id: clang-format
6 types_or: [c++, c]18 types_or: [c++, c]
@@ -14,9 +14,8 @@ if(POLICY CMP0116)
14 cmake_policy(SET CMP0116 OLD)14 cmake_policy(SET CMP0116 OLD)
15endif()15endif()
16 16 
17-project(AscIR LANGUAGES C CXX)17+project(AscIR LANGUAGES CXX)
18 18 
19-message(STATUS "C compiler: ${CMAKE_C_COMPILER}")
20message(STATUS "C++ compiler: ${CMAKE_CXX_COMPILER}")19message(STATUS "C++ compiler: ${CMAKE_CXX_COMPILER}")
21message(STATUS "Linker: ${CMAKE_LINKER}")20message(STATUS "Linker: ${CMAKE_LINKER}")
22message(STATUS "Selected build configuration: ${CMAKE_BUILD_TYPE}")21message(STATUS "Selected build configuration: ${CMAKE_BUILD_TYPE}")
@@ -95,7 +94,7 @@ if (ASCIR_COVERAGE)
95 add_link_options("-Wl,-z,nostart-stop-gc")94 add_link_options("-Wl,-z,nostart-stop-gc")
96 endif()95 endif()
97 elseif (CMAKE_CXX_COMPILER_ID MATCHES "GNU")96 elseif (CMAKE_CXX_COMPILER_ID MATCHES "GNU")
98- add_compile_options("-g" "-O0" "-fprofile-arcs" "-ftest-coverage" "--coverage")97+ add_compile_options("-g" "-O0" "-fprofile-arcs" "-ftest-coverage" "-fprofile-update=atomic" "--coverage")
99 add_link_options("-lgcov" "--coverage")98 add_link_options("-lgcov" "--coverage")
100 endif()99 endif()
101endif()100endif()
MOAT.xml+3-2
@@ -11,10 +11,11 @@
11 </policylist>11 </policylist>
12 <filefilterlist>12 <filefilterlist>
13 <filefilter name="defaultPolicyFilter" desc="Filters for compatibility,license header policies">13 <filefilter name="defaultPolicyFilter" desc="Filters for compatibility,license header policies">
14- <filteritem type="filename" name="OAT.xml" desc="Describe commite owner"/> 14+ <filteritem type="filename" name="OAT.xml" desc="Describe commite owner"/>
15 <filteritem type="filename" name=".gitmodules" desc="git submodule "/>15 <filteritem type="filename" name=".gitmodules" desc="git submodule "/>
16 <filteritem type="filename" name="pyproject.toml" desc="python config "/>16 <filteritem type="filename" name="pyproject.toml" desc="python config "/>
17 <filteritem type="filename" name="classify_rule.yaml" desc="test filter"/>17 <filteritem type="filename" name="classify_rule.yaml" desc="test filter"/>
18+ <filteritem type="filename" name=".pre-commit-config.yaml" desc="pre-commit tool configuration"/>
18 </filefilter>19 </filefilter>
19 <filefilter name="copyrightPolicyFilter" desc="Filters for copyright header policies">20 <filefilter name="copyrightPolicyFilter" desc="Filters for copyright header policies">
20 <filteritem type="filename" name=".gitmodules" desc="git submodule "/>21 <filteritem type="filename" name=".gitmodules" desc="git submodule "/>
@@ -44,4 +45,4 @@
44 </licensematcher>45 </licensematcher>
45 </licensematcherlist>46 </licensematcherlist>
46 </oatconfig>47 </oatconfig>
47-</configuration>48+</configuration>
@@ -76,7 +76,7 @@ mlir_check_all_link_libraries(${ASCIR_OPT_NAME})
76 76 
77set_target_properties(${ASCIR_OPT_NAME} PROPERTIES RUNTIME_OUTPUT_DIRECTORY "${CMAKE_BINARY_DIR}/bin")77set_target_properties(${ASCIR_OPT_NAME} PROPERTIES RUNTIME_OUTPUT_DIRECTORY "${CMAKE_BINARY_DIR}/bin")
78 78 
79-# ascir-translate 79+# ascir-translate
80 80 
81set(ASCIR_TRANSLATE_NAME ascir-translate)81set(ASCIR_TRANSLATE_NAME ascir-translate)
82 82 
@@ -159,16 +159,16 @@ def DataFormat : APIType<"DataFormat"> {
159 let apiName = "AscendC::DataFormat";159 let apiName = "AscendC::DataFormat";
160}160}
161 161 
162-def MatrixOffset : APIType<"MatrixOffset"> {
163- let mnemonic = "matrix_offset";
164- let apiName = "MatrixOffset";
165-}
166- 
167def DeqScale : APIType<"DeqScale"> {162def DeqScale : APIType<"DeqScale"> {
168 let mnemonic = "deq_scale";163 let mnemonic = "deq_scale";
169 let apiName = "AscendC::DeqScale";164 let apiName = "AscendC::DeqScale";
170}165}
171 166 
167+def Dn2NzParams : APIType<"Dn2NzParams"> {
168+ let mnemonic = "dn2nz_params";
169+ let apiName = "AscendC::Dn2NzParams";
170+}
171+ 
172def FixpipeConfig : APIType<"FixpipeConfig"> {172def FixpipeConfig : APIType<"FixpipeConfig"> {
173 let mnemonic = "fixpipe_config";173 let mnemonic = "fixpipe_config";
174 let apiName = "AscendC::FixpipeConfig";174 let apiName = "AscendC::FixpipeConfig";
@@ -229,6 +229,21 @@ def KfcServer : APIType<"KfcServer"> {
229 let apiName = "AscendC::KfcServer";229 let apiName = "AscendC::KfcServer";
230}230}
231 231 
232+def LayerNormTiling : APIType<"LayerNormTiling"> {
233+ let mnemonic = "layer_norm_tiling";
234+ let apiName = "LayerNormTiling";
235+}
236+ 
237+def LayerNormPara : APIType<"LayerNormPara"> {
238+ let mnemonic = "layer_norm_para";
239+ let apiName = "AscendC::LayerNormPara";
240+}
241+ 
242+def LayerNormSeparateTiling : APIType<"LayerNormSeparateTiling"> {
243+ let mnemonic = "layer_norm_separate_tiling";
244+ let apiName = "LayerNormSeparateTiling";
245+}
246+ 
232def ListTensorDesc : APIType<"ListTensorDesc"> {247def ListTensorDesc : APIType<"ListTensorDesc"> {
233 let mnemonic = "list_tensor_desc";248 let mnemonic = "list_tensor_desc";
234 let apiName = "AscendC::ListTensorDesc";249 let apiName = "AscendC::ListTensorDesc";
@@ -289,6 +304,11 @@ def LocalTensor : APIType {
289 let apiName = "AscendC::LocalTensor";304 let apiName = "AscendC::LocalTensor";
290}305}
291 306 
307+def LoopModeParams : APIType<"LoopModeParams"> {
308+ let mnemonic = "loop_mode_params";
309+ let apiName = "AscendC::LoopModeParams";
310+}
311+ 
292def MaskMode : APIType<"MaskMode"> {312def MaskMode : APIType<"MaskMode"> {
293 let mnemonic = "mask_mode";313 let mnemonic = "mask_mode";
294 let apiName = "AscendC::MaskMode";314 let apiName = "AscendC::MaskMode";
@@ -309,6 +329,11 @@ def MatmulConfig : APIType<"MatmulConfig"> {
309 let apiName = "MatmulConfig";329 let apiName = "MatmulConfig";
310}330}
311 331 
332+def MatrixOffset : APIType<"MatrixOffset"> {
333+ let mnemonic = "matrix_offset";
334+ let apiName = "MatrixOffset";
335+}
336+ 
312def MmadParams : APIType<"MmadParams"> {337def MmadParams : APIType<"MmadParams"> {
313 let mnemonic = "mmad_params";338 let mnemonic = "mmad_params";
314 let apiName = "AscendC::MmadParams";339 let apiName = "AscendC::MmadParams";
@@ -448,4 +473,5 @@ def VdeqInfo : APIType<"VdeqInfo"> {
448 let mnemonic = "vdeq_info";473 let mnemonic = "vdeq_info";
449 let apiName = "AscendC::VdeqInfo";474 let apiName = "AscendC::VdeqInfo";
450}475}
476+ 
451#endif // API_TYPES_TD477#endif // API_TYPES_TD
@@ -20,8 +20,9 @@ include "mlir/Interfaces/CastInterfaces.td"
20include "mlir/Interfaces/SideEffectInterfaces.td"20include "mlir/Interfaces/SideEffectInterfaces.td"
21include "mlir/IR/OpBase.td"21include "mlir/IR/OpBase.td"
22 22 
23-def AscendC_SimpleSoftMaxOp23+def AscendC_SimpleSoftMaxOp : VectorOp<"simple_softmax", "SimpleSoftMax", [
24- : APIOp<"simple_softmax", "SimpleSoftMax", [AttrSizedOperandSegments, OpWithDstInterface]> {24+ AttrSizedOperandSegments, OpWithDstInterface, OpWithSrcInterface
25+]> {
25 let arguments = (ins UnitAttr:$reuseSource, UnitAttr:$basicBlock,26 let arguments = (ins UnitAttr:$reuseSource, UnitAttr:$basicBlock,
26 UnitAttr:$dataFormatNZ, AscendC_LocalTensor:$dst,27 UnitAttr:$dataFormatNZ, AscendC_LocalTensor:$dst,
27 AscendC_LocalTensor:$sumTensor,28 AscendC_LocalTensor:$sumTensor,
@@ -30,17 +31,46 @@ def AscendC_SimpleSoftMaxOp
30 Optional<AscendC_LocalTensor>:$sharedTmpBuffer,31 Optional<AscendC_LocalTensor>:$sharedTmpBuffer,
31 AscendC_SoftMaxTiling:$tiling,32 AscendC_SoftMaxTiling:$tiling,
32 Optional<AnyType>:$softmaxShapeInfo);33 Optional<AnyType>:$softmaxShapeInfo);
34+ let extraClassDeclaration = extraClassDeclarationBase # [{
35+ SmallVector<Value> getDstTensors()
36+ {
37+ SmallVector<Value> dstTensors{getDst()};
38+ if (auto tensor = getSharedTmpBuffer())
39+ dstTensors.emplace_back(tensor);
40+ return dstTensors;
41+ }
42+ SmallVector<Value> getSrcTensors()
43+ {
44+ return {getSumTensor(), getMaxTensor(), getSrc()};
45+ }
46+ }];
33}47}
34 48 
35-def AscendC_SoftMaxOp : APIOp<"softmax", "SoftMax", [AttrSizedOperandSegments, OpWithDstInterface]> {49+def AscendC_SoftMaxOp : VectorOp<"softmax", "SoftMax", [
50+ AttrSizedOperandSegments, OpWithDstInterface, OpWithSrcInterface
51+]> {
36 let arguments = (ins UnitAttr:$reuseSource, UnitAttr:$basicBlock,52 let arguments = (ins UnitAttr:$reuseSource, UnitAttr:$basicBlock,
37 UnitAttr:$dataFormatNZ, AscendC_LocalTensor:$dst,53 UnitAttr:$dataFormatNZ, AscendC_LocalTensor:$dst,
38- AscendC_LocalTensor:$sumTensor,54+ Optional<AscendC_LocalTensor>:$sumTensor,
39- AscendC_LocalTensor:$maxTensor,55+ Optional<AscendC_LocalTensor>:$maxTensor,
40 AscendC_LocalTensor:$src,56 AscendC_LocalTensor:$src,
41 Optional<AscendC_LocalTensor>:$sharedTmpBuffer,57 Optional<AscendC_LocalTensor>:$sharedTmpBuffer,
42 AscendC_SoftMaxTiling:$tiling,58 AscendC_SoftMaxTiling:$tiling,
43 Optional<AnyType>:$softmaxShapeInfo);59 Optional<AnyType>:$softmaxShapeInfo);
60+ let extraClassDeclaration = extraClassDeclarationBase # [{
61+ SmallVector<Value> getDstTensors()
62+ {
63+ SmallVector<Value> dstTensors{getDst()};
64+ if (auto tensor = getSumTensor())
65+ dstTensors.emplace_back(tensor);
66+ if (auto tensor = getMaxTensor())
67+ dstTensors.emplace_back(tensor);
68+ if (auto tensor = getSharedTmpBuffer())
69+ dstTensors.emplace_back(tensor);
70+ return dstTensors;
71+ }
72+ SmallVector<Value> getSrcTensors() { return {getSrc()}; }
73+ }];
44}74}
45def AscendC_SwiGLUOp : APIOp<"swiglu", "SwiGLU", [AttrSizedOperandSegments, OpWithDstInterface]> {75def AscendC_SwiGLUOp : APIOp<"swiglu", "SwiGLU", [AttrSizedOperandSegments, OpWithDstInterface]> {
46 let arguments = (ins AscendC_LocalTensor:$dst,76 let arguments = (ins AscendC_LocalTensor:$dst,
@@ -50,4 +80,4 @@ def AscendC_SwiGLUOp : APIOp<"swiglu", "SwiGLU", [AttrSizedOperandSegments, OpWi
50 Optional<AscendC_LocalTensor>:$sharedTmpBuffer,80 Optional<AscendC_LocalTensor>:$sharedTmpBuffer,
51 Optional<AnyType>:$calCount);81 Optional<AnyType>:$calCount);
52}82}
53-#endif //ASC_ADV_ACTIVATION_TD83+#endif // ASC_ADV_ACTIVATION_TD
@@ -0,0 +1,30 @@
1+/*
2+ * Copyright (c) 2026 Huawei Technologies Co., Ltd.
3+ * This program is free software, you can redistribute it and/or modify it under the terms and conditions of
4+ * 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,
7+ * 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+#ifndef ASC_ADV_BROADCAST_TD
12+#define ASC_ADV_BROADCAST_TD
13+ 
14+include "Base.td"
15+include "Core/Interfaces.td"
16+include "Core/Types.td"
17+ 
18+include "mlir/IR/OpBase.td"
19+ 
20+def AscendC_BroadcastOp : VectorOp<"broadcast", "Broadcast", [
21+ AttrSizedOperandSegments, OpWithDstInterface, OpWithSrcInterface
22+]> {
23+ let arguments = (ins AscendC_LocalTensor:$dst, AscendC_LocalTensor:$src,
24+ Variadic<I32>:$dstShape, Variadic<I32>:$srcShape);
25+ let extraClassDeclaration = extraClassDeclarationBase # [{
26+ SmallVector<Value> getSrcTensors() { return {getSrc()}; }
27+ }];
28+}
29+ 
30+#endif // ASC_ADV_BROADCAST_TD
@@ -40,6 +40,7 @@ def FloorOp : UnaryMathOp<"floor", "Floor">;
40def FracOp : UnaryMathOp<"frac", "Frac">;40def FracOp : UnaryMathOp<"frac", "Frac">;
41def LgammaOp : UnaryMathOp<"lgamma", "Lgamma">;41def LgammaOp : UnaryMathOp<"lgamma", "Lgamma">;
42def LogOp : UnaryMathOp<"log", "Log">;42def LogOp : UnaryMathOp<"log", "Log">;
43+def Log2Op : UnaryMathOp<"log2", "Log2">;
43def RoundOp : UnaryMathOp<"round", "Round">;44def RoundOp : UnaryMathOp<"round", "Round">;
44def SignOp : UnaryMathOp<"sign", "Sign">;45def SignOp : UnaryMathOp<"sign", "Sign">;
45def SinhOp : UnaryMathOp<"sinh", "Sinh">;46def SinhOp : UnaryMathOp<"sinh", "Sinh">;
@@ -62,27 +63,77 @@ def XorOp : BinaryMathOp<"xor", "Xor">;
62def AxpyOp : MathLibraryOp<"axpy", "Axpy"> {63def AxpyOp : MathLibraryOp<"axpy", "Axpy"> {
63 let arguments = (ins AscendC_LocalTensor:$dst, AscendC_LocalTensor:$src, AnyType:$scalar,64 let arguments = (ins AscendC_LocalTensor:$dst, AscendC_LocalTensor:$src, AnyType:$scalar,
64 Optional<AnyType>:$sharedTmpBuffer, AnyType:$calCount, AnyType:$isReuseSource);65 Optional<AnyType>:$sharedTmpBuffer, AnyType:$calCount, AnyType:$isReuseSource);
66+ let extraClassDeclaration = extraClassDeclarationBase # [{
67+ SmallVector<Value> getDstTensors()
68+ {
69+ SmallVector<Value> dstTensors{getDst()};
70+ if (auto tensor = getSharedTmpBuffer())
71+ dstTensors.emplace_back(tensor);
72+ return dstTensors;
73+ }
74+ SmallVector<Value> getSrcTensors() { return {getSrc()}; }
75+ }];
65}76}
66 77 
67def ClampMaxOp : MathLibraryOp<"clamp_max", "ClampMax"> {78def ClampMaxOp : MathLibraryOp<"clamp_max", "ClampMax"> {
68 let arguments = (ins AscendC_LocalTensor:$dst, AnyType:$src, Optional<AnyType>:$sharedTmpBuffer,79 let arguments = (ins AscendC_LocalTensor:$dst, AnyType:$src, Optional<AnyType>:$sharedTmpBuffer,
69 AnyType:$scalar, AnyType:$calCount, AnyType:$isReuseSource);80 AnyType:$scalar, AnyType:$calCount, AnyType:$isReuseSource);
81+ let extraClassDeclaration = extraClassDeclarationBase # [{
82+ SmallVector<Value> getDstTensors()
83+ {
84+ SmallVector<Value> dstTensors{getDst()};
85+ if (auto tensor = getSharedTmpBuffer())
86+ dstTensors.emplace_back(tensor);
87+ return dstTensors;
88+ }
89+ SmallVector<Value> getSrcTensors() { return {getSrc()}; }
90+ }];
70}91}
71 92 
72def ClampMinOp : MathLibraryOp<"clamp_min", "ClampMin"> {93def ClampMinOp : MathLibraryOp<"clamp_min", "ClampMin"> {
73 let arguments = (ins AscendC_LocalTensor:$dst, AnyType:$src, Optional<AnyType>:$sharedTmpBuffer,94 let arguments = (ins AscendC_LocalTensor:$dst, AnyType:$src, Optional<AnyType>:$sharedTmpBuffer,
74 AnyType:$scalar, AnyType:$calCount, AnyType:$isReuseSource);95 AnyType:$scalar, AnyType:$calCount, AnyType:$isReuseSource);
96+ let extraClassDeclaration = extraClassDeclarationBase # [{
97+ SmallVector<Value> getDstTensors()
98+ {
99+ SmallVector<Value> dstTensors{getDst()};
100+ if (auto tensor = getSharedTmpBuffer())
101+ dstTensors.emplace_back(tensor);
102+ return dstTensors;
103+ }
104+ SmallVector<Value> getSrcTensors() { return {getSrc()}; }
105+ }];
75}106}
76 107 
77def CumSumOp : MathLibraryOp<"cumsum", "CumSum"> {108def CumSumOp : MathLibraryOp<"cumsum", "CumSum"> {
78 let arguments = (ins AnyType:$dst, AnyType:$lastRow, AnyType:$src,109 let arguments = (ins AnyType:$dst, AnyType:$lastRow, AnyType:$src,
79 Optional<AnyType>:$sharedTmpBuffer, AnyType:$lastAxis,110 Optional<AnyType>:$sharedTmpBuffer, AnyType:$lastAxis,
80 AnyType:$reuseSource, AnyType:$outputLastRow);111 AnyType:$reuseSource, AnyType:$outputLastRow);
112+ let extraClassDeclaration = extraClassDeclarationBase # [{
113+ SmallVector<Value> getDstTensors()
114+ {
115+ SmallVector<Value> dstTensors{getDst()};
116+ if (auto tensor = getSharedTmpBuffer())
117+ dstTensors.emplace_back(tensor);
118+ return dstTensors;
119+ }
120+ SmallVector<Value> getSrcTensors() { return {getSrc()}; }
121+ }];
81}122}
82 123 
83def ExpOp : MathLibraryOp<"exp", "Exp"> {124def ExpOp : MathLibraryOp<"exp", "Exp"> {
84 let arguments = (ins AscendC_LocalTensor:$dst, AscendC_LocalTensor:$src, AnyType:$calCount,125 let arguments = (ins AscendC_LocalTensor:$dst, AscendC_LocalTensor:$src, AnyType:$calCount,
85 AnyType:$taylorExpandLevel, Optional<AnyType>:$sharedTmpBuffer, AnyType:$isReuseSource);126 AnyType:$taylorExpandLevel, Optional<AnyType>:$sharedTmpBuffer, AnyType:$isReuseSource);
127+ let extraClassDeclaration = extraClassDeclarationBase # [{
128+ SmallVector<Value> getDstTensors()
129+ {
130+ SmallVector<Value> dstTensors{getDst()};
131+ if (auto tensor = getSharedTmpBuffer())
132+ dstTensors.emplace_back(tensor);
133+ return dstTensors;
134+ }
135+ SmallVector<Value> getSrcTensors() { return {getSrc()}; }
136+ }];
86}137}
87 138 
88-#endif //ASC_ADV_MATH_TD139+#endif // ASC_ADV_MATH_TD
@@ -86,7 +86,7 @@ def AscendC_MatmulSetSparseIndexOp
86 86 
87def AscendC_MatmulSetUserDefInfoOp87def AscendC_MatmulSetUserDefInfoOp
88 : APIOp<"matmul.set_user_def_info", "SetUserDefInfo", [AscMemberFunc]> {88 : APIOp<"matmul.set_user_def_info", "SetUserDefInfo", [AscMemberFunc]> {
89- let arguments = (ins AscendC_Matmul:$matmul, AnyMemRef:$tilingPtr);89+ let arguments = (ins AscendC_Matmul:$matmul, AnyRankedOrUnrankedMemRef:$tilingPtr);
90 let paramTypeLists = [0, -2];90 let paramTypeLists = [0, -2];
91}91}
92 92 
@@ -186,7 +186,7 @@ def AscendC_MatmulIterateBatchCubeOnlyOp
186 let paramTypeLists = [0, 0, 0, 0, 0, 0, 0, 0];186 let paramTypeLists = [0, 0, 0, 0, 0, 0, 0, 0];
187}187}
188 188 
189-def AscendC_MatmulWaitIterateBatchOp 189+def AscendC_MatmulWaitIterateBatchOp
190 : APIOp<"matmul.wait_iterate_batch", "WaitIterateBatch", [AscMemberFunc]> {190 : APIOp<"matmul.wait_iterate_batch", "WaitIterateBatch", [AscMemberFunc]> {
191 let arguments = (ins AscendC_Matmul:$matmul);191 let arguments = (ins AscendC_Matmul:$matmul);
192}192}
@@ -300,7 +300,7 @@ def AscendC_MatmulSetQuantVectorOp : APIOp<"matmul.set_quant_vector", "SetQuantV
300 let summary = "Call `AscendC::Matmul::SetQuantVector` method";300 let summary = "Call `AscendC::Matmul::SetQuantVector` method";
301 let arguments = (ins AscendC_Matmul:$matmul,301 let arguments = (ins AscendC_Matmul:$matmul,
302 AscendC_GlobalTensor:$quantvector);302 AscendC_GlobalTensor:$quantvector);
303- 303+ 
304}304}
305 305 
306def AscendC_MatmulSetOrgShapeOp : APIOp<"matmul.set_org_shape", "SetOrgShape", [AscMemberFunc, AttrSizedOperandSegments]> {306def AscendC_MatmulSetOrgShapeOp : APIOp<"matmul.set_org_shape", "SetOrgShape", [AscMemberFunc, AttrSizedOperandSegments]> {
@@ -20,10 +20,41 @@ include "mlir/Interfaces/CastInterfaces.td"
20include "mlir/Interfaces/SideEffectInterfaces.td"20include "mlir/Interfaces/SideEffectInterfaces.td"
21include "mlir/IR/OpBase.td"21include "mlir/IR/OpBase.td"
22 22 
23-def AscendC_RmsNormOp : APIOp<"rmsnorm", "RmsNorm"> {23+def AscendC_RmsNormOp : VectorOp<"rms_norm", "RmsNorm", [
24+ OpWithDstInterface, OpWithSrcInterface
25+]> {
24 let arguments = (ins UnitAttr:$basicBlock, AscendC_LocalTensor:$dst,26 let arguments = (ins UnitAttr:$basicBlock, AscendC_LocalTensor:$dst,
25 AscendC_LocalTensor:$src, AscendC_LocalTensor:$gamma,27 AscendC_LocalTensor:$src, AscendC_LocalTensor:$gamma,
26 AnyType:$epsilon, AscendC_RmsNormTiling:$tiling,28 AnyType:$epsilon, AscendC_RmsNormTiling:$tiling,
27 Optional<AscendC_LocalTensor>:$sharedTmpBuffer);29 Optional<AscendC_LocalTensor>:$sharedTmpBuffer);
30+ let extraClassDeclaration = extraClassDeclarationBase # [{
31+ SmallVector<Value> getDstTensors()
32+ {
33+ SmallVector<Value> dstTensors{getDst()};
34+ if (auto tensor = getSharedTmpBuffer())
35+ dstTensors.emplace_back(tensor);
36+ return dstTensors;
37+ }
38+ SmallVector<Value> getSrcTensors() { return {getSrc(), getGamma()}; }
39+ }];
40+}
41+ 
42+def AscendC_LayerNormOp : VectorOp<"layer_norm", "LayerNorm", [
43+ OpWithDstInterface, OpWithSrcInterface
44+]> {
45+ let arguments = (ins AscendC_LocalTensor:$dst,
46+ AscendC_LocalTensor:$dstMean, AscendC_LocalTensor:$dstVarRstd,
47+ AscendC_LocalTensor:$src, AscendC_LocalTensor:$gamma,
48+ AscendC_LocalTensor:$beta, AnyType:$epsilon,
49+ AscendC_LayerNormSeparateTiling:$separateTiling,
50+ AscendC_LayerNormPara:$para,
51+ AscendC_LocalTensor:$sharedTmpBuffer,
52+ UnitAttr:$outputRstd);
53+ let extraClassDeclaration = extraClassDeclarationBase # [{
54+ SmallVector<Value> getSrcTensors() { return {getSrc(), getGamma(), getBeta()}; }
55+ SmallVector<Value> getDstTensors() {
56+ return {getDst(), getDstMean(), getDstVarRstd(), getSharedTmpBuffer()};
57+ }
58+ }];
28}59}
29#endif //ASC_ADV_NORMALIZATION_TD60#endif //ASC_ADV_NORMALIZATION_TD
@@ -0,0 +1,41 @@
1+/*
2+ * Copyright (c) 2026 Huawei Technologies Co., Ltd.
3+ * This program is free software, you can redistribute it and/or modify it under the terms and conditions of
4+ * 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,
7+ * 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+#ifndef ASC_ADV_REDUCTION_TD
12+#define ASC_ADV_REDUCTION_TD
13+ 
14+include "Base.td"
15+include "Core/Attributes.td"
16+include "Core/Interfaces.td"
17+include "Core/Types.td"
18+ 
19+include "mlir/Interfaces/CastInterfaces.td"
20+include "mlir/Interfaces/SideEffectInterfaces.td"
21+include "mlir/IR/OpBase.td"
22+ 
23+class AscendC_ReduceOp<string mnemonic, string apiName> : VectorOp<mnemonic, apiName, [
24+ OpWithDstInterface, OpWithSrcInterface, ReduceOpInterface
25+]> {
26+ let arguments = (ins AscendC_LocalTensor:$dst, AscendC_LocalTensor:$src,
27+ AscendC_LocalTensor:$sharedTmpBuffer,
28+ Variadic<I32>:$srcShape, AscendC_ReducePatternAttr:$pattern,
29+ UnitAttr:$isReuseSource);
30+ let extraClassDeclaration = extraClassDeclarationBase # [{
31+ SmallVector<Value> getDstTensors() { return {getDst(), getSharedTmpBuffer()}; }
32+ SmallVector<Value> getSrcTensors() { return {getSrc()}; }
33+ }];
34+}
35+ 
36+def AscendC_ReduceProdOp: AscendC_ReduceOp<"reduce_prod", "ReduceProd">;
37+def AscendC_ReduceMinOp: AscendC_ReduceOp<"reduce_min", "ReduceMin">;
38+def AscendC_ReduceMaxOp: AscendC_ReduceOp<"reduce_max", "ReduceMax">;
39+def AscendC_ReduceSumOp: AscendC_ReduceOp<"reduce_sum", "ReduceSum">;
40+ 
41+#endif //ASC_ADV_REDUCTION_TD
@@ -75,6 +75,9 @@ class APIOp<string mnemonic, string apiName, list<Trait> traits = []>
75class DataCopyOp<string mnemonic, string apiName, list<Trait> traits = []>75class DataCopyOp<string mnemonic, string apiName, list<Trait> traits = []>
76 : APIOp<mnemonic, apiName, [DataCopyOpInterface] # traits>;76 : APIOp<mnemonic, apiName, [DataCopyOpInterface] # traits>;
77 77 
78+class CopyToL0Op<string mnemonic, string apiName, list<Trait> traits = []>
79+ : DataCopyOp<mnemonic, apiName, [CopyToL0OpInterface] # traits>;
80+ 
78class VectorOp<string mnemonic, string apiName, list<Trait> traits = []>81class VectorOp<string mnemonic, string apiName, list<Trait> traits = []>
79 : APIOp<mnemonic, apiName, [VectorOpInterface] # traits>;82 : APIOp<mnemonic, apiName, [VectorOpInterface] # traits>;
80 83 
@@ -85,6 +88,9 @@ class VectorOp<string mnemonic, string apiName, list<Trait> traits = []>
85class UnaryOp<string mnemonic, string apiName, list<Trait> traits = []>88class UnaryOp<string mnemonic, string apiName, list<Trait> traits = []>
86 : VectorOp<mnemonic, apiName, [UnaryOpInterface] # traits> {89 : VectorOp<mnemonic, apiName, [UnaryOpInterface] # traits> {
87 let description = "`AscendC::" # apiName # "` is a vector unary operation.\n";90 let description = "`AscendC::" # apiName # "` is a vector unary operation.\n";
91+ let extraClassDeclaration = extraClassDeclarationBase # [{
92+ SmallVector<Value> getSrcTensors() { return {getSrc()}; }
93+ }];
88}94}
89 95 
90class UnaryL0Op<string mnemonic, string apiName, list<Trait> traits = []>96class UnaryL0Op<string mnemonic, string apiName, list<Trait> traits = []>
@@ -122,13 +128,16 @@ multiclass UnaryL012Op<string mnemonic, string apiName, list<Trait> traits = []>
122class BinaryOp<string mnemonic, string apiName, list<Trait> traits = []>128class BinaryOp<string mnemonic, string apiName, list<Trait> traits = []>
123 : VectorOp<mnemonic, apiName, [BinaryOpInterface] # traits> {129 : VectorOp<mnemonic, apiName, [BinaryOpInterface] # traits> {
124 let description = "`AscendC::" # apiName # "` is a vector binary operation.\n";130 let description = "`AscendC::" # apiName # "` is a vector binary operation.\n";
131+ let extraClassDeclaration = extraClassDeclarationBase # [{
132+ SmallVector<Value> getSrcTensors() { return {getSrc0(), getSrc1()}; }
133+ }];
125}134}
126 135 
127class BinaryL0Op<string mnemonic, string apiName, list<Trait> traits = []>136class BinaryL0Op<string mnemonic, string apiName, list<Trait> traits = []>
128- : BinaryOp<mnemonic, apiName, traits> {137+ : BinaryOp<mnemonic, apiName, [BinaryL0OpInterface] # traits> {
129 let description = "`AscendC::" # apiName # "` is a vector binary operation (L0 API).\n";138 let description = "`AscendC::" # apiName # "` is a vector binary operation (L0 API).\n";
130 let arguments = (ins AnyType:$dst, AnyType:$src0, AnyType:$src1,139 let arguments = (ins AnyType:$dst, AnyType:$src0, AnyType:$src1,
131- AnyType:$mask, AnyType:$repeatTimes, 140+ AnyType:$mask, AnyType:$repeatTimes,
132 AscendC_BinaryRepeatParams:$repeatParams);141 AscendC_BinaryRepeatParams:$repeatParams);
133}142}
134 143 
@@ -136,22 +145,22 @@ class BinaryL1Op<string mnemonic, string apiName, list<Trait> traits = []>
136 : BinaryOp<mnemonic, apiName, traits> {145 : BinaryOp<mnemonic, apiName, traits> {
137 let description = "`AscendC::" # apiName # "` is a vector binary operation (L1 API).\n";146 let description = "`AscendC::" # apiName # "` is a vector binary operation (L1 API).\n";
138 let arguments = (ins AnyType:$dst, AnyType:$src0, AnyType:$src1,147 let arguments = (ins AnyType:$dst, AnyType:$src0, AnyType:$src1,
139- Variadic<UI64>:$mask, AnyType:$repeatTimes, 148+ Variadic<UI64>:$mask, AnyType:$repeatTimes,
140 AscendC_BinaryRepeatParams:$repeatParams);149 AscendC_BinaryRepeatParams:$repeatParams);
141}150}
142 151 
143class BinaryL2Op<string mnemonic, string apiName, list<Trait> traits = []>152class BinaryL2Op<string mnemonic, string apiName, list<Trait> traits = []>
144- : BinaryOp<mnemonic, apiName, traits> {153+ : BinaryOp<mnemonic, apiName, [BinaryL2OpInterface] # traits> {
145 let description = "`AscendC::" # apiName # "` is a vector binary operation (L2 API).\n";154 let description = "`AscendC::" # apiName # "` is a vector binary operation (L2 API).\n";
146 let arguments = (ins AnyType:$dst, AnyType:$src0, AnyType:$src1,155 let arguments = (ins AnyType:$dst, AnyType:$src0, AnyType:$src1,
147 AnyType:$calCount);156 AnyType:$calCount);
148}157}
149 158 
150class BinaryTemplateL0Op<string mnemonic, string apiName, list<Trait> traits = []>159class BinaryTemplateL0Op<string mnemonic, string apiName, list<Trait> traits = []>
151- : BinaryOp<mnemonic, apiName, traits> {160+ : BinaryOp<mnemonic, apiName, [BinaryL0OpInterface] # traits> {
152 let description = "`AscendC::" # apiName # "` is a vector binary operation (L0 API).\n";161 let description = "`AscendC::" # apiName # "` is a vector binary operation (L0 API).\n";
153 let arguments = (ins AnyType:$dst, AnyType:$src0, AnyType:$src1,162 let arguments = (ins AnyType:$dst, AnyType:$src0, AnyType:$src1,
154- AnyType:$mask, AnyType:$repeatTimes, 163+ AnyType:$mask, AnyType:$repeatTimes,
155 AscendC_BinaryRepeatParams:$repeatParams, UnitAttr:$isSetMask);164 AscendC_BinaryRepeatParams:$repeatParams, UnitAttr:$isSetMask);
156}165}
157 166 
@@ -159,22 +168,22 @@ class BinaryTemplateL1Op<string mnemonic, string apiName, list<Trait> traits = [
159 : BinaryOp<mnemonic, apiName, traits> {168 : BinaryOp<mnemonic, apiName, traits> {
160 let description = "`AscendC::" # apiName # "` is a vector binary operation (L1 API).\n";169 let description = "`AscendC::" # apiName # "` is a vector binary operation (L1 API).\n";
161 let arguments = (ins AnyType:$dst, AnyType:$src0, AnyType:$src1,170 let arguments = (ins AnyType:$dst, AnyType:$src0, AnyType:$src1,
162- Variadic<UI64>:$mask, AnyType:$repeatTimes, 171+ Variadic<UI64>:$mask, AnyType:$repeatTimes,
163 AscendC_BinaryRepeatParams:$repeatParams, UnitAttr:$isSetMask);172 AscendC_BinaryRepeatParams:$repeatParams, UnitAttr:$isSetMask);
164}173}
165 174 
166class BinaryTemplateL2Op<string mnemonic, string apiName, list<Trait> traits = []>175class BinaryTemplateL2Op<string mnemonic, string apiName, list<Trait> traits = []>
167- : BinaryOp<mnemonic, apiName, traits> {176+ : BinaryOp<mnemonic, apiName, [BinaryL2OpInterface] # traits> {
168 let description = "`AscendC::" # apiName # "` is a vector binary operation (L2 API).\n";177 let description = "`AscendC::" # apiName # "` is a vector binary operation (L2 API).\n";
169 let arguments = (ins AnyType:$dst, AnyType:$src0, AnyType:$src1,178 let arguments = (ins AnyType:$dst, AnyType:$src0, AnyType:$src1,
170 AnyType:$calCount, UnitAttr:$isSetMask);179 AnyType:$calCount, UnitAttr:$isSetMask);
171}180}
172 181 
173class BinaryCastL0Op<string mnemonic, string apiName, list<Trait> traits = []>182class BinaryCastL0Op<string mnemonic, string apiName, list<Trait> traits = []>
174- : BinaryOp<mnemonic, apiName, traits> {183+ : BinaryOp<mnemonic, apiName, [BinaryL0OpInterface] # traits> {
175 let description = "`AscendC::" # apiName # "` is a vector binary operation (L0 API).\n";184 let description = "`AscendC::" # apiName # "` is a vector binary operation (L0 API).\n";
176 let arguments = (ins AnyType:$dst, AnyType:$src0, AnyType:$src1,185 let arguments = (ins AnyType:$dst, AnyType:$src0, AnyType:$src1,
177- AnyType:$mask, AnyType:$repeatTimes, 186+ AnyType:$mask, AnyType:$repeatTimes,
178 AscendC_BinaryRepeatParams:$repeatParams, UnitAttr:$isSetMask);187 AscendC_BinaryRepeatParams:$repeatParams, UnitAttr:$isSetMask);
179}188}
180 189 
@@ -182,12 +191,12 @@ class BinaryCastL1Op<string mnemonic, string apiName, list<Trait> traits = []>
182 : BinaryOp<mnemonic, apiName, traits> {191 : BinaryOp<mnemonic, apiName, traits> {
183 let description = "`AscendC::" # apiName # "` is a vector binary operation (L1 API).\n";192 let description = "`AscendC::" # apiName # "` is a vector binary operation (L1 API).\n";
184 let arguments = (ins AnyType:$dst, AnyType:$src0, AnyType:$src1,193 let arguments = (ins AnyType:$dst, AnyType:$src0, AnyType:$src1,
185- Variadic<UI64>:$mask, AnyType:$repeatTimes, 194+ Variadic<UI64>:$mask, AnyType:$repeatTimes,
186 AscendC_BinaryRepeatParams:$repeatParams, UnitAttr:$isSetMask);195 AscendC_BinaryRepeatParams:$repeatParams, UnitAttr:$isSetMask);
187}196}
188 197 
189class BinaryCastL2Op<string mnemonic, string apiName, list<Trait> traits = []>198class BinaryCastL2Op<string mnemonic, string apiName, list<Trait> traits = []>
190- : BinaryOp<mnemonic, apiName, traits> {199+ : BinaryOp<mnemonic, apiName, [BinaryL2OpInterface] # traits> {
191 let description = "`AscendC::" # apiName # "` is a vector binary operation (L2 API).\n";200 let description = "`AscendC::" # apiName # "` is a vector binary operation (L2 API).\n";
192 let arguments = (ins AnyType:$dst, AnyType:$src0, AnyType:$src1,201 let arguments = (ins AnyType:$dst, AnyType:$src0, AnyType:$src1,
193 AnyType:$calCount, UnitAttr:$isSetMask);202 AnyType:$calCount, UnitAttr:$isSetMask);
@@ -237,6 +246,9 @@ multiclass BinaryTemplateL0123Op<string baseMnemonic, string apiName, string l3o
237class VecScalarOp<string mnemonic, string apiName, list<Trait> traits = []>246class VecScalarOp<string mnemonic, string apiName, list<Trait> traits = []>
238 : VectorOp<mnemonic, apiName, [VecScalarOpInterface] # traits> {247 : VectorOp<mnemonic, apiName, [VecScalarOpInterface] # traits> {
239 let description = "`AscendC::" # apiName # "` is a vector-scalar operation.\n";248 let description = "`AscendC::" # apiName # "` is a vector-scalar operation.\n";
249+ let extraClassDeclaration = extraClassDeclarationBase # [{
250+ SmallVector<Value> getSrcTensors() { return {getSrc()}; }
251+ }];
240}252}
241 253 
242class VecScalarL0Op<string mnemonic, string apiName, list<Trait> traits = []>254class VecScalarL0Op<string mnemonic, string apiName, list<Trait> traits = []>
@@ -268,6 +280,36 @@ multiclass VecScalarL012Op<string baseMnemonic, string apiName, list<Trait> trai
268 def L2Op : VecScalarL2Op<baseMnemonic # "_l2", apiName, traits>;280 def L2Op : VecScalarL2Op<baseMnemonic # "_l2", apiName, traits>;
269}281}
270 282 
283+//===----------------------------------------------------------------------===//
284+// Register API operations
285+//===----------------------------------------------------------------------===//
286+ 
287+class RegOp<string mnemonic, string apiName, list<Trait> traits = []>
288+ : APIOp<mnemonic, "Reg::" # apiName, [RegOpInterface] # traits> {
289+ let description = "`AscendC::Reg::" # apiName # "` is a register API operation.\n";
290+}
291+ 
292+class BinaryRegOp<string baseMnemonic, string apiName, list<Trait> traits = []>
293+ : RegOp<baseMnemonic # "_reg", apiName, [BinaryRegOpInterface] # traits> {
294+ let description = "`AscendC::Reg::" # apiName # "` is a binary register API operation.\n";
295+ let arguments = (ins AscendC_RegTensor:$dstReg, AscendC_RegTensor:$src0Reg,
296+ AscendC_RegTensor:$src1Reg, AscendC_MaskReg:$maskReg);
297+}
298+ 
299+class UnaryRegOp<string baseMnemonic, string apiName, list<Trait> traits = []>
300+ : RegOp<baseMnemonic # "_reg", apiName, [UnaryRegOpInterface] # traits> {
301+ let description = "`AscendC::Reg::" # apiName # "` is a unary register API operation.\n";
302+ let arguments = (ins AscendC_RegTensor:$dstReg, AscendC_RegTensor:$srcReg,
303+ AscendC_MaskReg:$maskReg);
304+}
305+ 
306+class VecScalarRegOp<string baseMnemonic, string apiName, list<Trait> traits = []>
307+ : RegOp<baseMnemonic # "_reg", apiName, [VecScalarRegOpInterface] # traits> {
308+ let description = "`AscendC::Reg::" # apiName # "` is a vector scalar register API operation.\n";
309+ let arguments = (ins AscendC_RegTensor:$dstReg, AscendC_RegTensor:$srcReg, AnyType: $scalar,
310+ AscendC_MaskReg:$maskReg);
311+}
312+ 
271//===----------------------------------------------------------------------===//313//===----------------------------------------------------------------------===//
272// Math library operations314// Math library operations
273//===----------------------------------------------------------------------===//315//===----------------------------------------------------------------------===//
@@ -282,6 +324,16 @@ class UnaryMathOp<string mnemonic, string apiName, list<Trait> traits = []>
282 let description = "`AscendC::" # apiName # "` is a unary math library operation.\n";324 let description = "`AscendC::" # apiName # "` is a unary math library operation.\n";
283 let arguments = (ins AscendC_LocalTensor:$dst, AnyType:$src, Optional<AnyType>:$sharedTmpBuffer,325 let arguments = (ins AscendC_LocalTensor:$dst, AnyType:$src, Optional<AnyType>:$sharedTmpBuffer,
284 Optional<AnyType>:$calCount, AnyType:$isReuseSource);326 Optional<AnyType>:$calCount, AnyType:$isReuseSource);
327+ let extraClassDeclaration = extraClassDeclarationBase # [{
328+ SmallVector<Value> getDstTensors()
329+ {
330+ SmallVector<Value> dstTensors{getDst()};
331+ if (auto tensor = getSharedTmpBuffer())
332+ dstTensors.emplace_back(tensor);
333+ return dstTensors;
334+ }
335+ SmallVector<Value> getSrcTensors() { return {getSrc()}; }
336+ }];
285}337}
286 338 
287class BinaryMathOp<string mnemonic, string apiName, list<Trait> traits = []>339class BinaryMathOp<string mnemonic, string apiName, list<Trait> traits = []>
@@ -289,6 +341,16 @@ class BinaryMathOp<string mnemonic, string apiName, list<Trait> traits = []>
289 let description = "`AscendC::" # apiName # "` is a binary math library operation.\n";341 let description = "`AscendC::" # apiName # "` is a binary math library operation.\n";
290 let arguments = (ins AscendC_LocalTensor:$dst, AnyType:$src0, AnyType:$src1,342 let arguments = (ins AscendC_LocalTensor:$dst, AnyType:$src0, AnyType:$src1,
291 Optional<AnyType>:$sharedTmpBuffer, Optional<AnyType>:$calCount, AnyType:$isReuseSource);343 Optional<AnyType>:$sharedTmpBuffer, Optional<AnyType>:$calCount, AnyType:$isReuseSource);
344+ let extraClassDeclaration = extraClassDeclarationBase # [{
345+ SmallVector<Value> getDstTensors()
346+ {
347+ SmallVector<Value> dstTensors{getDst()};
348+ if (auto tensor = getSharedTmpBuffer())
349+ dstTensors.emplace_back(tensor);
350+ return dstTensors;
351+ }
352+ SmallVector<Value> getSrcTensors() { return {getSrc0(), getSrc1()}; }
353+ }];
292}354}
293 355 
294//===----------------------------------------------------------------------===//356//===----------------------------------------------------------------------===//
@@ -102,4 +102,15 @@ def CrossCoreWaitFlagOp : APIOp<"cross_core_wait_flag", "CrossCoreWaitFlag"> {
102 let arguments = (ins AnyType:$flagId, UI8Attr:$modeId, AscendC_PipeAttr:$pipe);102 let arguments = (ins AnyType:$flagId, UI8Attr:$modeId, AscendC_PipeAttr:$pipe);
103 let assemblyFormat = "$flagId `,` $modeId `,` $pipe attr-dict `:` type($flagId)";103 let assemblyFormat = "$flagId `,` $modeId `,` $pipe attr-dict `:` type($flagId)";
104}104}
105+ 
106+def GetBufOp : APIOp<"get_buf", "get_buf"> {
107+ let arguments = (ins AscendC_PipeAttr:$pipe, I32Attr:$bufId, UnitAttr:$mode);
108+ let assemblyFormat = "$pipe `,` $bufId attr-dict";
109+}
110+ 
111+def RlsBufOp : APIOp<"rls_buf", "rls_buf"> {
112+ let arguments = (ins AscendC_PipeAttr:$pipe, I32Attr:$bufId, UnitAttr:$mode);
113+ let assemblyFormat = "$pipe `,` $bufId attr-dict";
114+}
115+ 
105#endif // ASC_BASIC_OP_BLOCKSYNC_TD116#endif // ASC_BASIC_OP_BLOCKSYNC_TD
@@ -26,4 +26,17 @@ def AscendC_TransDataTo5HDTensorListOp : TransDataTo5HDTensorListOp<"trans_data_
26def AscendC_TransDataTo5HDUintListOp : TransDataTo5HDUintListOp<"trans_data_to_5hd_uint_list", "TransDataTo5HD">;26def AscendC_TransDataTo5HDUintListOp : TransDataTo5HDUintListOp<"trans_data_to_5hd_uint_list", "TransDataTo5HD">;
27def AscendC_TransDataTo5HDOp : TransDataTo5HDSingleOp<"trans_data_to_5hd", "TransDataTo5HD">;27def AscendC_TransDataTo5HDOp : TransDataTo5HDSingleOp<"trans_data_to_5hd", "TransDataTo5HD">;
28 28 
29+def AscendC_TransDataTo5HDTensorOp: VectorOp<"trans_data_to_5hd_tensor","TransDataTo5HD",
30+ [OpWithDstInterface, OpWithSrcInterface]> {
31+ let summary = "Call `AscendC::" # apiName # "` with offsets in src and dst tensor";
32+ let arguments = (ins AscendC_LocalTensor:$dst,
33+ AscendC_LocalTensor:$src,
34+ DenseI32ArrayAttr:$dstOffsets,
35+ DenseI32ArrayAttr:$srcOffsets,
36+ AscendC_TransDataTo5HDParams:$params);
37+ let extraClassDeclaration = extraClassDeclarationBase # [{
38+ SmallVector<Value> getSrcTensors() { return {getSrc()}; }
39+ }];
40+}
41+ 
29#endif // ASC_BASIC_OP_DATA_CONVERSION_TD42#endif // ASC_BASIC_OP_DATA_CONVERSION_TD
@@ -73,18 +73,24 @@ def AscendC_DataCopySliceOp : DataCopyOp<"data_copy_slice", "DataCopy", [AttrSiz
73 Optional<UI32>:$dimValue);73 Optional<UI32>:$dimValue);
74}74}
75 75 
76-def AscendC_DataCopyL0Op : DataCopyOp<"data_copy_l0", "DataCopy", [AscFunc]> {76+def AscendC_DataCopyL0Op : CopyToL0Op<"data_copy_l0", "DataCopy", [AscFunc]> {
77 let description = "Perform copying between tensors (L0 API)";77 let description = "Perform copying between tensors (L0 API)";
78 let arguments = (ins AscendC_BaseTensorTypeInterface:$dst,78 let arguments = (ins AscendC_BaseTensorTypeInterface:$dst,
79 AscendC_BaseTensorTypeInterface:$src,79 AscendC_BaseTensorTypeInterface:$src,
80 AscendC_DataCopyParams:$repeatParams);80 AscendC_DataCopyParams:$repeatParams);
81+ let extraClassDeclaration = extraClassDeclarationBase # [{
82+ SmallVector<Value> getSrcTensors() { return {getSrc()}; }
83+ }];
81}84}
82 85 
83-def AscendC_DataCopyL2Op : DataCopyOp<"data_copy_l2", "DataCopy", [AscFunc]> {86+def AscendC_DataCopyL2Op : DataCopyOp<"data_copy_l2", "DataCopy", [AscFunc, OpWithSrcInterface]> {
84 let description = "Perform copying between tensors (L2 API)";87 let description = "Perform copying between tensors (L2 API)";
85 let arguments = (ins AscendC_BaseTensorTypeInterface:$dst,88 let arguments = (ins AscendC_BaseTensorTypeInterface:$dst,
86 AscendC_BaseTensorTypeInterface:$src,89 AscendC_BaseTensorTypeInterface:$src,
87 AnyType:$calCount);90 AnyType:$calCount);
91+ let extraClassDeclaration = extraClassDeclarationBase # [{
92+ SmallVector<Value> getSrcTensors() { return {getSrc()}; }
93+ }];
88}94}
89 95 
90def AscendC_DataCopyPadExtParamsOp : AscendC_Op<"data_copy_pad_ext_params", [AscConstructor]> {96def AscendC_DataCopyPadExtParamsOp : AscendC_Op<"data_copy_pad_ext_params", [AscConstructor]> {
@@ -139,12 +145,34 @@ def AscendC_DataCopyPadNd2NzOp : DataCopyOp<"data_copy_pad_nd2nz", "DataCopyPad"
139 AscendC_Nd2NzParams:$nd2nzParams);145 AscendC_Nd2NzParams:$nd2nzParams);
140}146}
141 147 
148+def AscendC_NdDmaParamsOp: AscendC_Op<"nd_dma_params", [AttrSizedOperandSegments]> {
149+ let summary = "Create NdDmaParams";
150+ let results = (outs AscendC_NdDmaParams:$params);
151+ let arguments = (ins I32Attr:$dim, AnyType:$padValue, Variadic<I32>:$size,
152+ Variadic<I32>:$srcStride, I32ArrayAttr:$dstStride,
153+ I32ArrayAttr:$padLeft, I32ArrayAttr:$padRight);
154+ let assemblyFormat = [{
155+ $dim `,` $padValue `:` type($padValue) `,` $size `,` $srcStride `,`
156+ $dstStride `,` $padLeft `,` $padRight `,` attr-dict `:` type($params)
157+ }];
158+}
159+ 
160+def AscendC_DataCopyNdDmaOp : DataCopyOp<"data_copy_nd_dma", "DataCopy", [
161+ AscFunc, OpWithSrcInterface
162+]> {
163+ let description = "DataCopy with NdDmaParams arguments";
164+ let arguments = (ins AscendC_BaseTensorTypeInterface:$dst,
165+ AscendC_BaseTensorTypeInterface:$src,
166+ AscendC_NdDmaParams:$params, I32Attr:$dims);
167+ let extraClassDeclaration = extraClassDeclarationBase # [{
168+ SmallVector<Value> getSrcTensors() { return {getSrc()}; }
169+ }];
170+}
171+ 
142def AscendC_LoadImageToLocalOp : APIOp<"load_image_to_local", "LoadImageToLocal", [AscFunc]> {172def AscendC_LoadImageToLocalOp : APIOp<"load_image_to_local", "LoadImageToLocal", [AscFunc]> {
143 let description = "Load image data from Global Memory to LocalTensor (supports TPosition A1/B1). ";173 let description = "Load image data from Global Memory to LocalTensor (supports TPosition A1/B1). ";
144- let arguments = (ins 174+ let arguments = (ins AscendC_LocalTensor:$dst,
145- AscendC_LocalTensor:$dst, 175+ AscendC_LoadImageToLocalParams:$loadDataParams);
146- AscendC_LoadImageToLocalParams:$loadDataParams
147- );
148}176}
149 177 
150def AscendC_SetPadValueOp : APIOp<"set_pad_value", "SetPadValue", [AscFunc]> {178def AscendC_SetPadValueOp : APIOp<"set_pad_value", "SetPadValue", [AscFunc]> {
@@ -153,4 +181,17 @@ def AscendC_SetPadValueOp : APIOp<"set_pad_value", "SetPadValue", [AscFunc]> {
153 AscendC_TPositionAttr:$pos);181 AscendC_TPositionAttr:$pos);
154 let paramTypeLists = [1, 3];182 let paramTypeLists = [1, 3];
155}183}
184+ 
185+def AscendC_SetLoopModeParaOp: APIOp<"set_loop_mode_para", "SetLoopModePara", [AscFunc]> {
186+ let description = "Set loop mode params for DataCopyPad";
187+ let arguments = (ins AscendC_LoopModeParams:$params,
188+ AscendC_DataCopyMVTypeAttr:$mvType);
189+ let paramTypeLists = [0, -1];
190+}
191+ 
192+def AscendC_ResetLoopModeParaOp: APIOp<"reset_loop_mode_para", "ResetLoopModePara", [AscFunc]> {
193+ let description = "Reset loop mode params for DataCopyPad";
194+ let arguments = (ins AscendC_DataCopyMVTypeAttr:$mvType);
195+ let paramTypeLists = [-1];
196+}
156#endif // ASC_BASIC_OP_DATACOPY_TD197#endif // ASC_BASIC_OP_DATACOPY_TD
@@ -20,11 +20,14 @@ include "mlir/Interfaces/CastInterfaces.td"
20include "mlir/Interfaces/SideEffectInterfaces.td"20include "mlir/Interfaces/SideEffectInterfaces.td"
21include "mlir/IR/OpBase.td"21include "mlir/IR/OpBase.td"
22 22 
23-def AscendC_FixpipeOp : APIOp<"fixpipe", "Fixpipe"> {23+def AscendC_FixpipeOp : DataCopyOp<"fixpipe", "Fixpipe", [OpWithDstInterface, OpWithSrcInterface]> {
24 let summary = "Call `AscendC::Fixpipe` function without cbufWorkspace";24 let summary = "Call `AscendC::Fixpipe` function without cbufWorkspace";
25- let arguments = (ins AscendC_GlobalTensor:$dst, AscendC_LocalTensor:$src, 25+ let arguments = (ins AscendC_BaseTensorTypeInterface:$dst, AscendC_LocalTensor:$src,
26- AscendC_FixpipeParamsV220:$intriParams,26+ AnyTypeOf<[AscendC_FixpipeParamsV220, AscendC_FixpipeParamsC310]>:$intriParams,
27 AscendC_FixpipeConfig:$fixpipeConfig);27 AscendC_FixpipeConfig:$fixpipeConfig);
28+ let extraClassDeclaration = extraClassDeclarationBase # [{
29+ SmallVector<Value> getSrcTensors() { return {getSrc()}; }
30+ }];
28}31}
29 32 
30def AscendC_FixpipeWithWorkspaceOp : APIOp<"fixpipe_with_workspace", "Fixpipe"> {33def AscendC_FixpipeWithWorkspaceOp : APIOp<"fixpipe_with_workspace", "Fixpipe"> {
@@ -19,6 +19,12 @@ include "mlir/Interfaces/CastInterfaces.td"
19include "mlir/Interfaces/SideEffectInterfaces.td"19include "mlir/Interfaces/SideEffectInterfaces.td"
20include "mlir/IR/OpBase.td"20include "mlir/IR/OpBase.td"
21 21 
22+def AscendC_FillOp : APIOp<"fill", "Fill", [AscFunc, OpWithDstInterface]> {
23+ let description = "Fill the LocalTensor with a specific value.";
24+ let arguments = (ins AscendC_LocalTensor:$dst,
25+ AscendC_InitConstValueParams:$initConstValueParams);
26+}
27+ 
22def AscendC_InitConstValueOp28def AscendC_InitConstValueOp
23 : APIOp<"init_const_value", "InitConstValue", [AscFunc]> {29 : APIOp<"init_const_value", "InitConstValue", [AscFunc]> {
24 let description = "Initialize the LocalTensor at a specific TPosition to a specific value.";30 let description = "Initialize the LocalTensor at a specific TPosition to a specific value.";
@@ -28,107 +34,139 @@ def AscendC_InitConstValueOp
28 );34 );
29}35}
30 36 
31-def AscendC_LoadDataOp37+def AscendC_LoadDataOp : CopyToL0Op<"load_data", "LoadData", [AscFunc]> {
32- : APIOp<"load_data", "LoadData", [AscFunc]> {
33 let arguments = (ins AscendC_LocalTensor:$dst, AscendC_LocalTensor:$src,38 let arguments = (ins AscendC_LocalTensor:$dst, AscendC_LocalTensor:$src,
34 AscendC_BaseLoadDataParamsTypeInterface:$loadDataParams);39 AscendC_BaseLoadDataParamsTypeInterface:$loadDataParams);
40+ let extraClassDeclaration = extraClassDeclarationBase # [{
41+ SmallVector<Value> getSrcTensors() { return {getSrc()}; }
42+ }];
35}43}
36 44 
37-def AscendC_LoadDataL0Op45+def AscendC_LoadDataL0Op : CopyToL0Op<"load_data_l0", "LoadData", [AscFunc]> {
38- : DataCopyOp<"load_data_l0", "LoadData", [AscFunc]> {
39 let description = "Perform 2D LoadData between LocalTensor and LocalTensor";46 let description = "Perform 2D LoadData between LocalTensor and LocalTensor";
40- let arguments = (ins AscendC_BaseTensorTypeInterface:$dst,47+ let arguments = (ins AscendC_LocalTensor:$dst,
41 AscendC_BaseTensorTypeInterface:$src,48 AscendC_BaseTensorTypeInterface:$src,
42 AscendC_LoadData2DParams:$loadDataParams);49 AscendC_LoadData2DParams:$loadDataParams);
50+ let extraClassDeclaration = extraClassDeclarationBase # [{
51+ SmallVector<Value> getSrcTensors() { return {getSrc()}; }
52+ }];
43}53}
44 54 
45-def AscendC_LoadDataG2LOp55+def AscendC_LoadDataG2LOp : CopyToL0Op<"load_data_g2l", "LoadData", [AscFunc]> {
46- : DataCopyOp<"load_data_g2l", "LoadData", [AscFunc]> {
47 let description = "Perform 2D LoadData from GlobalTensor to LocalTensor";56 let description = "Perform 2D LoadData from GlobalTensor to LocalTensor";
48- let arguments = (ins AscendC_BaseTensorTypeInterface:$dst,57+ let arguments = (ins AscendC_LocalTensor:$dst,
49 AscendC_BaseTensorTypeInterface:$src,58 AscendC_BaseTensorTypeInterface:$src,
50 AscendC_LoadData2DParams:$loadDataParams);59 AscendC_LoadData2DParams:$loadDataParams);
60+ let extraClassDeclaration = extraClassDeclarationBase # [{
61+ SmallVector<Value> getSrcTensors() { return {getSrc()}; }
62+ }];
51}63}
52 64 
53def AscendC_LoadDataL0V2Op65def AscendC_LoadDataL0V2Op
54- : DataCopyOp<"load_data_l0_v2", "LoadData", [AscFunc]> {66+ : CopyToL0Op<"load_data_l0_v2", "LoadData", [AscFunc]> {
55 let description = "Perform 2D LoadData V2 between LocalTensor and LocalTensor";67 let description = "Perform 2D LoadData V2 between LocalTensor and LocalTensor";
56- let arguments = (ins AscendC_BaseTensorTypeInterface:$dst,68+ let arguments = (ins AscendC_LocalTensor:$dst,
57 AscendC_BaseTensorTypeInterface:$src,69 AscendC_BaseTensorTypeInterface:$src,
58 AscendC_LoadData2DParamsV2:$loadDataParams);70 AscendC_LoadData2DParamsV2:$loadDataParams);
71+ let extraClassDeclaration = extraClassDeclarationBase # [{
72+ SmallVector<Value> getSrcTensors() { return {getSrc()}; }
73+ }];
59}74}
60 75 
61def AscendC_LoadDataG2LV2Op76def AscendC_LoadDataG2LV2Op
62- : DataCopyOp<"load_data_g2l_v2", "LoadData", [AscFunc]> {77+ : CopyToL0Op<"load_data_g2l_v2", "LoadData", [AscFunc]> {
63 let description = "Perform 2D LoadData V2 from GlobalTensor to LocalTensor";78 let description = "Perform 2D LoadData V2 from GlobalTensor to LocalTensor";
64- let arguments = (ins AscendC_BaseTensorTypeInterface:$dst,79+ let arguments = (ins AscendC_LocalTensor:$dst,
65 AscendC_BaseTensorTypeInterface:$src,80 AscendC_BaseTensorTypeInterface:$src,
66 AscendC_LoadData2DParamsV2:$loadDataParams);81 AscendC_LoadData2DParamsV2:$loadDataParams);
82+ let extraClassDeclaration = extraClassDeclarationBase # [{
83+ SmallVector<Value> getSrcTensors() { return {getSrc()}; }
84+ }];
67}85}
68 86 
69def AscendC_LoadData3DL0V1Op87def AscendC_LoadData3DL0V1Op
70- : DataCopyOp<"load_data_3d_l0_v1", "LoadData", [AscFunc]> {88+ : CopyToL0Op<"load_data_3d_l0_v1", "LoadData", [AscFunc]> {
71 let description = "Perform 3D LoadData V1 between LocalTensor and LocalTensor";89 let description = "Perform 3D LoadData V1 between LocalTensor and LocalTensor";
72- let arguments = (ins AscendC_BaseTensorTypeInterface:$dst,90+ let arguments = (ins AscendC_LocalTensor:$dst, AscendC_LocalTensor:$src,
73- AscendC_BaseTensorTypeInterface:$src,
74 AscendC_LoadData3DParamsV1:$loadDataParams);91 AscendC_LoadData3DParamsV1:$loadDataParams);
92+ let extraClassDeclaration = extraClassDeclarationBase # [{
93+ SmallVector<Value> getSrcTensors() { return {getSrc()}; }
94+ }];
75}95}
76 96 
77def AscendC_LoadData3DL0V2Op97def AscendC_LoadData3DL0V2Op
78- : DataCopyOp<"load_data_3d_l0_v2", "LoadData", [AscFunc]> {98+ : CopyToL0Op<"load_data_3d_l0_v2", "LoadData", [AscFunc]> {
79 let description = "Perform 3D LoadData V2 between LocalTensor and LocalTensor";99 let description = "Perform 3D LoadData V2 between LocalTensor and LocalTensor";
80- let arguments = (ins AscendC_BaseTensorTypeInterface:$dst,100+ let arguments = (ins AscendC_LocalTensor:$dst, AscendC_LocalTensor:$src,
81- AscendC_BaseTensorTypeInterface:$src,
82 AscendC_LoadData3DParamsV2:$loadDataParams);101 AscendC_LoadData3DParamsV2:$loadDataParams);
102+ let extraClassDeclaration = extraClassDeclarationBase # [{
103+ SmallVector<Value> getSrcTensors() { return {getSrc()}; }
104+ }];
83}105}
84 106 
85def AscendC_LoadData3DL0V2ProOp107def AscendC_LoadData3DL0V2ProOp
86- : DataCopyOp<"load_data_3d_l0_v2pro", "LoadData", [AscFunc]> {108+ : CopyToL0Op<"load_data_3d_l0_v2pro", "LoadData", [AscFunc]> {
87 let description = "Perform 3D LoadData V2Pro between LocalTensor and LocalTensor";109 let description = "Perform 3D LoadData V2Pro between LocalTensor and LocalTensor";
88- let arguments = (ins AscendC_BaseTensorTypeInterface:$dst,110+ let arguments = (ins AscendC_LocalTensor:$dst, AscendC_LocalTensor:$src,
89- AscendC_BaseTensorTypeInterface:$src,
90 AscendC_LoadData3DParamsV2Pro:$loadDataParams);111 AscendC_LoadData3DParamsV2Pro:$loadDataParams);
112+ let extraClassDeclaration = extraClassDeclarationBase # [{
113+ SmallVector<Value> getSrcTensors() { return {getSrc()}; }
114+ }];
91}115}
92 116 
93def AscendC_LoadDataWithSparseOp117def AscendC_LoadDataWithSparseOp
94- : APIOp<"load_data_with_sparse", "LoadDataWithSparse", [AscFunc]> {118+ : CopyToL0Op<"load_data_with_sparse", "LoadDataWithSparse", [AscFunc]> {
95 let arguments = (ins AscendC_LocalTensor:$dst, AscendC_LocalTensor:$src,119 let arguments = (ins AscendC_LocalTensor:$dst, AscendC_LocalTensor:$src,
96- AscendC_LocalTensor:$idx, AscendC_LoadData2DParams:$loadDataParam);120+ AscendC_LocalTensor:$idx,
121+ AscendC_LoadData2DParams:$loadDataParam);
122+ let extraClassDeclaration = extraClassDeclarationBase # [{
123+ SmallVector<Value> getSrcTensors() { return {getSrc(), getIdx()}; }
124+ }];
97}125}
98 126 
99def AscendC_LoadDataWithTransposeOp127def AscendC_LoadDataWithTransposeOp
100- : APIOp<"load_data_with_transpose", "LoadDataWithTranspose", [AscFunc]> {128+ : CopyToL0Op<"load_data_with_transpose", "LoadDataWithTranspose", [AscFunc]> {
101 let description = "Perform 2D LoadData with transpose between tensors";129 let description = "Perform 2D LoadData with transpose between tensors";
102- let arguments = (ins AscendC_BaseTensorTypeInterface:$dst,130+ let arguments = (ins AscendC_LocalTensor:$dst, AscendC_LocalTensor:$src,
103- AscendC_BaseTensorTypeInterface:$src,
104 AscendC_LoadData2dTransposeParams:$loadDataParams);131 AscendC_LoadData2dTransposeParams:$loadDataParams);
132+ let extraClassDeclaration = extraClassDeclarationBase # [{
133+ SmallVector<Value> getSrcTensors() { return {getSrc()}; }
134+ }];
105}135}
106 136 
107def AscendC_LoadDataWithTransposeV2Op137def AscendC_LoadDataWithTransposeV2Op
108- : APIOp<"load_data_with_transpose_v2", "LoadDataWithTranspose", [AscFunc]> {138+ : CopyToL0Op<"load_data_with_transpose_v2", "LoadDataWithTranspose", [AscFunc]> {
109 let description = "Perform 2D LoadData with transpose (V2 params) between tensors";139 let description = "Perform 2D LoadData with transpose (V2 params) between tensors";
110- let arguments = (ins AscendC_BaseTensorTypeInterface:$dst,140+ let arguments = (ins AscendC_LocalTensor:$dst, AscendC_LocalTensor:$src,
111- AscendC_BaseTensorTypeInterface:$src,
112 AscendC_LoadData2dTransposeParamsV2:$loadDataParams);141 AscendC_LoadData2dTransposeParamsV2:$loadDataParams);
142+ let extraClassDeclaration = extraClassDeclarationBase # [{
143+ SmallVector<Value> getSrcTensors() { return {getSrc()}; }
144+ }];
113}145}
114 146 
115def AscendC_MmadOp147def AscendC_MmadOp
116- : APIOp<"mmad", "Mmad", [AscFunc]> {148+ : APIOp<"mmad", "Mmad", [AscFunc, OpWithDstInterface, OpWithSrcInterface, OpWithReusableSrcInterface]> {
117 let description = "Matrix multiply-accumulate between LocalTensors (dst += fm * filter)";149 let description = "Matrix multiply-accumulate between LocalTensors (dst += fm * filter)";
118 let arguments = (ins AscendC_BaseTensorTypeInterface:$dst,150 let arguments = (ins AscendC_BaseTensorTypeInterface:$dst,
119 AscendC_BaseTensorTypeInterface:$fm,151 AscendC_BaseTensorTypeInterface:$fm,
120 AscendC_BaseTensorTypeInterface:$filter,152 AscendC_BaseTensorTypeInterface:$filter,
121 AscendC_MmadParams:$mmadParams);153 AscendC_MmadParams:$mmadParams);
154+ let extraClassDeclaration = extraClassDeclarationBase # [{
155+ SmallVector<Value> getSrcTensors() { return {getFm(), getFilter()}; }
156+ }];
122}157}
123 158 
124def AscendC_MmadWithBiasOp159def AscendC_MmadWithBiasOp
125- : APIOp<"mmad_with_bias", "Mmad", [AscFunc]> {160+ : APIOp<"mmad_with_bias", "Mmad", [AscFunc, OpWithDstInterface, OpWithSrcInterface, OpWithReusableSrcInterface]> {
126 let description = "Matrix multiply-accumulate with bias (dst += fm * filter + bias)";161 let description = "Matrix multiply-accumulate with bias (dst += fm * filter + bias)";
127 let arguments = (ins AscendC_BaseTensorTypeInterface:$dst,162 let arguments = (ins AscendC_BaseTensorTypeInterface:$dst,
128 AscendC_BaseTensorTypeInterface:$fm,163 AscendC_BaseTensorTypeInterface:$fm,
129 AscendC_BaseTensorTypeInterface:$filter,164 AscendC_BaseTensorTypeInterface:$filter,
130 AscendC_BaseTensorTypeInterface:$bias,165 AscendC_BaseTensorTypeInterface:$bias,
131 AscendC_MmadParams:$mmadParams);166 AscendC_MmadParams:$mmadParams);
167+ let extraClassDeclaration = extraClassDeclarationBase # [{
168+ SmallVector<Value> getSrcTensors() { return {getFm(), getFilter(), getBias()}; }
169+ }];
132}170}
133 171 
134def AscendC_MmadWithSparseOp172def AscendC_MmadWithSparseOp
@@ -0,0 +1,157 @@
1+/*
2+ * Copyright (c) 2026 Huawei Technologies Co., Ltd.
3+ * This program is free software, you can redistribute it and/or modify it under the terms and conditions of
4+ * 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,
7+ * 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+#ifndef ASC_BASIC_OP_REG_TD
12+#define ASC_BASIC_OP_REG_TD
13+ 
14+include "Base.td"
15+include "Core/Attributes.td"
16+include "Core/Interfaces.td"
17+include "Core/Types.td"
18+ 
19+include "mlir/Interfaces/CastInterfaces.td"
20+include "mlir/Interfaces/SideEffectInterfaces.td"
21+include "mlir/IR/OpBase.td"
22+ 
23+//===----------------------------------------------------------------------===//
24+// Binary register API operations
25+//===----------------------------------------------------------------------===//
26+ 
27+def AddRegOp : BinaryRegOp<"add", "Add">;
28+def AndRegOp : BinaryRegOp<"and", "And">;
29+def DivRegOp : BinaryRegOp<"div", "Div">;
30+def FusedAbsSubRegOp : BinaryRegOp<"fused_abs_sub", "FusedAbsSub">;
31+def FusedExpSubRegOp : BinaryRegOp<"fused_exp_sub", "FusedExpSub">;
32+def FusedMulDstAddRegOp : BinaryRegOp<"fused_mul_dst_add", "FusedMulDstAdd">;
33+def SubRegOp : BinaryRegOp<"sub", "Sub">;
34+def MaxRegOp : BinaryRegOp<"max", "Max">;
35+def MinRegOp : BinaryRegOp<"min", "Min">;
36+def MulRegOp : BinaryRegOp<"mul", "Mul">;
37+def MulAddDstRegOp : BinaryRegOp<"mul_add_dst", "MulAddDst">;
38+def OrRegOp : BinaryRegOp<"or", "Or">;
39+def PreluRegOp : BinaryRegOp<"prelu", "Prelu">;
40+def XorRegOp : BinaryRegOp<"xor", "Xor">;
41+ 
42+//===----------------------------------------------------------------------===//
43+// Unary register API operations
44+//===----------------------------------------------------------------------===//
45+ 
46+def AbsRegOp : UnaryRegOp<"abs", "Abs">;
47+def ExpRegOp : UnaryRegOp<"exp", "Exp">;
48+def LnRegOp : UnaryRegOp<"ln", "Ln">;
49+def LogRegOp : UnaryRegOp<"log", "Log">;
50+def Log10RegOp : UnaryRegOp<"log10", "Log10">;
51+def MaskNotRegOp : UnaryRegOp<"mask_not", "MaskNot">;
52+def NegRegOp : UnaryRegOp<"neg", "Neg">;
53+def NotRegOp : UnaryRegOp<"not", "Not">;
54+def ReluRegOp : UnaryRegOp<"relu", "Relu">;
55+def SqrtRegOp : UnaryRegOp<"sqrt", "Sqrt">;
56+ 
57+//===----------------------------------------------------------------------===//
58+// VecScalar register API operations
59+//===----------------------------------------------------------------------===//
60+ 
61+def AddsRegOp : VecScalarRegOp<"adds", "Adds">;
62+def LeakyReluRegOp : VecScalarRegOp<"leaky_relu", "LeakyRelu">;
63+def MulsRegOp : VecScalarRegOp<"muls", "Muls">;
64+def MaxsRegOp : VecScalarRegOp<"maxs", "Maxs">;
65+def MinsRegOp : VecScalarRegOp<"mins", "Mins">;
66+def ShiftLeftsRegOp : VecScalarRegOp<"shift_lefts", "ShiftLefts">;
67+def ShiftRightsRegOp : VecScalarRegOp<"shift_rights", "ShiftRights">;
68+ 
69+//===----------------------------------------------------------------------===//
70+// Reduce register API operations
71+//===----------------------------------------------------------------------===//
72+ 
73+def ReduceMaxRegOp : RegOp<"reduce_max_reg", "ReduceMax"> {
74+ let arguments = (ins AscendC_RegTensor:$dstReg, AscendC_RegTensor:$srcReg, AscendC_MaskReg:$maskReg);
75+}
76+ 
77+def ReduceSumRegOp : RegOp<"reduce_sum_reg", "ReduceSum"> {
78+ let arguments = (ins AscendC_RegTensor:$dstReg, AscendC_RegTensor:$srcReg, AscendC_MaskReg:$maskReg);
79+}
80+ 
81+def ReduceMinRegOp : RegOp<"reduce_min_reg", "ReduceMin"> {
82+ let arguments = (ins AscendC_RegTensor:$dstReg, AscendC_RegTensor:$srcReg, AscendC_MaskReg:$maskReg);
83+}
84+ 
85+def DuplicateRegOp : RegOp<"duplicate_reg", "Duplicate"> {
86+ let arguments = (ins AscendC_RegTensor:$dstReg, AnyType:$srcReg, AscendC_MaskReg:$maskReg);
87+}
88+ 
89+def DuplicateScalarRegOp : AscendC_Op<"duplicate", []> {
90+ let arguments = (ins AscendC_RegTensor:$dstReg, AnyType:$scalar);
91+ let assemblyFormat = "$dstReg `,` $scalar attr-dict `:` type($dstReg) `,` type($scalar)";
92+}
93+ 
94+//===----------------------------------------------------------------------===//
95+// Other register API operations
96+//===----------------------------------------------------------------------===//
97+ 
98+def SelectRegOp : RegOp<"select_reg", "Select"> {
99+ let arguments = (ins AscendC_RegTensor:$dstReg, AscendC_RegTensor:$src0Reg,
100+ AscendC_RegTensor:$src1Reg, AscendC_MaskReg:$maskReg);
101+}
102+ 
103+def DataCopyLoadOp : RegOp<"data_copy_vld_reg", "DataCopy"> {
104+ let description = "Perform copying between tensors (UB to RegTensor)";
105+ let arguments = (ins AscendC_RegTensor:$dstReg,
106+ AnyRankedOrUnrankedMemRef:$src);
107+}
108+ 
109+def DataCopyStoreOp : RegOp<"data_copy_vst_reg", "DataCopy"> {
110+ let description = "Perform copying between tensors (RegTensor to UB)";
111+ let arguments = (ins AnyRankedOrUnrankedMemRef:$dst,
112+ AscendC_RegTensor:$srcReg,
113+ AscendC_MaskReg:$maskReg);
114+}
115+ 
116+def UpdateMaskOp : RegOp<"update_mask", "UpdateMask"> {
117+ let summary = "Update reg mask";
118+ let comment = "The mask is updated on every iteration to match the total count";
119+ let arguments = (ins Arg<AnyRankedOrUnrankedMemRef, "counter", [MemWrite, MemRead]>:$count, TypeAttr:$type);
120+ let results = (outs AscendC_MaskReg:$maskReg);
121+ let assemblyFormat = [{
122+ $type `,` $count attr-dict `:` type($count)
123+ }];
124+}
125+ 
126+def RegTensorOp : RegOp<"reg_tensor", "RegTensor"> {
127+ let summary = "Create RegTensor";
128+ let results = (outs AscendC_RegTensor:$result);
129+ let assemblyFormat = [{
130+ attr-dict `:` type($result)
131+ }];
132+ let hasCanonicalizeMethod = 1;
133+}
134+ 
135+def LocalMemBarOp : RegOp<"local_mem_bar", "LocalMemBar"> {
136+ let arguments = (ins AscendC_MemType:$src, AscendC_MemType:$dst);
137+ let assemblyFormat = "$src `,` $dst attr-dict";
138+}
139+ 
140+def MaskRegOp : RegOp<"mask_reg", "MaskReg"> {
141+ let summary = "Create RegTensor";
142+ let results = (outs AscendC_MaskReg:$result);
143+ let assemblyFormat = [{
144+ attr-dict `:` type($result)
145+ }];
146+}
147+ 
148+def CreateMaskOp : RegOp<"create_mask", "CreateMask", [Pure]> {
149+ let summary = "Create RegTensor";
150+ let arguments = (ins TypeAttr:$dtype, AscendC_MaskPattern:$mask);
151+ let results = (outs AscendC_MaskReg:$result);
152+ let assemblyFormat = [{
153+ $dtype `,` $mask attr-dict `:` type($result)
154+ }];
155+}
156+ 
157+#endif // ASC_BASIC_OP_REG_TD
@@ -38,6 +38,13 @@ def AscendC_GetBlockNumOp : APIOp<"get_block_num", "GetBlockNum", [Pure]> {
38 let assemblyFormat = "attr-dict `:` type($value)";38 let assemblyFormat = "attr-dict `:` type($value)";
39}39}
40 40 
41+def AscendC_GetVecLenOp : APIOp<"get_vec_len", "GetVecLen"> {
42+ let results = (outs AnySignlessIntegerOrIndex:$oneRepeatSize);
43+ let assemblyFormat = [{
44+ attr-dict `:` qualified(type($oneRepeatSize))
45+ }];
46+}
47+ 
41def AscendC_GetDataBlockSizeInBytesOp : APIOp<"get_data_block_size_in_bytes", "GetDataBlockSizeInBytes", [AscFunc]> {48def AscendC_GetDataBlockSizeInBytesOp : APIOp<"get_data_block_size_in_bytes", "GetDataBlockSizeInBytes", [AscFunc]> {
42 let description = "Obtain the size of a datablock in bytes for the current chip version";49 let description = "Obtain the size of a datablock in bytes for the current chip version";
43 let results = (outs AnySignlessIntegerOrIndex:$value);50 let results = (outs AnySignlessIntegerOrIndex:$value);
@@ -21,11 +21,13 @@ include "mlir/Interfaces/SideEffectInterfaces.td"
21include "mlir/IR/OpBase.td"21include "mlir/IR/OpBase.td"
22 22 
23defm Adds : VecScalarL012Op<"adds", "Adds">;23defm Adds : VecScalarL012Op<"adds", "Adds">;
24+defm Divs : VecScalarL012Op<"divs", "Divs">;
24defm LeakyRelu : VecScalarL012Op<"leaky_relu", "LeakyRelu">;25defm LeakyRelu : VecScalarL012Op<"leaky_relu", "LeakyRelu">;
25defm Maxs : VecScalarL012Op<"maxs", "Maxs">;26defm Maxs : VecScalarL012Op<"maxs", "Maxs">;
26defm Mins : VecScalarL012Op<"mins", "Mins">;27defm Mins : VecScalarL012Op<"mins", "Mins">;
27defm Muls : VecScalarL012Op<"muls", "Muls">;28defm Muls : VecScalarL012Op<"muls", "Muls">;
28defm ShiftLeft : VecScalarL012Op<"shift_left", "ShiftLeft">;29defm ShiftLeft : VecScalarL012Op<"shift_left", "ShiftLeft">;
29defm ShiftRight : VecScalarL012Op<"shift_right", "ShiftRight">;30defm ShiftRight : VecScalarL012Op<"shift_right", "ShiftRight">;
31+defm Subs : VecScalarL012Op<"subs", "Subs">;
30 32 
31#endif //ASC_BASIC_OP_VEC_BINARY_SCALAR_TD33#endif //ASC_BASIC_OP_VEC_BINARY_SCALAR_TD
@@ -20,26 +20,36 @@ include "mlir/Interfaces/CastInterfaces.td"
20include "mlir/Interfaces/SideEffectInterfaces.td"20include "mlir/Interfaces/SideEffectInterfaces.td"
21include "mlir/IR/OpBase.td"21include "mlir/IR/OpBase.td"
22 22 
23-def AscendC_CompareL0Op : VectorOp<"compare_l0", "Compare", [OpWithDstInterface, AscFunc]> {23+def AscendC_CompareL0Op : VectorOp<"compare_l0", "Compare", [
24+ AscFunc, OpWithDstInterface, OpWithSrcInterface, OpWithReusableSrcInterface
25+]> {
24 let description = "`AscendC::Compare` is a vector binary operation (L0 API).\n";26 let description = "`AscendC::Compare` is a vector binary operation (L0 API).\n";
25 let arguments = (ins AnyType:$dst, AnyType:$src0, AnyType:$src1, AscendC_CMPMODEAttr:$cmpMode,27 let arguments = (ins AnyType:$dst, AnyType:$src0, AnyType:$src1, AscendC_CMPMODEAttr:$cmpMode,
26- AnyType:$mask, AnyType:$repeatTimes, 28+ AnyType:$mask, AnyType:$repeatTimes,
27 AscendC_BinaryRepeatParams:$repeatParams, UnitAttr:$isSetMask);29 AscendC_BinaryRepeatParams:$repeatParams, UnitAttr:$isSetMask);
28 let paramTypeLists = [0, 0, 0, -1, 0, 0, 0];30 let paramTypeLists = [0, 0, 0, -1, 0, 0, 0];
31+ let extraClassDeclaration = extraClassDeclarationBase # [{
32+ SmallVector<Value> getSrcTensors() { return {getSrc0(), getSrc1()}; }
33+ }];
29}34}
30 35 
31def AscendC_CompareL1Op : VectorOp<"compare_l1", "Compare", [OpWithDstInterface]> {36def AscendC_CompareL1Op : VectorOp<"compare_l1", "Compare", [OpWithDstInterface]> {
32 let description = "`AscendC::Compare` is a vector binary operation (L1 API).\n";37 let description = "`AscendC::Compare` is a vector binary operation (L1 API).\n";
33 let arguments = (ins AnyType:$dst, AnyType:$src0, AnyType:$src1, AscendC_CMPMODEAttr:$cmpMode,38 let arguments = (ins AnyType:$dst, AnyType:$src0, AnyType:$src1, AscendC_CMPMODEAttr:$cmpMode,
34- Variadic<UI64>:$mask, AnyType:$repeatTimes, 39+ Variadic<UI64>:$mask, AnyType:$repeatTimes,
35 AscendC_BinaryRepeatParams:$repeatParams, UnitAttr:$isSetMask);40 AscendC_BinaryRepeatParams:$repeatParams, UnitAttr:$isSetMask);
36}41}
37 42 
38-def AscendC_CompareL2Op : VectorOp<"compare_l2", "Compare", [OpWithDstInterface, AscFunc]> {43+def AscendC_CompareL2Op : VectorOp<"compare_l2", "Compare", [
44+ AscFunc, OpWithDstInterface, OpWithSrcInterface, OpWithReusableSrcInterface
45+]> {
39 let description = "`AscendC::Compare` is a vector binary operation (L2 API).\n";46 let description = "`AscendC::Compare` is a vector binary operation (L2 API).\n";
40 let arguments = (ins AnyType:$dst, AnyType:$src0, AnyType:$src1, AscendC_CMPMODEAttr:$cmpMode,47 let arguments = (ins AnyType:$dst, AnyType:$src0, AnyType:$src1, AscendC_CMPMODEAttr:$cmpMode,
41 AnyType:$calCount);48 AnyType:$calCount);
42 let paramTypeLists = [0, 0, 0, -1, 0];49 let paramTypeLists = [0, 0, 0, -1, 0];
50+ let extraClassDeclaration = extraClassDeclarationBase # [{
51+ SmallVector<Value> getSrcTensors() { return {getSrc0(), getSrc1()}; }
52+ }];
43}53}
44 54 
45def AscendC_CompareRL0Op : VectorOp<"compare_r_l0", "Compare", [AscFunc]> {55def AscendC_CompareRL0Op : VectorOp<"compare_r_l0", "Compare", [AscFunc]> {
@@ -55,26 +65,37 @@ def AscendC_CompareRL1Op : VectorOp<"compare_r_l1", "Compare"> {
55 Variadic<UI64>:$mask, AscendC_BinaryRepeatParams:$repeatParams, UnitAttr:$isSetMask);65 Variadic<UI64>:$mask, AscendC_BinaryRepeatParams:$repeatParams, UnitAttr:$isSetMask);
56}66}
57 67 
58-def AscendC_CompareScalarL0Op : VectorOp<"compare_scalar_l0", "CompareScalar", [OpWithDstInterface, AscFunc]> {68+def AscendC_CompareScalarL0Op : VectorOp<"compare_scalar_l0", "CompareScalar", [
69+ AscFunc, OpWithDstInterface, OpWithSrcInterface, OpWithReusableSrcInterface
70+]> {
59 let description = "`AscendC::CompareScalar` is a vector-scalar binary operation (L0 API).\n";71 let description = "`AscendC::CompareScalar` is a vector-scalar binary operation (L0 API).\n";
60 let arguments = (ins AnyType:$dst, AnyType:$src0, AnyType:$src1Scalar, AscendC_CMPMODEAttr:$cmpMode,72 let arguments = (ins AnyType:$dst, AnyType:$src0, AnyType:$src1Scalar, AscendC_CMPMODEAttr:$cmpMode,
61- AnyType:$mask, AnyType:$repeatTimes, 73+ AnyType:$mask, AnyType:$repeatTimes,
62 AscendC_UnaryRepeatParams:$repeatParams, UnitAttr:$isSetMask);74 AscendC_UnaryRepeatParams:$repeatParams, UnitAttr:$isSetMask);
63 let paramTypeLists = [0, 0, 0, -1, 0, 0, 0];75 let paramTypeLists = [0, 0, 0, -1, 0, 0, 0];
76+ let extraClassDeclaration = extraClassDeclarationBase # [{
77+ Value getSrc() { return {getSrc0()}; }
78+ SmallVector<Value> getSrcTensors() { return {getSrc0()}; }
79+ }];
64}80}
65 81 
66def AscendC_CompareScalarL1Op : VectorOp<"compare_scalar_l1", "CompareScalar", [OpWithDstInterface]> {82def AscendC_CompareScalarL1Op : VectorOp<"compare_scalar_l1", "CompareScalar", [OpWithDstInterface]> {
67 let description = "`AscendC::CompareScalar` is a vector-scalar binary operation (L1 API).\n";83 let description = "`AscendC::CompareScalar` is a vector-scalar binary operation (L1 API).\n";
68 let arguments = (ins AnyType:$dst, AnyType:$src0, AnyType:$src1Scalar, AscendC_CMPMODEAttr:$cmpMode,84 let arguments = (ins AnyType:$dst, AnyType:$src0, AnyType:$src1Scalar, AscendC_CMPMODEAttr:$cmpMode,
69- Variadic<UI64>:$mask, AnyType:$repeatTimes, 85+ Variadic<UI64>:$mask, AnyType:$repeatTimes,
70 AscendC_UnaryRepeatParams:$repeatParams, UnitAttr:$isSetMask);86 AscendC_UnaryRepeatParams:$repeatParams, UnitAttr:$isSetMask);
71}87}
72 88 
73-def AscendC_CompareScalarL2Op : VectorOp<"compare_scalar_l2", "CompareScalar", [OpWithDstInterface, AscFunc]> {89+def AscendC_CompareScalarL2Op : VectorOp<"compare_scalar_l2", "CompareScalar", [
90+ AscFunc, OpWithDstInterface, OpWithSrcInterface, OpWithReusableSrcInterface
91+]> {
74 let description = "`AscendC::CompareScalar` is a vector-scalar binary operation (L2 API).\n";92 let description = "`AscendC::CompareScalar` is a vector-scalar binary operation (L2 API).\n";
75 let arguments = (ins AnyType:$dst, AnyType:$src0, AnyType:$src1Scalar, AscendC_CMPMODEAttr:$cmpMode,93 let arguments = (ins AnyType:$dst, AnyType:$src0, AnyType:$src1Scalar, AscendC_CMPMODEAttr:$cmpMode,
76 AnyType:$calCount);94 AnyType:$calCount);
77 let paramTypeLists = [0, 0, 0, -1, 0];95 let paramTypeLists = [0, 0, 0, -1, 0];
96+ let extraClassDeclaration = extraClassDeclarationBase # [{
97+ SmallVector<Value> getSrcTensors() { return {getSrc0()}; }
98+ }];
78}99}
79 100 
80def AscendC_GetCmpMaskOp : VectorOp<"get_cmp_mask", "GetCmpMask", [OpWithDstInterface, AscFunc]> {101def AscendC_GetCmpMaskOp : VectorOp<"get_cmp_mask", "GetCmpMask", [OpWithDstInterface, AscFunc]> {
@@ -87,45 +108,66 @@ def AscendC_SetCmpMaskOp : VectorOp<"set_cmp_mask", "SetCmpMask", [AscFunc]> {
87 let arguments = (ins AscendC_LocalTensor:$src);108 let arguments = (ins AscendC_LocalTensor:$src);
88}109}
89 110 
90-def AscendC_SelectScalarL2Op : VectorOp<"select_scalar_l2", "Select", [OpWithDstInterface, AscFunc]> {111+def AscendC_SelectScalarL2Op : VectorOp<"select_scalar_l2", "Select", [
112+ AscFunc, OpWithDstInterface, OpWithSrcInterface, OpWithReusableSrcInterface
113+]> {
91 let description = "Select elements from source tensor and a scalar (L2 API)";114 let description = "Select elements from source tensor and a scalar (L2 API)";
92 let arguments = (ins AscendC_LocalTensor:$dst, AscendC_LocalTensor:$selMask,115 let arguments = (ins AscendC_LocalTensor:$dst, AscendC_LocalTensor:$selMask,
93 AscendC_LocalTensor:$src0, AnyType:$src1,116 AscendC_LocalTensor:$src0, AnyType:$src1,
94 AscendC_SelModeAttr:$selMode, AnyType:$calCount);117 AscendC_SelModeAttr:$selMode, AnyType:$calCount);
95 let paramTypeLists = [0, 0, 0, 0, -1, 0];118 let paramTypeLists = [0, 0, 0, 0, -1, 0];
119+ let extraClassDeclaration = extraClassDeclarationBase # [{
120+ SmallVector<Value> getSrcTensors() {
121+ return {getSelMask(), getSrc0()};
122+ }
123+ }];
96}124}
97 125 
98-def AscendC_SelectL2Op : VectorOp<"select_l2", "Select", [OpWithDstInterface, AscFunc]> {126+def AscendC_SelectL2Op : VectorOp<"select_l2", "Select", [
127+ AscFunc, OpWithDstInterface, OpWithSrcInterface, OpWithReusableSrcInterface
128+]> {
99 let description = "Select elements from two source tensors (L2 API)";129 let description = "Select elements from two source tensors (L2 API)";
100 let arguments = (ins AscendC_LocalTensor:$dst, AscendC_LocalTensor:$selMask,130 let arguments = (ins AscendC_LocalTensor:$dst, AscendC_LocalTensor:$selMask,
101 AscendC_LocalTensor:$src0, AscendC_LocalTensor:$src1,131 AscendC_LocalTensor:$src0, AscendC_LocalTensor:$src1,
102 AscendC_SelModeAttr:$selMode, AnyType:$calCount);132 AscendC_SelModeAttr:$selMode, AnyType:$calCount);
103 let paramTypeLists = [0, 0, 0, 0, -1, 0];133 let paramTypeLists = [0, 0, 0, 0, -1, 0];
134+ let extraClassDeclaration = extraClassDeclarationBase # [{
135+ SmallVector<Value> getSrcTensors() {
136+ return {getSelMask(), getSrc0(), getSrc1()};
137+ }
138+ }];
104}139}
105 140 
106def AscendC_SelectScalarL1Op : VectorOp<"select_scalar_l1", "Select", [OpWithDstInterface]> {141def AscendC_SelectScalarL1Op : VectorOp<"select_scalar_l1", "Select", [OpWithDstInterface]> {
107 let description = "Select elements from source tensor and a scalar (L1 API)";142 let description = "Select elements from source tensor and a scalar (L1 API)";
108 let arguments = (ins AscendC_LocalTensor:$dst, AscendC_LocalTensor:$selMask,143 let arguments = (ins AscendC_LocalTensor:$dst, AscendC_LocalTensor:$selMask,
109 AscendC_LocalTensor:$src0, AnyType:$src1,144 AscendC_LocalTensor:$src0, AnyType:$src1,
110- AscendC_SelModeAttr:$selMode, Variadic<UI64>:$mask, 145+ AscendC_SelModeAttr:$selMode, Variadic<UI64>:$mask,
111- AnyType:$repeatTimes, AscendC_BinaryRepeatParams:$repeatParams, 146+ AnyType:$repeatTimes, AscendC_BinaryRepeatParams:$repeatParams,
112 UnitAttr:$isSetMask);147 UnitAttr:$isSetMask);
113}148}
114 149 
115-def AscendC_SelectScalarL0Op : VectorOp<"select_scalar_l0", "Select", [OpWithDstInterface, AscFunc]> {150+def AscendC_SelectScalarL0Op : VectorOp<"select_scalar_l0", "Select", [
151+ AscFunc, OpWithDstInterface, OpWithSrcInterface, OpWithReusableSrcInterface
152+]> {
116 let description = "Select elements from source tensor and a scalar (L0 API)";153 let description = "Select elements from source tensor and a scalar (L0 API)";
117 let arguments = (ins AscendC_LocalTensor:$dst, AscendC_LocalTensor:$selMask,154 let arguments = (ins AscendC_LocalTensor:$dst, AscendC_LocalTensor:$selMask,
118 AscendC_LocalTensor:$src0, AnyType:$src1,155 AscendC_LocalTensor:$src0, AnyType:$src1,
119- AscendC_SelModeAttr:$selMode, AnyType:$mask, 156+ AscendC_SelModeAttr:$selMode, AnyType:$mask,
120- AnyType:$repeatTimes, AscendC_BinaryRepeatParams:$repeatParams, 157+ AnyType:$repeatTimes, AscendC_BinaryRepeatParams:$repeatParams,
121 UnitAttr:$isSetMask);158 UnitAttr:$isSetMask);
122 let paramTypeLists = [0, 0, 0, 0, -1, 0, 0, 0];159 let paramTypeLists = [0, 0, 0, 0, -1, 0, 0, 0];
160+ let extraClassDeclaration = extraClassDeclarationBase # [{
161+ SmallVector<Value> getSrcTensors() {
162+ return {getSelMask(), getSrc0()};
163+ }
164+ }];
123}165}
124 166 
125-def AscendC_SelectScalarRegOp : VectorOp<"select_scalar_reg", "Select", [OpWithDstInterface, AscFunc]> {167+def AscendC_SelectScalarRegMaskOp : VectorOp<"select_scalar_regmask", "Select", [OpWithDstInterface, AscFunc]> {
126 let description = "Select elements from source tensor and a scalar (register mask API)";168 let description = "Select elements from source tensor and a scalar (register mask API)";
127 let arguments = (ins AscendC_LocalTensor:$dst, AscendC_LocalTensor:$selMask,169 let arguments = (ins AscendC_LocalTensor:$dst, AscendC_LocalTensor:$selMask,
128- AscendC_LocalTensor:$src0, AnyType:$repeatTimes, 170+ AscendC_LocalTensor:$src0, AnyType:$repeatTimes,
129 AscendC_BinaryRepeatParams:$repeatParams);171 AscendC_BinaryRepeatParams:$repeatParams);
130 let paramTypeLists = [0, 0, 0, 0, 0];172 let paramTypeLists = [0, 0, 0, 0, 0];
131}173}
@@ -134,26 +176,34 @@ def AscendC_SelectL1Op : VectorOp<"select_l1", "Select", [OpWithDstInterface]> {
134 let description = "Select elements from two source tensors (L1 API)";176 let description = "Select elements from two source tensors (L1 API)";
135 let arguments = (ins AscendC_LocalTensor:$dst, AscendC_LocalTensor:$selMask,177 let arguments = (ins AscendC_LocalTensor:$dst, AscendC_LocalTensor:$selMask,
136 AscendC_LocalTensor:$src0, AscendC_LocalTensor:$src1,178 AscendC_LocalTensor:$src0, AscendC_LocalTensor:$src1,
137- AscendC_SelModeAttr:$selMode, Variadic<UI64>:$mask, 179+ AscendC_SelModeAttr:$selMode, Variadic<UI64>:$mask,
138- AnyType:$repeatTimes, AscendC_BinaryRepeatParams:$repeatParams, 180+ AnyType:$repeatTimes, AscendC_BinaryRepeatParams:$repeatParams,
139 UnitAttr:$isSetMask);181 UnitAttr:$isSetMask);
140}182}
141 183 
142-def AscendC_SelectL0Op : VectorOp<"select_l0", "Select", [OpWithDstInterface, AscFunc]> {184+def AscendC_SelectL0Op : VectorOp<"select_l0", "Select", [
185+ AscFunc, OpWithDstInterface, OpWithSrcInterface, OpWithReusableSrcInterface
186+]> {
143 let description = "Select elements from two source tensors (L0 API)";187 let description = "Select elements from two source tensors (L0 API)";
144 let arguments = (ins AscendC_LocalTensor:$dst, AscendC_LocalTensor:$selMask,188 let arguments = (ins AscendC_LocalTensor:$dst, AscendC_LocalTensor:$selMask,
145 AscendC_LocalTensor:$src0, AscendC_LocalTensor:$src1,189 AscendC_LocalTensor:$src0, AscendC_LocalTensor:$src1,
146- AscendC_SelModeAttr:$selMode, AnyType:$mask, 190+ AscendC_SelModeAttr:$selMode, AnyType:$mask,
147- AnyType:$repeatTimes, AscendC_BinaryRepeatParams:$repeatParams, 191+ AnyType:$repeatTimes, AscendC_BinaryRepeatParams:$repeatParams,
148 UnitAttr:$isSetMask);192 UnitAttr:$isSetMask);
149 let paramTypeLists = [0, 0, 0, 0, -1, 0, 0, 0];193 let paramTypeLists = [0, 0, 0, 0, -1, 0, 0, 0];
194+ let extraClassDeclaration = extraClassDeclarationBase # [{
195+ SmallVector<Value> getSrcTensors() {
196+ return {getSelMask(), getSrc0(), getSrc1()};
197+ }
198+ }];
150}199}
151 200 
152-def AscendC_SelectRegOp : VectorOp<"select_reg", "Select", [OpWithDstInterface, AscFunc]> {201+def AscendC_SelectRegMaskOp : VectorOp<"select_regmask", "Select", [OpWithDstInterface, AscFunc]> {
153 let description = "Select elements from two source tensors (register mask API)";202 let description = "Select elements from two source tensors (register mask API)";
154 let arguments = (ins AscendC_LocalTensor:$dst, AscendC_LocalTensor:$src0, AscendC_LocalTensor:$src1,203 let arguments = (ins AscendC_LocalTensor:$dst, AscendC_LocalTensor:$src0, AscendC_LocalTensor:$src1,
155 AnyType:$repeatTimes, AscendC_BinaryRepeatParams:$repeatParams,204 AnyType:$repeatTimes, AscendC_BinaryRepeatParams:$repeatParams,
156 AscendC_SelModeAttr:$selMode);205 AscendC_SelModeAttr:$selMode);
157 let paramTypeLists = [0, 0, 0, 0, 0];206 let paramTypeLists = [0, 0, 0, 0, 0];
158}207}
159-#endif //ASC_BASIC_OP_VEC_CMPSEL_TD208+ 
209+#endif // ASC_BASIC_OP_VEC_CMPSEL_TD
@@ -20,6 +20,6 @@ include "mlir/Interfaces/CastInterfaces.td"
20include "mlir/Interfaces/SideEffectInterfaces.td"20include "mlir/Interfaces/SideEffectInterfaces.td"
21include "mlir/IR/OpBase.td"21include "mlir/IR/OpBase.td"
22 22 
23-defm MulCast : BinaryL012Op<"mul_cast", "MulCast">;23+defm MulCast : BinaryCastL012Op<"mul_cast", "MulCast">;
24 24 
25#endif //ASC_BASIC_OP_MULCAST_TD25#endif //ASC_BASIC_OP_MULCAST_TD
@@ -26,7 +26,7 @@ include "mlir/IR/OpBase.td"
26 26 
27def AscendC_BlockReduceSumL0Op : VectorOp<"block_reduce_sum_l0", "BlockReduceSum", [AscFunc]> {27def AscendC_BlockReduceSumL0Op : VectorOp<"block_reduce_sum_l0", "BlockReduceSum", [AscFunc]> {
28 let description = "Sum all elements in each block (single mask version)";28 let description = "Sum all elements in each block (single mask version)";
29- let arguments = (ins 29+ let arguments = (ins
30 AscendC_LocalTensor:$dst,30 AscendC_LocalTensor:$dst,
31 AscendC_LocalTensor:$src,31 AscendC_LocalTensor:$src,
32 AnyType:$repeatTime,32 AnyType:$repeatTime,
@@ -40,7 +40,7 @@ def AscendC_BlockReduceSumL0Op : VectorOp<"block_reduce_sum_l0", "BlockReduceSum
40 40 
41def AscendC_BlockReduceSumL1Op : VectorOp<"block_reduce_sum_l1", "BlockReduceSum", [OpWithDstInterface]> {41def AscendC_BlockReduceSumL1Op : VectorOp<"block_reduce_sum_l1", "BlockReduceSum", [OpWithDstInterface]> {
42 let description = "Sum all elements in each block (mask array version)";42 let description = "Sum all elements in each block (mask array version)";
43- let arguments = (ins 43+ let arguments = (ins
44 AscendC_LocalTensor:$dst,44 AscendC_LocalTensor:$dst,
45 AscendC_LocalTensor:$src,45 AscendC_LocalTensor:$src,
46 AnyType:$repeatTime,46 AnyType:$repeatTime,
@@ -53,7 +53,7 @@ def AscendC_BlockReduceSumL1Op : VectorOp<"block_reduce_sum_l1", "BlockReduceSum
53 53 
54def AscendC_BlockReduceMaxL0Op : VectorOp<"block_reduce_max_l0", "BlockReduceMax", [AscFunc]> {54def AscendC_BlockReduceMaxL0Op : VectorOp<"block_reduce_max_l0", "BlockReduceMax", [AscFunc]> {
55 let description = "Find maximum value in each block (single mask version)";55 let description = "Find maximum value in each block (single mask version)";
56- let arguments = (ins 56+ let arguments = (ins
57 AscendC_LocalTensor:$dst,57 AscendC_LocalTensor:$dst,
58 AscendC_LocalTensor:$src,58 AscendC_LocalTensor:$src,
59 AnyType:$repeatTime,59 AnyType:$repeatTime,
@@ -67,7 +67,7 @@ def AscendC_BlockReduceMaxL0Op : VectorOp<"block_reduce_max_l0", "BlockReduceMax
67 67 
68def AscendC_BlockReduceMaxL1Op : VectorOp<"block_reduce_max_l1", "BlockReduceMax", [OpWithDstInterface]> {68def AscendC_BlockReduceMaxL1Op : VectorOp<"block_reduce_max_l1", "BlockReduceMax", [OpWithDstInterface]> {
69 let description = "Find maximum value in each block (mask array version)";69 let description = "Find maximum value in each block (mask array version)";
70- let arguments = (ins 70+ let arguments = (ins
71 AscendC_LocalTensor:$dst,71 AscendC_LocalTensor:$dst,
72 AscendC_LocalTensor:$src,72 AscendC_LocalTensor:$src,
73 AnyType:$repeatTime,73 AnyType:$repeatTime,
@@ -80,7 +80,7 @@ def AscendC_BlockReduceMaxL1Op : VectorOp<"block_reduce_max_l1", "BlockReduceMax
80 80 
81def AscendC_BlockReduceMinL0Op : VectorOp<"block_reduce_min_l0", "BlockReduceMin", [AscFunc]> {81def AscendC_BlockReduceMinL0Op : VectorOp<"block_reduce_min_l0", "BlockReduceMin", [AscFunc]> {
82 let description = "Find minimum value in each block (single mask version)";82 let description = "Find minimum value in each block (single mask version)";
83- let arguments = (ins 83+ let arguments = (ins
84 AscendC_LocalTensor:$dst,84 AscendC_LocalTensor:$dst,
85 AscendC_LocalTensor:$src,85 AscendC_LocalTensor:$src,
86 AnyType:$repeatTime,86 AnyType:$repeatTime,
@@ -94,7 +94,7 @@ def AscendC_BlockReduceMinL0Op : VectorOp<"block_reduce_min_l0", "BlockReduceMin
94 94 
95def AscendC_BlockReduceMinL1Op : VectorOp<"block_reduce_min_l1", "BlockReduceMin", [OpWithDstInterface]> {95def AscendC_BlockReduceMinL1Op : VectorOp<"block_reduce_min_l1", "BlockReduceMin", [OpWithDstInterface]> {
96 let description = "Find minimum value in each block (mask array version)";96 let description = "Find minimum value in each block (mask array version)";
97- let arguments = (ins 97+ let arguments = (ins
98 AscendC_LocalTensor:$dst,98 AscendC_LocalTensor:$dst,
99 AscendC_LocalTensor:$src,99 AscendC_LocalTensor:$src,
100 AnyType:$repeatTime,100 AnyType:$repeatTime,
@@ -107,11 +107,11 @@ def AscendC_BlockReduceMinL1Op : VectorOp<"block_reduce_min_l1", "BlockReduceMin
107 107 
108def AscendC_PairReduceSumL0Op : VectorOp<"pair_reduce_sum_l0", "PairReduceSum", [AscFunc]> {108def AscendC_PairReduceSumL0Op : VectorOp<"pair_reduce_sum_l0", "PairReduceSum", [AscFunc]> {
109 let description = "Pair reduce sum operation with continuous mask (L0 API)";109 let description = "Pair reduce sum operation with continuous mask (L0 API)";
110- let arguments = (ins 110+ let arguments = (ins
111 AscendC_LocalTensor:$dst,111 AscendC_LocalTensor:$dst,
112 AscendC_LocalTensor:$src,112 AscendC_LocalTensor:$src,
113 AnyType:$repeatTime,113 AnyType:$repeatTime,
114- AnyType:$mask, 114+ AnyType:$mask,
115 AnyType:$dstRepStride,115 AnyType:$dstRepStride,
116 AnyType:$srcBlkStride,116 AnyType:$srcBlkStride,
117 AnyType:$srcRepStride,117 AnyType:$srcRepStride,
@@ -122,11 +122,11 @@ def AscendC_PairReduceSumL0Op : VectorOp<"pair_reduce_sum_l0", "PairReduceSum",
122 122 
123def AscendC_PairReduceSumL1Op : VectorOp<"pair_reduce_sum_l1", "PairReduceSum", [OpWithDstInterface]> {123def AscendC_PairReduceSumL1Op : VectorOp<"pair_reduce_sum_l1", "PairReduceSum", [OpWithDstInterface]> {
124 let description = "Pair reduce sum operation with bit mask (L1 API)";124 let description = "Pair reduce sum operation with bit mask (L1 API)";
125- let arguments = (ins 125+ let arguments = (ins
126 AscendC_LocalTensor:$dst,126 AscendC_LocalTensor:$dst,
127 AscendC_LocalTensor:$src,127 AscendC_LocalTensor:$src,
128 AnyType:$repeatTime,128 AnyType:$repeatTime,
129- Variadic<UI64>:$mask, 129+ Variadic<UI64>:$mask,
130 AnyType:$dstRepStride,130 AnyType:$dstRepStride,
131 AnyType:$srcBlkStride,131 AnyType:$srcBlkStride,
132 AnyType:$srcRepStride,132 AnyType:$srcRepStride,
@@ -136,7 +136,7 @@ def AscendC_PairReduceSumL1Op : VectorOp<"pair_reduce_sum_l1", "PairReduceSum",
136 136 
137def AscendC_RepeatReduceSumL0Op : VectorOp<"repeat_reduce_sum_l0", "RepeatReduceSum", [AscFunc]> {137def AscendC_RepeatReduceSumL0Op : VectorOp<"repeat_reduce_sum_l0", "RepeatReduceSum", [AscFunc]> {
138 let description = "Repeat reduce sum operation (L0 API)";138 let description = "Repeat reduce sum operation (L0 API)";
139- let arguments = (ins 139+ let arguments = (ins
140 AscendC_LocalTensor:$dst,140 AscendC_LocalTensor:$dst,
141 AscendC_LocalTensor:$src,141 AscendC_LocalTensor:$src,
142 AnyType:$repeatTime,142 AnyType:$repeatTime,
@@ -152,10 +152,10 @@ def AscendC_RepeatReduceSumL0Op : VectorOp<"repeat_reduce_sum_l0", "RepeatReduce
152 152 
153def AscendC_WholeReduceMaxL0Op : VectorOp<"whole_reduce_max_l0", "WholeReduceMax", [AscFunc]> {153def AscendC_WholeReduceMaxL0Op : VectorOp<"whole_reduce_max_l0", "WholeReduceMax", [AscFunc]> {
154 let description = "Whole reduce max operation with continuous mask (L0 API)";154 let description = "Whole reduce max operation with continuous mask (L0 API)";
155- let arguments = (ins 155+ let arguments = (ins
156 AscendC_LocalTensor:$dst,156 AscendC_LocalTensor:$dst,
157 AscendC_LocalTensor:$src,157 AscendC_LocalTensor:$src,
158- AnyType:$mask, 158+ AnyType:$mask,
159 AnyType:$repeatTime,159 AnyType:$repeatTime,
160 AnyType:$dstRepStride,160 AnyType:$dstRepStride,
161 AnyType:$srcBlkStride,161 AnyType:$srcBlkStride,
@@ -168,10 +168,10 @@ def AscendC_WholeReduceMaxL0Op : VectorOp<"whole_reduce_max_l0", "WholeReduceMax
168 168 
169def AscendC_WholeReduceMaxL1Op : VectorOp<"whole_reduce_max_l1", "WholeReduceMax", [OpWithDstInterface]> {169def AscendC_WholeReduceMaxL1Op : VectorOp<"whole_reduce_max_l1", "WholeReduceMax", [OpWithDstInterface]> {
170 let description = "Whole reduce max operation with bit mask (L1 API)";170 let description = "Whole reduce max operation with bit mask (L1 API)";
171- let arguments = (ins 171+ let arguments = (ins
172 AscendC_LocalTensor:$dst,172 AscendC_LocalTensor:$dst,
173 AscendC_LocalTensor:$src,173 AscendC_LocalTensor:$src,
174- Variadic<UI64>:$mask, 174+ Variadic<UI64>:$mask,
175 AnyType:$repeatTime,175 AnyType:$repeatTime,
176 AnyType:$dstRepStride,176 AnyType:$dstRepStride,
177 AnyType:$srcBlkStride,177 AnyType:$srcBlkStride,
@@ -183,10 +183,10 @@ def AscendC_WholeReduceMaxL1Op : VectorOp<"whole_reduce_max_l1", "WholeReduceMax
183 183 
184def AscendC_WholeReduceMinL0Op : VectorOp<"whole_reduce_min_l0", "WholeReduceMin", [AscFunc]> {184def AscendC_WholeReduceMinL0Op : VectorOp<"whole_reduce_min_l0", "WholeReduceMin", [AscFunc]> {
185 let description = "Whole reduce min operation with continuous mask (L0 API)";185 let description = "Whole reduce min operation with continuous mask (L0 API)";
186- let arguments = (ins 186+ let arguments = (ins
187 AscendC_LocalTensor:$dst,187 AscendC_LocalTensor:$dst,
188 AscendC_LocalTensor:$src,188 AscendC_LocalTensor:$src,
189- AnyType:$mask, 189+ AnyType:$mask,
190 AnyType:$repeatTime,190 AnyType:$repeatTime,
191 AnyType:$dstRepStride,191 AnyType:$dstRepStride,
192 AnyType:$srcBlkStride,192 AnyType:$srcBlkStride,
@@ -199,10 +199,10 @@ def AscendC_WholeReduceMinL0Op : VectorOp<"whole_reduce_min_l0", "WholeReduceMin
199 199 
200def AscendC_WholeReduceMinL1Op : VectorOp<"whole_reduce_min_l1", "WholeReduceMin", [OpWithDstInterface]> {200def AscendC_WholeReduceMinL1Op : VectorOp<"whole_reduce_min_l1", "WholeReduceMin", [OpWithDstInterface]> {
201 let description = "Whole reduce min operation with bit mask (L1 API)";201 let description = "Whole reduce min operation with bit mask (L1 API)";
202- let arguments = (ins 202+ let arguments = (ins
203 AscendC_LocalTensor:$dst,203 AscendC_LocalTensor:$dst,
204 AscendC_LocalTensor:$src,204 AscendC_LocalTensor:$src,
205- Variadic<UI64>:$mask, 205+ Variadic<UI64>:$mask,
206 AnyType:$repeatTime,206 AnyType:$repeatTime,
207 AnyType:$dstRepStride,207 AnyType:$dstRepStride,
208 AnyType:$srcBlkStride,208 AnyType:$srcBlkStride,
@@ -214,10 +214,10 @@ def AscendC_WholeReduceMinL1Op : VectorOp<"whole_reduce_min_l1", "WholeReduceMin
214 214 
215def AscendC_WholeReduceSumL0Op : VectorOp<"whole_reduce_sum_l0", "WholeReduceSum", [AscFunc]> {215def AscendC_WholeReduceSumL0Op : VectorOp<"whole_reduce_sum_l0", "WholeReduceSum", [AscFunc]> {
216 let description = "Whole reduce sum operation with continuous mask (L0 API)";216 let description = "Whole reduce sum operation with continuous mask (L0 API)";
217- let arguments = (ins 217+ let arguments = (ins
218 AscendC_LocalTensor:$dst,218 AscendC_LocalTensor:$dst,
219 AscendC_LocalTensor:$src,219 AscendC_LocalTensor:$src,
220- AnyType:$mask, 220+ AnyType:$mask,
221 AnyType:$repeatTime,221 AnyType:$repeatTime,
222 AnyType:$dstRepStride,222 AnyType:$dstRepStride,
223 AnyType:$srcBlkStride,223 AnyType:$srcBlkStride,
@@ -229,10 +229,10 @@ def AscendC_WholeReduceSumL0Op : VectorOp<"whole_reduce_sum_l0", "WholeReduceSum
229 229 
230def AscendC_WholeReduceSumL1Op : VectorOp<"whole_reduce_sum_l1", "WholeReduceSum", [OpWithDstInterface]> {230def AscendC_WholeReduceSumL1Op : VectorOp<"whole_reduce_sum_l1", "WholeReduceSum", [OpWithDstInterface]> {
231 let description = "Whole reduce sum operation with bit mask (L1 API)";231 let description = "Whole reduce sum operation with bit mask (L1 API)";
232- let arguments = (ins 232+ let arguments = (ins
233 AscendC_LocalTensor:$dst,233 AscendC_LocalTensor:$dst,
234 AscendC_LocalTensor:$src,234 AscendC_LocalTensor:$src,
235- Variadic<UI64>:$mask, 235+ Variadic<UI64>:$mask,
236 AnyType:$repeatTime,236 AnyType:$repeatTime,
237 AnyType:$dstRepStride,237 AnyType:$dstRepStride,
238 AnyType:$srcBlkStride,238 AnyType:$srcBlkStride,
@@ -245,9 +245,11 @@ def AscendC_WholeReduceSumL1Op : VectorOp<"whole_reduce_sum_l1", "WholeReduceSum
245// ReduceMax operations245// ReduceMax operations
246//===----------------------------------------------------------------------===//246//===----------------------------------------------------------------------===//
247 247 
248-def AscendC_ReduceMaxL0Op : VectorOp<"reduce_max_l0", "ReduceMax", [OpWithDstInterface, AscFunc]> {248+def AscendC_ReduceMaxL0Op : VectorOp<"reduce_max_l0", "ReduceMax", [
249+ AscFunc, OpWithDstInterface, OpWithSrcInterface
250+]> {
249 let description = "Reduce maximum elements (L0 API). ";251 let description = "Reduce maximum elements (L0 API). ";
250- let arguments = (ins 252+ let arguments = (ins
251 AscendC_LocalTensor:$dst,253 AscendC_LocalTensor:$dst,
252 AscendC_LocalTensor:$src,254 AscendC_LocalTensor:$src,
253 AscendC_LocalTensor:$sharedTmpBuffer,255 AscendC_LocalTensor:$sharedTmpBuffer,
@@ -257,11 +259,15 @@ def AscendC_ReduceMaxL0Op : VectorOp<"reduce_max_l0", "ReduceMax", [OpWithDstInt
257 AnyType:$calIndex259 AnyType:$calIndex
258 );260 );
259 let paramTypeLists = [0, 0, 0, 0, 0, 0, 0];261 let paramTypeLists = [0, 0, 0, 0, 0, 0, 0];
262+ let extraClassDeclaration = extraClassDeclarationBase # [{
263+ SmallVector<Value> getDstTensors() { return {getDst(), getSharedTmpBuffer()}; }
264+ SmallVector<Value> getSrcTensors() { return {getSrc()}; }
265+ }];
260}266}
261 267 
262def AscendC_ReduceMaxL1Op : VectorOp<"reduce_max_l1", "ReduceMax", [OpWithDstInterface]> {268def AscendC_ReduceMaxL1Op : VectorOp<"reduce_max_l1", "ReduceMax", [OpWithDstInterface]> {
263 let description = "Reduce maximum elements (L1 API). ";269 let description = "Reduce maximum elements (L1 API). ";
264- let arguments = (ins 270+ let arguments = (ins
265 AscendC_LocalTensor:$dst,271 AscendC_LocalTensor:$dst,
266 AscendC_LocalTensor:$src,272 AscendC_LocalTensor:$src,
267 AscendC_LocalTensor:$sharedTmpBuffer,273 AscendC_LocalTensor:$sharedTmpBuffer,
@@ -272,9 +278,11 @@ def AscendC_ReduceMaxL1Op : VectorOp<"reduce_max_l1", "ReduceMax", [OpWithDstInt
272 );278 );
273}279}
274 280 
275-def AscendC_ReduceMaxL2Op : VectorOp<"reduce_max_l2", "ReduceMax", [OpWithDstInterface, AscFunc]> {281+def AscendC_ReduceMaxL2Op : VectorOp<"reduce_max_l2", "ReduceMax", [
282+ AscFunc, OpWithDstInterface, OpWithSrcInterface
283+]> {
276 let description = "Reduce maximum elements (L2 API). ";284 let description = "Reduce maximum elements (L2 API). ";
277- let arguments = (ins 285+ let arguments = (ins
278 AscendC_LocalTensor:$dst,286 AscendC_LocalTensor:$dst,
279 AscendC_LocalTensor:$src,287 AscendC_LocalTensor:$src,
280 AscendC_LocalTensor:$sharedTmpBuffer,288 AscendC_LocalTensor:$sharedTmpBuffer,
@@ -282,15 +290,21 @@ def AscendC_ReduceMaxL2Op : VectorOp<"reduce_max_l2", "ReduceMax", [OpWithDstInt
282 AnyType:$calIndex290 AnyType:$calIndex
283 );291 );
284 let paramTypeLists = [0, 0, 0, 0, 0];292 let paramTypeLists = [0, 0, 0, 0, 0];
293+ let extraClassDeclaration = extraClassDeclarationBase # [{
294+ SmallVector<Value> getDstTensors() { return {getDst(), getSharedTmpBuffer()}; }
295+ SmallVector<Value> getSrcTensors() { return {getSrc()}; }
296+ }];
285}297}
286 298 
287//===----------------------------------------------------------------------===//299//===----------------------------------------------------------------------===//
288// ReduceMin operations300// ReduceMin operations
289//===----------------------------------------------------------------------===//301//===----------------------------------------------------------------------===//
290 302 
291-def AscendC_ReduceMinL0Op : VectorOp<"reduce_min_l0", "ReduceMin", [OpWithDstInterface, AscFunc]> {303+def AscendC_ReduceMinL0Op : VectorOp<"reduce_min_l0", "ReduceMin", [
304+ AscFunc, OpWithDstInterface, OpWithSrcInterface
305+]> {
292 let description = "Reduce minimum elements (L0 API). ";306 let description = "Reduce minimum elements (L0 API). ";
293- let arguments = (ins 307+ let arguments = (ins
294 AscendC_LocalTensor:$dst,308 AscendC_LocalTensor:$dst,
295 AscendC_LocalTensor:$src,309 AscendC_LocalTensor:$src,
296 AscendC_LocalTensor:$sharedTmpBuffer,310 AscendC_LocalTensor:$sharedTmpBuffer,
@@ -300,11 +314,15 @@ def AscendC_ReduceMinL0Op : VectorOp<"reduce_min_l0", "ReduceMin", [OpWithDstInt
300 AnyType:$calIndex314 AnyType:$calIndex
301 );315 );
302 let paramTypeLists = [0, 0, 0, 0, 0, 0, 0];316 let paramTypeLists = [0, 0, 0, 0, 0, 0, 0];
317+ let extraClassDeclaration = extraClassDeclarationBase # [{
318+ SmallVector<Value> getDstTensors() { return {getDst(), getSharedTmpBuffer()}; }
319+ SmallVector<Value> getSrcTensors() { return {getSrc()}; }
320+ }];
303}321}
304 322 
305def AscendC_ReduceMinL1Op : VectorOp<"reduce_min_l1", "ReduceMin", [OpWithDstInterface]> {323def AscendC_ReduceMinL1Op : VectorOp<"reduce_min_l1", "ReduceMin", [OpWithDstInterface]> {
306 let description = "Reduce minimum elements (L1 API). ";324 let description = "Reduce minimum elements (L1 API). ";
307- let arguments = (ins 325+ let arguments = (ins
308 AscendC_LocalTensor:$dst,326 AscendC_LocalTensor:$dst,
309 AscendC_LocalTensor:$src,327 AscendC_LocalTensor:$src,
310 AscendC_LocalTensor:$sharedTmpBuffer,328 AscendC_LocalTensor:$sharedTmpBuffer,
@@ -315,9 +333,11 @@ def AscendC_ReduceMinL1Op : VectorOp<"reduce_min_l1", "ReduceMin", [OpWithDstInt
315 );333 );
316}334}
317 335 
318-def AscendC_ReduceMinL2Op : VectorOp<"reduce_min_l2", "ReduceMin", [OpWithDstInterface, AscFunc]> {336+def AscendC_ReduceMinL2Op : VectorOp<"reduce_min_l2", "ReduceMin", [
337+ AscFunc, OpWithDstInterface, OpWithSrcInterface
338+]> {
319 let description = "Reduce minimum elements (L2 API). ";339 let description = "Reduce minimum elements (L2 API). ";
320- let arguments = (ins 340+ let arguments = (ins
321 AscendC_LocalTensor:$dst,341 AscendC_LocalTensor:$dst,
322 AscendC_LocalTensor:$src,342 AscendC_LocalTensor:$src,
323 AscendC_LocalTensor:$sharedTmpBuffer,343 AscendC_LocalTensor:$sharedTmpBuffer,
@@ -325,45 +345,62 @@ def AscendC_ReduceMinL2Op : VectorOp<"reduce_min_l2", "ReduceMin", [OpWithDstInt
325 AnyType:$calIndex345 AnyType:$calIndex
326 );346 );
327 let paramTypeLists = [0, 0, 0, 0, 0];347 let paramTypeLists = [0, 0, 0, 0, 0];
348+ let extraClassDeclaration = extraClassDeclarationBase # [{
349+ SmallVector<Value> getDstTensors() { return {getDst(), getSharedTmpBuffer()}; }
350+ SmallVector<Value> getSrcTensors() { return {getSrc()}; }
351+ }];
328}352}
329 353 
330//===----------------------------------------------------------------------===//354//===----------------------------------------------------------------------===//
331// ReduceSum operations355// ReduceSum operations
332//===----------------------------------------------------------------------===//356//===----------------------------------------------------------------------===//
333 357 
334-def AscendC_ReduceSumL0Op : VectorOp<"reduce_sum_l0", "ReduceSum", [OpWithDstInterface, AscFunc]> {358+def AscendC_ReduceSumL0Op : VectorOp<"reduce_sum_l0", "ReduceSum", [
359+ AscFunc, OpWithDstInterface, OpWithSrcInterface
360+]> {
335 let description = "Reduce sum elements (L0 API, int32_t mask). ";361 let description = "Reduce sum elements (L0 API, int32_t mask). ";
336- let arguments = (ins 362+ let arguments = (ins
337 AscendC_LocalTensor:$dst,363 AscendC_LocalTensor:$dst,
338 AscendC_LocalTensor:$src,364 AscendC_LocalTensor:$src,
339 AscendC_LocalTensor:$sharedTmpBuffer,365 AscendC_LocalTensor:$sharedTmpBuffer,
340- AnyType:$mask, 366+ AnyType:$mask,
341- AnyType:$repeatTime, 367+ AnyType:$repeatTime,
342- AnyType:$srcRepStride 368+ AnyType:$srcRepStride
343 );369 );
344 let paramTypeLists = [0, 0, 0, 0, 0, 0];370 let paramTypeLists = [0, 0, 0, 0, 0, 0];
371+ let extraClassDeclaration = extraClassDeclarationBase # [{
372+ SmallVector<Value> getDstTensors() { return {getDst(), getSharedTmpBuffer()}; }
373+ SmallVector<Value> getSrcTensors() { return {getSrc()}; }
374+ }];
345}375}
346 376 
347def AscendC_ReduceSumL1Op : VectorOp<"reduce_sum_l1", "ReduceSum", [OpWithDstInterface]> {377def AscendC_ReduceSumL1Op : VectorOp<"reduce_sum_l1", "ReduceSum", [OpWithDstInterface]> {
348 let description = "Reduce sum elements (L1 API, uint64_t[] mask). ";378 let description = "Reduce sum elements (L1 API, uint64_t[] mask). ";
349- let arguments = (ins 379+ let arguments = (ins
350 AscendC_LocalTensor:$dst,380 AscendC_LocalTensor:$dst,
351 AscendC_LocalTensor:$src,381 AscendC_LocalTensor:$src,
352 AscendC_LocalTensor:$sharedTmpBuffer,382 AscendC_LocalTensor:$sharedTmpBuffer,
353- Variadic<AnyType>:$mask, 383+ Variadic<AnyType>:$mask,
354- AnyType:$repeatTime, 384+ AnyType:$repeatTime,
355- AnyType:$srcRepStride 385+ AnyType:$srcRepStride
356 );386 );
357}387}
358 388 
359-def AscendC_ReduceSumL2Op : VectorOp<"reduce_sum_l2", "ReduceSum", [OpWithDstInterface, AscFunc]> {389+def AscendC_ReduceSumL2Op : VectorOp<"reduce_sum_l2", "ReduceSum", [
390+ AscFunc, OpWithDstInterface, OpWithSrcInterface
391+]> {
360 let description = "Reduce sum elements (L2 API, int32_t count). ";392 let description = "Reduce sum elements (L2 API, int32_t count). ";
361- let arguments = (ins 393+ let arguments = (ins
362 AscendC_LocalTensor:$dst,394 AscendC_LocalTensor:$dst,
363 AscendC_LocalTensor:$src,395 AscendC_LocalTensor:$src,
364 AscendC_LocalTensor:$sharedTmpBuffer,396 AscendC_LocalTensor:$sharedTmpBuffer,
365 AnyType:$count397 AnyType:$count
366 );398 );
367 let paramTypeLists = [0, 0, 0, 0];399 let paramTypeLists = [0, 0, 0, 0];
400+ let extraClassDeclaration = extraClassDeclarationBase # [{
401+ SmallVector<Value> getDstTensors() { return {getDst(), getSharedTmpBuffer()}; }
402+ SmallVector<Value> getSrcTensors() { return {getSrc()}; }
403+ }];
368}404}
369-#endif // ASC_BASIC_OP_VEC_REDUCE_TD405+ 
406+#endif // ASC_BASIC_OP_VEC_REDUCE_TD
@@ -20,12 +20,17 @@ include "mlir/Interfaces/CastInterfaces.td"
20include "mlir/Interfaces/SideEffectInterfaces.td"20include "mlir/Interfaces/SideEffectInterfaces.td"
21include "mlir/IR/OpBase.td"21include "mlir/IR/OpBase.td"
22 22 
23-def AscendC_CastL0Op : VectorOp<"cast_l0", "Cast", [OpWithDstInterface]> {23+def AscendC_CastL0Op : VectorOp<"cast_l0", "Cast", [
24+ OpWithDstInterface, OpWithSrcInterface
25+]> {
24 let description = "Convert the input tensor to the specified data type (L0 API)";26 let description = "Convert the input tensor to the specified data type (L0 API)";
25 let arguments = (ins AscendC_LocalTensor:$dst, AscendC_LocalTensor:$src,27 let arguments = (ins AscendC_LocalTensor:$dst, AscendC_LocalTensor:$src,
26 AscendC_RoundModeAttr:$roundMode, AnyType:$mask,28 AscendC_RoundModeAttr:$roundMode, AnyType:$mask,
27 AnyType:$repeatTimes, AscendC_UnaryRepeatParams:$repeatParams,29 AnyType:$repeatTimes, AscendC_UnaryRepeatParams:$repeatParams,
28 UnitAttr:$isSetMask);30 UnitAttr:$isSetMask);
31+ let extraClassDeclaration = extraClassDeclarationBase # [{
32+ SmallVector<Value> getSrcTensors() { return {getSrc()}; }
33+ }];
29}34}
30 35 
31def AscendC_CastL1Op : VectorOp<"cast_l1", "Cast", [OpWithDstInterface]> {36def AscendC_CastL1Op : VectorOp<"cast_l1", "Cast", [OpWithDstInterface]> {
@@ -36,10 +41,15 @@ def AscendC_CastL1Op : VectorOp<"cast_l1", "Cast", [OpWithDstInterface]> {
36 UnitAttr:$isSetMask);41 UnitAttr:$isSetMask);
37}42}
38 43 
39-def AscendC_CastL2Op : VectorOp<"cast_l2", "Cast", [OpWithDstInterface]> {44+def AscendC_CastL2Op : VectorOp<"cast_l2", "Cast", [
45+ OpWithDstInterface, OpWithSrcInterface
46+]> {
40 let description = "Convert the input tensor to the specified data type (L2 API)";47 let description = "Convert the input tensor to the specified data type (L2 API)";
41 let arguments = (ins AscendC_LocalTensor:$dst, AscendC_LocalTensor:$src,48 let arguments = (ins AscendC_LocalTensor:$dst, AscendC_LocalTensor:$src,
42 AscendC_RoundModeAttr:$roundMode, AnyType:$calCount);49 AscendC_RoundModeAttr:$roundMode, AnyType:$calCount);
50+ let extraClassDeclaration = extraClassDeclarationBase # [{
51+ SmallVector<Value> getSrcTensors() { return {getSrc()}; }
52+ }];
43}53}
44 54 
45def AscendC_CastDeqL0Op : VectorOp<"cast_deq_l0", "CastDeq"> {55def AscendC_CastDeqL0Op : VectorOp<"cast_deq_l0", "CastDeq"> {
@@ -133,17 +133,6 @@ def AscendC_CacheLineAttr : I32EnumAttr<"CacheLine", "", [
133 let underlyingType = "uint8_t";133 let underlyingType = "uint8_t";
134}134}
135 135 
136-def AscendC_CopyDirectionAttr : I32EnumAttr<"CopyDirection", "", [
137- I32EnumAttrCase<"Unknown", 0>,
138- I32EnumAttrCase<"gm_ubuf", 1>,
139- I32EnumAttrCase<"ubuf_gm", 2>,
140- I32EnumAttrCase<"ubuf_ubuf", 3>,
141- I32EnumAttrCase<"gm_gm", 4>,
142-]> {
143- let cppNamespace = "::mlir::ascendc";
144- let description = "Copy direction";
145-}
146- 
147def AscendC_CubeFormatAttr : I32EnumAttr<"CubeFormat", "matmul tensor format", [136def AscendC_CubeFormatAttr : I32EnumAttr<"CubeFormat", "matmul tensor format", [
148 I32EnumAttrCase<"ND", 0, "nd">,137 I32EnumAttrCase<"ND", 0, "nd">,
149 I32EnumAttrCase<"NZ", 1, "nz">,138 I32EnumAttrCase<"NZ", 1, "nz">,
@@ -343,6 +332,27 @@ def AscendC_MatmulConfigAttr : AscendC_Attr<"MatmulConfig", "matmul_config">{
343 let assemblyFormat = "`<` struct(params) `>`";332 let assemblyFormat = "`<` struct(params) `>`";
344}333}
345 334 
335+def AscendC_MaskPattern : I32EnumAttr<"MaskPattern", "", [
336+ I32EnumAttrCase<"ALL", 0>,
337+ I32EnumAttrCase<"VL1", 1>,
338+ I32EnumAttrCase<"VL2", 2>,
339+ I32EnumAttrCase<"VL3", 3>,
340+ I32EnumAttrCase<"VL4", 4>,
341+ I32EnumAttrCase<"VL8", 5>,
342+ I32EnumAttrCase<"VL16", 6>,
343+ I32EnumAttrCase<"VL32", 7>,
344+ I32EnumAttrCase<"VL64", 8>,
345+ I32EnumAttrCase<"VL128", 9>,
346+ I32EnumAttrCase<"M3", 10>,
347+ I32EnumAttrCase<"M4", 11>,
348+ I32EnumAttrCase<"H", 12>,
349+ I32EnumAttrCase<"Q", 13>,
350+ I32EnumAttrCase<"ALLF", 15>,
351+]>{
352+ let cppNamespace = "::mlir::ascendc";
353+ let description = "Represents AscendC::Reg::MaskPattern";
354+ let underlyingType = "uint8_t";
355+}
346 356 
347def AscendC_MatmulPolicyAttr : I32EnumAttr<"MatmulPolicy", "", [357def AscendC_MatmulPolicyAttr : I32EnumAttr<"MatmulPolicy", "", [
348 I32EnumAttrCase<"MatmulPolicy", 0>,358 I32EnumAttrCase<"MatmulPolicy", 0>,
@@ -355,6 +365,19 @@ def AscendC_MatmulPolicyAttr : I32EnumAttr<"MatmulPolicy", "", [
355 let underlyingType = "uint8_t";365 let underlyingType = "uint8_t";
356}366}
357 367 
368+def AscendC_MemType : I32EnumAttr<"MemType", "", [
369+ I32EnumAttrCase<"VEC_STORE", 0>,
370+ I32EnumAttrCase<"VEC_LOAD", 1>,
371+ I32EnumAttrCase<"SCALAR_STORE", 2>,
372+ I32EnumAttrCase<"SCALAR_LOAD", 3>,
373+ I32EnumAttrCase<"VEC_ALL", 4>,
374+ I32EnumAttrCase<"SCALAR_ALL", 5>,
375+]>{
376+ let cppNamespace = "::mlir::ascendc";
377+ let description = "Represents AscendC::Reg::MemType";
378+ let underlyingType = "uint8_t";
379+}
380+ 
358def AscendC_PipeAttr : I32EnumAttr<"Pipe", "", [381def AscendC_PipeAttr : I32EnumAttr<"Pipe", "", [
359 I32EnumAttrCase<"PIPE_S", 0, "pipe_s">,382 I32EnumAttrCase<"PIPE_S", 0, "pipe_s">,
360 I32EnumAttrCase<"PIPE_V", 1, "pipe_v">,383 I32EnumAttrCase<"PIPE_V", 1, "pipe_v">,
@@ -539,4 +562,41 @@ def AscendC_AtomicOpAttr : I32EnumAttr<"AtomicOp", "", [
539 let underlyingType = "uint8_t";562 let underlyingType = "uint8_t";
540}563}
541 564 
565+def AscendC_ReducePatternAttr : I32EnumAttr<"ReducePattern", "", [
566+ I32EnumAttrCase<"R", 0>,
567+ I32EnumAttrCase<"AR", 1>,
568+ I32EnumAttrCase<"RA", 2>,
569+ I32EnumAttrCase<"ARA", 3>,
570+ I32EnumAttrCase<"ARAR", 4>,
571+ I32EnumAttrCase<"ARARA", 5>,
572+ I32EnumAttrCase<"ARARAR", 6>,
573+ I32EnumAttrCase<"ARARARA", 7>,
574+ I32EnumAttrCase<"ARARARAR", 8>,
575+ I32EnumAttrCase<"ARARARARA", 9>,
576+ I32EnumAttrCase<"RAR", 10>,
577+ I32EnumAttrCase<"RARA", 11>,
578+ I32EnumAttrCase<"RARAR", 12>,
579+ I32EnumAttrCase<"RARARA", 13>,
580+ I32EnumAttrCase<"RARARAR", 14>,
581+ I32EnumAttrCase<"RARARARA", 15>,
582+]>{
583+ let cppNamespace = "::mlir::ascendc";
584+ let description = "Describes reduction pattern from reduce_common.h";
585+ let underlyingType = "uint8_t";
586+}
587+ 
588+def AscendC_DataCopyMVTypeAttr: I32EnumAttr<"DataCopyMVType", "", [
589+ I32EnumAttrCase<"UB_TO_OUT", 0>,
590+ I32EnumAttrCase<"OUT_TO_UB", 1>,
591+]>{
592+ let cppNamespace = "::mlir::ascendc";
593+ let description = "Used in SetLoopModePara/ResetLoopModePara";
594+ let underlyingType = "uint8_t";
595+}
596+ 
597+def AscendC_CopyDirectionAttr : AscendC_Attr<"CopyDirection", "copy_direction"> {
598+ let parameters = (ins "TPosition":$src, "TPosition":$dst);
599+ let assemblyFormat = "`<` $src `,` $dst `>`";
600+}
601+ 
542#endif // ASC_CORE_ATTRIBUTES_TD602#endif // ASC_CORE_ATTRIBUTES_TD
@@ -69,7 +69,7 @@ def AscendC_GlobalTensorGetPhyAddrOp : APIOp<"global_tensor.get_phy_addr", "GetP
69 let arguments = (ins AscendC_GlobalTensor:$tensor, Optional<UI64>:$offset);69 let arguments = (ins AscendC_GlobalTensor:$tensor, Optional<UI64>:$offset);
70 let results = (outs AnyRankedOrUnrankedMemRef:$result);70 let results = (outs AnyRankedOrUnrankedMemRef:$result);
71 let assemblyFormat = [{71 let assemblyFormat = [{
72- $tensor attr-dict (`,` $offset^)? `:` qualified(type($tensor)) 72+ $tensor attr-dict (`,` $offset^)? `:` qualified(type($tensor))
73 `,` qualified(type($result)) (`,` type($offset)^)?73 `,` qualified(type($result)) (`,` type($offset)^)?
74 }];74 }];
75}75}
@@ -109,7 +109,7 @@ def AscendC_GlobalTensorSetL2CacheHintOp : APIOp<"global_tensor.set_l2_cache_hin
109def AscendC_GlobalTensorSetShapeInfoOp : APIOp<"global_tensor.set_shape_info", "SetShapeInfo", [AscMemberFunc]> {109def AscendC_GlobalTensorSetShapeInfoOp : APIOp<"global_tensor.set_shape_info", "SetShapeInfo", [AscMemberFunc]> {
110 let summary = "Call `AscendC::GlobalTensor::SetShapeInfo` method";110 let summary = "Call `AscendC::GlobalTensor::SetShapeInfo` method";
111 let arguments = (ins AscendC_GlobalTensor:$tensor, AscendC_ShapeInfo:$shapeInfo);111 let arguments = (ins AscendC_GlobalTensor:$tensor, AscendC_ShapeInfo:$shapeInfo);
112-} 112+}
113 113 
114def AscendC_GlobalTensorSetShapeInfoV2Op : APIOp<"global_tensor.set_shape_info_v2", "SetShapeInfo"> {114def AscendC_GlobalTensorSetShapeInfoV2Op : APIOp<"global_tensor.set_shape_info_v2", "SetShapeInfo"> {
115 let summary = "Call `AscendC::GlobalTensor::SetShapeInfo` method";115 let summary = "Call `AscendC::GlobalTensor::SetShapeInfo` method";
@@ -134,6 +134,7 @@ def AscendC_GlobalTensorSubIndexOp
134 $tensor `[` $index `]` attr-dict `:` qualified(type($tensor)) `,`134 $tensor `[` $index `]` attr-dict `:` qualified(type($tensor)) `,`
135 type($index) `,` qualified(type($result))135 type($index) `,` qualified(type($result))
136 }];136 }];
137+ let hasFolder = 1;
137}138}
138 139 
139def AscendC_GlobalTensorSetGlobalBufferOp : APIOp<"global_tensor.set_global_buffer", "SetGlobalBuffer", [AscMemberFunc]> {140def AscendC_GlobalTensorSetGlobalBufferOp : APIOp<"global_tensor.set_global_buffer", "SetGlobalBuffer", [AscMemberFunc]> {
@@ -165,6 +166,13 @@ def AscendC_LocalTensorV2Op : AscendC_Op<"local_tensor_v2"> {
165 let assemblyFormat = "$pos `,` $addr `,` $tileSize attr-dict `:` qualified(type($result))";166 let assemblyFormat = "$pos `,` $addr `,` $tileSize attr-dict `:` qualified(type($result))";
166}167}
167 168 
169+def AscendC_LocalTensorV3Op : AscendC_Op<"local_tensor_v3"> {
170+ let summary = "Instantiate local tensor";
171+ let arguments = (ins AscendC_TPositionAttr:$pos, I32Attr:$addr, I32Attr:$tileSize);
172+ let results = (outs AscendC_LocalTensor:$result);
173+ let assemblyFormat = "$pos `,` $addr `,` $tileSize attr-dict `:` qualified(type($result))";
174+}
175+ 
168def AscendC_LocalTensorBracketOp176def AscendC_LocalTensorBracketOp
169 : AscendC_Op<"local_tensor.bracket", [Pure]> {177 : AscendC_Op<"local_tensor.bracket", [Pure]> {
170 let summary = "Call `AscendC::LocalTensor::operator()` method";178 let summary = "Call `AscendC::LocalTensor::operator()` method";
@@ -190,11 +198,21 @@ def AscendC_LocalTensorGetPhyAddrOp : APIOp<"local_tensor.get_phy_addr", "GetPhy
190 let arguments = (ins AscendC_LocalTensor:$tensor, Optional<UI32>:$offset);198 let arguments = (ins AscendC_LocalTensor:$tensor, Optional<UI32>:$offset);
191 let results = (outs AnyType:$result);199 let results = (outs AnyType:$result);
192 let assemblyFormat = [{200 let assemblyFormat = [{
193- $tensor attr-dict (`,` $offset^)? `:` qualified(type($tensor)) 201+ $tensor attr-dict (`,` $offset^)? `:` qualified(type($tensor))
194 `,` qualified(type($result)) (`,` type($offset)^)?202 `,` qualified(type($result)) (`,` type($offset)^)?
195 }];203 }];
196}204}
197 205 
206+def AscendC_LocalTensorGetPhyAddrV2Op : APIOp<"local_tensor.get_phy_addr_v2", "GetPhyAddr"> {
207+ let summary = "Call `AscendC::LocalTensor::GetPhyAddr` method";
208+ let arguments = (ins AscendC_LocalTensor:$tensor);
209+ let results = (outs AnyRankedOrUnrankedMemRef:$result);
210+ let assemblyFormat = [{
211+ $tensor attr-dict `:` qualified(type($tensor))
212+ `,` qualified(type($result))
213+ }];
214+}
215+ 
198def AscendC_LocalTensorGetPositionOp : APIOp<"local_tensor.get_position", "GetPosition", [AscMemberFunc]> {216def AscendC_LocalTensorGetPositionOp : APIOp<"local_tensor.get_position", "GetPosition", [AscMemberFunc]> {
199 let summary = "Call `AscendC::LocalTensor::GetPosition` method";217 let summary = "Call `AscendC::LocalTensor::GetPosition` method";
200 let arguments = (ins AscendC_LocalTensor:$tensor);218 let arguments = (ins AscendC_LocalTensor:$tensor);
@@ -231,7 +249,7 @@ def AscendC_LocalTensorGetUserTagOp : APIOp<"local_tensor.get_user_tag", "GetUse
231 }];249 }];
232}250}
233 251 
234-def AscendC_LocalTensorGetValueOp : APIOp<"local_tensor.get_value", "GetValue", [AscMemberFunc]> {252+def AscendC_LocalTensorGetValueOp : APIOp<"local_tensor.get_value", "GetValue", [AscMemberFunc, OpWithSrcInterface]> {
235 let summary = "Call `AscendC::LocalTensor::GetValue` method";253 let summary = "Call `AscendC::LocalTensor::GetValue` method";
236 let arguments = (ins AscendC_LocalTensor:$tensor, AnyType:$index);254 let arguments = (ins AscendC_LocalTensor:$tensor, AnyType:$index);
237 let results = (outs AnyType:$value);255 let results = (outs AnyType:$value);
@@ -239,6 +257,9 @@ def AscendC_LocalTensorGetValueOp : APIOp<"local_tensor.get_value", "GetValue",
239 $tensor `,` $index attr-dict `:` qualified(type($tensor)) `,`257 $tensor `,` $index attr-dict `:` qualified(type($tensor)) `,`
240 qualified(type($index)) `,` qualified(type($value))258 qualified(type($index)) `,` qualified(type($value))
241 }];259 }];
260+ let extraClassDeclaration = extraClassDeclarationBase # [{
261+ SmallVector<Value> getSrcTensors() { return {getTensor()}; }
262+ }];
242}263}
243 264 
244def AscendC_LocalTensorPrintOp265def AscendC_LocalTensorPrintOp
@@ -285,13 +306,16 @@ def AscendC_LocalTensorSetUserTagOp : APIOp<"local_tensor.set_user_tag", "SetUse
285 let arguments = (ins AscendC_LocalTensor:$tensor, I32:$tag);306 let arguments = (ins AscendC_LocalTensor:$tensor, I32:$tag);
286}307}
287 308 
288-def AscendC_LocalTensorSetValueOp : APIOp<"local_tensor.set_value", "SetValue", [AscMemberFunc]> {309+def AscendC_LocalTensorSetValueOp : APIOp<"local_tensor.set_value", "SetValue", [AscMemberFunc, OpWithSrcInterface]> {
289 let summary = "Call `AscendC::LocalTensor::SetValue` method";310 let summary = "Call `AscendC::LocalTensor::SetValue` method";
290 let arguments = (ins AscendC_LocalTensor:$tensor, AnyType:$index, AnyType:$value);311 let arguments = (ins AscendC_LocalTensor:$tensor, AnyType:$index, AnyType:$value);
291 let assemblyFormat = [{312 let assemblyFormat = [{
292 $tensor `,` $index `,` $value attr-dict `:` qualified(type($tensor)) `,`313 $tensor `,` $index `,` $value attr-dict `:` qualified(type($tensor)) `,`
293 type($index) `,` type($value)314 type($index) `,` type($value)
294 }];315 }];
316+ let extraClassDeclaration = extraClassDeclarationBase # [{
317+ SmallVector<Value> getSrcTensors() { return {getTensor()}; }
318+ }];
295}319}
296 320 
297def AscendC_LocalTensorSubIndexOp321def AscendC_LocalTensorSubIndexOp
@@ -303,6 +327,7 @@ def AscendC_LocalTensorSubIndexOp
303 $tensor `[` $index `]` attr-dict `:` qualified(type($tensor)) `,`327 $tensor `[` $index `]` attr-dict `:` qualified(type($tensor)) `,`
304 type($index) `,` qualified(type($result))328 type($index) `,` qualified(type($result))
305 }];329 }];
330+ let hasFolder = 1;
306}331}
307 332 
308def AscendC_LocalTensorToFileOp333def AscendC_LocalTensorToFileOp
@@ -49,6 +49,15 @@ def AscendC_FixpipeParams : AscendC_Type<"FixpipeParams", "fixpipe_params"> {
49 let assemblyFormat = "`<` $type `>`";49 let assemblyFormat = "`<` $type `>`";
50}50}
51 51 
52+def AscendC_FixpipeParamsC310 : AscendC_Type<"FixpipeParamsC310", "fixpipe_params_c310"> {
53+ let description = "Represents AscendC::FixpipeParamsC310";
54+ let parameters = (ins "CO2LayoutAttr":$cO2LayoutAttr);
55+ let assemblyFormat = "`<` $cO2LayoutAttr `>`";
56+ let extraClassDeclaration = [{
57+ CO2Layout getCO2Layout() { return getCO2LayoutAttr().getValue(); }
58+ }];
59+}
60+ 
52def AscendC_GlobalTensor : AscendC_BaseTensorType<"GlobalTensor", "global_tensor"> {61def AscendC_GlobalTensor : AscendC_BaseTensorType<"GlobalTensor", "global_tensor"> {
53 let summary = "Global tensor from GM (OUT)";62 let summary = "Global tensor from GM (OUT)";
54}63}
@@ -67,6 +76,10 @@ def AscendC_Mask : AscendC_Type<"Mask", "mask"> {
67 let summary = "Represents vector mask (bit mode)";76 let summary = "Represents vector mask (bit mode)";
68}77}
69 78 
79+def AscendC_MaskReg : AscendC_Type<"MaskReg", "mask_reg"> {
80+ let description = "Represents AscendC::Reg::MaskReg";
81+}
82+ 
70def AscendC_Matmul : AscendC_Type<"Matmul", "matmul"> {83def AscendC_Matmul : AscendC_Type<"Matmul", "matmul"> {
71 let description = "Represents matmul::Matmul";84 let description = "Represents matmul::Matmul";
72 let parameters = (ins "TPositionAttr":$srcAPositionAttr,85 let parameters = (ins "TPositionAttr":$srcAPositionAttr,
@@ -170,6 +183,13 @@ def AscendC_Queue : AscendC_BaseQueueType<"Queue", "queue"> {
170 }];183 }];
171}184}
172 185 
186+def AscendC_RegTensor : AscendC_Type<"RegTensor", "reg_tensor"> {
187+ let summary = "Reg tensor from UB";
188+ let description = "Represents AscendC::Reg::RegTensor";
189+ let parameters = (ins "Type":$elementType);
190+ let assemblyFormat = "`<` $elementType `>`";
191+}
192+ 
173def AscendC_TBufPool : AscendC_Type<"TBufPool", "tbuf_pool"> {193def AscendC_TBufPool : AscendC_Type<"TBufPool", "tbuf_pool"> {
174 let description = "Represents AscendC::TBufPool";194 let description = "Represents AscendC::TBufPool";
175 let parameters = (ins "TPositionAttr":$tPositionAttr, "uint32_t": $bufIDSize);195 let parameters = (ins "TPositionAttr":$tPositionAttr, "uint32_t": $bufIDSize);
@@ -184,6 +204,12 @@ def AscendC_TBufPool : AscendC_Type<"TBufPool", "tbuf_pool"> {
184 }];204 }];
185}205}
186 206 
207+def AscendC_NdDmaParams : AscendC_Type<"NdDmaParams","nd_dma_params"> {
208+ let description = "Represents AscendC::NdDmaParams";
209+ let parameters = (ins "Type":$type, "int32_t":$dims);
210+ let assemblyFormat = "`<` $type `,` $dims `>`";
211+}
212+ 
187include "ascir/API/Types.td.inc"213include "ascir/API/Types.td.inc"
188 214 
189#endif // ASC_CORE_TYPES_TD215#endif // ASC_CORE_TYPES_TD
@@ -0,0 +1,20 @@
1+/*
2+ * Copyright (c) 2026 Huawei Technologies Co., Ltd.
3+ * This program is free software, you can redistribute it and/or modify it under the terms and conditions of
4+ * 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,
7+ * 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+#ifndef ASCENDC_DOC_TD
12+#define ASCENDC_DOC_TD
13+ 
14+include "Dialect.td"
15+ 
16+include "Base.td"
17+include "Interfaces.td"
18+include "Ops.td"
19+ 
20+#endif // ASCENDC_DOC_TD
@@ -18,34 +18,52 @@ def AscendC_CalCount : AnyTypeOf<[I8, I16, I32, I64]>;
18class GetMethod<string fieldName, string methodName, string returnType = "::mlir::Value">18class GetMethod<string fieldName, string methodName, string returnType = "::mlir::Value">
19 : InterfaceMethod<"Obtain `" # fieldName # "` parameter", returnType, methodName>;19 : InterfaceMethod<"Obtain `" # fieldName # "` parameter", returnType, methodName>;
20 20 
21+class GetMutableMethod<string fieldName, string methodName>
22+ : InterfaceMethod<"Obtain the mutable `" # fieldName # "` operand", "::mlir::OpOperand&", methodName>;
23+ 
21class AscendC_InterfaceMethods {24class AscendC_InterfaceMethods {
22 // Method list should be kept sorted by name.25 // Method list should be kept sorted by name.
23 InterfaceMethod getAPIName = InterfaceMethod<"Obtain Ascend C library name", "::llvm::StringRef", "getAPIName">;26 InterfaceMethod getAPIName = InterfaceMethod<"Obtain Ascend C library name", "::llvm::StringRef", "getAPIName">;
24 InterfaceMethod getCalCount = GetMethod<"calCount", "getCalCount">;27 InterfaceMethod getCalCount = GetMethod<"calCount", "getCalCount">;
28+ InterfaceMethod getCalCountMutable = GetMutableMethod<"calCount", "getCalCountMutable">;
25 InterfaceMethod getComment = InterfaceMethod<"Obtain describing comment", "::llvm::StringRef", "getComment">;29 InterfaceMethod getComment = InterfaceMethod<"Obtain describing comment", "::llvm::StringRef", "getComment">;
26 InterfaceMethod getDst = GetMethod<"dst", "getDst">;30 InterfaceMethod getDst = GetMethod<"dst", "getDst">;
27 InterfaceMethod getDstBlkStride = GetMethod<"dstBlkStride", "getDstBlkStride">;31 InterfaceMethod getDstBlkStride = GetMethod<"dstBlkStride", "getDstBlkStride">;
32+ InterfaceMethod getDstReg = GetMethod<"dstReg", "getDstReg">;
33+ InterfaceMethod getDstBlkStrideMutable = GetMutableMethod<"dstBlkStride", "getDstBlkStrideMutable">;
28 InterfaceMethod getDstRepStride = GetMethod<"dstRepStride", "getDstRepStride">;34 InterfaceMethod getDstRepStride = GetMethod<"dstRepStride", "getDstRepStride">;
35+ InterfaceMethod getDstRepStrideMutable = GetMutableMethod<"dstRepStride", "getDstRepStrideMutable">;
29 InterfaceMethod getMask = GetMethod<"mask", "getMask">;36 InterfaceMethod getMask = GetMethod<"mask", "getMask">;
37+ InterfaceMethod getMaskMutable = GetMutableMethod<"mask", "getMaskMutable">;
38+ InterfaceMethod getMaskReg = GetMethod<"maskReg", "getMaskReg">;
30 InterfaceMethod getParams = GetMethod<"params", "getParams">;39 InterfaceMethod getParams = GetMethod<"params", "getParams">;
40+ InterfaceMethod getPattern = GetMethod<"pattern", "getPattern", "::mlir::ascendc::ReducePattern">;
31 InterfaceMethod getRepeatParams = GetMethod<"repeatParams", "getRepeatParams">;41 InterfaceMethod getRepeatParams = GetMethod<"repeatParams", "getRepeatParams">;
42+ InterfaceMethod getRepeatParamsMutable = GetMutableMethod<"repeatParams", "getRepeatParamsMutable">;
32 InterfaceMethod getRepeatTimes = GetMethod<"repeatTimes", "getRepeatTimes">;43 InterfaceMethod getRepeatTimes = GetMethod<"repeatTimes", "getRepeatTimes">;
44+ InterfaceMethod getRepeatTimesMutable = GetMutableMethod<"repeatTimes", "getRepeatTimesMutable">;
33 InterfaceMethod getIsSetMask = GetMethod<"isSetMask", "getIsSetMask", "bool">;45 InterfaceMethod getIsSetMask = GetMethod<"isSetMask", "getIsSetMask", "bool">;
46+ InterfaceMethod setIsSetMask = InterfaceMethod<"Sets isSetMask parameter", "void", "setIsSetMask", (ins "bool":$attrValue)>;
34 InterfaceMethod getScalar = GetMethod<"scalar", "getScalar">;47 InterfaceMethod getScalar = GetMethod<"scalar", "getScalar">;
48+ InterfaceMethod getScalarMutable = GetMutableMethod<"scalar", "getScalarMutable">;
35 InterfaceMethod getSharedTmpBuffer = GetMethod<"sharedTmpBuffer", "getSharedTmpBuffer">;49 InterfaceMethod getSharedTmpBuffer = GetMethod<"sharedTmpBuffer", "getSharedTmpBuffer">;
36 InterfaceMethod getSrc = GetMethod<"src", "getSrc">;50 InterfaceMethod getSrc = GetMethod<"src", "getSrc">;
37 InterfaceMethod getSrcBlkStride = GetMethod<"srcBlkStride", "getSrcBlkStride">;51 InterfaceMethod getSrcBlkStride = GetMethod<"srcBlkStride", "getSrcBlkStride">;
52+ InterfaceMethod getSrcReg = GetMethod<"srcReg", "getSrcReg">;
38 InterfaceMethod getSrcRepStride = GetMethod<"srcRepStride", "getSrcRepStride">;53 InterfaceMethod getSrcRepStride = GetMethod<"srcRepStride", "getSrcRepStride">;
39 InterfaceMethod getSrc0 = GetMethod<"src0", "getSrc0">;54 InterfaceMethod getSrc0 = GetMethod<"src0", "getSrc0">;
40 InterfaceMethod getSrc0BlkStride = GetMethod<"src0BlkStride", "getSrc0BlkStride">;55 InterfaceMethod getSrc0BlkStride = GetMethod<"src0BlkStride", "getSrc0BlkStride">;
56+ InterfaceMethod getSrc0Reg = GetMethod<"src0Reg", "getSrc0Reg">;
41 InterfaceMethod getSrc0RepStride = GetMethod<"src0RepStride", "getSrc0RepStride">;57 InterfaceMethod getSrc0RepStride = GetMethod<"src0RepStride", "getSrc0RepStride">;
42 InterfaceMethod getSrc1 = GetMethod<"src1", "getSrc1">;58 InterfaceMethod getSrc1 = GetMethod<"src1", "getSrc1">;
43 InterfaceMethod getSrc1BlkStride = GetMethod<"src1BlkStride", "getSrc1BlkStride">;59 InterfaceMethod getSrc1BlkStride = GetMethod<"src1BlkStride", "getSrc1BlkStride">;
60+ InterfaceMethod getSrc1Reg = GetMethod<"src1Reg", "getSrc1Reg">;
44 InterfaceMethod getSrc1RepStride = GetMethod<"src1RepStride", "getSrc1RepStride">;61 InterfaceMethod getSrc1RepStride = GetMethod<"src1RepStride", "getSrc1RepStride">;
62+ InterfaceMethod getSrcShape = GetMethod<"srcShape", "getSrcShape", "mlir::OperandRange">;
45 InterfaceMethod getIsReuseSource = GetMethod<"isReuseSource", "getIsReuseSource">;63 InterfaceMethod getIsReuseSource = GetMethod<"isReuseSource", "getIsReuseSource">;
46- InterfaceMethod getIsExhaustedSuspension = GetMethod<"isExhaustedSuspension", 64+ InterfaceMethod getIsExhaustedSuspension = GetMethod<"isExhaustedSuspension",
47 "getIsExhaustedSuspension", "bool">;65 "getIsExhaustedSuspension", "bool">;
48- InterfaceMethod getIsFullSort = GetMethod<"isFullSort", 66+ InterfaceMethod getIsFullSort = GetMethod<"isFullSort",
49 "getIsFullSort", "bool">;67 "getIsFullSort", "bool">;
50 InterfaceMethod getVdeq = GetMethod<"vdeq", "getVdeq">;68 InterfaceMethod getVdeq = GetMethod<"vdeq", "getVdeq">;
51 InterfaceMethod getVdeqInfo = GetMethod<"vdeqInfo", "getVdeqInfo">;69 InterfaceMethod getVdeqInfo = GetMethod<"vdeqInfo", "getVdeqInfo">;
@@ -67,71 +85,140 @@ def APIOpInterface : AscendC_OpInterface<"APIOp"> {
67 85 
68def OpWithDstInterface : AscendC_OpInterface<"OpWithDst"> {86def OpWithDstInterface : AscendC_OpInterface<"OpWithDst"> {
69 let description = "Operation with `dst` operand (typically tensor)";87 let description = "Operation with `dst` operand (typically tensor)";
70- let methods = [getDst];88+ let methods = [
89+ getDst,
90+ InterfaceMethod<"Get dst tensors", "::llvm::SmallVector<::mlir::Value>",
91+ "getDstTensors", /*args*/ (ins), /*methodBody*/ "",
92+ /*defaultImplementation*/ "return {$_op.getDst()};">,
93+ ];
94+}
95+ 
96+def OpWithSrcInterface : AscendC_OpInterface<"OpWithSrc"> {
97+ let description = "Operation with source tensor operands";
98+ let methods = [InterfaceMethod<"Get source tensors", "::llvm::SmallVector<::mlir::Value>", "getSrcTensors">];
99+}
100+ 
101+def OpWithReusableSrcInterface : AscendC_OpInterface<"OpWithReusableSrc"> {
102+ let description = [{
103+ Marker interface for operations where hardware reads all source operands
104+ completely before writing the destination. This property allows the source
105+ and destination to safely share the same buffer (same-op reuse in the
106+ ReuseTensorAllocation pass). Operations like BroadcastOp that read sources
107+ multiple times during execution must NOT have this interface.
108+ }];
71}109}
72 110 
73def VectorOpInterface : AscendC_OpInterface<"VectorOp", [APIOpInterface]> {111def VectorOpInterface : AscendC_OpInterface<"VectorOp", [APIOpInterface]> {
74 let description = "Base interface for vector operations";112 let description = "Base interface for vector operations";
75}113}
76 114 
77-def DataCopyOpInterface :115+class SimpleDirMethod<string src, string dst> : InterfaceMethod<
78- AscendC_OpInterface<"DataCopyOp", [APIOpInterface, OpWithDstInterface]> {116+ "Check if " # src # "-to-" # dst # " copy", "bool", "is" # src # "To" # dst,
79- InterfaceMethod getDirection = InterfaceMethod<"Get copy direction",117+ (ins), "", "return isa<" # src # "TensorType>($_op.getSrc().getType()) && "
80- "::mlir::ascendc::CopyDirection", "getDirection", (ins), "", [{118+ "isa<" # dst # "TensorType>($_op.getDst().getType());">;
81- auto dstType = $_op.getDst().getType();
82- auto srcType = $_op.getSrc().getType();
83- if (isa<GlobalTensorType>(srcType)) {
84- if (isa<GlobalTensorType>(dstType))
85- return CopyDirection::gm_gm;
86- if (isa<LocalTensorType>(dstType))
87- return CopyDirection::gm_ubuf;
88- } else if (isa<LocalTensorType>(srcType)) {
89- if (isa<GlobalTensorType>(dstType))
90- return CopyDirection::ubuf_gm;
91- if (isa<LocalTensorType>(dstType))
92- return CopyDirection::ubuf_ubuf;
93- }
94- return CopyDirection::Unknown;
95- }]>;
96- InterfaceMethod setSrc = InterfaceMethod<"Set src operand", "void", "setSrc",
97- (ins "Value":$src), "", [{
98- $_op.getSrcMutable().assign(src);
99- }]>;
100 119 
120+def DataCopyOpInterface
121+ : AscendC_OpInterface<"DataCopyOp", [APIOpInterface, OpWithDstInterface]> {
101 let description = "Data copy operation";122 let description = "Data copy operation";
102- let methods = [getSrc, getDirection, setSrc];123+ let methods = [
124+ getSrc,
125+ StaticInterfaceMethod<"Get attr name", "::llvm::StringLiteral",
126+ "getDirectionAttrName", (ins), "", [{ return "direction"; }]>,
127+ InterfaceMethod<"Get copy direction", "::mlir::ascendc::CopyDirectionAttr",
128+ "getDirectionAttr", (ins), "", [{
129+ auto attr = $_op->getAttr(getDirectionAttrName());
130+ return dyn_cast_if_present<CopyDirectionAttr>(attr);
131+ }]>,
132+ InterfaceMethod<"Get copy direction",
133+ "std::optional<std::pair<::mlir::ascendc::TPosition, ::mlir::ascendc::TPosition>>",
134+ "getDirection", (ins), "", [{
135+ if (auto attr = $_op.getDirectionAttr())
136+ return std::pair(attr.getSrc(), attr.getDst());
137+ return std::nullopt;
138+ }]>,
139+ InterfaceMethod<"Set copy direction", "void", "setDirection",
140+ (ins "TPosition":$src, "TPosition":$dst), "", [{
141+ auto attr = CopyDirectionAttr::get($_op.getContext(), src, dst);
142+ $_op->setAttr(getDirectionAttrName(), attr);
143+ }]>,
144+ SimpleDirMethod<"Global", "Local">,
145+ SimpleDirMethod<"Local", "Global">,
146+ SimpleDirMethod<"Local", "Local">,
147+ InterfaceMethod<"Set src operand", "void", "setSrc", (ins "Value":$src), "",
148+ "$_op.getSrcMutable().assign(src);">,
149+ ];
150+}
151+ 
152+def CopyToL0OpInterface : AscendC_OpInterface<"CopyToL0Op", [
153+ APIOpInterface, DataCopyOpInterface, OpWithSrcInterface
154+]> {
155+ let description = "Copy from GM or L1 (TPosition A1/B1) to L0 (A2/B2)";
156+}
157+ 
158+def RegOpInterface : AscendC_OpInterface<"RegOp", [APIOpInterface]> {
159+ let description = "Base interface for operations representing register API";
103}160}
104 161 
105//===----------------------------------------------------------------------===//162//===----------------------------------------------------------------------===//
106// Vector unary operations163// Vector unary operations
107//===----------------------------------------------------------------------===//164//===----------------------------------------------------------------------===//
108 165 
109-def UnaryOpInterface166+def UnaryOpInterface : AscendC_OpInterface<"UnaryOp", [
110- : AscendC_OpInterface<"UnaryOp", [VectorOpInterface, OpWithDstInterface]> {167+ VectorOpInterface, OpWithDstInterface, OpWithSrcInterface, OpWithReusableSrcInterface
168+]> {
111 let description = "Vector unary operation";169 let description = "Vector unary operation";
112 let methods = [getSrc];170 let methods = [getSrc];
113}171}
114 172 
115def UnaryL0OpInterface : AscendC_OpInterface<"UnaryL0Op", [UnaryOpInterface]> {173def UnaryL0OpInterface : AscendC_OpInterface<"UnaryL0Op", [UnaryOpInterface]> {
116 let description = "Vector unary operation (L0 API)";174 let description = "Vector unary operation (L0 API)";
117- let methods = [getMask, getRepeatTimes, getRepeatParams];175+ let methods = [
176+ getMask,
177+ getRepeatTimes,
178+ getRepeatParams,
179+ getMaskMutable,
180+ getRepeatTimesMutable,
181+ getRepeatParamsMutable,
182+ getIsSetMask,
183+ setIsSetMask,
184+ ];
118}185}
119 186 
120def UnaryL2OpInterface : AscendC_OpInterface<"UnaryL2Op", [UnaryOpInterface]> {187def UnaryL2OpInterface : AscendC_OpInterface<"UnaryL2Op", [UnaryOpInterface]> {
121 let description = "Vector unary operation (L2 API)";188 let description = "Vector unary operation (L2 API)";
122- let methods = [getCalCount];189+ let methods = [getCalCount, getCalCountMutable];
123}190}
124 191 
125//===----------------------------------------------------------------------===//192//===----------------------------------------------------------------------===//
126// Vector binary operations193// Vector binary operations
127//===----------------------------------------------------------------------===//194//===----------------------------------------------------------------------===//
128 195 
129-def BinaryOpInterface196+def BinaryOpInterface : AscendC_OpInterface<"BinaryOp", [
130- : AscendC_OpInterface<"BinaryOp", [VectorOpInterface, OpWithDstInterface]> {197+ VectorOpInterface, OpWithDstInterface, OpWithSrcInterface, OpWithReusableSrcInterface
198+]> {
131 let description = "Vector binary operation";199 let description = "Vector binary operation";
132 let methods = [getSrc0, getSrc1];200 let methods = [getSrc0, getSrc1];
133}201}
134 202 
203+def BinaryL0OpInterface : AscendC_OpInterface<"BinaryL0Op", [BinaryOpInterface]> {
204+ let description = "Vector binary operation (L0 API)";
205+ let methods = [
206+ getMask,
207+ getRepeatTimes,
208+ getRepeatParams,
209+ getMaskMutable,
210+ getRepeatTimesMutable,
211+ getRepeatParamsMutable,
212+ getIsSetMask,
213+ setIsSetMask,
214+ ];
215+}
216+ 
217+def BinaryL2OpInterface : AscendC_OpInterface<"BinaryL2Op", [BinaryOpInterface]> {
218+ let description = "Vector binary operation (L2 API)";
219+ let methods = [getCalCount, getCalCountMutable];
220+}
221+ 
135def BinaryL3OpInterface : AscendC_OpInterface<"BinaryL3Op", [BinaryOpInterface]> {222def BinaryL3OpInterface : AscendC_OpInterface<"BinaryL3Op", [BinaryOpInterface]> {
136 let description = "Vector binary operation (L3 API)";223 let description = "Vector binary operation (L3 API)";
137}224}
@@ -140,28 +227,58 @@ def BinaryL3OpInterface : AscendC_OpInterface<"BinaryL3Op", [BinaryOpInterface]>
140// Vector-scalar operations227// Vector-scalar operations
141//===----------------------------------------------------------------------===//228//===----------------------------------------------------------------------===//
142 229 
143-def VecScalarOpInterface230+def VecScalarOpInterface : AscendC_OpInterface<"VecScalarOp", [
144- : AscendC_OpInterface<"VecScalarOp", [VectorOpInterface, OpWithDstInterface]> {231+ VectorOpInterface, OpWithDstInterface, OpWithSrcInterface, OpWithReusableSrcInterface
232+]> {
145 let description = "Vector-scalar operation";233 let description = "Vector-scalar operation";
146- let methods = [getSrc, getScalar];234+ let methods = [getSrc, getScalar, getScalarMutable];
147}235}
148 236 
149def VecScalarL0OpInterface : AscendC_OpInterface<"VecScalarL0Op", [VecScalarOpInterface]> {237def VecScalarL0OpInterface : AscendC_OpInterface<"VecScalarL0Op", [VecScalarOpInterface]> {
150 let description = "Vector-scalar operation (L0 API)";238 let description = "Vector-scalar operation (L0 API)";
151- let methods = [getMask, getRepeatTimes, getRepeatParams];239+ let methods = [
240+ getMask,
241+ getRepeatTimes,
242+ getRepeatParams,
243+ getMaskMutable,
244+ getRepeatTimesMutable,
245+ getRepeatParamsMutable,
246+ getIsSetMask,
247+ setIsSetMask,
248+ ];
152}249}
153 250 
154def VecScalarL2OpInterface : AscendC_OpInterface<"VecScalarL2Op", [VecScalarOpInterface]> {251def VecScalarL2OpInterface : AscendC_OpInterface<"VecScalarL2Op", [VecScalarOpInterface]> {
155 let description = "Vector-scalar operation (L2 API)";252 let description = "Vector-scalar operation (L2 API)";
156- let methods = [getCalCount];253+ let methods = [getCalCount, getCalCountMutable];
254+}
255+ 
256+//===----------------------------------------------------------------------===//
257+// Register API operations
258+//===----------------------------------------------------------------------===//
259+ 
260+def BinaryRegOpInterface : AscendC_OpInterface<"BinaryRegOp", [RegOpInterface]> {
261+ let description = "Binary operation (register API)";
262+ let methods = [getDstReg, getSrc0Reg, getSrc1Reg, getMaskReg];
263+}
264+ 
265+def UnaryRegOpInterface : AscendC_OpInterface<"UnaryRegOp", [RegOpInterface]> {
266+ let description = "Unary operation (register API)";
267+ let methods = [getDstReg, getSrcReg, getMaskReg];
268+}
269+ 
270+def VecScalarRegOpInterface : AscendC_OpInterface<"VecScalarRegOp", [RegOpInterface]> {
271+ let description = "Vector scalar operation (register API)";
272+ let methods = [getDstReg, getSrcReg, getScalar, getMaskReg];
157}273}
158 274 
159//===----------------------------------------------------------------------===//275//===----------------------------------------------------------------------===//
160// Math library operations276// Math library operations
161//===----------------------------------------------------------------------===//277//===----------------------------------------------------------------------===//
162 278 
163-def MathLibraryOpInterface279+def MathLibraryOpInterface : AscendC_OpInterface<"MathLibraryOp", [
164- : AscendC_OpInterface<"MathLibraryOp", [VectorOpInterface, OpWithDstInterface]> {280+ VectorOpInterface, OpWithDstInterface, OpWithSrcInterface
281+]> {
165 let description = "Math library operation";282 let description = "Math library operation";
166}283}
167 284 
@@ -184,4 +301,15 @@ def DataConversionOpInterface
184 let description = "Data conversion operation";301 let description = "Data conversion operation";
185}302}
186 303 
304+//===----------------------------------------------------------------------===//
305+// Reduction operations
306+//===----------------------------------------------------------------------===//
307+ 
308+def ReduceOpInterface : AscendC_OpInterface<"ReduceOp", [
309+ VectorOpInterface, OpWithDstInterface, OpWithSrcInterface
310+]> {
311+ let methods = [getPattern, getSharedTmpBuffer, getSrcShape, getSrc,
312+ GetMethod<"isReuseSource", "getIsReuseSource", "bool">];
313+}
314+ 
187#endif // ASC_INTERFACES_TD315#endif // ASC_INTERFACES_TD
@@ -12,11 +12,13 @@
12#define ASC_OPS_TD12#define ASC_OPS_TD
13 13 
14include "Adv/Activation.td"14include "Adv/Activation.td"
15+include "Adv/Broadcast.td"
15include "Adv/Kfc.td"16include "Adv/Kfc.td"
16include "Adv/Math.td"17include "Adv/Math.td"
17include "Adv/Matmul.td"18include "Adv/Matmul.td"
18include "Adv/Normalization.td"19include "Adv/Normalization.td"
19include "Adv/Quantization.td"20include "Adv/Quantization.td"
21+include "Adv/Reduction.td"
20include "Adv/Sort.td"22include "Adv/Sort.td"
21include "Base.td"23include "Base.td"
22include "Basic/Common.td"24include "Basic/Common.td"
@@ -37,6 +39,7 @@ include "Basic/OpLimits.td"
37include "Basic/OpListTensor.td"39include "Basic/OpListTensor.td"
38include "Basic/OpMm.td"40include "Basic/OpMm.td"
39include "Basic/OpProposal.td"41include "Basic/OpProposal.td"
42+include "Basic/OpReg.td"
40include "Basic/OpScalar.td"43include "Basic/OpScalar.td"
41include "Basic/OpSetAtomic.td"44include "Basic/OpSetAtomic.td"
42include "Basic/OpSwapMem.td"45include "Basic/OpSwapMem.td"
@@ -111,12 +114,37 @@ def AscendC_AscendIsAIVOp : APIOp<"ascend_is_aiv", "AscendIsAIV"> {
111 let assemblyFormat = "attr-dict `:` type($result)";114 let assemblyFormat = "attr-dict `:` type($result)";
112}115}
113 116 
117+def AscendC_YieldOp : AscendC_Op<"yield", [Terminator]> {
118+ let summary = "Terminator for operations with region";
119+ let arguments = (ins Variadic<AnyType>:$operands);
120+ let assemblyFormat = "($operands^ `:` type($operands))? attr-dict";
121+ let builders = [OpBuilder<(ins), "/* empty yield */">];
122+}
123+ 
124+def AscendC_IfAICOp : AscendC_Op<"if_aic", [SingleBlockImplicitTerminator<"YieldOp">]> {
125+ let summary = "Conditional execution block targeting AIC (AI Cube) compute unit";
126+ let arguments = (ins Variadic<AnyType>:$srcTensors);
127+ let results = (outs Variadic<AnyType>:$results);
128+ let regions = (region SizedRegion<1>:$region);
129+ let assemblyFormat = "(`(` $srcTensors^ `:` type($srcTensors) `)`)? (`->` type($results)^)? $region attr-dict";
130+ let hasCanonicalizer = 1;
131+}
132+ 
133+def AscendC_IfAIVOp : AscendC_Op<"if_aiv", [SingleBlockImplicitTerminator<"YieldOp">]> {
134+ let summary = "Conditional execution block targeting AIV (AI Vector) compute unit";
135+ let arguments = (ins Variadic<AnyType>:$srcTensors);
136+ let results = (outs Variadic<AnyType>:$results);
137+ let regions = (region SizedRegion<1>:$region);
138+ let assemblyFormat = "(`(` $srcTensors^ `:` type($srcTensors) `)`)? (`->` type($results)^)? $region attr-dict";
139+ let hasCanonicalizer = 1;
140+}
141+ 
114def AscendC_FftsCrossCoreSyncOp : APIOp<"ffts_cross_core_sync", "FftsCrossCoreSync"> {142def AscendC_FftsCrossCoreSyncOp : APIOp<"ffts_cross_core_sync", "FftsCrossCoreSync"> {
115 let arguments = (ins AscendC_PipeAttr:$pipe, AnyType:$config);143 let arguments = (ins AscendC_PipeAttr:$pipe, AnyType:$config);
116}144}
117 145 
118def AscendC_SetFftsBaseAddrOp : APIOp<"set_ffts_base_addr", "SetFftsBaseAddr"> {146def AscendC_SetFftsBaseAddrOp : APIOp<"set_ffts_base_addr", "SetFftsBaseAddr"> {
119- let arguments = (ins AnyMemRef);147+ let arguments = (ins AnyRankedOrUnrankedMemRef);
120}148}
121 149 
122def AscendC_PopStackBufferOp : APIOp<"pop_stack_buffer", "PopStackBuffer"> {150def AscendC_PopStackBufferOp : APIOp<"pop_stack_buffer", "PopStackBuffer"> {
@@ -126,17 +154,19 @@ def AscendC_PopStackBufferOp : APIOp<"pop_stack_buffer", "PopStackBuffer"> {
126 154 
127def AscendC_LocalTensorAutoOp : APIOp<"local_tensor_auto", "LocalTensorAuto"> {155def AscendC_LocalTensorAutoOp : APIOp<"local_tensor_auto", "LocalTensorAuto"> {
128 let summary = "Create virtual tensor with automatic allocation semantic";156 let summary = "Create virtual tensor with automatic allocation semantic";
129- let arguments = (ins UnitAttr:$input, UnitAttr:$output, Variadic<I64>:$dynamicShape);157+ let arguments = (ins UnitAttr:$input, UnitAttr:$output, Variadic<I64>:$dynamicShape,
158+ DefaultValuedAttr<AscendC_TPositionAttr, "TPosition::VECCALC">:$position);
130 let results = (outs AscendC_LocalTensor:$result);159 let results = (outs AscendC_LocalTensor:$result);
131 let assemblyFormat = [{160 let assemblyFormat = [{
132- `(` $dynamicShape `)` (`input` $input^)? (`output` $output^)? attr-dict `:`161+ $position `(` $dynamicShape `)` (`input` $input^)? (`output` $output^)? attr-dict `:`
133 type($result)162 type($result)
134 }];163 }];
135 let builders = [164 let builders = [
136- OpBuilder<(ins "Type":$result), [{165+ OpBuilder<(ins "Type":$result, CArg<"TPosition", "TPosition::VECCALC">:$position), [{
137- build($_builder, $_state, result, false, false, ValueRange{});166+ build($_builder, $_state, result, false, false, ValueRange{}, position);
138 }]>,167 }]>,
139 ];168 ];
169+ let hasCanonicalizeMethod = 1;
140}170}
141 171 
142#endif // ASC_OPS_TD172#endif // ASC_OPS_TD
@@ -25,11 +25,11 @@ std::unique_ptr<Pass> createDetectKernelTypePass();
25std::unique_ptr<Pass> createEraseSyncPass();25std::unique_ptr<Pass> createEraseSyncPass();
26std::unique_ptr<Pass> createGenerateBoilerplatePass();26std::unique_ptr<Pass> createGenerateBoilerplatePass();
27std::unique_ptr<Pass> createHoistQueBindPass();27std::unique_ptr<Pass> createHoistQueBindPass();
28-std::unique_ptr<Pass> createHoistUBAllocationPass();28+std::unique_ptr<Pass> createHoistTensorAllocationPass(bool excludeInOut = false);
29std::unique_ptr<Pass> createInputOutputTensorPass();29std::unique_ptr<Pass> createInputOutputTensorPass();
30-std::unique_ptr<Pass> createInsertSyncPass();30+std::unique_ptr<Pass> createInsertQueSyncPass();
31-std::unique_ptr<Pass> createLegalizeKernelArgsPass();31+std::unique_ptr<Pass> createLegalizeKernelArgsPass(bool setFftsAddr = false);
32-std::unique_ptr<Pass> createMaterializeTensorPass();32+std::unique_ptr<Pass> createMaterializeTensorPass(bool alwaysBuf = false);
33std::unique_ptr<Pass> createNoopPass();33std::unique_ptr<Pass> createNoopPass();
34std::unique_ptr<Pass> createPrivatizeFuncPass();34std::unique_ptr<Pass> createPrivatizeFuncPass();
35std::unique_ptr<Pass> createUnifyPipePass();35std::unique_ptr<Pass> createUnifyPipePass();
@@ -14,92 +14,285 @@
14include "mlir/Pass/PassBase.td"14include "mlir/Pass/PassBase.td"
15 15 
16def DeclarePyStruct : Pass<"ascendc-declare-py-struct", "ModuleOp"> {16def DeclarePyStruct : Pass<"ascendc-declare-py-struct", "ModuleOp"> {
17- let summary = "Insert emitasc.declare_py_struct";17+ let summary = "Insert emitasc.declare_py_struct operations for all Python struct types used in the module";
18 let constructor = "mlir::ascendc::createDeclarePyStructPass()";18 let constructor = "mlir::ascendc::createDeclarePyStructPass()";
19 let dependentDialects = ["emitasc::EmitAscDialect"];19 let dependentDialects = ["emitasc::EmitAscDialect"];
20+ let description = [{
21+ Collects all `emitasc.py_struct` types referenced in function arguments, block arguments, and operation results,
22+ then inserts `emitasc.declare_py_struct` operations at the module beginning to declare them for C++ code emission.
23+ 
24+ The pass deduplicates struct types while preserving the order of first occurrence. Nested struct types
25+ (structs containing other structs) are also declared.
26+ }];
20}27}
21 28 
22def DefineCubeOnly : Pass<"ascendc-define-cube-only", "ModuleOp"> {29def DefineCubeOnly : Pass<"ascendc-define-cube-only", "ModuleOp"> {
23- let summary = "Insert emitc.define if CUBE-ONLY";30+ let summary = "Define ASCENDC_CUBE_ONLY macro for cube-only (matrix multiplication) kernels";
24 let constructor = "mlir::ascendc::createDefineCubeOnlyPass()";31 let constructor = "mlir::ascendc::createDefineCubeOnlyPass()";
25 let dependentDialects = ["emitc::EmitCDialect"];32 let dependentDialects = ["emitc::EmitCDialect"];
26-}33+ let description = [{
27- 34+ Inserts `#define ASCENDC_CUBE_ONLY` verbatim at module start and sets `matmul_cube_only` attribute.
28-def DetectKernelType : Pass<"ascendc-detect-kernel-type", "ModuleOp"> {35+ This macro enables cube-only optimizations in the Ascend C runtime for kernels that only perform
29- let summary = "Check whether the kernel is vector-only or mixed";36+ matrix multiplication without vector operations.
30- let constructor = "mlir::ascendc::createDetectKernelTypePass()";37+ }];
31-}
32- 
33-def EraseSync : Pass<"ascendc-erase-sync", "func::FuncOp"> {
34- let summary = "Erase intra-core synchronization operations";
35- let constructor = "mlir::ascendc::createEraseSyncPass()";
36-}
37- 
38-def GenerateBoilerplate : Pass<"ascendc-generate-boilerplate", "ModuleOp"> {
39- let summary = "Insert emitc.include and additional boilerplate code";
40- let constructor = "mlir::ascendc::createGenerateBoilerplatePass()";
41- let dependentDialects = ["emitc::EmitCDialect"];
42-}
43- 
44-def HoistQueBind : Pass<"ascendc-hoist-que-bind", "func::FuncOp"> {
45- let summary = "Hoist TQueBind, TQue, TBuf initialization operations";
46- let constructor = "mlir::ascendc::createHoistQueBindPass()";
47-}
48- 
49-def HoistUBAllocation : Pass<"ascendc-hoist-ub-allocation", "func::FuncOp"> {
50- let summary = "Hoist tensor allocations to the function root";
51- let constructor = "mlir::ascendc::createHoistUBAllocationPass()";
52-}
53- 
54-def InputOutputTensor : Pass<"ascendc-input-output-tensor", "func::FuncOp"> {
55- let summary = "Set input and output operand for local_tensor_auto";
56- let constructor = "mlir::ascendc::createInputOutputTensorPass()";
57-}
58- 
59-def InsertSync : Pass<"ascendc-insert-sync", "func::FuncOp"> {
60- let summary = "Insert intra-core synchronization operations";
61- let constructor = "mlir::ascendc::createInsertSyncPass()";
62- let dependentDialects = ["ascendc::AscendCDialect"];
63-}
64- 
65-def MaterializeTensor : Pass<"ascendc-materialize-tensor", "func::FuncOp"> {
66- let summary = "Insert ascendc.tbuf, ascendc.queue and ascendc.alloca for local_tensor_auto";
67- let constructor = "mlir::ascendc::createMaterializeTensorPass()";
68- let dependentDialects = ["arith::ArithDialect", "ascendc::AscendCDialect"];
69-}
70- 
71-def LegalizeKernelArgs : Pass<"ascendc-legalize-kernel-args", "ModuleOp"> {
72- let summary = "Attach emitasc.kernel_arg attributes and insert operations";
73- let constructor = "mlir::ascendc::createLegalizeKernelArgsPass()";
74- let dependentDialects = [
75- "arith::ArithDialect", "ascendc::AscendCDialect", "emitasc::EmitAscDialect",
76- "scf::SCFDialect",
77- ];
78-}
79- 
80-def Noop : Pass<"ascendc-noop", "func::FuncOp"> {
81- let summary = "This pass does nothing";
82- let constructor = "mlir::ascendc::createNoopPass()";
83-}
84- 
85-def PrivatizeFunc : Pass<"ascendc-privatize-func", "ModuleOp"> {
86- let summary = "Mark functions without ascendc.global attribute as private";
87- let constructor = "mlir::ascendc::createPrivatizeFuncPass()";
88-}
89- 
90-def UnifyPipe : Pass<"ascendc-unify-pipe", "func::FuncOp"> {
91- let summary = "Unify pipe opertation";
92- let constructor = "mlir::ascendc::createUnifyPipePass()";
93-}
94- 
95-def VerifySync : Pass<"ascendc-verify-sync", "func::FuncOp"> {
96- let summary = "Verify TQue synchronization";
97- let constructor = "mlir::ascendc::createVerifySyncPass()";
98}38}
99 39 
100def DetectEnableDebug : Pass<"ascendc-detect-enable-debug", "ModuleOp"> {40def DetectEnableDebug : Pass<"ascendc-detect-enable-debug", "ModuleOp"> {
101- let summary = "Check whether the kernel is using debug utils";41+ let summary = "Detect debug utility usage (printf, dump_tensor) and set enable_debug attribute";
102 let constructor = "mlir::ascendc::createDetectEnableDebugPass()";42 let constructor = "mlir::ascendc::createDetectEnableDebugPass()";
43+ let description = [{
44+ Scans the module for `PrintfOp` or `DumpTensorOp` operations. If any are found, sets the `enable_debug`
45+ unit attribute on the module to enable debug runtime support during kernel execution.
46+ }];
47+}
48+ 
49+def DetectKernelType : Pass<"ascendc-detect-kernel-type", "ModuleOp"> {
50+ let summary = "Classify kernel as vector, cube, or mixed based on operation types present";
51+ let constructor = "mlir::ascendc::createDetectKernelTypePass()";
52+ let description = [{
53+ Analyzes operations in the module to determine kernel type and sets `kernel_type` string attribute:
54+ - "vector": Only vector operations present (no `MmadOp` or `RegistMatmulObjOp`)
55+ - "cube": Only matrix multiplication operations present (no `VectorOp`)
56+ - "mixed": Both vector and cube operations present
57+ 
58+ This classification affects synchronization strategy and runtime configuration.
59+ }];
60+}
61+ 
62+def EraseSync : Pass<"ascendc-erase-sync", "func::FuncOp"> {
63+ let summary = "Remove TQueBind synchronization operations and replace deque_tensor with allocated tensors";
64+ let constructor = "mlir::ascendc::createEraseSyncPass()";
65+ let description = [{
66+ Removes intra-core synchronization infrastructure for static allocation mode:
67+ - Erases `TQueBindEnqueTensorOp`, `SetFlagOp`, `WaitFlagOp`, `PipeBarrierOp`
68+ - Replaces `TQueBindDequeTensorOp` with the tensor from corresponding `TQueBindAllocTensorOp`
69+ 
70+ This pass is used when tensors are statically allocated and queue-based synchronization is not needed.
71+ }];
72+}
73+ 
74+def GenerateBoilerplate : Pass<"ascendc-generate-boilerplate", "ModuleOp"> {
75+ let summary = "Insert required C++ include headers based on operations used in the kernel";
76+ let constructor = "mlir::ascendc::createGenerateBoilerplatePass()";
77+ let dependentDialects = ["emitc::EmitCDialect"];
78+ let description = [{
79+ Adds necessary Ascend C header includes at module start:
80+ - Always includes `kernel_operator.h` (core runtime)
81+ - If `ListTensorDescOp` or `ListTensorDescV2Op` present: includes `kernel_operator_list_tensor_intf.h`
82+ - If `RegistMatmulObjOp` present: includes `lib/matmul_intf.h`
83+ - If `TensorDescOp` present: includes `kernel_operator_list_tensor_intf.h`
84+ }];
85+}
86+ 
87+def HoistQueBind : Pass<"ascendc-hoist-que-bind", "func::FuncOp"> {
88+ let summary = "Move queue and buffer initialization operations to the function entry block";
89+ let constructor = "mlir::ascendc::createHoistQueBindPass()";
90+ let description = [{
91+ Hoists TQueBind-related initialization ops from loops to function root:
92+ - `QueueOp`, `QueBindOp`, `TBufOp` declarations
93+ - `TPipeInitQueueOp`, `TPipeInitBufferOp` initialization
94+ - `TBufGetTensorOp` tensor retrieval
95+ 
96+ This optimization ensures queue/buffer setup happens once at kernel entry rather than repeatedly in loops.
97+ }];
98+}
99+ 
100+def HoistTensorAllocation : Pass<"ascendc-hoist-tensor-allocation", "func::FuncOp"> {
101+ let summary = "Move LocalTensorAutoOp allocations from loops to the function entry block";
102+ let constructor = "mlir::ascendc::createHoistTensorAllocationPass()";
103+ let options = [
104+ Option<"excludeInOut", "exclude-in-out", "bool", "false",
105+ "Keep input/output tensors inside loops, only hoist temporary tensors">,
106+ ];
107+ let description = [{
108+ Hoists `LocalTensorAutoOp` tensor allocations to the function root to reduce repeated allocations in loops.
109+ When `exclude-in-out=true`: only hoists tensors without `input` or `output` flags (pure temporary tensors).
110+ When `exclude-in-out=false`: hoists all tensors regardless of input/output status.
111+ 
112+ Input/output tensors are those connected to GM via `DataCopyOp` with `GlobalToLocal`/`LocalToGlobal` direction.
113+ }];
114+}
115+ 
116+def InputOutputTensor : Pass<"ascendc-input-output-tensor", "func::FuncOp"> {
117+ let summary = "Mark tensors as input/output based on GM data copy direction and handle input-only to output copies";
118+ let constructor = "mlir::ascendc::createInputOutputTensorPass()";
119+ let description = [{
120+ Analyzes `DataCopyOp` to set `input`/`output` flags on `LocalTensorAutoOp`:
121+ - `gm_ubuf` direction → tensor is input (loaded from GM)
122+ - `ubuf_gm` direction → tensor is output (stored to GM)
123+ 
124+ Also handles special case: if an input-only tensor is used for `ubuf_gm` copy, creates a separate output tensor
125+ and inserts `DataCopyL2Op` to copy input tensor to output tensor before the GM store.
126+ 
127+ Example transformation:
128+ ```mlir
129+ // Before: unmarked tensors
130+ %0 = ascendc.local_tensor_auto vecin() : <64xf32>
131+ ascendc.data_copy_l2 %0, %gm_in, %c64 // gm→ub (input)
132+ %1 = ascendc.local_tensor_auto vecout() : <64xf32>
133+ ascendc.data_copy_l2 %gm_out, %1, %c64 // ub→gm (output)
134+ 
135+ // After: tensors marked with input/output flags
136+ %0 = ascendc.local_tensor_auto vecin() input : <64xf32>
137+ %1 = ascendc.local_tensor_auto vecout() output : <64xf32>
138+ ```
139+ }];
140+}
141+ 
142+def InsertQueSync : Pass<"ascendc-insert-que-sync", "func::FuncOp"> {
143+ let summary = "Insert TQueBind enqueue/dequeue synchronization and scalar get/set_value barriers";
144+ let constructor = "mlir::ascendc::createInsertQueSyncPass()";
145+ let dependentDialects = ["ascendc::AscendCDialect"];
146+ let description = [{
147+ Inserts intra-core synchronization for TQueBind-managed tensors and scalar operations:
148+ 1. **Enqueue tensors**: After each `OpWithDst`, inserts `TQueBindEnqueTensorOp` or `PipeBarrierOp(pipe_v)`
149+ 2. **Dequeue tensors**: For enqueued tensors, inserts `TQueBindDequeTensorOp` before first user
150+ 3. **Get/Set value sync**: Wraps `get_value`/`set_value` with V_S/S_V event sync (fetch_event_id, set_flag, wait_flag)
151+ 4. **Loop handling**: Special sync placement for `set_value` inside single-op loops
152+ 5. **Canonicalization**: Adds `PipeBarrierOp(pipe_all)` at function end and runs canonicalization
153+ 
154+ TQueBind synchronization enables double-buffering and overlapping compute with data transfer.
155+ 
156+ Example: enqueue/dequeue synchronization:
157+ ```mlir
158+ // Before: tensor used directly after allocation
159+ %tensor = ascendc.que_bind.alloc_tensor %queue
160+ ascendc.add_l2 %dst, %tensor, %tensor, %count
161+ 
162+ // After: tensor dequeued before use, enqueued after compute
163+ %tensor = ascendc.que_bind.alloc_tensor %queue
164+ %dequeued = ascendc.que_bind.deque_tensor %queue
165+ ascendc.add_l2 %dst, %dequeued, %dequeued, %count
166+ ascendc.que_bind.enque_tensor %queue, %dst
167+ ```
168+ 
169+ Example: scalar get/set_value with V_S/S_V sync:
170+ ```mlir
171+ // Before: direct scalar access
172+ %val = ascendc.local_tensor.get_value %tensor, %offset
173+ 
174+ // After: wrapped with event sync
175+ %event = ascendc.pipe.fetch_event_id %pipe, v_s
176+ ascendc.set_flag v_s, %event
177+ ascendc.wait_flag v_s, %event
178+ %val = ascendc.local_tensor.get_value %tensor, %offset
179+ %event2 = ascendc.pipe.fetch_event_id %pipe, s_v
180+ ascendc.set_flag s_v, %event2
181+ ascendc.wait_flag s_v, %event2
182+ ```
183+ }];
184+}
185+ 
186+def MaterializeTensor : Pass<"ascendc-materialize-tensor", "func::FuncOp"> {
187+ let summary = "Convert LocalTensorAutoOp placeholders to concrete TBuf or TQueBind allocations";
188+ let constructor = "mlir::ascendc::createMaterializeTensorPass()";
189+ let dependentDialects = ["arith::ArithDialect", "ascendc::AscendCDialect"];
190+ let options = [
191+ Option<"alwaysBuf", "always-buf", "bool", "false",
192+ "Use TBuf for all tensors; otherwise use TQueBind for input/output tensors">,
193+ ];
194+ let description = [{
195+ Materializes tensor allocation placeholders into concrete Ascend C buffer objects:
196+ - **TQueBind path** (for input/output tensors when `alwaysBuf=false`):
197+ Creates `QueueOp` + `TPipeInitQueueOp` + `TQueBindAllocTensorOp` + `TQueBindFreeTensorOp` at function end
198+ - **TBuf path** (for temporary tensors or when `alwaysBuf=true`):
199+ Creates `TBufOp` + `TPipeInitBufferOp` + `TBufGetTensorOp`
200+ 
201+ Input tensors use `VECIN` queue position, output tensors use `VECOUT`. Temporary tensors use `VECCALC` TBuf.
202+ Buffer size computed from tensor shape (static or dynamic via shape operands).
203+ 
204+ Example transformation (TQueBind path for input tensor):
205+ ```mlir
206+ // Before: placeholder
207+ %0 = ascendc.local_tensor_auto vecin() input : <64xf32>
208+ 
209+ // After: queue-based allocation
210+ %pipe = ascendc.pipe
211+ %queue = ascendc.queue : <vecin, 1>
212+ ascendc.pipe.init_queue %pipe, %queue, %c1, %c256
213+ %0 = ascendc.que_bind.alloc_tensor %queue : !ascendc.queue<vecin, 1>, !ascendc.local_tensor<64xf32>
214+ // ... at function end:
215+ ascendc.que_bind.free_tensor %queue, %0
216+ ```
217+ 
218+ Example transformation (TBuf path for temporary tensor):
219+ ```mlir
220+ // Before: placeholder
221+ %0 = ascendc.local_tensor_auto veccalc() : <64xf32>
222+ 
223+ // After: buffer-based allocation
224+ %pipe = ascendc.pipe
225+ %tbuf = ascendc.tbuf : <veccalc>
226+ ascendc.pipe.init_buffer %pipe, %tbuf, %c256
227+ %0 = ascendc.tbuf.get_tensor %tbuf : !ascendc.tbuf<veccalc>, !ascendc.local_tensor<64xf32>
228+ ```
229+ }];
230+}
231+ 
232+def LegalizeKernelArgs : Pass<"ascendc-legalize-kernel-args", "ModuleOp"> {
233+ let summary = "Attach kernel argument attributes and insert ffts_addr handling for multi-core kernels";
234+ let constructor = "mlir::ascendc::createLegalizeKernelArgsPass()";
235+ let dependentDialects = ["arith::ArithDialect", "ascendc::AscendCDialect", "emitasc::EmitAscDialect", "scf::SCFDialect"];
236+ let options = [
237+ Option<"setFftsAddr", "set-ffts-addr", "bool", "false",
238+ "Append ffts_addr kernel argument and call set_ffts_base_addr for cross-core sync">,
239+ ];
240+ let description = [{
241+ Processes kernel entry functions (marked with `ascendc.global` attribute):
242+ 1. Marks all existing arguments as `Explicit` kernel arguments
243+ 2. If `setFftsAddr=true`: adds `ffts_addr` memref argument and calls `SetFftsBaseAddrOp`
244+ 3. If matmul present and not cube-only: inserts `AscendIsAICOp` check with `FftsCrossCoreSyncOp` for AIC mode
245+ 
246+ Kernel argument attributes (`emitasc.kernel_arg`) control how arguments are handled by the runtime.
247+ }];
248+}
249+ 
250+def Noop : Pass<"ascendc-noop", "func::FuncOp"> {
251+ let summary = "Placeholder pass that performs no transformations";
252+ let constructor = "mlir::ascendc::createNoopPass()";
253+ let description = [{
254+ A no-operation pass useful for testing, pipeline debugging, or as a placeholder in pass schedules.
255+ The pass walks the function but performs no modifications to the IR.
256+ }];
257+}
258+ 
259+def PrivatizeFunc : Pass<"ascendc-privatize-func", "ModuleOp"> {
260+ let summary = "Mark non-kernel functions as private and kernel functions as public for emission";
261+ let constructor = "mlir::ascendc::createPrivatizeFuncPass()";
262+ let description = [{
263+ Adjusts function visibility based on kernel status:
264+ - Functions without `ascendc.global` attribute → set as `private` (internal helper functions)
265+ - Functions with `ascendc.global` and body → set as `public` (kernel entry points)
266+ - Functions with `ascendc.global` but no body (declarations) → remain as-is
267+ 
268+ This ensures only kernel entry points are exported while helper functions remain internal.
269+ }];
270+}
271+ 
272+def UnifyPipe : Pass<"ascendc-unify-pipe", "func::FuncOp"> {
273+ let summary = "Replace multiple PipeOp instances with a single unified pipe at function entry";
274+ let constructor = "mlir::ascendc::createUnifyPipePass()";
275+ let description = [{
276+ Consolidates multiple `PipeOp` declarations into a single pipe object at function entry.
277+ All uses of individual pipes are replaced with the unified pipe, and original `PipeOp` instances are erased.
278+ 
279+ This ensures consistent pipe management across the kernel and reduces redundant TPipe object creation.
280+ }];
281+}
282+ 
283+def VerifySync : Pass<"ascendc-verify-sync", "func::FuncOp"> {
284+ let summary = "Validate TQueBind synchronization correctness and emit warnings for mismatches";
285+ let constructor = "mlir::ascendc::createVerifySyncPass()";
286+ let description = [{
287+ Verifies proper pairing of TQueBind operations and emits warnings for incorrect usage:
288+ - `alloc_tensor` without corresponding `free_tensor`
289+ - `free_tensor` for already-freed tensor or without `alloc_tensor`
290+ - `enque_tensor` without corresponding `deque_tensor`
291+ - `deque_tensor` without `enque_tensor` in queue
292+ - Unexpected tensor uses between `enque_tensor` and `deque_tensor`
293+ 
294+ This is a verification-only pass that helps detect synchronization bugs during development.
295+ }];
103}296}
104 297 
105#endif // ASC_PASSES_TD298#endif // ASC_PASSES_TD
@@ -17,12 +17,19 @@ namespace ascendc {
17 17 
18namespace attr {18namespace attr {
19LITERAL aicore = "ascendc.aicore";19LITERAL aicore = "ascendc.aicore";
20-LITERAL api = "ascendc.api";20+LITERAL compilationArch = "asc.compilation_arch";
21-LITERAL compile_mix = "asc.compile_mix";
22LITERAL emitAsUnsigned = "ascendc.emit_as_unsigned";21LITERAL emitAsUnsigned = "ascendc.emit_as_unsigned";
23LITERAL global = "ascendc.global";22LITERAL global = "ascendc.global";
24-LITERAL enable_debug = "asc.enable_debug";23+LITERAL enableDebug = "asc.enable_debug";
24+LITERAL kernelType = "asc.kernel_type";
25LITERAL matmulCubeOnly = "asc.matmul_cube_only";25LITERAL matmulCubeOnly = "asc.matmul_cube_only";
26+LITERAL memoryConsumed = "asc.memory_consumed";
27+LITERAL socVersion = "asc.soc_version";
28+LITERAL vfVecLen = "asc.vf_vec_len";
29+ 
30+LITERAL kernelCube = "cube";
31+LITERAL kernelMixed = "mixed";
32+LITERAL kernelVector = "vector";
26} // namespace attr33} // namespace attr
27 34 
28} // namespace ascendc35} // namespace ascendc
@@ -11,6 +11,7 @@
11#ifndef ASCIR_DIALECT_ASC_UTILS_UTILS_H11#ifndef ASCIR_DIALECT_ASC_UTILS_UTILS_H
12#define ASCIR_DIALECT_ASC_UTILS_UTILS_H12#define ASCIR_DIALECT_ASC_UTILS_UTILS_H
13 13 
14+#include "ascir/Dialect/Asc/IR/Asc.h"
14#include "mlir/Dialect/Func/IR/FuncOps.h"15#include "mlir/Dialect/Func/IR/FuncOps.h"
15#include "mlir/IR/DialectRegistry.h"16#include "mlir/IR/DialectRegistry.h"
16#include "mlir/IR/Dominance.h"17#include "mlir/IR/Dominance.h"
@@ -20,6 +21,12 @@
20namespace mlir {21namespace mlir {
21namespace ascendc {22namespace ascendc {
22 23 
24+constexpr unsigned cubeBlockSize = 16; // In elements
25+constexpr unsigned ubBlockSize = 32; // In bytes
26+constexpr unsigned cubeKBlockBytes = ubBlockSize; // In bytes
27+constexpr unsigned repeatBlockSize = 256; // In bytes
28+constexpr unsigned bitmaskSize = 64;
29+ 
23template <typename OpT>30template <typename OpT>
24struct HoistOpPattern : public OpRewritePattern<OpT> {31struct HoistOpPattern : public OpRewritePattern<OpT> {
25 using OpRewritePattern<OpT>::OpRewritePattern;32 using OpRewritePattern<OpT>::OpRewritePattern;
@@ -44,12 +51,34 @@ struct HoistOpPattern : public OpRewritePattern<OpT> {
44 }51 }
45};52};
46 53 
54+int64_t getTypeSize(Type type);
55+ 
56+int64_t getTypeSizeCubeBlockAlign(ShapedType type, TPosition position);
57+ 
58+int64_t getElementTypeSize(ShapedType type);
59+ 
47bool opPrecedes(Operation* lhs, Operation* rhs);60bool opPrecedes(Operation* lhs, Operation* rhs);
48 61 
49bool opPrecedes(Operation* lhs, Operation* rhs, DominanceInfo& di);62bool opPrecedes(Operation* lhs, Operation* rhs, DominanceInfo& di);
50 63 
51void registerInlinerInterfaces(DialectRegistry& registry);64void registerInlinerInterfaces(DialectRegistry& registry);
52 65 
66+ModuleOp getModule(Operation* op);
67+ 
68+StringRef getCompilationArch(Operation* op);
69+ 
70+StringRef getSocVersion(Operation* op);
71+ 
72+std::optional<int64_t> getVecLen(Operation* op);
73+ 
74+bool isTargetArchC310(Operation* op);
75+ 
76+SmallVector<Operation*> collectAllUsers(ascendc::LocalTensorAutoOp tensorOp);
77+ 
78+ascendc::LocalTensorAutoOp getAllocationRoot(Value v);
79+ 
80+Pipe getOpPipe(Operation* op, Pipe defaultPipe = Pipe::PIPE_S);
81+ 
53} // namespace ascendc82} // namespace ascendc
54} // namespace mlir83} // namespace mlir
55 84 
@@ -0,0 +1,20 @@
1+/*
2+ * Copyright (c) 2026 Huawei Technologies Co., Ltd.
3+ * This program is free software, you can redistribute it and/or modify it under the terms and conditions of
4+ * 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,
7+ * 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+#ifndef EMITASC_DOC_TD
12+#define EMITASC_DOC_TD
13+ 
14+include "Dialect.td"
15+ 
16+include "Attributes.td"
17+include "Ops.td"
18+include "Types.td"
19+ 
20+#endif // EMITASC_DOC_TD
@@ -15,6 +15,7 @@
15#include "mlir/IR/Dialect.h"15#include "mlir/IR/Dialect.h"
16#include "mlir/IR/OpDefinition.h"16#include "mlir/IR/OpDefinition.h"
17#include "mlir/Interfaces/CastInterfaces.h"17#include "mlir/Interfaces/CastInterfaces.h"
18+#include "mlir/Interfaces/LoopLikeInterface.h"
18#include "mlir/Interfaces/SideEffectInterfaces.h"19#include "mlir/Interfaces/SideEffectInterfaces.h"
19#include "mlir/Interfaces/ViewLikeInterface.h"20#include "mlir/Interfaces/ViewLikeInterface.h"
20 21 
@@ -15,8 +15,10 @@ include "Dialect.td"
15include "Types.td"15include "Types.td"
16 16 
17include "mlir/Interfaces/CastInterfaces.td"17include "mlir/Interfaces/CastInterfaces.td"
18+include "mlir/Interfaces/LoopLikeInterface.td"
18include "mlir/Interfaces/SideEffectInterfaces.td"19include "mlir/Interfaces/SideEffectInterfaces.td"
19include "mlir/Interfaces/ViewLikeInterface.td"20include "mlir/Interfaces/ViewLikeInterface.td"
21+include "mlir/IR/AttrTypeBase.td"
20include "mlir/IR/BuiltinAttributeInterfaces.td"22include "mlir/IR/BuiltinAttributeInterfaces.td"
21 23 
22class EmitAsc_Op<string mnemonic, list<Trait> traits = []>24class EmitAsc_Op<string mnemonic, list<Trait> traits = []>
@@ -40,7 +42,7 @@ def EmitAsc_CallOpaqueOp : EmitAsc_Op<"call_opaque"> {
40 42 
41def EmitAsc_CopyStructOp : EmitAsc_Op<"copy_struct"> {43def EmitAsc_CopyStructOp : EmitAsc_Op<"copy_struct"> {
42 let summary = "Create local structure and perform memcpy from another one";44 let summary = "Create local structure and perform memcpy from another one";
43- let arguments = (ins AnyMemRef:$base);45+ let arguments = (ins AnyRankedOrUnrankedMemRef:$base);
44 let results = (outs AnyType:$result);46 let results = (outs AnyType:$result);
45 let assemblyFormat = "$base attr-dict `:` type($base) `,` type($result)";47 let assemblyFormat = "$base attr-dict `:` type($base) `,` type($result)";
46}48}
@@ -53,11 +55,45 @@ def EmitAsc_DeclarePyStructOp : EmitAsc_Op<"declare_py_struct", [Pure]> {
53 55 
54def EmitAsc_DereferenceOp : EmitAsc_Op<"dereference", [Pure]> {56def EmitAsc_DereferenceOp : EmitAsc_Op<"dereference", [Pure]> {
55 let summary = "Dereference pointer to underlying object (unary `operator*`)";57 let summary = "Dereference pointer to underlying object (unary `operator*`)";
56- let arguments = (ins AnyMemRef:$base);58+ let arguments = (ins AnyRankedOrUnrankedMemRef:$base);
57 let results = (outs AnyType:$result);59 let results = (outs AnyType:$result);
58 let assemblyFormat = "$base attr-dict `:` type($base) `,` type($result)";60 let assemblyFormat = "$base attr-dict `:` type($base) `,` type($result)";
59}61}
60 62 
63+def EmitAsc_InitStructOp : EmitAsc_Op<"init_struct", [Pure]> {
64+ let summary = "Create a struct instance and initialize its fields";
65+ let arguments = (ins StrArrayAttr:$fieldNames, Variadic<AnyType>:$fieldValues);
66+ let results = (outs AnyType:$result);
67+ let builders = [
68+ OpBuilder<(ins "Type":$result), [{
69+ build($_builder, $_state, result, $_builder.getArrayAttr(std::nullopt),
70+ ValueRange{});
71+ }]>,
72+ ];
73+ let hasCustomAssemblyFormat = 1;
74+ let hasVerifier = 1;
75+ let extraClassDeclaration = [{
76+ size_t getNumFields();
77+ bool hasField(::llvm::StringRef name);
78+ Value getField(::llvm::StringRef name);
79+ std::optional<size_t> getFieldOperandIndex(::llvm::StringRef name);
80+ void setField(::llvm::StringRef name, Value value);
81+ void addField(::llvm::StringRef name, Value value);
82+ }];
83+}
84+ 
85+def EmitAsc_MaskOp : EmitAsc_Op<"mask", [Pure]> {
86+ let summary = "Creates a mask in bit-by-bit mode (pair of 64-bit integers)";
87+ let arguments = (ins I64 : $maskH, I64 : $maskL);
88+ let results = (outs EmitAsc_Mask : $mask);
89+ let assemblyFormat = "$maskH `,` $maskL attr-dict";
90+ let builders = [
91+ OpBuilder<(ins "Value":$maskH, "Value":$maskL), [{
92+ build($_builder, $_state, MaskType::get($_builder.getContext()), maskH, maskL);
93+ }]>,
94+ ];
95+}
96+ 
61def EmitAsc_MemberOp : EmitAsc_Op<"member"> {97def EmitAsc_MemberOp : EmitAsc_Op<"member"> {
62 let arguments = (ins AnyType:$base, StrAttr:$field);98 let arguments = (ins AnyType:$base, StrAttr:$field);
63 let results = (outs AnyType:$result);99 let results = (outs AnyType:$result);
@@ -76,8 +112,8 @@ def EmitAsc_MemberPtrOp : EmitAsc_Op<"member_ptr", [Pure]> {
76 auto *result = (ResultT *)&base->member;112 auto *result = (ResultT *)&base->member;
77 ```113 ```
78 }];114 }];
79- let arguments = (ins AnyMemRef:$base, IndexAttr:$index, OptionalAttr<StrAttr>:$field);115+ let arguments = (ins AnyRankedOrUnrankedMemRef:$base, IndexAttr:$index, OptionalAttr<StrAttr>:$field);
80- let results = (outs AnyMemRef:$result);116+ let results = (outs AnyRankedOrUnrankedMemRef:$result);
81 let assemblyFormat = [{117 let assemblyFormat = [{
82 $base `[` $index `]` ($field^)? attr-dict `:` type($base) `,` type($result)118 $base `[` $index `]` ($field^)? attr-dict `:` type($base) `,` type($result)
83 }];119 }];
@@ -93,7 +129,7 @@ def EmitAsc_MemberRefOp : EmitAsc_Op<"member_ref", [Pure]> {
93 ResultT &result = base->member;129 ResultT &result = base->member;
94 ```130 ```
95 }];131 }];
96- let arguments = (ins AnyMemRef:$base, IndexAttr:$index, OptionalAttr<StrAttr>:$field);132+ let arguments = (ins AnyRankedOrUnrankedMemRef:$base, IndexAttr:$index, OptionalAttr<StrAttr>:$field);
97 let results = (outs AnyType:$result);133 let results = (outs AnyType:$result);
98 let assemblyFormat = [{134 let assemblyFormat = [{
99 $base `[` $index `]` ($field^)? attr-dict `:` type($base) `,` type($result)135 $base `[` $index `]` ($field^)? attr-dict `:` type($base) `,` type($result)
@@ -111,9 +147,9 @@ def EmitAsc_PtrOffsetOp : EmitAsc_Op<"ptr_offset", [
111 Pure, DeclareOpInterfaceMethods<ViewLikeOpInterface>147 Pure, DeclareOpInterfaceMethods<ViewLikeOpInterface>
112]> {148]> {
113 let summary = "Apply `operator+` between pointer and offset";149 let summary = "Apply `operator+` between pointer and offset";
114- let arguments = (ins AnyMemRef:$base, OptionalAttr<IndexAttr>:$staticOffset,150+ let arguments = (ins AnyRankedOrUnrankedMemRef:$base, OptionalAttr<IndexAttr>:$staticOffset,
115 Optional<Index>:$dynamicOffset);151 Optional<Index>:$dynamicOffset);
116- let results = (outs AnyMemRef:$result);152+ let results = (outs AnyRankedOrUnrankedMemRef:$result);
117 let assemblyFormat = [{153 let assemblyFormat = [{
118 $base `[` ($dynamicOffset^):($staticOffset ` `)? `]` attr-dict `:`154 $base `[` ($dynamicOffset^):($staticOffset ` `)? `]` attr-dict `:`
119 type($base) `,` type($result)155 type($base) `,` type($result)
@@ -130,7 +166,7 @@ def EmitAsc_ReinterpretCastOp : EmitAsc_Op<"reinterpret_cast", [
130 let assemblyFormat = "$source attr-dict `:` type($source) `to` type($result)";166 let assemblyFormat = "$source attr-dict `:` type($source) `to` type($result)";
131}167}
132 168 
133-def EmitAsc_VariableOp : EmitAsc_Op<"variable", [Pure]> {169+def EmitAsc_VariableOp : EmitAsc_Op<"variable"> {
134 let summary = "Define a mutable variable";170 let summary = "Define a mutable variable";
135 let description = [{171 let description = [{
136 This operation combines an allocation and storing a value to `memref<1x...>`.172 This operation combines an allocation and storing a value to `memref<1x...>`.
@@ -162,6 +198,7 @@ def EmitAsc_VariableOp : EmitAsc_Op<"variable", [Pure]> {
162 bool isStatic();198 bool isStatic();
163 OpFoldResult getInit(bool fold = true);199 OpFoldResult getInit(bool fold = true);
164 }];200 }];
201+ let hasCanonicalizeMethod = 1;
165}202}
166 203 
167def EmitAsc_VerbatimOp : EmitAsc_Op<"verbatim"> {204def EmitAsc_VerbatimOp : EmitAsc_Op<"verbatim"> {
@@ -19,6 +19,13 @@ class EmitAsc_Type<string name, string typeMnemonic, list<Trait> traits = []>
19 let mnemonic = typeMnemonic;19 let mnemonic = typeMnemonic;
20}20}
21 21 
22+def EmitAsc_Mask : EmitAsc_Type<"Mask", "mask"> {
23+ let description = "Represents 2-element array with high and low mask";
24+ let extraClassDeclaration = [{
25+ IntegerType getValueType() { return IntegerType::get(getContext(), 64U); }
26+ }];
27+}
28+ 
22def PyStruct : EmitAsc_Type<"PyStruct", "py_struct", [MemRefElementTypeInterface]> {29def PyStruct : EmitAsc_Type<"PyStruct", "py_struct", [MemRefElementTypeInterface]> {
23 let parameters = (ins "StringAttr":$nameAttr, "ArrayAttr":$typesAttr, "ArrayAttr":$namesAttr);30 let parameters = (ins "StringAttr":$nameAttr, "ArrayAttr":$typesAttr, "ArrayAttr":$namesAttr);
24 let assemblyFormat = [{ `<` $nameAttr `,` $typesAttr `,` $namesAttr `>` }];31 let assemblyFormat = [{ `<` $nameAttr `,` $typesAttr `,` $namesAttr `>` }];
@@ -0,0 +1,52 @@
1+/*
2+ * Copyright (c) 2026 Huawei Technologies Co., Ltd.
3+ * This program is free software, you can redistribute it and/or modify it under the terms and conditions of
4+ * 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,
7+ * 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+#ifndef ASCIR_DIALECT_EMITASC_UTILS_INITSTRUCTBUILDER_H
12+#define ASCIR_DIALECT_EMITASC_UTILS_INITSTRUCTBUILDER_H
13+ 
14+#include "ascir/Dialect/EmitAsc/IR/EmitAsc.h"
15+ 
16+#include "mlir/IR/Builders.h"
17+ 
18+namespace mlir {
19+namespace emitasc {
20+ 
21+class InitStructBuilder {
22+ SmallVector<StringRef> fieldNames;
23+ SmallVector<Value> fieldValues;
24+ 
25+public:
26+ Type result;
27+ 
28+ InitStructBuilder(Type result) : result(result) {}
29+ InitStructBuilder(const InitStructBuilder&) = default;
30+ InitStructBuilder(InitStructBuilder&&) = default;
31+ ~InitStructBuilder() = default;
32+ 
33+ ArrayRef<StringRef> names() const { return fieldNames; }
34+ ValueRange values() const { return fieldValues; }
35+ 
36+ InitStructBuilder& addField(StringRef name, Value value)
37+ {
38+ fieldNames.push_back(name);
39+ fieldValues.push_back(value);
40+ return *this;
41+ }
42+ 
43+ InitStructOp create(OpBuilder& builder, Location loc) const
44+ {
45+ return builder.create<InitStructOp>(loc, result, builder.getStrArrayAttr(names()), values());
46+ }
47+};
48+ 
49+} // namespace emitasc
50+} // namespace mlir
51+ 
52+#endif // ASCIR_DIALECT_EMITASC_UTILS_INITSTRUCTBUILDER_H
@@ -0,0 +1,108 @@
1+/*
2+ * Copyright (c) 2026 Huawei Technologies Co., Ltd.
3+ * This program is free software, you can redistribute it and/or modify it under the terms and conditions of
4+ * 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,
7+ * 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+#ifndef ASCIR_DIALECT_UTILS_CVGROUPCANONICALIZATION_H
12+#define ASCIR_DIALECT_UTILS_CVGROUPCANONICALIZATION_H
13+ 
14+#include "mlir/IR/PatternMatch.h"
15+#include "llvm/ADT/SmallVector.h"
16+#include "llvm/ADT/STLExtras.h"
17+ 
18+namespace mlir {
19+namespace ascir {
20+ 
21+template <typename CVGroupOp, typename YieldOp>
22+struct EraseEmptyGroup : public OpRewritePattern<CVGroupOp> {
23+ using OpRewritePattern<CVGroupOp>::OpRewritePattern;
24+ 
25+ LogicalResult matchAndRewrite(CVGroupOp op, PatternRewriter& rewriter) const override
26+ {
27+ Block* body = op.getBody();
28+ if (!body->without_terminator().empty())
29+ return failure();
30+ auto yieldOperands = cast<YieldOp>(body->getTerminator()).getOperands();
31+ rewriter.replaceOp(op, yieldOperands);
32+ return success();
33+ };
34+};
35+ 
36+template <typename CVGroupOp, typename YieldOp>
37+struct EraseUnusedOperands : public OpRewritePattern<CVGroupOp> {
38+ using OpRewritePattern<CVGroupOp>::OpRewritePattern;
39+ 
40+ LogicalResult matchAndRewrite(CVGroupOp op, PatternRewriter& rewriter) const override
41+ {
42+ BitVector unusedOperands(op.getNumOperands());
43+ Block* body = op.getBody();
44+ for (unsigned i = 0; i < op.getNumOperands(); i++) {
45+ auto userInsideGroup = [op](Operation* user) { return op->isProperAncestor(user) && !isa<YieldOp>(user); };
46+ if (llvm::none_of(op.getOperand(i).getUsers(), userInsideGroup))
47+ unusedOperands.set(i);
48+ }
49+ if (unusedOperands.none())
50+ return failure();
51+ rewriter.modifyOpInPlace(op, [&] { op->eraseOperands(unusedOperands); });
52+ return success();
53+ };
54+};
55+ 
56+template <typename CVGroupOp, typename YieldOp>
57+struct EraseUnusedResults : public OpRewritePattern<CVGroupOp> {
58+ enum struct ResultKind { Used, Forwarded, Unused };
59+ 
60+ using OpRewritePattern<CVGroupOp>::OpRewritePattern;
61+ 
62+ LogicalResult matchAndRewrite(CVGroupOp op, PatternRewriter& rewriter) const override
63+ {
64+ unsigned numResults = op.getNumResults();
65+ if (numResults == 0)
66+ return failure();
67+ SmallVector<ResultKind, 4> results(numResults, ResultKind::Used);
68+ auto yieldOp = cast<YieldOp>(op.getBody()->getTerminator());
69+ for (unsigned i = 0; i < numResults; i++) {
70+ auto result = op.getResult(i);
71+ if (result.use_empty()) {
72+ results[i] = ResultKind::Unused;
73+ } else {
74+ auto* forwardDef = yieldOp.getOperand(i).getDefiningOp();
75+ if (!forwardDef || !op->isProperAncestor(forwardDef))
76+ results[i] = ResultKind::Forwarded;
77+ }
78+ }
79+ if (results.front() == ResultKind::Used && llvm::all_equal(results))
80+ return failure();
81+ SmallVector<Type, 4> newTypes;
82+ for (auto [kind, type] : llvm::zip_equal(results, op.getResultTypes()))
83+ if (kind == ResultKind::Used)
84+ newTypes.push_back(type);
85+ auto newOp = rewriter.create<CVGroupOp>(op.getLoc(), newTypes, op.getOperands());
86+ rewriter.inlineRegionBefore(op.getRegion(), newOp.getRegion(), newOp.getRegion().end());
87+ SmallVector<Value, 4> newYields, newResults;
88+ unsigned resultIdx = 0;
89+ for (auto [kind, result, yield] : llvm::zip_equal(results, op.getResults(), yieldOp.getOperands())) {
90+ if (kind == ResultKind::Used) {
91+ newYields.push_back(yield);
92+ newResults.push_back(newOp.getResult(resultIdx++));
93+ } else if (kind == ResultKind::Forwarded) {
94+ newResults.push_back(yield);
95+ } else if (kind == ResultKind::Unused) {
96+ newResults.push_back(Value{});
97+ }
98+ }
99+ rewriter.modifyOpInPlace(yieldOp, [&] { yieldOp->setOperands(newYields); });
100+ rewriter.replaceOp(op, newResults);
101+ return success();
102+ }
103+};
104+ 
105+} // namespace ascir
106+} // namespace mlir
107+ 
108+#endif // ASCIR_DIALECT_UTILS_CVGROUPCANONICALIZATION_H
@@ -39,11 +39,13 @@ struct ConstantOpBuilder {
39 39 
40 Value i64(int64_t value) { return create(builder.getI64Type(), value); }40 Value i64(int64_t value) { return create(builder.getI64Type(), value); }
41 41 
42- Value i32(int32_t value) { return create(builder.getI32Type(), value); }42+ Value i32(int64_t value) { return create(builder.getI32Type(), value); }
43 43 
44- Value i16(int16_t value) { return create(builder.getI16Type(), value); }44+ Value i16(int64_t value) { return create(builder.getI16Type(), value); }
45 45 
46- Value i8(int8_t value) { return create(builder.getI8Type(), value); }46+ Value i8(int64_t value) { return create(builder.getI8Type(), value); }
47+ 
48+ Value i1(bool value) { return create(builder.getI1Type(), value); }
47 49 
48 Value f64(double value) { return create(builder.getF64Type(), value); }50 Value f64(double value) { return create(builder.getF64Type(), value); }
49 51 
@@ -0,0 +1,24 @@
1+/*
2+ * Copyright (c) 2026 Huawei Technologies Co., Ltd.
3+ * This program is free software, you can redistribute it and/or modify it under the terms and conditions of
4+ * 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,
7+ * 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+#ifndef ASCIR_TARGET_ASC_ADV_BROADCAST_H
12+#define ASCIR_TARGET_ASC_ADV_BROADCAST_H
13+ 
14+#include "ascir/Target/Asc/Common.h"
15+ 
16+namespace mlir {
17+namespace ascendc {
18+ 
19+LogicalResult printOperation(CodeEmitter& emitter, ascendc::BroadcastOp op);
20+ 
21+} // namespace ascendc
22+} // namespace mlir
23+ 
24+#endif // ASCIR_TARGET_ASC_ADV_BROADCAST_H
@@ -24,8 +24,9 @@ template <typename UnaryMathOp>
24auto printOperation(CodeEmitter& emitter, UnaryMathOp op) -> LogicalResultForT<24auto printOperation(CodeEmitter& emitter, UnaryMathOp op) -> LogicalResultForT<
25 UnaryMathOp, ascendc::AcoshOp, ascendc::AcosOp, ascendc::AsinhOp, ascendc::AsinOp, ascendc::AtanhOp,25 UnaryMathOp, ascendc::AcoshOp, ascendc::AcosOp, ascendc::AsinhOp, ascendc::AsinOp, ascendc::AtanhOp,
26 ascendc::AtanOp, ascendc::CeilOp, ascendc::CoshOp, ascendc::CosOp, ascendc::DigammaOp, ascendc::ErfcOp,26 ascendc::AtanOp, ascendc::CeilOp, ascendc::CoshOp, ascendc::CosOp, ascendc::DigammaOp, ascendc::ErfcOp,
27- ascendc::ErfOp, ascendc::FloorOp, ascendc::FracOp, ascendc::LgammaOp, ascendc::LogOp, ascendc::RoundOp,27+ ascendc::ErfOp, ascendc::FloorOp, ascendc::FracOp, ascendc::LgammaOp, ascendc::LogOp, ascendc::Log2Op,
28- ascendc::SignOp, ascendc::SinhOp, ascendc::SinOp, ascendc::TanhOp, ascendc::TanOp, ascendc::TruncOp>28+ ascendc::RoundOp, ascendc::SignOp, ascendc::SinhOp, ascendc::SinOp, ascendc::TanhOp, ascendc::TanOp,
29+ ascendc::TruncOp>
29{30{
30 auto& os = emitter.ostream();31 auto& os = emitter.ostream();
31 os << ascNamespace << "::" << op.getAPIName();32 os << ascNamespace << "::" << op.getAPIName();
@@ -17,6 +17,7 @@ namespace mlir {
17namespace ascendc {17namespace ascendc {
18 18 
19LogicalResult printOperation(CodeEmitter& emitter, ascendc::RmsNormOp op);19LogicalResult printOperation(CodeEmitter& emitter, ascendc::RmsNormOp op);
20+LogicalResult printOperation(CodeEmitter& emitter, ascendc::LayerNormOp op);
20 21 
21} // namespace ascendc22} // namespace ascendc
22} // namespace mlir23} // namespace mlir
@@ -0,0 +1,24 @@
1+/*
2+ * Copyright (c) 2026 Huawei Technologies Co., Ltd.
3+ * This program is free software, you can redistribute it and/or modify it under the terms and conditions of
4+ * 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,
7+ * 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+#ifndef ASCIR_TARGET_ASC_ADV_REDUCTION_H
12+#define ASCIR_TARGET_ASC_ADV_REDUCTION_H
13+ 
14+#include "ascir/Target/Asc/Common.h"
15+ 
16+namespace mlir {
17+namespace ascendc {
18+ 
19+LogicalResult printOperation(CodeEmitter& emitter, ascendc::ReduceOp op);
20+ 
21+} // namespace ascendc
22+} // namespace mlir
23+ 
24+#endif // ASCIR_TARGET_ASC_ADV_REDUCTION_H
@@ -28,6 +28,10 @@ LogicalResult printOperation(CodeEmitter& emitter, ascendc::CrossCoreSetFlagOp o
28 28 
29LogicalResult printOperation(CodeEmitter& emitter, ascendc::CrossCoreWaitFlagOp op);29LogicalResult printOperation(CodeEmitter& emitter, ascendc::CrossCoreWaitFlagOp op);
30 30 
31+LogicalResult printOperation(CodeEmitter& emitter, ascendc::GetBufOp op);
32+ 
33+LogicalResult printOperation(CodeEmitter& emitter, ascendc::RlsBufOp op);
34+ 
31} // namespace ascendc35} // namespace ascendc
32} // namespace mlir36} // namespace mlir
33 37 
@@ -26,6 +26,8 @@ LogicalResult printOperation(CodeEmitter& emitter, ascendc::TransDataTo5HDUintLi
26 26 
27LogicalResult printOperation(CodeEmitter& emitter, ascendc::TransDataTo5HDOp op);27LogicalResult printOperation(CodeEmitter& emitter, ascendc::TransDataTo5HDOp op);
28 28 
29+LogicalResult printOperation(CodeEmitter& emitter, ascendc::TransDataTo5HDTensorOp op);
30+ 
29} // namespace ascendc31} // namespace ascendc
30} // namespace mlir32} // namespace mlir
31 33 
@@ -26,6 +26,8 @@ LogicalResult printOperation(CodeEmitter& emitter, ascendc::CopyL0Op op);
26 26 
27LogicalResult printOperation(CodeEmitter& emitter, ascendc::CopyL1Op op);27LogicalResult printOperation(CodeEmitter& emitter, ascendc::CopyL1Op op);
28 28 
29+LogicalResult printOperation(CodeEmitter& emitter, ascendc::NdDmaParamsOp op);
30+ 
29} // namespace ascendc31} // namespace ascendc
30} // namespace mlir32} // namespace mlir
31 33 
@@ -56,21 +56,24 @@ LogicalResult printOperation(CodeEmitter& emitter, ascendc::FixpipeWithWorkspace
56 56 
57LogicalResult printOperation(CodeEmitter& emitter, ascendc::GetStoreAtomicConfigOp op);57LogicalResult printOperation(CodeEmitter& emitter, ascendc::GetStoreAtomicConfigOp op);
58 58 
59+LogicalResult printOperation(CodeEmitter& emitter, ascendc::IfAICOp op);
60+ 
61+LogicalResult printOperation(CodeEmitter& emitter, ascendc::IfAIVOp op);
62+ 
63+LogicalResult printOperation(CodeEmitter& emitter, ascendc::YieldOp op);
64+ 
59template <typename FixpipeOp>65template <typename FixpipeOp>
60auto printFixpipeTemplate(CodeEmitter& emitter, FixpipeOp op)66auto printFixpipeTemplate(CodeEmitter& emitter, FixpipeOp op)
61{67{
62 auto& os = emitter.ostream();68 auto& os = emitter.ostream();
63- auto dstType = cast<ascendc::GlobalTensorType>(op.getDst().getType()).getElementType();69+ auto dstType = getElementTypeOrSelf(op.getDst());
64 auto srcType = cast<ascendc::LocalTensorType>(op.getSrc().getType()).getElementType();70 auto srcType = cast<ascendc::LocalTensorType>(op.getSrc().getType()).getElementType();
65 os << ascNamespace << "::" << op.getAPIName() << "<";71 os << ascNamespace << "::" << op.getAPIName() << "<";
66 FAIL_OR(emitter.emitType(op.getLoc(), dstType));72 FAIL_OR(emitter.emitType(op.getLoc(), dstType));
67 os << ", ";73 os << ", ";
68 FAIL_OR(emitter.emitType(op.getLoc(), srcType));74 FAIL_OR(emitter.emitType(op.getLoc(), srcType));
69 os << ", ";75 os << ", ";
70- auto constructOp = cast<ascendc::ConstructOp>(op.getFixpipeConfig().getDefiningOp());76+ os << emitter.getOrCreateName(op.getFixpipeConfig());
71- auto constOp = cast<arith::ConstantOp>(constructOp->getOperand(0).getDefiningOp());
72- int64_t value = cast<IntegerAttr>(constOp.getValue()).getInt();
73- os << ascNamespace << "::" << (value == 0 ? "CFG_NZ" : "CFG_ROW_MAJOR");
74 os << ">";77 os << ">";
75 return success();78 return success();
76}79}
@@ -0,0 +1,96 @@
1+/*
2+ * Copyright (c) 2026 Huawei Technologies Co., Ltd.
3+ * This program is free software, you can redistribute it and/or modify it under the terms and conditions of
4+ * 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,
7+ * 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+#ifndef ASCIR_TARGET_ASC_BASIC_REG_H
12+#define ASCIR_TARGET_ASC_BASIC_REG_H
13+ 
14+#include "ascir/Target/Asc/Common.h"
15+ 
16+namespace mlir {
17+namespace ascendc {
18+ 
19+//===----------------------------------------------------------------------===//
20+// Binary register API operations
21+//===----------------------------------------------------------------------===//
22+ 
23+template <typename BinaryRegOp>
24+LogicalResultForT<
25+ BinaryRegOp, ascendc::AddRegOp, ascendc::AndRegOp, ascendc::DivRegOp, ascendc::FusedAbsSubRegOp,
26+ ascendc::FusedExpSubRegOp, ascendc::FusedMulDstAddRegOp, ascendc::SubRegOp, ascendc::MaxRegOp, ascendc::MinRegOp,
27+ ascendc::MulRegOp, ascendc::MulAddDstRegOp, ascendc::OrRegOp, ascendc::PreluRegOp, ascendc::XorRegOp>
28+printOperation(CodeEmitter& emitter, BinaryRegOp op)
29+{
30+ auto& os = emitter.ostream();
31+ os << ascNamespace << "::" << op.getAPIName() << "(" << emitter.getOrCreateName(op.getDstReg()) << ", "
32+ << emitter.getOrCreateName(op.getSrc0Reg()) << ", " << emitter.getOrCreateName(op.getSrc1Reg()) << ", "
33+ << emitter.getOrCreateName(op.getMaskReg()) << ")";
34+ return success();
35+}
36+ 
37+//===----------------------------------------------------------------------===//
38+// Unary, Reduction, Duplicate register API operations
39+//===----------------------------------------------------------------------===//
40+ 
41+template <typename RegOp>
42+LogicalResultForT<
43+ RegOp, ascendc::AbsRegOp, ascendc::ExpRegOp, ascendc::LnRegOp, ascendc::LogRegOp, ascendc::Log10RegOp,
44+ ascendc::MaskNotRegOp, ascendc::NegRegOp, ascendc::NotRegOp, ascendc::ReluRegOp, ascendc::SqrtRegOp,
45+ ascendc::ReduceMaxRegOp, ascendc::ReduceMinRegOp, ascendc::ReduceSumRegOp, ascendc::DuplicateRegOp>
46+printOperation(CodeEmitter& emitter, RegOp op)
47+{
48+ auto& os = emitter.ostream();
49+ os << ascNamespace << "::" << op.getAPIName() << "(" << emitter.getOrCreateName(op.getDstReg()) << ", "
50+ << emitter.getOrCreateName(op.getSrcReg()) << ", " << emitter.getOrCreateName(op.getMaskReg()) << ")";
51+ return success();
52+}
53+ 
54+//===----------------------------------------------------------------------===//
55+// VecScalar register API operations
56+//===----------------------------------------------------------------------===//
57+ 
58+template <typename VecScalarOp>
59+LogicalResultForT<
60+ VecScalarOp, ascendc::AddsRegOp, ascendc::MulsRegOp, ascendc::MaxsRegOp, ascendc::MinsRegOp,
61+ ascendc::LeakyReluRegOp, ascendc::ShiftLeftsRegOp, ascendc::ShiftRightsRegOp>
62+printOperation(CodeEmitter& emitter, VecScalarOp op)
63+{
64+ auto& os = emitter.ostream();
65+ os << ascNamespace << "::" << op.getAPIName() << "(" << emitter.getOrCreateName(op.getDstReg()) << ", "
66+ << emitter.getOrCreateName(op.getSrcReg()) << ", " << emitter.getOrCreateName(op.getScalar()) << ", "
67+ << emitter.getOrCreateName(op.getMaskReg()) << ")";
68+ return success();
69+}
70+ 
71+//===----------------------------------------------------------------------===//
72+// Other register API operations
73+//===----------------------------------------------------------------------===//
74+ 
75+LogicalResult printOperation(CodeEmitter& emitter, ascendc::DataCopyLoadOp op);
76+ 
77+LogicalResult printOperation(CodeEmitter& emitter, ascendc::DataCopyStoreOp op);
78+ 
79+LogicalResult printOperation(CodeEmitter& emitter, ascendc::UpdateMaskOp op);
80+ 
81+LogicalResult printOperation(CodeEmitter& emitter, ascendc::RegTensorOp op);
82+ 
83+LogicalResult printOperation(CodeEmitter& emitter, ascendc::DuplicateScalarRegOp op);
84+ 
85+LogicalResult printOperation(CodeEmitter& emitter, ascendc::GetVecLenOp op);
86+ 
87+LogicalResult printOperation(CodeEmitter& emitter, ascendc::LocalMemBarOp op);
88+ 
89+LogicalResult printOperation(CodeEmitter& emitter, ascendc::SelectRegOp op);
90+ 
91+LogicalResult printOperation(CodeEmitter& emitter, ascendc::CreateMaskOp op);
92+ 
93+} // namespace ascendc
94+} // namespace mlir
95+ 
96+#endif // ASCIR_TARGET_ASC_BASIC_REG_H
@@ -46,25 +46,6 @@ auto printBinaryL2Params(CodeEmitter& emitter, BinaryOp op)
46 << emitter.getOrCreateName(op.getSrc1()) << ", " << emitter.getOrCreateName(op.getCalCount()) << ")";46 << emitter.getOrCreateName(op.getSrc1()) << ", " << emitter.getOrCreateName(op.getCalCount()) << ")";
47}47}
48 48 
49-template <typename BinaryL0Op>
50-auto printOperation(CodeEmitter& emitter, BinaryL0Op op) -> LogicalResultForT<BinaryL0Op, ascendc::MulCastL0Op>
51-{
52- auto& os = emitter.ostream();
53- os << ascNamespace << "::" << op.getAPIName();
54- printBinaryL0Params(emitter, op);
55- return success();
56-}
57- 
58-template <typename BinaryL1Op>
59-auto printOperation(CodeEmitter& emitter, BinaryL1Op op) -> LogicalResultForT<BinaryL1Op, ascendc::MulCastL1Op>
60-{
61- auto& os = emitter.ostream();
62- auto maskName = printMask(emitter, op);
63- os << ascNamespace << "::" << op.getAPIName();
64- printBinaryL1Params(emitter, op, maskName);
65- return success();
66-}
67- 
68template <typename BinaryL2Op>49template <typename BinaryL2Op>
69auto printOperation(CodeEmitter& emitter, BinaryL2Op op) -> LogicalResultForT<50auto printOperation(CodeEmitter& emitter, BinaryL2Op op) -> LogicalResultForT<
70 BinaryL2Op, ascendc::AddL2Op, ascendc::AddDeqReluL2Op, ascendc::AddReluL2Op, ascendc::AddReluCastL2Op,51 BinaryL2Op, ascendc::AddL2Op, ascendc::AddDeqReluL2Op, ascendc::AddReluL2Op, ascendc::AddReluCastL2Op,
@@ -106,7 +87,8 @@ auto printOperation(CodeEmitter& emitter, BinaryTemplateL1Op op) -> LogicalResul
106 87 
107template <typename BinaryCastL0Op>88template <typename BinaryCastL0Op>
108auto printOperation(CodeEmitter& emitter, BinaryCastL0Op op) -> LogicalResultForT<89auto printOperation(CodeEmitter& emitter, BinaryCastL0Op op) -> LogicalResultForT<
109- BinaryCastL0Op, ascendc::AddDeqReluL0Op, ascendc::AddReluCastL0Op, ascendc::SubReluCastL0Op, ascendc::MulAddDstL0Op>90+ BinaryCastL0Op, ascendc::AddDeqReluL0Op, ascendc::AddReluCastL0Op, ascendc::SubReluCastL0Op, ascendc::MulAddDstL0Op,
91+ ascendc::MulCastL0Op>
110{92{
111 auto& os = emitter.ostream();93 auto& os = emitter.ostream();
112 FAIL_OR(printIsSetMaskCastTemplate(emitter, op));94 FAIL_OR(printIsSetMaskCastTemplate(emitter, op));
@@ -116,7 +98,8 @@ auto printOperation(CodeEmitter& emitter, BinaryCastL0Op op) -> LogicalResultFor
116 98 
117template <typename BinaryCastL1Op>99template <typename BinaryCastL1Op>
118auto printOperation(CodeEmitter& emitter, BinaryCastL1Op op) -> LogicalResultForT<100auto printOperation(CodeEmitter& emitter, BinaryCastL1Op op) -> LogicalResultForT<
119- BinaryCastL1Op, ascendc::AddDeqReluL1Op, ascendc::AddReluCastL1Op, ascendc::SubReluCastL1Op, ascendc::MulAddDstL1Op>101+ BinaryCastL1Op, ascendc::AddDeqReluL1Op, ascendc::AddReluCastL1Op, ascendc::SubReluCastL1Op, ascendc::MulAddDstL1Op,
102+ ascendc::MulCastL1Op>
120{103{
121 auto& os = emitter.ostream();104 auto& os = emitter.ostream();
122 auto maskName = printMask(emitter, op);105 auto maskName = printMask(emitter, op);
@@ -22,8 +22,8 @@ namespace ascendc {
22 22 
23template <typename VecScalarL0Op>23template <typename VecScalarL0Op>
24auto printOperation(CodeEmitter& emitter, VecScalarL0Op op) -> LogicalResultForT<24auto printOperation(CodeEmitter& emitter, VecScalarL0Op op) -> LogicalResultForT<
25- VecScalarL0Op, ascendc::AddsL0Op, ascendc::LeakyReluL0Op, ascendc::MaxsL0Op, ascendc::MinsL0Op, ascendc::MulsL0Op,25+ VecScalarL0Op, ascendc::AddsL0Op, ascendc::SubsL0Op, ascendc::LeakyReluL0Op, ascendc::MaxsL0Op, ascendc::MinsL0Op,
26- ascendc::ShiftLeftL0Op, ascendc::ShiftRightL0Op>26+ ascendc::MulsL0Op, ascendc::DivsL0Op, ascendc::ShiftLeftL0Op, ascendc::ShiftRightL0Op>
27{27{
28 auto& os = emitter.ostream();28 auto& os = emitter.ostream();
29 FAIL_OR(printIsSetMaskTemplate(emitter, op));29 FAIL_OR(printIsSetMaskTemplate(emitter, op));
@@ -35,8 +35,8 @@ auto printOperation(CodeEmitter& emitter, VecScalarL0Op op) -> LogicalResultForT
35 35 
36template <typename VecScalarL1Op>36template <typename VecScalarL1Op>
37auto printOperation(CodeEmitter& emitter, VecScalarL1Op op) -> LogicalResultForT<37auto printOperation(CodeEmitter& emitter, VecScalarL1Op op) -> LogicalResultForT<
38- VecScalarL1Op, ascendc::AddsL1Op, ascendc::LeakyReluL1Op, ascendc::MaxsL1Op, ascendc::MinsL1Op, ascendc::MulsL1Op,38+ VecScalarL1Op, ascendc::AddsL1Op, ascendc::SubsL1Op, ascendc::LeakyReluL1Op, ascendc::MaxsL1Op, ascendc::MinsL1Op,
39- ascendc::ShiftLeftL1Op, ascendc::ShiftRightL1Op>39+ ascendc::MulsL1Op, ascendc::DivsL1Op, ascendc::ShiftLeftL1Op, ascendc::ShiftRightL1Op>
40{40{
41 auto& os = emitter.ostream();41 auto& os = emitter.ostream();
42 auto maskName = printMask(emitter, op);42 auto maskName = printMask(emitter, op);
@@ -49,8 +49,8 @@ auto printOperation(CodeEmitter& emitter, VecScalarL1Op op) -> LogicalResultForT
49 49 
50template <typename VecScalarL2Op>50template <typename VecScalarL2Op>
51auto printOperation(CodeEmitter& emitter, VecScalarL2Op op) -> LogicalResultForT<51auto printOperation(CodeEmitter& emitter, VecScalarL2Op op) -> LogicalResultForT<
52- VecScalarL2Op, ascendc::AddsL2Op, ascendc::LeakyReluL2Op, ascendc::MaxsL2Op, ascendc::MinsL2Op, ascendc::MulsL2Op,52+ VecScalarL2Op, ascendc::AddsL2Op, ascendc::SubsL2Op, ascendc::LeakyReluL2Op, ascendc::MaxsL2Op, ascendc::MinsL2Op,
53- ascendc::ShiftLeftL2Op, ascendc::ShiftRightL2Op>53+ ascendc::MulsL2Op, ascendc::DivsL2Op, ascendc::ShiftLeftL2Op, ascendc::ShiftRightL2Op>
54{54{
55 auto& os = emitter.ostream();55 auto& os = emitter.ostream();
56 FAIL_OR(printIsSetMaskTemplate(emitter, op));56 FAIL_OR(printIsSetMaskTemplate(emitter, op));
@@ -31,6 +31,8 @@ struct CodeEmitter {
31 31 
32 static void emitTPosition(raw_ostream& os, ascendc::TPosition pos);32 static void emitTPosition(raw_ostream& os, ascendc::TPosition pos);
33 33 
34+ static void emitCO2Layout(raw_ostream& os, ascendc::CO2Layout layout);
35+ 
34 static void emitLayoutMode(raw_ostream& os, ascendc::LayoutMode layout);36 static void emitLayoutMode(raw_ostream& os, ascendc::LayoutMode layout);
35 37 
36 explicit CodeEmitter(raw_ostream& os);38 explicit CodeEmitter(raw_ostream& os);
@@ -165,6 +167,12 @@ private:
165 167 
166 LogicalResult emitAscLocalTensorType(Location loc, Type type, bool emitAsUnsigned);168 LogicalResult emitAscLocalTensorType(Location loc, Type type, bool emitAsUnsigned);
167 169 
170+ LogicalResult emitAscFixpipeParamsC310Type(Location loc, Type type, bool emitAsUnsigned);
171+ 
172+ LogicalResult emitAscRegTensorType(Location loc, Type type, bool emitAsUnsigned);
173+ 
174+ LogicalResult emitAscMaskRegType(Location loc, Type type, bool emitAsUnsigned);
175+ 
168 LogicalResult emitAscLocalMemAllocatorType(Location loc, Type type, bool emitAsUnsigned);176 LogicalResult emitAscLocalMemAllocatorType(Location loc, Type type, bool emitAsUnsigned);
169 177 
170 LogicalResult emitAscPyStructType(Location loc, Type type, bool emitAsUnsigned);178 LogicalResult emitAscPyStructType(Location loc, Type type, bool emitAsUnsigned);
@@ -189,6 +197,8 @@ private:
189 197 
190 LogicalResult emitTypeAttr(Location loc, Attribute attr);198 LogicalResult emitTypeAttr(Location loc, Attribute attr);
191 199 
200+ LogicalResult emitAscNdDmaParams(Location loc, Type type, bool emitAsUnsigned);
201+ 
192 void printInt(const APInt& value, bool isUnsigned);202 void printInt(const APInt& value, bool isUnsigned);
193 203 
194 void printFloat(const APFloat& value);204 void printFloat(const APFloat& value);
@@ -22,12 +22,16 @@ namespace ascendc {
22 22 
23LogicalResult printOperation(CodeEmitter& emitter, ascendc::LocalTensorV2Op op);23LogicalResult printOperation(CodeEmitter& emitter, ascendc::LocalTensorV2Op op);
24 24 
25+LogicalResult printOperation(CodeEmitter& emitter, ascendc::LocalTensorV3Op op);
26+ 
25LogicalResult printOperation(CodeEmitter& emitter, ascendc::LocalTensorReinterpretCastOp op);27LogicalResult printOperation(CodeEmitter& emitter, ascendc::LocalTensorReinterpretCastOp op);
26 28 
27LogicalResult printOperation(CodeEmitter& emitter, ascendc::LocalTensorSubIndexOp op);29LogicalResult printOperation(CodeEmitter& emitter, ascendc::LocalTensorSubIndexOp op);
28 30 
29LogicalResult printOperation(CodeEmitter& emitter, ascendc::LocalTensorBracketOp op);31LogicalResult printOperation(CodeEmitter& emitter, ascendc::LocalTensorBracketOp op);
30 32 
33+LogicalResult printOperation(CodeEmitter& emitter, ascendc::LocalTensorGetPhyAddrV2Op op);
34+ 
31} // namespace ascendc35} // namespace ascendc
32} // namespace mlir36} // namespace mlir
33 37 
@@ -17,10 +17,6 @@
17namespace mlir {17namespace mlir {
18namespace emitasc {18namespace emitasc {
19 19 
20-//===----------------------------------------------------------------------===//
21-// EmitAsc operations
22-//===----------------------------------------------------------------------===//
23- 
24LogicalResult printOperation(CodeEmitter& emitter, emitasc::CallOpaqueOp op);20LogicalResult printOperation(CodeEmitter& emitter, emitasc::CallOpaqueOp op);
25 21 
26LogicalResult printOperation(CodeEmitter& emitter, emitasc::CopyStructOp op);22LogicalResult printOperation(CodeEmitter& emitter, emitasc::CopyStructOp op);
@@ -29,6 +25,10 @@ LogicalResult printOperation(CodeEmitter& emitter, emitasc::DeclarePyStructOp op
29 25 
30LogicalResult printOperation(CodeEmitter& emitter, emitasc::DereferenceOp op);26LogicalResult printOperation(CodeEmitter& emitter, emitasc::DereferenceOp op);
31 27 
28+LogicalResult printOperation(CodeEmitter& emitter, emitasc::InitStructOp op);
29+ 
30+LogicalResult printOperation(CodeEmitter& emitter, emitasc::MaskOp op);
31+ 
32LogicalResult printOperation(CodeEmitter& emitter, emitasc::MemberOp op);32LogicalResult printOperation(CodeEmitter& emitter, emitasc::MemberOp op);
33 33 
34LogicalResult printOperation(CodeEmitter& emitter, emitasc::MemberPtrOp op);34LogicalResult printOperation(CodeEmitter& emitter, emitasc::MemberPtrOp op);
@@ -100,6 +100,8 @@ LogicalResult printOperation(CodeEmitter& emitter, arith::SelectOp op);
100 100 
101LogicalResult printOperation(CodeEmitter& emitter, arith::IndexCastOp op);101LogicalResult printOperation(CodeEmitter& emitter, arith::IndexCastOp op);
102 102 
103+LogicalResult printOperation(CodeEmitter& emitter, arith::NegFOp op);
104+ 
103} // namespace mlir105} // namespace mlir
104 106 
105#endif // ASCIR_TARGET_ASC_MLIR_ARITH_H107#endif // ASCIR_TARGET_ASC_MLIR_ARITH_H
@@ -7,5 +7,4 @@
7# See LICENSE in the root of the software repository for the full text of the License.7# See LICENSE in the root of the software repository for the full text of the License.
8 8 
9add_subdirectory(Dialect)9add_subdirectory(Dialect)
10-# add_subdirectory(TableGen)
11add_subdirectory(Target)10add_subdirectory(Target)
@@ -9,10 +9,11 @@
9 */9 */
10 10 
11#include "ascir/Dialect/Asc/IR/Asc.h"11#include "ascir/Dialect/Asc/IR/Asc.h"
12+#include "ascir/Dialect/Utils/CVGroupCanonicalization.h"
12 13 
13#include "mlir/IR/Builders.h"14#include "mlir/IR/Builders.h"
15+#include "mlir/IR/Matchers.h"
14#include "mlir/IR/OpImplementation.h"16#include "mlir/IR/OpImplementation.h"
15-#include "mlir/IR/PatternMatch.h"
16 17 
17#define GET_OP_CLASSES18#define GET_OP_CLASSES
18#include "ascir/Dialect/Asc/IR/AscendCOps.cpp.inc"19#include "ascir/Dialect/Asc/IR/AscendCOps.cpp.inc"
@@ -42,6 +43,39 @@ LogicalResult GlobalTensorOp::canonicalize(GlobalTensorOp op, PatternRewriter& r
42 return eraseUnusedOp(op, rewriter);43 return eraseUnusedOp(op, rewriter);
43}44}
44 45 
46+//===----------------------------------------------------------------------===//
47+// IfAICOp
48+//===----------------------------------------------------------------------===//
49+ 
50+void IfAICOp::getCanonicalizationPatterns(RewritePatternSet& results, MLIRContext* context)
51+{
52+ results.add<
53+ ascir::EraseEmptyGroup<IfAICOp, YieldOp>, ascir::EraseUnusedOperands<IfAICOp, YieldOp>,
54+ ascir::EraseUnusedResults<IfAICOp, YieldOp>>(context);
55+}
56+ 
57+//===----------------------------------------------------------------------===//
58+// IfAIVOp
59+//===----------------------------------------------------------------------===//
60+ 
61+void IfAIVOp::getCanonicalizationPatterns(RewritePatternSet& results, MLIRContext* context)
62+{
63+ results.add<
64+ ascir::EraseEmptyGroup<IfAIVOp, YieldOp>, ascir::EraseUnusedOperands<IfAIVOp, YieldOp>,
65+ ascir::EraseUnusedResults<IfAIVOp, YieldOp>>(context);
66+}
67+ 
68+//===----------------------------------------------------------------------===//
69+// GlobalTensorSubIndexOp
70+//===----------------------------------------------------------------------===//
71+ 
72+OpFoldResult GlobalTensorSubIndexOp::fold([[maybe_unused]] FoldAdaptor adaptor)
73+{
74+ if (matchPattern(getIndex(), m_Zero()) && getType() == getTensor().getType())
75+ return getTensor();
76+ return {};
77+}
78+ 
45//===----------------------------------------------------------------------===//79//===----------------------------------------------------------------------===//
46// LocalTensorOp80// LocalTensorOp
47//===----------------------------------------------------------------------===//81//===----------------------------------------------------------------------===//
@@ -51,6 +85,17 @@ LogicalResult LocalTensorOp::canonicalize(LocalTensorOp op, PatternRewriter& rew
51 return eraseUnusedOp(op, rewriter);85 return eraseUnusedOp(op, rewriter);
52}86}
53 87 
88+//===----------------------------------------------------------------------===//
89+// LocalTensorSubIndexOp
90+//===----------------------------------------------------------------------===//
91+ 
92+OpFoldResult LocalTensorSubIndexOp::fold([[maybe_unused]] FoldAdaptor adaptor)
93+{
94+ if (matchPattern(getIndex(), m_Zero()) && getType() == getTensor().getType())
95+ return getTensor();
96+ return {};
97+}
98+ 
54//===----------------------------------------------------------------------===//99//===----------------------------------------------------------------------===//
55// PipeBarrierOp100// PipeBarrierOp
56//===----------------------------------------------------------------------===//101//===----------------------------------------------------------------------===//
@@ -75,7 +120,7 @@ LogicalResult PipeBarrierOp::canonicalize(PipeBarrierOp op, PatternRewriter& rew
75}120}
76 121 
77//===----------------------------------------------------------------------===//122//===----------------------------------------------------------------------===//
78-// ReinterpretCastOp123+// LocalTensorReinterpretCastOp
79//===----------------------------------------------------------------------===//124//===----------------------------------------------------------------------===//
80 125 
81bool LocalTensorReinterpretCastOp::areCastCompatible(TypeRange inputs, TypeRange outputs)126bool LocalTensorReinterpretCastOp::areCastCompatible(TypeRange inputs, TypeRange outputs)
@@ -87,7 +132,35 @@ bool LocalTensorReinterpretCastOp::areCastCompatible(TypeRange inputs, TypeRange
87OpFoldResult LocalTensorReinterpretCastOp::fold([[maybe_unused]] FoldAdaptor adaptor)132OpFoldResult LocalTensorReinterpretCastOp::fold([[maybe_unused]] FoldAdaptor adaptor)
88{133{
89 Value in = getIn();134 Value in = getIn();
90- return in.getType() == getType() ? in : nullptr;135+ Type resultType = getResult().getType();
136+ if (in.getType() == resultType)
137+ return in;
138+ if (auto defOp = in.getDefiningOp<LocalTensorReinterpretCastOp>()) {
139+ Value defIn = defOp.getIn();
140+ if (resultType == defIn.getType())
141+ return defIn;
142+ setOperand(defIn);
143+ return getResult();
144+ }
145+ return {};
146+}
147+ 
148+//===----------------------------------------------------------------------===//
149+// LocalTensorAutoOp
150+//===----------------------------------------------------------------------===//
151+ 
152+LogicalResult LocalTensorAutoOp::canonicalize(LocalTensorAutoOp op, PatternRewriter& rewriter)
153+{
154+ return eraseUnusedOp(op, rewriter);
155+}
156+ 
157+//===----------------------------------------------------------------------===//
158+// RegTensorOp
159+//===----------------------------------------------------------------------===//
160+ 
161+LogicalResult RegTensorOp::canonicalize(RegTensorOp op, PatternRewriter& rewriter)
162+{
163+ return eraseUnusedOp(op, rewriter);
91}164}
92 165 
93//===----------------------------------------------------------------------===//166//===----------------------------------------------------------------------===//
@@ -13,9 +13,9 @@ add_mlir_dialect_library(MLIRAscTransforms
13 EraseSync.cpp13 EraseSync.cpp
14 GenerateBoilerplatePass.cpp14 GenerateBoilerplatePass.cpp
15 HoistQueBind.cpp15 HoistQueBind.cpp
16- HoistUBAllocation.cpp16+ HoistTensorAllocation.cpp
17 InputOutputTensor.cpp17 InputOutputTensor.cpp
18- InsertSync.cpp18+ InsertQueSync.cpp
19 LegalizeKernelArgs.cpp19 LegalizeKernelArgs.cpp
20 MaterializeTensor.cpp20 MaterializeTensor.cpp
21 Noop.cpp21 Noop.cpp
@@ -28,7 +28,6 @@ namespace ascendc {
28} // namespace mlir28} // namespace mlir
29 29 
30using namespace mlir;30using namespace mlir;
31-using namespace mlir::ascendc;
32 31 
33namespace {32namespace {
34 33 
@@ -92,8 +91,4 @@ public:
92 91 
93} // namespace92} // namespace
94 93 
95-namespace mlir {94+std::unique_ptr<Pass> mlir::ascendc::createDeclarePyStructPass() { return std::make_unique<DeclarePyStructPass>(); }
96-namespace ascendc {
97-std::unique_ptr<Pass> createDeclarePyStructPass() { return std::make_unique<DeclarePyStructPass>(); }
98-} // namespace ascendc
99-} // namespace mlir
@@ -25,7 +25,6 @@ namespace ascendc {
25} // namespace mlir25} // namespace mlir
26 26 
27using namespace mlir;27using namespace mlir;
28-using namespace mlir::ascendc;
29 28 
30namespace {29namespace {
31 30 
@@ -42,8 +41,4 @@ public:
42 41 
43} // namespace42} // namespace
44 43 
45-namespace mlir {44+std::unique_ptr<Pass> mlir::ascendc::createDefineCubeOnlyPass() { return std::make_unique<DefineCubeOnlyPass>(); }
46-namespace ascendc {
47-std::unique_ptr<Pass> createDefineCubeOnlyPass() { return std::make_unique<DefineCubeOnlyPass>(); }
48-} // namespace ascendc
49-} // namespace mlir
@@ -24,7 +24,6 @@ namespace ascendc {
24} // namespace mlir24} // namespace mlir
25 25 
26using namespace mlir;26using namespace mlir;
27-using namespace mlir::ascendc;
28 27 
29namespace {28namespace {
30 29 
@@ -34,10 +33,10 @@ public:
34 {33 {
35 ModuleOp op = getOperation();34 ModuleOp op = getOperation();
36 if (op.walk([](ascendc::PrintfOp) { return WalkResult::interrupt(); }).wasInterrupted()) {35 if (op.walk([](ascendc::PrintfOp) { return WalkResult::interrupt(); }).wasInterrupted()) {
37- op->setAttr(attr::enable_debug, UnitAttr::get(op->getContext()));36+ op->setAttr(ascendc::attr::enableDebug, UnitAttr::get(op->getContext()));
38 }37 }
39 if (op.walk([](ascendc::DumpTensorOp) { return WalkResult::interrupt(); }).wasInterrupted()) {38 if (op.walk([](ascendc::DumpTensorOp) { return WalkResult::interrupt(); }).wasInterrupted()) {
40- op->setAttr(attr::enable_debug, UnitAttr::get(op->getContext()));39+ op->setAttr(ascendc::attr::enableDebug, UnitAttr::get(op->getContext()));
41 }40 }
42 }41 }
43};42};
@@ -32,16 +32,33 @@ class DetectKernelTypePass : public ascendc::impl::DetectKernelTypeBase<DetectKe
32public:32public:
33 void runOnOperation() override33 void runOnOperation() override
34 {34 {
35- ModuleOp op = getOperation();35+ auto moduleOp = getOperation();
36- if (op.walk([](ascendc::RegistMatmulObjOp) { return WalkResult::interrupt(); }).wasInterrupted())36+ if (moduleOp->hasAttr(attr::kernelType))
37- op->setAttr(attr::compile_mix, UnitAttr::get(op->getContext()));37+ return;
38+ auto hasVectorOps = false;
39+ auto hasCubeOps = false;
40+ moduleOp.walk([&hasVectorOps, &hasCubeOps](Operation* op) {
41+ if (!hasVectorOps && isa<VectorOp>(op))
42+ hasVectorOps = true;
43+ if (!hasCubeOps && isa<MmadOp, MmadWithBiasOp>(op))
44+ hasCubeOps = true;
45+ if (isa<RegistMatmulObjOp>(op))
46+ hasVectorOps = hasCubeOps = true;
47+ if (hasVectorOps && hasCubeOps)
48+ return WalkResult::interrupt();
49+ return WalkResult::advance();
50+ });
51+ StringRef kernelType;
52+ if (hasVectorOps && hasCubeOps)
53+ kernelType = attr::kernelMixed;
54+ else if (hasCubeOps)
55+ kernelType = attr::kernelCube;
56+ else
57+ kernelType = attr::kernelVector;
58+ moduleOp->setAttr(attr::kernelType, StringAttr::get(moduleOp.getContext(), kernelType));
38 }59 }
39};60};
40 61 
41} // namespace62} // namespace
42 63 
43-namespace mlir {64+std::unique_ptr<Pass> mlir::ascendc::createDetectKernelTypePass() { return std::make_unique<DetectKernelTypePass>(); }
44-namespace ascendc {
45-std::unique_ptr<Pass> createDetectKernelTypePass() { return std::make_unique<DetectKernelTypePass>(); }
46-} // namespace ascendc
47-} // namespace mlir
@@ -60,8 +60,4 @@ struct EraseSyncPass : public ascendc::impl::EraseSyncBase<EraseSyncPass> {
60 60 
61} // namespace61} // namespace
62 62 
63-namespace mlir {63+std::unique_ptr<Pass> mlir::ascendc::createEraseSyncPass() { return std::make_unique<EraseSyncPass>(); }
64-namespace ascendc {
65-std::unique_ptr<Pass> createEraseSyncPass() { return std::make_unique<EraseSyncPass>(); }
66-} // namespace ascendc
67-} // namespace mlir
@@ -24,7 +24,6 @@ namespace ascendc {
24} // namespace mlir24} // namespace mlir
25 25 
26using namespace mlir;26using namespace mlir;
27-using namespace mlir::ascendc;
28 27 
29namespace {28namespace {
30 29 
@@ -55,8 +54,7 @@ public:
55 54 
56} // namespace55} // namespace
57 56 
58-namespace mlir {57+std::unique_ptr<Pass> mlir::ascendc::createGenerateBoilerplatePass()
59-namespace ascendc {58+{
60-std::unique_ptr<Pass> createGenerateBoilerplatePass() { return std::make_unique<GenerateBoilerplatePass>(); }59+ return std::make_unique<GenerateBoilerplatePass>();
61-} // namespace ascendc60+}
62-} // namespace mlir
@@ -47,8 +47,4 @@ public:
47 47 
48} // namespace48} // namespace
49 49 
50-namespace mlir {50+std::unique_ptr<Pass> mlir::ascendc::createHoistQueBindPass() { return std::make_unique<HoistQueBindPass>(); }
51-namespace ascendc {
52-std::unique_ptr<Pass> createHoistQueBindPass() { return std::make_unique<HoistQueBindPass>(); }
53-} // namespace ascendc
54-} // namespace mlir
Rlib/Dialect/Asc/Transforms/HoistUBAllocation.cpp→lib/Dialect/Asc/Transforms/HoistTensorAllocation.cpp+18-9
@@ -17,7 +17,7 @@
17 17 
18namespace mlir {18namespace mlir {
19namespace ascendc {19namespace ascendc {
20-#define GEN_PASS_DEF_HOISTUBALLOCATION20+#define GEN_PASS_DEF_HOISTTENSORALLOCATION
21#include "ascir/Dialect/Asc/Transforms/Passes.h.inc"21#include "ascir/Dialect/Asc/Transforms/Passes.h.inc"
22} // namespace ascendc22} // namespace ascendc
23} // namespace mlir23} // namespace mlir
@@ -26,18 +26,26 @@ using namespace mlir;
26 26 
27namespace {27namespace {
28 28 
29-struct HoistTensor : ascendc::HoistOpPattern<ascendc::LocalTensorAutoOp> {29+using HoistTensor = ascendc::HoistOpPattern<ascendc::LocalTensorAutoOp>;
30+ 
31+struct HoistTensorExceptInOut : HoistTensor {
30 using HoistOpPattern::HoistOpPattern;32 using HoistOpPattern::HoistOpPattern;
31 33 
32 bool hoistable(ascendc::LocalTensorAutoOp op) const override { return !op.getInput() && !op.getOutput(); }34 bool hoistable(ascendc::LocalTensorAutoOp op) const override { return !op.getInput() && !op.getOutput(); }
33};35};
34 36 
35-struct HoistUBAllocationPass : public ascendc::impl::HoistUBAllocationBase<HoistUBAllocationPass> {37+struct HoistTensorAllocationPass : public ascendc::impl::HoistTensorAllocationBase<HoistTensorAllocationPass> {
38+ HoistTensorAllocationPass(const ascendc::HoistTensorAllocationOptions& opt) : HoistTensorAllocationBase(opt) {}
39+ 
36 void runOnOperation() override40 void runOnOperation() override
37 {41 {
38 MLIRContext* context = &getContext();42 MLIRContext* context = &getContext();
39 RewritePatternSet patterns(context);43 RewritePatternSet patterns(context);
40- patterns.add<HoistTensor>(context);44+ if (excludeInOut) {
45+ patterns.add<HoistTensorExceptInOut>(context);
46+ } else {
47+ patterns.add<HoistTensor>(context);
48+ }
41 if (applyPatternsAndFoldGreedily(getOperation(), std::move(patterns)).failed()) {49 if (applyPatternsAndFoldGreedily(getOperation(), std::move(patterns)).failed()) {
42 signalPassFailure();50 signalPassFailure();
43 }51 }
@@ -46,8 +54,9 @@ struct HoistUBAllocationPass : public ascendc::impl::HoistUBAllocationBase<Hoist
46 54 
47} // namespace55} // namespace
48 56 
49-namespace mlir {57+std::unique_ptr<Pass> mlir::ascendc::createHoistTensorAllocationPass(bool excludeInOut)
50-namespace ascendc {58+{
51-std::unique_ptr<Pass> createHoistUBAllocationPass() { return std::make_unique<HoistUBAllocationPass>(); }59+ HoistTensorAllocationOptions options;
52-} // namespace ascendc60+ options.excludeInOut = excludeInOut;
53-} // namespace mlir61+ return std::make_unique<HoistTensorAllocationPass>(options);
62+}
@@ -25,57 +25,70 @@ namespace ascendc {
25} // namespace mlir25} // namespace mlir
26 26 
27using namespace mlir;27using namespace mlir;
28-using namespace mlir::ascendc;
29 28 
30namespace {29namespace {
31 30 
32-void createDataCopyIfNeeded(Operation* op)31+using TensorOp = ascendc::LocalTensorAutoOp;
32+ 
33+template <typename ControlFlowOp>
34+void createDataCopyIfNeeded(ControlFlowOp op)
33{35{
34 for (auto& use : op->getUses()) {36 for (auto& use : op->getUses()) {
35 auto copyOp = dyn_cast<ascendc::DataCopyOp>(use.getOwner());37 auto copyOp = dyn_cast<ascendc::DataCopyOp>(use.getOwner());
36- if (!copyOp || copyOp.getDirection() != ascendc::CopyDirection::ubuf_gm)38+ if (!copyOp || !copyOp.isLocalToGlobal())
37 return;39 return;
38 OpBuilder builder(op);40 OpBuilder builder(op);
39 ascir::ConstantOpBuilder consts(builder);41 ascir::ConstantOpBuilder consts(builder);
40 auto type = cast<ascendc::BaseTensorType>(use.get().getType());42 auto type = cast<ascendc::BaseTensorType>(use.get().getType());
41- auto dst = builder.create<ascendc::LocalTensorAutoOp>(43+ auto dst = builder.create<TensorOp>(op->getLoc(), type, /*input*/ false, /*output*/ true, ValueRange{});
42- op->getLoc(), type, /*input*/ false,
43- /*output*/ true, ValueRange{});
44 builder.setInsertionPointAfter(op);44 builder.setInsertionPointAfter(op);
45 Value calCount = consts.i64(type.getNumElements());45 Value calCount = consts.i64(type.getNumElements());
46- builder.create<ascendc::DataCopyL2Op>(op->getLoc(), dst, use.get(), calCount);46+ auto extraOp = builder.create<ascendc::DataCopyL2Op>(op->getLoc(), dst, use.get(), calCount);
47+ extraOp.setDirection(ascendc::TPosition::VECCALC, ascendc::TPosition::VECCALC);
47 copyOp.setSrc(dst);48 copyOp.setSrc(dst);
48 }49 }
49}50}
50 51 
52+TensorOp findTensorOrigin(Value tensor)
53+{
54+ auto* defOp = tensor.getDefiningOp();
55+ if (!defOp)
56+ return {};
57+ if (auto op = dyn_cast<TensorOp>(defOp))
58+ return op;
59+ if (auto op = dyn_cast<ascendc::LocalTensorReinterpretCastOp>(defOp))
60+ return findTensorOrigin(op.getIn());
61+ if (auto op = dyn_cast<ascendc::LocalTensorSubIndexOp>(defOp))
62+ return findTensorOrigin(op.getTensor());
63+ return {};
64+}
65+ 
51void setInOutTensors(func::FuncOp funcOp)66void setInOutTensors(func::FuncOp funcOp)
52{67{
53- funcOp.walk([](ascendc::LocalTensorAutoOp op) {68+ funcOp.walk([](ascendc::DataCopyOp op) {
54- bool input = false;69+ auto src = findTensorOrigin(op.getSrc());
55- bool output = false;70+ if (src && op.isLocalToGlobal())
56- for (Operation* user : op->getUsers()) {71+ src.setOutput(true);
57- if (auto copyOp = dyn_cast<ascendc::DataCopyOp>(user)) {72+ auto dst = findTensorOrigin(op.getDst());
58- auto dir = copyOp.getDirection();73+ if (dst && op.isGlobalToLocal())
59- if (dir == ascendc::CopyDirection::gm_ubuf) {74+ dst.setInput(true);
60- input = true;75+ if (!op.isLocalToLocal())
61- continue;76+ return;
62- }77+ if (isa<ascendc::CopyToL0Op, ascendc::FixpipeOp>(*op)) {
63- if (dir == ascendc::CopyDirection::ubuf_gm) {78+ if (src)
64- output = true;79+ src.setOutput(true);
65- continue;80+ if (dst)
66- }81+ dst.setInput(true);
67- }
68 }82 }
69- op.setInput(input);
70- op.setOutput(output);
71 });83 });
72- funcOp.walk([](scf::ForOp op) { createDataCopyIfNeeded(op); });84+ funcOp.walk(createDataCopyIfNeeded<scf::ForOp>);
73- funcOp.walk([](scf::IfOp op) { createDataCopyIfNeeded(op); });85+ funcOp.walk(createDataCopyIfNeeded<scf::IfOp>);
86+ funcOp.walk(createDataCopyIfNeeded<scf::WhileOp>);
74}87}
75 88 
76void fixInOutTensor(func::FuncOp& funcOp)89void fixInOutTensor(func::FuncOp& funcOp)
77{90{
78- funcOp.walk([](ascendc::LocalTensorAutoOp inTensor) {91+ funcOp.walk([](TensorOp inTensor) {
79 if (!inTensor.getInput() || inTensor.getOutput())92 if (!inTensor.getInput() || inTensor.getOutput())
80 return;93 return;
81 auto loc = inTensor.getLoc();94 auto loc = inTensor.getLoc();
@@ -85,14 +98,13 @@ void fixInOutTensor(func::FuncOp& funcOp)
85 for (auto& use : inTensor->getUses()) {98 for (auto& use : inTensor->getUses()) {
86 auto* owner = use.getOwner();99 auto* owner = use.getOwner();
87 auto copyOp = dyn_cast<ascendc::DataCopyOp>(owner);100 auto copyOp = dyn_cast<ascendc::DataCopyOp>(owner);
88- if (!copyOp || copyOp.getDirection() != ascendc::CopyDirection::ubuf_gm)101+ if (!copyOp || !copyOp.isLocalToGlobal())
89 return builder.setInsertionPoint(owner);102 return builder.setInsertionPoint(owner);
90 ascir::ConstantOpBuilder consts(builder);103 ascir::ConstantOpBuilder consts(builder);
91 Value calCount = consts.i64(tensorType.getNumElements());104 Value calCount = consts.i64(tensorType.getNumElements());
92- auto outTensor = builder.create<ascendc::LocalTensorAutoOp>(105+ auto outTensor = builder.create<TensorOp>(loc, tensorType, /*input*/ false, /*output*/ true, ValueRange{});
93- loc, tensorType, /*input*/ false,106+ auto extraOp = builder.create<ascendc::DataCopyL2Op>(loc, outTensor, inTensor, calCount);
94- /*output*/ true, ValueRange{});107+ extraOp.setDirection(ascendc::TPosition::VECCALC, ascendc::TPosition::VECCALC);
95- builder.create<ascendc::DataCopyL2Op>(loc, outTensor, inTensor, calCount);
96 owner->setOperand(use.getOperandNumber(), outTensor);108 owner->setOperand(use.getOperandNumber(), outTensor);
97 }109 }
98 });110 });
@@ -102,25 +114,15 @@ struct InputOutputTensorPass : public ascendc::impl::InputOutputTensorBase<Input
102 void runOnOperation() override114 void runOnOperation() override
103 {115 {
104 func::FuncOp funcOp = getOperation();116 func::FuncOp funcOp = getOperation();
105- if (funcOp.isDeclaration()) {
106- return;
107- }
108 setInOutTensors(funcOp);117 setInOutTensors(funcOp);
109 fixInOutTensor(funcOp);118 fixInOutTensor(funcOp);
110 MLIRContext* context = &getContext();119 MLIRContext* context = &getContext();
111 RewritePatternSet patterns(context);120 RewritePatternSet patterns(context);
112- ascendc::LocalTensorAutoOp::getCanonicalizationPatterns(patterns, context);121+ if (applyPatternsAndFoldGreedily(funcOp, std::move(patterns)).failed())
113- if (applyPatternsAndFoldGreedily(funcOp, std::move(patterns)).failed()) {
114 signalPassFailure();122 signalPassFailure();
115- return;
116- }
117 }123 }
118};124};
119 125 
120} // namespace126} // namespace
121 127 
122-namespace mlir {128+std::unique_ptr<Pass> mlir::ascendc::createInputOutputTensorPass() { return std::make_unique<InputOutputTensorPass>(); }
123-namespace ascendc {
124-std::unique_ptr<Pass> createInputOutputTensorPass() { return std::make_unique<InputOutputTensorPass>(); }
125-} // namespace ascendc
126-} // namespace mlir
Rlib/Dialect/Asc/Transforms/InsertSync.cpp→lib/Dialect/Asc/Transforms/InsertQueSync.cpp+3-7
@@ -21,7 +21,7 @@
21 21 
22namespace mlir {22namespace mlir {
23namespace ascendc {23namespace ascendc {
24-#define GEN_PASS_DEF_INSERTSYNC24+#define GEN_PASS_DEF_INSERTQUESYNC
25#include "ascir/Dialect/Asc/Transforms/Passes.h.inc"25#include "ascir/Dialect/Asc/Transforms/Passes.h.inc"
26} // namespace ascendc26} // namespace ascendc
27} // namespace mlir27} // namespace mlir
@@ -171,7 +171,7 @@ void canonicalizeBarriers(func::FuncOp funcOp)
171 (void)applyPatternsAndFoldGreedily(funcOp, std::move(patterns));171 (void)applyPatternsAndFoldGreedily(funcOp, std::move(patterns));
172}172}
173 173 
174-struct InsertSyncPass : public ascendc::impl::InsertSyncBase<InsertSyncPass> {174+struct InsertQueSyncPass : public ascendc::impl::InsertQueSyncBase<InsertQueSyncPass> {
175public:175public:
176 void runOnOperation() override176 void runOnOperation() override
177 {177 {
@@ -192,8 +192,4 @@ public:
192 192 
193} // namespace193} // namespace
194 194 
195-namespace mlir {195+std::unique_ptr<Pass> mlir::ascendc::createInsertQueSyncPass() { return std::make_unique<InsertQueSyncPass>(); }
196-namespace ascendc {
197-std::unique_ptr<Pass> createInsertSyncPass() { return std::make_unique<InsertSyncPass>(); }
198-} // namespace ascendc
199-} // namespace mlir
@@ -42,7 +42,7 @@ BlockArgument appendKernelArgument(func::FuncOp op, emitasc::KernelArgument kind
42 return op.getArgument(idx);42 return op.getArgument(idx);
43}43}
44 44 
45-void processKernel(func::FuncOp op)45+void processKernel(func::FuncOp op, bool setFftsAddr)
46{46{
47 auto builder = OpBuilder::atBlockBegin(&op.getFunctionBody().front());47 auto builder = OpBuilder::atBlockBegin(&op.getFunctionBody().front());
48 for (unsigned i = 0; i < op.getNumArguments(); i++) {48 for (unsigned i = 0; i < op.getNumArguments(); i++) {
@@ -50,12 +50,14 @@ void processKernel(func::FuncOp op)
50 i, emitasc::attr::kernelArg,50 i, emitasc::attr::kernelArg,
51 builder.getAttr<emitasc::KernelArgumentAttr>(emitasc::KernelArgument::Explicit));51 builder.getAttr<emitasc::KernelArgumentAttr>(emitasc::KernelArgument::Explicit));
52 }52 }
53- auto as = builder.getI64IntegerAttr(static_cast<int64_t>(ascendc::AddressSpace::gm));
54 auto loc = builder.getUnknownLoc();53 auto loc = builder.getUnknownLoc();
55- auto fftsAddr = appendKernelArgument(54+ if (setFftsAddr) {
56- op, emitasc::KernelArgument::FftsAddr, "ffts_addr",55+ auto as = builder.getI64IntegerAttr(static_cast<int64_t>(ascendc::AddressSpace::gm));
57- MemRefType::get(ShapedType::kDynamic, builder.getIntegerType(64, false), AffineMap(), as));56+ auto fftsAddr = appendKernelArgument(
58- builder.create<ascendc::SetFftsBaseAddrOp>(loc, fftsAddr);57+ op, emitasc::KernelArgument::FftsAddr, "ffts_addr",
58+ MemRefType::get(ShapedType::kDynamic, builder.getIntegerType(64, false), AffineMap(), as));
59+ builder.create<ascendc::SetFftsBaseAddrOp>(loc, fftsAddr);
60+ }
59 bool hasMatmul = op.walk([](ascendc::RegistMatmulObjOp) { return WalkResult::interrupt(); }).wasInterrupted();61 bool hasMatmul = op.walk([](ascendc::RegistMatmulObjOp) { return WalkResult::interrupt(); }).wasInterrupted();
60 bool matmulCubeOnly = op->getParentOfType<ModuleOp>()->hasAttrOfType<UnitAttr>(ascendc::attr::matmulCubeOnly);62 bool matmulCubeOnly = op->getParentOfType<ModuleOp>()->hasAttrOfType<UnitAttr>(ascendc::attr::matmulCubeOnly);
61 if (hasMatmul && !matmulCubeOnly) {63 if (hasMatmul && !matmulCubeOnly) {
@@ -69,12 +71,14 @@ void processKernel(func::FuncOp op)
69}71}
70 72 
71struct LegalizeKernelArgsPass : public ascendc::impl::LegalizeKernelArgsBase<LegalizeKernelArgsPass> {73struct LegalizeKernelArgsPass : public ascendc::impl::LegalizeKernelArgsBase<LegalizeKernelArgsPass> {
74+ LegalizeKernelArgsPass(const ascendc::LegalizeKernelArgsOptions& options) : LegalizeKernelArgsBase(options) {}
75+ 
72 void runOnOperation() override76 void runOnOperation() override
73 {77 {
74 auto mod = getOperation();78 auto mod = getOperation();
75- mod.walk([](func::FuncOp op) {79+ mod.walk([this](func::FuncOp op) {
76 if (op->hasAttrOfType<UnitAttr>(ascendc::attr::global)) {80 if (op->hasAttrOfType<UnitAttr>(ascendc::attr::global)) {
77- processKernel(op);81+ processKernel(op, setFftsAddr);
78 }82 }
79 });83 });
80 }84 }
@@ -82,8 +86,9 @@ struct LegalizeKernelArgsPass : public ascendc::impl::LegalizeKernelArgsBase<Leg
82 86 
83} // namespace87} // namespace
84 88 
85-namespace mlir {89+std::unique_ptr<Pass> mlir::ascendc::createLegalizeKernelArgsPass(bool setFftsAddr)
86-namespace ascendc {90+{
87-std::unique_ptr<Pass> createLegalizeKernelArgsPass() { return std::make_unique<LegalizeKernelArgsPass>(); }91+ LegalizeKernelArgsOptions options;
88-} // namespace ascendc92+ options.setFftsAddr = setFftsAddr;
89-} // namespace mlir93+ return std::make_unique<LegalizeKernelArgsPass>(options);
94+}
@@ -8,10 +8,9 @@
8 * See LICENSE in the root of the software repository for the full text of the License.8 * See LICENSE in the root of the software repository for the full text of the License.
9 */9 */
10 10 
11-#include <climits>
12- 
13#include "ascir/Dialect/Asc/IR/Asc.h"11#include "ascir/Dialect/Asc/IR/Asc.h"
14#include "ascir/Dialect/Asc/Transforms/Passes.h"12#include "ascir/Dialect/Asc/Transforms/Passes.h"
13+#include "ascir/Dialect/Asc/Utils/Utils.h"
15#include "ascir/Dialect/Utils/ConstantOpBuilder.h"14#include "ascir/Dialect/Utils/ConstantOpBuilder.h"
16 15 
17#include "mlir/Dialect/Func/IR/FuncOps.h"16#include "mlir/Dialect/Func/IR/FuncOps.h"
@@ -32,7 +31,9 @@ using namespace mlir;
32namespace {31namespace {
33 32 
34struct MaterializeLocalTensor : OpRewritePattern<ascendc::LocalTensorAutoOp> {33struct MaterializeLocalTensor : OpRewritePattern<ascendc::LocalTensorAutoOp> {
35- using OpRewritePattern::OpRewritePattern;34+ bool alwaysBuf;
35+ 
36+ MaterializeLocalTensor(MLIRContext* context, bool alwaysBuf) : alwaysBuf(alwaysBuf), OpRewritePattern(context) {}
36 37 
37 static ascendc::TPosition getPosition(ascendc::LocalTensorAutoOp op)38 static ascendc::TPosition getPosition(ascendc::LocalTensorAutoOp op)
38 {39 {
@@ -49,18 +50,29 @@ struct MaterializeLocalTensor : OpRewritePattern<ascendc::LocalTensorAutoOp> {
49 auto loc = op.getLoc();50 auto loc = op.getLoc();
50 ascir::ConstantOpBuilder consts(rewriter);51 ascir::ConstantOpBuilder consts(rewriter);
51 Value length;52 Value length;
53+ auto position = op.getPosition();
52 if (type.hasStaticShape()) {54 if (type.hasStaticShape()) {
53- length = consts.i64(type.getNumElements() * type.getElementTypeBitWidth() / CHAR_BIT);55+ if (position == ascendc::TPosition::A1 || position == ascendc::TPosition::B1 ||
56+ position == ascendc::TPosition::A2 || position == ascendc::TPosition::B2)
57+ length = consts.i64(ascendc::getTypeSizeCubeBlockAlign(type, position));
58+ else
59+ length = consts.i64(ascendc::getTypeSize(type));
54 } else {60 } else {
55 assert(op->getNumOperands() != 0 && "must have operands for dynamic shape");61 assert(op->getNumOperands() != 0 && "must have operands for dynamic shape");
56- length = consts.i64(type.getElementTypeBitWidth() / CHAR_BIT);62+ auto elementTypeSize = ascendc::getElementTypeSize(type);
63+ length = consts.i64(elementTypeSize);
64+ auto align = consts.i64(ascendc::ubBlockSize / elementTypeSize);
57 for (auto dim : op.getDynamicShape()) {65 for (auto dim : op.getDynamicShape()) {
66+ if (position == ascendc::TPosition::A1) {
67+ auto ceilDim = rewriter.create<arith::CeilDivSIOp>(loc, dim, align);
68+ dim = rewriter.create<arith::MulIOp>(loc, ceilDim, align);
69+ }
58 length = rewriter.create<arith::MulIOp>(loc, length, dim);70 length = rewriter.create<arith::MulIOp>(loc, length, dim);
59 }71 }
60 }72 }
61 Value pipe = rewriter.create<ascendc::PipeOp>(loc);73 Value pipe = rewriter.create<ascendc::PipeOp>(loc);
62- if (!op.getInput() && !op.getOutput()) {74+ if (alwaysBuf || !op.getInput() && !op.getOutput()) {
63- auto bufferTy = ascendc::TBufType::get(op.getContext(), ascendc::TPosition::VECCALC);75+ auto bufferTy = ascendc::TBufType::get(op.getContext(), op.getPosition());
64 Value buffer = rewriter.create<ascendc::TBufOp>(loc, bufferTy);76 Value buffer = rewriter.create<ascendc::TBufOp>(loc, bufferTy);
65 rewriter.create<ascendc::TPipeInitBufferOp>(loc, pipe, buffer, length);77 rewriter.create<ascendc::TPipeInitBufferOp>(loc, pipe, buffer, length);
66 rewriter.replaceOpWithNewOp<ascendc::TBufGetTensorOp>(op, type, buffer);78 rewriter.replaceOpWithNewOp<ascendc::TBufGetTensorOp>(op, type, buffer);
@@ -77,7 +89,9 @@ struct MaterializeLocalTensor : OpRewritePattern<ascendc::LocalTensorAutoOp> {
77 }89 }
78};90};
79 91 
80-class MaterializeTensorPass : public ascendc::impl::MaterializeTensorBase<MaterializeTensorPass> {92+struct MaterializeTensorPass : public ascendc::impl::MaterializeTensorBase<MaterializeTensorPass> {
93+ MaterializeTensorPass(const ascendc::MaterializeTensorOptions& options) : MaterializeTensorBase(options) {}
94+ 
81 void runOnOperation() override95 void runOnOperation() override
82 {96 {
83 func::FuncOp funcOp = getOperation();97 func::FuncOp funcOp = getOperation();
@@ -86,7 +100,7 @@ class MaterializeTensorPass : public ascendc::impl::MaterializeTensorBase<Materi
86 }100 }
87 MLIRContext* context = &getContext();101 MLIRContext* context = &getContext();
88 RewritePatternSet patterns(context);102 RewritePatternSet patterns(context);
89- patterns.add<MaterializeLocalTensor>(context);103+ patterns.add<MaterializeLocalTensor>(context, alwaysBuf);
90 if (applyPatternsAndFoldGreedily(funcOp, std::move(patterns)).failed()) {104 if (applyPatternsAndFoldGreedily(funcOp, std::move(patterns)).failed()) {
91 signalPassFailure();105 signalPassFailure();
92 }106 }
@@ -94,8 +108,9 @@ class MaterializeTensorPass : public ascendc::impl::MaterializeTensorBase<Materi
94};108};
95} // namespace109} // namespace
96 110 
97-namespace mlir {111+std::unique_ptr<Pass> mlir::ascendc::createMaterializeTensorPass(bool alwaysBuf)
98-namespace ascendc {112+{
99-std::unique_ptr<Pass> createMaterializeTensorPass() { return std::make_unique<MaterializeTensorPass>(); }113+ MaterializeTensorOptions options;
100-} // namespace ascendc114+ options.alwaysBuf = alwaysBuf;
101-} // namespace mlir115+ return std::make_unique<MaterializeTensorPass>(options);
116+}
@@ -32,8 +32,4 @@ struct NoopPass : public ascendc::impl::NoopBase<NoopPass> {
32 32 
33} // namespace33} // namespace
34 34 
35-namespace mlir {35+std::unique_ptr<Pass> mlir::ascendc::createNoopPass() { return std::make_unique<NoopPass>(); }
36-namespace ascendc {
37-std::unique_ptr<Pass> createNoopPass() { return std::make_unique<NoopPass>(); }
38-} // namespace ascendc
39-} // namespace mlir
@@ -40,8 +40,4 @@ struct PrivatizeFuncPass : public ascendc::impl::PrivatizeFuncBase<PrivatizeFunc
40 40 
41} // namespace41} // namespace
42 42 
43-namespace mlir {43+std::unique_ptr<Pass> mlir::ascendc::createPrivatizeFuncPass() { return std::make_unique<PrivatizeFuncPass>(); }
44-namespace ascendc {
45-std::unique_ptr<Pass> createPrivatizeFuncPass() { return std::make_unique<PrivatizeFuncPass>(); }
46-} // namespace ascendc
47-} // namespace mlir
@@ -10,7 +10,6 @@
10 10 
11#include "ascir/Dialect/Asc/IR/Asc.h"11#include "ascir/Dialect/Asc/IR/Asc.h"
12#include "ascir/Dialect/Asc/Transforms/Passes.h"12#include "ascir/Dialect/Asc/Transforms/Passes.h"
13-#include "ascir/Dialect/Utils/ConstantOpBuilder.h"
14 13 
15#include "mlir/Dialect/Func/IR/FuncOps.h"14#include "mlir/Dialect/Func/IR/FuncOps.h"
16#include "mlir/IR/Builders.h"15#include "mlir/IR/Builders.h"
@@ -30,7 +29,7 @@ void unifyPipe(func::FuncOp root)
30{29{
31 SmallVector<ascendc::PipeOp> pipes;30 SmallVector<ascendc::PipeOp> pipes;
32 root.walk([&pipes](ascendc::PipeOp op) { pipes.push_back(op); });31 root.walk([&pipes](ascendc::PipeOp op) { pipes.push_back(op); });
33- if (pipes.size() <= 1) {32+ if (pipes.size() == 1) {
34 return;33 return;
35 }34 }
36 auto builder = OpBuilder::atBlockBegin(&root.getBody().front());35 auto builder = OpBuilder::atBlockBegin(&root.getBody().front());
@@ -47,8 +46,4 @@ class UnifyPipePass : public ascendc::impl::UnifyPipeBase<UnifyPipePass> {
47 46 
48} // namespace47} // namespace
49 48 
50-namespace mlir {49+std::unique_ptr<Pass> mlir::ascendc::createUnifyPipePass() { return std::make_unique<UnifyPipePass>(); }
51-namespace ascendc {
52-std::unique_ptr<Pass> createUnifyPipePass() { return std::make_unique<UnifyPipePass>(); }
53-} // namespace ascendc
54-} // namespace mlir
@@ -157,8 +157,4 @@ struct VerifySyncPass : public ascendc::impl::VerifySyncBase<VerifySyncPass> {
157 157 
158} // namespace158} // namespace
159 159 
160-namespace mlir {160+std::unique_ptr<Pass> mlir::ascendc::createVerifySyncPass() { return std::make_unique<VerifySyncPass>(); }
161-namespace ascendc {
162-std::unique_ptr<Pass> createVerifySyncPass() { return std::make_unique<VerifySyncPass>(); }
163-} // namespace ascendc
164-} // namespace mlir
@@ -9,13 +9,16 @@
9 */9 */
10 10 
11#include "ascir/Dialect/Asc/Utils/Utils.h"11#include "ascir/Dialect/Asc/Utils/Utils.h"
12+#include "ascir/Dialect/Asc/Utils/Attributes.h"
12#include "ascir/Dialect/Utils/Inlining.h"13#include "ascir/Dialect/Utils/Inlining.h"
13 14 
14#include "mlir/Dialect/EmitC/IR/EmitC.h"15#include "mlir/Dialect/EmitC/IR/EmitC.h"
16+#include "mlir/Dialect/SCF/IR/SCF.h"
15#include "mlir/IR/BuiltinDialect.h"17#include "mlir/IR/BuiltinDialect.h"
16#include "mlir/IR/BuiltinOps.h"18#include "mlir/IR/BuiltinOps.h"
17#include "mlir/IR/Dominance.h"19#include "mlir/IR/Dominance.h"
18#include "mlir/IR/Operation.h"20#include "mlir/IR/Operation.h"
21+#include "llvm/ADT/TypeSwitch.h"
19 22 
20namespace mlir {23namespace mlir {
21 24 
@@ -24,6 +27,73 @@ using AllowInline = ascir::AllowlistInlinerInterface<T...>;
24 27 
25namespace ascendc {28namespace ascendc {
26 29 
30+namespace {
31+ 
32+void appendImplicitUsers(Value value, SmallVectorImpl<Operation*>& allUsers)
33+{
34+ llvm::copy(value.getUsers(), std::back_inserter(allUsers));
35+ for (auto* user : value.getUsers()) {
36+ if (isa<CastOpInterface>(user) || isa<LocalTensorSubIndexOp>(user)) {
37+ auto users = user->getUsers();
38+ if (!users.empty()) {
39+ allUsers.append(users.begin(), users.end());
40+ appendImplicitUsers(user->getResult(0), allUsers);
41+ }
42+ }
43+ // if value use as init value then this memory use for first iteration and life interval must include
44+ // lifeInterval iterArg
45+ if (auto forOp = dyn_cast<scf::ForOp>(user)) {
46+ auto inits = forOp.getInits();
47+ auto iterArgs = forOp.getRegionIterArgs();
48+ for (int i = 0; i < inits.size(); ++i) {
49+ if (inits[i] == value) {
50+ appendImplicitUsers(iterArgs[i], allUsers);
51+ }
52+ }
53+ }
54+ // if value return in yield op then it is accumulator and he used as return value in forOp
55+ if (auto yieldOp = dyn_cast<scf::YieldOp>(user)) {
56+ auto forOp = cast<scf::ForOp>(yieldOp->getParentOp());
57+ allUsers.push_back(forOp);
58+ auto opnds = yieldOp.getOperands();
59+ auto iterArgs = forOp.getRegionIterArgs();
60+ for (int i = 0; i < opnds.size(); ++i) {
61+ if (opnds[i] == value) {
62+ appendImplicitUsers(iterArgs[i], allUsers);
63+ appendImplicitUsers(forOp->getResult(i), allUsers);
64+ }
65+ }
66+ }
67+ }
68+}
69+ 
70+} // namespace
71+ 
72+int64_t getTypeSize(Type type)
73+{
74+ if (auto shaped = dyn_cast<ShapedType>(type))
75+ return shaped.getNumElements() * getTypeSize(shaped.getElementType());
76+ return type.getIntOrFloatBitWidth() / CHAR_BIT;
77+}
78+ 
79+int64_t getTypeSizeCubeBlockAlign(ShapedType type, TPosition position)
80+{
81+ auto shape = type.getShape();
82+ int64_t elemSize = getElementTypeSize(type);
83+ int64_t elemAlign = cubeKBlockBytes / elemSize;
84+ int64_t size = 1;
85+ for (size_t i = 0; i < shape.size(); ++i) {
86+ int64_t align = cubeBlockSize;
87+ if (((position == TPosition::A1 || position == TPosition::A2) && i == 1) ||
88+ ((position == TPosition::B1 || position == TPosition::B2) && i == 0))
89+ align = elemAlign;
90+ size *= static_cast<int64_t>(llvm::alignTo(shape[i], align));
91+ }
92+ return size * elemSize;
93+}
94+ 
95+int64_t getElementTypeSize(ShapedType type) { return getTypeSize(type.getElementType()); }
96+ 
27bool opPrecedes(Operation* lhs, Operation* rhs) { return lhs != rhs && lhs->isBeforeInBlock(rhs); }97bool opPrecedes(Operation* lhs, Operation* rhs) { return lhs != rhs && lhs->isBeforeInBlock(rhs); }
28 98 
29bool opPrecedes(Operation* lhs, Operation* rhs, DominanceInfo& di)99bool opPrecedes(Operation* lhs, Operation* rhs, DominanceInfo& di)
@@ -39,7 +109,17 @@ bool opPrecedes(Operation* lhs, Operation* rhs, DominanceInfo& di)
39 Block* dtr = di.findNearestCommonDominator(lhsBlk, rhsBlk);109 Block* dtr = di.findNearestCommonDominator(lhsBlk, rhsBlk);
40 Operation* lhsAnc = dtr->findAncestorOpInBlock(*lhs);110 Operation* lhsAnc = dtr->findAncestorOpInBlock(*lhs);
41 Operation* rhsAnc = dtr->findAncestorOpInBlock(*rhs);111 Operation* rhsAnc = dtr->findAncestorOpInBlock(*rhs);
42- return lhsAnc->isBeforeInBlock(rhsAnc);112+ if (lhsAnc != rhsAnc)
113+ return lhsAnc->isBeforeInBlock(rhsAnc);
114+ if (lhs->isAncestor(rhs))
115+ return true;
116+ if (rhs->isAncestor(lhs))
117+ return false;
118+ if (auto ifOp = dyn_cast<scf::IfOp>(lhsAnc))
119+ return ifOp.thenBlock()->findAncestorOpInBlock(*lhs);
120+ if (auto whileOp = dyn_cast<scf::WhileOp>(lhsAnc))
121+ return whileOp.getBeforeBody()->findAncestorOpInBlock(*lhs);
122+ return di.properlyDominates(lhs, rhs);
43}123}
44 124 
45void registerInlinerInterfaces(DialectRegistry& registry)125void registerInlinerInterfaces(DialectRegistry& registry)
@@ -52,5 +132,85 @@ void registerInlinerInterfaces(DialectRegistry& registry)
52 });132 });
53}133}
54 134 
135+ModuleOp getModule(Operation* op)
136+{
137+ if (isa<ModuleOp>(op))
138+ return cast<ModuleOp>(op);
139+ auto mod = op->getParentOfType<ModuleOp>();
140+ assert(mod && "operation must be within a module");
141+ return mod;
142+}
143+ 
144+StringRef getCompilationArch(Operation* op)
145+{
146+ if (auto attr = getModule(op)->getAttrOfType<StringAttr>(attr::compilationArch))
147+ return attr.getValue();
148+ return {};
149+}
150+ 
151+StringRef getSocVersion(Operation* op)
152+{
153+ if (auto attr = getModule(op)->getAttrOfType<StringAttr>(attr::socVersion))
154+ return attr.getValue();
155+ return {};
156+}
157+ 
158+std::optional<int64_t> getVecLen(Operation* op)
159+{
160+ if (auto attr = getModule(op)->getAttrOfType<IntegerAttr>(attr::vfVecLen))
161+ return attr.getValue().getSExtValue();
162+ return std::nullopt;
163+}
164+ 
165+bool isTargetArchC310(Operation* op) { return getCompilationArch(op) == "c310"; }
166+ 
167+SmallVector<Operation*> collectAllUsers(LocalTensorAutoOp tensorOp)
168+{
169+ SmallVector<Operation*> users;
170+ appendImplicitUsers(tensorOp, users);
171+ return users;
172+}
173+ 
174+LocalTensorAutoOp getAllocationRoot(Value v)
175+{
176+ auto* defOp = v.getDefiningOp();
177+ if (!defOp)
178+ return {};
179+ if (auto op = dyn_cast<LocalTensorAutoOp>(defOp))
180+ return op;
181+ if (auto op = dyn_cast<LocalTensorReinterpretCastOp>(defOp))
182+ return getAllocationRoot(op.getIn());
183+ if (auto op = dyn_cast<LocalTensorSubIndexOp>(defOp))
184+ return getAllocationRoot(op.getTensor());
185+ return {};
186+}
187+ 
188+Pipe getOpPipe(Operation* op, Pipe defaultPipe)
189+{
190+ return llvm::TypeSwitch<Operation*, Pipe>(op)
191+ .Case<VectorOp>([](auto) { return Pipe::PIPE_V; })
192+ .Case<MmadOp, MmadWithBiasOp>([](auto) { return Pipe::PIPE_M; })
193+ .Case([](FixpipeOp) { return Pipe::PIPE_FIX; })
194+ .Case([](CopyToL0Op) { return Pipe::PIPE_MTE1; })
195+ .Case([](FillOp) { return Pipe::PIPE_MTE2; })
196+ .Case([defaultPipe](DataCopyOp copyOp) {
197+ if (auto direction = copyOp.getDirection()) {
198+ auto [src, dst] = *direction;
199+ if (src == TPosition::A1 && dst == TPosition::VECCALC ||
200+ src == TPosition::A1 && (dst == TPosition::A2 || dst == TPosition::B2 || dst == TPosition::CO1))
201+ return Pipe::PIPE_MTE1;
202+ if (src == TPosition::GM)
203+ return Pipe::PIPE_MTE2;
204+ if (dst == TPosition::GM || src == TPosition::VECCALC && dst == TPosition::A1)
205+ return Pipe::PIPE_MTE3;
206+ if (src == TPosition::VECCALC && dst == TPosition::VECCALC)
207+ return Pipe::PIPE_V;
208+ }
209+ return defaultPipe;
210+ })
211+ .Case<LocalTensorGetValueOp, LocalTensorSetValueOp>([](auto) { return Pipe::PIPE_S; })
212+ .Default(defaultPipe);
213+}
214+ 
55} // namespace ascendc215} // namespace ascendc
56} // namespace mlir216} // namespace mlir
@@ -20,6 +20,96 @@
20using namespace mlir;20using namespace mlir;
21using namespace mlir::emitasc;21using namespace mlir::emitasc;
22 22 
23+//===----------------------------------------------------------------------===//
24+// InitStructOp
25+//===----------------------------------------------------------------------===//
26+ 
27+void InitStructOp::addField(StringRef name, Value value)
28+{
29+ SmallVector<Attribute> names(getFieldNames().getValue());
30+ names.push_back(StringAttr::get(getContext(), name));
31+ setFieldNamesAttr(ArrayAttr::get(getContext(), names));
32+ getFieldValuesMutable().append(value);
33+}
34+ 
35+Value InitStructOp::getField(StringRef name)
36+{
37+ auto index = getFieldOperandIndex(name);
38+ return index ? getFieldValues()[*index] : Value{};
39+}
40+ 
41+std::optional<size_t> InitStructOp::getFieldOperandIndex(StringRef name)
42+{
43+ auto names = getFieldNames().getValue();
44+ const auto* it = llvm::find_if(names, [name](Attribute attr) { return cast<StringAttr>(attr).getValue() == name; });
45+ if (it == names.end())
46+ return std::nullopt;
47+ return std::distance(names.begin(), it);
48+}
49+ 
50+size_t InitStructOp::getNumFields() { return getFieldNames().size(); }
51+ 
52+bool InitStructOp::hasField(StringRef name) { return getFieldOperandIndex(name).has_value(); }
53+ 
54+ParseResult InitStructOp::parse(OpAsmParser& parser, OperationState& result)
55+{
56+ auto& builder = parser.getBuilder();
57+ Type resultType;
58+ if (parser.parseType(resultType) || parser.parseLParen())
59+ return ParseResult::failure();
60+ result.addTypes(resultType);
61+ SmallVector<Attribute> names;
62+ SmallVector<Value> values;
63+ bool first = true;
64+ while (parser.parseOptionalRParen().failed()) {
65+ if (first) {
66+ first = false;
67+ } else {
68+ if (parser.parseComma())
69+ return ParseResult::failure();
70+ }
71+ std::string name;
72+ OpAsmParser::UnresolvedOperand operand;
73+ Type type;
74+ if (parser.parseString(&name) || parser.parseEqual() || parser.parseOperand(operand) ||
75+ parser.parseColonType(type) || parser.resolveOperand(operand, type, values))
76+ return ParseResult::failure();
77+ names.push_back(builder.getStringAttr(name));
78+ }
79+ result.addAttribute(getAttributeNames()[0], builder.getArrayAttr(names));
80+ result.addOperands(values);
81+ return ParseResult::success();
82+}
83+ 
84+void InitStructOp::print(OpAsmPrinter& printer)
85+{
86+ printer << ' ' << getType() << '(';
87+ bool first = true;
88+ for (auto [name, value] : llvm::zip_equal(getFieldNames(), getFieldValues())) {
89+ if (first)
90+ first = false;
91+ else
92+ printer << ", ";
93+ printer << cast<StringAttr>(name) << " = " << value << " : " << value.getType();
94+ }
95+ printer << ')';
96+}
97+ 
98+void InitStructOp::setField(StringRef name, Value value)
99+{
100+ auto index = getFieldOperandIndex(name);
101+ if (!index)
102+ return;
103+ getFieldValuesMutable().slice(*index, 1).assign(value);
104+}
105+ 
106+LogicalResult InitStructOp::verify()
107+{
108+ if (getFieldNames().size() != getFieldValues().size())
109+ return emitOpError("must have number of field names equal to number of field values");
110+ return success();
111+}
112+ 
23//===----------------------------------------------------------------------===//113//===----------------------------------------------------------------------===//
24// PtrOffsetOp114// PtrOffsetOp
25//===----------------------------------------------------------------------===//115//===----------------------------------------------------------------------===//
@@ -60,6 +150,15 @@ OpFoldResult VariableOp::getInit(bool fold)
60 return staticInit.value();150 return staticInit.value();
61}151}
62 152 
153+LogicalResult VariableOp::canonicalize(VariableOp op, PatternRewriter& rewriter)
154+{
155+ if (op->getUses().empty()) {
156+ rewriter.eraseOp(op);
157+ return success();
158+ }
159+ return failure();
160+}
161+ 
63//===----------------------------------------------------------------------===//162//===----------------------------------------------------------------------===//
64// EmitAscDialect163// EmitAscDialect
65//===----------------------------------------------------------------------===//164//===----------------------------------------------------------------------===//
@@ -46,7 +46,7 @@ void printMethod(raw_indented_ostream& os, const Record* def)
46 auto name = mlir::asc::fetchOpClass(def->getName());46 auto name = mlir::asc::fetchOpClass(def->getName());
47 std::vector<mlir::asc::VirtualArg> args;47 std::vector<mlir::asc::VirtualArg> args;
48 mlir::asc::fetchResults(def->getValueAsDag("results"), args);48 mlir::asc::fetchResults(def->getValueAsDag("results"), args);
49- bool retVal = args.size() == 1;49+ bool hasSingleResult = args.size() == 1 && args[0].cppType == "::mlir::Type";
50 mlir::asc::fetchArguments(def->getValueAsDag("arguments"), args);50 mlir::asc::fetchArguments(def->getValueAsDag("arguments"), args);
51 os << ".def(\"create_";51 os << ".def(\"create_";
52 auto dialectName = def->getValueAsDef("opDialect")->getValueAsString("name");52 auto dialectName = def->getValueAsDef("opDialect")->getValueAsString("name");
@@ -60,18 +60,21 @@ void printMethod(raw_indented_ostream& os, const Record* def)
60 os << ", const " << arg.cppType << " &" << arg.name;60 os << ", const " << arg.cppType << " &" << arg.name;
61 }61 }
62 os << ") ";62 os << ") ";
63- if (retVal) {63+ if (hasSingleResult) {
64 os << "-> Value ";64 os << "-> Value ";
65+ } else {
66+ os << "-> Operation* ";
65 }67 }
66 os << "{\n";68 os << "{\n";
67 os.indent();69 os.indent();
68- if (retVal) {70+ os << "return ";
69- os << "return ";
70- }
71 os << "self.create<" << def->getValueAsString("cppNamespace") << "::" << name << ">(";71 os << "self.create<" << def->getValueAsString("cppNamespace") << "::" << name << ">(";
72 interleaveComma(args, os, [&os](const auto& arg) { os << arg.substitution; });72 interleaveComma(args, os, [&os](const auto& arg) { os << arg.substitution; });
73 os << ");\n";73 os << ");\n";
74 os.unindent() << "}";74 os.unindent() << "}";
75+ if (!hasSingleResult) {
76+ os << ", py::return_value_policy::reference";
77+ }
75 auto lastRequired =78 auto lastRequired =
76 std::find_if(args.rbegin(), args.rend(), [](const mlir::asc::VirtualArg& arg) { return !arg.optional; });79 std::find_if(args.rbegin(), args.rend(), [](const mlir::asc::VirtualArg& arg) { return !arg.optional; });
77 std::for_each(lastRequired, args.rend(), [](auto& a) { a.optional = false; });80 std::for_each(lastRequired, args.rend(), [](auto& a) { a.optional = false; });
@@ -40,6 +40,8 @@ void fetchResults(const DagInit* resultsDag, std::vector<VirtualArg>& dest)
40 const auto* init = dyn_cast<DefInit>(resultsDag->getArg(i));40 const auto* init = dyn_cast<DefInit>(resultsDag->getArg(i));
41 assert(init && "argument must have defined types");41 assert(init && "argument must have defined types");
42 const auto* resultDef = init->getDef();42 const auto* resultDef = init->getDef();
43+ if (resultDef->isSubClassOf("Res"))
44+ resultDef = resultDef->getValueAsDef("constraint");
43 if (resultDef->isSubClassOf("Variadic")) {45 if (resultDef->isSubClassOf("Variadic")) {
44 result.cppType = "::std::vector< ::mlir::Type >";46 result.cppType = "::std::vector< ::mlir::Type >";
45 } else {47 } else {
@@ -65,6 +67,8 @@ void fetchArguments(const DagInit* argsDag, std::vector<VirtualArg>& dest)
65 const auto* init = dyn_cast<DefInit>(argsDag->getArg(i));67 const auto* init = dyn_cast<DefInit>(argsDag->getArg(i));
66 assert(init && "argument must have defined types");68 assert(init && "argument must have defined types");
67 const auto* argDef = init->getDef();69 const auto* argDef = init->getDef();
70+ if (argDef->isSubClassOf("Arg"))
71+ argDef = argDef->getValueAsDef("constraint");
68 if (argDef->isSubClassOf("TypeConstraint")) {72 if (argDef->isSubClassOf("TypeConstraint")) {
69 if (argDef->isSubClassOf("Variadic")) {73 if (argDef->isSubClassOf("Variadic")) {
70 arg.cppType = "::std::vector< ::mlir::Value >";74 arg.cppType = "::std::vector< ::mlir::Value >";
@@ -33,10 +33,19 @@ LogicalResult mlir::ascendc::printOperation(CodeEmitter& emitter, ascendc::SoftM
33 if (failed(emitter.emitType(op.getLoc(), op.getDst().getType().getElementType()))) {33 if (failed(emitter.emitType(op.getLoc(), op.getDst().getType().getElementType()))) {
34 return failure();34 return failure();
35 }35 }
36- os << ", " << op.getReuseSource() << ", " << op.getBasicBlock() << ", " << op.getDataFormatNZ() << ">("36+ os << ", " << op.getReuseSource() << ", " << op.getBasicBlock() << ">(";
37- << emitter.getOrCreateName(op.getDst()) << ", " << emitter.getOrCreateName(op.getSumTensor()) << ", "37+ os << emitter.getOrCreateName(op.getDst()) << ", ";
38- << emitter.getOrCreateName(op.getMaxTensor()) << ", " << emitter.getOrCreateName(op.getSrc()) << ", "38+ if (auto sumTensor = op.getSumTensor()) {
39- << emitter.getOrCreateName(op.getTiling());39+ os << emitter.getOrCreateName(sumTensor) << ", ";
40+ }
41+ if (auto maxTensor = op.getMaxTensor()) {
42+ os << emitter.getOrCreateName(maxTensor) << ", ";
43+ }
44+ os << emitter.getOrCreateName(op.getSrc()) << ", ";
45+ if (auto sharedTmpBuffer = op.getSharedTmpBuffer()) {
46+ os << emitter.getOrCreateName(sharedTmpBuffer) << ", ";
47+ }
48+ os << emitter.getOrCreateName(op.getTiling());
40 if (auto ssi = op.getSoftmaxShapeInfo()) {49 if (auto ssi = op.getSoftmaxShapeInfo()) {
41 os << ", " << emitter.getOrCreateName(ssi);50 os << ", " << emitter.getOrCreateName(ssi);
42 }51 }
@@ -0,0 +1,37 @@
1+/*
2+ * Copyright (c) 2026 Huawei Technologies Co., Ltd.
3+ * This program is free software, you can redistribute it and/or modify it under the terms and conditions of
4+ * 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,
7+ * 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+#include "ascir/Target/Asc/Adv/Broadcast.h"
12+#include "ascir/Target/Asc/Common.h"
13+ 
14+using namespace mlir;
15+using namespace mlir::ascendc;
16+ 
17+LogicalResult mlir::ascendc::printOperation(CodeEmitter& emitter, ascendc::BroadcastOp op)
18+{
19+ auto& os = emitter.ostream();
20+ assert(op.getSrcShape().size() == op.getDstShape().size());
21+ auto dims = op.getSrcShape().size();
22+ os << "{\n";
23+ os.indent() << "const uint32_t dstShape[" << dims << "] = {";
24+ llvm::interleaveComma(op.getDstShape(), os, [&](Value operand) { os << emitter.getOrCreateName(operand); });
25+ os << "};\n";
26+ os << "const uint32_t srcShape[" << dims << "] = {";
27+ llvm::interleaveComma(op.getSrcShape(), os, [&](Value operand) { os << emitter.getOrCreateName(operand); });
28+ os << "};\n";
29+ os << ascNamespace << "::BroadcastTiling bcTiling;\n";
30+ os << ascNamespace << "::GetBroadcastTilingInfo<";
31+ FAIL_OR(emitter.emitType(op.getLoc(), op.getDst().getType().getElementType()));
32+ os << ", " << dims << ">(" << dims << ", dstShape, srcShape, false, bcTiling);\n";
33+ os << ascNamespace << "::" << op.getAPIName() << "(" << emitter.getOrCreateName(op.getDst()) << ", "
34+ << emitter.getOrCreateName(op.getSrc()) << ", dstShape, srcShape, &bcTiling);\n";
35+ os.unindent() << "}";
36+ return success();
37+}
@@ -26,3 +26,25 @@ LogicalResult mlir::ascendc::printOperation(CodeEmitter& emitter, ascendc::RmsNo
26 os << ", " << emitter.getOrCreateName(op.getEpsilon()) << ", " << emitter.getOrCreateName(op.getTiling()) << ")";26 os << ", " << emitter.getOrCreateName(op.getEpsilon()) << ", " << emitter.getOrCreateName(op.getTiling()) << ")";
27 return success();27 return success();
28}28}
29+ 
30+LogicalResult mlir::ascendc::printOperation(CodeEmitter& emitter, ascendc::LayerNormOp op)
31+{
32+ auto& os = emitter.ostream();
33+ os << "{\n";
34+ os.indent();
35+ os << "static constexpr " << ascNamespace << "::LayerNormConfig lnConfig"
36+ << " = {.isNoBeta = false, .isNoGamma = false, .isOnlyOutput = false, .isOutputRstd = "
37+ << (op.getOutputRstd() ? "true" : "false") << "};\n";
38+ os << ascNamespace << "::" << op.getAPIName() << "<";
39+ FAIL_OR(emitter.emitType(op.getLoc(), op.getGamma().getType().getElementType()));
40+ os << ", ";
41+ FAIL_OR(emitter.emitType(op.getLoc(), op.getDst().getType().getElementType()));
42+ os << ", false, lnConfig>(" << emitter.getOrCreateName(op.getDst()) << ", "
43+ << emitter.getOrCreateName(op.getDstMean()) << ", " << emitter.getOrCreateName(op.getDstVarRstd()) << ", "
44+ << emitter.getOrCreateName(op.getSrc()) << ", " << emitter.getOrCreateName(op.getGamma()) << ", "
45+ << emitter.getOrCreateName(op.getBeta()) << ", " << emitter.getOrCreateName(op.getEpsilon()) << ", "
46+ << emitter.getOrCreateName(op.getSharedTmpBuffer()) << ", " << emitter.getOrCreateName(op.getPara()) << ", "
47+ << emitter.getOrCreateName(op.getSeparateTiling()) << ");\n";
48+ os.unindent() << "}\n";
49+ return success();
50+}
Msetup.py+76-146
Rtest/Dialect/AscendC/Transforms/hoist-ub-allocation.mlir→test/Dialect/AscendC/Transforms/hoist-tensor-allocation.mlir+13-13
Rtest/Dialect/AscendC/Transforms/insert-sync.mlir→test/Dialect/AscendC/Transforms/insert-que-sync.mlir+1-1