已合并
Upgrade the infrastructure to support future extensions #157
Maksim Vlasov创建于 8月11日
Upgrade the infrastructure to support future extensions #157
已合并
共 162 个文件变更+5279-1086
| @@ -8,13 +8,12 @@ __pycache__/ | |||
| 8 | 8 | ||
| 9 | # Distribution / packaging | 9 | # Distribution / packaging |
| 10 | .Python | 10 | .Python |
| 11 | -build/ | 11 | +build*/ |
| 12 | develop-eggs/ | 12 | develop-eggs/ |
| 13 | dist/ | 13 | dist/ |
| 14 | downloads/ | 14 | downloads/ |
| 15 | eggs/ | 15 | eggs/ |
| 16 | .eggs/ | 16 | .eggs/ |
| 17 | -lib/ | ||
| 18 | lib64/ | 17 | lib64/ |
| 19 | parts/ | 18 | parts/ |
| 20 | sdist/ | 19 | sdist/ |
| @@ -197,9 +196,9 @@ cython_debug/ | |||
| 197 | .abstra/ | 196 | .abstra/ |
| 198 | 197 | ||
| 199 | # Visual Studio Code | 198 | # 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.gitignore | 200 | # 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 folder | 202 | # you could uncomment the following to ignore the entire vscode folder |
| 204 | # .vscode/ | 203 | # .vscode/ |
| 205 | # Temporary file for partial code execution | 204 | # Temporary file for partial code execution |
| @@ -227,4 +226,4 @@ docs/python-api/rst/**/generated/ | |||
| 227 | .spec-workflow/ | 226 | .spec-workflow/ |
| 228 | docs/superpowers/ | 227 | docs/superpowers/ |
| 229 | AGENTS.md | 228 | AGENTS.md |
| 230 | -CLAUDE.md | 229 | +CLAUDE.md |
| @@ -1,6 +1,18 @@ | |||
| 1 | repos: | 1 | repos: |
| 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-format | 14 | - repo: https://github.com/pre-commit/mirrors-clang-format |
| 3 | - rev: v16.0.0 | 15 | + rev: v20.1.8 |
| 4 | hooks: | 16 | hooks: |
| 5 | - id: clang-format | 17 | - 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) |
| 15 | endif() | 15 | endif() |
| 16 | 16 | ||
| 17 | -project(AscIR LANGUAGES C CXX) | 17 | +project(AscIR LANGUAGES CXX) |
| 18 | 18 | ||
| 19 | -message(STATUS "C compiler: ${CMAKE_C_COMPILER}") | ||
| 20 | message(STATUS "C++ compiler: ${CMAKE_CXX_COMPILER}") | 19 | message(STATUS "C++ compiler: ${CMAKE_CXX_COMPILER}") |
| 21 | message(STATUS "Linker: ${CMAKE_LINKER}") | 20 | message(STATUS "Linker: ${CMAKE_LINKER}") |
| 22 | message(STATUS "Selected build configuration: ${CMAKE_BUILD_TYPE}") | 21 | message(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() |
| 101 | endif() | 100 | endif() |
| @@ -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 | ||
| 77 | set_target_properties(${ASCIR_OPT_NAME} PROPERTIES RUNTIME_OUTPUT_DIRECTORY "${CMAKE_BINARY_DIR}/bin") | 77 | set_target_properties(${ASCIR_OPT_NAME} PROPERTIES RUNTIME_OUTPUT_DIRECTORY "${CMAKE_BINARY_DIR}/bin") |
| 78 | 78 | ||
| 79 | -# ascir-translate | 79 | +# ascir-translate |
| 80 | 80 | ||
| 81 | set(ASCIR_TRANSLATE_NAME ascir-translate) | 81 | set(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 | - | ||
| 167 | def DeqScale : APIType<"DeqScale"> { | 162 | def 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 | + | ||
| 172 | def FixpipeConfig : APIType<"FixpipeConfig"> { | 172 | def 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 | + | ||
| 232 | def ListTensorDesc : APIType<"ListTensorDesc"> { | 247 | def 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 | + | ||
| 292 | def MaskMode : APIType<"MaskMode"> { | 312 | def 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 | + | ||
| 312 | def MmadParams : APIType<"MmadParams"> { | 337 | def 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_TD | 477 | #endif // API_TYPES_TD |
| @@ -20,8 +20,9 @@ include "mlir/Interfaces/CastInterfaces.td" | |||
| 20 | include "mlir/Interfaces/SideEffectInterfaces.td" | 20 | include "mlir/Interfaces/SideEffectInterfaces.td" |
| 21 | include "mlir/IR/OpBase.td" | 21 | include "mlir/IR/OpBase.td" |
| 22 | 22 | ||
| 23 | -def AscendC_SimpleSoftMaxOp | 23 | +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 | } |
| 45 | def AscendC_SwiGLUOp : APIOp<"swiglu", "SwiGLU", [AttrSizedOperandSegments, OpWithDstInterface]> { | 75 | def 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_TD | 83 | +#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">; | |||
| 40 | def FracOp : UnaryMathOp<"frac", "Frac">; | 40 | def FracOp : UnaryMathOp<"frac", "Frac">; |
| 41 | def LgammaOp : UnaryMathOp<"lgamma", "Lgamma">; | 41 | def LgammaOp : UnaryMathOp<"lgamma", "Lgamma">; |
| 42 | def LogOp : UnaryMathOp<"log", "Log">; | 42 | def LogOp : UnaryMathOp<"log", "Log">; |
| 43 | +def Log2Op : UnaryMathOp<"log2", "Log2">; | ||
| 43 | def RoundOp : UnaryMathOp<"round", "Round">; | 44 | def RoundOp : UnaryMathOp<"round", "Round">; |
| 44 | def SignOp : UnaryMathOp<"sign", "Sign">; | 45 | def SignOp : UnaryMathOp<"sign", "Sign">; |
| 45 | def SinhOp : UnaryMathOp<"sinh", "Sinh">; | 46 | def SinhOp : UnaryMathOp<"sinh", "Sinh">; |
| @@ -62,27 +63,77 @@ def XorOp : BinaryMathOp<"xor", "Xor">; | |||
| 62 | def AxpyOp : MathLibraryOp<"axpy", "Axpy"> { | 63 | def 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 | ||
| 67 | def ClampMaxOp : MathLibraryOp<"clamp_max", "ClampMax"> { | 78 | def 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 | ||
| 72 | def ClampMinOp : MathLibraryOp<"clamp_min", "ClampMin"> { | 93 | def 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 | ||
| 77 | def CumSumOp : MathLibraryOp<"cumsum", "CumSum"> { | 108 | def 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 | ||
| 83 | def ExpOp : MathLibraryOp<"exp", "Exp"> { | 124 | def 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_TD | 139 | +#endif // ASC_ADV_MATH_TD |
| @@ -86,7 +86,7 @@ def AscendC_MatmulSetSparseIndexOp | |||
| 86 | 86 | ||
| 87 | def AscendC_MatmulSetUserDefInfoOp | 87 | def 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 | ||
| 306 | def AscendC_MatmulSetOrgShapeOp : APIOp<"matmul.set_org_shape", "SetOrgShape", [AscMemberFunc, AttrSizedOperandSegments]> { | 306 | def AscendC_MatmulSetOrgShapeOp : APIOp<"matmul.set_org_shape", "SetOrgShape", [AscMemberFunc, AttrSizedOperandSegments]> { |
| @@ -20,10 +20,41 @@ include "mlir/Interfaces/CastInterfaces.td" | |||
| 20 | include "mlir/Interfaces/SideEffectInterfaces.td" | 20 | include "mlir/Interfaces/SideEffectInterfaces.td" |
| 21 | include "mlir/IR/OpBase.td" | 21 | include "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_TD | 60 | #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 = []> | |||
| 75 | class DataCopyOp<string mnemonic, string apiName, list<Trait> traits = []> | 75 | class 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 | + | ||
| 78 | class VectorOp<string mnemonic, string apiName, list<Trait> traits = []> | 81 | class 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 = []> | |||
| 85 | class UnaryOp<string mnemonic, string apiName, list<Trait> traits = []> | 88 | class 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 | ||
| 90 | class UnaryL0Op<string mnemonic, string apiName, list<Trait> traits = []> | 96 | class UnaryL0Op<string mnemonic, string apiName, list<Trait> traits = []> |
| @@ -122,13 +128,16 @@ multiclass UnaryL012Op<string mnemonic, string apiName, list<Trait> traits = []> | |||
| 122 | class BinaryOp<string mnemonic, string apiName, list<Trait> traits = []> | 128 | class 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 | ||
| 127 | class BinaryL0Op<string mnemonic, string apiName, list<Trait> traits = []> | 136 | class 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 | ||
| 143 | class BinaryL2Op<string mnemonic, string apiName, list<Trait> traits = []> | 152 | class 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 | ||
| 150 | class BinaryTemplateL0Op<string mnemonic, string apiName, list<Trait> traits = []> | 159 | class 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 | ||
| 166 | class BinaryTemplateL2Op<string mnemonic, string apiName, list<Trait> traits = []> | 175 | class 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 | ||
| 173 | class BinaryCastL0Op<string mnemonic, string apiName, list<Trait> traits = []> | 182 | class 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 | ||
| 189 | class BinaryCastL2Op<string mnemonic, string apiName, list<Trait> traits = []> | 198 | class 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 | |||
| 237 | class VecScalarOp<string mnemonic, string apiName, list<Trait> traits = []> | 246 | class 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 | ||
| 242 | class VecScalarL0Op<string mnemonic, string apiName, list<Trait> traits = []> | 254 | class 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 operations | 314 | // 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 | ||
| 287 | class BinaryMathOp<string mnemonic, string apiName, list<Trait> traits = []> | 339 | class 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_TD | 116 | #endif // ASC_BASIC_OP_BLOCKSYNC_TD |
| @@ -26,4 +26,17 @@ def AscendC_TransDataTo5HDTensorListOp : TransDataTo5HDTensorListOp<"trans_data_ | |||
| 26 | def AscendC_TransDataTo5HDUintListOp : TransDataTo5HDUintListOp<"trans_data_to_5hd_uint_list", "TransDataTo5HD">; | 26 | def AscendC_TransDataTo5HDUintListOp : TransDataTo5HDUintListOp<"trans_data_to_5hd_uint_list", "TransDataTo5HD">; |
| 27 | def AscendC_TransDataTo5HDOp : TransDataTo5HDSingleOp<"trans_data_to_5hd", "TransDataTo5HD">; | 27 | def 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_TD | 42 | #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 | ||
| 90 | def AscendC_DataCopyPadExtParamsOp : AscendC_Op<"data_copy_pad_ext_params", [AscConstructor]> { | 96 | def 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 | + | ||
| 142 | def AscendC_LoadImageToLocalOp : APIOp<"load_image_to_local", "LoadImageToLocal", [AscFunc]> { | 172 | def 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 | ||
| 150 | def AscendC_SetPadValueOp : APIOp<"set_pad_value", "SetPadValue", [AscFunc]> { | 178 | def 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_TD | 197 | #endif // ASC_BASIC_OP_DATACOPY_TD |
| @@ -20,11 +20,14 @@ include "mlir/Interfaces/CastInterfaces.td" | |||
| 20 | include "mlir/Interfaces/SideEffectInterfaces.td" | 20 | include "mlir/Interfaces/SideEffectInterfaces.td" |
| 21 | include "mlir/IR/OpBase.td" | 21 | include "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 | ||
| 30 | def AscendC_FixpipeWithWorkspaceOp : APIOp<"fixpipe_with_workspace", "Fixpipe"> { | 33 | def AscendC_FixpipeWithWorkspaceOp : APIOp<"fixpipe_with_workspace", "Fixpipe"> { |
| @@ -19,6 +19,12 @@ include "mlir/Interfaces/CastInterfaces.td" | |||
| 19 | include "mlir/Interfaces/SideEffectInterfaces.td" | 19 | include "mlir/Interfaces/SideEffectInterfaces.td" |
| 20 | include "mlir/IR/OpBase.td" | 20 | include "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 | + | ||
| 22 | def AscendC_InitConstValueOp | 28 | def 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_LoadDataOp | 37 | +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_LoadDataL0Op | 45 | +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_LoadDataG2LOp | 55 | +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 | ||
| 53 | def AscendC_LoadDataL0V2Op | 65 | def 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 | ||
| 61 | def AscendC_LoadDataG2LV2Op | 76 | def 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 | ||
| 69 | def AscendC_LoadData3DL0V1Op | 87 | def 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 | ||
| 77 | def AscendC_LoadData3DL0V2Op | 97 | def 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 | ||
| 85 | def AscendC_LoadData3DL0V2ProOp | 107 | def 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 | ||
| 93 | def AscendC_LoadDataWithSparseOp | 117 | def 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 | ||
| 99 | def AscendC_LoadDataWithTransposeOp | 127 | def 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 | ||
| 107 | def AscendC_LoadDataWithTransposeV2Op | 137 | def 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 | ||
| 115 | def AscendC_MmadOp | 147 | def 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 | ||
| 124 | def AscendC_MmadWithBiasOp | 159 | def 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 | ||
| 134 | def AscendC_MmadWithSparseOp | 172 | def 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 | + | ||
| 41 | def AscendC_GetDataBlockSizeInBytesOp : APIOp<"get_data_block_size_in_bytes", "GetDataBlockSizeInBytes", [AscFunc]> { | 48 | def 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" | |||
| 21 | include "mlir/IR/OpBase.td" | 21 | include "mlir/IR/OpBase.td" |
| 22 | 22 | ||
| 23 | defm Adds : VecScalarL012Op<"adds", "Adds">; | 23 | defm Adds : VecScalarL012Op<"adds", "Adds">; |
| 24 | +defm Divs : VecScalarL012Op<"divs", "Divs">; | ||
| 24 | defm LeakyRelu : VecScalarL012Op<"leaky_relu", "LeakyRelu">; | 25 | defm LeakyRelu : VecScalarL012Op<"leaky_relu", "LeakyRelu">; |
| 25 | defm Maxs : VecScalarL012Op<"maxs", "Maxs">; | 26 | defm Maxs : VecScalarL012Op<"maxs", "Maxs">; |
| 26 | defm Mins : VecScalarL012Op<"mins", "Mins">; | 27 | defm Mins : VecScalarL012Op<"mins", "Mins">; |
| 27 | defm Muls : VecScalarL012Op<"muls", "Muls">; | 28 | defm Muls : VecScalarL012Op<"muls", "Muls">; |
| 28 | defm ShiftLeft : VecScalarL012Op<"shift_left", "ShiftLeft">; | 29 | defm ShiftLeft : VecScalarL012Op<"shift_left", "ShiftLeft">; |
| 29 | defm ShiftRight : VecScalarL012Op<"shift_right", "ShiftRight">; | 30 | defm ShiftRight : VecScalarL012Op<"shift_right", "ShiftRight">; |
| 31 | +defm Subs : VecScalarL012Op<"subs", "Subs">; | ||
| 30 | 32 | ||
| 31 | #endif //ASC_BASIC_OP_VEC_BINARY_SCALAR_TD | 33 | #endif //ASC_BASIC_OP_VEC_BINARY_SCALAR_TD |
| @@ -20,26 +20,36 @@ include "mlir/Interfaces/CastInterfaces.td" | |||
| 20 | include "mlir/Interfaces/SideEffectInterfaces.td" | 20 | include "mlir/Interfaces/SideEffectInterfaces.td" |
| 21 | include "mlir/IR/OpBase.td" | 21 | include "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 | ||
| 31 | def AscendC_CompareL1Op : VectorOp<"compare_l1", "Compare", [OpWithDstInterface]> { | 36 | def 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 | ||
| 45 | def AscendC_CompareRL0Op : VectorOp<"compare_r_l0", "Compare", [AscFunc]> { | 55 | def 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 | ||
| 66 | def AscendC_CompareScalarL1Op : VectorOp<"compare_scalar_l1", "CompareScalar", [OpWithDstInterface]> { | 82 | def 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 | ||
| 80 | def AscendC_GetCmpMaskOp : VectorOp<"get_cmp_mask", "GetCmpMask", [OpWithDstInterface, AscFunc]> { | 101 | def 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 | ||
| 106 | def AscendC_SelectScalarL1Op : VectorOp<"select_scalar_l1", "Select", [OpWithDstInterface]> { | 141 | def 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_TD | 208 | + |
| 209 | +#endif // ASC_BASIC_OP_VEC_CMPSEL_TD | ||
| @@ -20,6 +20,6 @@ include "mlir/Interfaces/CastInterfaces.td" | |||
| 20 | include "mlir/Interfaces/SideEffectInterfaces.td" | 20 | include "mlir/Interfaces/SideEffectInterfaces.td" |
| 21 | include "mlir/IR/OpBase.td" | 21 | include "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_TD | 25 | #endif //ASC_BASIC_OP_MULCAST_TD |
| @@ -26,7 +26,7 @@ include "mlir/IR/OpBase.td" | |||
| 26 | 26 | ||
| 27 | def AscendC_BlockReduceSumL0Op : VectorOp<"block_reduce_sum_l0", "BlockReduceSum", [AscFunc]> { | 27 | def 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 | ||
| 41 | def AscendC_BlockReduceSumL1Op : VectorOp<"block_reduce_sum_l1", "BlockReduceSum", [OpWithDstInterface]> { | 41 | def 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 | ||
| 54 | def AscendC_BlockReduceMaxL0Op : VectorOp<"block_reduce_max_l0", "BlockReduceMax", [AscFunc]> { | 54 | def 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 | ||
| 68 | def AscendC_BlockReduceMaxL1Op : VectorOp<"block_reduce_max_l1", "BlockReduceMax", [OpWithDstInterface]> { | 68 | def 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 | ||
| 81 | def AscendC_BlockReduceMinL0Op : VectorOp<"block_reduce_min_l0", "BlockReduceMin", [AscFunc]> { | 81 | def 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 | ||
| 95 | def AscendC_BlockReduceMinL1Op : VectorOp<"block_reduce_min_l1", "BlockReduceMin", [OpWithDstInterface]> { | 95 | def 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 | ||
| 108 | def AscendC_PairReduceSumL0Op : VectorOp<"pair_reduce_sum_l0", "PairReduceSum", [AscFunc]> { | 108 | def 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 | ||
| 123 | def AscendC_PairReduceSumL1Op : VectorOp<"pair_reduce_sum_l1", "PairReduceSum", [OpWithDstInterface]> { | 123 | def 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 | ||
| 137 | def AscendC_RepeatReduceSumL0Op : VectorOp<"repeat_reduce_sum_l0", "RepeatReduceSum", [AscFunc]> { | 137 | def 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 | ||
| 153 | def AscendC_WholeReduceMaxL0Op : VectorOp<"whole_reduce_max_l0", "WholeReduceMax", [AscFunc]> { | 153 | def 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 | ||
| 169 | def AscendC_WholeReduceMaxL1Op : VectorOp<"whole_reduce_max_l1", "WholeReduceMax", [OpWithDstInterface]> { | 169 | def 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 | ||
| 184 | def AscendC_WholeReduceMinL0Op : VectorOp<"whole_reduce_min_l0", "WholeReduceMin", [AscFunc]> { | 184 | def 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 | ||
| 200 | def AscendC_WholeReduceMinL1Op : VectorOp<"whole_reduce_min_l1", "WholeReduceMin", [OpWithDstInterface]> { | 200 | def 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 | ||
| 215 | def AscendC_WholeReduceSumL0Op : VectorOp<"whole_reduce_sum_l0", "WholeReduceSum", [AscFunc]> { | 215 | def 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 | ||
| 230 | def AscendC_WholeReduceSumL1Op : VectorOp<"whole_reduce_sum_l1", "WholeReduceSum", [OpWithDstInterface]> { | 230 | def 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 operations | 245 | // 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:$calIndex | 259 | 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 | ||
| 262 | def AscendC_ReduceMaxL1Op : VectorOp<"reduce_max_l1", "ReduceMax", [OpWithDstInterface]> { | 268 | def 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:$calIndex | 290 | 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 operations | 300 | // 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:$calIndex | 314 | 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 | ||
| 305 | def AscendC_ReduceMinL1Op : VectorOp<"reduce_min_l1", "ReduceMin", [OpWithDstInterface]> { | 323 | def 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:$calIndex | 345 | 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 operations | 355 | // 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 | ||
| 347 | def AscendC_ReduceSumL1Op : VectorOp<"reduce_sum_l1", "ReduceSum", [OpWithDstInterface]> { | 377 | def 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:$count | 397 | 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_TD | 405 | + |
| 406 | +#endif // ASC_BASIC_OP_VEC_REDUCE_TD | ||
| @@ -20,12 +20,17 @@ include "mlir/Interfaces/CastInterfaces.td" | |||
| 20 | include "mlir/Interfaces/SideEffectInterfaces.td" | 20 | include "mlir/Interfaces/SideEffectInterfaces.td" |
| 21 | include "mlir/IR/OpBase.td" | 21 | include "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 | ||
| 31 | def AscendC_CastL1Op : VectorOp<"cast_l1", "Cast", [OpWithDstInterface]> { | 36 | def 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 | ||
| 45 | def AscendC_CastDeqL0Op : VectorOp<"cast_deq_l0", "CastDeq"> { | 55 | def 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 | - | ||
| 147 | def AscendC_CubeFormatAttr : I32EnumAttr<"CubeFormat", "matmul tensor format", [ | 136 | def 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 | ||
| 347 | def AscendC_MatmulPolicyAttr : I32EnumAttr<"MatmulPolicy", "", [ | 357 | def 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 | + | ||
| 358 | def AscendC_PipeAttr : I32EnumAttr<"Pipe", "", [ | 381 | def 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_TD | 602 | #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 | |||
| 109 | def AscendC_GlobalTensorSetShapeInfoOp : APIOp<"global_tensor.set_shape_info", "SetShapeInfo", [AscMemberFunc]> { | 109 | def 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 | ||
| 114 | def AscendC_GlobalTensorSetShapeInfoV2Op : APIOp<"global_tensor.set_shape_info_v2", "SetShapeInfo"> { | 114 | def 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 | ||
| 139 | def AscendC_GlobalTensorSetGlobalBufferOp : APIOp<"global_tensor.set_global_buffer", "SetGlobalBuffer", [AscMemberFunc]> { | 140 | def 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 | + | ||
| 168 | def AscendC_LocalTensorBracketOp | 176 | def 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 | + | ||
| 198 | def AscendC_LocalTensorGetPositionOp : APIOp<"local_tensor.get_position", "GetPosition", [AscMemberFunc]> { | 216 | def 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 | ||
| 244 | def AscendC_LocalTensorPrintOp | 265 | def 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 | ||
| 297 | def AscendC_LocalTensorSubIndexOp | 321 | def 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 | ||
| 308 | def AscendC_LocalTensorToFileOp | 333 | def 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 | + | ||
| 52 | def AscendC_GlobalTensor : AscendC_BaseTensorType<"GlobalTensor", "global_tensor"> { | 61 | def 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 | + | ||
| 70 | def AscendC_Matmul : AscendC_Type<"Matmul", "matmul"> { | 83 | def 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 | + | ||
| 173 | def AscendC_TBufPool : AscendC_Type<"TBufPool", "tbuf_pool"> { | 193 | def 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 | + | ||
| 187 | include "ascir/API/Types.td.inc" | 213 | include "ascir/API/Types.td.inc" |
| 188 | 214 | ||
| 189 | #endif // ASC_CORE_TYPES_TD | 215 | #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]>; | |||
| 18 | class GetMethod<string fieldName, string methodName, string returnType = "::mlir::Value"> | 18 | class 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 | + | ||
| 21 | class AscendC_InterfaceMethods { | 24 | class 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 | ||
| 68 | def OpWithDstInterface : AscendC_OpInterface<"OpWithDst"> { | 86 | def 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 | ||
| 73 | def VectorOpInterface : AscendC_OpInterface<"VectorOp", [APIOpInterface]> { | 111 | def 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 operations | 163 | // Vector unary operations |
| 107 | //===----------------------------------------------------------------------===// | 164 | //===----------------------------------------------------------------------===// |
| 108 | 165 | ||
| 109 | -def UnaryOpInterface | 166 | +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 | ||
| 115 | def UnaryL0OpInterface : AscendC_OpInterface<"UnaryL0Op", [UnaryOpInterface]> { | 173 | def 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 | ||
| 120 | def UnaryL2OpInterface : AscendC_OpInterface<"UnaryL2Op", [UnaryOpInterface]> { | 187 | def 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 operations | 193 | // Vector binary operations |
| 127 | //===----------------------------------------------------------------------===// | 194 | //===----------------------------------------------------------------------===// |
| 128 | 195 | ||
| 129 | -def BinaryOpInterface | 196 | +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 | + | ||
| 135 | def BinaryL3OpInterface : AscendC_OpInterface<"BinaryL3Op", [BinaryOpInterface]> { | 222 | def 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 operations | 227 | // Vector-scalar operations |
| 141 | //===----------------------------------------------------------------------===// | 228 | //===----------------------------------------------------------------------===// |
| 142 | 229 | ||
| 143 | -def VecScalarOpInterface | 230 | +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 | ||
| 149 | def VecScalarL0OpInterface : AscendC_OpInterface<"VecScalarL0Op", [VecScalarOpInterface]> { | 237 | def 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 | ||
| 154 | def VecScalarL2OpInterface : AscendC_OpInterface<"VecScalarL2Op", [VecScalarOpInterface]> { | 251 | def 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 operations | 276 | // Math library operations |
| 161 | //===----------------------------------------------------------------------===// | 277 | //===----------------------------------------------------------------------===// |
| 162 | 278 | ||
| 163 | -def MathLibraryOpInterface | 279 | +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_TD | 315 | #endif // ASC_INTERFACES_TD |
| @@ -12,11 +12,13 @@ | |||
| 12 | #define ASC_OPS_TD | 12 | #define ASC_OPS_TD |
| 13 | 13 | ||
| 14 | include "Adv/Activation.td" | 14 | include "Adv/Activation.td" |
| 15 | +include "Adv/Broadcast.td" | ||
| 15 | include "Adv/Kfc.td" | 16 | include "Adv/Kfc.td" |
| 16 | include "Adv/Math.td" | 17 | include "Adv/Math.td" |
| 17 | include "Adv/Matmul.td" | 18 | include "Adv/Matmul.td" |
| 18 | include "Adv/Normalization.td" | 19 | include "Adv/Normalization.td" |
| 19 | include "Adv/Quantization.td" | 20 | include "Adv/Quantization.td" |
| 21 | +include "Adv/Reduction.td" | ||
| 20 | include "Adv/Sort.td" | 22 | include "Adv/Sort.td" |
| 21 | include "Base.td" | 23 | include "Base.td" |
| 22 | include "Basic/Common.td" | 24 | include "Basic/Common.td" |
| @@ -37,6 +39,7 @@ include "Basic/OpLimits.td" | |||
| 37 | include "Basic/OpListTensor.td" | 39 | include "Basic/OpListTensor.td" |
| 38 | include "Basic/OpMm.td" | 40 | include "Basic/OpMm.td" |
| 39 | include "Basic/OpProposal.td" | 41 | include "Basic/OpProposal.td" |
| 42 | +include "Basic/OpReg.td" | ||
| 40 | include "Basic/OpScalar.td" | 43 | include "Basic/OpScalar.td" |
| 41 | include "Basic/OpSetAtomic.td" | 44 | include "Basic/OpSetAtomic.td" |
| 42 | include "Basic/OpSwapMem.td" | 45 | include "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 | + | ||
| 114 | def AscendC_FftsCrossCoreSyncOp : APIOp<"ffts_cross_core_sync", "FftsCrossCoreSync"> { | 142 | def 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 | ||
| 118 | def AscendC_SetFftsBaseAddrOp : APIOp<"set_ffts_base_addr", "SetFftsBaseAddr"> { | 146 | def AscendC_SetFftsBaseAddrOp : APIOp<"set_ffts_base_addr", "SetFftsBaseAddr"> { |
| 119 | - let arguments = (ins AnyMemRef); | 147 | + let arguments = (ins AnyRankedOrUnrankedMemRef); |
| 120 | } | 148 | } |
| 121 | 149 | ||
| 122 | def AscendC_PopStackBufferOp : APIOp<"pop_stack_buffer", "PopStackBuffer"> { | 150 | def AscendC_PopStackBufferOp : APIOp<"pop_stack_buffer", "PopStackBuffer"> { |
| @@ -126,17 +154,19 @@ def AscendC_PopStackBufferOp : APIOp<"pop_stack_buffer", "PopStackBuffer"> { | |||
| 126 | 154 | ||
| 127 | def AscendC_LocalTensorAutoOp : APIOp<"local_tensor_auto", "LocalTensorAuto"> { | 155 | def 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_TD | 172 | #endif // ASC_OPS_TD |
| @@ -25,11 +25,11 @@ std::unique_ptr<Pass> createDetectKernelTypePass(); | |||
| 25 | std::unique_ptr<Pass> createEraseSyncPass(); | 25 | std::unique_ptr<Pass> createEraseSyncPass(); |
| 26 | std::unique_ptr<Pass> createGenerateBoilerplatePass(); | 26 | std::unique_ptr<Pass> createGenerateBoilerplatePass(); |
| 27 | std::unique_ptr<Pass> createHoistQueBindPass(); | 27 | std::unique_ptr<Pass> createHoistQueBindPass(); |
| 28 | -std::unique_ptr<Pass> createHoistUBAllocationPass(); | 28 | +std::unique_ptr<Pass> createHoistTensorAllocationPass(bool excludeInOut = false); |
| 29 | std::unique_ptr<Pass> createInputOutputTensorPass(); | 29 | std::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); |
| 33 | std::unique_ptr<Pass> createNoopPass(); | 33 | std::unique_ptr<Pass> createNoopPass(); |
| 34 | std::unique_ptr<Pass> createPrivatizeFuncPass(); | 34 | std::unique_ptr<Pass> createPrivatizeFuncPass(); |
| 35 | std::unique_ptr<Pass> createUnifyPipePass(); | 35 | std::unique_ptr<Pass> createUnifyPipePass(); |
| @@ -14,92 +14,285 @@ | |||
| 14 | include "mlir/Pass/PassBase.td" | 14 | include "mlir/Pass/PassBase.td" |
| 15 | 15 | ||
| 16 | def DeclarePyStruct : Pass<"ascendc-declare-py-struct", "ModuleOp"> { | 16 | def 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 | ||
| 22 | def DefineCubeOnly : Pass<"ascendc-define-cube-only", "ModuleOp"> { | 29 | def 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 | ||
| 100 | def DetectEnableDebug : Pass<"ascendc-detect-enable-debug", "ModuleOp"> { | 40 | def 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_TD | 298 | #endif // ASC_PASSES_TD |
| @@ -17,12 +17,19 @@ namespace ascendc { | |||
| 17 | 17 | ||
| 18 | namespace attr { | 18 | namespace attr { |
| 19 | LITERAL aicore = "ascendc.aicore"; | 19 | LITERAL aicore = "ascendc.aicore"; |
| 20 | -LITERAL api = "ascendc.api"; | 20 | +LITERAL compilationArch = "asc.compilation_arch"; |
| 21 | -LITERAL compile_mix = "asc.compile_mix"; | ||
| 22 | LITERAL emitAsUnsigned = "ascendc.emit_as_unsigned"; | 21 | LITERAL emitAsUnsigned = "ascendc.emit_as_unsigned"; |
| 23 | LITERAL global = "ascendc.global"; | 22 | LITERAL global = "ascendc.global"; |
| 24 | -LITERAL enable_debug = "asc.enable_debug"; | 23 | +LITERAL enableDebug = "asc.enable_debug"; |
| 24 | +LITERAL kernelType = "asc.kernel_type"; | ||
| 25 | LITERAL matmulCubeOnly = "asc.matmul_cube_only"; | 25 | LITERAL 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 attr | 33 | } // namespace attr |
| 27 | 34 | ||
| 28 | } // namespace ascendc | 35 | } // namespace ascendc |
| @@ -11,6 +11,7 @@ | |||
| 11 | 11 | ||
| 12 | 12 | ||
| 13 | 13 | ||
| 14 | + | ||
| 14 | 15 | ||
| 15 | 16 | ||
| 16 | 17 | ||
| @@ -20,6 +21,12 @@ | |||
| 20 | namespace mlir { | 21 | namespace mlir { |
| 21 | namespace ascendc { | 22 | namespace 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 | + | ||
| 23 | template <typename OpT> | 30 | template <typename OpT> |
| 24 | struct HoistOpPattern : public OpRewritePattern<OpT> { | 31 | struct 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 | + | ||
| 47 | bool opPrecedes(Operation* lhs, Operation* rhs); | 60 | bool opPrecedes(Operation* lhs, Operation* rhs); |
| 48 | 61 | ||
| 49 | bool opPrecedes(Operation* lhs, Operation* rhs, DominanceInfo& di); | 62 | bool opPrecedes(Operation* lhs, Operation* rhs, DominanceInfo& di); |
| 50 | 63 | ||
| 51 | void registerInlinerInterfaces(DialectRegistry& registry); | 64 | void 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 ascendc | 82 | } // namespace ascendc |
| 54 | } // namespace mlir | 83 | } // 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 | 15 | ||
| 16 | 16 | ||
| 17 | 17 | ||
| 18 | + | ||
| 18 | 19 | ||
| 19 | 20 | ||
| 20 | 21 | ||
| @@ -15,8 +15,10 @@ include "Dialect.td" | |||
| 15 | include "Types.td" | 15 | include "Types.td" |
| 16 | 16 | ||
| 17 | include "mlir/Interfaces/CastInterfaces.td" | 17 | include "mlir/Interfaces/CastInterfaces.td" |
| 18 | +include "mlir/Interfaces/LoopLikeInterface.td" | ||
| 18 | include "mlir/Interfaces/SideEffectInterfaces.td" | 19 | include "mlir/Interfaces/SideEffectInterfaces.td" |
| 19 | include "mlir/Interfaces/ViewLikeInterface.td" | 20 | include "mlir/Interfaces/ViewLikeInterface.td" |
| 21 | +include "mlir/IR/AttrTypeBase.td" | ||
| 20 | include "mlir/IR/BuiltinAttributeInterfaces.td" | 22 | include "mlir/IR/BuiltinAttributeInterfaces.td" |
| 21 | 23 | ||
| 22 | class EmitAsc_Op<string mnemonic, list<Trait> traits = []> | 24 | class EmitAsc_Op<string mnemonic, list<Trait> traits = []> |
| @@ -40,7 +42,7 @@ def EmitAsc_CallOpaqueOp : EmitAsc_Op<"call_opaque"> { | |||
| 40 | 42 | ||
| 41 | def EmitAsc_CopyStructOp : EmitAsc_Op<"copy_struct"> { | 43 | def 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 | ||
| 54 | def EmitAsc_DereferenceOp : EmitAsc_Op<"dereference", [Pure]> { | 56 | def 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 | + | ||
| 61 | def EmitAsc_MemberOp : EmitAsc_Op<"member"> { | 97 | def 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 | ||
| 167 | def EmitAsc_VerbatimOp : EmitAsc_Op<"verbatim"> { | 204 | def 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 | + | ||
| 22 | def PyStruct : EmitAsc_Type<"PyStruct", "py_struct", [MemRefElementTypeInterface]> { | 29 | def 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 | + | ||
| 12 | + | ||
| 13 | + | ||
| 14 | + | ||
| 15 | + | ||
| 16 | + | ||
| 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 | + | ||
| @@ -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 | + | ||
| 12 | + | ||
| 13 | + | ||
| 14 | + | ||
| 15 | + | ||
| 16 | + | ||
| 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 | + | ||
| @@ -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 | + | ||
| 12 | + | ||
| 13 | + | ||
| 14 | + | ||
| 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 | + | ||
| @@ -24,8 +24,9 @@ template <typename UnaryMathOp> | |||
| 24 | auto printOperation(CodeEmitter& emitter, UnaryMathOp op) -> LogicalResultForT< | 24 | auto 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 { | |||
| 17 | namespace ascendc { | 17 | namespace ascendc { |
| 18 | 18 | ||
| 19 | LogicalResult printOperation(CodeEmitter& emitter, ascendc::RmsNormOp op); | 19 | LogicalResult printOperation(CodeEmitter& emitter, ascendc::RmsNormOp op); |
| 20 | +LogicalResult printOperation(CodeEmitter& emitter, ascendc::LayerNormOp op); | ||
| 20 | 21 | ||
| 21 | } // namespace ascendc | 22 | } // namespace ascendc |
| 22 | } // namespace mlir | 23 | } // 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 | + | ||
| 12 | + | ||
| 13 | + | ||
| 14 | + | ||
| 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 | + | ||
| @@ -28,6 +28,10 @@ LogicalResult printOperation(CodeEmitter& emitter, ascendc::CrossCoreSetFlagOp o | |||
| 28 | 28 | ||
| 29 | LogicalResult printOperation(CodeEmitter& emitter, ascendc::CrossCoreWaitFlagOp op); | 29 | LogicalResult 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 ascendc | 35 | } // namespace ascendc |
| 32 | } // namespace mlir | 36 | } // namespace mlir |
| 33 | 37 | ||
| @@ -26,6 +26,8 @@ LogicalResult printOperation(CodeEmitter& emitter, ascendc::TransDataTo5HDUintLi | |||
| 26 | 26 | ||
| 27 | LogicalResult printOperation(CodeEmitter& emitter, ascendc::TransDataTo5HDOp op); | 27 | LogicalResult printOperation(CodeEmitter& emitter, ascendc::TransDataTo5HDOp op); |
| 28 | 28 | ||
| 29 | +LogicalResult printOperation(CodeEmitter& emitter, ascendc::TransDataTo5HDTensorOp op); | ||
| 30 | + | ||
| 29 | } // namespace ascendc | 31 | } // namespace ascendc |
| 30 | } // namespace mlir | 32 | } // namespace mlir |
| 31 | 33 | ||
| @@ -26,6 +26,8 @@ LogicalResult printOperation(CodeEmitter& emitter, ascendc::CopyL0Op op); | |||
| 26 | 26 | ||
| 27 | LogicalResult printOperation(CodeEmitter& emitter, ascendc::CopyL1Op op); | 27 | LogicalResult printOperation(CodeEmitter& emitter, ascendc::CopyL1Op op); |
| 28 | 28 | ||
| 29 | +LogicalResult printOperation(CodeEmitter& emitter, ascendc::NdDmaParamsOp op); | ||
| 30 | + | ||
| 29 | } // namespace ascendc | 31 | } // namespace ascendc |
| 30 | } // namespace mlir | 32 | } // namespace mlir |
| 31 | 33 | ||
| @@ -56,21 +56,24 @@ LogicalResult printOperation(CodeEmitter& emitter, ascendc::FixpipeWithWorkspace | |||
| 56 | 56 | ||
| 57 | LogicalResult printOperation(CodeEmitter& emitter, ascendc::GetStoreAtomicConfigOp op); | 57 | LogicalResult 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 | + | ||
| 59 | template <typename FixpipeOp> | 65 | template <typename FixpipeOp> |
| 60 | auto printFixpipeTemplate(CodeEmitter& emitter, FixpipeOp op) | 66 | auto 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 | + | ||
| 12 | + | ||
| 13 | + | ||
| 14 | + | ||
| 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 | + | ||
| @@ -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 | - | ||
| 68 | template <typename BinaryL2Op> | 49 | template <typename BinaryL2Op> |
| 69 | auto printOperation(CodeEmitter& emitter, BinaryL2Op op) -> LogicalResultForT< | 50 | auto 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 | ||
| 107 | template <typename BinaryCastL0Op> | 88 | template <typename BinaryCastL0Op> |
| 108 | auto printOperation(CodeEmitter& emitter, BinaryCastL0Op op) -> LogicalResultForT< | 89 | auto 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 | ||
| 117 | template <typename BinaryCastL1Op> | 99 | template <typename BinaryCastL1Op> |
| 118 | auto printOperation(CodeEmitter& emitter, BinaryCastL1Op op) -> LogicalResultForT< | 100 | auto 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 | ||
| 23 | template <typename VecScalarL0Op> | 23 | template <typename VecScalarL0Op> |
| 24 | auto printOperation(CodeEmitter& emitter, VecScalarL0Op op) -> LogicalResultForT< | 24 | auto 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 | ||
| 36 | template <typename VecScalarL1Op> | 36 | template <typename VecScalarL1Op> |
| 37 | auto printOperation(CodeEmitter& emitter, VecScalarL1Op op) -> LogicalResultForT< | 37 | auto 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 | ||
| 50 | template <typename VecScalarL2Op> | 50 | template <typename VecScalarL2Op> |
| 51 | auto printOperation(CodeEmitter& emitter, VecScalarL2Op op) -> LogicalResultForT< | 51 | auto 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 | ||
| 23 | LogicalResult printOperation(CodeEmitter& emitter, ascendc::LocalTensorV2Op op); | 23 | LogicalResult printOperation(CodeEmitter& emitter, ascendc::LocalTensorV2Op op); |
| 24 | 24 | ||
| 25 | +LogicalResult printOperation(CodeEmitter& emitter, ascendc::LocalTensorV3Op op); | ||
| 26 | + | ||
| 25 | LogicalResult printOperation(CodeEmitter& emitter, ascendc::LocalTensorReinterpretCastOp op); | 27 | LogicalResult printOperation(CodeEmitter& emitter, ascendc::LocalTensorReinterpretCastOp op); |
| 26 | 28 | ||
| 27 | LogicalResult printOperation(CodeEmitter& emitter, ascendc::LocalTensorSubIndexOp op); | 29 | LogicalResult printOperation(CodeEmitter& emitter, ascendc::LocalTensorSubIndexOp op); |
| 28 | 30 | ||
| 29 | LogicalResult printOperation(CodeEmitter& emitter, ascendc::LocalTensorBracketOp op); | 31 | LogicalResult printOperation(CodeEmitter& emitter, ascendc::LocalTensorBracketOp op); |
| 30 | 32 | ||
| 33 | +LogicalResult printOperation(CodeEmitter& emitter, ascendc::LocalTensorGetPhyAddrV2Op op); | ||
| 34 | + | ||
| 31 | } // namespace ascendc | 35 | } // namespace ascendc |
| 32 | } // namespace mlir | 36 | } // namespace mlir |
| 33 | 37 | ||
| @@ -17,10 +17,6 @@ | |||
| 17 | namespace mlir { | 17 | namespace mlir { |
| 18 | namespace emitasc { | 18 | namespace emitasc { |
| 19 | 19 | ||
| 20 | -//===----------------------------------------------------------------------===// | ||
| 21 | -// EmitAsc operations | ||
| 22 | -//===----------------------------------------------------------------------===// | ||
| 23 | - | ||
| 24 | LogicalResult printOperation(CodeEmitter& emitter, emitasc::CallOpaqueOp op); | 20 | LogicalResult printOperation(CodeEmitter& emitter, emitasc::CallOpaqueOp op); |
| 25 | 21 | ||
| 26 | LogicalResult printOperation(CodeEmitter& emitter, emitasc::CopyStructOp op); | 22 | LogicalResult printOperation(CodeEmitter& emitter, emitasc::CopyStructOp op); |
| @@ -29,6 +25,10 @@ LogicalResult printOperation(CodeEmitter& emitter, emitasc::DeclarePyStructOp op | |||
| 29 | 25 | ||
| 30 | LogicalResult printOperation(CodeEmitter& emitter, emitasc::DereferenceOp op); | 26 | LogicalResult 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 | + | ||
| 32 | LogicalResult printOperation(CodeEmitter& emitter, emitasc::MemberOp op); | 32 | LogicalResult printOperation(CodeEmitter& emitter, emitasc::MemberOp op); |
| 33 | 33 | ||
| 34 | LogicalResult printOperation(CodeEmitter& emitter, emitasc::MemberPtrOp op); | 34 | LogicalResult printOperation(CodeEmitter& emitter, emitasc::MemberPtrOp op); |
| @@ -100,6 +100,8 @@ LogicalResult printOperation(CodeEmitter& emitter, arith::SelectOp op); | |||
| 100 | 100 | ||
| 101 | LogicalResult printOperation(CodeEmitter& emitter, arith::IndexCastOp op); | 101 | LogicalResult printOperation(CodeEmitter& emitter, arith::IndexCastOp op); |
| 102 | 102 | ||
| 103 | +LogicalResult printOperation(CodeEmitter& emitter, arith::NegFOp op); | ||
| 104 | + | ||
| 103 | } // namespace mlir | 105 | } // namespace mlir |
| 104 | 106 | ||
| 105 | 107 | ||
| @@ -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 | ||
| 9 | add_subdirectory(Dialect) | 9 | add_subdirectory(Dialect) |
| 10 | -# add_subdirectory(TableGen) | ||
| 11 | add_subdirectory(Target) | 10 | add_subdirectory(Target) |
| @@ -9,10 +9,11 @@ | |||
| 9 | */ | 9 | */ |
| 10 | 10 | ||
| 11 | 11 | ||
| 12 | + | ||
| 12 | 13 | ||
| 13 | 14 | ||
| 15 | + | ||
| 14 | 16 | ||
| 15 | - | ||
| 16 | 17 | ||
| 17 | 18 | ||
| 18 | 19 | ||
| @@ -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 | // LocalTensorOp | 80 | // 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 | // PipeBarrierOp | 100 | // PipeBarrierOp |
| 56 | //===----------------------------------------------------------------------===// | 101 | //===----------------------------------------------------------------------===// |
| @@ -75,7 +120,7 @@ LogicalResult PipeBarrierOp::canonicalize(PipeBarrierOp op, PatternRewriter& rew | |||
| 75 | } | 120 | } |
| 76 | 121 | ||
| 77 | //===----------------------------------------------------------------------===// | 122 | //===----------------------------------------------------------------------===// |
| 78 | -// ReinterpretCastOp | 123 | +// LocalTensorReinterpretCastOp |
| 79 | //===----------------------------------------------------------------------===// | 124 | //===----------------------------------------------------------------------===// |
| 80 | 125 | ||
| 81 | bool LocalTensorReinterpretCastOp::areCastCompatible(TypeRange inputs, TypeRange outputs) | 126 | bool LocalTensorReinterpretCastOp::areCastCompatible(TypeRange inputs, TypeRange outputs) |
| @@ -87,7 +132,35 @@ bool LocalTensorReinterpretCastOp::areCastCompatible(TypeRange inputs, TypeRange | |||
| 87 | OpFoldResult LocalTensorReinterpretCastOp::fold([[maybe_unused]] FoldAdaptor adaptor) | 132 | OpFoldResult 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.cpp | 13 | EraseSync.cpp |
| 14 | GenerateBoilerplatePass.cpp | 14 | GenerateBoilerplatePass.cpp |
| 15 | HoistQueBind.cpp | 15 | HoistQueBind.cpp |
| 16 | - HoistUBAllocation.cpp | 16 | + HoistTensorAllocation.cpp |
| 17 | InputOutputTensor.cpp | 17 | InputOutputTensor.cpp |
| 18 | - InsertSync.cpp | 18 | + InsertQueSync.cpp |
| 19 | LegalizeKernelArgs.cpp | 19 | LegalizeKernelArgs.cpp |
| 20 | MaterializeTensor.cpp | 20 | MaterializeTensor.cpp |
| 21 | Noop.cpp | 21 | Noop.cpp |
| @@ -28,7 +28,6 @@ namespace ascendc { | |||
| 28 | } // namespace mlir | 28 | } // namespace mlir |
| 29 | 29 | ||
| 30 | using namespace mlir; | 30 | using namespace mlir; |
| 31 | -using namespace mlir::ascendc; | ||
| 32 | 31 | ||
| 33 | namespace { | 32 | namespace { |
| 34 | 33 | ||
| @@ -92,8 +91,4 @@ public: | |||
| 92 | 91 | ||
| 93 | } // namespace | 92 | } // 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 mlir | 25 | } // namespace mlir |
| 26 | 26 | ||
| 27 | using namespace mlir; | 27 | using namespace mlir; |
| 28 | -using namespace mlir::ascendc; | ||
| 29 | 28 | ||
| 30 | namespace { | 29 | namespace { |
| 31 | 30 | ||
| @@ -42,8 +41,4 @@ public: | |||
| 42 | 41 | ||
| 43 | } // namespace | 42 | } // 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 mlir | 24 | } // namespace mlir |
| 25 | 25 | ||
| 26 | using namespace mlir; | 26 | using namespace mlir; |
| 27 | -using namespace mlir::ascendc; | ||
| 28 | 27 | ||
| 29 | namespace { | 28 | namespace { |
| 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 | |||
| 32 | public: | 32 | public: |
| 33 | void runOnOperation() override | 33 | 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 | } // namespace | 62 | } // 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 | } // namespace | 61 | } // 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 mlir | 24 | } // namespace mlir |
| 25 | 25 | ||
| 26 | using namespace mlir; | 26 | using namespace mlir; |
| 27 | -using namespace mlir::ascendc; | ||
| 28 | 27 | ||
| 29 | namespace { | 28 | namespace { |
| 30 | 29 | ||
| @@ -55,8 +54,7 @@ public: | |||
| 55 | 54 | ||
| 56 | } // namespace | 55 | } // 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 ascendc | 60 | +} |
| 62 | -} // namespace mlir | ||
| @@ -47,8 +47,4 @@ public: | |||
| 47 | 47 | ||
| 48 | } // namespace | 48 | } // 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 | ||
| 18 | namespace mlir { | 18 | namespace mlir { |
| 19 | namespace ascendc { | 19 | namespace ascendc { |
| 20 | -#define GEN_PASS_DEF_HOISTUBALLOCATION | 20 | +#define GEN_PASS_DEF_HOISTTENSORALLOCATION |
| 21 | 21 | ||
| 22 | } // namespace ascendc | 22 | } // namespace ascendc |
| 23 | } // namespace mlir | 23 | } // namespace mlir |
| @@ -26,18 +26,26 @@ using namespace mlir; | |||
| 26 | 26 | ||
| 27 | namespace { | 27 | namespace { |
| 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() override | 40 | 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 | } // namespace | 55 | } // 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 ascendc | 60 | + options.excludeInOut = excludeInOut; |
| 53 | -} // namespace mlir | 61 | + return std::make_unique<HoistTensorAllocationPass>(options); |
| 62 | +} | ||
| @@ -25,57 +25,70 @@ namespace ascendc { | |||
| 25 | } // namespace mlir | 25 | } // namespace mlir |
| 26 | 26 | ||
| 27 | using namespace mlir; | 27 | using namespace mlir; |
| 28 | -using namespace mlir::ascendc; | ||
| 29 | 28 | ||
| 30 | namespace { | 29 | namespace { |
| 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 | + | ||
| 51 | void setInOutTensors(func::FuncOp funcOp) | 66 | void 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 | ||
| 76 | void fixInOutTensor(func::FuncOp& funcOp) | 89 | void 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() override | 114 | 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 | } // namespace | 126 | } // 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 | ||
| @@ -21,7 +21,7 @@ | |||
| 21 | 21 | ||
| 22 | namespace mlir { | 22 | namespace mlir { |
| 23 | namespace ascendc { | 23 | namespace ascendc { |
| 24 | -#define GEN_PASS_DEF_INSERTSYNC | 24 | +#define GEN_PASS_DEF_INSERTQUESYNC |
| 25 | 25 | ||
| 26 | } // namespace ascendc | 26 | } // namespace ascendc |
| 27 | } // namespace mlir | 27 | } // 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> { |
| 175 | public: | 175 | public: |
| 176 | void runOnOperation() override | 176 | void runOnOperation() override |
| 177 | { | 177 | { |
| @@ -192,8 +192,4 @@ public: | |||
| 192 | 192 | ||
| 193 | } // namespace | 193 | } // 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 | ||
| 71 | struct LegalizeKernelArgsPass : public ascendc::impl::LegalizeKernelArgsBase<LegalizeKernelArgsPass> { | 73 | struct LegalizeKernelArgsPass : public ascendc::impl::LegalizeKernelArgsBase<LegalizeKernelArgsPass> { |
| 74 | + LegalizeKernelArgsPass(const ascendc::LegalizeKernelArgsOptions& options) : LegalizeKernelArgsBase(options) {} | ||
| 75 | + | ||
| 72 | void runOnOperation() override | 76 | 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 | } // namespace | 87 | } // 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 ascendc | 92 | + options.setFftsAddr = setFftsAddr; |
| 89 | -} // namespace mlir | 93 | + 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 | - | ||
| 12 | - | ||
| 13 | 11 | ||
| 14 | 12 | ||
| 13 | + | ||
| 15 | 14 | ||
| 16 | 15 | ||
| 17 | 16 | ||
| @@ -32,7 +31,9 @@ using namespace mlir; | |||
| 32 | namespace { | 31 | namespace { |
| 33 | 32 | ||
| 34 | struct MaterializeLocalTensor : OpRewritePattern<ascendc::LocalTensorAutoOp> { | 33 | struct 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() override | 95 | 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 | } // namespace | 109 | } // 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 ascendc | 114 | + options.alwaysBuf = alwaysBuf; |
| 101 | -} // namespace mlir | 115 | + return std::make_unique<MaterializeTensorPass>(options); |
| 116 | +} | ||
| @@ -32,8 +32,4 @@ struct NoopPass : public ascendc::impl::NoopBase<NoopPass> { | |||
| 32 | 32 | ||
| 33 | } // namespace | 33 | } // 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 | } // namespace | 41 | } // 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 | 11 | ||
| 12 | 12 | ||
| 13 | - | ||
| 14 | 13 | ||
| 15 | 14 | ||
| 16 | 15 | ||
| @@ -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 | } // namespace | 47 | } // 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 | } // namespace | 158 | } // 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 | 11 | ||
| 12 | + | ||
| 12 | 13 | ||
| 13 | 14 | ||
| 14 | 15 | ||
| 16 | + | ||
| 15 | 17 | ||
| 16 | 18 | ||
| 17 | 19 | ||
| 18 | 20 | ||
| 21 | + | ||
| 19 | 22 | ||
| 20 | namespace mlir { | 23 | namespace mlir { |
| 21 | 24 | ||
| @@ -24,6 +27,73 @@ using AllowInline = ascir::AllowlistInlinerInterface<T...>; | |||
| 24 | 27 | ||
| 25 | namespace ascendc { | 28 | namespace 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 | + | ||
| 27 | bool opPrecedes(Operation* lhs, Operation* rhs) { return lhs != rhs && lhs->isBeforeInBlock(rhs); } | 97 | bool opPrecedes(Operation* lhs, Operation* rhs) { return lhs != rhs && lhs->isBeforeInBlock(rhs); } |
| 28 | 98 | ||
| 29 | bool opPrecedes(Operation* lhs, Operation* rhs, DominanceInfo& di) | 99 | bool 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 | ||
| 45 | void registerInlinerInterfaces(DialectRegistry& registry) | 125 | void 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 ascendc | 215 | } // namespace ascendc |
| 56 | } // namespace mlir | 216 | } // namespace mlir |
| @@ -20,6 +20,96 @@ | |||
| 20 | using namespace mlir; | 20 | using namespace mlir; |
| 21 | using namespace mlir::emitasc; | 21 | using 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 | // PtrOffsetOp | 114 | // 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 | // EmitAscDialect | 163 | // 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 | + | ||
| 12 | + | ||
| 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 | +} | ||
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