diff --git a/lib/PTO/Transforms/CMakeLists.txt b/lib/PTO/Transforms/CMakeLists.txt index 9419efe865..5c19c403cc 100644 --- a/lib/PTO/Transforms/CMakeLists.txt +++ b/lib/PTO/Transforms/CMakeLists.txt @@ -40,6 +40,15 @@ add_mlir_dialect_library(PTOTransforms TileFusion/PTOFlattenFusionRegion.cpp VPTOLLVMEmitter.cpp VPTOCANN900LLVMEmitter.cpp + VPTOCANN900LLVMEmitterArithmeticPatterns.cpp + VPTOCANN900LLVMEmitterCalleeCore.cpp + VPTOCANN900LLVMEmitterCalleeMemory.cpp + VPTOCANN900LLVMEmitterMemoryPatterns.cpp + VPTOCANN900LLVMEmitterPacking.cpp + VPTOCANN900LLVMEmitterPipeline.cpp + VPTOCANN900LLVMEmitterScalarPatterns.cpp + VPTOCANN900LLVMEmitterTypeHelpers.cpp + VPTOCANN900LLVMEmitterTypePatterns.cpp VPTOLLVMEmitterDispatcher.cpp VPTOLLVMEmitterHelper.cpp VPTOPtrNormalize.cpp diff --git a/lib/PTO/Transforms/VPTOBufferMaterialization.cpp b/lib/PTO/Transforms/VPTOBufferMaterialization.cpp index 37fe569f6d..80b000f0f6 100644 --- a/lib/PTO/Transforms/VPTOBufferMaterialization.cpp +++ b/lib/PTO/Transforms/VPTOBufferMaterialization.cpp @@ -28,11 +28,8 @@ static AddressSpaceAttr getNormalizedPtrMemorySpace(Attribute memorySpace, return AddressSpaceAttr::get(context, AddressSpace::GM); } -static Value materializeMemRefView(Value value, ArrayRef shape, - Type elementType, Attribute memorySpace, +static Value materializeMemRefView(Value value, MemRefType memrefType, PatternRewriter &rewriter, Location loc) { - auto memrefType = - MemRefType::get(shape, elementType, AffineMap(), memorySpace); if (value.getType() == memrefType) { return value; } @@ -53,9 +50,10 @@ static Value materializeTileBufferView(Value value, PatternRewriter &rewriter, return {}; } - return materializeMemRefView(value, tileType.getShape(), - tileType.getElementType(), - tileType.getMemorySpace(), rewriter, loc); + auto memrefType = + MemRefType::get(tileType.getShape(), tileType.getElementType(), + AffineMap(), tileType.getMemorySpace()); + return materializeMemRefView(value, memrefType, rewriter, loc); } } // namespace diff --git a/lib/PTO/Transforms/VPTOCANN900LLVMEmitter.cpp b/lib/PTO/Transforms/VPTOCANN900LLVMEmitter.cpp index db85c0cfef..cd5093778c 100644 --- a/lib/PTO/Transforms/VPTOCANN900LLVMEmitter.cpp +++ b/lib/PTO/Transforms/VPTOCANN900LLVMEmitter.cpp @@ -6,11994 +6,17 @@ // INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE. // See LICENSE in the root of the software repository for the full text of the License. -// https://discourse.llvm.org/t/matchandrewrite-hiding-virtual-functions/84933/8 -#pragma GCC diagnostic ignored "-Woverloaded-virtual" -#include "PTO/Transforms/VPTOLLVMEmitter.h" -#include "PTO/Transforms/VPTOLLVMEmitterHelper.h" +#include "VPTOCANN900LLVMEmitterInternal.h" -#include "PTO/IR/PTO.h" -#include "PTO/IR/PTOTypeUtils.h" -#include "PTO/IR/PTOSyncUtils.h" -#include "PTO/Transforms/Passes.h" - -#include "mlir/Conversion/Passes.h" -#include "mlir/Conversion/ReconcileUnrealizedCasts/ReconcileUnrealizedCasts.h" -#include "mlir/Conversion/SCFToControlFlow/SCFToControlFlow.h" -#include "mlir/Dialect/Arith/IR/Arith.h" -#include "mlir/Dialect/Arith/Transforms/Passes.h" -#include "mlir/Dialect/ControlFlow/IR/ControlFlowOps.h" -#include "mlir/Dialect/Func/IR/FuncOps.h" -#include "mlir/Dialect/Func/Transforms/FuncConversions.h" -#include "mlir/Dialect/LLVMIR/LLVMDialect.h" -#include "mlir/Dialect/SCF/Transforms/Patterns.h" -#include "mlir/IR/Builders.h" -#include "mlir/IR/BuiltinOps.h" -#include "mlir/IR/PatternMatch.h" -#include "mlir/Pass/Pass.h" -#include "mlir/Pass/PassManager.h" -#include "mlir/Transforms/DialectConversion.h" -#include "mlir/Target/LLVMIR/Dialect/Builtin/BuiltinToLLVMIRTranslation.h" -#include "mlir/Target/LLVMIR/Dialect/LLVMIR/LLVMToLLVMIRTranslation.h" -#include "mlir/Target/LLVMIR/Export.h" -#include "llvm/ADT/DenseMap.h" -#include "llvm/ADT/STLExtras.h" -#include "llvm/ADT/StringSet.h" -#include "llvm/Bitcode/BitcodeWriter.h" -#include "llvm/IR/Constants.h" -#include "llvm/IR/Function.h" -#include "llvm/IR/Instructions.h" -#include "llvm/IR/LLVMContext.h" -#include "llvm/IR/GlobalVariable.h" -#include "llvm/Support/raw_ostream.h" -#include "llvm/Transforms/Utils/ModuleUtils.h" - -namespace mlir::pto { - -void materializeVecScopeCarrierLoops(ModuleOp module); -LogicalResult applyQueriedTargetAttrs(ModuleOp module, - const VPTOEmissionOptions &options, - llvm::raw_ostream &diagOS); -LogicalResult attachAIVectorScopeMetadata(llvm::Module &llvmModule, - llvm::raw_ostream &diagOS); -void attachHIVMKernelAnnotations(llvm::Module &llvmModule, - ModuleOp sourceModule); - -namespace { - -constexpr llvm::StringLiteral kVectorSuffix = "_mix_aiv"; -constexpr llvm::StringLiteral kCubeSuffix = "_mix_aic"; - -static std::string getElementTypeFragment(Type type); -static Type getElementTypeFromVectorLike(Type type); -static std::optional getElementCountFromVectorLike(Type type); - -static Type getLowPrecisionLLVMType(Type type, MLIRContext *context) { - if (pto::isPTOHiFloat8Type(type)) { - return LLVM::LLVMHiFloat8Type::get(context); - } - if (isa(type)) { - return LLVM::LLVMFloat4E1M2x2Type::get(context); - } - if (isa(type)) { - return LLVM::LLVMFloat4E2M1x2Type::get(context); - } - if (pto::isPTOFloat8E4M3LikeType(type)) { - return LLVM::LLVMFloat8E4M3Type::get(context); - } - if (pto::isPTOFloat8E5M2LikeType(type)) { - return LLVM::LLVMFloat8E5M2Type::get(context); - } - return {}; -} - -static bool isLLVMExtensionVectorElementType(Type type) { - return isa(type); -} - -static Type getLLVMCompatibleVectorType(ArrayRef shape, - Type elementType, - ArrayRef scalableDims = {}) { - if (shape.size() == 1 && isLLVMExtensionVectorElementType(elementType)) - return LLVM::LLVMFixedVectorType::get(elementType, shape.front()); - return VectorType::get(shape, elementType, scalableDims); -} - -static Type normalizePayloadTypeForLLVMLowering(Type type, Builder &builder) { - if (pto::isPTOHiFloat8x2Type(type)) - return getLLVMCompatibleVectorType( - {2}, LLVM::LLVMHiFloat8Type::get(builder.getContext())); - // bf16x2 is a 4-byte packed pair; lower it as an opaque i32 so vregs whose - // element type is bf16x2 get a valid LLVM type. - if (pto::isPTOBF16x2Type(type)) { - return builder.getI32Type(); - } - if (Type lowpType = getLowPrecisionLLVMType(type, builder.getContext())) { - return lowpType; - } - - if (auto intType = dyn_cast(type)) { - if (!intType.isSignless()) { - return builder.getIntegerType(intType.getWidth()); - } - return type; - } - - if (auto vecType = dyn_cast(type)) { - Type normalizedElement = - normalizePayloadTypeForLLVMLowering(vecType.getElementType(), builder); - if (normalizedElement == vecType.getElementType()) { - return type; - } - return getLLVMCompatibleVectorType(vecType.getShape(), normalizedElement, - vecType.getScalableDims()); - } - - return type; -} - -static Type normalizeGEPElementTypeForLLVMLowering(Type type, - Builder &builder) { - if (pto::isPTOHiFloat8x2Type(type)) { - return builder.getI16Type(); - } - // bf16x2 is 4 bytes, not an 8-bit low-precision type. - if (pto::isPTOBF16x2Type(type)) { - return builder.getI32Type(); - } - if (pto::isPTOLowPrecisionType(type)) { - return builder.getI8Type(); - } - if (isa(type)) - return builder.getI8Type(); - - if (auto vecType = dyn_cast(type)) { - Type normalizedElement = - normalizeGEPElementTypeForLLVMLowering(vecType.getElementType(), - builder); - if (normalizedElement == vecType.getElementType()) { - return normalizePayloadTypeForLLVMLowering(type, builder); - } - return getLLVMCompatibleVectorType(vecType.getShape(), normalizedElement, - vecType.getScalableDims()); - } - - if (auto vecType = dyn_cast(type)) { - Type normalizedElement = - normalizeGEPElementTypeForLLVMLowering(vecType.getElementType(), - builder); - if (normalizedElement == vecType.getElementType()) - return normalizePayloadTypeForLLVMLowering(type, builder); - return getLLVMCompatibleVectorType({vecType.getNumElements()}, - normalizedElement); - } - - return normalizePayloadTypeForLLVMLowering(type, builder); -} - -static Type convertVPTOType(Type type, Builder &builder) { - if (auto vecType = dyn_cast(type)) { - Type elementType = - normalizePayloadTypeForLLVMLowering(vecType.getElementType(), builder); - return getLLVMCompatibleVectorType({vecType.getElementCount()}, - elementType); - } - if (isa(type)) { - return VectorType::get({256}, builder.getI1Type()); - } - if (isa(type)) { - return VectorType::get({32}, builder.getI8Type()); - } - if (isa(type)) { - return LLVM::LLVMPointerType::get(builder.getContext()); - } - if (auto ptrType = dyn_cast(type)) { - return LLVM::LLVMPointerType::get( - builder.getContext(), - static_cast(ptrType.getMemorySpace().getAddressSpace())); - } - return normalizePayloadTypeForLLVMLowering(type, builder); -} - -static unsigned getNaturalByteAlignment(Type type) { - if (auto vecType = dyn_cast(type)) { - unsigned elemAlign = getNaturalByteAlignment(vecType.getElementType()); - if (!elemAlign) { - return 0; - } - int64_t elems = 1; - for (int64_t dim : vecType.getShape()) { - elems *= dim; - } - return elemAlign * static_cast(elems); - } - if (auto vecType = dyn_cast(type)) { - unsigned elemAlign = getNaturalByteAlignment(vecType.getElementType()); - if (!elemAlign) { - return 0; - } - return elemAlign * vecType.getNumElements(); - } - if (auto intType = dyn_cast(type)) { - return llvm::divideCeil(static_cast(intType.getWidth()), 8U); - } - if (pto::isPTOHiFloat8x2Type(type)) { - return 2; - } - if (pto::isPTOBF16x2Type(type)) { - return 4; - } - if (pto::isPTOLowPrecisionType(type)) { - return 1; - } - if (type.isF16() || type.isBF16()) { - return 2; - } - if (type.isF32()) { - return 4; - } - if (type.isF64()) { - return 8; - } - return 0; -} - -static bool hasVPTOConvertibleType(Type type) { - if (!type) { - return false; - } - if (isa(type) || - pto::isPTOLowPrecisionType(type)) - return true; - if (auto vecType = dyn_cast(type)) { - return hasVPTOConvertibleType(vecType.getElementType()); - } - return false; -} - -static bool hasVPTOConvertibleType(TypeRange types) { - return llvm::any_of(types, [](Type type) { return hasVPTOConvertibleType(type); }); -} - -static Value materializeVPTOCast(OpBuilder &builder, Type resultType, - ValueRange inputs, Location loc) { - if (inputs.size() != 1) { - return {}; - } - return builder - .create(loc, TypeRange{resultType}, inputs) - .getResult(0); -} - -class VPTOTypeConverter final : public TypeConverter { -public: - explicit VPTOTypeConverter(MLIRContext *context) { - addConversion([](Type type) { return type; }); - addConversion([](Type type) -> Type { - // The conversion callback outlives this constructor, so build on demand - // from the current type context instead of capturing a local Builder. - Builder builder(type.getContext()); - return convertVPTOType(type, builder); - }); - addSourceMaterialization(materializeVPTOCast); - addTargetMaterialization(materializeVPTOCast); - } -}; - -// Struct values carry the address of stack-local storage. Keep the pointee -// type local to struct access lowering so the public type conversion remains -// an opaque LLVM pointer, consistent with other pointer-like PTO handles. -static LLVM::LLVMStructType getVPTOStructStorageType(pto::StructType structType, - Builder &builder) { - struct Frame { - pto::StructType type; - bool materialize; - }; - - // PTO structs form an acyclic type tree. Build literal LLVM struct types in - // explicit post-order so deeply nested legal structs do not consume the C++ - // call stack during lowering. - SmallVector worklist{{structType, false}}; - llvm::DenseMap storageTypes; - while (!worklist.empty()) { - Frame frame = worklist.pop_back_val(); - if (!frame.materialize) { - worklist.push_back({frame.type, true}); - for (Type fieldType : frame.type.getFieldTypes()) { - if (auto nestedStruct = dyn_cast(fieldType)) { - worklist.push_back({nestedStruct, false}); - } - } - continue; - } - - SmallVector fieldTypes; - fieldTypes.reserve(frame.type.getNumFields()); - for (Type fieldType : frame.type.getFieldTypes()) { - if (auto nestedStruct = dyn_cast(fieldType)) { - fieldTypes.push_back(storageTypes.find(nestedStruct)->second); - continue; - } - fieldTypes.push_back(convertVPTOType(fieldType, builder)); - } - storageTypes[frame.type] = - LLVM::LLVMStructType::getLiteral(builder.getContext(), fieldTypes); - } - return storageTypes.find(structType)->second; -} - -static FailureOr -getVPTOStructFieldAddress(ConversionPatternRewriter &rewriter, Location loc, - Value root, pto::StructType rootType, - ArrayRef path) { - auto pointerType = LLVM::LLVMPointerType::get(rewriter.getContext()); - Value address = root; - pto::StructType currentType = rootType; - for (auto [depth, index] : llvm::enumerate(path)) { - if (index < 0 || index >= static_cast(currentType.getNumFields())) { - return failure(); - } - Type storageType = getVPTOStructStorageType(currentType, rewriter); - address = rewriter.create( - loc, pointerType, storageType, address, - ArrayRef{0, static_cast(index)}); - Type fieldType = currentType.getFieldType(static_cast(index)); - if (depth + 1 == path.size()) { - continue; - } - auto nestedStruct = dyn_cast(fieldType); - if (!nestedStruct) { - return failure(); - } - currentType = nestedStruct; - } - return address; -} - -struct PlannedDecl { - std::string name; - FunctionType type; -}; - -struct LoweringState { - SmallVector plannedDecls; -}; - -class LowerTrapOpPattern final : public OpConversionPattern { -public: - explicit LowerTrapOpPattern(TypeConverter &typeConverter, - MLIRContext *context, LoweringState &state) - : OpConversionPattern(typeConverter, context), - state(state) {} - - LogicalResult - matchAndRewrite(pto::TrapOp op, pto::TrapOp::Adaptor, - ConversionPatternRewriter &rewriter) const override { - constexpr StringLiteral calleeName = "llvm.hivm.TRAP"; - auto funcType = rewriter.getFunctionType({}, {}); - rewriter.create(op.getLoc(), calleeName, TypeRange{}, - ValueRange{}); - state.plannedDecls.push_back(PlannedDecl{calleeName.str(), funcType}); - rewriter.eraseOp(op); - return success(); - } - -private: - LoweringState &state; -}; - -enum class VcvtElemKind { - Invalid, - F16, - BF16, - F32, - F8E4M3, - F8E5M2, - HiF8, - F4E1M2x2, - F4E2M1x2, - S8, - U8, - S16, - U16, - S32, - U32, - S64, -}; - -struct VcvtContract { - const char *intrinsic; - bool requiresRnd; - bool requiresSat; - bool requiresPart; - unsigned maskBitWidth; - bool satBeforeRnd = false; -}; - -static Value getI64Constant(OpBuilder &builder, Location loc, uint64_t value) { - return builder.create(loc, builder.getI64IntegerAttr(value)) - .getResult(); -} - -static Value getI32Constant(OpBuilder &builder, Location loc, uint64_t value) { - return builder.create(loc, builder.getI32IntegerAttr(value)) - .getResult(); -} - -[[maybe_unused]] static Value getI1Constant(OpBuilder &builder, Location loc, - bool value) { - return builder - .create( - loc, builder.getIntegerAttr(builder.getI1Type(), value ? 1 : 0)) - .getResult(); -} - -static bool isMxElementType(Type ty) { - if (auto floatType = dyn_cast(ty)) { - return floatType.getWidth() == 8; - } - if (isa(ty)) { - return true; - } - std::string typeText; - llvm::raw_string_ostream os(typeText); - ty.print(os); - os.flush(); - return StringRef(typeText).starts_with("f8"); -} - -static std::string getMadMxElementFragment(Type type) { - if (type.isF16()) { - return "f16"; - } - if (type.isBF16()) { - return "bf16"; - } - - std::string typeText; - llvm::raw_string_ostream os(typeText); - type.print(os); - os.flush(); - - std::string lower = StringRef(typeText).lower(); - if (StringRef(lower).contains("e4m3")) { - return "e4m3"; - } - if (StringRef(lower).contains("e5m2")) { - return "e5m2"; - } - if (StringRef(lower).contains("hif4")) { - return "hif4"; - } - if (StringRef(lower).contains("e2m1x2")) { - return "e2m1x2"; - } - if (StringRef(lower).contains("e1m2x2")) { - return "e1m2x2"; - } - return {}; -} - -static FailureOr buildMadMxCalleeName(MLIRContext *context, - Type lhsElem, Type rhsElem) { - std::string lhs = getMadMxElementFragment(lhsElem); - std::string rhs = getMadMxElementFragment(rhsElem); - if (lhs.empty() || rhs.empty()) { - return failure(); - } - return StringAttr::get(context, "llvm.hivm.MMAD.MX." + lhs + rhs).getValue(); -} - -static bool isSignedOrSignlessInteger(IntegerType intType, unsigned width) { - return intType && intType.getWidth() == width && - (intType.isSigned() || intType.isSignless()); -} - -static std::string getMadRhsFragment(Type type) { - if (type.isF16()) { - return "f16"; - } - if (type.isBF16()) { - return "bf16"; - } - if (type.isF32()) { - return "f32"; - } - if (auto intType = dyn_cast(type)) { - if (isSignedOrSignlessInteger(intType, 4)) { - return "s4"; - } - if (isSignedOrSignlessInteger(intType, 8)) { - return "s8"; - } - if (intType.isUnsigned() && intType.getWidth() == 2) { - return "u2"; - } - } - - std::string typeText; - llvm::raw_string_ostream os(typeText); - type.print(os); - os.flush(); - std::string lower = StringRef(typeText).lower(); - if (StringRef(lower).contains("e8m0")) { - return "e8m0"; - } - return {}; -} - -static bool isMadE4M3ElementType(Type type) { - return pto::isPTOFloat8E4M3LikeType(type); -} - -static bool isMadE5M2ElementType(Type type) { - return pto::isPTOFloat8E5M2LikeType(type); -} - -static std::string getMadDstFragment(Type type) { - if (type.isF16()) { - return "f16"; - } - if (type.isF32()) { - return "f32"; - } - if (auto intType = dyn_cast(type)) { - if (isSignedOrSignlessInteger(intType, 32)) { - return "s32"; - } - } - return {}; -} - -static FailureOr buildMadTypedCalleeName(MLIRContext *context, - Type lhsElem, Type rhsElem, - Type dstElem) { - std::string rhs = getMadRhsFragment(rhsElem); - std::string dst = getMadDstFragment(dstElem); - if (lhsElem.isF16() && rhs == "f16" && dst == "f32") - return StringAttr::get(context, "llvm.hivm.MAD.f162f32.c310").getValue(); - if (lhsElem.isF16() && rhs == "f16" && dst == "f16") - return StringAttr::get(context, "llvm.hivm.MAD.f162f16").getValue(); - if (lhsElem.isF16() && rhs == "f16" && dst == "s32") - return StringAttr::get(context, "llvm.hivm.MAD.f162s32.1952").getValue(); - if (lhsElem.isBF16() && rhs == "bf16" && dst == "f32") - return StringAttr::get(context, "llvm.hivm.MAD.bf162f32.c310").getValue(); - if (lhsElem.isF32() && rhs == "f32" && dst == "f32") - return StringAttr::get(context, "llvm.hivm.MAD.f322f32.c310").getValue(); - if (isSignedOrSignlessInteger(dyn_cast(lhsElem), 8) && - rhs == "s8" && dst == "s32") - return StringAttr::get(context, "llvm.hivm.MAD.s8.c310").getValue(); - if (isMadE4M3ElementType(lhsElem) && isMadE4M3ElementType(rhsElem) && - dst == "f32") - return StringAttr::get(context, "llvm.hivm.MAD.e4m3e4m3.c310").getValue(); - if (isMadE4M3ElementType(lhsElem) && isMadE5M2ElementType(rhsElem) && - dst == "f32") - return StringAttr::get(context, "llvm.hivm.MAD.e4m3e5m2.c310").getValue(); - if (isMadE5M2ElementType(lhsElem) && isMadE4M3ElementType(rhsElem) && - dst == "f32") - return StringAttr::get(context, "llvm.hivm.MAD.e5m2e4m3.c310").getValue(); - if (isMadE5M2ElementType(lhsElem) && isMadE5M2ElementType(rhsElem) && - dst == "f32") - return StringAttr::get(context, "llvm.hivm.MAD.e5m2e5m2.c310").getValue(); - if (pto::isPTOHiFloat8Type(lhsElem) && pto::isPTOHiFloat8Type(rhsElem) && - dst == "f32") - return StringAttr::get(context, "llvm.hivm.MAD.e4m3e4m3.c310").getValue(); - if (lhsElem.isF16() && rhs == "s4") - return StringAttr::get(context, "llvm.hivm.MAD.f16s4.c310").getValue(); - if (lhsElem.isF16() && rhs == "s8") - return StringAttr::get(context, "llvm.hivm.MAD.f16s8.c310").getValue(); - if (lhsElem.isF16() && rhs == "u2") - return StringAttr::get(context, "llvm.hivm.MAD.f16u2").getValue(); - if (lhsElem.isF16() && rhs == "e8m0") - return StringAttr::get(context, "llvm.hivm.MAD.f16e8m0.c310").getValue(); - return failure(); -} - -static FailureOr buildLaneTypedCallee(MLIRContext *context, - Type resultType, - StringRef stem, - StringRef suffix) { - std::string vec = - getElementTypeFragment(getElementTypeFromVectorLike(resultType)); - auto lanes = getElementCountFromVectorLike(resultType); - if (vec.empty() || !lanes) - return failure(); - - return StringAttr::get(context, "llvm.hivm." + stem.str() + ".v" + - std::to_string(*lanes) + vec + - suffix.str()) - .getValue(); -} - -static std::string getLowPrecisionElementFragment(Type type); - -static std::string getCANN900VectorElementFragment(Type type) { - if (type.isF16()) - return "f16"; - if (type.isBF16()) - return "bf16"; - if (type.isF32()) - return "f32"; - if (std::string lowPrecision = getLowPrecisionElementFragment(type); - !lowPrecision.empty()) - return lowPrecision; - if (auto intType = dyn_cast(type)) - return "i" + std::to_string(intType.getWidth()); - return {}; -} - -static std::string getCANN900VectorTypeFragment(Type vectorType) { - std::string elem = - getCANN900VectorElementFragment(getElementTypeFromVectorLike(vectorType)); - auto lanes = getElementCountFromVectorLike(vectorType); - if (elem.empty() || !lanes) - return {}; - return "v" + std::to_string(*lanes) + elem; -} - -static std::string getCANN900SignednessFragment(Type elemType) { - if (elemType.isF16() || elemType.isBF16() || elemType.isF32()) - return "s"; - if (auto intType = dyn_cast(elemType)) - return intType.isUnsigned() ? "u" : "s"; - return {}; -} - -static FailureOr buildCANN900ModeTypedCallee(MLIRContext *context, - Type vectorType, - StringRef stem, - StringRef mode) { - std::string vec = getCANN900VectorTypeFragment(vectorType); - if (vec.empty()) - return failure(); - return StringAttr::get(context, "llvm.hivm." + stem.str() + "." + - mode.str() + "." + vec) - .getValue(); -} - -static FailureOr -buildCANN900SignedModeTypedCallee(MLIRContext *context, Type vectorType, - StringRef stem, StringRef mode) { - std::string vec = getCANN900VectorTypeFragment(vectorType); - std::string signedness = - getCANN900SignednessFragment(getElementTypeFromVectorLike(vectorType)); - if (vec.empty() || signedness.empty()) - return failure(); - return StringAttr::get(context, "llvm.hivm." + stem.str() + "." + - signedness + "." + mode.str() + "." + - vec) - .getValue(); -} - -static FailureOr -buildCANN900WideningReductionCallee(MLIRContext *context, Type inputType, - Type resultType, StringRef stem, - StringRef mode) { - std::string inputVec = getCANN900VectorTypeFragment(inputType); - std::string resultVec = getCANN900VectorTypeFragment(resultType); - std::string signedness = - getCANN900SignednessFragment(getElementTypeFromVectorLike(inputType)); - if (inputVec.empty() || resultVec.empty() || signedness.empty()) - return failure(); - return StringAttr::get(context, "llvm.hivm." + stem.str() + "." + - signedness + "." + mode.str() + "." + - resultVec + "." + inputVec) - .getValue(); -} - -static std::string getElementTypeFragment(Type type) { - if (type.isF16()) - return "f16"; - if (type.isBF16()) - return "bf16"; - if (type.isF32()) - return "f32"; - if (auto intType = dyn_cast(type)) - return (intType.isUnsigned() ? "u" : "s") + std::to_string(intType.getWidth()); - return {}; -} - -static std::string getLowPrecisionElementFragment(Type type) { - if (pto::isPTOHiFloat8x2Type(type)) - return "hif8x2"; - if (pto::isPTOHiFloat8Type(type)) - return "hif8"; - if (isa(type)) - return "f4e1m2x2"; - if (isa(type)) - return "f4e2m1x2"; - if (pto::isPTOBF16x2Type(type)) - return "bf16x2"; - if (pto::isPTOFloat8E4M3LikeType(type)) - return "f8e4m3"; - if (pto::isPTOFloat8E5M2LikeType(type)) - return "f8e5m2"; - return {}; -} - -static std::string getMemoryElementTypeFragment(Type type) { - if (auto intType = dyn_cast(type)) - return "i" + std::to_string(intType.getWidth()); - if (pto::isPTOHiFloat8Type(type)) - return "s8"; - if (std::string elem = getElementTypeFragment(type); !elem.empty()) - return elem; - return getLowPrecisionElementFragment(type); -} - -static bool isLowpPayloadElementType(Type type) { - return pto::isPTOFloat8Type(type) || pto::isPTOHiFloat8Type(type) || - pto::isPTOFloat4PackedType(type); -} - -struct LowpPayloadABI { - Type llvmElementType; - StringRef intrinsicElementFragment; -}; - -static std::optional -getLowpPayloadABI(Type elementType, MLIRContext *context) { - if (!isLowpPayloadElementType(elementType)) - return std::nullopt; - return LowpPayloadABI{IntegerType::get(context, 8), "u8"}; -} - -static std::string getDirectLowpVLogicElementFragment(Type type) { - if (pto::isPTOFloat8E4M3LikeType(type)) - return "fp8e4m3"; - if (pto::isPTOFloat8E5M2LikeType(type)) - return "fp8e5m2"; - return {}; -} - -static FailureOr -buildDirectLowpVLogicCallee(MLIRContext *context, Type vectorType, - StringRef stem, StringRef mode) { - Type elementType = getElementTypeFromVectorLike(vectorType); - auto lanes = getElementCountFromVectorLike(vectorType); - std::string elem = getDirectLowpVLogicElementFragment(elementType); - if (elem.empty() || !lanes) - return failure(); - return StringAttr::get(context, "llvm.hivm." + stem.str() + "." + - mode.str() + ".v" + - std::to_string(*lanes) + elem) - .getValue(); -} - -static FailureOr -buildLowpPayloadVLogicCallee(MLIRContext *context, Type vectorType, - StringRef stem, StringRef mode) { - Type elementType = getElementTypeFromVectorLike(vectorType); - auto lanes = getElementCountFromVectorLike(vectorType); - std::optional abi = getLowpPayloadABI(elementType, context); - if (!abi || !lanes) - return failure(); - return StringAttr::get(context, "llvm.hivm." + stem.str() + "." + - mode.str() + ".v" + - std::to_string(*lanes) + - abi->intrinsicElementFragment.str()) - .getValue(); -} - -static Type getLowpPayloadCarrierType(Type vectorLikeType, - MLIRContext *context) { - Type elementType = getElementTypeFromVectorLike(vectorLikeType); - std::optional abi = - getLowpPayloadABI(elementType, context); - if (!abi) - return {}; - auto lanes = getElementCountFromVectorLike(vectorLikeType); - if (!lanes) - return {}; - return VectorType::get({*lanes}, abi->llvmElementType); -} - -static Type getPayloadABIType(Type semanticType, Type convertedType, - MLIRContext *context) { - if (Type carrierType = getLowpPayloadCarrierType(semanticType, context)) - return carrierType; - return convertedType; -} - -static Value castToPayloadABI(Location loc, Value value, - Type semanticType, - ConversionPatternRewriter &rewriter) { - Type carrierType = - getLowpPayloadCarrierType(semanticType, rewriter.getContext()); - if (!carrierType || carrierType == value.getType()) - return value; - return rewriter.create(loc, carrierType, value); -} - -static Value castFromPayloadABI( - Location loc, Value value, Type semanticType, Type convertedType, - ConversionPatternRewriter &rewriter) { - Type carrierType = - getLowpPayloadCarrierType(semanticType, rewriter.getContext()); - if (!carrierType || carrierType == convertedType) - return value; - return rewriter.create(loc, convertedType, value); -} - -static std::string getAtomicElementTypeFragment(Type type, - Attribute signednessAttr) { - if (auto vecType = dyn_cast(type)) { - if (vecType.getRank() != 1 || vecType.getDimSize(0) != 2) - return {}; - if (vecType.getElementType().isF16()) - return "f16x2"; - if (vecType.getElementType().isBF16()) - return "bf16x2"; - return {}; - } - if (type.isF16()) - return "fp16"; - if (type.isBF16()) - return "bf16"; - if (type.isF32()) - return "fp32"; - auto intType = dyn_cast(type); - if (!intType) - return {}; - if (intType.getWidth() != 32 && intType.getWidth() != 64) - return {}; - if (signednessAttr) { - auto signedness = cast(signednessAttr).getValue(); - return std::string(signedness == pto::Signedness::Unsigned ? "u" : "s") + - std::to_string(intType.getWidth()); - } - return std::string(intType.isUnsigned() ? "u" : "s") + - std::to_string(intType.getWidth()); -} - -static std::string getL0LoadElementFragment(Type type) { - std::string elem = getElementTypeFragment(type); - if (!elem.empty()) - return elem; - - std::string typeText; - llvm::raw_string_ostream os(typeText); - type.print(os); - os.flush(); - std::string lower = StringRef(typeText).lower(); - if (StringRef(lower).contains("e4m3") || - StringRef(lower).contains("e5m2") || - StringRef(lower).contains("e8m0") || - StringRef(lower).contains("hif8") || - StringRef(lower).contains("e1m2x2") || - StringRef(lower).contains("e2m1x2")) - return "s8"; - return {}; -} - -static std::string getShuffleIntrinsicTypeFragment(Type type) { - if (auto intType = dyn_cast(type)) { - switch (intType.getWidth()) { - case 32: - return "i32"; - case 64: - return "i64"; - default: - return {}; - } - } - if (type.isF16()) - return "f16"; - if (type.isF32()) - return "f32"; - if (auto vecType = dyn_cast(type)) { - if (vecType.getRank() == 1 && vecType.getDimSize(0) == 2 && - vecType.getElementType().isF16()) - return "v2f16"; - } - return {}; -} - -static std::string getReduxIntrinsicTypeFragment(Type type, - Attribute signednessAttr) { - if (auto intType = dyn_cast(type)) { - if (intType.getWidth() != 32) - return {}; - bool isUnsigned = false; - if (signednessAttr) { - isUnsigned = cast(signednessAttr).getValue() == - pto::Signedness::Unsigned; - } - return isUnsigned ? "u32" : "s32"; - } - if (type.isF16()) - return "f16"; - if (type.isF32()) - return "f32"; - return {}; -} - -static Type getElementTypeFromVectorLike(Type type) { - if (auto vecType = dyn_cast(type)) - return vecType.getElementType(); - if (auto vecType = dyn_cast(type)) - return vecType.getElementType(); - if (auto vecType = dyn_cast(type)) - return vecType.getElementType(); - return {}; -} - -static std::optional getElementCountFromVectorLike(Type type) { - if (auto vecType = dyn_cast(type)) - return vecType.getElementCount(); - if (auto vecType = dyn_cast(type)) { - if (vecType.getRank() != 1) - return std::nullopt; - return vecType.getShape().front(); - } - if (auto vecType = dyn_cast(type)) - return vecType.getNumElements(); - return std::nullopt; -} - -static Value castIntegerLikeTo(Operation *anchor, Value value, Type targetType) { - OpBuilder builder(anchor); - builder.setInsertionPoint(anchor); - - if (value.getType() == targetType) - return value; - - auto targetInt = dyn_cast(targetType); - if (value.getType().isIndex() && targetInt) - return builder.create(anchor->getLoc(), targetType, value); - if (auto sourceInt = dyn_cast(value.getType())) { - if (targetInt) { - if (sourceInt.getWidth() < targetInt.getWidth()) - return builder.create(anchor->getLoc(), targetType, value); - if (sourceInt.getWidth() > targetInt.getWidth()) - return builder.create(anchor->getLoc(), targetType, value); - return value; - } - if (targetType.isIndex()) - return builder.create(anchor->getLoc(), targetType, value); - } - - return {}; -} - -static FailureOr reinterpretPointerToAddrSpace(Operation *anchor, - Value value, - unsigned targetAddressSpace) { - auto sourcePtrType = dyn_cast(value.getType()); - if (!sourcePtrType) - return failure(); - if (sourcePtrType.getAddressSpace() == targetAddressSpace) - return value; - - OpBuilder builder(anchor); - builder.setInsertionPoint(anchor); - Location loc = anchor->getLoc(); - Value asInt = builder.create(loc, builder.getI64Type(), value); - Type targetPtrType = - LLVM::LLVMPointerType::get(anchor->getContext(), targetAddressSpace); - return builder.create(loc, targetPtrType, asInt).getResult(); -} - -static FailureOr normalizeVdupScalarOperand(OpBuilder &builder, Location loc, - Value input, - Type resultType) { - auto intType = dyn_cast(input.getType()); - if (!intType || intType.getWidth() != 8) - return input; - - Type resultElemType = getElementTypeFromVectorLike(resultType); - std::string resultElemFragment = getElementTypeFragment(resultElemType); - if (resultElemFragment != "s8" && resultElemFragment != "u8") - return input; - - if (intType.isSignless()) - return input; - - Type signlessType = builder.getIntegerType(intType.getWidth()); - return builder - .create(loc, TypeRange{signlessType}, input) - .getResult(0); -} - -static Value normalizeByteScalarOperandForCANN900VectorCall( - OpBuilder &builder, Location loc, Value input, Type semanticElementType) { - (void)semanticElementType; - auto intType = dyn_cast(input.getType()); - if (!intType || intType.getWidth() != 8 || intType.isSignless()) - return input; - - Type signlessType = builder.getIntegerType(8); - return builder - .create(loc, TypeRange{signlessType}, input) - .getResult(0); -} - -static bool isCompatibleScalarForSemanticType(Type semanticType, - Type scalarType) { - if (semanticType == scalarType) - return true; - - auto semanticInt = dyn_cast(semanticType); - auto scalarInt = dyn_cast(scalarType); - if (!semanticInt || !scalarInt || semanticInt.getWidth() != scalarInt.getWidth()) - return false; - - if (semanticInt.isSigned()) - return scalarInt.isSigned() || scalarInt.isSignless(); - if (semanticInt.isUnsigned()) - return scalarInt.isUnsigned() || scalarInt.isSignless(); - return scalarInt.isSignless(); -} - -static std::string getCopyElementFragment(Type elementType) { - if (!elementType) - return {}; - if (elementType.isF16()) - return "f16"; - if (elementType.isBF16()) - return "bf16"; - if (elementType.isF32()) - return "f32"; - // Handle FP8 family (e4m3/e5m2/e8m0/hif8) used by cube-matmul/mad_mx. - std::string typeText; - llvm::raw_string_ostream os(typeText); - elementType.print(os); - os.flush(); - std::string lower = StringRef(typeText).lower(); - if (StringRef(lower).contains("e4m3")) - return "e4m3"; - if (StringRef(lower).contains("e5m2")) - return "e5m2"; - if (StringRef(lower).contains("e8m0")) - return "e8m0"; - if (StringRef(lower).contains("hif8")) - return "hif8"; - if (StringRef(lower).contains("e1m2x2") || StringRef(lower).contains("e2m1x2")) - return "u8"; - if (auto intType = dyn_cast(elementType)) { - switch (intType.getWidth()) { - case 8: - return intType.isUnsigned() ? "u8" : "s8"; - case 16: - return intType.isUnsigned() ? "u16" : "s16"; - case 32: - return intType.isUnsigned() ? "u32" : "s32"; - default: - return {}; - } - } - return {}; -} - -static std::string getNd2NzCopyElementFragment(Type elementType) { - if (!elementType) - return {}; - std::string typeText; - llvm::raw_string_ostream os(typeText); - elementType.print(os); - os.flush(); - std::string lower = StringRef(typeText).lower(); - if (StringRef(lower).contains("e4m3") || StringRef(lower).contains("e5m2") || - StringRef(lower).contains("e8m0") || StringRef(lower).contains("hif8")) - return "U8"; - if (StringRef(lower).contains("e1m2x2") || StringRef(lower).contains("e2m1x2")) - return "U8"; - - if (elementType.isF16() || elementType.isBF16()) - return "U16"; - if (elementType.isF32()) - return "U32"; - if (auto intType = dyn_cast(elementType)) { - switch (intType.getWidth()) { - case 8: - return "U8"; - case 16: - return "U16"; - case 32: - return "U32"; - default: - return {}; - } - } - return {}; -} - -static std::optional parsePredicatePatternImmediate(StringRef pattern) { - if (pattern == "PAT_ALL") - return 0; - if (pattern == "PAT_VL1") - return 1; - if (pattern == "PAT_VL2") - return 2; - if (pattern == "PAT_VL3") - return 3; - if (pattern == "PAT_VL4") - return 4; - if (pattern == "PAT_VL8") - return 5; - if (pattern == "PAT_VL16") - return 6; - if (pattern == "PAT_VL32") - return 7; - if (pattern == "PAT_VL64") - return 8; - if (pattern == "PAT_VL128") - return 9; - if (pattern == "PAT_M3") - return 10; - if (pattern == "PAT_M4") - return 11; - if (pattern == "PAT_H") - return 12; - if (pattern == "PAT_Q") - return 13; - if (pattern == "PAT_ALLF") - return 15; - return std::nullopt; -} - -static std::optional parseHiLoPartImmediate(StringRef part) { - if (part == "LOWER") - return 0; - if (part == "HIGHER") - return 1; - return std::nullopt; -} - -static std::optional parseRoundModeImmediate(StringRef roundMode) { - if (roundMode == "R" || roundMode == "ROUND_R") - return 0; - if (roundMode == "A" || roundMode == "ROUND_A") - return 1; - if (roundMode == "F" || roundMode == "ROUND_F") - return 2; - if (roundMode == "C" || roundMode == "ROUND_C") - return 3; - if (roundMode == "Z" || roundMode == "ROUND_Z") - return 4; - if (roundMode == "O" || roundMode == "ROUND_O") - return 5; - if (roundMode == "H" || roundMode == "ROUND_H") - return 6; - return std::nullopt; -} - -static std::optional parseSaturationImmediate(StringRef sat) { - if (sat == "SAT") - return 1; - if (sat == "NOSAT") - return 0; - return std::nullopt; -} - -static std::optional parsePartImmediate(StringRef part) { - if (part == "EVEN" || part == "PART_EVEN") - return 0; - if (part == "ODD" || part == "PART_ODD") - return 1; - return std::nullopt; -} - -static std::optional parseVcvtPartImmediate(StringRef part) { - if (part == "EVEN" || part == "PART_EVEN" || part == "P0" || - part == "PART_P0") - return 0; - if (part == "ODD" || part == "PART_ODD" || part == "P1" || - part == "PART_P1") - return 1; - if (part == "P2" || part == "PART_P2") - return 2; - if (part == "P3" || part == "PART_P3") - return 3; - return std::nullopt; -} - -static std::optional parsePredicateStoreDistImmediate(StringRef dist) { - if (dist == "NORM") - return 0; - if (dist == "PK") - return 1; - return std::nullopt; -} - -static std::optional parsePredicateLoadDistImmediate(StringRef dist) { - if (dist.empty() || dist == "NORM") - return 0; - if (dist == "US") - return 1; - if (dist == "DS") - return 2; - return std::nullopt; -} - -static std::optional parsePostModeImmediate(StringRef mode) { - if (mode == "NO_POST_UPDATE") - return 0; - if (mode == "POST_UPDATE") - return 1; - return std::nullopt; -} - -static std::optional parsePipeImmediate(StringRef pipe) { - if (pipe == "PIPE_S") - return 0; - if (pipe == "PIPE_V") - return 1; - if (pipe == "PIPE_M") - return 2; - if (pipe == "PIPE_MTE1") - return 3; - if (pipe == "PIPE_MTE2") - return 4; - if (pipe == "PIPE_MTE3") - return 5; - if (pipe == "PIPE_ALL") - return 6; - if (pipe == "PIPE_MTE4") - return 7; - if (pipe == "PIPE_MTE5") - return 8; - if (pipe == "PIPE_V2") - return 9; - if (pipe == "PIPE_FIX") - return 10; - if (pipe == "VIRTUAL_PIPE_MTE2_L1A") - return 11; - if (pipe == "VIRTUAL_PIPE_MTE2_L1B") - return 12; - return std::nullopt; -} - -static std::optional parseEventImmediate(StringRef event) { - if (!event.consume_front("EVENT_ID")) - return std::nullopt; - uint64_t value = 0; - if (event.getAsInteger(10, value)) - return std::nullopt; - return value; -} - -static std::optional parseSprImmediate(StringRef spr) { - if (spr == "AR") - return 74; - return std::nullopt; -} - -static std::optional getDistElementWidth(Type type) { - if (auto intType = dyn_cast(type)) - return intType.getWidth(); - if (isLowpPayloadElementType(type)) - return 8; - if (type.isF16() || type.isBF16()) - return 16; - if (type.isF32()) - return 32; - if (type.isF64()) - return 64; - // bf16x2 is a 32-bit packed pair; its dist width is 32 (i32/align4 ABI). - if (pto::isPTOBF16x2Type(type)) - return 32; - return std::nullopt; -} - -static VcvtElemKind classifyVcvtElemType(Type type) { - if (type.isF16()) - return VcvtElemKind::F16; - if (type.isBF16()) - return VcvtElemKind::BF16; - if (type.isF32()) - return VcvtElemKind::F32; - if (pto::isPTOFloat8E4M3LikeType(type)) - return VcvtElemKind::F8E4M3; - if (pto::isPTOFloat8E5M2LikeType(type)) - return VcvtElemKind::F8E5M2; - if (pto::isPTOHiFloat8Type(type)) - return VcvtElemKind::HiF8; - if (isa(type)) - return VcvtElemKind::F4E1M2x2; - if (isa(type)) - return VcvtElemKind::F4E2M1x2; - if (auto intType = dyn_cast(type)) { - switch (intType.getWidth()) { - case 8: - return intType.isUnsigned() ? VcvtElemKind::U8 : VcvtElemKind::S8; - case 16: - return intType.isUnsigned() ? VcvtElemKind::U16 : VcvtElemKind::S16; - case 32: - return intType.isUnsigned() ? VcvtElemKind::U32 : VcvtElemKind::S32; - case 64: - return intType.isUnsigned() ? VcvtElemKind::Invalid : VcvtElemKind::S64; - default: - return VcvtElemKind::Invalid; - } - } - return VcvtElemKind::Invalid; -} - -static std::optional lookupVcvtContract(VcvtElemKind src, - VcvtElemKind dst) { - switch (src) { - case VcvtElemKind::F32: - switch (dst) { - case VcvtElemKind::F8E4M3: - return VcvtContract{"llvm.hivm.vcvtff.f322f8e4m3.x", true, true, true, 32}; - case VcvtElemKind::F8E5M2: - return VcvtContract{"llvm.hivm.vcvtff.f322f8e5m2.x", true, true, true, 32}; - case VcvtElemKind::HiF8: - return VcvtContract{"llvm.hivm.vcvtff.f322hif8.x", true, true, true, 32}; - case VcvtElemKind::F16: - return VcvtContract{"llvm.hivm.vcvtff.f322f16.x", true, true, true, 32}; - case VcvtElemKind::BF16: - return VcvtContract{"llvm.hivm.vcvtff.f322bf16.x", true, true, true, 32}; - case VcvtElemKind::S16: - return VcvtContract{"llvm.hivm.vcvtfi.f322s16.x", true, true, true, 32}; - case VcvtElemKind::S32: - return VcvtContract{"llvm.hivm.vcvtfi.f322s32.x", true, true, false, 32}; - case VcvtElemKind::S64: - return VcvtContract{"llvm.hivm.vcvtfi.f322s64.x", true, true, true, 32}; - default: - return std::nullopt; - } - case VcvtElemKind::F16: - switch (dst) { - case VcvtElemKind::F8E4M3: - return VcvtContract{"llvm.hivm.vcvtff.f162f8e4m3.x", true, true, true, 16}; - case VcvtElemKind::F8E5M2: - return VcvtContract{"llvm.hivm.vcvtff.f162f8e5m2.x", true, true, true, 16}; - case VcvtElemKind::HiF8: - return VcvtContract{"llvm.hivm.vcvtff.f162hif8.x", true, true, true, 16}; - case VcvtElemKind::F32: - return VcvtContract{"llvm.hivm.vcvtff.f162f32.x", false, false, true, 16}; - case VcvtElemKind::S32: - return VcvtContract{"llvm.hivm.vcvtfi.f162s32.x", true, false, true, 16}; - case VcvtElemKind::S16: - return VcvtContract{"llvm.hivm.vcvtfi.f162s16.x", true, true, false, 16}; - case VcvtElemKind::S8: - return VcvtContract{"llvm.hivm.vcvtfi.f162s8.x", true, true, true, 16}; - case VcvtElemKind::U8: - return VcvtContract{"llvm.hivm.vcvtfi.f162u8.x", true, true, true, 16}; - default: - return std::nullopt; - } - case VcvtElemKind::BF16: - switch (dst) { - case VcvtElemKind::F8E4M3: - return VcvtContract{"llvm.hivm.vcvtff.bf162f8e4m3.x", true, true, true, 16}; - case VcvtElemKind::F8E5M2: - return VcvtContract{"llvm.hivm.vcvtff.bf162f8e5m2.x", true, true, true, 16}; - case VcvtElemKind::F4E1M2x2: - return VcvtContract{"llvm.hivm.vcvtff2.bf162f4e1m2x2.x", true, false, true, 16}; - case VcvtElemKind::F4E2M1x2: - return VcvtContract{"llvm.hivm.vcvtff2.bf162f4e2m1x2.x", true, false, true, 16}; - case VcvtElemKind::F16: - return VcvtContract{"llvm.hivm.vcvtff.bf162f16.x", true, true, false, 16, - true}; - case VcvtElemKind::F32: - return VcvtContract{"llvm.hivm.vcvtff.bf162f32.x", false, false, true, 16}; - case VcvtElemKind::S32: - return VcvtContract{"llvm.hivm.vcvtfi.bf162s32.x", true, true, true, 16}; - default: - return std::nullopt; - } - case VcvtElemKind::U8: - switch (dst) { - case VcvtElemKind::F16: - return VcvtContract{"llvm.hivm.vcvtif.u82f16.x", false, false, true, 8}; - case VcvtElemKind::U16: - return VcvtContract{"llvm.hivm.vcvtii.u82u16.x", false, false, true, 8}; - case VcvtElemKind::U32: - return VcvtContract{"llvm.hivm.vcvtii.u82u32.x", false, false, true, 8}; - default: - return std::nullopt; - } - case VcvtElemKind::S8: - switch (dst) { - case VcvtElemKind::F16: - return VcvtContract{"llvm.hivm.vcvtif.s82f16.x", false, false, true, 8}; - case VcvtElemKind::S16: - return VcvtContract{"llvm.hivm.vcvtii.s82s16.x", false, false, true, 8}; - case VcvtElemKind::S32: - return VcvtContract{"llvm.hivm.vcvtii.s82s32.x", false, false, true, 8}; - default: - return std::nullopt; - } - case VcvtElemKind::U16: - switch (dst) { - case VcvtElemKind::U8: - return VcvtContract{"llvm.hivm.vcvtii.u162u8.x", false, true, true, 16}; - case VcvtElemKind::U32: - return VcvtContract{"llvm.hivm.vcvtii.u162u32.x", false, false, true, 16}; - default: - return std::nullopt; - } - case VcvtElemKind::S16: - switch (dst) { - case VcvtElemKind::F16: - return VcvtContract{"llvm.hivm.vcvtif.s162f16.x", true, false, false, 16}; - case VcvtElemKind::F32: - return VcvtContract{"llvm.hivm.vcvtif.s162f32.x", false, false, true, 16}; - case VcvtElemKind::U8: - return VcvtContract{"llvm.hivm.vcvtii.s162u8.x", false, true, true, 16}; - case VcvtElemKind::U32: - return VcvtContract{"llvm.hivm.vcvtii.s162u32.x", false, false, true, 16}; - case VcvtElemKind::S32: - return VcvtContract{"llvm.hivm.vcvtii.s162s32.x", false, false, true, 16}; - default: - return std::nullopt; - } - case VcvtElemKind::U32: - switch (dst) { - case VcvtElemKind::U8: - return VcvtContract{"llvm.hivm.vcvtii.u322u8.x", false, true, true, 32}; - case VcvtElemKind::U16: - return VcvtContract{"llvm.hivm.vcvtii.u322u16.x", false, true, true, 32}; - case VcvtElemKind::S16: - return VcvtContract{"llvm.hivm.vcvtii.u322s16.x", false, true, true, 32}; - default: - return std::nullopt; - } - case VcvtElemKind::S32: - switch (dst) { - case VcvtElemKind::F32: - return VcvtContract{"llvm.hivm.vcvtif.s322f32.x", true, false, false, 32}; - case VcvtElemKind::U8: - return VcvtContract{"llvm.hivm.vcvtii.s322u8.x", false, true, true, 32}; - case VcvtElemKind::U16: - return VcvtContract{"llvm.hivm.vcvtii.s322u16.x", false, true, true, 32}; - case VcvtElemKind::S16: - return VcvtContract{"llvm.hivm.vcvtii.s322s16.x", false, true, true, 32}; - case VcvtElemKind::S64: - return VcvtContract{"llvm.hivm.vcvtii.s322s64.x", false, false, true, 32}; - default: - return std::nullopt; - } - case VcvtElemKind::S64: - switch (dst) { - case VcvtElemKind::F32: - return VcvtContract{"llvm.hivm.vcvtif.s642f32.x", true, false, true, 32}; - case VcvtElemKind::S32: - return VcvtContract{"llvm.hivm.vcvtii.s642s32.x", false, true, true, 32}; - default: - return std::nullopt; - } - case VcvtElemKind::F8E4M3: - switch (dst) { - case VcvtElemKind::F32: - return VcvtContract{"llvm.hivm.vcvtff.f8e4m32f32.x", false, false, true, 8}; - default: - return std::nullopt; - } - case VcvtElemKind::F8E5M2: - switch (dst) { - case VcvtElemKind::F32: - return VcvtContract{"llvm.hivm.vcvtff.f8e5m22f32.x", false, false, true, 8}; - default: - return std::nullopt; - } - case VcvtElemKind::HiF8: - switch (dst) { - case VcvtElemKind::F32: - return VcvtContract{"llvm.hivm.vcvtff.hif82f32.x", false, false, true, 8}; - default: - return std::nullopt; - } - case VcvtElemKind::F4E1M2x2: - switch (dst) { - case VcvtElemKind::BF16: - return VcvtContract{"llvm.hivm.vcvtff2.f4e1m2x22bf16.x", false, false, true, 8}; - default: - return std::nullopt; - } - case VcvtElemKind::F4E2M1x2: - switch (dst) { - case VcvtElemKind::BF16: - return VcvtContract{"llvm.hivm.vcvtff2.f4e2m1x22bf16.x", false, false, true, 8}; - default: - return std::nullopt; - } - case VcvtElemKind::Invalid: - return std::nullopt; - } - return std::nullopt; -} - -// VSQZ #st hint must only be set when the compacted vector feeds VSTUR. -// Emitting #st=1 without a matching VSTUR consumer can deadlock hardware queues. -static uint64_t determineVsqzStoreHint(pto::VsqzOp vsqz) { - Value result = vsqz.getResult(); - for (Operation *user : result.getUsers()) { - auto vstur = dyn_cast(user); - if (!vstur) - continue; - if (vstur.getValue() == result) - return 1; - } - return 0; -} - -static std::optional parseLoadDistImmediate(StringRef dist, - Type elementType) { - auto width = getDistElementWidth(elementType); - if (dist.empty() || dist == "NORM") - return 0; - if (!width) - return std::nullopt; - if (dist == "BRC_B8") - return std::optional(1); - if (dist == "BRC_B16") - return std::optional(2); - if (dist == "BRC_B32") - return std::optional(3); - if (dist == "US_B8") - return std::optional(6); - if (dist == "US_B16") - return std::optional(7); - if (dist == "DS_B8") - return std::optional(8); - if (dist == "DS_B16") - return std::optional(9); - if (dist == "UNPK_B8") - return std::optional(13); - if (dist == "UNPK_B16") - return std::optional(14); - if (dist == "UNPK_B32") - return std::optional(18); - if (dist == "BRC_BLK") - return 15; - if (dist == "E2B_B16") - return std::optional(16); - if (dist == "E2B_B32") - return std::optional(17); - if (dist == "UNPK4") - return *width == 8 ? std::optional(20) : std::nullopt; - if (dist == "SPLT4CHN") - return *width == 8 ? std::optional(21) : std::nullopt; - if (dist == "SPLT2CHN_B8") - return std::optional(22); - if (dist == "SPLT2CHN_B16") - return std::optional(23); - return std::nullopt; -} - -static std::optional parseLoadX2DistImmediate(StringRef dist, - Type elementType) { - auto width = getDistElementWidth(elementType); - if (dist == "BDINTLV") - return 10; - if (!width) - return std::nullopt; - if (dist == "DINTLV_B8") - return std::optional(11); - if (dist == "DINTLV_B16") - return std::optional(12); - if (dist == "DINTLV_B32") - return std::optional(19); - return std::nullopt; -} - -static std::optional parseStoreDistImmediate(StringRef dist, - Type elementType) { - auto width = getDistElementWidth(elementType); - if (dist.empty()) { - if (!width) - return std::nullopt; - if (*width == 8) - return 0; - if (*width == 16) - return 1; - if (*width == 32) - return 2; - return std::nullopt; - } - if (dist == "NORM_B8") - return std::optional(0); - if (dist == "NORM_B16") - return std::optional(1); - if (dist == "NORM_B32") - return std::optional(2); - if (dist == "1PT_B8") - return std::optional(3); - if (dist == "1PT_B16") - return std::optional(4); - if (dist == "1PT_B32") - return std::optional(5); - if (dist == "PK_B16") - return std::optional(6); - if (dist == "PK_B32") - return std::optional(7); - if (dist == "PK_B64") - return std::optional(10); - if (dist == "PK4_B32") - return std::optional(12); - if (dist == "MRG4CHN_B8") - return std::optional(13); - if (dist == "MRG2CHN_B8") - return std::optional(14); - if (dist == "MRG2CHN_B16") - return std::optional(15); - return std::nullopt; -} - -static bool isOnePointStoreDist(StringRef dist) { - return dist == "1PT_B8" || dist == "1PT_B16" || dist == "1PT_B32"; -} - -static bool isMaskOnlyUsedByOnePointStores(Value mask) { - return !mask.use_empty() && llvm::all_of(mask.getUsers(), [](Operation *user) { - auto store = dyn_cast(user); - return store && store.getDist() && isOnePointStoreDist(*store.getDist()); - }); -} - -static std::optional parseStoreX2DistImmediate(StringRef dist, - Type elementType) { - auto width = getDistElementWidth(elementType); - if (!width) - return std::nullopt; - if (dist == "INTLV_B8") - return std::optional(8); - if (dist == "INTLV_B16") - return std::optional(9); - if (dist == "INTLV_B32") - return std::optional(11); - return std::nullopt; -} - -static Value packBlockRepeatStride(Operation *anchor, Value blockStride, - Value repeatStride) { - OpBuilder builder(anchor); - builder.setInsertionPoint(anchor); - - Value blockI32 = castIntegerLikeTo(anchor, blockStride, builder.getI32Type()); - Value repeatI32 = - castIntegerLikeTo(anchor, repeatStride, builder.getI32Type()); - if (!blockI32 || !repeatI32) - return {}; - - auto c16 = builder.create(anchor->getLoc(), 16, 32); - auto blockShifted = - builder.create(anchor->getLoc(), blockI32, c16); - return builder - .create(anchor->getLoc(), blockShifted, repeatI32) - .getResult(); -} - -static std::optional parseOrderImmediate(StringRef order) { - if (order.empty() || order == "ASC") - return 0; - if (order == "DESC") - return 1; - return std::nullopt; -} - -static FailureOr packLoopPair(Operation *anchor, Value low, Value high) { - OpBuilder builder(anchor); - builder.setInsertionPoint(anchor); - - Value lowI64 = castIntegerLikeTo(anchor, low, builder.getI64Type()); - Value highI64 = castIntegerLikeTo(anchor, high, builder.getI64Type()); - if (!lowI64 || !highI64) - return failure(); - - Value shift = getI64Constant(builder, anchor->getLoc(), 40); - Value highShifted = - builder.create(anchor->getLoc(), highI64, shift).getResult(); - return builder.create(anchor->getLoc(), highShifted, lowI64) - .getResult(); -} - -static FailureOr packLoopSize(Operation *anchor, Value loop2, Value loop1) { - OpBuilder builder(anchor); - builder.setInsertionPoint(anchor); - - Value loop2I64 = castIntegerLikeTo(anchor, loop2, builder.getI64Type()); - Value loop1I64 = castIntegerLikeTo(anchor, loop1, builder.getI64Type()); - if (!loop2I64 || !loop1I64) - return failure(); - - Value shift = getI64Constant(builder, anchor->getLoc(), 21); - Value loop2Shifted = - builder.create(anchor->getLoc(), loop2I64, shift).getResult(); - return builder.create(anchor->getLoc(), loop2Shifted, loop1I64) - .getResult(); -} - -static FailureOr -packCopyGmToUbConfig0(Operation *anchor, ValueRange operands) { - if (operands.size() != 11) - return failure(); - - OpBuilder builder(anchor); - builder.setInsertionPoint(anchor); - Location loc = anchor->getLoc(); - - auto getI64Operand = [&](unsigned idx) -> Value { - return castIntegerLikeTo(anchor, operands[idx], builder.getI64Type()); - }; - - Value sid = getI64Operand(2); - Value nBurst = getI64Operand(3); - Value lenBurst = getI64Operand(4); - Value leftPadding = getI64Operand(5); - Value rightPadding = getI64Operand(6); - Value dataSelect = castIntegerLikeTo(anchor, operands[7], builder.getI64Type()); - Value cacheCtl = getI64Operand(8); - if (!sid || !nBurst || !lenBurst || !leftPadding || !rightPadding || - !dataSelect || !cacheCtl) - return failure(); - - auto shl = [&](Value value, uint64_t amount) -> Value { - return builder.create(loc, value, - getI64Constant(builder, loc, amount)); - }; - auto bitOr = [&](Value lhs, Value rhs) -> Value { - return builder.create(loc, lhs, rhs); - }; - - Value config = sid; - config = bitOr(config, shl(nBurst, 4)); - config = bitOr(config, shl(lenBurst, 25)); - config = bitOr(config, shl(leftPadding, 46)); - config = bitOr(config, shl(rightPadding, 52)); - config = bitOr(config, shl(dataSelect, 58)); - config = bitOr(config, shl(cacheCtl, 60)); - return config; -} - -static FailureOr -packCopyGmToUbConfig1(Operation *anchor, ValueRange operands) { - if (operands.size() != 11) - return failure(); - return packLoopPair(anchor, operands[9], operands[10]); -} - -[[maybe_unused]] static FailureOr -packCopyGmToUbConfig0(Operation *anchor, Value sid, Value nBurst, - Value lenBurst, Value leftPadding, Value rightPadding, - Value dataSelect, Value cacheCtl) { - SmallVector operands(11); - operands[2] = sid; - operands[3] = nBurst; - operands[4] = lenBurst; - operands[5] = leftPadding; - operands[6] = rightPadding; - operands[7] = dataSelect; - operands[8] = cacheCtl; - return packCopyGmToUbConfig0(anchor, operands); -} - -static FailureOr -packCopyUbToGmConfig0(Operation *anchor, ValueRange operands) { - if (operands.size() != 8) - return failure(); - - OpBuilder builder(anchor); - builder.setInsertionPoint(anchor); - Location loc = anchor->getLoc(); - - auto getI64Operand = [&](unsigned idx) -> Value { - return castIntegerLikeTo(anchor, operands[idx], builder.getI64Type()); - }; - - Value sid = getI64Operand(2); - Value nBurst = getI64Operand(3); - Value lenBurst = getI64Operand(4); - Value l2CacheCtl = getI64Operand(5); - if (!sid || !nBurst || !lenBurst || !l2CacheCtl) - return failure(); - - auto shl = [&](Value value, uint64_t amount) -> Value { - return builder.create(loc, value, - getI64Constant(builder, loc, amount)); - }; - auto bitOr = [&](Value lhs, Value rhs) -> Value { - return builder.create(loc, lhs, rhs); - }; - - Value config = sid; - config = bitOr(config, shl(nBurst, 4)); - config = bitOr(config, shl(lenBurst, 25)); - config = bitOr(config, shl(l2CacheCtl, 60)); - return config; -} - -static FailureOr -packCopyUbToGmConfig1(Operation *anchor, ValueRange operands) { - if (operands.size() != 8) - return failure(); - return packLoopPair(anchor, operands[6], operands[7]); -} - -[[maybe_unused]] static FailureOr -packCopyUbToGmConfig0(Operation *anchor, Value sid, Value nBurst, - Value lenBurst, Value l2CacheCtl) { - SmallVector operands(8); - operands[2] = sid; - operands[3] = nBurst; - operands[4] = lenBurst; - operands[5] = l2CacheCtl; - return packCopyUbToGmConfig0(anchor, operands); -} - -static FailureOr -packCopyUbToUbConfig(Operation *anchor, ValueRange operands) { - if (operands.size() != 7) - return failure(); - OpBuilder builder(anchor); - builder.setInsertionPoint(anchor); - Location loc = anchor->getLoc(); - - auto getI64Operand = [&](unsigned idx) -> Value { - return castIntegerLikeTo(anchor, operands[idx], builder.getI64Type()); - }; - - Value nBurst = getI64Operand(3); - Value lenBurst = getI64Operand(4); - Value srcStride = getI64Operand(5); - Value dstStride = getI64Operand(6); - if (!nBurst || !lenBurst || !srcStride || !dstStride) - return failure(); - - auto shl = [&](Value value, uint64_t amount) -> Value { - return builder.create(loc, value, - getI64Constant(builder, loc, amount)); - }; - auto bitOr = [&](Value lhs, Value rhs) -> Value { - return builder.create(loc, lhs, rhs); - }; - - Value config = nBurst; - config = bitOr(config, shl(lenBurst, 16)); - config = bitOr(config, shl(srcStride, 32)); - config = bitOr(config, shl(dstStride, 48)); - return config; -} - -static FailureOr -packCopyCbufToUbConfig(Operation *anchor, ValueRange operands) { - if (operands.size() != 7) - return failure(); - OpBuilder builder(anchor); - builder.setInsertionPoint(anchor); - Location loc = anchor->getLoc(); - - auto getI64Operand = [&](unsigned idx) -> Value { - return castIntegerLikeTo(anchor, operands[idx], builder.getI64Type()); - }; - - Value sid = getI64Operand(2); - Value nBurst = getI64Operand(3); - Value lenBurst = getI64Operand(4); - Value srcStride = getI64Operand(5); - Value dstStride = getI64Operand(6); - if (!sid || !nBurst || !lenBurst || !srcStride || !dstStride) - return failure(); - - auto shl = [&](Value value, uint64_t amount) -> Value { - return builder.create(loc, value, - getI64Constant(builder, loc, amount)); - }; - auto bitOr = [&](Value lhs, Value rhs) -> Value { - return builder.create(loc, lhs, rhs); - }; - - Value config = sid; - config = bitOr(config, shl(nBurst, 4)); - config = bitOr(config, shl(lenBurst, 16)); - config = bitOr(config, shl(srcStride, 32)); - config = bitOr(config, shl(dstStride, 48)); - return config; -} - -static FailureOr -packCopyUbToCbufConfig(Operation *anchor, ValueRange operands) { - if (operands.size() != 7) - return failure(); - OpBuilder builder(anchor); - builder.setInsertionPoint(anchor); - Location loc = anchor->getLoc(); - - auto getI64Operand = [&](unsigned idx) -> Value { - return castIntegerLikeTo(anchor, operands[idx], builder.getI64Type()); - }; - - Value sid = getI64Operand(2); - Value nBurst = getI64Operand(3); - Value lenBurst = getI64Operand(4); - Value srcStride = getI64Operand(5); - Value dstStride = getI64Operand(6); - if (!sid || !nBurst || !lenBurst || !srcStride || !dstStride) - return failure(); - - auto shl = [&](Value value, uint64_t amount) -> Value { - return builder.create(loc, value, - getI64Constant(builder, loc, amount)); - }; - auto bitOr = [&](Value lhs, Value rhs) -> Value { - return builder.create(loc, lhs, rhs); - }; - - Value config = sid; - config = bitOr(config, shl(nBurst, 4)); - config = bitOr(config, shl(lenBurst, 16)); - config = bitOr(config, shl(srcStride, 32)); - config = bitOr(config, shl(dstStride, 48)); - return config; -} - -static FailureOr -packCopyGmToCbufConfig0(Operation *anchor, Value nBurst, Value lenBurst) { - OpBuilder builder(anchor); - builder.setInsertionPoint(anchor); - Location loc = anchor->getLoc(); - - Value nBurstI64 = castIntegerLikeTo(anchor, nBurst, builder.getI64Type()); - Value lenBurstI64 = castIntegerLikeTo(anchor, lenBurst, builder.getI64Type()); - if (!nBurstI64 || !lenBurstI64) - return failure(); - - auto shl = [&](Value value, uint64_t amount) -> Value { - return builder.create(loc, value, - getI64Constant(builder, loc, amount)); - }; - auto bitOr = [&](Value lhs, Value rhs) -> Value { - return builder.create(loc, lhs, rhs); - }; - - Value config0 = getI64Constant(builder, loc, 0); // sid - config0 = bitOr(config0, shl(nBurstI64, 4)); // burst_num[24:4] - config0 = bitOr(config0, shl(lenBurstI64, 25)); // burst_len[45:25] - return config0; -} - -static FailureOr -packCopyGmToCbufConfig1(Operation *anchor, Value srcStride, - Value dstStride) { - OpBuilder builder(anchor); - builder.setInsertionPoint(anchor); - Location loc = anchor->getLoc(); - - Value srcStrideI64 = castIntegerLikeTo(anchor, srcStride, builder.getI64Type()); - Value dstStrideI64 = castIntegerLikeTo(anchor, dstStride, builder.getI64Type()); - if (!srcStrideI64 || !dstStrideI64) - return failure(); - - auto shl = [&](Value value, uint64_t amount) -> Value { - return builder.create(loc, value, - getI64Constant(builder, loc, amount)); - }; - auto bitOr = [&](Value lhs, Value rhs) -> Value { - return builder.create(loc, lhs, rhs); - }; - - // config1 packs burst_src_stride[39:0] and burst_dst_stride[60:40]. - return bitOr(srcStrideI64, shl(dstStrideI64, 40)); -} - -static FailureOr -packCopyGmToCbufMultiConfig0(Operation *anchor, Value sid, - Value loop1SrcStride, Value l2CacheCtl, - Value nValue) { - OpBuilder builder(anchor); - builder.setInsertionPoint(anchor); - Location loc = anchor->getLoc(); - - Value sidI64 = castIntegerLikeTo(anchor, sid, builder.getI64Type()); - Value loop1SrcStrideI64 = - castIntegerLikeTo(anchor, loop1SrcStride, builder.getI64Type()); - Value l2CacheCtlI64 = castIntegerLikeTo(anchor, l2CacheCtl, builder.getI64Type()); - Value nValueI64 = castIntegerLikeTo(anchor, nValue, builder.getI64Type()); - if (!sidI64 || !loop1SrcStrideI64 || !l2CacheCtlI64 || !nValueI64) - return failure(); - - auto shl = [&](Value value, uint64_t amount) -> Value { - return builder.create(loc, value, - getI64Constant(builder, loc, amount)); - }; - auto bitOr = [&](Value lhs, Value rhs) -> Value { - return builder.create(loc, lhs, rhs); - }; - - Value config0 = sidI64; - config0 = bitOr(config0, shl(loop1SrcStrideI64, 4)); - config0 = bitOr(config0, shl(l2CacheCtlI64, 44)); - config0 = bitOr(config0, shl(nValueI64, 48)); - return config0; -} - -static FailureOr -packCopyGmToCbufMultiConfig1(Operation *anchor, Value dValue, - Value loop4SrcStride, Value smallC0En) { - OpBuilder builder(anchor); - builder.setInsertionPoint(anchor); - Location loc = anchor->getLoc(); - - Value dValueI64 = castIntegerLikeTo(anchor, dValue, builder.getI64Type()); - Value loop4SrcStrideI64 = - castIntegerLikeTo(anchor, loop4SrcStride, builder.getI64Type()); - Value smallC0EnI64 = castIntegerLikeTo(anchor, smallC0En, builder.getI64Type()); - if (!dValueI64 || !loop4SrcStrideI64 || !smallC0EnI64) - return failure(); - - auto shl = [&](Value value, uint64_t amount) -> Value { - return builder.create(loc, value, - getI64Constant(builder, loc, amount)); - }; - auto bitOr = [&](Value lhs, Value rhs) -> Value { - return builder.create(loc, lhs, rhs); - }; - - Value config1 = dValueI64; - config1 = bitOr(config1, shl(loop4SrcStrideI64, 21)); - config1 = bitOr(config1, shl(smallC0EnI64, 61)); - return config1; -} - -static FailureOr packCopyCbufToBtConfig(Operation *anchor, - Value convControl, - Value nBurst, Value lenBurst, - Value sourceGap, - Value dstGap) { - OpBuilder builder(anchor); - builder.setInsertionPoint(anchor); - Location loc = anchor->getLoc(); - - Value convControlI64 = - castIntegerLikeTo(anchor, convControl, builder.getI64Type()); - Value nBurstI64 = castIntegerLikeTo(anchor, nBurst, builder.getI64Type()); - Value lenBurstI64 = castIntegerLikeTo(anchor, lenBurst, builder.getI64Type()); - Value sourceGapI64 = castIntegerLikeTo(anchor, sourceGap, builder.getI64Type()); - Value dstGapI64 = castIntegerLikeTo(anchor, dstGap, builder.getI64Type()); - if (!convControlI64 || !nBurstI64 || !lenBurstI64 || !sourceGapI64 || - !dstGapI64) - return failure(); - - auto shl = [&](Value value, uint64_t amount) -> Value { - return builder.create(loc, value, - getI64Constant(builder, loc, amount)); - }; - auto bitOr = [&](Value lhs, Value rhs) -> Value { - return builder.create(loc, lhs, rhs); - }; - - Value config = shl(convControlI64, 3); - config = bitOr(config, shl(nBurstI64, 4)); - config = bitOr(config, shl(lenBurstI64, 16)); - config = bitOr(config, shl(sourceGapI64, 32)); - config = bitOr(config, shl(dstGapI64, 48)); - return config; -} - -static FailureOr packCopyCbufToFbufConfig(Operation *anchor, Value nBurst, - Value lenBurst, - Value sourceGap, - Value dstGap) { - OpBuilder builder(anchor); - builder.setInsertionPoint(anchor); - Location loc = anchor->getLoc(); - - Value nBurstI64 = castIntegerLikeTo(anchor, nBurst, builder.getI64Type()); - Value lenBurstI64 = castIntegerLikeTo(anchor, lenBurst, builder.getI64Type()); - Value sourceGapI64 = castIntegerLikeTo(anchor, sourceGap, builder.getI64Type()); - Value dstGapI64 = castIntegerLikeTo(anchor, dstGap, builder.getI64Type()); - if (!nBurstI64 || !lenBurstI64 || !sourceGapI64 || !dstGapI64) - return failure(); - - auto shl = [&](Value value, uint64_t amount) -> Value { - return builder.create(loc, value, - getI64Constant(builder, loc, amount)); - }; - auto bitOr = [&](Value lhs, Value rhs) -> Value { - return builder.create(loc, lhs, rhs); - }; - - Value config = shl(nBurstI64, 4); - config = bitOr(config, shl(lenBurstI64, 16)); - config = bitOr(config, shl(sourceGapI64, 32)); - config = bitOr(config, shl(dstGapI64, 48)); - return config; -} - -static FailureOr -packLoadCbufToS4Config0(Operation *anchor, Value mStart, Value kStart, - Value mStep, Value kStep) { - OpBuilder builder(anchor); - builder.setInsertionPoint(anchor); - Location loc = anchor->getLoc(); - - Value mStartI64 = castIntegerLikeTo(anchor, mStart, builder.getI64Type()); - Value kStartI64 = castIntegerLikeTo(anchor, kStart, builder.getI64Type()); - Value mStepI64 = castIntegerLikeTo(anchor, mStep, builder.getI64Type()); - Value kStepI64 = castIntegerLikeTo(anchor, kStep, builder.getI64Type()); - if (!mStartI64 || !kStartI64 || !mStepI64 || !kStepI64) - return failure(); - - auto shl = [&](Value value, uint64_t amount) -> Value { - return builder.create(loc, value, - getI64Constant(builder, loc, amount)); - }; - auto bitOr = [&](Value lhs, Value rhs) -> Value { - return builder.create(loc, lhs, rhs); - }; - - Value config0 = mStartI64; - config0 = bitOr(config0, shl(kStartI64, 16)); - config0 = bitOr(config0, shl(mStepI64, 32)); - config0 = bitOr(config0, shl(kStepI64, 40)); - return config0; -} - -static FailureOr -packLoadCbufToS4Config1(Operation *anchor, Value srcStride, Value dstStride) { - OpBuilder builder(anchor); - builder.setInsertionPoint(anchor); - Location loc = anchor->getLoc(); - - Value srcStrideI64 = castIntegerLikeTo(anchor, srcStride, builder.getI64Type()); - Value dstStrideI64 = castIntegerLikeTo(anchor, dstStride, builder.getI64Type()); - if (!srcStrideI64 || !dstStrideI64) - return failure(); - - auto shl = [&](Value value, uint64_t amount) -> Value { - return builder.create(loc, value, - getI64Constant(builder, loc, amount)); - }; - return builder.create(loc, srcStrideI64, shl(dstStrideI64, 16)) - .getResult(); -} - -static FailureOr -packLoadCbufToCaConfig0(Operation *anchor, Value mStart, Value kStart, - Value mStep, Value kStep) { - OpBuilder builder(anchor); - builder.setInsertionPoint(anchor); - Location loc = anchor->getLoc(); - - Value mStartI64 = castIntegerLikeTo(anchor, mStart, builder.getI64Type()); - Value kStartI64 = castIntegerLikeTo(anchor, kStart, builder.getI64Type()); - Value mStepI64 = castIntegerLikeTo(anchor, mStep, builder.getI64Type()); - Value kStepI64 = castIntegerLikeTo(anchor, kStep, builder.getI64Type()); - if (!mStartI64 || !kStartI64 || !mStepI64 || !kStepI64) - return failure(); - - auto shl = [&](Value value, uint64_t amount) -> Value { - return builder.create(loc, value, - getI64Constant(builder, loc, amount)); - }; - auto bitOr = [&](Value lhs, Value rhs) -> Value { - return builder.create(loc, lhs, rhs); - }; - - Value config0 = mStartI64; - config0 = bitOr(config0, shl(kStartI64, 16)); - config0 = bitOr(config0, shl(mStepI64, 32)); - config0 = bitOr(config0, shl(kStepI64, 40)); - return config0; -} - -static FailureOr -packLoadCbufToCaConfig1(Operation *anchor, Value srcStride, Value dstStride) { - OpBuilder builder(anchor); - builder.setInsertionPoint(anchor); - Location loc = anchor->getLoc(); - - Value srcStrideI64 = - castIntegerLikeTo(anchor, srcStride, builder.getI64Type()); - Value dstStrideI64 = - castIntegerLikeTo(anchor, dstStride, builder.getI64Type()); - if (!srcStrideI64 || !dstStrideI64) - return failure(); - - auto shl = [&](Value value, uint64_t amount) -> Value { - return builder.create(loc, value, - getI64Constant(builder, loc, amount)); - }; - return builder.create(loc, srcStrideI64, shl(dstStrideI64, 16)) - .getResult(); -} - -static FailureOr -packLoadCbufToCbConfig0(Operation *anchor, Value mStart, Value kStart, - Value mStep, Value kStep) { - OpBuilder builder(anchor); - builder.setInsertionPoint(anchor); - Location loc = anchor->getLoc(); - - Value mStartI64 = castIntegerLikeTo(anchor, mStart, builder.getI64Type()); - Value kStartI64 = castIntegerLikeTo(anchor, kStart, builder.getI64Type()); - Value mStepI64 = castIntegerLikeTo(anchor, mStep, builder.getI64Type()); - Value kStepI64 = castIntegerLikeTo(anchor, kStep, builder.getI64Type()); - if (!mStartI64 || !kStartI64 || !mStepI64 || !kStepI64) - return failure(); - - auto shl = [&](Value value, uint64_t amount) -> Value { - return builder.create(loc, value, - getI64Constant(builder, loc, amount)); - }; - auto bitOr = [&](Value lhs, Value rhs) -> Value { - return builder.create(loc, lhs, rhs); - }; - - Value config0 = mStartI64; - config0 = bitOr(config0, shl(kStartI64, 16)); - config0 = bitOr(config0, shl(mStepI64, 32)); - config0 = bitOr(config0, shl(kStepI64, 40)); - return config0; -} - -static FailureOr -packLoadCbufToCbConfig1(Operation *anchor, Value srcStride, Value dstStride) { - OpBuilder builder(anchor); - builder.setInsertionPoint(anchor); - Location loc = anchor->getLoc(); - - Value srcStrideI64 = - castIntegerLikeTo(anchor, srcStride, builder.getI64Type()); - Value dstStrideI64 = - castIntegerLikeTo(anchor, dstStride, builder.getI64Type()); - if (!srcStrideI64 || !dstStrideI64) - return failure(); - - auto shl = [&](Value value, uint64_t amount) -> Value { - return builder.create(loc, value, - getI64Constant(builder, loc, amount)); - }; - return builder.create(loc, srcStrideI64, shl(dstStrideI64, 16)) - .getResult(); -} - -static Value buildMadBiasDestination(Operation *anchor, - ConversionPatternRewriter &rewriter, - Value dst, Value bias) { - Type i64Ty = rewriter.getI64Type(); - Value dstAddr = rewriter.create(anchor->getLoc(), i64Ty, dst); - Value biasAddr = - rewriter.create(anchor->getLoc(), i64Ty, bias); - Value lowMask = getI64Constant(rewriter, anchor->getLoc(), 0xffffffffULL); - Value dstLow = rewriter.create(anchor->getLoc(), dstAddr, lowMask); - Value biasLow = rewriter.create(anchor->getLoc(), biasAddr, lowMask); - Value biasHigh = rewriter.create( - anchor->getLoc(), biasLow, getI64Constant(rewriter, anchor->getLoc(), 32)); - Value packed = rewriter.create(anchor->getLoc(), dstLow, biasHigh); - return rewriter.create(anchor->getLoc(), dst.getType(), packed); -} - -static FailureOr packVbitsortConfig(Operation *anchor, Value repeatTimes) { - OpBuilder builder(anchor); - builder.setInsertionPoint(anchor); - Location loc = anchor->getLoc(); - - Value repeatI64 = castIntegerLikeTo(anchor, repeatTimes, builder.getI64Type()); - if (!repeatI64) - return failure(); - return builder - .create(loc, repeatI64, getI64Constant(builder, loc, 56)) - .getResult(); -} - -[[maybe_unused]] static FailureOr -materializeDynamicPltMask(ConversionPatternRewriter &rewriter, - LoweringState &state, Location loc, Value laneCount, - Type vectorElemType) { - Type i32Type = rewriter.getI32Type(); - Value laneCountI32 = laneCount; - if (laneCountI32.getType() != i32Type) { - laneCountI32 = castIntegerLikeTo(rewriter.getInsertionBlock()->getParentOp(), - laneCountI32, i32Type); - if (!laneCountI32) - return failure(); - } - - StringRef calleeName; - if (vectorElemType.isF32()) { - calleeName = StringRef("llvm.hivm.plt.b32.v300"); - } else if (vectorElemType.isF16() || vectorElemType.isBF16()) { - calleeName = StringRef("llvm.hivm.plt.b16.v300"); - } else if (auto intType = dyn_cast(vectorElemType)) { - if (intType.getWidth() == 32) - calleeName = StringRef("llvm.hivm.plt.b32.v300"); - else if (intType.getWidth() == 16) - calleeName = StringRef("llvm.hivm.plt.b16.v300"); - else if (intType.getWidth() == 8) - calleeName = StringRef("llvm.hivm.plt.b8.v300"); - } - if (calleeName.empty()) - return failure(); - - Type maskType = VectorType::get({256}, rewriter.getI1Type()); - auto funcType = - rewriter.getFunctionType(TypeRange{i32Type}, TypeRange{maskType, i32Type}); - auto call = rewriter.create(loc, calleeName, funcType.getResults(), - ValueRange{laneCountI32}); - state.plannedDecls.push_back(PlannedDecl{calleeName.str(), funcType}); - return call.getResult(0); -} - -static FailureOr buildCarryBinaryCallee(MLIRContext *context, - Type resultType, - StringRef stem) { - std::string vec = - getElementTypeFragment(cast(resultType).getElementType()); - auto lanes = getElementCountFromVectorLike(resultType); - if (vec.empty() || !lanes) - return failure(); - return StringAttr::get(context, "llvm.hivm." + stem.str() + ".v" + - std::to_string(*lanes) + vec) - .getValue(); -} - -template -static StringRef getUnaryMaskedStem() { - if constexpr (std::is_same_v) - return "vabs"; - if constexpr (std::is_same_v) - return "vexp"; - if constexpr (std::is_same_v) - return "vln"; - if constexpr (std::is_same_v) - return "vneg"; - if constexpr (std::is_same_v) - return "vsqrt"; - if constexpr (std::is_same_v) - return "vrelu"; - if constexpr (std::is_same_v) - return "vnot"; - return {}; -} - -template -static FailureOr buildUnaryMaskedCallee(MLIRContext *context, - Type resultType) { - StringRef stem = getUnaryMaskedStem(); - if (stem.empty()) - return failure(); - return buildCANN900ModeTypedCallee(context, resultType, stem, "x"); -} - -template -static StringRef getBinaryMaskedStem() { - if constexpr (std::is_same_v) - return "vadd"; - if constexpr (std::is_same_v) - return "vsub"; - if constexpr (std::is_same_v) - return "vmul"; - if constexpr (std::is_same_v) - return "vdiv"; - if constexpr (std::is_same_v) - return "vmax"; - if constexpr (std::is_same_v) - return "vmin"; - if constexpr (std::is_same_v) - return "vand"; - if constexpr (std::is_same_v) - return "vor"; - if constexpr (std::is_same_v) - return "vxor"; - if constexpr (std::is_same_v) - return "vshl"; - if constexpr (std::is_same_v) - return "vshr"; - if constexpr (std::is_same_v) - return "vprelu"; - return {}; -} - -template -static StringRef getTernaryMaskedStem() { - if constexpr (std::is_same_v) - return "vmadd"; - return {}; -} - -template -static constexpr bool usesSignedBinaryCANN900Callee() { - return !std::is_same_v && - !std::is_same_v && - !std::is_same_v && - !std::is_same_v; -} - -template -static constexpr bool usesSignedTernaryCANN900Callee() { - return false; -} - -template -static StringRef getCarryBinaryStem() { - if constexpr (std::is_same_v) - return "vaddc"; - if constexpr (std::is_same_v) - return "vsubc"; - if constexpr (std::is_same_v) - return "vaddcs"; - if constexpr (std::is_same_v) - return "vsubcs"; - return {}; -} - -template -static constexpr bool hasCarryInput() { - return std::is_same_v || - std::is_same_v; -} - -static FailureOr buildVselCallee(MLIRContext *context, - Type resultType) { - std::string vec = getCANN900VectorTypeFragment(resultType); - if (vec.empty()) - return failure(); - return StringAttr::get(context, "llvm.hivm.vsel." + vec) - .getValue(); -} - -static FailureOr buildVselrCallee(MLIRContext *context, - Type resultType) { - Type elementType = getElementTypeFromVectorLike(resultType); - auto lanes = getElementCountFromVectorLike(resultType); - if (!elementType || !lanes) - return failure(); - - std::optional abi = - getLowpPayloadABI(elementType, context); - std::string vec = abi ? "v" + std::to_string(*lanes) + - abi->intrinsicElementFragment.str() - : getCANN900VectorTypeFragment(resultType); - if (vec.empty()) - return failure(); - return StringAttr::get(context, "llvm.hivm.vselr." + vec) - .getValue(); -} - -static FailureOr buildVdupCallee(MLIRContext *context, pto::VdupOp op) { - Type inputType = op.getInput().getType(); - Type resultType = op.getResult().getType(); - std::string vec = getCANN900VectorTypeFragment(resultType); - if (vec.empty()) - return failure(); - - if (isa(inputType)) { - StringRef position = op.getPosition().value_or("LOWEST"); - StringRef family = position == "HIGHEST" ? "vdupm" : "vdup"; - return StringAttr::get(context, "llvm.hivm." + family.str() + ".z." + vec) - .getValue(); - } - - return StringAttr::get(context, "llvm.hivm.vdups.z." + vec) - .getValue(); -} - -static FailureOr buildVbrCallee(MLIRContext *context, - Type resultType) { - std::string vec = getCANN900VectorTypeFragment(resultType); - if (vec.empty()) - return failure(); - return StringAttr::get(context, "llvm.hivm.vbr." + vec).getValue(); -} - -static FailureOr buildPstuCallee(MLIRContext *context, pto::PstuOp op) { - if (auto maskType = dyn_cast(op.getValue().getType())) { - if (maskType.isB16()) - return StringAttr::get(context, "llvm.hivm.pstu.b16").getValue(); - if (maskType.isB32()) - return StringAttr::get(context, "llvm.hivm.pstu.b32").getValue(); - } - return failure(); -} - -static FailureOr buildVstusCallee(MLIRContext *context, - Type valueType) { - std::string vec = - getMemoryElementTypeFragment(getElementTypeFromVectorLike(valueType)); - auto lanes = getElementCountFromVectorLike(valueType); - if (vec.empty() || !lanes) - return failure(); - return StringAttr::get(context, "llvm.hivm.vstus.v" + - std::to_string(*lanes) + vec) - .getValue(); -} - -static FailureOr buildVstusPostCallee(MLIRContext *context, - Type valueType) { - std::string vec = - getMemoryElementTypeFragment(getElementTypeFromVectorLike(valueType)); - auto lanes = getElementCountFromVectorLike(valueType); - if (vec.empty() || !lanes) - return failure(); - return StringAttr::get(context, "llvm.hivm.vstus.post.v" + - std::to_string(*lanes) + vec) - .getValue(); -} - -static StringRef buildVsturCallee(MLIRContext *context) { - return StringAttr::get(context, "llvm.hivm.vstur").getValue(); -} - -static StringRef buildInitAlignCallee(MLIRContext *context) { - return StringAttr::get(context, "llvm.hivm.init.vector.align.data").getValue(); -} - -template -static StringRef buildRuntimeQueryCallee(MLIRContext *context); - -template <> -StringRef buildRuntimeQueryCallee(MLIRContext *context) { - return StringAttr::get(context, "llvm.hivm.GET.CTRL").getValue(); -} - -template <> -StringRef buildRuntimeQueryCallee(MLIRContext *context) { - return StringAttr::get(context, "llvm.hivm.GET.VMS4.SR").getValue(); -} - -template <> -StringRef buildRuntimeQueryCallee(MLIRContext *context) { - return StringAttr::get(context, "llvm.hivm.get.TID.X").getValue(); -} - -template <> -StringRef buildRuntimeQueryCallee(MLIRContext *context) { - return StringAttr::get(context, "llvm.hivm.get.TID.Y").getValue(); -} - -template <> -StringRef buildRuntimeQueryCallee(MLIRContext *context) { - return StringAttr::get(context, "llvm.hivm.get.TID.Z").getValue(); -} - -template <> -StringRef buildRuntimeQueryCallee(MLIRContext *context) { - return StringAttr::get(context, "llvm.hivm.get.BLOCK.DIM.X").getValue(); -} - -template <> -StringRef buildRuntimeQueryCallee(MLIRContext *context) { - return StringAttr::get(context, "llvm.hivm.get.BLOCK.DIM.Y").getValue(); -} - -template <> -StringRef buildRuntimeQueryCallee(MLIRContext *context) { - return StringAttr::get(context, "llvm.hivm.get.BLOCK.DIM.Z").getValue(); -} - -template <> -StringRef buildRuntimeQueryCallee(MLIRContext *context) { - return StringAttr::get(context, "llvm.hivm.get.GRID.DIM.X").getValue(); -} - -template <> -StringRef buildRuntimeQueryCallee(MLIRContext *context) { - return StringAttr::get(context, "llvm.hivm.get.GRID.DIM.Y").getValue(); -} - -template <> -StringRef buildRuntimeQueryCallee(MLIRContext *context) { - return StringAttr::get(context, "llvm.hivm.get.GRID.DIM.Z").getValue(); -} - -template <> -StringRef buildRuntimeQueryCallee(MLIRContext *context) { - return StringAttr::get(context, "llvm.hivm.get.BLOCK.IDX.X").getValue(); -} - -template <> -StringRef buildRuntimeQueryCallee(MLIRContext *context) { - return StringAttr::get(context, "llvm.hivm.get.BLOCK.IDX.Y").getValue(); -} - -template <> -StringRef buildRuntimeQueryCallee(MLIRContext *context) { - return StringAttr::get(context, "llvm.hivm.get.BLOCK.IDX.Z").getValue(); -} - -template <> -StringRef buildRuntimeQueryCallee(MLIRContext *context) { - return StringAttr::get(context, "llvm.hivm.tpe.get.VECCOREID").getValue(); -} - -template <> -StringRef buildRuntimeQueryCallee(MLIRContext *context) { - return StringAttr::get(context, "llvm.hivm.get.laneID").getValue(); -} - -template <> -StringRef buildRuntimeQueryCallee(MLIRContext *context) { - return StringAttr::get(context, "llvm.hivm.get.CLOCK32").getValue(); -} - -template <> -StringRef buildRuntimeQueryCallee(MLIRContext *context) { - return StringAttr::get(context, "llvm.hivm.get.CLOCK64").getValue(); -} - -template <> -StringRef buildRuntimeQueryCallee(MLIRContext *context) { - return StringAttr::get(context, "llvm.hivm.get.LANEMASK.EQ").getValue(); -} - -template <> -StringRef buildRuntimeQueryCallee(MLIRContext *context) { - return StringAttr::get(context, "llvm.hivm.get.LANEMASK.LE").getValue(); -} - -template <> -StringRef buildRuntimeQueryCallee(MLIRContext *context) { - return StringAttr::get(context, "llvm.hivm.get.LANEMASK.LT").getValue(); -} - -template <> -StringRef buildRuntimeQueryCallee(MLIRContext *context) { - return StringAttr::get(context, "llvm.hivm.get.LANEMASK.GE").getValue(); -} - -template <> -StringRef buildRuntimeQueryCallee(MLIRContext *context) { - return StringAttr::get(context, "llvm.hivm.get.LANEMASK.GT").getValue(); -} - -static StringRef buildSprclrCallee(MLIRContext *context) { - return StringAttr::get(context, "llvm.hivm.sprclr").getValue(); -} - -static StringRef buildSprstiCallee(MLIRContext *context, bool post) { - return StringAttr::get(context, - post ? "llvm.hivm.sprsti.post" - : "llvm.hivm.sprsti") - .getValue(); -} - -static StringRef buildSprstsCallee(MLIRContext *context, bool post) { - return StringAttr::get(context, - post ? "llvm.hivm.sprsts.post" - : "llvm.hivm.sprsts") - .getValue(); -} - -template -static StringRef buildSprStoreCallee(MLIRContext *context, bool post); - -template <> -StringRef buildSprStoreCallee(MLIRContext *context, bool post) { - return buildSprstiCallee(context, post); -} - -template <> -StringRef buildSprStoreCallee(MLIRContext *context, bool post) { - return buildSprstsCallee(context, post); -} - -template -static StringRef buildUnaryConfigCallee(MLIRContext *context); - -template <> -StringRef buildUnaryConfigCallee(MLIRContext *context) { - return StringAttr::get(context, "llvm.hivm.SET.CTRL").getValue(); -} - -static StringRef buildStoreVfSimtInfoCallee(MLIRContext *context) { - return StringAttr::get(context, "llvm.hivm.store.vfsimt.info").getValue(); -} - -static StringRef buildSyncthreadsCallee(MLIRContext *context) { - return StringAttr::get(context, "llvm.hivm.sync.workitems").getValue(); -} - -static StringRef buildThreadfenceCallee(MLIRContext *context) { - return StringAttr::get(context, "llvm.hivm.fence.workitems").getValue(); -} - -static StringRef buildThreadfenceBlockCallee(MLIRContext *context) { - return StringAttr::get(context, "llvm.hivm.fenceblock.workitems").getValue(); -} - -static StringRef buildVstarCallee(MLIRContext *context) { - return StringAttr::get(context, "llvm.hivm.vstar").getValue(); -} - -static StringRef buildVstasCallee(MLIRContext *context, bool post) { - return StringAttr::get(context, - post ? "llvm.hivm.vstas.post" - : "llvm.hivm.vstas") - .getValue(); -} - -template -static StringRef buildVoteCallee(MLIRContext *context); - -template <> -StringRef buildVoteCallee(MLIRContext *context) { - return StringAttr::get(context, "llvm.hivm.vote.all").getValue(); -} - -template <> -StringRef buildVoteCallee(MLIRContext *context) { - return StringAttr::get(context, "llvm.hivm.vote.any").getValue(); -} - -template <> -StringRef buildVoteCallee(MLIRContext *context) { - return StringAttr::get(context, "llvm.hivm.vote.uni").getValue(); -} - -template <> -StringRef buildVoteCallee(MLIRContext *context) { - return StringAttr::get(context, "llvm.hivm.vote.ballot").getValue(); -} - -template -static StringRef buildBinaryI64PureCallee(MLIRContext *context); - -template <> -StringRef buildBinaryI64PureCallee(MLIRContext *context) { - return StringAttr::get(context, "llvm.hivm.SBITSET0").getValue(); -} - -template <> -StringRef buildBinaryI64PureCallee(MLIRContext *context) { - return StringAttr::get(context, "llvm.hivm.SBITSET1").getValue(); -} - -template -static FailureOr buildShuffleCallee(MLIRContext *context, - Type valueType); - -template <> -FailureOr buildShuffleCallee(MLIRContext *context, - Type valueType) { - std::string elem = getShuffleIntrinsicTypeFragment(valueType); - if (elem.empty()) - return failure(); - return StringAttr::get(context, "llvm.hivm.shfl.idx." + elem).getValue(); -} - -template <> -FailureOr buildShuffleCallee(MLIRContext *context, - Type valueType) { - std::string elem = getShuffleIntrinsicTypeFragment(valueType); - if (elem.empty()) - return failure(); - return StringAttr::get(context, "llvm.hivm.shfl.up." + elem).getValue(); -} - -template <> -FailureOr buildShuffleCallee(MLIRContext *context, - Type valueType) { - std::string elem = getShuffleIntrinsicTypeFragment(valueType); - if (elem.empty()) - return failure(); - return StringAttr::get(context, "llvm.hivm.shfl.down." + elem).getValue(); -} - -template <> -FailureOr buildShuffleCallee(MLIRContext *context, - Type valueType) { - std::string elem = getShuffleIntrinsicTypeFragment(valueType); - if (elem.empty()) - return failure(); - return StringAttr::get(context, "llvm.hivm.shfl.bfly." + elem).getValue(); -} - -static Value buildShuffleControlValue(OpBuilder &builder, Location loc, - Value controlValue, int64_t widthValue, - unsigned controlMask) { - Value lowBits = builder.create( - loc, controlValue, getI32Constant(builder, loc, 0x1f)); - Value encodedWidth = - getI32Constant(builder, loc, static_cast(32 - widthValue) << 16); - Value encodedMask = - getI32Constant(builder, loc, static_cast(controlMask) << 8); - Value highBits = builder.create(loc, encodedWidth, encodedMask); - return builder.create(loc, highBits, lowBits); -} - -template -static FailureOr buildReduxCallee(MLIRContext *context, - Type valueType, - Attribute signednessAttr); - -template <> -FailureOr buildReduxCallee(MLIRContext *context, - Type valueType, - Attribute signednessAttr) { - std::string elem = getReduxIntrinsicTypeFragment(valueType, signednessAttr); - if (elem.empty()) - return failure(); - return StringAttr::get(context, "llvm.hivm.redux.add." + elem).getValue(); -} - -template <> -FailureOr buildReduxCallee(MLIRContext *context, - Type valueType, - Attribute signednessAttr) { - std::string elem = getReduxIntrinsicTypeFragment(valueType, signednessAttr); - if (elem.empty()) - return failure(); - return StringAttr::get(context, "llvm.hivm.redux.max." + elem).getValue(); -} - -template <> -FailureOr buildReduxCallee(MLIRContext *context, - Type valueType, - Attribute signednessAttr) { - std::string elem = getReduxIntrinsicTypeFragment(valueType, signednessAttr); - if (elem.empty()) - return failure(); - return StringAttr::get(context, "llvm.hivm.redux.min." + elem).getValue(); -} - -template -static FailureOr buildAtomicCallee(MLIRContext *context, - Type ptrType, Type valueType, - Attribute signednessAttr); - -static FailureOr buildAtomicCalleeName(MLIRContext *context, - Type ptrType, Type valueType, - Attribute signednessAttr, - StringRef opName) { - std::string elem = getAtomicElementTypeFragment(valueType, signednessAttr); - if (elem.empty()) - return failure(); - auto ptrTy = dyn_cast(ptrType); - if (!ptrTy) - return failure(); - - StringRef space; - switch (ptrTy.getMemorySpace().getAddressSpace()) { - case pto::AddressSpace::GM: - space = "G"; - break; - case pto::AddressSpace::VEC: - if (valueType.isInteger(64)) - return failure(); - space = "S"; - break; - default: - return failure(); - } - - return StringAttr::get(context, "llvm.hivm.atom." + opName.str() + "." + - space.str() + "." + elem) - .getValue(); -} - -#define PTO_BUILD_ATOMIC_CALLEE(OP, NAME) \ - template <> \ - [[maybe_unused]] FailureOr buildAtomicCallee( \ - MLIRContext *context, Type ptrType, Type valueType, \ - Attribute signednessAttr) { \ - return buildAtomicCalleeName(context, ptrType, valueType, signednessAttr, \ - NAME); \ - } - -PTO_BUILD_ATOMIC_CALLEE(AtomicCasOp, "CAS") -PTO_BUILD_ATOMIC_CALLEE(AtomicExchOp, "EXCH") -PTO_BUILD_ATOMIC_CALLEE(AtomicAddOp, "ADD") -PTO_BUILD_ATOMIC_CALLEE(AtomicSubOp, "SUB") -PTO_BUILD_ATOMIC_CALLEE(AtomicMinOp, "MIN") -PTO_BUILD_ATOMIC_CALLEE(AtomicMaxOp, "MAX") -PTO_BUILD_ATOMIC_CALLEE(AtomicAndOp, "AND") -PTO_BUILD_ATOMIC_CALLEE(AtomicOrOp, "OR") -PTO_BUILD_ATOMIC_CALLEE(AtomicXorOp, "XOR") - -#undef PTO_BUILD_ATOMIC_CALLEE - -static FailureOr buildL1CacheLoadCallee(MLIRContext *context, - Type resultType, - pto::L1Cache l1cache) { - std::string elem; - if (auto intType = dyn_cast(resultType)) { - if (intType.getWidth() == 8) - elem = "s8"; - else if (intType.getWidth() == 16) - elem = "s16"; - else if (intType.getWidth() == 32) - elem = "s32"; - else if (intType.getWidth() == 64) - elem = "s64"; - } else if (resultType.isF16() || resultType.isBF16()) { - elem = "s16"; - } else if (resultType.isF32()) { - elem = "s32"; - } else if (resultType.isF64()) { - elem = "s64"; - } else if (pto::isPTOFloat8Type(resultType) || - pto::isPTOHiFloat8Type(resultType)) { - elem = "s8"; - } else if (pto::isPTOPackedLdgStgVectorType(resultType)) { - unsigned totalBits = pto::getPTOPackedLdgStgTotalBits(resultType); - if (totalBits == 16) - elem = "s16"; - else if (totalBits == 32) - elem = "s32"; - else if (totalBits == 64) - elem = "s64"; - } - if (elem.empty()) - return failure(); - StringRef l1cacheName = - l1cache == pto::L1Cache::Cache ? "cache" : "uncache"; - return StringAttr::get(context, - "llvm.hivm.ldg." + l1cacheName.str() + "." + elem) - .getValue(); -} - -static FailureOr buildL1CacheStoreCallee(MLIRContext *context, - Type valueType, - pto::L1Cache l1cache) { - std::string elem; - if (auto intType = dyn_cast(valueType)) { - if (intType.getWidth() == 8) - elem = "b8"; - else if (intType.getWidth() == 16) - elem = "b16"; - else if (intType.getWidth() == 32) - elem = "b32"; - else if (intType.getWidth() == 64) - elem = "b64"; - } else if (valueType.isF16() || valueType.isBF16()) { - elem = "b16"; - } else if (valueType.isF32()) { - elem = "b32"; - } else if (valueType.isF64()) { - elem = "b64"; - } else if (pto::isPTOFloat8Type(valueType) || - pto::isPTOHiFloat8Type(valueType)) { - elem = "b8"; - } else if (pto::isPTOPackedLdgStgVectorType(valueType)) { - unsigned totalBits = pto::getPTOPackedLdgStgTotalBits(valueType); - if (totalBits == 16) - elem = "b16"; - else if (totalBits == 32) - elem = "b32"; - else if (totalBits == 64) - elem = "b64"; - } - if (elem.empty()) - return failure(); - StringRef l1cacheName = - l1cache == pto::L1Cache::Cache ? "cache" : "uncache"; - return StringAttr::get(context, - "llvm.hivm.stg." + l1cacheName.str() + "." + elem) - .getValue(); -} - -template -static StringRef buildScalarIntrinsicCallee(MLIRContext *context); - -template <> -StringRef buildScalarIntrinsicCallee(MLIRContext *context) { - return StringAttr::get(context, "llvm.hivm.prmt").getValue(); -} - -static FailureOr -buildMulhiCallee(MLIRContext *context, Type resultType, - pto::Signedness signedness) { - if (resultType.isInteger(32)) { - return StringAttr::get( - context, signedness == pto::Signedness::Unsigned - ? "llvm.hivm.mulhi.ui" - : "llvm.hivm.mulhi.i") - .getValue(); - } - if (resultType.isInteger(64) && signedness == pto::Signedness::Unsigned) - return StringAttr::get(context, "llvm.hivm.mul64hi.ui").getValue(); - return failure(); -} - -static FailureOr -buildMulI32ToI64Callee(MLIRContext *context, pto::Signedness signedness) { - return StringAttr::get( - context, signedness == pto::Signedness::Unsigned - ? "llvm.hivm.mul.i32toi64.ui" - : "llvm.hivm.mul.i32toi64.i") - .getValue(); -} - -static std::string getScalarFloatBuiltinFragment(Type type) { - if (type.isF32()) - return "f32"; - if (type.isF16()) - return "f16"; - if (type.isBF16()) - return "bf16"; - return {}; -} - -static std::string getLLVMFloatBuiltinFragment(Type type) { - std::string scalar = getScalarFloatBuiltinFragment(type); - if (!scalar.empty()) - return scalar; - - auto vecType = dyn_cast(type); - if (!vecType || vecType.getRank() != 1 || vecType.getDimSize(0) != 2) - return {}; - Type elementType = vecType.getElementType(); - if (elementType.isF16()) - return "v2f16"; - if (elementType.isBF16()) - return "v2bf16"; - return {}; -} - -static std::string getHIVMFloatBuiltinFragment(Type type) { - std::string scalar = getScalarFloatBuiltinFragment(type); - if (!scalar.empty()) - return scalar; - - auto vecType = dyn_cast(type); - if (!vecType || vecType.getRank() != 1 || vecType.getDimSize(0) != 2) - return {}; - Type elementType = vecType.getElementType(); - if (elementType.isF16()) - return "f16x2"; - if (elementType.isBF16()) - return "bf16x2"; - return {}; -} - -static FailureOr buildSqrtCallee(MLIRContext *context, Type valueType) { - std::string elem = getLLVMFloatBuiltinFragment(valueType); - if (elem != "f32" && elem != "f16" && elem != "v2f16") - return failure(); - return StringAttr::get(context, "llvm.sqrt." + elem).getValue(); -} - -static std::string getScalarHIVMFloatShortFragment(Type type) { - if (type.isF32()) - return "f"; - if (type.isF16()) - return "h"; - if (type.isBF16()) - return "y"; - return {}; -} - -template -static FailureOr buildUnaryScalarMathCallee(MLIRContext *context, - Type valueType); - -template <> -FailureOr buildUnaryScalarMathCallee(MLIRContext *context, - Type valueType) { - std::string elem = getLLVMFloatBuiltinFragment(valueType); - if (elem != "f16" && elem != "f32" && elem != "v2f16" && elem != "v2bf16") - return failure(); - return StringAttr::get(context, "llvm.fabs." + elem).getValue(); -} - -template <> -FailureOr buildUnaryScalarMathCallee(MLIRContext *context, - Type valueType) { - std::string elem = getLLVMFloatBuiltinFragment(valueType); - if (elem != "f32" && elem != "f16" && elem != "v2f16") - return failure(); - return StringAttr::get(context, "llvm.exp." + elem).getValue(); -} - -template <> -FailureOr buildUnaryScalarMathCallee(MLIRContext *context, - Type valueType) { - std::string elem = getLLVMFloatBuiltinFragment(valueType); - if (elem != "f32" && elem != "f16" && elem != "v2f16") - return failure(); - return StringAttr::get(context, "llvm.log." + elem).getValue(); -} - -template <> -FailureOr buildUnaryScalarMathCallee(MLIRContext *context, - Type valueType) { - std::string elem = getScalarHIVMFloatShortFragment(valueType); - if (elem.empty()) - return failure(); - return StringAttr::get(context, "llvm.hivm.ceil." + elem).getValue(); -} - -template <> -FailureOr buildUnaryScalarMathCallee(MLIRContext *context, - Type valueType) { - std::string elem = getScalarHIVMFloatShortFragment(valueType); - if (elem.empty()) - return failure(); - return StringAttr::get(context, "llvm.hivm.floor." + elem).getValue(); -} - -template <> -FailureOr buildUnaryScalarMathCallee(MLIRContext *context, - Type valueType) { - std::string elem = getScalarHIVMFloatShortFragment(valueType); - if (elem.empty()) - return failure(); - return StringAttr::get(context, "llvm.hivm.rint." + elem).getValue(); -} - -template <> -FailureOr buildUnaryScalarMathCallee(MLIRContext *context, - Type valueType) { - std::string elem = getScalarHIVMFloatShortFragment(valueType); - if (elem.empty()) - return failure(); - return StringAttr::get(context, "llvm.hivm.round." + elem).getValue(); -} - -template -static FailureOr buildBinaryScalarMathCallee(MLIRContext *context, - Type valueType); - -template <> -FailureOr buildBinaryScalarMathCallee(MLIRContext *context, - Type valueType) { - std::string elem = getLLVMFloatBuiltinFragment(valueType); - if (elem != "f16" && elem != "f32" && elem != "bf16" && - elem != "v2f16" && elem != "v2bf16") - return failure(); - return StringAttr::get(context, "llvm.minnum." + elem).getValue(); -} - -template <> -FailureOr buildBinaryScalarMathCallee(MLIRContext *context, - Type valueType) { - std::string elem = getLLVMFloatBuiltinFragment(valueType); - if (elem != "f16" && elem != "f32" && elem != "bf16" && - elem != "v2f16" && elem != "v2bf16") - return failure(); - return StringAttr::get(context, "llvm.maxnum." + elem).getValue(); -} - -template <> -FailureOr buildBinaryScalarMathCallee(MLIRContext *context, - Type valueType) { - std::string elem = getLLVMFloatBuiltinFragment(valueType); - if (elem != "f32" && elem != "f16" && elem != "v2f16") - return failure(); - return StringAttr::get(context, "llvm.pow." + elem).getValue(); -} - -static FailureOr buildFmaCallee(MLIRContext *context, Type valueType) { - std::string elem = getHIVMFloatBuiltinFragment(valueType); - if (elem.empty()) - return failure(); - return StringAttr::get(context, "llvm.hivm.ffma." + elem + ".rrr").getValue(); -} - -static std::string getConvertScalarFragment(Type type, - Attribute signednessAttr) { - if (auto vecType = dyn_cast(type)) { - if (vecType.getRank() != 1 || vecType.getDimSize(0) != 2) - return {}; - Type elementType = vecType.getElementType(); - if (std::string elem = getLowPrecisionElementFragment(elementType); - !elem.empty() && !pto::isPTOFloat4PackedType(elementType)) - return elem + "x2"; - if (elementType.isF32()) - return "f32x2"; - if (elementType.isF16()) - return "f16x2"; - if (elementType.isBF16()) - return "bf16x2"; - return {}; - } - if (type.isF32()) - return "fp32"; - if (type.isF16()) - return "fp16"; - if (type.isBF16()) - return "bf16"; - if (std::string elem = getLowPrecisionElementFragment(type); !elem.empty()) - return elem; - auto intType = dyn_cast(type); - if (!intType || (intType.getWidth() != 32 && intType.getWidth() != 64) || - !signednessAttr) - return {}; - auto signedness = cast(signednessAttr).getValue(); - return std::string(signedness == pto::Signedness::Unsigned ? "u" : "s") + - std::to_string(intType.getWidth()); -} - -static FailureOr buildConvertCallee(MLIRContext *context, - Type srcType, Type dstType, - Attribute signednessAttr) { - std::string src = getConvertScalarFragment(srcType, signednessAttr); - std::string dst = getConvertScalarFragment(dstType, signednessAttr); - if (src.empty() || dst.empty()) - return failure(); - return StringAttr::get(context, - "llvm.hivm." + src + ".to." + dst) - .getValue(); -} - -static FailureOr buildVldsPostCallee(MLIRContext *context, - Type resultType) { - std::string vec = - getMemoryElementTypeFragment(getElementTypeFromVectorLike(resultType)); - auto lanes = getElementCountFromVectorLike(resultType); - if (vec.empty() || !lanes) - return failure(); - return StringAttr::get(context, "llvm.hivm.vldsx1.post.v" + - std::to_string(*lanes) + vec) - .getValue(); -} - -static FailureOr buildVstsPostCallee(MLIRContext *context, - Type valueType) { - std::string vec = - getMemoryElementTypeFragment(getElementTypeFromVectorLike(valueType)); - auto lanes = getElementCountFromVectorLike(valueType); - if (vec.empty() || !lanes) - return failure(); - return StringAttr::get(context, "llvm.hivm.vstsx1.post.v" + - std::to_string(*lanes) + vec) - .getValue(); -} - -static StringRef buildVldasCallee(MLIRContext *context) { - return StringAttr::get(context, "llvm.hivm.vldas").getValue(); -} - -static FailureOr buildVldusCallee(MLIRContext *context, - Type resultType) { - std::string vec = - getMemoryElementTypeFragment(getElementTypeFromVectorLike(resultType)); - auto lanes = getElementCountFromVectorLike(resultType); - if (vec.empty() || !lanes) - return failure(); - return StringAttr::get(context, "llvm.hivm.vldus.v" + - std::to_string(*lanes) + vec) - .getValue(); -} - -static FailureOr buildVldusPostCallee(MLIRContext *context, - Type resultType) { - std::string vec = - getMemoryElementTypeFragment(getElementTypeFromVectorLike(resultType)); - auto lanes = getElementCountFromVectorLike(resultType); - if (vec.empty() || !lanes) - return failure(); - return StringAttr::get(context, "llvm.hivm.vldus.post.v" + - std::to_string(*lanes) + vec) - .getValue(); -} - -static FailureOr buildVcmpCallee(MLIRContext *context, Type inputType, - StringRef cmpMode, - bool isScalarCompare) { - std::string vec = getCANN900VectorTypeFragment(inputType); - std::string signedness = - getCANN900SignednessFragment(getElementTypeFromVectorLike(inputType)); - if (vec.empty() || signedness.empty()) - return failure(); - StringRef stem = isScalarCompare ? "vcmps" : "vcmp"; - return StringAttr::get(context, "llvm.hivm." + stem.str() + "." + - cmpMode.str() + "." + signedness + - ".z." + vec) - .getValue(); -} - -template -static StringRef getVecScalarMaskedStem() { - if constexpr (std::is_same_v) - return "vmuls"; - if constexpr (std::is_same_v) - return "vadds"; - if constexpr (std::is_same_v) - return "vmaxs"; - if constexpr (std::is_same_v) - return "vmins"; - if constexpr (std::is_same_v) - return "vlrelu"; - if constexpr (std::is_same_v) - return "vshls"; - if constexpr (std::is_same_v) - return "vshrs"; - return {}; -} - -template -static constexpr bool usesSignedVecScalarCANN900Callee() { - return !std::is_same_v; -} - -template -static StringRef getReductionUnaryStem() { - if constexpr (std::is_same_v) - return "vcadd"; - if constexpr (std::is_same_v) - return "vcmax"; - if constexpr (std::is_same_v) - return "vcmin"; - if constexpr (std::is_same_v) - return "vcgadd"; - if constexpr (std::is_same_v) - return "vcgmax"; - if constexpr (std::is_same_v) - return "vcgmin"; - if constexpr (std::is_same_v) - return "vcpadd"; - return {}; -} - -template -static StringRef getHistogramCallee(MLIRContext *context) { - if constexpr (std::is_same_v) - return StringAttr::get(context, "llvm.hivm.chistv2.m").getValue(); - if constexpr (std::is_same_v) - return StringAttr::get(context, "llvm.hivm.dhistv2.m").getValue(); - return {}; -} - -template -static StringRef getExtremaPredicateStem() { - if constexpr (std::is_same_v) - return "vcbmax"; - if constexpr (std::is_same_v) - return "vcbmin"; - return {}; -} - -template -static FailureOr buildExtremaPredicateCallee(MLIRContext *context, - Type resultType) { - return buildCANN900SignedModeTypedCallee( - context, resultType, getExtremaPredicateStem(), "x"); -} - -template -static constexpr bool usesSignedReductionCANN900Callee() { - return !std::is_same_v; -} - -static FailureOr buildCopyGmToUbCallee(MLIRContext *context, - Type sourceType) { - auto ptrType = dyn_cast(sourceType); - if (!ptrType) - return failure(); - Type elementType = ptrType.getElementType(); - if ((isa(elementType) && - cast(elementType).getWidth() == 64) || - elementType.isF64()) { - return StringAttr::get(context, "llvm.hivm.MOV.OUT.TO.UB.ALIGN.V2.s32.DV") - .getValue(); - } - std::string elem = getCopyElementFragment(elementType); - if (elem.empty()) - return failure(); - return StringAttr::get(context, "llvm.hivm.MOV.OUT.TO.UB.ALIGN.V2." + elem + - ".DV") - .getValue(); -} - -static StringRef buildCopyUbToGmCallee(MLIRContext *context) { - return StringAttr::get(context, "llvm.hivm.MOV.UB.TO.OUT.ALIGN.V2.DV") - .getValue(); -} - -static StringRef buildCopyUbToUbCallee(MLIRContext *context) { - return StringAttr::get(context, "llvm.hivm.MOV.UB.TO.UB.v310").getValue(); -} - -static StringRef buildCopyCbufToUbCallee(MLIRContext *context) { - return StringAttr::get(context, "llvm.hivm.MOV.L1.TO.UB.v310").getValue(); -} - -static StringRef buildCopyUbToCbufCallee(MLIRContext *context) { - return StringAttr::get(context, "llvm.hivm.MOV.UB.TO.L1.v310").getValue(); -} - -static FailureOr buildOrdinaryMadCallee(MLIRContext *context, - pto::MadRawOpInterface op) { - auto lhsType = dyn_cast(op.getLhs().getType()); - auto rhsType = dyn_cast(op.getRhs().getType()); - auto dstType = dyn_cast(op.getDst().getType()); - if (!lhsType || !rhsType || !dstType) - return failure(); - - return buildMadTypedCalleeName(context, lhsType.getElementType(), - rhsType.getElementType(), - dstType.getElementType()); -} - -static FailureOr buildMxMadCallee(MLIRContext *context, - pto::MadRawOpInterface op) { - auto lhsType = dyn_cast(op.getLhs().getType()); - auto rhsType = dyn_cast(op.getRhs().getType()); - if (!lhsType || !rhsType) - return failure(); - if (isMxElementType(lhsType.getElementType()) && - isMxElementType(rhsType.getElementType())) { - return buildMadMxCalleeName(context, lhsType.getElementType(), - rhsType.getElementType()); - } - return failure(); -} - -static FailureOr buildCopyGmToCbufCallee(MLIRContext *context, - Type sourceType) { - auto ptrType = dyn_cast(sourceType); - if (!ptrType) - return failure(); - std::string elem = getCopyElementFragment(ptrType.getElementType()); - if (elem.empty()) - return failure(); - return StringAttr::get(context, "llvm.hivm.MOV.OUT.TO.L1.ALIGN.V2." + elem + - ".DV") - .getValue(); -} - -static FailureOr -buildCopyGmToCbufMultiNd2NzCallee(MLIRContext *context, Type sourceType) { - auto ptrType = dyn_cast(sourceType); - if (!ptrType) - return failure(); - std::string elem = getNd2NzCopyElementFragment(ptrType.getElementType()); - if (elem.empty()) - return failure(); - return StringAttr::get(context, "llvm.hivm.MOV.OUT.TO.L1.MULTI.ND2NZ." + - elem + ".V310") - .getValue(); -} - -static std::string getDn2NzCopyElementFragment(Type type) { - auto ptrType = dyn_cast(type); - if (!ptrType) - return {}; - - Type elementType = ptrType.getElementType(); - std::string typeText; - llvm::raw_string_ostream os(typeText); - elementType.print(os); - os.flush(); - std::string lower = StringRef(typeText).lower(); - if (StringRef(lower).contains("e4m3") || StringRef(lower).contains("e5m2") || - StringRef(lower).contains("e8m0") || StringRef(lower).contains("hif8")) - return "u8"; - - if (elementType.isF16() || elementType.isBF16()) - return "u16"; - if (elementType.isF32()) - return "u32"; - - if (auto intType = dyn_cast(elementType)) { - switch (intType.getWidth()) { - case 8: - return "u8"; - case 16: - return "u16"; - case 32: - return "u32"; - default: - return {}; - } - } - return {}; -} - -static FailureOr -buildCopyGmToCbufMultiDn2NzCallee(MLIRContext *context, Type sourceType) { - auto ptrType = dyn_cast(sourceType); - if (!ptrType) - return failure(); - std::string elem = getDn2NzCopyElementFragment(sourceType); - if (elem.empty()) - return failure(); - return StringAttr::get(context, - "llvm.hivm.MOV.OUT.TO.L1.MULTI.DN2NZ." + elem) - .getValue(); -} - -static FailureOr buildLoadCbufToCaCallee(MLIRContext *context, - Type sourceType) { - auto ptrType = dyn_cast(sourceType); - if (!ptrType) - return failure(); - std::string elem = getL0LoadElementFragment(ptrType.getElementType()); - if (elem.empty()) - return failure(); - return StringAttr::get(context, "llvm.hivm.LOAD.L1.TO.L0A.2Dv2." + elem) - .getValue(); -} - -static FailureOr buildLoadCbufToCbCallee(MLIRContext *context, - Type sourceType) { - auto ptrType = dyn_cast(sourceType); - if (!ptrType) - return failure(); - std::string elem = getL0LoadElementFragment(ptrType.getElementType()); - if (elem.empty()) - return failure(); - return StringAttr::get(context, "llvm.hivm.LOAD.L1.TO.L0B.2Dv2." + elem) - .getValue(); -} - -static FailureOr buildLoadCbufToCaS4Callee(MLIRContext *context, - Type sourceType) { - auto ptrType = dyn_cast(sourceType); - if (!ptrType) - return failure(); - Type elementType = ptrType.getElementType(); - if (!isa(elementType)) - return failure(); - return StringAttr::get(context, "llvm.hivm.LOAD.L1.TO.L0A.2Dv2.s4") - .getValue(); -} - -static FailureOr buildLoadCbufToCbS4Callee(MLIRContext *context, - Type sourceType) { - auto ptrType = dyn_cast(sourceType); - if (!ptrType) - return failure(); - Type elementType = ptrType.getElementType(); - if (!isa(elementType)) - return failure(); - return StringAttr::get(context, "llvm.hivm.LOAD.L1.TO.L0B.2Dv2.s4") - .getValue(); -} - -static StringRef buildLoadCbufToCaMxCallee(MLIRContext *context) { - return StringAttr::get(context, "llvm.hivm.LOAD.L1.TO.L0A.MX.2Dv2.v") - .getValue(); -} - -[[maybe_unused]] static StringRef buildLoadCbufToCbMxCallee( - MLIRContext *context) { - return StringAttr::get(context, "llvm.hivm.LOAD.L1.TO.L0B.MX.2Dv2.v") - .getValue(); -} - -static StringRef buildCopyMatrixCcToGmCallee(MLIRContext *context) { - return StringAttr::get(context, "llvm.hivm.FIX.L0C.TO.OUT.f32.EXT") - .getValue(); -} - -static StringRef buildCopyMatrixCcToCbufCallee(MLIRContext *context) { - return StringAttr::get(context, "llvm.hivm.FIX.L0C.TO.L1.f32.EXT") - .getValue(); -} - -static FailureOr buildCopyMatrixCcToUbCallee(MLIRContext *context, - Type destinationType) { - auto ptrType = dyn_cast(destinationType); - if (!ptrType) - return failure(); - Type dstElem = ptrType.getElementType(); - if (dstElem.isF16()) - return StringAttr::get(context, "llvm.hivm.FIX.L0C.TO.UB.f322f16.EXT") - .getValue(); - if (dstElem.isF32()) - return StringAttr::get(context, "llvm.hivm.FIX.L0C.TO.UB.f32.EXT") - .getValue(); - return failure(); -} - -static FailureOr buildCopyCbufToBtCallee(pto::CopyCbufToBtOp op) { - auto ptrType = dyn_cast(op.getSource().getType()); - if (!ptrType) - return failure(); - Type srcElem = ptrType.getElementType(); - if (srcElem.isF16()) - return StringAttr::get(op.getContext(), "llvm.hivm.MOV.L1.TO.BT.f16") - .getValue(); - if (srcElem.isBF16()) - return StringAttr::get(op.getContext(), "llvm.hivm.MOV.L1.TO.BT.bf16") - .getValue(); - if (srcElem.isF32()) - return StringAttr::get(op.getContext(), "llvm.hivm.MOV.L1.TO.BT.f32") - .getValue(); - if (auto intType = dyn_cast(srcElem); - intType && intType.getWidth() == 32) { - return StringAttr::get(op.getContext(), "llvm.hivm.MOV.L1.TO.BT.s32") - .getValue(); - } - return failure(); -} - -static StringRef buildCopyCbufToFbufCallee(MLIRContext *context) { - return StringAttr::get(context, "llvm.hivm.MOV.L1.TO.FB.v220").getValue(); -} - -static StringRef buildPstiCallee(MLIRContext *context, bool post) { - return StringAttr::get(context, - post ? "llvm.hivm.psti.post.b8" - : "llvm.hivm.psti.b8") - .getValue(); -} - -static StringRef buildPstsCallee(MLIRContext *context, bool post) { - return StringAttr::get(context, - post ? "llvm.hivm.psts.post.b8" - : "llvm.hivm.psts.b8") - .getValue(); -} - -static StringRef buildPldiCallee(MLIRContext *context, bool post) { - return StringAttr::get(context, - post ? "llvm.hivm.pldi.post.b8" - : "llvm.hivm.pldi.b8") - .getValue(); -} - -static StringRef buildPldsCallee(MLIRContext *context, bool post) { - return StringAttr::get(context, - post ? "llvm.hivm.plds.post.b8" - : "llvm.hivm.plds.b8") - .getValue(); -} - -static StringRef buildPnotCallee(MLIRContext *context) { - return StringAttr::get(context, "llvm.hivm.pnot.z").getValue(); -} - -static StringRef buildPselCallee(MLIRContext *context) { - return StringAttr::get(context, "llvm.hivm.psel").getValue(); -} - -static StringRef buildPandCallee(MLIRContext *context) { - return StringAttr::get(context, "llvm.hivm.pand.z").getValue(); -} - -static StringRef buildPorCallee(MLIRContext *context) { - return StringAttr::get(context, "llvm.hivm.por.z").getValue(); -} - -static StringRef buildPxorCallee(MLIRContext *context) { - return StringAttr::get(context, "llvm.hivm.pxor.z").getValue(); -} - -static StringRef buildPpackCallee(MLIRContext *context) { - return StringAttr::get(context, "llvm.hivm.ppack.z").getValue(); -} - -static StringRef buildPunpackCallee(MLIRContext *context) { - return StringAttr::get(context, "llvm.hivm.punpack").getValue(); -} - -template -static StringRef buildPredicatePairReorderCallee(MLIRContext *context); - -template <> -StringRef buildPredicatePairReorderCallee(MLIRContext *context) { - return StringAttr::get(context, "llvm.hivm.pdintlv.b8").getValue(); -} - -template <> -StringRef buildPredicatePairReorderCallee(MLIRContext *context) { - return StringAttr::get(context, "llvm.hivm.pdintlv.b16").getValue(); -} - -template <> -StringRef buildPredicatePairReorderCallee(MLIRContext *context) { - return StringAttr::get(context, "llvm.hivm.pdintlv.b32").getValue(); -} - -template <> -StringRef buildPredicatePairReorderCallee(MLIRContext *context) { - return StringAttr::get(context, "llvm.hivm.pintlv.b8").getValue(); -} - -template <> -StringRef buildPredicatePairReorderCallee(MLIRContext *context) { - return StringAttr::get(context, "llvm.hivm.pintlv.b16").getValue(); -} - -template <> -StringRef buildPredicatePairReorderCallee(MLIRContext *context) { - return StringAttr::get(context, "llvm.hivm.pintlv.b32").getValue(); -} - -static FailureOr buildInterleaveCallee(MLIRContext *context, - Type resultType, - StringRef stem) { - // bf16x2 has no dedicated vintlv/vdintlv intrinsic. It is a 32-bit packed - // pair lowered to i32 at the LLVM ABI, and (de)interleave is a bit-level - // lane shuffle, so the intrinsic serves the type. - if (pto::isPTOBF16x2Type(getElementTypeFromVectorLike(resultType))) { - auto lanes = getElementCountFromVectorLike(resultType); - if (lanes) - return StringAttr::get(context, "llvm.hivm." + stem.str() + ".v" + - std::to_string(*lanes) + "i32") - .getValue(); - } - std::string vec = getCANN900VectorTypeFragment(resultType); - if (vec.empty()) - return failure(); - return StringAttr::get(context, "llvm.hivm." + stem.str() + "." + vec) - .getValue(); -} - -static FailureOr buildUnpackCallee(MLIRContext *context, - Type inputType, - Type resultType, - StringRef stem) { - (void)inputType; - std::string vec = getCANN900VectorTypeFragment(resultType); - if (vec.empty()) - return failure(); - return StringAttr::get(context, "llvm.hivm." + stem.str() + "." + vec) - .getValue(); -} - -static FailureOr buildVpackCallee(MLIRContext *context, Type inputType, - Type resultType) { - (void)resultType; - std::string vec = getCANN900VectorTypeFragment(inputType); - if (vec.empty()) - return failure(); - - return StringAttr::get(context, "llvm.hivm.vpack.x." + vec) - .getValue(); -} - -static FailureOr buildVsqzCallee(MLIRContext *context, - Type resultType) { - return buildCANN900ModeTypedCallee(context, resultType, "vsqz", "x"); -} - -static FailureOr buildVusqzCallee(MLIRContext *context, - Type resultType) { - return buildCANN900ModeTypedCallee(context, resultType, "vusqz", "m"); -} - -static FailureOr buildVmulaCallee(MLIRContext *context, - Type resultType) { - return buildCANN900SignedModeTypedCallee(context, resultType, "vmula", "m"); -} - -static FailureOr buildVmullCallee(MLIRContext *context, - Type resultType) { - return buildLaneTypedCallee(context, resultType, "vmull", ""); -} - -template -static StringRef getPredicateStoreCallee(MLIRContext *context, bool post); - -template <> -StringRef getPredicateStoreCallee(MLIRContext *context, - bool post) { - return buildPstiCallee(context, post); -} - -template <> -StringRef getPredicateStoreCallee(MLIRContext *context, - bool post) { - return buildPstsCallee(context, post); -} - -template -static StringRef getPredicateLoadCallee(MLIRContext *context, bool post); - -template <> -StringRef getPredicateLoadCallee(MLIRContext *context, bool post) { - return buildPldiCallee(context, post); -} - -template <> -StringRef getPredicateLoadCallee(MLIRContext *context, bool post) { - return buildPldsCallee(context, post); -} - -template -static StringRef getPredicateMaskCallee(MLIRContext *context); - -template <> -StringRef getPredicateMaskCallee(MLIRContext *context) { - return buildPnotCallee(context); -} - -template <> -StringRef getPredicateMaskCallee(MLIRContext *context) { - return buildPselCallee(context); -} - -template <> -StringRef getPredicateMaskCallee(MLIRContext *context) { - return buildPandCallee(context); -} - -template <> -StringRef getPredicateMaskCallee(MLIRContext *context) { - return buildPorCallee(context); -} - -template <> -StringRef getPredicateMaskCallee(MLIRContext *context) { - return buildPxorCallee(context); -} - -template -static StringRef getPredicatePackCallee(MLIRContext *context); - -template <> -StringRef getPredicatePackCallee(MLIRContext *context) { - return buildPpackCallee(context); -} - -template <> -StringRef getPredicatePackCallee(MLIRContext *context) { - return buildPunpackCallee(context); -} - -template -static StringRef buildPltCallee(MLIRContext *context); - -template <> -StringRef buildPltCallee(MLIRContext *context) { - return StringAttr::get(context, "llvm.hivm.plt.b8.v300").getValue(); -} - -template <> -StringRef buildPltCallee(MLIRContext *context) { - return StringAttr::get(context, "llvm.hivm.plt.b16.v300").getValue(); -} - -template <> -StringRef buildPltCallee(MLIRContext *context) { - return StringAttr::get(context, "llvm.hivm.plt.b32.v300").getValue(); -} - -template -static StringRef buildPltmCallee(MLIRContext *context); - -template <> -StringRef buildPltmCallee(MLIRContext *context) { - return StringAttr::get(context, "llvm.hivm.pltm.b8.v300").getValue(); -} - -template <> -StringRef buildPltmCallee(MLIRContext *context) { - return StringAttr::get(context, "llvm.hivm.pltm.b16.v300").getValue(); -} - -template <> -StringRef buildPltmCallee(MLIRContext *context) { - return StringAttr::get(context, "llvm.hivm.pltm.b32.v300").getValue(); -} - -template -static StringRef buildPsetCallee(MLIRContext *context); - -template <> -StringRef buildPsetCallee(MLIRContext *context) { - return StringAttr::get(context, "llvm.hivm.pset.b8").getValue(); -} - -template <> -StringRef buildPsetCallee(MLIRContext *context) { - return StringAttr::get(context, "llvm.hivm.pset.b16").getValue(); -} - -template <> -StringRef buildPsetCallee(MLIRContext *context) { - return StringAttr::get(context, "llvm.hivm.pset.b32").getValue(); -} - -template -static StringRef buildPgeCallee(MLIRContext *context); - -template <> -StringRef buildPgeCallee(MLIRContext *context) { - return StringAttr::get(context, "llvm.hivm.pge.b8").getValue(); -} - -template <> -StringRef buildPgeCallee(MLIRContext *context) { - return StringAttr::get(context, "llvm.hivm.pge.b16").getValue(); -} - -template <> -StringRef buildPgeCallee(MLIRContext *context) { - return StringAttr::get(context, "llvm.hivm.pge.b32").getValue(); -} - -static FailureOr buildVldsCallee(MLIRContext *context, Type resultType) { - std::string vec = - getMemoryElementTypeFragment(getElementTypeFromVectorLike(resultType)); - auto lanes = getElementCountFromVectorLike(resultType); - if (vec.empty() || !lanes) - return failure(); - return StringAttr::get(context, "llvm.hivm.vldsx1.v" + std::to_string(*lanes) + - vec) - .getValue(); -} - -static FailureOr buildVldsx2Callee(MLIRContext *context, - Type resultType, bool post) { - std::string vec = - getMemoryElementTypeFragment(getElementTypeFromVectorLike(resultType)); - auto lanes = getElementCountFromVectorLike(resultType); - if (vec.empty() || !lanes) - return failure(); - return StringAttr::get( - context, "llvm.hivm.vldsx2" + - std::string(post ? ".post" : "") + ".v" + - std::to_string(*lanes) + vec) - .getValue(); -} - -static FailureOr -buildBlockStridedMemoryCallee(MLIRContext *context, Type vectorType, - StringRef stem, bool post) { - Type elementType = getElementTypeFromVectorLike(vectorType); - auto lanes = getElementCountFromVectorLike(vectorType); - if (!elementType || !lanes) - return failure(); - - std::string element; - if (auto intType = dyn_cast(elementType)) - element = "i" + std::to_string(intType.getWidth()); - else if (isLowpPayloadElementType(elementType)) - element = "i8"; - else - element = getMemoryElementTypeFragment(elementType); - if (element.empty()) - return failure(); - - return StringAttr::get(context, - "llvm.hivm." + stem.str() + - std::string(post ? ".post" : "") + ".v" + - std::to_string(*lanes) + element) - .getValue(); -} - -static FailureOr buildVsldbCallee(MLIRContext *context, - Type resultType, bool post) { - return buildBlockStridedMemoryCallee(context, resultType, "vsldb", - post); -} - -static FailureOr buildVstsCallee(MLIRContext *context, Type valueType) { - std::string vec = - getMemoryElementTypeFragment(getElementTypeFromVectorLike(valueType)); - auto lanes = getElementCountFromVectorLike(valueType); - if (vec.empty() || !lanes) - return failure(); - return StringAttr::get(context, "llvm.hivm.vstsx1.v" + std::to_string(*lanes) + - vec) - .getValue(); -} - -static FailureOr buildVstsx2Callee(MLIRContext *context, Type valueType) { - Type elementType = getElementTypeFromVectorLike(valueType); - auto lanes = getElementCountFromVectorLike(valueType); - if (!elementType || !lanes) - return failure(); - - std::string element = getMemoryElementTypeFragment(elementType); - if (element.empty()) - return failure(); - - return StringAttr::get(context, "llvm.hivm.vstsx2.v" + - std::to_string(*lanes) + element) - .getValue(); -} - -static FailureOr buildVsstbCallee(MLIRContext *context, - Type valueType, bool post) { - return buildBlockStridedMemoryCallee(context, valueType, "vsstb", post); -} - -static Type getVgather2SourceElementType(Type sourceType) { - if (auto ptrType = dyn_cast(sourceType)) - return ptrType.getElementType(); - if (auto memrefType = dyn_cast(sourceType)) - return memrefType.getElementType(); - return {}; -} - -static FailureOr buildVgather2Callee(MLIRContext *context, - Type sourceType, - Type resultType) { - Type sourceElemType = getVgather2SourceElementType(sourceType); - Type resultElemType = getElementTypeFromVectorLike(resultType); - auto lanes = getElementCountFromVectorLike(resultType); - if (!sourceElemType || !resultElemType || !lanes) - return failure(); - - std::string vec; - int64_t intrinsicLanes = *lanes; - if (pto::getPTOStorageElemBitWidth(sourceElemType) == 8) { - vec = getElementTypeFragment(sourceElemType); - intrinsicLanes *= 2; - } else { - vec = getElementTypeFragment(resultElemType); - } - if (vec.empty()) - return failure(); - - return StringAttr::get(context, "llvm.hivm.vgather2.v300.v" + - std::to_string(intrinsicLanes) + vec) - .getValue(); -} - -static std::optional getFixedVectorBitWidth(Type type) { - auto vectorType = dyn_cast(type); - if (!vectorType || vectorType.getRank() != 1 || vectorType.isScalable()) - return std::nullopt; - int64_t lanes = vectorType.getDimSize(0); - if (lanes <= 0) - return std::nullopt; - auto elementType = dyn_cast(vectorType.getElementType()); - if (!elementType) - return std::nullopt; - return static_cast(lanes) * elementType.getWidth(); -} - -static FailureOr getVgather2OffsetsCarrierType(PatternRewriter &rewriter, - Type sourceType, - Type resultType, - Type offsetsType) { - Type sourceElemType = getVgather2SourceElementType(sourceType); - Type elementType = getElementTypeFromVectorLike(resultType); - auto lanes = getElementCountFromVectorLike(resultType); - if (!sourceElemType || !elementType || !lanes || *lanes <= 0) - return failure(); - - Type carrierType = offsetsType; - if (pto::getPTOStorageElemBitWidth(elementType) == 16) { - if (*lanes % 2 != 0) - return failure(); - carrierType = VectorType::get({*lanes / 2}, rewriter.getI32Type()); - } - - std::optional offsetsBits = getFixedVectorBitWidth(offsetsType); - std::optional carrierBits = getFixedVectorBitWidth(carrierType); - if (!offsetsBits || !carrierBits || *offsetsBits != *carrierBits) - return failure(); - return carrierType; -} - -static FailureOr buildVgather2BcCallee(MLIRContext *context, - Type resultType) { - return buildLaneTypedCallee(context, resultType, "vgather2.bc", ""); -} - -static FailureOr buildVgatherbCallee(MLIRContext *context, - Type resultType) { - return buildLaneTypedCallee(context, resultType, "vgatherb.v310", ""); -} - -static FailureOr buildVscatterCallee(MLIRContext *context, - Type valueType) { - return buildLaneTypedCallee(context, valueType, "vscatter", ".v300"); -} - -static FailureOr getVscatterOffsetsCarrierType(Type offsetsType) { - return offsetsType; -} - -static FailureOr buildVaxpyCallee(MLIRContext *context, - Type resultType) { - return buildCANN900ModeTypedCallee(context, resultType, "vaxpy", "m"); -} - -static FailureOr buildVmulscvtCallee(MLIRContext *context, - Type inputType, - Type resultType) { - auto inputElemType = getElementTypeFromVectorLike(inputType); - auto resultElemType = getElementTypeFromVectorLike(resultType); - auto inputLanes = getElementCountFromVectorLike(inputType); - auto resultLanes = getElementCountFromVectorLike(resultType); - if (!inputElemType || !resultElemType || !inputLanes || !resultLanes) - return failure(); - if (!inputElemType.isF32() || !resultElemType.isF16() || *inputLanes != 64 || - *resultLanes != 128) - return failure(); - return StringAttr::get(context, "llvm.hivm.vmulscvt.v128f16").getValue(); -} - -static FailureOr buildVciCallee(MLIRContext *context, Type resultType) { - std::string vec = getCANN900VectorTypeFragment(resultType); - if (vec.empty()) - return failure(); - return StringAttr::get(context, "llvm.hivm.vci." + vec) - .getValue(); -} - -static FailureOr buildVtrcCallee(MLIRContext *context, Type resultType) { - std::string vec = - getElementTypeFragment(getElementTypeFromVectorLike(resultType)); - auto lanes = getElementCountFromVectorLike(resultType); - if (vec.empty() || !lanes) - return failure(); - return StringAttr::get(context, "llvm.hivm.vtrc." + vec + ".x").getValue(); -} - -static FailureOr buildVexpdifCallee(MLIRContext *context, - Type inputType, - Type resultType) { - Type inputElem = getElementTypeFromVectorLike(inputType); - Type resultElem = getElementTypeFromVectorLike(resultType); - auto srcLanes = getElementCountFromVectorLike(inputType); - if (!srcLanes) - return failure(); - if (inputElem.isF16() && resultElem.isF32() && *srcLanes == 128) - return StringAttr::get(context, - "llvm.hivm.vexpdif.interleave.v128f16") - .getValue(); - if (inputElem.isF32() && resultElem.isF32() && *srcLanes == 64) - return StringAttr::get(context, "llvm.hivm.vexpdif.v64f32").getValue(); - return failure(); -} - -static FailureOr buildVbitsortCallee(MLIRContext *context, - pto::VbitsortOp op) { - Type sourceElemType = cast(op.getSource().getType()).getElementType(); - if (sourceElemType.isF16()) - return StringAttr::get(context, "llvm.hivm.VBS32.V300.f16").getValue(); - if (sourceElemType.isF32()) - return StringAttr::get(context, "llvm.hivm.VBS32.V300.f32").getValue(); - return failure(); -} - -static FailureOr buildVmrgsort4Callee(MLIRContext *context, - pto::Vmrgsort4Op op) { - Type elemType = - cast(op.getDestination().getType()).getElementType(); - if (elemType.isF16()) - return StringAttr::get(context, "llvm.hivm.VMRGSORT.f16.V300").getValue(); - if (elemType.isF32()) - return StringAttr::get(context, "llvm.hivm.VMRGSORT.f32.V300").getValue(); - return failure(); -} - -static FailureOr packVmrgsort4SourceAddr(Operation *anchor, Value source0, - Value source1, Value source2, - Value source3, Type elemType) { - OpBuilder builder(anchor); - builder.setInsertionPoint(anchor); - Location loc = anchor->getLoc(); - unsigned addrShift = 0; - if (elemType.isF16()) - addrShift = 3; - else if (elemType.isF32()) - addrShift = 3; - else - return failure(); - - auto packOne = [&](Value source, uint64_t laneShift) -> FailureOr { - FailureOr ubPtr = reinterpretPointerToAddrSpace(anchor, source, 6); - if (failed(ubPtr)) - return failure(); - Value asInt = - builder.create(loc, builder.getI64Type(), *ubPtr); - Value shifted = builder.create( - loc, asInt, getI64Constant(builder, loc, addrShift)); - Value masked = builder.create( - loc, shifted, getI64Constant(builder, loc, 0xFFFFULL)); - if (laneShift == 0) - return masked; - return builder - .create(loc, masked, - getI64Constant(builder, loc, laneShift)) - .getResult(); - }; - - FailureOr low0 = packOne(source0, 0); - FailureOr low1 = packOne(source1, 16); - FailureOr low2 = packOne(source2, 32); - FailureOr low3 = packOne(source3, 48); - if (failed(low0) || failed(low1) || failed(low2) || failed(low3)) - return failure(); - - Value packed01 = builder.create(loc, *low0, *low1); - Value packed23 = builder.create(loc, *low2, *low3); - Value packed = builder.create(loc, packed01, packed23); - Type ubPtrTy = LLVM::LLVMPointerType::get(anchor->getContext(), 6); - return builder.create(loc, ubPtrTy, packed).getResult(); -} - -static FailureOr buildVcvtContract(pto::VcvtOp op) { - Type inputElemType = getElementTypeFromVectorLike(op.getInput().getType()); - Type resultElemType = getElementTypeFromVectorLike(op.getResult().getType()); - if (!inputElemType || !resultElemType) - return failure(); - auto contract = lookupVcvtContract(classifyVcvtElemType(inputElemType), - classifyVcvtElemType(resultElemType)); - if (!contract) - return failure(); - return *contract; -} - -static bool needsV300CtrlModeForVPTOFunc(func::FuncOp funcOp) { - if (!pto::isPTOEntryFunction(funcOp) || funcOp.getBlocks().empty()) - return false; - - bool needsCtrlSetup = false; - funcOp.walk([&](pto::VcvtOp vcvtOp) { - FailureOr contract = buildVcvtContract(vcvtOp); - if (succeeded(contract) && (*contract).requiresSat) { - needsCtrlSetup = true; - return WalkResult::interrupt(); - } - return WalkResult::advance(); - }); - return needsCtrlSetup; -} - -template -static StringRef buildSetLoopCallee(MLIRContext *context); - -template -static StringRef buildUnaryConfigCallee(MLIRContext *context); - -template -static StringRef buildNullaryConfigCallee(MLIRContext *context); - -template <> -StringRef buildSetLoopCallee(MLIRContext *context) { - return StringAttr::get(context, "llvm.hivm.SET.LOOP2.STRIDE.OUTTOUB") - .getValue(); -} - -template <> -StringRef buildSetLoopCallee(MLIRContext *context) { - return StringAttr::get(context, "llvm.hivm.SET.LOOP1.STRIDE.OUTTOUB") - .getValue(); -} - -template <> -StringRef buildSetLoopCallee(MLIRContext *context) { - return StringAttr::get(context, "llvm.hivm.SET.LOOP.SIZE.OUTTOUB") - .getValue(); -} - -template <> -StringRef buildSetLoopCallee(MLIRContext *context) { - return StringAttr::get(context, "llvm.hivm.SET.LOOP2.STRIDE.UBTOOUT") - .getValue(); -} - -template <> -StringRef buildSetLoopCallee(MLIRContext *context) { - return StringAttr::get(context, "llvm.hivm.SET.LOOP1.STRIDE.UBTOOUT") - .getValue(); -} - -template <> -StringRef buildSetLoopCallee(MLIRContext *context) { - return StringAttr::get(context, "llvm.hivm.SET.LOOP.SIZE.UBTOOUT") - .getValue(); -} - -template <> -StringRef buildSetLoopCallee(MLIRContext *context) { - return StringAttr::get(context, "llvm.hivm.SET.LOOP3.PARA").getValue(); -} - -template <> -StringRef buildSetLoopCallee(MLIRContext *context) { - return StringAttr::get(context, "llvm.hivm.SET.CHANNEL.PARA").getValue(); -} - -template <> -StringRef buildUnaryConfigCallee(MLIRContext *context) { - return StringAttr::get(context, "llvm.hivm.SET.MOV.PAD.VAL").getValue(); -} - -template <> -StringRef buildUnaryConfigCallee(MLIRContext *context) { - return StringAttr::get(context, "llvm.hivm.SET.QUANT.PRE.v300").getValue(); -} - -template <> -StringRef buildUnaryConfigCallee(MLIRContext *context) { - return StringAttr::get(context, "llvm.hivm.SET.RELU.ALPHA").getValue(); -} - -template <> -StringRef buildUnaryConfigCallee(MLIRContext *context) { - return StringAttr::get(context, "llvm.hivm.SET.FIX.CLIP.RELU").getValue(); -} - -template <> -StringRef buildUnaryConfigCallee( - MLIRContext *context) { - return StringAttr::get(context, "llvm.hivm.SET.LOOP2.STRIDE.OUTTOL1") - .getValue(); -} - -template <> -StringRef buildUnaryConfigCallee( - MLIRContext *context) { - return StringAttr::get(context, "llvm.hivm.SET.LOOP1.STRIDE.OUTTOL1") - .getValue(); -} - -template <> -StringRef buildUnaryConfigCallee( - MLIRContext *context) { - return StringAttr::get(context, "llvm.hivm.SET.LOOP.SIZE.OUTTOL1") - .getValue(); -} - -template <> -StringRef buildUnaryConfigCallee(MLIRContext *context) { - return StringAttr::get(context, "llvm.hivm.SET.MTE2.NZ.PARA").getValue(); -} - -template <> -StringRef buildUnaryConfigCallee( - MLIRContext *context) { - return StringAttr::get(context, "llvm.hivm.SET.PAD.VAL.OUTTOL1") - .getValue(); -} - -template <> -StringRef buildUnaryConfigCallee(MLIRContext *context) { - return StringAttr::get(context, "llvm.hivm.SET.FPC").getValue(); -} - -template <> -StringRef buildUnaryConfigCallee( - MLIRContext *context) { - return StringAttr::get(context, "llvm.hivm.SET.ST.ATOMIC.CFG").getValue(); -} - -template <> -StringRef buildNullaryConfigCallee(MLIRContext *context) { - return StringAttr::get(context, "llvm.hivm.SET.ATOMIC.S32").getValue(); -} - -template <> -StringRef buildNullaryConfigCallee(MLIRContext *context) { - return StringAttr::get(context, "llvm.hivm.SET.ATOMIC.S8").getValue(); -} - -static FailureOr encodeMovPadValue(Location loc, Value value, - ConversionPatternRewriter &rewriter) { - Type type = value.getType(); - Value payload = value; - unsigned bitWidth = 0; - - if (auto intType = dyn_cast(type)) { - bitWidth = intType.getWidth(); - } else if (auto floatType = dyn_cast(type)) { - bitWidth = floatType.getWidth(); - auto intType = rewriter.getIntegerType(bitWidth); - payload = rewriter.create(loc, intType, value); - } else { - return failure(); - } - - if (bitWidth != 8 && bitWidth != 16 && bitWidth != 32) - return failure(); - - return rewriter.create(loc, rewriter.getI64Type(), payload) - .getResult(); -} - -template -static StringRef buildSyncCallee(MLIRContext *context); - -template <> -StringRef buildSyncCallee(MLIRContext *context) { - return StringAttr::get(context, "llvm.hivm.SET.FLAG.IMM").getValue(); -} - -template <> -StringRef buildSyncCallee(MLIRContext *context) { - return StringAttr::get(context, "llvm.hivm.WAIT.FLAG.IMM").getValue(); -} - -template <> -StringRef buildSyncCallee(MLIRContext *context) { - return StringAttr::get(context, "llvm.hivm.SET.FLAG.REG").getValue(); -} - -template <> -StringRef buildSyncCallee(MLIRContext *context) { - return StringAttr::get(context, "llvm.hivm.WAIT.FLAG.REG").getValue(); -} - -template <> -StringRef buildSyncCallee(MLIRContext *context) { - return StringAttr::get(context, "llvm.hivm.BARRIER").getValue(); -} - -template <> -StringRef buildSyncCallee(MLIRContext *context) { - return StringAttr::get(context, "llvm.hivm.SET.CROSS.CORE").getValue(); -} - -template <> -StringRef buildSyncCallee(MLIRContext *context) { - return StringAttr::get(context, "llvm.hivm.WAIT.FLAG.DEV.REG").getValue(); -} - -template <> -StringRef buildSyncCallee(MLIRContext *context) { - return StringAttr::get(context, "llvm.hivm.SET.INTRA.BLOCK.mode").getValue(); -} - -template <> -StringRef buildSyncCallee(MLIRContext *context) { - return StringAttr::get(context, "llvm.hivm.WAIT.INTRA.BLOCK.mode").getValue(); -} - -static StringRef buildMemBarCallee(MemBarKind kind, MLIRContext *context) { - switch (kind) { - case MemBarKind::VV_ALL: - return StringAttr::get(context, "llvm.hivm.mem.bar.vv.all").getValue(); - case MemBarKind::VST_VLD: - return StringAttr::get(context, "llvm.hivm.mem.bar.vst.vld").getValue(); - case MemBarKind::VLD_VST: - return StringAttr::get(context, "llvm.hivm.mem.bar.vld.vst").getValue(); - case MemBarKind::VST_VST: - return StringAttr::get(context, "llvm.hivm.mem.bar.vst.vst").getValue(); - case MemBarKind::VS_ALL: - return StringAttr::get(context, "llvm.hivm.mem.bar.vs.all").getValue(); - case MemBarKind::VST_LD: - return StringAttr::get(context, "llvm.hivm.mem.bar.vst.ld").getValue(); - case MemBarKind::VLD_ST: - return StringAttr::get(context, "llvm.hivm.mem.bar.vld.st").getValue(); - case MemBarKind::VST_ST: - return StringAttr::get(context, "llvm.hivm.mem.bar.vst.st").getValue(); - case MemBarKind::SV_ALL: - return StringAttr::get(context, "llvm.hivm.mem.bar.sv.all").getValue(); - case MemBarKind::ST_VLD: - return StringAttr::get(context, "llvm.hivm.mem.bar.st.vld").getValue(); - case MemBarKind::LD_VST: - return StringAttr::get(context, "llvm.hivm.mem.bar.ld.vst").getValue(); - case MemBarKind::ST_VST: - return StringAttr::get(context, "llvm.hivm.mem.bar.st.vst").getValue(); - case MemBarKind::SS_ALL: - return StringAttr::get(context, "llvm.hivm.mem.bar.ss.all").getValue(); - case MemBarKind::ST_LD: - return StringAttr::get(context, "llvm.hivm.mem.bar.st.ld").getValue(); - case MemBarKind::LD_ST: - return StringAttr::get(context, "llvm.hivm.mem.bar.ld.st").getValue(); - case MemBarKind::ST_ST: - return StringAttr::get(context, "llvm.hivm.mem.bar.st.st").getValue(); - } - llvm_unreachable("unexpected membar kind"); -} - -static uint64_t getDsbMemImmediate(DsbMem kind) { - return static_cast(kind); -} - -static uint64_t getDcciCacheLineImmediate(DcciCacheLine kind) { - return static_cast(kind); -} - -static uint64_t getDcciDstImmediate(DcciDst kind) { - return static_cast(kind); -} - -static StringRef buildDcciCallee(unsigned addressSpace, bool hasDst, - MLIRContext *context) { - if (addressSpace == static_cast(pto::AddressSpace::GM)) { - return StringAttr::get(context, hasDst ? "llvm.hivm.DCCI.DST" - : "llvm.hivm.DCCI") - .getValue(); - } - if (addressSpace == static_cast(pto::AddressSpace::VEC)) { - return StringAttr::get(context, hasDst ? "llvm.hivm.DCCI.DST.UB" - : "llvm.hivm.DCCI.UB") - .getValue(); - } - llvm_unreachable("unexpected dcci address space"); -} - -template <> -StringRef buildSyncCallee(MLIRContext *context) { - return StringAttr::get(context, "llvm.hivm.GET.BUFI.mode").getValue(); -} - -template <> -StringRef buildSyncCallee(MLIRContext *context) { - return StringAttr::get(context, "llvm.hivm.RLS.BUFI.mode").getValue(); -} - -static StringRef buildBufDynSyncCallee(MLIRContext *context, bool isGetBuf) { - return StringAttr::get(context, - isGetBuf ? "llvm.hivm.GET.BUF.mode" - : "llvm.hivm.RLS.BUF.mode") - .getValue(); -} - -template -static StringRef buildRuntimeQueryCallee(MLIRContext *context); - -template <> -StringRef buildRuntimeQueryCallee(MLIRContext *context) { - return StringAttr::get(context, "llvm.hivm.GET.BLOCK.IDX").getValue(); -} - -template <> -StringRef buildRuntimeQueryCallee(MLIRContext *context) { - return StringAttr::get(context, "llvm.hivm.GET.SUBBLOCKID").getValue(); -} - -template <> -StringRef buildRuntimeQueryCallee(MLIRContext *context) { - return StringAttr::get(context, "llvm.hivm.GET.BLOCK.NUM").getValue(); -} - -template <> -StringRef buildRuntimeQueryCallee(MLIRContext *context) { - return StringAttr::get(context, "llvm.hivm.GET.SUBBLOCKDIM").getValue(); -} - -template -static StringRef buildSimtBlockQueryCallee(MLIRContext *context); - -template <> -StringRef -buildSimtBlockQueryCallee(MLIRContext *context) { - return StringAttr::get(context, "llvm.hivm.tpe.get.BLOCK.IDX").getValue(); -} - -template <> -StringRef -buildSimtBlockQueryCallee(MLIRContext *context) { - return StringAttr::get(context, "llvm.hivm.tpe.get.BLOCK.NUM").getValue(); -} - -static LogicalResult -materializeDecls(ModuleOp module, ArrayRef plannedDecls, - llvm::raw_ostream &diagOS) { - OpBuilder builder(module.getBodyRegion()); - builder.setInsertionPointToStart(&module.getBodyRegion().front()); - for (const PlannedDecl &decl : plannedDecls) { - if (func::FuncOp existing = module.lookupSymbol(decl.name)) { - if (existing.getFunctionType() != decl.type) { - diagOS << "VPTO LLVM emission failed: conflicting declaration for " - << decl.name << "\n"; - return failure(); - } - continue; - } - auto func = - builder.create(module.getLoc(), decl.name, decl.type); - func.setPrivate(); - } - return success(); -} - -template -class LowerUnaryMaskedOpPattern final : public OpConversionPattern { -public: - explicit LowerUnaryMaskedOpPattern(TypeConverter &typeConverter, - MLIRContext *context, - LoweringState &state) - : OpConversionPattern(typeConverter, context), state(state) {} - - LogicalResult - matchAndRewrite(UnaryOp op, typename UnaryOp::Adaptor adaptor, - ConversionPatternRewriter &rewriter) const override { - FailureOr calleeName = - buildUnaryMaskedCallee(op.getContext(), - op.getResult().getType()); - if (failed(calleeName)) - return rewriter.notifyMatchFailure(op, "unsupported unary VPTO signature"); - - Type resultType = - this->getTypeConverter()->convertType(op.getResult().getType()); - if (!resultType) - return rewriter.notifyMatchFailure(op, - "failed to convert unary result type"); - - Value input = adaptor.getOperands()[0]; - Value mask = adaptor.getOperands()[1]; - Type expectedMaskType = - this->getTypeConverter()->convertType(op->getOperand(1).getType()); - if (!input || !mask || input.getType() != resultType || - mask.getType() != expectedMaskType) { - return rewriter.notifyMatchFailure( - op, "unexpected converted unary VPTO operand types"); - } - - auto call = rewriter.create(op.getLoc(), *calleeName, - TypeRange{resultType}, - ValueRange{input, mask}); - state.plannedDecls.push_back( - PlannedDecl{calleeName->str(), call.getCalleeType()}); - rewriter.replaceOp(op, call.getResults()); - return success(); - } - -private: - LoweringState &state; -}; - -class LowerVsqzOpPattern final : public OpConversionPattern { -public: - explicit LowerVsqzOpPattern(TypeConverter &typeConverter, MLIRContext *context, - LoweringState &state) - : OpConversionPattern(typeConverter, context), state(state) {} - - LogicalResult - matchAndRewrite(pto::VsqzOp op, pto::VsqzOp::Adaptor adaptor, - ConversionPatternRewriter &rewriter) const override { - FailureOr calleeName = - buildVsqzCallee(op.getContext(), op.getResult().getType()); - if (failed(calleeName)) - return rewriter.notifyMatchFailure(op, "unsupported vsqz VPTO signature"); - - Type resultType = - this->getTypeConverter()->convertType(op.getResult().getType()); - Type maskType = this->getTypeConverter()->convertType(op.getMask().getType()); - if (!resultType || !maskType) - return rewriter.notifyMatchFailure(op, "failed to convert vsqz types"); - - Value input = adaptor.getInput(); - Value mask = adaptor.getMask(); - if (!input || !mask || input.getType() != resultType || - mask.getType() != maskType) { - return rewriter.notifyMatchFailure(op, - "unexpected converted vsqz operand types"); - } - - Value storeHint = - getI32Constant(rewriter, op.getLoc(), determineVsqzStoreHint(op)); - auto funcType = rewriter.getFunctionType( - TypeRange{resultType, maskType, storeHint.getType()}, TypeRange{resultType}); - auto call = rewriter.create( - op.getLoc(), *calleeName, TypeRange{resultType}, - ValueRange{input, mask, storeHint}); - state.plannedDecls.push_back(PlannedDecl{calleeName->str(), funcType}); - rewriter.replaceOp(op, call.getResults()); - return success(); - } - -private: - LoweringState &state; -}; - -class LowerVusqzOpPattern final : public OpConversionPattern { -public: - explicit LowerVusqzOpPattern(TypeConverter &typeConverter, - MLIRContext *context, LoweringState &state) - : OpConversionPattern(typeConverter, context), state(state) {} - - LogicalResult - matchAndRewrite(pto::VusqzOp op, pto::VusqzOp::Adaptor adaptor, - ConversionPatternRewriter &rewriter) const override { - FailureOr calleeName = - buildVusqzCallee(op.getContext(), op.getResult().getType()); - if (failed(calleeName)) - return rewriter.notifyMatchFailure(op, "unsupported vusqz VPTO signature"); - - Type resultType = - this->getTypeConverter()->convertType(op.getResult().getType()); - Type maskType = this->getTypeConverter()->convertType(op.getMask().getType()); - if (!resultType || !maskType) - return rewriter.notifyMatchFailure(op, "failed to convert vusqz types"); - - Value src = adaptor.getSrc(); - Value mask = adaptor.getMask(); - if (!src || !mask || src.getType() != resultType || mask.getType() != maskType) { - return rewriter.notifyMatchFailure(op, - "unexpected converted vusqz operand types"); - } - - auto funcType = - rewriter.getFunctionType(TypeRange{resultType, maskType}, TypeRange{resultType}); - auto call = rewriter.create( - op.getLoc(), *calleeName, TypeRange{resultType}, ValueRange{src, mask}); - state.plannedDecls.push_back(PlannedDecl{calleeName->str(), funcType}); - rewriter.replaceOp(op, call.getResults()); - return success(); - } - -private: - LoweringState &state; -}; - -class LowerVmulaOpPattern final : public OpConversionPattern { -public: - explicit LowerVmulaOpPattern(TypeConverter &typeConverter, MLIRContext *context, - LoweringState &state) - : OpConversionPattern(typeConverter, context), state(state) {} - - LogicalResult - matchAndRewrite(pto::VmulaOp op, pto::VmulaOp::Adaptor adaptor, - ConversionPatternRewriter &rewriter) const override { - FailureOr calleeName = - buildVmulaCallee(op.getContext(), op.getResult().getType()); - if (failed(calleeName)) - return rewriter.notifyMatchFailure(op, "unsupported vmula VPTO signature"); - - Type resultType = - this->getTypeConverter()->convertType(op.getResult().getType()); - Type maskType = this->getTypeConverter()->convertType(op.getMask().getType()); - if (!resultType || !maskType) - return rewriter.notifyMatchFailure(op, "failed to convert vmula types"); - - Value acc = adaptor.getAcc(); - Value lhs = adaptor.getLhs(); - Value rhs = adaptor.getRhs(); - Value mask = adaptor.getMask(); - if (!acc || !lhs || !rhs || !mask || acc.getType() != resultType || - lhs.getType() != resultType || rhs.getType() != resultType || - mask.getType() != maskType) { - return rewriter.notifyMatchFailure(op, - "unexpected converted vmula operand types"); - } - - auto funcType = rewriter.getFunctionType( - TypeRange{resultType, resultType, resultType, maskType}, - TypeRange{resultType}); - auto call = rewriter.create( - op.getLoc(), *calleeName, TypeRange{resultType}, - ValueRange{acc, lhs, rhs, mask}); - state.plannedDecls.push_back(PlannedDecl{calleeName->str(), funcType}); - rewriter.replaceOp(op, call.getResults()); - return success(); - } - -private: - LoweringState &state; -}; - -class LowerVmullOpPattern final : public OpConversionPattern { -public: - explicit LowerVmullOpPattern(TypeConverter &typeConverter, MLIRContext *context, - LoweringState &state) - : OpConversionPattern(typeConverter, context), state(state) {} - - LogicalResult - matchAndRewrite(pto::VmullOp op, pto::VmullOp::Adaptor adaptor, - ConversionPatternRewriter &rewriter) const override { - FailureOr calleeName = - buildVmullCallee(op.getContext(), op.getLow().getType()); - if (failed(calleeName)) - return rewriter.notifyMatchFailure(op, "unsupported vmull VPTO signature"); - - Type inputType = this->getTypeConverter()->convertType(op.getLhs().getType()); - Type maskType = this->getTypeConverter()->convertType(op.getMask().getType()); - SmallVector resultTypes; - if (!inputType || !maskType || - failed(this->getTypeConverter()->convertTypes(op->getResultTypes(), resultTypes))) { - return rewriter.notifyMatchFailure(op, "failed to convert vmull types"); - } - if (resultTypes.size() != 2 || resultTypes[0] != resultTypes[1]) - return rewriter.notifyMatchFailure(op, "unexpected converted vmull results"); - - Value lhs = adaptor.getLhs(); - Value rhs = adaptor.getRhs(); - Value mask = adaptor.getMask(); - if (!lhs || !rhs || !mask || lhs.getType() != inputType || - rhs.getType() != inputType || mask.getType() != maskType) { - return rewriter.notifyMatchFailure(op, - "unexpected converted vmull operand types"); - } - - auto funcType = rewriter.getFunctionType(TypeRange{inputType, inputType, maskType}, - resultTypes); - auto call = rewriter.create(op.getLoc(), *calleeName, resultTypes, - ValueRange{lhs, rhs, mask}); - state.plannedDecls.push_back(PlannedDecl{calleeName->str(), funcType}); - rewriter.replaceOp(op, call.getResults()); - return success(); - } - -private: - LoweringState &state; -}; - -template -class LowerBinaryMaskedOpPattern final : public OpConversionPattern { -public: - explicit LowerBinaryMaskedOpPattern(TypeConverter &typeConverter, - MLIRContext *context, - LoweringState &state) - : OpConversionPattern(typeConverter, context), state(state) {} - - LogicalResult - matchAndRewrite(BinaryOp op, typename BinaryOp::Adaptor adaptor, - ConversionPatternRewriter &rewriter) const override { - StringRef stem = getBinaryMaskedStem(); - Type resultType = - this->getTypeConverter()->convertType(op.getResult().getType()); - if (!resultType) - return rewriter.notifyMatchFailure(op, - "failed to convert binary result type"); - - Value lhs = adaptor.getOperands()[0]; - Value rhs = adaptor.getOperands()[1]; - Value mask = adaptor.getOperands()[2]; - Type expectedMaskType = - this->getTypeConverter()->convertType(op->getOperand(2).getType()); - if (!lhs || !rhs || !mask || lhs.getType() != resultType || - rhs.getType() != resultType || mask.getType() != expectedMaskType) { - return rewriter.notifyMatchFailure( - op, "unexpected converted binary VPTO operand types"); - } - - Type callResultType = resultType; - Value callLhs = lhs; - Value callRhs = rhs; - FailureOr calleeName = - usesSignedBinaryCANN900Callee() - ? buildCANN900SignedModeTypedCallee( - op.getContext(), op.getResult().getType(), stem, "x") - : buildCANN900ModeTypedCallee(op.getContext(), - op.getResult().getType(), stem, "x"); - - if constexpr (std::is_same_v || - std::is_same_v || - std::is_same_v) { - Type elementType = getElementTypeFromVectorLike(op.getResult().getType()); - if (elementType && pto::isPTOLowPrecisionType(elementType)) { - calleeName = buildDirectLowpVLogicCallee( - op.getContext(), op.getResult().getType(), stem, "x"); - if (failed(calleeName)) { - Type carrierType = getLowpPayloadCarrierType( - op.getResult().getType(), rewriter.getContext()); - if (!carrierType) - return rewriter.notifyMatchFailure( - op, "unsupported low-precision binary payload ABI"); - callResultType = carrierType; - callLhs = castToPayloadABI(op.getLoc(), lhs, - op.getResult().getType(), rewriter); - callRhs = castToPayloadABI(op.getLoc(), rhs, - op.getResult().getType(), rewriter); - calleeName = buildLowpPayloadVLogicCallee( - op.getContext(), op.getResult().getType(), stem, "x"); - } - } - } - if (failed(calleeName)) - return rewriter.notifyMatchFailure(op, "unsupported binary VPTO signature"); - - auto call = rewriter.create(op.getLoc(), *calleeName, - TypeRange{callResultType}, - ValueRange{callLhs, callRhs, mask}); - state.plannedDecls.push_back( - PlannedDecl{calleeName->str(), call.getCalleeType()}); - Value result = castFromPayloadABI(op.getLoc(), call.getResult(0), - op.getResult().getType(), resultType, - rewriter); - rewriter.replaceOp(op, result); - return success(); - } - -private: - LoweringState &state; -}; - -template -class LowerTernaryMaskedOpPattern final - : public OpConversionPattern { -public: - explicit LowerTernaryMaskedOpPattern(TypeConverter &typeConverter, - MLIRContext *context, - LoweringState &state) - : OpConversionPattern(typeConverter, context), state(state) {} - - LogicalResult - matchAndRewrite(TernaryOp op, typename TernaryOp::Adaptor adaptor, - ConversionPatternRewriter &rewriter) const override { - StringRef stem = getTernaryMaskedStem(); - FailureOr calleeName = - usesSignedTernaryCANN900Callee() - ? buildCANN900SignedModeTypedCallee( - op.getContext(), op.getResult().getType(), stem, "m") - : buildCANN900ModeTypedCallee(op.getContext(), - op.getResult().getType(), stem, "m"); - if (failed(calleeName)) - return rewriter.notifyMatchFailure(op, - "unsupported ternary VPTO signature"); - - Type resultType = - this->getTypeConverter()->convertType(op.getResult().getType()); - Type expectedMaskType = - this->getTypeConverter()->convertType(op.getMask().getType()); - if (!resultType || !expectedMaskType) - return rewriter.notifyMatchFailure( - op, "failed to convert ternary VPTO types"); - - Value acc = adaptor.getAcc(); - Value lhs = adaptor.getLhs(); - Value rhs = adaptor.getRhs(); - Value mask = adaptor.getMask(); - if (!acc || !lhs || !rhs || !mask || acc.getType() != resultType || - lhs.getType() != resultType || rhs.getType() != resultType || - mask.getType() != expectedMaskType) { - return rewriter.notifyMatchFailure( - op, "unexpected converted ternary VPTO operand types"); - } - - auto call = rewriter.create( - op.getLoc(), *calleeName, TypeRange{resultType}, - ValueRange{acc, lhs, rhs, mask}); - state.plannedDecls.push_back( - PlannedDecl{calleeName->str(), call.getCalleeType()}); - rewriter.replaceOp(op, call.getResults()); - return success(); - } - -private: - LoweringState &state; -}; - -template -class LowerCarryBinaryOpPattern final : public OpConversionPattern { -public: - explicit LowerCarryBinaryOpPattern(TypeConverter &typeConverter, - MLIRContext *context, LoweringState &state) - : OpConversionPattern(typeConverter, context), state(state) {} - - LogicalResult - matchAndRewrite(CarryOp op, typename CarryOp::Adaptor adaptor, - ConversionPatternRewriter &rewriter) const override { - StringRef stem = getCarryBinaryStem(); - FailureOr calleeName = - buildCarryBinaryCallee(op.getContext(), op.getResult().getType(), stem); - if (failed(calleeName)) - return rewriter.notifyMatchFailure(op, "unsupported carry VPTO signature"); - - Type resultType = - this->getTypeConverter()->convertType(op.getResult().getType()); - Type carryType = - this->getTypeConverter()->convertType(op->getResult(1).getType()); - if (!resultType || !carryType) - return rewriter.notifyMatchFailure(op, - "failed to convert carry result types"); - - SmallVector callArgs; - callArgs.append(adaptor.getOperands().begin(), adaptor.getOperands().end()); - const size_t expectedArgCount = hasCarryInput() ? 4 : 3; - if (callArgs.size() != expectedArgCount || callArgs[0].getType() != resultType || - callArgs[1].getType() != resultType || callArgs.back().getType() != carryType) - return rewriter.notifyMatchFailure(op, - "unexpected converted carry operand types"); - if constexpr (hasCarryInput()) { - if (callArgs[2].getType() != carryType) - return rewriter.notifyMatchFailure( - op, "unexpected converted carry input operand type"); - } - - auto call = rewriter.create( - op.getLoc(), *calleeName, TypeRange{resultType, carryType}, callArgs); - state.plannedDecls.push_back( - PlannedDecl{calleeName->str(), call.getCalleeType()}); - rewriter.replaceOp(op, call.getResults()); - return success(); - } - -private: - LoweringState &state; -}; - -template -class LowerCopyOpPattern final : public OpConversionPattern { -public: - explicit LowerCopyOpPattern(TypeConverter &typeConverter, MLIRContext *context, - LoweringState &state) - : OpConversionPattern(typeConverter, context), state(state) {} - - LogicalResult - matchAndRewrite(CopyOp op, typename CopyOp::Adaptor adaptor, - ConversionPatternRewriter &rewriter) const override { - FailureOr calleeName = failure(); - if constexpr (std::is_same_v) - calleeName = buildCopyGmToUbCallee(op.getContext(), op.getSource().getType()); - else - calleeName = buildCopyUbToGmCallee(op.getContext()); - if (failed(calleeName)) - return rewriter.notifyMatchFailure(op, "unsupported copy VPTO signature"); - - auto llvmSourceType = - dyn_cast(adaptor.getOperands()[0].getType()); - auto llvmDestType = - dyn_cast(adaptor.getOperands()[1].getType()); - if (!llvmSourceType || !llvmDestType) - return rewriter.notifyMatchFailure(op, "expected LLVM pointer copy operands"); - - FailureOr config0 = failure(); - FailureOr config1 = failure(); - if constexpr (std::is_same_v) { - config0 = packCopyGmToUbConfig0(op, adaptor.getOperands()); - config1 = packCopyGmToUbConfig1(op, adaptor.getOperands()); - } else { - config0 = packCopyUbToGmConfig0(op, adaptor.getOperands()); - config1 = packCopyUbToGmConfig1(op, adaptor.getOperands()); - } - if (failed(config0) || failed(config1)) - return rewriter.notifyMatchFailure(op, "failed to materialize copy config"); - - SmallVector args{adaptor.getOperands()[1], adaptor.getOperands()[0], - *config0, *config1}; - auto funcType = rewriter.getFunctionType( - TypeRange{llvmDestType, llvmSourceType, rewriter.getI64Type(), - rewriter.getI64Type()}, - TypeRange{}); - auto call = rewriter.create(op.getLoc(), *calleeName, - TypeRange{}, args); - state.plannedDecls.push_back(PlannedDecl{calleeName->str(), funcType}); - rewriter.eraseOp(op); - (void)call; - return success(); - } - -private: - LoweringState &state; -}; - -class LowerCopyUbufToUbufOpPattern final - : public OpConversionPattern { -public: - explicit LowerCopyUbufToUbufOpPattern(TypeConverter &typeConverter, - MLIRContext *context, - LoweringState &state) - : OpConversionPattern(typeConverter, context), - state(state) {} - - LogicalResult - matchAndRewrite(pto::CopyUbufToUbufOp op, - pto::CopyUbufToUbufOp::Adaptor adaptor, - ConversionPatternRewriter &rewriter) const override { - auto llvmSourceType = - dyn_cast(adaptor.getOperands()[0].getType()); - auto llvmDestType = - dyn_cast(adaptor.getOperands()[1].getType()); - if (!llvmSourceType || !llvmDestType) - return rewriter.notifyMatchFailure(op, "expected LLVM pointer copy operands"); - - FailureOr config = packCopyUbToUbConfig(op, adaptor.getOperands()); - if (failed(config)) - return rewriter.notifyMatchFailure(op, "failed to materialize copy config"); - - StringRef calleeName = buildCopyUbToUbCallee(op.getContext()); - SmallVector args{adaptor.getOperands()[1], adaptor.getOperands()[0], - *config}; - auto funcType = rewriter.getFunctionType( - TypeRange{llvmDestType, llvmSourceType, rewriter.getI64Type()}, - TypeRange{}); - auto call = rewriter.create(op.getLoc(), calleeName, - TypeRange{}, args); - state.plannedDecls.push_back(PlannedDecl{calleeName.str(), funcType}); - rewriter.eraseOp(op); - (void)call; - return success(); - } - -private: - LoweringState &state; -}; - -class LowerCopyCbufToUbufOpPattern final - : public OpConversionPattern { -public: - explicit LowerCopyCbufToUbufOpPattern(TypeConverter &typeConverter, - MLIRContext *context, - LoweringState &state) - : OpConversionPattern(typeConverter, context), - state(state) {} - - LogicalResult - matchAndRewrite(pto::CopyCbufToUbufOp op, - pto::CopyCbufToUbufOp::Adaptor adaptor, - ConversionPatternRewriter &rewriter) const override { - Value sourceRaw = adaptor.getSource(); - Value destinationRaw = adaptor.getDestination(); - if (!sourceRaw || !destinationRaw) - return rewriter.notifyMatchFailure(op, "expected converted operands"); - if (!isa(sourceRaw.getType()) || - !isa(destinationRaw.getType())) - return rewriter.notifyMatchFailure(op, "expected LLVM pointer src/dst"); - - constexpr unsigned cbufAddressSpace = - static_cast(pto::AddressSpace::MAT); - constexpr unsigned ubufAddressSpace = - static_cast(pto::AddressSpace::VEC); - FailureOr source = - reinterpretPointerToAddrSpace(op, sourceRaw, cbufAddressSpace); - FailureOr destination = - reinterpretPointerToAddrSpace(op, destinationRaw, ubufAddressSpace); - if (failed(source) || failed(destination)) - return rewriter.notifyMatchFailure(op, "failed to map cbuf/ubuf pointer spaces"); - - FailureOr config = packCopyCbufToUbConfig(op, adaptor.getOperands()); - if (failed(config)) - return rewriter.notifyMatchFailure(op, "failed to materialize copy config"); - - StringRef calleeName = buildCopyCbufToUbCallee(op.getContext()); - auto funcType = rewriter.getFunctionType( - TypeRange{destination->getType(), source->getType(), - rewriter.getI64Type()}, - TypeRange{}); - rewriter.create(op.getLoc(), calleeName, TypeRange{}, - ValueRange{*destination, *source, *config}); - state.plannedDecls.push_back(PlannedDecl{calleeName.str(), funcType}); - rewriter.eraseOp(op); - return success(); - } - -private: - LoweringState &state; -}; - -class LowerCopyUbufToCbufOpPattern final - : public OpConversionPattern { -public: - explicit LowerCopyUbufToCbufOpPattern(TypeConverter &typeConverter, - MLIRContext *context, - LoweringState &state) - : OpConversionPattern(typeConverter, context), - state(state) {} - - LogicalResult - matchAndRewrite(pto::CopyUbufToCbufOp op, - pto::CopyUbufToCbufOp::Adaptor adaptor, - ConversionPatternRewriter &rewriter) const override { - Value sourceRaw = adaptor.getSource(); - Value destinationRaw = adaptor.getDestination(); - if (!sourceRaw || !destinationRaw) - return rewriter.notifyMatchFailure(op, "expected converted operands"); - if (!isa(sourceRaw.getType()) || - !isa(destinationRaw.getType())) - return rewriter.notifyMatchFailure(op, "expected LLVM pointer src/dst"); - - constexpr unsigned ubufAddressSpace = - static_cast(pto::AddressSpace::VEC); - constexpr unsigned cbufAddressSpace = - static_cast(pto::AddressSpace::MAT); - FailureOr source = - reinterpretPointerToAddrSpace(op, sourceRaw, ubufAddressSpace); - FailureOr destination = - reinterpretPointerToAddrSpace(op, destinationRaw, cbufAddressSpace); - if (failed(source) || failed(destination)) - return rewriter.notifyMatchFailure(op, "failed to map ubuf/cbuf pointer spaces"); - - FailureOr config = packCopyUbToCbufConfig(op, adaptor.getOperands()); - if (failed(config)) - return rewriter.notifyMatchFailure(op, "failed to materialize copy config"); - - StringRef calleeName = buildCopyUbToCbufCallee(op.getContext()); - auto funcType = rewriter.getFunctionType( - TypeRange{destination->getType(), source->getType(), - rewriter.getI64Type()}, - TypeRange{}); - rewriter.create(op.getLoc(), calleeName, TypeRange{}, - ValueRange{*destination, *source, *config}); - state.plannedDecls.push_back(PlannedDecl{calleeName.str(), funcType}); - rewriter.eraseOp(op); - return success(); - } - -private: - LoweringState &state; -}; - -class LowerCreateCbufMatrixOpPattern final - : public OpConversionPattern { -public: - explicit LowerCreateCbufMatrixOpPattern(TypeConverter &typeConverter, - MLIRContext *context, - LoweringState &state) - : OpConversionPattern(typeConverter, context), - state(state) {} - - LogicalResult - matchAndRewrite(pto::CreateCbufMatrixOp op, - pto::CreateCbufMatrixOp::Adaptor adaptor, - ConversionPatternRewriter &rewriter) const override { - Value destinationRaw = adaptor.getDst(); - Value rawValue = adaptor.getRawValue(); - Value repeatTimes = adaptor.getRepeatTimes(); - Value blockNum32b = adaptor.getBlockNum_32b(); - Value dstGap32b = adaptor.getDstGap_32b(); - if (!destinationRaw || !rawValue || !repeatTimes || !blockNum32b || - !dstGap32b) { - return rewriter.notifyMatchFailure(op, "expected converted operands"); - } - if (!isa(destinationRaw.getType())) { - return rewriter.notifyMatchFailure(op, "expected LLVM pointer destination"); - } - - Type i32Ty = rewriter.getI32Type(); - Type i64Ty = rewriter.getI64Type(); - const bool validControlTypes = - rawValue.getType() == i32Ty && repeatTimes.getType() == i64Ty && - blockNum32b.getType() == i64Ty && dstGap32b.getType() == i64Ty; - if (!validControlTypes) { - return rewriter.notifyMatchFailure(op, "expected i32 value and i64 controls"); - } - - constexpr unsigned cbufAddressSpace = - static_cast(pto::AddressSpace::MAT); - FailureOr destination = reinterpretPointerToAddrSpace( - op, destinationRaw, cbufAddressSpace); - if (failed(destination)) { - return rewriter.notifyMatchFailure(op, - "failed to map destination to mat/l1"); - } - - Location loc = op.getLoc(); - const uint64_t fillWordWidth = static_cast(op.getFillWordBits()); - StringRef calleeName; - Value fillPattern; - if (fillWordWidth == 16) { - Value wordMask = getI32Constant(rewriter, loc, 0xFFFFU); - Value lowWord = rewriter.create(loc, rawValue, wordMask); - Value wordBits = rewriter.create( - loc, rewriter.getI16Type(), lowWord); - fillPattern = - rewriter.create(loc, rewriter.getF16Type(), wordBits); - calleeName = "llvm.hivm.CREATE.CBUF.MATRIX.v3.u16.h"; - } else if (fillWordWidth == 32) { - fillPattern = rewriter.create(loc, i64Ty, rawValue); - calleeName = "llvm.hivm.CREATE.CBUF.MATRIX.v3.u32"; - } else { - return rewriter.notifyMatchFailure(op, "expected a 16-bit or 32-bit fill word"); - } - - Value fieldMask = getI64Constant(rewriter, loc, 0x7FFFU); - auto maskField = [&](Value value) -> Value { - return rewriter.create(loc, value, fieldMask); - }; - auto shiftField = [&](Value value, uint64_t amount) -> Value { - return rewriter.create( - loc, value, getI64Constant(rewriter, loc, amount)); - }; - - Value config = maskField(repeatTimes); - config = rewriter.create( - loc, config, shiftField(maskField(blockNum32b), 16)); - config = rewriter.create( - loc, config, shiftField(maskField(dstGap32b), 32)); - - auto funcType = rewriter.getFunctionType( - TypeRange{destination->getType(), i64Ty, fillPattern.getType()}, TypeRange{}); - rewriter.create(loc, calleeName, TypeRange{}, - ValueRange{*destination, config, fillPattern}); - state.plannedDecls.push_back(PlannedDecl{calleeName.str(), funcType}); - rewriter.eraseOp(op); - return success(); - } - -private: - LoweringState &state; -}; - -static LogicalResult lowerMadRawOp(pto::MadRawOpInterface op, - ValueRange convertedOperands, - ConversionPatternRewriter &rewriter, - LoweringState &state) { - Value lhsRaw = convertedOperands[0]; - Value rhsRaw = convertedOperands[1]; - Value dstRaw = convertedOperands[2]; - Value biasRaw = op.hasBiasOperand() ? convertedOperands[3] : Value(); - Value xt = convertedOperands[op.hasBiasOperand() ? 4 : 3]; - if (!lhsRaw || !rhsRaw || !dstRaw || !xt || - (op.hasBiasOperand() && !biasRaw)) - return rewriter.notifyMatchFailure(op, "expected converted mad raw operands"); - - if (!isa(lhsRaw.getType()) || - !isa(rhsRaw.getType()) || - !isa(dstRaw.getType()) || - (biasRaw && !isa(biasRaw.getType()))) { - return rewriter.notifyMatchFailure( - op, "expected LLVM pointer lhs/rhs/dst/bias operands"); - } - - Type i64Ty = rewriter.getI64Type(); - constexpr unsigned caAddressSpace = - static_cast(pto::AddressSpace::LEFT); - constexpr unsigned cbAddressSpace = - static_cast(pto::AddressSpace::RIGHT); - constexpr unsigned ccAddressSpace = - static_cast(pto::AddressSpace::ACC); - constexpr unsigned btAddressSpace = - static_cast(pto::AddressSpace::BIAS); - FailureOr lhs = - reinterpretPointerToAddrSpace(op, lhsRaw, caAddressSpace); - FailureOr rhs = - reinterpretPointerToAddrSpace(op, rhsRaw, cbAddressSpace); - FailureOr dst = - reinterpretPointerToAddrSpace(op, dstRaw, ccAddressSpace); - FailureOr bias; - if (biasRaw) - bias = reinterpretPointerToAddrSpace(op, biasRaw, btAddressSpace); - if (failed(lhs) || failed(rhs) || failed(dst) || - (biasRaw && failed(bias))) { - return rewriter.notifyMatchFailure(op, "failed to map cube pointer spaces"); - } - - FailureOr calleeName = - op.isMadMxFamily() ? buildMxMadCallee(op.getContext(), op) - : buildOrdinaryMadCallee(op.getContext(), op); - if (failed(calleeName)) - return rewriter.notifyMatchFailure( - op, "unsupported mad element types for raw dispatch"); - - Value callDst = *dst; - if (biasRaw) - callDst = buildMadBiasDestination(op, rewriter, *dst, *bias); - auto funcType = rewriter.getFunctionType( - TypeRange{dst->getType(), lhs->getType(), rhs->getType(), i64Ty}, - TypeRange{}); - auto call = rewriter.create( - op->getLoc(), *calleeName, TypeRange{}, - ValueRange{callDst, *lhs, *rhs, xt}); - state.plannedDecls.push_back(PlannedDecl{calleeName->str(), funcType}); - rewriter.replaceOp(op, call.getResults()); - return success(); -} - -template -class LowerMadRawPattern final : public OpConversionPattern { -public: - explicit LowerMadRawPattern(TypeConverter &typeConverter, - MLIRContext *context, LoweringState &state) - : OpConversionPattern(typeConverter, context), state(state) {} - - LogicalResult - matchAndRewrite(RawOp op, typename RawOp::Adaptor adaptor, - ConversionPatternRewriter &rewriter) const override { - auto raw = dyn_cast(op.getOperation()); - if (!raw) - return failure(); - return lowerMadRawOp(raw, adaptor.getOperands(), rewriter, state); - } - -private: - LoweringState &state; -}; - -class LowerCopyGmToCbufOpPattern final - : public OpConversionPattern { -public: - explicit LowerCopyGmToCbufOpPattern(TypeConverter &typeConverter, - MLIRContext *context, - LoweringState &state) - : OpConversionPattern(typeConverter, context), - state(state) {} - - LogicalResult matchAndRewrite( - pto::CopyGmToCbufOp op, - pto::CopyGmToCbufOp::Adaptor adaptor, - ConversionPatternRewriter &rewriter) const override { - Value sourceRaw = adaptor.getSource(); - Value destinationRaw = adaptor.getDestination(); - Value nBurst = adaptor.getNBurst(); - Value lenBurst = adaptor.getLenBurst(); - Value srcStride = adaptor.getSrcStride(); - Value dstStride = adaptor.getDstStride(); - if (!sourceRaw || !destinationRaw || !nBurst || !lenBurst || !srcStride || - !dstStride) { - return rewriter.notifyMatchFailure(op, "expected converted operands"); - } - - if (!isa(sourceRaw.getType()) || - !isa(destinationRaw.getType())) { - return rewriter.notifyMatchFailure(op, "expected LLVM pointer src/dst"); - } - - Type i64Ty = rewriter.getI64Type(); - if (nBurst.getType() != i64Ty || lenBurst.getType() != i64Ty || - srcStride.getType() != i64Ty || dstStride.getType() != i64Ty) { - return rewriter.notifyMatchFailure(op, "expected i64 config operands"); - } - - constexpr unsigned gmAddressSpace = - static_cast(pto::AddressSpace::GM); - constexpr unsigned cbufAddressSpace = - static_cast(pto::AddressSpace::MAT); - FailureOr source = reinterpretPointerToAddrSpace(op, sourceRaw, gmAddressSpace); - FailureOr destination = - reinterpretPointerToAddrSpace(op, destinationRaw, cbufAddressSpace); - if (failed(source) || failed(destination)) - return rewriter.notifyMatchFailure(op, "failed to map cbuf/gm pointer spaces"); - - FailureOr calleeName = - buildCopyGmToCbufCallee(op.getContext(), op.getSource().getType()); - if (failed(calleeName)) - return rewriter.notifyMatchFailure(op, - "unsupported copy_gm_to_cbuf element type"); - FailureOr config0 = - packCopyGmToCbufConfig0(op, nBurst, lenBurst); - FailureOr config1 = - packCopyGmToCbufConfig1(op, srcStride, dstStride); - if (failed(config0) || failed(config1)) - return rewriter.notifyMatchFailure(op, - "failed to pack copy_gm_to_cbuf config"); - - auto funcType = rewriter.getFunctionType( - TypeRange{destination->getType(), source->getType(), i64Ty, i64Ty}, - TypeRange{}); - rewriter.create( - op.getLoc(), *calleeName, TypeRange{}, - ValueRange{*destination, *source, *config0, *config1}); - state.plannedDecls.push_back(PlannedDecl{calleeName->str(), funcType}); - rewriter.eraseOp(op); - return success(); - } - -private: - LoweringState &state; -}; - -template -class LowerCopyGmToCbufMultiOpPattern final - : public OpConversionPattern { -public: - explicit LowerCopyGmToCbufMultiOpPattern(TypeConverter &typeConverter, - MLIRContext *context, - LoweringState &state) - : OpConversionPattern(typeConverter, context), state(state) {} - - LogicalResult - matchAndRewrite(CopyOp op, typename CopyOp::Adaptor adaptor, - ConversionPatternRewriter &rewriter) const override { - Value sourceRaw = adaptor.getSource(); - Value destinationRaw = adaptor.getDestination(); - if (!sourceRaw || !destinationRaw) - return rewriter.notifyMatchFailure(op, "expected converted operands"); - if (!isa(sourceRaw.getType()) || - !isa(destinationRaw.getType())) - return rewriter.notifyMatchFailure(op, "expected LLVM pointer src/dst"); - - constexpr unsigned gmAddressSpace = - static_cast(pto::AddressSpace::GM); - constexpr unsigned cbufAddressSpace = - static_cast(pto::AddressSpace::MAT); - FailureOr source = - reinterpretPointerToAddrSpace(op, sourceRaw, gmAddressSpace); - FailureOr destination = - reinterpretPointerToAddrSpace(op, destinationRaw, cbufAddressSpace); - if (failed(source) || failed(destination)) - return rewriter.notifyMatchFailure(op, "failed to map cbuf/gm pointer spaces"); - - FailureOr config0 = packCopyGmToCbufMultiConfig0( - op, adaptor.getSid(), adaptor.getLoop1SrcStride(), - adaptor.getL2CacheCtrl(), adaptor.getNValue()); - FailureOr config1 = - packCopyGmToCbufMultiConfig1(op, adaptor.getDValue(), - adaptor.getLoop4SrcStride(), - adaptor.getSmallc0En()); - if (failed(config0) || failed(config1)) - return rewriter.notifyMatchFailure(op, "failed to pack multi copy config"); - - FailureOr calleeName = [&] (MLIRContext *ctx, Type sourceType) - -> FailureOr { - if constexpr (std::is_same_v) - return buildCopyGmToCbufMultiNd2NzCallee(ctx, op.getSource().getType()); - return buildCopyGmToCbufMultiDn2NzCallee(ctx, sourceType); - }(op.getContext(), op.getSource().getType()); - if (failed(calleeName)) - return rewriter.notifyMatchFailure( - op, "unsupported copy_gm_to_cbuf_multi element type"); - - Type i64Ty = rewriter.getI64Type(); - auto funcType = rewriter.getFunctionType( - TypeRange{destination->getType(), source->getType(), i64Ty, i64Ty}, - TypeRange{}); - rewriter.create( - op.getLoc(), *calleeName, TypeRange{}, - ValueRange{*destination, *source, *config0, *config1}); - state.plannedDecls.push_back(PlannedDecl{calleeName->str(), funcType}); - rewriter.eraseOp(op); - return success(); - } - -private: - LoweringState &state; -}; - -class LowerCopyCbufToBtOpPattern final - : public OpConversionPattern { -public: - explicit LowerCopyCbufToBtOpPattern(TypeConverter &typeConverter, - MLIRContext *context, LoweringState &state) - : OpConversionPattern(typeConverter, context), - state(state) {} - - LogicalResult matchAndRewrite(pto::CopyCbufToBtOp op, - pto::CopyCbufToBtOp::Adaptor adaptor, - ConversionPatternRewriter &rewriter) const override { - Value sourceRaw = adaptor.getSource(); - Value destinationRaw = adaptor.getDestination(); - if (!sourceRaw || !destinationRaw) - return rewriter.notifyMatchFailure(op, "expected converted operands"); - if (!isa(sourceRaw.getType()) || - !isa(destinationRaw.getType())) - return rewriter.notifyMatchFailure(op, "expected LLVM pointer src/dst"); - - constexpr unsigned cbufAddressSpace = - static_cast(pto::AddressSpace::MAT); - constexpr unsigned btAddressSpace = - static_cast(pto::AddressSpace::BIAS); - FailureOr source = - reinterpretPointerToAddrSpace(op, sourceRaw, cbufAddressSpace); - FailureOr destinationPtr = - reinterpretPointerToAddrSpace(op, destinationRaw, btAddressSpace); - if (failed(source) || failed(destinationPtr)) - return rewriter.notifyMatchFailure(op, "failed to map cbuf/bt pointer spaces"); - - FailureOr config = packCopyCbufToBtConfig( - op, adaptor.getConvControl(), adaptor.getNBurst(), adaptor.getLenBurst(), - adaptor.getSourceGap(), adaptor.getDstGap()); - if (failed(config)) - return rewriter.notifyMatchFailure(op, "failed to pack copy_cbuf_to_bt config"); - - Type i64Ty = rewriter.getI64Type(); - Value destination = - rewriter.create(op.getLoc(), i64Ty, *destinationPtr); - FailureOr calleeName = buildCopyCbufToBtCallee(op); - if (failed(calleeName)) - return rewriter.notifyMatchFailure(op, "unsupported copy_cbuf_to_bt source element type"); - auto funcType = rewriter.getFunctionType( - TypeRange{i64Ty, source->getType(), i64Ty}, TypeRange{}); - rewriter.create(op.getLoc(), *calleeName, TypeRange{}, - ValueRange{destination, *source, *config}); - state.plannedDecls.push_back(PlannedDecl{calleeName->str(), funcType}); - rewriter.eraseOp(op); - return success(); - } - -private: - LoweringState &state; -}; - -class LowerCopyCbufToFbufOpPattern final - : public OpConversionPattern { -public: - explicit LowerCopyCbufToFbufOpPattern(TypeConverter &typeConverter, - MLIRContext *context, - LoweringState &state) - : OpConversionPattern(typeConverter, context), - state(state) {} - - LogicalResult matchAndRewrite(pto::CopyCbufToFbufOp op, - pto::CopyCbufToFbufOp::Adaptor adaptor, - ConversionPatternRewriter &rewriter) const override { - Value sourceRaw = adaptor.getSource(); - Value destinationRaw = adaptor.getDestination(); - if (!sourceRaw || !destinationRaw) - return rewriter.notifyMatchFailure(op, "expected converted operands"); - if (!isa(sourceRaw.getType()) || - !isa(destinationRaw.getType())) - return rewriter.notifyMatchFailure(op, "expected LLVM pointer src/dst"); - - constexpr unsigned cbufAddressSpace = - static_cast(pto::AddressSpace::MAT); - constexpr unsigned fbufAddressSpace = 7; - FailureOr source = - reinterpretPointerToAddrSpace(op, sourceRaw, cbufAddressSpace); - FailureOr destination = - reinterpretPointerToAddrSpace(op, destinationRaw, fbufAddressSpace); - if (failed(source) || failed(destination)) - return rewriter.notifyMatchFailure(op, "failed to map cbuf/fbuf pointer spaces"); - - FailureOr config = packCopyCbufToFbufConfig( - op, adaptor.getNBurst(), adaptor.getLenBurst(), adaptor.getSourceGap(), - adaptor.getDstGap()); - if (failed(config)) - return rewriter.notifyMatchFailure(op, "failed to pack copy_cbuf_to_fbuf config"); - - Type i64Ty = rewriter.getI64Type(); - StringRef calleeName = buildCopyCbufToFbufCallee(op.getContext()); - auto funcType = rewriter.getFunctionType( - TypeRange{destination->getType(), source->getType(), i64Ty}, TypeRange{}); - rewriter.create(op.getLoc(), calleeName, TypeRange{}, - ValueRange{*destination, *source, *config}); - state.plannedDecls.push_back(PlannedDecl{calleeName.str(), funcType}); - rewriter.eraseOp(op); - return success(); - } - -private: - LoweringState &state; -}; - -class LowerLoadCbufToCaOpPattern final - : public OpConversionPattern { -public: - explicit LowerLoadCbufToCaOpPattern(TypeConverter &typeConverter, - MLIRContext *context, - LoweringState &state) - : OpConversionPattern(typeConverter, context), - state(state) {} - - LogicalResult matchAndRewrite(pto::LoadCbufToCaOp op, - pto::LoadCbufToCaOp::Adaptor adaptor, - ConversionPatternRewriter &rewriter) const override { - Value sourceRaw = adaptor.getSource(); - Value destinationRaw = adaptor.getDestination(); - Value mStart = adaptor.getMStart(); - Value kStart = adaptor.getKStart(); - Value mStep = adaptor.getMStep(); - Value kStep = adaptor.getKStep(); - Value srcStride = adaptor.getSrcStride(); - Value dstStride = adaptor.getDstStride(); - if (!sourceRaw || !destinationRaw || !mStart || !kStart || !mStep || - !kStep || !srcStride || !dstStride) - return rewriter.notifyMatchFailure(op, "expected converted operands"); - - if (!isa(sourceRaw.getType()) || - !isa(destinationRaw.getType())) { - return rewriter.notifyMatchFailure(op, "expected LLVM pointer src/dst"); - } - - Type i64Ty = rewriter.getI64Type(); - - constexpr unsigned cbufAddressSpace = - static_cast(pto::AddressSpace::MAT); - constexpr unsigned caAddressSpace = - static_cast(pto::AddressSpace::LEFT); - FailureOr source = reinterpretPointerToAddrSpace(op, sourceRaw, cbufAddressSpace); - FailureOr destination = - reinterpretPointerToAddrSpace(op, destinationRaw, caAddressSpace); - if (failed(source) || failed(destination)) - return rewriter.notifyMatchFailure(op, "failed to map cbuf/ca pointer spaces"); - - FailureOr config0 = - packLoadCbufToCaConfig0(op, mStart, kStart, mStep, kStep); - FailureOr config1 = - packLoadCbufToCaConfig1(op, srcStride, dstStride); - if (failed(config0) || failed(config1)) - return rewriter.notifyMatchFailure(op, "failed to pack load_cbuf_to_ca config"); - Value transpose = - getI64Constant(rewriter, op.getLoc(), op.getTranspose() ? 1 : 0); - - FailureOr calleeName = - buildLoadCbufToCaCallee(op.getContext(), op.getSource().getType()); - if (failed(calleeName)) - return rewriter.notifyMatchFailure(op, - "unsupported load_cbuf_to_ca element type"); - auto funcType = rewriter.getFunctionType( - TypeRange{destination->getType(), source->getType(), i64Ty, i64Ty, - i64Ty}, - TypeRange{}); - rewriter.create(op.getLoc(), *calleeName, TypeRange{}, - ValueRange{*destination, *source, *config0, - *config1, transpose}); - state.plannedDecls.push_back(PlannedDecl{calleeName->str(), funcType}); - rewriter.eraseOp(op); - return success(); - } - -private: - LoweringState &state; -}; - -template -class LowerLoadCbufToS4OpPattern final : public OpConversionPattern { -public: - explicit LowerLoadCbufToS4OpPattern(TypeConverter &typeConverter, - MLIRContext *context, - LoweringState &state) - : OpConversionPattern(typeConverter, context), state(state) {} - - LogicalResult - matchAndRewrite(LoadOp op, typename LoadOp::Adaptor adaptor, - ConversionPatternRewriter &rewriter) const override { - Value sourceRaw = adaptor.getSource(); - Value destinationRaw = adaptor.getDestination(); - if (!sourceRaw || !destinationRaw) - return rewriter.notifyMatchFailure(op, "expected converted operands"); - if (!isa(sourceRaw.getType()) || - !isa(destinationRaw.getType())) - return rewriter.notifyMatchFailure(op, "expected LLVM pointer src/dst"); - - constexpr unsigned cbufAddressSpace = - static_cast(pto::AddressSpace::MAT); - constexpr unsigned targetAddressSpace = - std::is_same_v - ? static_cast(pto::AddressSpace::LEFT) - : static_cast(pto::AddressSpace::RIGHT); - FailureOr source = - reinterpretPointerToAddrSpace(op, sourceRaw, cbufAddressSpace); - FailureOr destination = - reinterpretPointerToAddrSpace(op, destinationRaw, targetAddressSpace); - if (failed(source) || failed(destination)) - return rewriter.notifyMatchFailure(op, "failed to map cbuf/cube pointer spaces"); - - FailureOr config0 = packLoadCbufToS4Config0( - op, adaptor.getMStart(), adaptor.getKStart(), adaptor.getMStep(), - adaptor.getKStep()); - FailureOr config1 = - packLoadCbufToS4Config1(op, adaptor.getSrcStride(), - adaptor.getDstStride()); - if (failed(config0) || failed(config1)) - return rewriter.notifyMatchFailure(op, "failed to pack load_cbuf_to_*_s4 config"); - - Value transpose = - castIntegerLikeTo(op, adaptor.getTranspose(), rewriter.getI64Type()); - if (!transpose) - return rewriter.notifyMatchFailure(op, "failed to cast transpose to i64"); - - FailureOr calleeName = - std::is_same_v - ? buildLoadCbufToCaS4Callee(op.getContext(), - op.getSource().getType()) - : buildLoadCbufToCbS4Callee(op.getContext(), - op.getSource().getType()); - if (failed(calleeName)) - return rewriter.notifyMatchFailure( - op, "unsupported load_cbuf_to_*_s4 element type"); - Type i64Ty = rewriter.getI64Type(); - auto funcType = rewriter.getFunctionType( - TypeRange{destination->getType(), source->getType(), i64Ty, i64Ty, - i64Ty}, - TypeRange{}); - rewriter.create( - op.getLoc(), *calleeName, TypeRange{}, - ValueRange{*destination, *source, *config0, *config1, transpose}); - state.plannedDecls.push_back(PlannedDecl{calleeName->str(), funcType}); - rewriter.eraseOp(op); - return success(); - } - -private: - LoweringState &state; -}; - -class LowerLoadCbufToCbOpPattern final - : public OpConversionPattern { -public: - explicit LowerLoadCbufToCbOpPattern(TypeConverter &typeConverter, - MLIRContext *context, - LoweringState &state) - : OpConversionPattern(typeConverter, context), - state(state) {} - - LogicalResult matchAndRewrite(pto::LoadCbufToCbOp op, - pto::LoadCbufToCbOp::Adaptor adaptor, - ConversionPatternRewriter &rewriter) const override { - Value sourceRaw = adaptor.getSource(); - Value destinationRaw = adaptor.getDestination(); - Value mStart = adaptor.getMStart(); - Value kStart = adaptor.getKStart(); - Value mStep = adaptor.getMStep(); - Value kStep = adaptor.getKStep(); - Value srcStride = adaptor.getSrcStride(); - Value dstStride = adaptor.getDstStride(); - if (!sourceRaw || !destinationRaw || !mStart || !kStart || !mStep || - !kStep || !srcStride || !dstStride) - return rewriter.notifyMatchFailure(op, "expected converted operands"); - - if (!isa(sourceRaw.getType()) || - !isa(destinationRaw.getType())) { - return rewriter.notifyMatchFailure(op, "expected LLVM pointer src/dst"); - } - - Type i64Ty = rewriter.getI64Type(); - - constexpr unsigned cbufAddressSpace = - static_cast(pto::AddressSpace::MAT); - constexpr unsigned cbAddressSpace = - static_cast(pto::AddressSpace::RIGHT); - FailureOr source = reinterpretPointerToAddrSpace(op, sourceRaw, cbufAddressSpace); - FailureOr destination = - reinterpretPointerToAddrSpace(op, destinationRaw, cbAddressSpace); - if (failed(source) || failed(destination)) - return rewriter.notifyMatchFailure(op, "failed to map cbuf/cb pointer spaces"); - - bool transpose = op.getTranspose(); - FailureOr config0 = - packLoadCbufToCbConfig0(op, mStart, kStart, mStep, kStep); - FailureOr config1 = - packLoadCbufToCbConfig1(op, srcStride, dstStride); - if (failed(config0) || failed(config1)) - return rewriter.notifyMatchFailure(op, "failed to pack load_cbuf_to_cb config"); - Value transposeValue = - getI64Constant(rewriter, op.getLoc(), transpose ? 1 : 0); - - FailureOr calleeName = - buildLoadCbufToCbCallee(op.getContext(), op.getSource().getType()); - if (failed(calleeName)) - return rewriter.notifyMatchFailure(op, - "unsupported load_cbuf_to_cb element type"); - auto funcType = rewriter.getFunctionType( - TypeRange{destination->getType(), source->getType(), i64Ty, i64Ty, - i64Ty}, - TypeRange{}); - rewriter.create(op.getLoc(), *calleeName, TypeRange{}, - ValueRange{*destination, *source, *config0, - *config1, transposeValue}); - state.plannedDecls.push_back(PlannedDecl{calleeName->str(), funcType}); - rewriter.eraseOp(op); - return success(); - } - -private: - LoweringState &state; -}; - -class LowerLoadCbufToCaMxOpPattern final - : public OpConversionPattern { -public: - explicit LowerLoadCbufToCaMxOpPattern(TypeConverter &typeConverter, - MLIRContext *context, - LoweringState &state) - : OpConversionPattern(typeConverter, context), - state(state) {} - - LogicalResult - matchAndRewrite(pto::LoadCbufToCaMxOp op, OpAdaptor adaptor, - ConversionPatternRewriter &rewriter) const override { - Value srcRaw = adaptor.getSource(); - Value dstRaw = adaptor.getDestination(); - if (!srcRaw || !dstRaw || !adaptor.getXStartPosition() || - !adaptor.getYStartPosition() || !adaptor.getXStep() || - !adaptor.getYStep() || !adaptor.getSrcStride() || - !adaptor.getDstStride()) - return rewriter.notifyMatchFailure(op, "expected converted operands"); - if (!isa(srcRaw.getType()) || - !isa(dstRaw.getType())) - return rewriter.notifyMatchFailure(op, "expected LLVM pointer src/dst"); - - constexpr unsigned cbufAddressSpace = - static_cast(pto::AddressSpace::MAT); - constexpr unsigned caAddressSpace = - static_cast(pto::AddressSpace::LEFT); - FailureOr src = reinterpretPointerToAddrSpace(op, srcRaw, cbufAddressSpace); - FailureOr dst = reinterpretPointerToAddrSpace(op, dstRaw, caAddressSpace); - if (failed(src) || failed(dst)) - return rewriter.notifyMatchFailure(op, "failed to map cbuf/ca pointer spaces"); - - Type sourceElemType = cast(op.getSource().getType()).getElementType(); - unsigned elemBitWidth = pto::getPTOStorageElemBitWidth(sourceElemType); - if (elemBitWidth == 0 || (elemBitWidth % 8) != 0) - return rewriter.notifyMatchFailure(op, - "unsupported load_cbuf_to_ca_mx element type"); - FailureOr config0 = - packLoadCbufToCaConfig0(op, adaptor.getXStartPosition(), - adaptor.getYStartPosition(), adaptor.getXStep(), - adaptor.getYStep()); - FailureOr config1 = - packLoadCbufToCaConfig1(op, adaptor.getSrcStride(), - adaptor.getDstStride()); - if (failed(config0) || failed(config1)) - return rewriter.notifyMatchFailure(op, - "failed to pack load_cbuf_to_ca_mx config"); - auto i64Ty = rewriter.getI64Type(); - Value dstAddr = rewriter.create(op.getLoc(), i64Ty, *dst); - - StringRef calleeName = buildLoadCbufToCaMxCallee(op.getContext()); - auto funcType = rewriter.getFunctionType( - TypeRange{i64Ty, src->getType(), i64Ty, i64Ty}, - TypeRange{}); - rewriter.create(op.getLoc(), calleeName, TypeRange{}, - ValueRange{dstAddr, *src, *config0, *config1}); - state.plannedDecls.push_back(PlannedDecl{calleeName.str(), funcType}); - rewriter.eraseOp(op); - return success(); - } - -private: - LoweringState &state; -}; - -class LowerLoadCbufToCbMxOpPattern final - : public OpConversionPattern { -public: - explicit LowerLoadCbufToCbMxOpPattern(TypeConverter &typeConverter, - MLIRContext *context, - LoweringState &state) - : OpConversionPattern(typeConverter, context), - state(state) {} - - LogicalResult - matchAndRewrite(pto::LoadCbufToCbMxOp op, OpAdaptor adaptor, - ConversionPatternRewriter &rewriter) const override { - Value srcRaw = adaptor.getSource(); - Value dstRaw = adaptor.getDestination(); - if (!srcRaw || !dstRaw || !adaptor.getXStartPosition() || - !adaptor.getYStartPosition() || !adaptor.getXStep() || - !adaptor.getYStep() || !adaptor.getSrcStride() || - !adaptor.getDstStride()) - return rewriter.notifyMatchFailure(op, "expected converted operands"); - if (!isa(srcRaw.getType()) || - !isa(dstRaw.getType())) - return rewriter.notifyMatchFailure(op, "expected LLVM pointer src/dst"); - - constexpr unsigned cbufAddressSpace = - static_cast(pto::AddressSpace::MAT); - constexpr unsigned cbAddressSpace = - static_cast(pto::AddressSpace::RIGHT); - FailureOr src = reinterpretPointerToAddrSpace(op, srcRaw, cbufAddressSpace); - FailureOr dst = reinterpretPointerToAddrSpace(op, dstRaw, cbAddressSpace); - if (failed(src) || failed(dst)) - return rewriter.notifyMatchFailure(op, "failed to map cbuf/cb pointer spaces"); - - Type sourceElemType = cast(op.getSource().getType()).getElementType(); - unsigned elemBitWidth = pto::getPTOStorageElemBitWidth(sourceElemType); - if (elemBitWidth == 0 || (elemBitWidth % 8) != 0) - return rewriter.notifyMatchFailure(op, - "unsupported load_cbuf_to_cb_mx element type"); - FailureOr config0 = - packLoadCbufToCbConfig0(op, adaptor.getXStartPosition(), - adaptor.getYStartPosition(), adaptor.getXStep(), - adaptor.getYStep()); - FailureOr config1 = - packLoadCbufToCbConfig1(op, adaptor.getSrcStride(), - adaptor.getDstStride()); - if (failed(config0) || failed(config1)) - return rewriter.notifyMatchFailure(op, - "failed to pack load_cbuf_to_cb_mx config"); - auto i64Ty = rewriter.getI64Type(); - Value dstAddr = rewriter.create(op.getLoc(), i64Ty, *dst); - - StringRef calleeName = buildLoadCbufToCbMxCallee(op.getContext()); - auto funcType = rewriter.getFunctionType( - TypeRange{i64Ty, src->getType(), i64Ty, i64Ty}, - TypeRange{}); - rewriter.create(op.getLoc(), calleeName, TypeRange{}, - ValueRange{dstAddr, *src, *config0, *config1}); - state.plannedDecls.push_back(PlannedDecl{calleeName.str(), funcType}); - rewriter.eraseOp(op); - return success(); - } - -private: - LoweringState &state; -}; - -class LowerCopyMatrixCcToGmOpPattern final - : public OpConversionPattern { -public: - explicit LowerCopyMatrixCcToGmOpPattern(TypeConverter &typeConverter, - MLIRContext *context, - LoweringState &state) - : OpConversionPattern(typeConverter, context), - state(state) {} - - LogicalResult matchAndRewrite( - pto::CopyMatrixCcToGmOp op, pto::CopyMatrixCcToGmOp::Adaptor adaptor, - ConversionPatternRewriter &rewriter) const override { - Value sourceRaw = adaptor.getSource(); - Value destinationRaw = adaptor.getDestination(); - Value xm = adaptor.getXm(); - Value xt = adaptor.getXt(); - if (!sourceRaw || !destinationRaw || !xm || !xt) - return rewriter.notifyMatchFailure(op, "expected converted operands"); - - if (!isa(sourceRaw.getType()) || - !isa(destinationRaw.getType())) { - return rewriter.notifyMatchFailure(op, "expected LLVM pointer src/dst"); - } - - Type i64Ty = rewriter.getI64Type(); - if (xm.getType() != i64Ty || xt.getType() != i64Ty) - return rewriter.notifyMatchFailure(op, "expected i64 xm/xt operands"); - - constexpr unsigned gmAddressSpace = - static_cast(pto::AddressSpace::GM); - constexpr unsigned ccAddressSpace = - static_cast(pto::AddressSpace::ACC); - FailureOr source = reinterpretPointerToAddrSpace(op, sourceRaw, ccAddressSpace); - FailureOr destination = - reinterpretPointerToAddrSpace(op, destinationRaw, gmAddressSpace); - if (failed(source) || failed(destination)) - return rewriter.notifyMatchFailure(op, "failed to map cc/gm pointer spaces"); - - StringRef calleeName = buildCopyMatrixCcToGmCallee(op.getContext()); - auto funcType = rewriter.getFunctionType( - TypeRange{destination->getType(), source->getType(), i64Ty, i64Ty}, - TypeRange{}); - rewriter.create(op.getLoc(), calleeName, TypeRange{}, - ValueRange{*destination, *source, xm, xt}); - state.plannedDecls.push_back(PlannedDecl{calleeName.str(), funcType}); - rewriter.eraseOp(op); - return success(); - } - -private: - LoweringState &state; -}; - -template -class LowerCopyMatrixCcToBufOpPattern final - : public OpConversionPattern { -public: - explicit LowerCopyMatrixCcToBufOpPattern(TypeConverter &typeConverter, - MLIRContext *context, - LoweringState &state) - : OpConversionPattern(typeConverter, context), state(state) {} - - LogicalResult - matchAndRewrite(CopyOp op, typename CopyOp::Adaptor adaptor, - ConversionPatternRewriter &rewriter) const override { - Value sourceRaw = adaptor.getSource(); - Value destinationRaw = adaptor.getDestination(); - if (!sourceRaw || !destinationRaw) - return rewriter.notifyMatchFailure(op, "expected converted operands"); - if (!isa(sourceRaw.getType()) || - !isa(destinationRaw.getType())) - return rewriter.notifyMatchFailure(op, "expected LLVM pointer src/dst"); - - constexpr unsigned ccAddressSpace = - static_cast(pto::AddressSpace::ACC); - constexpr unsigned targetAddressSpace = - std::is_same_v - ? static_cast(pto::AddressSpace::MAT) - : static_cast(pto::AddressSpace::VEC); - FailureOr source = - reinterpretPointerToAddrSpace(op, sourceRaw, ccAddressSpace); - FailureOr destination = - reinterpretPointerToAddrSpace(op, destinationRaw, targetAddressSpace); - if (failed(source) || failed(destination)) - return rewriter.notifyMatchFailure(op, "failed to map cc->buf pointer spaces"); - - Type i64Ty = rewriter.getI64Type(); - Value config0 = castIntegerLikeTo(op, adaptor.getConfig0(), i64Ty); - Value config1 = castIntegerLikeTo(op, adaptor.getConfig1(), i64Ty); - if (!config0 || !config1) - return rewriter.notifyMatchFailure(op, "failed to cast config operands to i64"); - - FailureOr calleeName = - std::is_same_v - ? FailureOr(buildCopyMatrixCcToCbufCallee(op.getContext())) - : buildCopyMatrixCcToUbCallee(op.getContext(), - op.getDestination().getType()); - if (failed(calleeName)) - return rewriter.notifyMatchFailure( - op, "unsupported copy_matrix_cc_to_{cbuf,ub} element type"); - auto funcType = rewriter.getFunctionType( - TypeRange{destination->getType(), source->getType(), i64Ty, i64Ty}, - TypeRange{}); - rewriter.create(op.getLoc(), *calleeName, TypeRange{}, - ValueRange{*destination, *source, config0, - config1}); - state.plannedDecls.push_back(PlannedDecl{calleeName->str(), funcType}); - rewriter.eraseOp(op); - return success(); - } - -private: - LoweringState &state; -}; - -template -class LowerVecScalarMaskedOpPattern final - : public OpConversionPattern { -public: - explicit LowerVecScalarMaskedOpPattern(TypeConverter &typeConverter, - MLIRContext *context, - LoweringState &state) - : OpConversionPattern(typeConverter, context), state(state) {} - - LogicalResult - matchAndRewrite(VecScalarOp op, typename VecScalarOp::Adaptor adaptor, - ConversionPatternRewriter &rewriter) const override { - StringRef stem = getVecScalarMaskedStem(); - FailureOr calleeName = - usesSignedVecScalarCANN900Callee() - ? buildCANN900SignedModeTypedCallee( - op.getContext(), op.getResult().getType(), stem, "x") - : buildCANN900ModeTypedCallee(op.getContext(), - op.getResult().getType(), stem, "x"); - if (failed(calleeName)) - return rewriter.notifyMatchFailure(op, - "unsupported vec-scalar VPTO signature"); - - Type resultType = - this->getTypeConverter()->convertType(op.getResult().getType()); - if (!resultType) - return rewriter.notifyMatchFailure( - op, "failed to convert vec-scalar result type"); - - Value input = adaptor.getOperands()[0]; - Value scalar = adaptor.getOperands()[1]; - Value mask = adaptor.getOperands()[2]; - Type expectedMaskType = - this->getTypeConverter()->convertType(op->getOperand(2).getType()); - if (!input || !scalar || !mask || input.getType() != resultType || - mask.getType() != expectedMaskType) { - return rewriter.notifyMatchFailure( - op, "unexpected converted vec-scalar VPTO operand types"); - } - - auto call = rewriter.create( - op.getLoc(), *calleeName, TypeRange{resultType}, - ValueRange{input, scalar, mask}); - state.plannedDecls.push_back( - PlannedDecl{calleeName->str(), call.getCalleeType()}); - rewriter.replaceOp(op, call.getResults()); - return success(); - } - -private: - LoweringState &state; -}; - -template -class LowerReductionUnaryOpPattern final - : public OpConversionPattern { -public: - explicit LowerReductionUnaryOpPattern(TypeConverter &typeConverter, - MLIRContext *context, - LoweringState &state) - : OpConversionPattern(typeConverter, context), state(state) {} - - LogicalResult - matchAndRewrite(ReductionOp op, typename ReductionOp::Adaptor adaptor, - ConversionPatternRewriter &rewriter) const override { - StringRef stem = getReductionUnaryStem(); - FailureOr calleeName = - usesSignedReductionCANN900Callee() - ? buildCANN900SignedModeTypedCallee( - op.getContext(), op.getResult().getType(), stem, "x") - : buildCANN900ModeTypedCallee(op.getContext(), - op.getResult().getType(), stem, "x"); - if (failed(calleeName)) - return rewriter.notifyMatchFailure(op, - "unsupported reduction VPTO signature"); - - Type resultType = - this->getTypeConverter()->convertType(op.getResult().getType()); - Type maskType = this->getTypeConverter()->convertType(op.getMask().getType()); - if (!resultType || !maskType) { - return rewriter.notifyMatchFailure( - op, "failed to convert reduction result type"); - } - - Value input = adaptor.getInput(); - Value mask = adaptor.getMask(); - if (!input || !mask || input.getType() != resultType || - mask.getType() != maskType) { - return rewriter.notifyMatchFailure( - op, "unexpected converted reduction operand types"); - } - - auto call = rewriter.create(op.getLoc(), *calleeName, - TypeRange{resultType}, - ValueRange{input, mask}); - state.plannedDecls.push_back( - PlannedDecl{calleeName->str(), call.getCalleeType()}); - rewriter.replaceOp(op, call.getResults()); - return success(); - } - -private: - LoweringState &state; -}; - -template -class LowerHistogramOpPattern final : public OpConversionPattern { -public: - explicit LowerHistogramOpPattern(TypeConverter &typeConverter, - MLIRContext *context, LoweringState &state) - : OpConversionPattern(typeConverter, context), state(state) {} - - LogicalResult - matchAndRewrite(HistOp op, typename HistOp::Adaptor adaptor, - ConversionPatternRewriter &rewriter) const override { - StringRef calleeName = getHistogramCallee(op.getContext()); - if (calleeName.empty()) - return rewriter.notifyMatchFailure(op, "unsupported histogram op"); - - Type resultType = - this->getTypeConverter()->convertType(op.getResult().getType()); - Type sourceType = - this->getTypeConverter()->convertType(op.getSource().getType()); - Type maskType = this->getTypeConverter()->convertType(op.getMask().getType()); - if (!resultType || !sourceType || !maskType) - return rewriter.notifyMatchFailure(op, "failed to convert histogram types"); - - Value acc = adaptor.getAcc(); - Value source = adaptor.getSource(); - Value mask = adaptor.getMask(); - Value bin = adaptor.getBin(); - if (!acc || !source || !mask || !bin || acc.getType() != resultType || - source.getType() != sourceType || mask.getType() != maskType || - !bin.getType().isInteger(32)) { - return rewriter.notifyMatchFailure( - op, "unexpected converted histogram operand types"); - } - - auto funcType = rewriter.getFunctionType( - TypeRange{resultType, sourceType, maskType, rewriter.getI32Type()}, - TypeRange{resultType}); - auto call = rewriter.create( - op.getLoc(), calleeName, TypeRange{resultType}, - ValueRange{acc, source, mask, bin}); - state.plannedDecls.push_back(PlannedDecl{calleeName.str(), funcType}); - rewriter.replaceOp(op, call.getResults()); - return success(); - } - -private: - LoweringState &state; -}; - -template -class LowerExtremaPredicateOpPattern final - : public OpConversionPattern { -public: - explicit LowerExtremaPredicateOpPattern(TypeConverter &typeConverter, - MLIRContext *context, - LoweringState &state) - : OpConversionPattern(typeConverter, context), state(state) {} - - LogicalResult - matchAndRewrite(ExtremaOp op, typename ExtremaOp::Adaptor adaptor, - ConversionPatternRewriter &rewriter) const override { - FailureOr calleeName = - buildExtremaPredicateCallee(op.getContext(), - op.getValue().getType()); - if (failed(calleeName)) - return rewriter.notifyMatchFailure( - op, "unsupported extrema-predicate VPTO signature"); - - Type valueType = - this->getTypeConverter()->convertType(op.getValue().getType()); - Type predicateType = - this->getTypeConverter()->convertType(op.getPredicate().getType()); - if (!valueType || !predicateType) - return rewriter.notifyMatchFailure( - op, "failed to convert extrema-predicate result types"); - - Value input = adaptor.getInput(); - Value mask = adaptor.getMask(); - if (!input || !mask || input.getType() != valueType || - mask.getType() != predicateType) { - return rewriter.notifyMatchFailure( - op, "unexpected converted extrema-predicate operand types"); - } - - auto call = rewriter.create( - op.getLoc(), *calleeName, TypeRange{valueType, predicateType}, - ValueRange{input, mask}); - state.plannedDecls.push_back( - PlannedDecl{calleeName->str(), call.getCalleeType()}); - rewriter.replaceOp(op, call.getResults()); - return success(); - } - -private: - LoweringState &state; -}; - -template -class LowerWideningReductionUnaryOpPattern final - : public OpConversionPattern { -public: - explicit LowerWideningReductionUnaryOpPattern(TypeConverter &typeConverter, - MLIRContext *context, - LoweringState &state) - : OpConversionPattern(typeConverter, context), state(state) {} - - LogicalResult - matchAndRewrite(ReductionOp op, typename ReductionOp::Adaptor adaptor, - ConversionPatternRewriter &rewriter) const override { - StringRef stem = getReductionUnaryStem(); - FailureOr calleeName = buildCANN900WideningReductionCallee( - op.getContext(), op.getInput().getType(), op.getResult().getType(), - stem, "x"); - if (failed(calleeName)) - return rewriter.notifyMatchFailure(op, - "unsupported widening reduction VPTO signature"); - - Type inputType = - this->getTypeConverter()->convertType(op.getInput().getType()); - Type resultType = - this->getTypeConverter()->convertType(op.getResult().getType()); - Type maskType = this->getTypeConverter()->convertType(op.getMask().getType()); - if (!inputType || !resultType || !maskType) - return rewriter.notifyMatchFailure(op, - "failed to convert widening reduction types"); - - Value input = adaptor.getInput(); - Value mask = adaptor.getMask(); - if (!input || !mask || input.getType() != inputType || - mask.getType() != maskType) { - return rewriter.notifyMatchFailure( - op, "unexpected converted widening reduction operand types"); - } - - auto call = rewriter.create(op.getLoc(), *calleeName, - TypeRange{resultType}, - ValueRange{input, mask}); - state.plannedDecls.push_back( - PlannedDecl{calleeName->str(), call.getCalleeType()}); - rewriter.replaceOp(op, call.getResults()); - return success(); - } - -private: - LoweringState &state; -}; - -class LowerVselOpPattern final : public OpConversionPattern { -public: - explicit LowerVselOpPattern(TypeConverter &typeConverter, MLIRContext *context, - LoweringState &state) - : OpConversionPattern(typeConverter, context), state(state) {} - - LogicalResult - matchAndRewrite(pto::VselOp op, pto::VselOp::Adaptor adaptor, - ConversionPatternRewriter &rewriter) const override { - FailureOr calleeName = - buildVselCallee(op.getContext(), op.getResult().getType()); - if (failed(calleeName)) - return rewriter.notifyMatchFailure(op, "unsupported vsel VPTO signature"); - - Type resultType = this->getTypeConverter()->convertType(op.getResult().getType()); - Type maskType = this->getTypeConverter()->convertType(op.getMask().getType()); - if (!resultType || !maskType) - return rewriter.notifyMatchFailure(op, "failed to convert vsel result type"); - - Value src0 = adaptor.getSrc0(); - Value src1 = adaptor.getSrc1(); - Value mask = adaptor.getMask(); - if (!src0 || !src1 || !mask || src0.getType() != resultType || - src1.getType() != resultType || mask.getType() != maskType) { - return rewriter.notifyMatchFailure(op, - "unexpected converted vsel operand types"); - } - - auto call = rewriter.create(op.getLoc(), *calleeName, - TypeRange{resultType}, - ValueRange{src0, src1, mask}); - state.plannedDecls.push_back( - PlannedDecl{calleeName->str(), call.getCalleeType()}); - rewriter.replaceOp(op, call.getResults()); - return success(); - } - -private: - LoweringState &state; -}; - -class LowerVdupOpPattern final : public OpConversionPattern { -public: - explicit LowerVdupOpPattern(TypeConverter &typeConverter, MLIRContext *context, - LoweringState &state) - : OpConversionPattern(typeConverter, context), state(state) {} - - LogicalResult - matchAndRewrite(pto::VdupOp op, pto::VdupOp::Adaptor adaptor, - ConversionPatternRewriter &rewriter) const override { - FailureOr calleeName = buildVdupCallee(op.getContext(), op); - if (failed(calleeName)) - return rewriter.notifyMatchFailure(op, "unsupported vdup VPTO signature"); - - Type resultType = this->getTypeConverter()->convertType(op.getResult().getType()); - Type maskType = this->getTypeConverter()->convertType(op.getMask().getType()); - if (!resultType || !maskType) - return rewriter.notifyMatchFailure(op, "failed to convert vdup result type"); - - Value mask = adaptor.getMask(); - if (!mask || mask.getType() != maskType) - return rewriter.notifyMatchFailure(op, - "unexpected converted vdup mask type"); - - SmallVector callArgs; - bool vectorInput = isa(op.getInput().getType()); - if (vectorInput) { - Value input = adaptor.getInput(); - if (!input || input.getType() != resultType) { - return rewriter.notifyMatchFailure( - op, "vector-input vdup requires matching result type"); - } - callArgs.push_back(input); - } else { - Type scalarType = getElementTypeFromVectorLike(op.getResult().getType()); - if (!scalarType || - (op.getInput().getType() != scalarType && - !isCompatibleScalarForSemanticType(scalarType, - op.getInput().getType()))) { - return rewriter.notifyMatchFailure(op, - "unexpected scalar-input vdup type"); - } - FailureOr normalizedScalar = - normalizeVdupScalarOperand(rewriter, op.getLoc(), adaptor.getInput(), - op.getResult().getType()); - if (failed(normalizedScalar)) - return rewriter.notifyMatchFailure(op, - "failed to normalize scalar vdup input"); - Value scalarForCall = normalizeByteScalarOperandForCANN900VectorCall( - rewriter, op.getLoc(), *normalizedScalar, scalarType); - callArgs.push_back(scalarForCall); - } - - callArgs.push_back(mask); - callArgs.push_back(getI32Constant(rewriter, op.getLoc(), 1)); - - auto call = rewriter.create( - op.getLoc(), *calleeName, TypeRange{resultType}, callArgs); - state.plannedDecls.push_back( - PlannedDecl{calleeName->str(), call.getCalleeType()}); - rewriter.replaceOp(op, call.getResults()); - return success(); - } - -private: - LoweringState &state; -}; - -class LowerVbrOpPattern final : public OpConversionPattern { -public: - explicit LowerVbrOpPattern(TypeConverter &typeConverter, MLIRContext *context, - LoweringState &state) - : OpConversionPattern(typeConverter, context), state(state) {} - - LogicalResult - matchAndRewrite(pto::VbrOp op, pto::VbrOp::Adaptor adaptor, - ConversionPatternRewriter &rewriter) const override { - FailureOr calleeName = - buildVbrCallee(op.getContext(), op.getResult().getType()); - if (failed(calleeName)) - return rewriter.notifyMatchFailure(op, "unsupported vbr VPTO signature"); - - Type resultType = this->getTypeConverter()->convertType(op.getResult().getType()); - if (!resultType) - return rewriter.notifyMatchFailure(op, "failed to convert vbr result type"); - - Value scalar = adaptor.getValue(); - Type expectedScalarType = - this->getTypeConverter()->convertType(op.getValue().getType()); - if (!scalar || !expectedScalarType || scalar.getType() != expectedScalarType) - return rewriter.notifyMatchFailure(op, - "unexpected converted vbr operand type"); - - scalar = normalizeByteScalarOperandForCANN900VectorCall( - rewriter, op.getLoc(), scalar, - cast(op.getResult().getType()).getElementType()); - - auto funcType = rewriter.getFunctionType(TypeRange{scalar.getType()}, - TypeRange{resultType}); - auto call = rewriter.create(op.getLoc(), *calleeName, - TypeRange{resultType}, - ValueRange{scalar}); - state.plannedDecls.push_back(PlannedDecl{calleeName->str(), funcType}); - rewriter.replaceOp(op, call.getResults()); - return success(); - } - -private: - LoweringState &state; -}; - -class LowerVselrOpPattern final : public OpConversionPattern { -public: - explicit LowerVselrOpPattern(TypeConverter &typeConverter, MLIRContext *context, - LoweringState &state) - : OpConversionPattern(typeConverter, context), state(state) {} - - LogicalResult - matchAndRewrite(pto::VselrOp op, pto::VselrOp::Adaptor adaptor, - ConversionPatternRewriter &rewriter) const override { - FailureOr calleeName = - buildVselrCallee(op.getContext(), op.getResult().getType()); - if (failed(calleeName)) - return rewriter.notifyMatchFailure(op, "unsupported vselr VPTO signature"); - - Type resultType = this->getTypeConverter()->convertType(op.getResult().getType()); - if (!resultType) - return rewriter.notifyMatchFailure(op, - "unexpected converted vselr result type"); - auto resultVectorType = dyn_cast(resultType); - if (!resultVectorType) - return rewriter.notifyMatchFailure(op, - "unexpected converted vselr result type"); - - Type intrinsicResultType = resultType; - if (auto floatType = dyn_cast(resultVectorType.getElementType()); - floatType && floatType.isF32()) { - intrinsicResultType = VectorType::get( - resultVectorType.getShape(), rewriter.getI32Type(), - resultVectorType.getScalableDims()); - } - if (Type carrierType = getLowpPayloadCarrierType( - op.getResult().getType(), rewriter.getContext())) - intrinsicResultType = carrierType; - - Type indexType = this->getTypeConverter()->convertType(op.getSrc1().getType()); - if (!indexType) - return rewriter.notifyMatchFailure(op, - "failed to convert vselr index type"); - - Value src0 = adaptor.getSrc0(); - Value src1 = adaptor.getSrc1(); - if (!src0 || !src1 || src1.getType() != indexType) - return rewriter.notifyMatchFailure(op, - "unexpected converted vselr operand types"); - - if (src0.getType() != intrinsicResultType) { - if (src0.getType() != resultType) - return rewriter.notifyMatchFailure(op, - "unexpected converted vselr source type"); - src0 = rewriter.create(op.getLoc(), intrinsicResultType, src0); - } - - auto funcType = rewriter.getFunctionType( - TypeRange{intrinsicResultType, indexType}, TypeRange{intrinsicResultType}); - auto call = rewriter.create( - op.getLoc(), *calleeName, TypeRange{intrinsicResultType}, - ValueRange{src0, src1}); - state.plannedDecls.push_back(PlannedDecl{calleeName->str(), funcType}); - - Value result = call.getResult(0); - if (intrinsicResultType != resultType) - result = rewriter.create(op.getLoc(), resultType, result); - rewriter.replaceOp(op, ValueRange{result}); - return success(); - } - -private: - LoweringState &state; -}; - -class LowerPnotOpPattern final : public OpConversionPattern { -public: - explicit LowerPnotOpPattern(TypeConverter &typeConverter, MLIRContext *context, - LoweringState &state) - : OpConversionPattern(typeConverter, context), state(state) {} - - LogicalResult - matchAndRewrite(pto::PnotOp op, pto::PnotOp::Adaptor adaptor, - ConversionPatternRewriter &rewriter) const override { - Type resultType = this->getTypeConverter()->convertType(op.getResult().getType()); - if (!resultType) - return rewriter.notifyMatchFailure(op, - "failed to convert pnot result type"); - - Value input = adaptor.getInput(); - Value mask = adaptor.getMask(); - if (!input || !mask || input.getType() != resultType || - mask.getType() != resultType) { - return rewriter.notifyMatchFailure( - op, "unexpected converted pnot operand types"); - } - - StringRef calleeName = getPredicateMaskCallee(op.getContext()); - auto call = rewriter.create(op.getLoc(), calleeName, - TypeRange{resultType}, - ValueRange{input, mask}); - state.plannedDecls.push_back( - PlannedDecl{calleeName.str(), call.getCalleeType()}); - rewriter.replaceOp(op, call.getResults()); - return success(); - } - -private: - LoweringState &state; -}; - -template -class LowerInterleaveOpPattern final - : public OpConversionPattern { -public: - explicit LowerInterleaveOpPattern(TypeConverter &typeConverter, - MLIRContext *context, LoweringState &state) - : OpConversionPattern(typeConverter, context), state(state) {} - - LogicalResult - matchAndRewrite(InterleaveOp op, typename InterleaveOp::Adaptor adaptor, - ConversionPatternRewriter &rewriter) const override { - StringRef stem = std::is_same_v ? "vintlv" : "vdintlv"; - FailureOr calleeName = - buildInterleaveCallee(op.getContext(), op.getLow().getType(), stem); - if (failed(calleeName)) - return rewriter.notifyMatchFailure(op, - "unsupported interleave VPTO signature"); - - Type lowType = this->getTypeConverter()->convertType(op.getLow().getType()); - Type highType = this->getTypeConverter()->convertType(op.getHigh().getType()); - if (!lowType || !highType || lowType != highType) { - return rewriter.notifyMatchFailure( - op, "failed to convert interleave result types"); - } - - Value lhs = adaptor.getLhs(); - Value rhs = adaptor.getRhs(); - if (!lhs || !rhs || lhs.getType() != lowType || rhs.getType() != lowType) { - return rewriter.notifyMatchFailure( - op, "unexpected converted interleave operand types"); - } - - auto funcType = rewriter.getFunctionType(TypeRange{lowType, lowType}, - TypeRange{lowType, highType}); - auto call = rewriter.create( - op.getLoc(), *calleeName, TypeRange{lowType, highType}, ValueRange{lhs, rhs}); - state.plannedDecls.push_back(PlannedDecl{calleeName->str(), funcType}); - rewriter.replaceOp(op, call.getResults()); - return success(); - } - -private: - LoweringState &state; -}; - -template -class LowerPredicatePackOpPattern final : public OpConversionPattern { -public: - explicit LowerPredicatePackOpPattern(TypeConverter &typeConverter, - MLIRContext *context, - LoweringState &state) - : OpConversionPattern(typeConverter, context), state(state) {} - - LogicalResult - matchAndRewrite(PackOp op, typename PackOp::Adaptor adaptor, - ConversionPatternRewriter &rewriter) const override { - Type resultType = - this->getTypeConverter()->convertType(op.getResult().getType()); - if (!resultType) - return rewriter.notifyMatchFailure( - op, "failed to convert predicate-pack result type"); - - auto part = parseHiLoPartImmediate(op.getPart()); - if (!part) - return rewriter.notifyMatchFailure( - op, "unsupported predicate-pack part immediate"); - - Value input = adaptor.getInput(); - if (!input || input.getType() != resultType) - return rewriter.notifyMatchFailure( - op, "unexpected converted predicate-pack operand type"); - - Value partValue = rewriter.create( - op.getLoc(), rewriter.getI32IntegerAttr(*part)); - StringRef calleeName = getPredicatePackCallee(op.getContext()); - auto funcType = rewriter.getFunctionType( - TypeRange{resultType, rewriter.getI32Type()}, TypeRange{resultType}); - auto call = rewriter.create( - op.getLoc(), calleeName, TypeRange{resultType}, ValueRange{input, partValue}); - state.plannedDecls.push_back(PlannedDecl{calleeName.str(), funcType}); - rewriter.replaceOp(op, call.getResults()); - return success(); - } - -private: - LoweringState &state; -}; - -template -class LowerUnpackOpPattern final : public OpConversionPattern { -public: - explicit LowerUnpackOpPattern(TypeConverter &typeConverter, - MLIRContext *context, LoweringState &state) - : OpConversionPattern(typeConverter, context), state(state) {} - - LogicalResult - matchAndRewrite(UnpackOp op, typename UnpackOp::Adaptor adaptor, - ConversionPatternRewriter &rewriter) const override { - StringRef stem = std::is_same_v ? "vsunpack" - : "vzunpack"; - FailureOr calleeName = buildUnpackCallee( - op.getContext(), op.getSrc().getType(), op.getResult().getType(), stem); - if (failed(calleeName)) - return rewriter.notifyMatchFailure(op, - "unsupported unpack VPTO signature"); - - Type srcType = this->getTypeConverter()->convertType(op.getSrc().getType()); - Type resultType = - this->getTypeConverter()->convertType(op.getResult().getType()); - if (!srcType || !resultType) - return rewriter.notifyMatchFailure(op, - "failed to convert unpack types"); - - Value src = adaptor.getSrc(); - if (!src || src.getType() != srcType) { - return rewriter.notifyMatchFailure( - op, "unexpected converted unpack source type"); - } - - Value part = castIntegerLikeTo(op, adaptor.getPart(), rewriter.getI32Type()); - if (!part) - return rewriter.notifyMatchFailure(op, "failed to materialize unpack part"); - - auto funcType = rewriter.getFunctionType(TypeRange{srcType, part.getType()}, - TypeRange{resultType}); - auto call = rewriter.create( - op.getLoc(), *calleeName, TypeRange{resultType}, ValueRange{src, part}); - state.plannedDecls.push_back(PlannedDecl{calleeName->str(), funcType}); - rewriter.replaceOp(op, call.getResults()); - return success(); - } - -private: - LoweringState &state; -}; - -class LowerVpackOpPattern final : public OpConversionPattern { -public: - explicit LowerVpackOpPattern(TypeConverter &typeConverter, - MLIRContext *context, LoweringState &state) - : OpConversionPattern(typeConverter, context), state(state) {} - - LogicalResult - matchAndRewrite(pto::VpackOp op, pto::VpackOp::Adaptor adaptor, - ConversionPatternRewriter &rewriter) const override { - FailureOr calleeName = - buildVpackCallee(op.getContext(), op.getSrc().getType(), - op.getResult().getType()); - if (failed(calleeName)) - return rewriter.notifyMatchFailure(op, "unsupported vpack VPTO signature"); - - Type srcType = this->getTypeConverter()->convertType(op.getSrc().getType()); - Type resultType = - this->getTypeConverter()->convertType(op.getResult().getType()); - if (!srcType || !resultType) - return rewriter.notifyMatchFailure(op, "failed to convert vpack types"); - - auto partImm = parseHiLoPartImmediate(op.getPart()); - if (!partImm) - return rewriter.notifyMatchFailure(op, "unsupported vpack part immediate"); - - Value src = adaptor.getSrc(); - if (!src || src.getType() != srcType) { - return rewriter.notifyMatchFailure( - op, "unexpected converted vpack source type"); - } - - Value part = getI32Constant(rewriter, op.getLoc(), *partImm); - auto funcType = rewriter.getFunctionType(TypeRange{srcType, part.getType()}, - TypeRange{resultType}); - auto call = rewriter.create( - op.getLoc(), *calleeName, TypeRange{resultType}, ValueRange{src, part}); - state.plannedDecls.push_back(PlannedDecl{calleeName->str(), funcType}); - rewriter.replaceOp(op, call.getResults()); - return success(); - } - -private: - LoweringState &state; -}; - -template -class LowerPredicateMaskBinaryOpPattern final - : public OpConversionPattern { -public: - explicit LowerPredicateMaskBinaryOpPattern(TypeConverter &typeConverter, - MLIRContext *context, - LoweringState &state) - : OpConversionPattern(typeConverter, context), - state(state) {} - - LogicalResult - matchAndRewrite(PredicateMaskOp op, typename PredicateMaskOp::Adaptor adaptor, - ConversionPatternRewriter &rewriter) const override { - Type resultType = - this->getTypeConverter()->convertType(op.getResult().getType()); - if (!resultType) - return rewriter.notifyMatchFailure( - op, "failed to convert predicate-mask result type"); - - Value src0 = adaptor.getSrc0(); - Value src1 = adaptor.getSrc1(); - Value mask = adaptor.getMask(); - if (!src0 || !src1 || !mask || src0.getType() != resultType || - src1.getType() != resultType || mask.getType() != resultType) { - return rewriter.notifyMatchFailure( - op, "unexpected converted predicate-mask operand types"); - } - - StringRef calleeName = getPredicateMaskCallee(op.getContext()); - auto call = rewriter.create(op.getLoc(), calleeName, - TypeRange{resultType}, - ValueRange{src0, src1, mask}); - state.plannedDecls.push_back( - PlannedDecl{calleeName.str(), call.getCalleeType()}); - rewriter.replaceOp(op, call.getResults()); - return success(); - } - -private: - LoweringState &state; -}; - -template -class LowerPredicatePairReorderOpPattern final - : public OpConversionPattern { -public: - explicit LowerPredicatePairReorderOpPattern(TypeConverter &typeConverter, - MLIRContext *context, - LoweringState &state) - : OpConversionPattern(typeConverter, context), state(state) {} - - LogicalResult - matchAndRewrite(ReorderOp op, typename ReorderOp::Adaptor adaptor, - ConversionPatternRewriter &rewriter) const override { - SmallVector resultTypes; - if (failed(this->getTypeConverter()->convertTypes(op->getResultTypes(), resultTypes))) - return rewriter.notifyMatchFailure( - op, "failed to convert predicate-pair-reorder result types"); - if (resultTypes.size() != 2 || resultTypes[0] != resultTypes[1]) - return rewriter.notifyMatchFailure( - op, "unexpected predicate-pair-reorder converted result types"); - - Value lhs = adaptor.getLhs(); - Value rhs = adaptor.getRhs(); - if (!lhs || !rhs || lhs.getType() != resultTypes[0] || - rhs.getType() != resultTypes[0]) { - return rewriter.notifyMatchFailure( - op, "unexpected converted predicate-pair-reorder operand types"); - } - - StringRef calleeName = - buildPredicatePairReorderCallee(op.getContext()); - auto call = rewriter.create(op.getLoc(), calleeName, resultTypes, - ValueRange{lhs, rhs}); - state.plannedDecls.push_back( - PlannedDecl{calleeName.str(), call.getCalleeType()}); - rewriter.replaceOp(op, call.getResults()); - return success(); - } - -private: - LoweringState &state; -}; - -template -class LowerCmpOpPattern final : public OpConversionPattern { -public: - explicit LowerCmpOpPattern(TypeConverter &typeConverter, MLIRContext *context, - LoweringState &state) - : OpConversionPattern(typeConverter, context), state(state) {} - - LogicalResult - matchAndRewrite(CmpOp op, typename CmpOp::Adaptor adaptor, - ConversionPatternRewriter &rewriter) const override { - constexpr bool isScalarCompare = std::is_same_v; - Type inputType = Type(); - if constexpr (isScalarCompare) - inputType = op.getSrc().getType(); - else - inputType = op.getSrc0().getType(); - FailureOr calleeName = - buildVcmpCallee(op.getContext(), inputType, op.getCmpMode(), - isScalarCompare); - if (failed(calleeName)) - return rewriter.notifyMatchFailure(op, - "unsupported compare VPTO signature"); - - Type resultType = this->getTypeConverter()->convertType(op.getResult().getType()); - Type maskType = - this->getTypeConverter()->convertType(op.getMask().getType()); - if (!resultType || !maskType) - return rewriter.notifyMatchFailure(op, - "failed to convert compare result type"); - if (resultType != maskType) - return rewriter.notifyMatchFailure(op, - "unexpected compare mask conversion"); - - SmallVector callArgs; - callArgs.append(adaptor.getOperands().begin(), adaptor.getOperands().end()); - if constexpr (isScalarCompare) { - if (callArgs.size() != 3 || !callArgs[0] || !callArgs[1] || !callArgs[2] || - callArgs[2].getType() != maskType) { - return rewriter.notifyMatchFailure( - op, "unexpected converted scalar-compare operand types"); - } - callArgs[1] = normalizeByteScalarOperandForCANN900VectorCall( - rewriter, op.getLoc(), callArgs[1], - cast(op.getSrc().getType()).getElementType()); - } else { - if (callArgs.size() != 3 || !callArgs[0] || !callArgs[1] || !callArgs[2] || - callArgs[0].getType() != callArgs[1].getType() || - callArgs[2].getType() != maskType) { - return rewriter.notifyMatchFailure( - op, "unexpected converted compare operand types"); - } - } - - auto call = rewriter.create(op.getLoc(), *calleeName, - TypeRange{resultType}, callArgs); - state.plannedDecls.push_back( - PlannedDecl{calleeName->str(), call.getCalleeType()}); - rewriter.replaceOp(op, call.getResults()); - return success(); - } - -private: - LoweringState &state; -}; - -template -class LowerPltOpPattern final : public OpConversionPattern { -public: - explicit LowerPltOpPattern(TypeConverter &typeConverter, MLIRContext *context, - LoweringState &state) - : OpConversionPattern(typeConverter, context), state(state) {} - - LogicalResult - matchAndRewrite(PltOp op, typename PltOp::Adaptor adaptor, - ConversionPatternRewriter &rewriter) const override { - Value laneCount = castIntegerLikeTo(op, adaptor.getScalar(), rewriter.getI32Type()); - if (!laneCount) - return rewriter.notifyMatchFailure(op, "failed to materialize plt lane count"); - - SmallVector resultTypes; - if (failed(this->getTypeConverter()->convertTypes(op->getResultTypes(), resultTypes))) - return rewriter.notifyMatchFailure(op, "failed to convert plt result types"); - - StringRef calleeName = buildPltCallee(op.getContext()); - auto funcType = rewriter.getFunctionType(TypeRange{rewriter.getI32Type()}, - resultTypes); - auto call = rewriter.create(op.getLoc(), calleeName, - resultTypes, ValueRange{laneCount}); - state.plannedDecls.push_back(PlannedDecl{calleeName.str(), funcType}); - rewriter.replaceOp(op, call.getResults()); - return success(); - } - -private: - LoweringState &state; -}; - -template -class LowerPltmOpPattern final : public OpConversionPattern { -public: - explicit LowerPltmOpPattern(TypeConverter &typeConverter, MLIRContext *context, - LoweringState &state) - : OpConversionPattern(typeConverter, context), state(state) {} - - LogicalResult - matchAndRewrite(PltmOp op, typename PltmOp::Adaptor adaptor, - ConversionPatternRewriter &rewriter) const override { - SmallVector resultTypes; - if (failed(this->getTypeConverter()->convertTypes(op->getResultTypes(), - resultTypes))) - return rewriter.notifyMatchFailure(op, "failed to convert pltm result type"); - - Value loop = adaptor.getLoop(); - Value bound = adaptor.getBound(); - if (!loop || !bound || !loop.getType().isInteger(16) || - !bound.getType().isInteger(32)) - return rewriter.notifyMatchFailure(op, - "unexpected converted pltm operand types"); - - StringRef calleeName = buildPltmCallee(op.getContext()); - auto funcType = rewriter.getFunctionType( - TypeRange{rewriter.getI16Type(), rewriter.getI32Type()}, resultTypes); - auto call = rewriter.create(op.getLoc(), calleeName, - resultTypes, ValueRange{loop, bound}); - state.plannedDecls.push_back(PlannedDecl{calleeName.str(), funcType}); - rewriter.replaceOp(op, call.getResults()); - return success(); - } - -private: - LoweringState &state; -}; - -template -class LowerPsetOpPattern final : public OpConversionPattern { -public: - explicit LowerPsetOpPattern(TypeConverter &typeConverter, MLIRContext *context, - LoweringState &state) - : OpConversionPattern(typeConverter, context), state(state) {} - - LogicalResult - matchAndRewrite(PsetOp op, typename PsetOp::Adaptor adaptor, - ConversionPatternRewriter &rewriter) const override { - (void)adaptor; - auto pattern = parsePredicatePatternImmediate(op.getPattern()); - if (!pattern) - return rewriter.notifyMatchFailure(op, "unsupported pset pattern"); - - SmallVector resultTypes; - if (failed(this->getTypeConverter()->convertTypes(op->getResultTypes(), resultTypes))) - return rewriter.notifyMatchFailure(op, "failed to convert pset result types"); - - if (isMaskOnlyUsedByOnePointStores(op.getResult())) { - auto undef = rewriter.create(op.getLoc(), resultTypes.front()); - rewriter.replaceOp(op, undef.getResult()); - return success(); - } - - StringRef calleeName = buildPsetCallee(op.getContext()); - Value patternValue = rewriter.create( - op.getLoc(), rewriter.getI32IntegerAttr(*pattern)); - auto funcType = rewriter.getFunctionType(TypeRange{rewriter.getI32Type()}, - resultTypes); - auto call = rewriter.create(op.getLoc(), calleeName, - resultTypes, ValueRange{patternValue}); - state.plannedDecls.push_back(PlannedDecl{calleeName.str(), funcType}); - rewriter.replaceOp(op, call.getResults()); - return success(); - } - -private: - LoweringState &state; -}; - -template -class LowerPgeOpPattern final : public OpConversionPattern { -public: - explicit LowerPgeOpPattern(TypeConverter &typeConverter, MLIRContext *context, - LoweringState &state) - : OpConversionPattern(typeConverter, context), state(state) {} - - LogicalResult - matchAndRewrite(PgeOp op, typename PgeOp::Adaptor adaptor, - ConversionPatternRewriter &rewriter) const override { - (void)adaptor; - auto pattern = parsePredicatePatternImmediate(op.getPattern()); - if (!pattern) - return rewriter.notifyMatchFailure(op, "unsupported pge pattern"); - - SmallVector resultTypes; - if (failed(this->getTypeConverter()->convertTypes(op->getResultTypes(), resultTypes))) - return rewriter.notifyMatchFailure(op, "failed to convert pge result types"); - - if (isMaskOnlyUsedByOnePointStores(op.getResult())) { - auto undef = rewriter.create(op.getLoc(), resultTypes.front()); - rewriter.replaceOp(op, undef.getResult()); - return success(); - } - - StringRef calleeName = buildPgeCallee(op.getContext()); - Value patternValue = rewriter.create( - op.getLoc(), rewriter.getI32IntegerAttr(*pattern)); - Value zero = rewriter.create(op.getLoc(), - rewriter.getI32IntegerAttr(0)); - auto funcType = rewriter.getFunctionType( - TypeRange{rewriter.getI32Type(), rewriter.getI32Type()}, resultTypes); - auto call = - rewriter.create(op.getLoc(), calleeName, resultTypes, - ValueRange{patternValue, zero}); - state.plannedDecls.push_back(PlannedDecl{calleeName.str(), funcType}); - rewriter.replaceOp(op, call.getResults()); - return success(); - } - -private: - LoweringState &state; -}; - -class LowerVldsOpPattern final : public OpConversionPattern { -public: - explicit LowerVldsOpPattern(TypeConverter &typeConverter, MLIRContext *context, - LoweringState &state) - : OpConversionPattern(typeConverter, context), state(state) {} - - LogicalResult - matchAndRewrite(pto::VldsOp op, OpAdaptor adaptor, - ConversionPatternRewriter &rewriter) const override { - Type ptoResultType = op.getResult().getType(); - Type elementType = getElementTypeFromVectorLike(ptoResultType); - if (!elementType) - return rewriter.notifyMatchFailure(op, "unsupported vlds element type"); - bool usePostIntrinsic = static_cast(op.getUpdatedBase()); - auto loweredOffset = lowerVPTOElementOffsetForIntrinsic( - op, adaptor.getSource(), adaptor.getOffset(), elementType, - usePostIntrinsic, rewriter); - auto dist = - parseLoadDistImmediate(op.getDist().value_or("NORM"), elementType); - bool invalidAddress = failed(loweredOffset) || !dist; - if (invalidAddress) { - return rewriter.notifyMatchFailure(op, "failed to materialize vlds operands"); - } - - SmallVector resultTypes; - if (failed(this->getTypeConverter()->convertTypes(op->getResultTypes(), resultTypes))) - return rewriter.notifyMatchFailure(op, "failed to convert vlds result types"); - - if (usePostIntrinsic) { - if (resultTypes.size() != 2 || resultTypes[1] != adaptor.getSource().getType()) - return rewriter.notifyMatchFailure(op, - "unsupported vlds post-update results"); - } else if (resultTypes.size() != 1) { - return rewriter.notifyMatchFailure(op, "unsupported vlds result count"); - } - Type callValueType = getPayloadABIType( - ptoResultType, resultTypes[0], rewriter.getContext()); - SmallVector callResultTypes; - callResultTypes.push_back(callValueType); - if (usePostIntrinsic) - callResultTypes.push_back(resultTypes[1]); - - FailureOr calleeName = - usePostIntrinsic ? buildVldsPostCallee(op.getContext(), ptoResultType) - : buildVldsCallee(op.getContext(), ptoResultType); - if (failed(calleeName)) - return rewriter.notifyMatchFailure(op, "unsupported vlds signature"); - - Value distValue = getI32Constant(rewriter, op.getLoc(), *dist); - Value postValue = getI32Constant(rewriter, op.getLoc(), usePostIntrinsic ? 1 : 0); - SmallVector args{loweredOffset->base, - loweredOffset->intrinsicOffset, distValue, - postValue}; - auto funcType = rewriter.getFunctionType( - TypeRange{loweredOffset->base.getType(), - loweredOffset->intrinsicOffset.getType(), - distValue.getType(), postValue.getType()}, - callResultTypes); - auto call = rewriter.create(op.getLoc(), *calleeName, - callResultTypes, args); - state.plannedDecls.push_back(PlannedDecl{calleeName->str(), funcType}); - Value loaded = castFromPayloadABI( - op.getLoc(), call.getResult(0), ptoResultType, resultTypes[0], - rewriter); - if (usePostIntrinsic) { - Value updatedBase = loweredOffset->updatedBase - ? loweredOffset->updatedBase - : call.getResult(1); - rewriter.replaceOp(op, ValueRange{loaded, updatedBase}); - } else { - rewriter.replaceOp(op, ValueRange{loaded}); - } - return success(); - } - -private: - LoweringState &state; -}; - -class LowerVldsx2OpPattern final : public OpConversionPattern { -public: - explicit LowerVldsx2OpPattern(TypeConverter &typeConverter, - MLIRContext *context, LoweringState &state) - : OpConversionPattern(typeConverter, context), - state(state) {} - - LogicalResult - matchAndRewrite(pto::Vldsx2Op op, pto::Vldsx2Op::Adaptor adaptor, - ConversionPatternRewriter &rewriter) const override { - Type elementType = getElementTypeFromVectorLike(op.getLow().getType()); - if (!elementType) - return rewriter.notifyMatchFailure(op, "unsupported vldsx2 element type"); - - bool usePostIntrinsic = op.getUpdatedBase() != nullptr; - auto loweredOffset = lowerVPTOElementOffsetForIntrinsic( - op, adaptor.getSource(), adaptor.getOffset(), elementType, - usePostIntrinsic, rewriter); - auto dist = parseLoadX2DistImmediate(op.getDist(), elementType); - bool invalidAddress = failed(loweredOffset) || !dist; - if (invalidAddress) { - return rewriter.notifyMatchFailure(op, - "failed to materialize vldsx2 operands"); - } - - SmallVector resultTypes; - if (failed(this->getTypeConverter()->convertTypes(op->getResultTypes(), - resultTypes)) || - resultTypes.size() != (usePostIntrinsic ? 3U : 2U)) { - return rewriter.notifyMatchFailure(op, - "failed to convert vldsx2 result types"); - } - Type lowCallType = getPayloadABIType( - op.getLow().getType(), resultTypes[0], rewriter.getContext()); - Type highCallType = getPayloadABIType( - op.getHigh().getType(), resultTypes[1], rewriter.getContext()); - SmallVector callResultTypes{lowCallType, highCallType}; - if (usePostIntrinsic) - callResultTypes.push_back(resultTypes[2]); - - FailureOr calleeName = - buildVldsx2Callee(op.getContext(), op.getLow().getType(), - usePostIntrinsic); - if (failed(calleeName)) - return rewriter.notifyMatchFailure(op, "unsupported vldsx2 signature"); - - Value distValue = getI32Constant(rewriter, op.getLoc(), *dist); - Value postValue = - getI32Constant(rewriter, op.getLoc(), usePostIntrinsic ? 1 : 0); - SmallVector args{loweredOffset->base, - loweredOffset->intrinsicOffset, distValue, - postValue}; - auto funcType = rewriter.getFunctionType( - TypeRange{loweredOffset->base.getType(), - loweredOffset->intrinsicOffset.getType(), - distValue.getType(), postValue.getType()}, - callResultTypes); - auto call = rewriter.create(op.getLoc(), *calleeName, - callResultTypes, args); - state.plannedDecls.push_back(PlannedDecl{calleeName->str(), funcType}); - Value low = castFromPayloadABI( - op.getLoc(), call.getResult(0), op.getLow().getType(), resultTypes[0], - rewriter); - Value high = castFromPayloadABI( - op.getLoc(), call.getResult(1), op.getHigh().getType(), resultTypes[1], - rewriter); - if (usePostIntrinsic) { - Value updatedBase = loweredOffset->updatedBase - ? loweredOffset->updatedBase - : call.getResult(2); - rewriter.replaceOp(op, ValueRange{low, high, updatedBase}); - } else { - rewriter.replaceOp(op, ValueRange{low, high}); - } - return success(); - } - -private: - LoweringState &state; -}; - -class LowerVsldbOpPattern final : public OpConversionPattern { -public: - explicit LowerVsldbOpPattern(TypeConverter &typeConverter, - MLIRContext *context, LoweringState &state) - : OpConversionPattern(typeConverter, context), state(state) {} - - LogicalResult - matchAndRewrite(pto::VsldbOp op, pto::VsldbOp::Adaptor adaptor, - ConversionPatternRewriter &rewriter) const override { - auto basePtr = dyn_cast(adaptor.getSource().getType()); - Value packedStride = - packBlockRepeatStride(op, adaptor.getBlockStride(), adaptor.getRepeatStride()); - if (!basePtr || !packedStride) - return rewriter.notifyMatchFailure(op, "failed to materialize vsldb operands"); - - bool usePostIntrinsic = op.getUpdatedBase() != nullptr; - SmallVector resultTypes; - if (failed(this->getTypeConverter()->convertTypes(op->getResultTypes(), - resultTypes)) || - resultTypes.size() != (usePostIntrinsic ? 2U : 1U)) - return rewriter.notifyMatchFailure(op, "failed to convert vsldb result type"); - - Type callResultType = getPayloadABIType( - op.getResult().getType(), resultTypes[0], rewriter.getContext()); - SmallVector callResultTypes{callResultType}; - if (usePostIntrinsic) - callResultTypes.push_back(resultTypes[1]); - - FailureOr calleeName = - buildVsldbCallee(op.getContext(), op.getResult().getType(), - usePostIntrinsic); - if (failed(calleeName)) - return rewriter.notifyMatchFailure(op, "unsupported vsldb signature"); - Value postValue = - getI32Constant(rewriter, op.getLoc(), usePostIntrinsic ? 1 : 0); - SmallVector args{adaptor.getSource(), packedStride, postValue, - adaptor.getMask()}; - auto funcType = rewriter.getFunctionType( - TypeRange{adaptor.getSource().getType(), packedStride.getType(), - postValue.getType(), adaptor.getMask().getType()}, - callResultTypes); - auto call = rewriter.create(op.getLoc(), *calleeName, - callResultTypes, args); - state.plannedDecls.push_back(PlannedDecl{calleeName->str(), funcType}); - Value result = castFromPayloadABI( - op.getLoc(), call.getResult(0), op.getResult().getType(), resultTypes[0], - rewriter); - if (usePostIntrinsic) - rewriter.replaceOp(op, ValueRange{result, call.getResult(1)}); - else - rewriter.replaceOp(op, ValueRange{result}); - return success(); - } - -private: - LoweringState &state; -}; - -class LowerInitAlignOpPattern final - : public OpConversionPattern { -public: - explicit LowerInitAlignOpPattern(TypeConverter &typeConverter, - MLIRContext *context, LoweringState &state) - : OpConversionPattern(typeConverter, context), - state(state) {} - - LogicalResult - matchAndRewrite(pto::InitAlignOp op, pto::InitAlignOp::Adaptor adaptor, - ConversionPatternRewriter &rewriter) const override { - (void)adaptor; - Type resultType = this->getTypeConverter()->convertType(op.getResult().getType()); - if (!resultType) - return rewriter.notifyMatchFailure(op, "failed to convert init_align result type"); - - StringRef calleeName = buildInitAlignCallee(op.getContext()); - auto funcType = rewriter.getFunctionType(TypeRange{}, TypeRange{resultType}); - auto call = - rewriter.create(op.getLoc(), calleeName, TypeRange{resultType}); - state.plannedDecls.push_back(PlannedDecl{calleeName.str(), funcType}); - rewriter.replaceOp(op, call.getResults()); - return success(); - } - -private: - LoweringState &state; -}; - -class LowerVldasOpPattern final : public OpConversionPattern { -public: - explicit LowerVldasOpPattern(TypeConverter &typeConverter, - MLIRContext *context, LoweringState &state) - : OpConversionPattern(typeConverter, context), state(state) {} - - LogicalResult - matchAndRewrite(pto::VldasOp op, pto::VldasOp::Adaptor adaptor, - ConversionPatternRewriter &rewriter) const override { - auto sourceType = dyn_cast(adaptor.getSource().getType()); - Type resultType = this->getTypeConverter()->convertType(op.getResult().getType()); - if (!sourceType || !resultType) - return rewriter.notifyMatchFailure(op, - "expected converted vldas operand/result types"); - - StringRef calleeName = buildVldasCallee(op.getContext()); - auto funcType = - rewriter.getFunctionType(TypeRange{adaptor.getSource().getType()}, - TypeRange{resultType}); - auto call = rewriter.create(op.getLoc(), calleeName, - TypeRange{resultType}, - ValueRange{adaptor.getSource()}); - state.plannedDecls.push_back(PlannedDecl{calleeName.str(), funcType}); - rewriter.replaceOp(op, call.getResults()); - return success(); - } - -private: - LoweringState &state; -}; - -class LowerVldusOpPattern final : public OpConversionPattern { -public: - explicit LowerVldusOpPattern(TypeConverter &typeConverter, - MLIRContext *context, LoweringState &state) - : OpConversionPattern(typeConverter, context), state(state) {} - - LogicalResult - matchAndRewrite(pto::VldusOp op, pto::VldusOp::Adaptor adaptor, - ConversionPatternRewriter &rewriter) const override { - auto sourceType = dyn_cast(adaptor.getSource().getType()); - SmallVector resultTypes; - bool usePostIntrinsic = static_cast(op.getUpdatedBase()); - if (!sourceType || - failed(this->getTypeConverter()->convertTypes(op->getResultTypes(), resultTypes)) || - resultTypes.size() != (usePostIntrinsic ? 3U : 2U) || - adaptor.getAlign().getType() != resultTypes[1] || - (usePostIntrinsic && resultTypes[2] != adaptor.getSource().getType())) { - return rewriter.notifyMatchFailure(op, - "expected converted vldus operand/result types"); - } - - FailureOr calleeName = - usePostIntrinsic - ? buildVldusPostCallee(op.getContext(), op.getResult().getType()) - : buildVldusCallee(op.getContext(), op.getResult().getType()); - if (failed(calleeName)) - return rewriter.notifyMatchFailure(op, "unsupported vldus signature"); - - Type callValueType = getPayloadABIType( - op.getResult().getType(), resultTypes[0], rewriter.getContext()); - SmallVector intrinsicResultTypes{callValueType, resultTypes[1]}; - // The installed no-post A5 vldus intrinsic returns an extra hidden base ptr. - intrinsicResultTypes.push_back(adaptor.getSource().getType()); - - SmallVector args{adaptor.getSource(), adaptor.getAlign()}; - Value explicitUpdatedBase; - if (usePostIntrinsic) { - Type elementType = getElementTypeFromVectorLike(op.getResult().getType()); - auto loweredIncrement = lowerVPTOElementOffsetForIntrinsic( - op, adaptor.getSource(), adaptor.getIncrement(), elementType, - /*isPostUpdate=*/true, rewriter); - if (failed(loweredIncrement)) { - return rewriter.notifyMatchFailure(op, - "failed to convert vldus increment"); - } - args.front() = loweredIncrement->base; - args.push_back(loweredIncrement->intrinsicOffset); - explicitUpdatedBase = loweredIncrement->updatedBase; - } - SmallVector argTypes; - for (Value arg : args) - argTypes.push_back(arg.getType()); - auto funcType = rewriter.getFunctionType(argTypes, intrinsicResultTypes); - auto call = rewriter.create( - op.getLoc(), *calleeName, intrinsicResultTypes, args); - state.plannedDecls.push_back(PlannedDecl{calleeName->str(), funcType}); - Value loaded = castFromPayloadABI( - op.getLoc(), call.getResult(0), op.getResult().getType(), - resultTypes[0], rewriter); - SmallVector replacements{loaded, call.getResult(1)}; - if (usePostIntrinsic) { - replacements.push_back(explicitUpdatedBase ? explicitUpdatedBase - : call.getResult(2)); - } - rewriter.replaceOp(op, replacements); - return success(); - } - -private: - LoweringState &state; -}; - -class LowerSprclrOpPattern final : public OpConversionPattern { -public: - explicit LowerSprclrOpPattern(TypeConverter &typeConverter, MLIRContext *context, - LoweringState &state) - : OpConversionPattern(typeConverter, context), state(state) {} - - LogicalResult - matchAndRewrite(pto::SprclrOp op, pto::SprclrOp::Adaptor adaptor, - ConversionPatternRewriter &rewriter) const override { - (void)adaptor; - auto spr = parseSprImmediate(op.getSpr()); - if (!spr) - return rewriter.notifyMatchFailure(op, "unsupported sprclr target"); - - StringRef calleeName = buildSprclrCallee(op.getContext()); - Value sprValue = rewriter.create( - op.getLoc(), rewriter.getI16IntegerAttr(*spr)); - auto funcType = rewriter.getFunctionType(TypeRange{sprValue.getType()}, TypeRange{}); - rewriter.create(op.getLoc(), calleeName, TypeRange{}, ValueRange{sprValue}); - state.plannedDecls.push_back(PlannedDecl{calleeName.str(), funcType}); - rewriter.eraseOp(op); - return success(); - } - -private: - LoweringState &state; -}; - -template -class LowerSprStoreOpPattern final : public OpConversionPattern { -public: - explicit LowerSprStoreOpPattern(TypeConverter &typeConverter, - MLIRContext *context, LoweringState &state) - : OpConversionPattern(typeConverter, context), state(state) {} - - LogicalResult - matchAndRewrite(SprStoreOp op, typename SprStoreOp::Adaptor adaptor, - ConversionPatternRewriter &rewriter) const override { - auto spr = parseSprImmediate(op.getSpr()); - if (!spr) - return rewriter.notifyMatchFailure(op, "unsupported spr store target"); - auto destType = - dyn_cast(adaptor.getDestination().getType()); - if (!destType || !adaptor.getOffset().getType().isInteger(32)) - return rewriter.notifyMatchFailure(op, - "expected converted spr store operands"); - - bool usePostIntrinsic = op.getUpdatedBase() != nullptr; - SmallVector resultTypes; - if (failed(this->getTypeConverter()->convertTypes(op->getResultTypes(), - resultTypes)) || - resultTypes.size() != (usePostIntrinsic ? 1U : 0U)) - return rewriter.notifyMatchFailure( - op, "failed to convert spr store result types"); - - StringRef calleeName = - buildSprStoreCallee(op.getContext(), usePostIntrinsic); - Value sprValue = rewriter.create( - op.getLoc(), rewriter.getI16IntegerAttr(*spr)); - Value postValue = rewriter.create( - op.getLoc(), rewriter.getI32IntegerAttr(usePostIntrinsic ? 1 : 0)); - SmallVector args{sprValue, adaptor.getDestination(), - adaptor.getOffset(), postValue}; - auto funcType = rewriter.getFunctionType( - TypeRange{sprValue.getType(), adaptor.getDestination().getType(), - adaptor.getOffset().getType(), postValue.getType()}, - resultTypes); - auto call = rewriter.create(op.getLoc(), calleeName, - resultTypes, args); - state.plannedDecls.push_back(PlannedDecl{calleeName.str(), funcType}); - if (usePostIntrinsic) - rewriter.replaceOp(op, call.getResults()); - else - rewriter.eraseOp(op); - return success(); - } - -private: - LoweringState &state; -}; - -class LowerVstsOpPattern final : public OpConversionPattern { -public: - explicit LowerVstsOpPattern(TypeConverter &typeConverter, MLIRContext *context, - LoweringState &state) - : OpConversionPattern(typeConverter, context), state(state) {} - - LogicalResult - matchAndRewrite(pto::VstsOp op, OpAdaptor adaptor, - ConversionPatternRewriter &rewriter) const override { - Type elementType = getElementTypeFromVectorLike(op.getValue().getType()); - if (!elementType) - return rewriter.notifyMatchFailure(op, "unsupported vsts element type"); - Type offsetElementType = elementType; - if (auto ptrType = dyn_cast(op.getDestination().getType())) - offsetElementType = ptrType.getElementType(); - else if (auto memrefType = dyn_cast(op.getDestination().getType())) - offsetElementType = memrefType.getElementType(); - bool usePostIntrinsic = static_cast(op.getUpdatedBase()); - auto loweredOffset = lowerVPTOElementOffsetForIntrinsic( - op, adaptor.getDestination(), adaptor.getOffset(), offsetElementType, - usePostIntrinsic, rewriter); - auto dist = - parseStoreDistImmediate(op.getDist().value_or(""), elementType); - bool invalidAddress = failed(loweredOffset) || !dist; - if (invalidAddress) { - return rewriter.notifyMatchFailure(op, "failed to materialize vsts operands"); - } - - FailureOr calleeName = - op.getUpdatedBase() - ? buildVstsPostCallee(op.getContext(), op.getValue().getType()) - : buildVstsCallee(op.getContext(), op.getValue().getType()); - if (failed(calleeName)) - return rewriter.notifyMatchFailure(op, "unsupported vsts signature"); - - SmallVector resultTypes; - if (failed(this->getTypeConverter()->convertTypes(op->getResultTypes(), - resultTypes))) - return rewriter.notifyMatchFailure(op, "failed to convert vsts result types"); - if (usePostIntrinsic) { - if (resultTypes.size() != 1 || - resultTypes[0] != adaptor.getDestination().getType()) - return rewriter.notifyMatchFailure(op, - "unsupported vsts post-update result"); - } else if (!resultTypes.empty()) { - return rewriter.notifyMatchFailure(op, "unsupported vsts result count"); - } - - Value distValue = rewriter.create( - op.getLoc(), rewriter.getI32IntegerAttr(*dist)); - Value zero = rewriter.create(op.getLoc(), - rewriter.getI32IntegerAttr( - usePostIntrinsic ? 1 : 0)); - Value value = castToPayloadABI( - op.getLoc(), adaptor.getValue(), op.getValue().getType(), rewriter); - Value mask = adaptor.getMask(); - // The 1PT store forms keep a mask operand in the LLVM ABI, but the - // hardware ignores it. Do not materialize a pset/pge mask solely for - // this dead operand; an LLVM undef is sufficient at this boundary. - StringRef distToken = op.getDist().value_or(""); - if (isOnePointStoreDist(distToken)) - mask = rewriter.create(op.getLoc(), mask.getType()); - SmallVector args{value, loweredOffset->base, - loweredOffset->intrinsicOffset, distValue, zero, - mask}; - auto funcType = rewriter.getFunctionType( - TypeRange{value.getType(), loweredOffset->base.getType(), - rewriter.getI32Type(), rewriter.getI32Type(), - rewriter.getI32Type(), mask.getType()}, - resultTypes); - auto call = rewriter.create(op.getLoc(), *calleeName, - resultTypes, args); - state.plannedDecls.push_back(PlannedDecl{calleeName->str(), funcType}); - if (usePostIntrinsic) { - Value updatedBase = loweredOffset->updatedBase - ? loweredOffset->updatedBase - : call.getResult(0); - rewriter.replaceOp(op, updatedBase); - } else { - rewriter.eraseOp(op); - } - return success(); - } - -private: - LoweringState &state; -}; - -class LowerVsstbOpPattern final : public OpConversionPattern { -public: - explicit LowerVsstbOpPattern(TypeConverter &typeConverter, - MLIRContext *context, LoweringState &state) - : OpConversionPattern(typeConverter, context), state(state) {} - - LogicalResult - matchAndRewrite(pto::VsstbOp op, pto::VsstbOp::Adaptor adaptor, - ConversionPatternRewriter &rewriter) const override { - auto basePtr = - dyn_cast(adaptor.getDestination().getType()); - Value packedStride = - packBlockRepeatStride(op, adaptor.getBlockStride(), adaptor.getRepeatStride()); - if (!basePtr || !packedStride) - return rewriter.notifyMatchFailure(op, "failed to materialize vsstb operands"); - - SmallVector resultTypes; - if (failed(this->getTypeConverter()->convertTypes(op->getResultTypes(), - resultTypes))) - return rewriter.notifyMatchFailure(op, - "failed to convert vsstb result types"); - bool usePostIntrinsic = static_cast(op.getUpdatedBase()); - if (usePostIntrinsic) { - if (resultTypes.size() != 1 || - resultTypes[0] != adaptor.getDestination().getType()) - return rewriter.notifyMatchFailure( - op, "unsupported vsstb post-update result"); - } else if (!resultTypes.empty()) { - return rewriter.notifyMatchFailure(op, "unsupported vsstb result count"); - } - - FailureOr calleeName = buildVsstbCallee( - op.getContext(), op.getValue().getType(), usePostIntrinsic); - if (failed(calleeName)) - return rewriter.notifyMatchFailure(op, "unsupported vsstb signature"); - Value zeroValue = getI32Constant(rewriter, op.getLoc(), usePostIntrinsic ? 1 : 0); - Value value = castToPayloadABI( - op.getLoc(), adaptor.getValue(), op.getValue().getType(), rewriter); - SmallVector args{value, adaptor.getDestination(), - packedStride, zeroValue, adaptor.getMask()}; - auto funcType = rewriter.getFunctionType( - TypeRange{value.getType(), adaptor.getDestination().getType(), - packedStride.getType(), zeroValue.getType(), - adaptor.getMask().getType()}, - resultTypes); - auto call = rewriter.create(op.getLoc(), *calleeName, - resultTypes, args); - state.plannedDecls.push_back(PlannedDecl{calleeName->str(), funcType}); - if (usePostIntrinsic) - rewriter.replaceOp(op, call.getResults()); - else - rewriter.eraseOp(op); - return success(); - } - -private: - LoweringState &state; -}; - -class LowerVstsx2OpPattern final : public OpConversionPattern { -public: - explicit LowerVstsx2OpPattern(TypeConverter &typeConverter, - MLIRContext *context, LoweringState &state) - : OpConversionPattern(typeConverter, context), - state(state) {} - - LogicalResult - matchAndRewrite(pto::Vstsx2Op op, pto::Vstsx2Op::Adaptor adaptor, - ConversionPatternRewriter &rewriter) const override { - Type elementType = getElementTypeFromVectorLike(op.getLow().getType()); - if (!elementType) - return rewriter.notifyMatchFailure(op, "unsupported vstsx2 element type"); - - auto loweredOffset = lowerVPTOElementOffsetForIntrinsic( - op, adaptor.getDestination(), adaptor.getOffset(), elementType, - /*isPostUpdate=*/false, rewriter); - auto dist = parseStoreX2DistImmediate(op.getDist(), elementType); - bool invalidAddress = failed(loweredOffset) || !dist; - if (invalidAddress) { - return rewriter.notifyMatchFailure(op, - "failed to materialize vstsx2 operands"); - } - - FailureOr calleeName = - buildVstsx2Callee(op.getContext(), op.getLow().getType()); - if (failed(calleeName)) - return rewriter.notifyMatchFailure(op, "unsupported vstsx2 signature"); - - Value distValue = getI32Constant(rewriter, op.getLoc(), *dist); - Value zeroValue = getI32Constant(rewriter, op.getLoc(), 0); - Value low = castToPayloadABI( - op.getLoc(), adaptor.getLow(), op.getLow().getType(), rewriter); - Value high = castToPayloadABI( - op.getLoc(), adaptor.getHigh(), op.getHigh().getType(), rewriter); - SmallVector args{low, high, loweredOffset->base, - loweredOffset->intrinsicOffset, distValue, - zeroValue, adaptor.getMask()}; - auto funcType = rewriter.getFunctionType( - TypeRange{low.getType(), high.getType(), - loweredOffset->base.getType(), - loweredOffset->intrinsicOffset.getType(), - distValue.getType(), zeroValue.getType(), - adaptor.getMask().getType()}, - TypeRange{}); - rewriter.create(op.getLoc(), *calleeName, TypeRange{}, args); - state.plannedDecls.push_back(PlannedDecl{calleeName->str(), funcType}); - rewriter.eraseOp(op); - return success(); - } - -private: - LoweringState &state; -}; - -class LowerPstuOpPattern final : public OpConversionPattern { -public: - explicit LowerPstuOpPattern(TypeConverter &typeConverter, MLIRContext *context, - LoweringState &state) - : OpConversionPattern(typeConverter, context), state(state) {} - - LogicalResult - matchAndRewrite(pto::PstuOp op, pto::PstuOp::Adaptor adaptor, - ConversionPatternRewriter &rewriter) const override { - FailureOr calleeName = buildPstuCallee(op.getContext(), op); - if (failed(calleeName)) - return rewriter.notifyMatchFailure(op, "unsupported pstu signature"); - - SmallVector resultTypes; - if (failed(this->getTypeConverter()->convertTypes(op->getResultTypes(), resultTypes))) - return rewriter.notifyMatchFailure(op, "failed to convert pstu result types"); - if (resultTypes.size() != 2) - return rewriter.notifyMatchFailure(op, "unexpected converted pstu result arity"); - - auto baseType = dyn_cast(adaptor.getBase().getType()); - if (!baseType || adaptor.getAlignIn().getType() != resultTypes[0] || - adaptor.getBase().getType() != resultTypes[1]) { - return rewriter.notifyMatchFailure(op, - "unexpected converted pstu operand/result types"); - } - - SmallVector args{adaptor.getValue(), adaptor.getBase(), adaptor.getAlignIn()}; - auto funcType = rewriter.getFunctionType( - TypeRange{adaptor.getValue().getType(), adaptor.getBase().getType(), - adaptor.getAlignIn().getType()}, - resultTypes); - auto call = rewriter.create(op.getLoc(), *calleeName, resultTypes, - args); - state.plannedDecls.push_back(PlannedDecl{calleeName->str(), funcType}); - rewriter.replaceOp(op, call.getResults()); - return success(); - } - -private: - LoweringState &state; -}; - -class LowerVstusOpPattern final : public OpConversionPattern { -public: - explicit LowerVstusOpPattern(TypeConverter &typeConverter, MLIRContext *context, - LoweringState &state) - : OpConversionPattern(typeConverter, context), state(state) {} - - LogicalResult - matchAndRewrite(pto::VstusOp op, pto::VstusOp::Adaptor adaptor, - ConversionPatternRewriter &rewriter) const override { - Type elementType = getElementTypeFromVectorLike(op.getValue().getType()); - if (!elementType) - return rewriter.notifyMatchFailure(op, "unsupported vstus element type"); - - bool usePostIntrinsic = static_cast(op.getBaseOut()); - auto loweredOffset = lowerVPTOElementOffsetForIntrinsic( - op, adaptor.getBase(), adaptor.getOffset(), elementType, - usePostIntrinsic, rewriter); - if (failed(loweredOffset)) { - return rewriter.notifyMatchFailure(op, "failed to convert vstus offset"); - } - - SmallVector resultTypes; - if (failed(this->getTypeConverter()->convertTypes(op->getResultTypes(), - resultTypes))) - return rewriter.notifyMatchFailure(op, - "failed to convert vstus result types"); - auto baseType = dyn_cast(adaptor.getBase().getType()); - if (!baseType || resultTypes.size() != (usePostIntrinsic ? 2U : 1U) || - adaptor.getAlignIn().getType() != resultTypes[0] || - (usePostIntrinsic && resultTypes[1] != adaptor.getBase().getType())) { - return rewriter.notifyMatchFailure(op, - "unexpected converted vstus operand/result types"); - } - - FailureOr calleeName = - buildVstusCallee(op.getContext(), op.getValue().getType()); - if (usePostIntrinsic) - calleeName = - buildVstusPostCallee(op.getContext(), op.getValue().getType()); - if (failed(calleeName)) - return rewriter.notifyMatchFailure(op, "unsupported vstus signature"); - Value value = castToPayloadABI( - op.getLoc(), adaptor.getValue(), op.getValue().getType(), rewriter); - SmallVector args{value, loweredOffset->base, - loweredOffset->intrinsicOffset, - adaptor.getAlignIn()}; - auto funcType = rewriter.getFunctionType( - TypeRange{value.getType(), loweredOffset->base.getType(), - loweredOffset->intrinsicOffset.getType(), - adaptor.getAlignIn().getType()}, - resultTypes); - auto call = rewriter.create(op.getLoc(), *calleeName, - resultTypes, args); - state.plannedDecls.push_back(PlannedDecl{calleeName->str(), funcType}); - if (usePostIntrinsic && loweredOffset->updatedBase) { - rewriter.replaceOp( - op, ValueRange{call.getResult(0), loweredOffset->updatedBase}); - } else { - rewriter.replaceOp(op, call.getResults()); - } - return success(); - } - -private: - LoweringState &state; -}; - -class LowerVsturOpPattern final : public OpConversionPattern { -public: - explicit LowerVsturOpPattern(TypeConverter &typeConverter, MLIRContext *context, - LoweringState &state) - : OpConversionPattern(typeConverter, context), state(state) {} - - LogicalResult - matchAndRewrite(pto::VsturOp op, pto::VsturOp::Adaptor adaptor, - ConversionPatternRewriter &rewriter) const override { - auto postMode = parsePostModeImmediate(op.getMode()); - if (!postMode) - return rewriter.notifyMatchFailure(op, "unsupported vstur mode immediate"); - - Type resultType = this->getTypeConverter()->convertType(op.getAlignOut().getType()); - auto baseType = dyn_cast(adaptor.getBase().getType()); - if (!resultType || !baseType || adaptor.getAlignIn().getType() != resultType) { - return rewriter.notifyMatchFailure(op, - "unexpected converted vstur operand/result types"); - } - - StringRef calleeName = buildVsturCallee(op.getContext()); - Value modeValue = getI32Constant(rewriter, op.getLoc(), *postMode); - Value zeroValue = getI32Constant(rewriter, op.getLoc(), 0); - Value value = castToPayloadABI( - op.getLoc(), adaptor.getValue(), op.getValue().getType(), rewriter); - SmallVector args{value, adaptor.getBase(), adaptor.getAlignIn(), - modeValue, zeroValue}; - auto funcType = rewriter.getFunctionType( - TypeRange{value.getType(), adaptor.getBase().getType(), - adaptor.getAlignIn().getType(), modeValue.getType(), - zeroValue.getType()}, - TypeRange{resultType}); - auto call = - rewriter.create(op.getLoc(), calleeName, TypeRange{resultType}, args); - state.plannedDecls.push_back(PlannedDecl{calleeName.str(), funcType}); - rewriter.replaceOp(op, call.getResults()); - return success(); - } - -private: - LoweringState &state; -}; - -class LowerVstarOpPattern final : public OpConversionPattern { -public: - explicit LowerVstarOpPattern(TypeConverter &typeConverter, MLIRContext *context, - LoweringState &state) - : OpConversionPattern(typeConverter, context), state(state) {} - - LogicalResult - matchAndRewrite(pto::VstarOp op, pto::VstarOp::Adaptor adaptor, - ConversionPatternRewriter &rewriter) const override { - auto baseType = dyn_cast(adaptor.getDestination().getType()); - Type alignType = this->getTypeConverter()->convertType(op.getValue().getType()); - if (!baseType || !alignType || adaptor.getValue().getType() != alignType) { - return rewriter.notifyMatchFailure(op, - "unexpected converted vstar operand types"); - } - - StringRef calleeName = buildVstarCallee(op.getContext()); - Value zeroValue = getI32Constant(rewriter, op.getLoc(), 0); - SmallVector args{adaptor.getValue(), adaptor.getDestination(), zeroValue}; - auto funcType = rewriter.getFunctionType( - TypeRange{adaptor.getValue().getType(), adaptor.getDestination().getType(), - zeroValue.getType()}, - TypeRange{}); - rewriter.create(op.getLoc(), calleeName, TypeRange{}, args); - state.plannedDecls.push_back(PlannedDecl{calleeName.str(), funcType}); - rewriter.eraseOp(op); - return success(); - } - -private: - LoweringState &state; -}; - -class LowerVstasOpPattern final : public OpConversionPattern { -public: - explicit LowerVstasOpPattern(TypeConverter &typeConverter, MLIRContext *context, - LoweringState &state) - : OpConversionPattern(typeConverter, context), state(state) {} - - LogicalResult - matchAndRewrite(pto::VstasOp op, pto::VstasOp::Adaptor adaptor, - ConversionPatternRewriter &rewriter) const override { - auto baseType = dyn_cast(adaptor.getDestination().getType()); - Type alignType = this->getTypeConverter()->convertType(op.getValue().getType()); - auto dstType = dyn_cast(op.getDestination().getType()); - if (!baseType || !alignType || adaptor.getValue().getType() != alignType || !dstType) { - return rewriter.notifyMatchFailure(op, - "unexpected converted vstas operand types"); - } - - bool usePostIntrinsic = op.getUpdatedBase() != nullptr; - auto loweredOffset = lowerVPTOElementOffsetForIntrinsic( - op, adaptor.getDestination(), adaptor.getOffset(), - dstType.getElementType(), usePostIntrinsic, rewriter); - if (failed(loweredOffset)) { - return rewriter.notifyMatchFailure(op, "failed to convert vstas offset"); - } - - SmallVector resultTypes; - if (failed(this->getTypeConverter()->convertTypes(op->getResultTypes(), - resultTypes)) || - resultTypes.size() != (usePostIntrinsic ? 1U : 0U)) - return rewriter.notifyMatchFailure( - op, "failed to convert vstas result types"); - - StringRef calleeName = - buildVstasCallee(op.getContext(), usePostIntrinsic); - Value postValue = - getI32Constant(rewriter, op.getLoc(), usePostIntrinsic ? 1 : 0); - SmallVector args{adaptor.getValue(), loweredOffset->base, - loweredOffset->intrinsicOffset, postValue}; - auto funcType = rewriter.getFunctionType( - TypeRange{adaptor.getValue().getType(), loweredOffset->base.getType(), - loweredOffset->intrinsicOffset.getType(), - postValue.getType()}, - resultTypes); - auto call = rewriter.create(op.getLoc(), calleeName, - resultTypes, args); - state.plannedDecls.push_back(PlannedDecl{calleeName.str(), funcType}); - if (usePostIntrinsic) { - Value updatedBase = loweredOffset->updatedBase - ? loweredOffset->updatedBase - : call.getResult(0); - rewriter.replaceOp(op, updatedBase); - } else { - rewriter.eraseOp(op); - } - return success(); - } - -private: - LoweringState &state; -}; - -class LowerVgather2OpPattern final - : public OpConversionPattern { -public: - explicit LowerVgather2OpPattern(TypeConverter &typeConverter, - MLIRContext *context, LoweringState &state) - : OpConversionPattern(typeConverter, context), - state(state) {} - - LogicalResult - matchAndRewrite(pto::Vgather2Op op, pto::Vgather2Op::Adaptor adaptor, - ConversionPatternRewriter &rewriter) const override { - Type elemType = getElementTypeFromVectorLike(op.getResult().getType()); - auto basePtr = dyn_cast(adaptor.getSource().getType()); - if (!elemType || !basePtr) - return rewriter.notifyMatchFailure(op, - "unexpected converted vgather2 operand types"); - - Type resultType = this->getTypeConverter()->convertType(op.getResult().getType()); - if (!resultType) - return rewriter.notifyMatchFailure(op, "failed to convert vgather2 result type"); - - FailureOr calleeName = - buildVgather2Callee(op.getContext(), op.getSource().getType(), - op.getResult().getType()); - if (failed(calleeName)) - return rewriter.notifyMatchFailure(op, "unsupported vgather2 signature"); - - Value offsets = adaptor.getOffsets(); - FailureOr offsetsCarrierType = getVgather2OffsetsCarrierType( - rewriter, op.getSource().getType(), op.getResult().getType(), - offsets.getType()); - if (failed(offsetsCarrierType)) - return rewriter.notifyMatchFailure(op, "unsupported vgather2 offsets carrier"); - if (offsets.getType() != *offsetsCarrierType) - offsets = rewriter.create(op.getLoc(), *offsetsCarrierType, - offsets); - - auto funcType = rewriter.getFunctionType( - TypeRange{adaptor.getSource().getType(), *offsetsCarrierType, - adaptor.getMask().getType()}, - TypeRange{resultType}); - auto call = rewriter.create( - op.getLoc(), *calleeName, TypeRange{resultType}, - ValueRange{adaptor.getSource(), offsets, adaptor.getMask()}); - state.plannedDecls.push_back(PlannedDecl{calleeName->str(), funcType}); - rewriter.replaceOp(op, call.getResults()); - return success(); - } - -private: - LoweringState &state; -}; - -class LowerVgather2BcOpPattern final - : public OpConversionPattern { -public: - explicit LowerVgather2BcOpPattern(TypeConverter &typeConverter, - MLIRContext *context, LoweringState &state) - : OpConversionPattern(typeConverter, context), - state(state) {} - - LogicalResult - matchAndRewrite(pto::Vgather2BcOp op, pto::Vgather2BcOp::Adaptor adaptor, - ConversionPatternRewriter &rewriter) const override { - auto basePtr = dyn_cast(adaptor.getSource().getType()); - Type resultType = this->getTypeConverter()->convertType(op.getResult().getType()); - if (!basePtr || !resultType) - return rewriter.notifyMatchFailure(op, - "unexpected converted vgather2_bc operand/result types"); - - FailureOr calleeName = - buildVgather2BcCallee(op.getContext(), op.getResult().getType()); - if (failed(calleeName)) - return rewriter.notifyMatchFailure(op, "unsupported vgather2_bc signature"); - - auto funcType = rewriter.getFunctionType( - TypeRange{adaptor.getSource().getType(), adaptor.getOffsets().getType(), - adaptor.getMask().getType()}, - TypeRange{resultType}); - auto call = rewriter.create( - op.getLoc(), *calleeName, TypeRange{resultType}, - ValueRange{adaptor.getSource(), adaptor.getOffsets(), adaptor.getMask()}); - state.plannedDecls.push_back(PlannedDecl{calleeName->str(), funcType}); - rewriter.replaceOp(op, call.getResults()); - return success(); - } - -private: - LoweringState &state; -}; - -class LowerVgatherbOpPattern final - : public OpConversionPattern { -public: - explicit LowerVgatherbOpPattern(TypeConverter &typeConverter, - MLIRContext *context, LoweringState &state) - : OpConversionPattern(typeConverter, context), - state(state) {} - - LogicalResult - matchAndRewrite(pto::VgatherbOp op, pto::VgatherbOp::Adaptor adaptor, - ConversionPatternRewriter &rewriter) const override { - auto basePtr = dyn_cast(adaptor.getSource().getType()); - Type resultType = this->getTypeConverter()->convertType(op.getResult().getType()); - if (!basePtr || !resultType) - return rewriter.notifyMatchFailure(op, - "unexpected converted vgatherb operand/result types"); - - FailureOr calleeName = - buildVgatherbCallee(op.getContext(), op.getResult().getType()); - if (failed(calleeName)) - return rewriter.notifyMatchFailure(op, "unsupported vgatherb signature"); - - auto funcType = rewriter.getFunctionType( - TypeRange{adaptor.getSource().getType(), adaptor.getOffsets().getType(), - adaptor.getMask().getType()}, - TypeRange{resultType}); - auto call = rewriter.create( - op.getLoc(), *calleeName, TypeRange{resultType}, - ValueRange{adaptor.getSource(), adaptor.getOffsets(), adaptor.getMask()}); - state.plannedDecls.push_back(PlannedDecl{calleeName->str(), funcType}); - rewriter.replaceOp(op, call.getResults()); - return success(); - } - -private: - LoweringState &state; -}; - -class LowerVscatterOpPattern final - : public OpConversionPattern { -public: - explicit LowerVscatterOpPattern(TypeConverter &typeConverter, - MLIRContext *context, LoweringState &state) - : OpConversionPattern(typeConverter, context), - state(state) {} - - LogicalResult - matchAndRewrite(pto::VscatterOp op, pto::VscatterOp::Adaptor adaptor, - ConversionPatternRewriter &rewriter) const override { - Type elemType = getElementTypeFromVectorLike(op.getValue().getType()); - auto basePtr = - dyn_cast(adaptor.getDestination().getType()); - if (!elemType || !basePtr) - return rewriter.notifyMatchFailure(op, - "unexpected converted vscatter operand types"); - - FailureOr calleeName = - buildVscatterCallee(op.getContext(), op.getValue().getType()); - if (failed(calleeName)) - return rewriter.notifyMatchFailure(op, "unsupported vscatter signature"); - - FailureOr offsetsCarrierType = getVscatterOffsetsCarrierType( - adaptor.getOffsets().getType()); - if (failed(offsetsCarrierType)) - return rewriter.notifyMatchFailure(op, "unsupported vscatter offsets carrier"); - - auto funcType = rewriter.getFunctionType( - TypeRange{adaptor.getValue().getType(), adaptor.getDestination().getType(), - *offsetsCarrierType, adaptor.getMask().getType()}, - TypeRange{}); - rewriter.create( - op.getLoc(), *calleeName, TypeRange{}, - ValueRange{adaptor.getValue(), adaptor.getDestination(), - adaptor.getOffsets(), adaptor.getMask()}); - state.plannedDecls.push_back(PlannedDecl{calleeName->str(), funcType}); - rewriter.eraseOp(op); - return success(); - } - -private: - LoweringState &state; -}; - -class LowerVaxpyOpPattern final : public OpConversionPattern { -public: - explicit LowerVaxpyOpPattern(TypeConverter &typeConverter, - MLIRContext *context, LoweringState &state) - : OpConversionPattern(typeConverter, context), state(state) {} - - LogicalResult - matchAndRewrite(pto::VaxpyOp op, pto::VaxpyOp::Adaptor adaptor, - ConversionPatternRewriter &rewriter) const override { - Type elemType = getElementTypeFromVectorLike(op.getResult().getType()); - if (!elemType) - return rewriter.notifyMatchFailure(op, "unsupported vaxpy signature"); - - Type resultType = this->getTypeConverter()->convertType(op.getResult().getType()); - if (!resultType) - return rewriter.notifyMatchFailure(op, "failed to convert vaxpy result type"); - - FailureOr calleeName = - buildVaxpyCallee(op.getContext(), op.getResult().getType()); - if (failed(calleeName)) - return rewriter.notifyMatchFailure(op, "unsupported vaxpy callee"); - - auto funcType = rewriter.getFunctionType( - TypeRange{adaptor.getSrc1().getType(), adaptor.getSrc0().getType(), - adaptor.getAlpha().getType(), adaptor.getMask().getType()}, - TypeRange{resultType}); - auto call = rewriter.create( - op.getLoc(), *calleeName, TypeRange{resultType}, - ValueRange{adaptor.getSrc1(), adaptor.getSrc0(), adaptor.getAlpha(), - adaptor.getMask()}); - state.plannedDecls.push_back(PlannedDecl{calleeName->str(), funcType}); - rewriter.replaceOp(op, call.getResults()); - return success(); - } - -private: - LoweringState &state; -}; - -class LowerVmulscvtOpPattern final - : public OpConversionPattern { -public: - explicit LowerVmulscvtOpPattern(TypeConverter &typeConverter, - MLIRContext *context, LoweringState &state) - : OpConversionPattern(typeConverter, context), - state(state) {} - - LogicalResult - matchAndRewrite(pto::VmulscvtOp op, pto::VmulscvtOp::Adaptor adaptor, - ConversionPatternRewriter &rewriter) const override { - auto roundMode = parseRoundModeImmediate(op.getRnd()); - if (!roundMode) - return rewriter.notifyMatchFailure(op, "vmulscvt requires valid rnd attr"); - if (*roundMode != 1) - return rewriter.notifyMatchFailure( - op, "current vmulscvt lowering only supports rnd A"); - - auto part = parsePartImmediate(op.getPart()); - if (!part) - return rewriter.notifyMatchFailure(op, "unsupported vmulscvt part"); - - Type resultType = - this->getTypeConverter()->convertType(op.getResult().getType()); - if (!resultType) - return rewriter.notifyMatchFailure(op, - "failed to convert vmulscvt result type"); - - FailureOr calleeName = - buildVmulscvtCallee(op.getContext(), op.getInput().getType(), - op.getResult().getType()); - if (failed(calleeName)) - return rewriter.notifyMatchFailure(op, "unsupported vmulscvt signature"); - - Value partValue = getI32Constant(rewriter, op.getLoc(), *part); - auto funcType = rewriter.getFunctionType( - TypeRange{adaptor.getInput().getType(), adaptor.getScalar().getType(), - adaptor.getMask().getType(), partValue.getType()}, - TypeRange{resultType}); - auto call = rewriter.create( - op.getLoc(), *calleeName, TypeRange{resultType}, - ValueRange{adaptor.getInput(), adaptor.getScalar(), adaptor.getMask(), - partValue}); - state.plannedDecls.push_back(PlannedDecl{calleeName->str(), funcType}); - rewriter.replaceOp(op, call.getResults()); - return success(); - } - -private: - LoweringState &state; -}; - -class LowerVciOpPattern final : public OpConversionPattern { -public: - explicit LowerVciOpPattern(TypeConverter &typeConverter, MLIRContext *context, - LoweringState &state) - : OpConversionPattern(typeConverter, context), state(state) {} - - LogicalResult - matchAndRewrite(pto::VciOp op, pto::VciOp::Adaptor adaptor, - ConversionPatternRewriter &rewriter) const override { - auto order = parseOrderImmediate(op.getOrder().value_or("ASC")); - if (!order) - return rewriter.notifyMatchFailure(op, "unsupported vci order"); - - Type resultType = this->getTypeConverter()->convertType(op.getResult().getType()); - if (!resultType) - return rewriter.notifyMatchFailure(op, "failed to convert vci result type"); - - FailureOr calleeName = - buildVciCallee(op.getContext(), op.getResult().getType()); - if (failed(calleeName)) - return rewriter.notifyMatchFailure(op, "unsupported vci callee"); - - Value indexValue = adaptor.getIndex(); - - Value orderValue = getI32Constant(rewriter, op.getLoc(), *order); - auto funcType = rewriter.getFunctionType( - TypeRange{indexValue.getType(), orderValue.getType()}, - TypeRange{resultType}); - auto call = rewriter.create( - op.getLoc(), *calleeName, TypeRange{resultType}, - ValueRange{indexValue, orderValue}); - state.plannedDecls.push_back(PlannedDecl{calleeName->str(), funcType}); - rewriter.replaceOp(op, call.getResults()); - return success(); - } - -private: - LoweringState &state; -}; - -class LowerVexpdifOpPattern final - : public OpConversionPattern { -public: - explicit LowerVexpdifOpPattern(TypeConverter &typeConverter, - MLIRContext *context, LoweringState &state) - : OpConversionPattern(typeConverter, context), - state(state) {} - - LogicalResult - matchAndRewrite(pto::VexpdifOp op, pto::VexpdifOp::Adaptor adaptor, - ConversionPatternRewriter &rewriter) const override { - auto part = parsePartImmediate(op.getPart()); - if (!part) - return rewriter.notifyMatchFailure(op, "unsupported vexpdif signature"); - - Type resultType = this->getTypeConverter()->convertType(op.getResult().getType()); - if (!resultType) - return rewriter.notifyMatchFailure(op, "failed to convert vexpdif result type"); - - FailureOr calleeName = - buildVexpdifCallee(op.getContext(), op.getInput().getType(), - op.getResult().getType()); - if (failed(calleeName)) - return rewriter.notifyMatchFailure(op, "unsupported vexpdif callee"); - - Value partValue = getI32Constant(rewriter, op.getLoc(), *part); - auto funcType = rewriter.getFunctionType( - TypeRange{adaptor.getInput().getType(), adaptor.getMax().getType(), - adaptor.getMask().getType(), partValue.getType()}, - TypeRange{resultType}); - auto call = rewriter.create( - op.getLoc(), *calleeName, TypeRange{resultType}, - ValueRange{adaptor.getInput(), adaptor.getMax(), adaptor.getMask(), - partValue}); - state.plannedDecls.push_back(PlannedDecl{calleeName->str(), funcType}); - rewriter.replaceOp(op, call.getResults()); - return success(); - } - -private: - LoweringState &state; -}; - -class LowerVbitsortOpPattern final - : public OpConversionPattern { -public: - explicit LowerVbitsortOpPattern(TypeConverter &typeConverter, - MLIRContext *context, LoweringState &state) - : OpConversionPattern(typeConverter, context), - state(state) {} - - LogicalResult - matchAndRewrite(pto::VbitsortOp op, pto::VbitsortOp::Adaptor adaptor, - ConversionPatternRewriter &rewriter) const override { - auto dstType = - dyn_cast(adaptor.getDestination().getType()); - auto srcType = dyn_cast(adaptor.getSource().getType()); - auto idxType = - dyn_cast(adaptor.getIndices().getType()); - if (!dstType || !srcType || !idxType) - return rewriter.notifyMatchFailure(op, - "unexpected converted vbitsort operand types"); - - FailureOr config = packVbitsortConfig(op, adaptor.getRepeatTimes()); - if (failed(config)) - return rewriter.notifyMatchFailure(op, "failed to pack vbitsort config"); - - FailureOr calleeName = buildVbitsortCallee(op.getContext(), op); - if (failed(calleeName)) - return rewriter.notifyMatchFailure(op, "unsupported vbitsort signature"); - - auto funcType = rewriter.getFunctionType( - TypeRange{adaptor.getDestination().getType(), adaptor.getSource().getType(), - adaptor.getIndices().getType(), (*config).getType()}, - TypeRange{}); - rewriter.create( - op.getLoc(), *calleeName, TypeRange{}, - ValueRange{adaptor.getDestination(), adaptor.getSource(), - adaptor.getIndices(), *config}); - state.plannedDecls.push_back(PlannedDecl{calleeName->str(), funcType}); - rewriter.eraseOp(op); - return success(); - } - -private: - LoweringState &state; -}; - -class LowerVmrgsort4OpPattern final - : public OpConversionPattern { -public: - explicit LowerVmrgsort4OpPattern(TypeConverter &typeConverter, - MLIRContext *context, LoweringState &state) - : OpConversionPattern(typeConverter, context), - state(state) {} - - LogicalResult - matchAndRewrite(pto::Vmrgsort4Op op, pto::Vmrgsort4Op::Adaptor adaptor, - ConversionPatternRewriter &rewriter) const override { - auto dstType = - dyn_cast(adaptor.getDestination().getType()); - auto src0Type = - dyn_cast(adaptor.getSource0().getType()); - auto src1Type = - dyn_cast(adaptor.getSource1().getType()); - auto src2Type = - dyn_cast(adaptor.getSource2().getType()); - auto src3Type = - dyn_cast(adaptor.getSource3().getType()); - if (!dstType || !src0Type || !src1Type || !src2Type || !src3Type) - return rewriter.notifyMatchFailure( - op, "unexpected converted vmrgsort4 operand types"); - - Type elemType = - cast(op.getDestination().getType()).getElementType(); - FailureOr packedSrc = packVmrgsort4SourceAddr( - op, adaptor.getSource0(), adaptor.getSource1(), adaptor.getSource2(), - adaptor.getSource3(), elemType); - if (failed(packedSrc)) - return rewriter.notifyMatchFailure( - op, "failed to pack vmrgsort4 source addresses"); - - FailureOr dst = reinterpretPointerToAddrSpace(op, adaptor.getDestination(), 6); - if (failed(dst)) - return rewriter.notifyMatchFailure(op, "failed to normalize vmrgsort4 destination"); - - FailureOr calleeName = buildVmrgsort4Callee(op.getContext(), op); - if (failed(calleeName)) - return rewriter.notifyMatchFailure(op, "unsupported vmrgsort4 signature"); - - auto funcType = rewriter.getFunctionType( - TypeRange{(*dst).getType(), (*packedSrc).getType(), - adaptor.getCount().getType(), adaptor.getConfig().getType()}, - TypeRange{}); - rewriter.create( - op.getLoc(), *calleeName, TypeRange{}, - ValueRange{*dst, *packedSrc, adaptor.getCount(), adaptor.getConfig()}); - state.plannedDecls.push_back(PlannedDecl{calleeName->str(), funcType}); - rewriter.eraseOp(op); - return success(); - } - -private: - LoweringState &state; -}; - -class LowerVcvtOpPattern final : public OpConversionPattern { -public: - explicit LowerVcvtOpPattern(TypeConverter &typeConverter, - MLIRContext *context, LoweringState &state) - : OpConversionPattern(typeConverter, context), state(state) {} - - LogicalResult - matchAndRewrite(pto::VcvtOp op, pto::VcvtOp::Adaptor adaptor, - ConversionPatternRewriter &rewriter) const override { - FailureOr contract = buildVcvtContract(op); - if (failed(contract)) - return rewriter.notifyMatchFailure(op, "unsupported vcvt type pair"); - - Type resultType = this->getTypeConverter()->convertType(op.getResult().getType()); - if (!resultType) - return rewriter.notifyMatchFailure(op, "failed to convert vcvt result type"); - - SmallVector callArgs; - SmallVector argTypes; - callArgs.push_back(adaptor.getInput()); - argTypes.push_back(adaptor.getInput().getType()); - callArgs.push_back(adaptor.getMask()); - argTypes.push_back(adaptor.getMask().getType()); - - auto appendRndArg = [&]() -> LogicalResult { - auto roundMode = - op.getRndAttr() ? parseRoundModeImmediate(*op.getRnd()) : std::nullopt; - if (!roundMode) - return rewriter.notifyMatchFailure(op, "vcvt requires valid rnd attr"); - Value roundValue = getI32Constant(rewriter, op.getLoc(), *roundMode); - callArgs.push_back(roundValue); - argTypes.push_back(roundValue.getType()); - return success(); - }; - - auto appendSatArg = [&]() -> LogicalResult { - auto saturation = - op.getSatAttr() ? parseSaturationImmediate(*op.getSat()) : std::nullopt; - if (!saturation) - return rewriter.notifyMatchFailure(op, "vcvt requires valid sat attr"); - Value satValue = getI32Constant(rewriter, op.getLoc(), *saturation); - callArgs.push_back(satValue); - argTypes.push_back(satValue.getType()); - return success(); - }; - - if ((*contract).satBeforeRnd) { - if ((*contract).requiresSat && failed(appendSatArg())) - return failure(); - if ((*contract).requiresRnd && failed(appendRndArg())) - return failure(); - } else { - if ((*contract).requiresRnd && failed(appendRndArg())) - return failure(); - if ((*contract).requiresSat && failed(appendSatArg())) - return failure(); - } - - if ((*contract).requiresPart) { - auto part = - op.getPartAttr() ? parseVcvtPartImmediate(*op.getPart()) : std::nullopt; - if (!part) - return rewriter.notifyMatchFailure(op, "vcvt requires valid part attr"); - Value partValue = getI32Constant(rewriter, op.getLoc(), *part); - callArgs.push_back(partValue); - argTypes.push_back(partValue.getType()); - } - - auto funcType = rewriter.getFunctionType(argTypes, TypeRange{resultType}); - auto call = rewriter.create( - op.getLoc(), StringRef((*contract).intrinsic), TypeRange{resultType}, callArgs); - state.plannedDecls.push_back( - PlannedDecl{std::string((*contract).intrinsic), funcType}); - rewriter.replaceOp(op, call.getResults()); - return success(); - } - -private: - LoweringState &state; -}; - -class LowerVbitcastOpPattern final - : public OpConversionPattern { -public: - explicit LowerVbitcastOpPattern(TypeConverter &typeConverter, - MLIRContext *context, LoweringState &state) - : OpConversionPattern(typeConverter, context) {} - - LogicalResult - matchAndRewrite(pto::VbitcastOp op, pto::VbitcastOp::Adaptor adaptor, - ConversionPatternRewriter &rewriter) const override { - // A vbitcast whose result has no users is a dead noop (Pure). Erase it - // instead of emitting an LLVM bitcast the device compiler may not lower - // (e.g. bf16x2 <-> bf16 physical views). - if (op->use_empty()) { - rewriter.eraseOp(op); - return success(); - } - Type resultType = - this->getTypeConverter()->convertType(op.getResult().getType()); - if (!resultType) - return rewriter.notifyMatchFailure(op, - "failed to convert vbitcast result type"); - rewriter.replaceOpWithNewOp(op, resultType, - adaptor.getInput()); - return success(); - } -}; - -class LowerPbitcastOpPattern final - : public OpConversionPattern { -public: - explicit LowerPbitcastOpPattern(TypeConverter &typeConverter, - MLIRContext *context, LoweringState &state) - : OpConversionPattern(typeConverter, context) {} - - LogicalResult - matchAndRewrite(pto::PbitcastOp op, pto::PbitcastOp::Adaptor adaptor, - ConversionPatternRewriter &rewriter) const override { - Type resultType = - this->getTypeConverter()->convertType(op.getResult().getType()); - if (!resultType) - return rewriter.notifyMatchFailure(op, - "failed to convert pbitcast result type"); - if (adaptor.getInput().getType() != resultType) { - return rewriter.notifyMatchFailure( - op, "pbitcast expects identical lowered input/result types"); - } - rewriter.replaceOp(op, adaptor.getInput()); - return success(); - } -}; - -class LowerVtrcOpPattern final : public OpConversionPattern { -public: - explicit LowerVtrcOpPattern(TypeConverter &typeConverter, MLIRContext *context, - LoweringState &state) - : OpConversionPattern(typeConverter, context), state(state) {} - - LogicalResult - matchAndRewrite(pto::VtrcOp op, pto::VtrcOp::Adaptor adaptor, - ConversionPatternRewriter &rewriter) const override { - auto roundMode = parseRoundModeImmediate(op.getRoundMode()); - if (!roundMode) - return rewriter.notifyMatchFailure(op, "unsupported vtrc signature"); - - Type resultType = this->getTypeConverter()->convertType(op.getResult().getType()); - if (!resultType) - return rewriter.notifyMatchFailure(op, "failed to convert vtrc result type"); - - FailureOr calleeName = - buildVtrcCallee(op.getContext(), op.getResult().getType()); - if (failed(calleeName)) - return rewriter.notifyMatchFailure(op, "unsupported vtrc callee"); - - Value roundValue = getI32Constant(rewriter, op.getLoc(), *roundMode); - auto funcType = rewriter.getFunctionType( - TypeRange{adaptor.getInput().getType(), roundValue.getType(), - adaptor.getMask().getType()}, - TypeRange{resultType}); - auto call = rewriter.create( - op.getLoc(), *calleeName, TypeRange{resultType}, - ValueRange{adaptor.getInput(), roundValue, adaptor.getMask()}); - state.plannedDecls.push_back(PlannedDecl{calleeName->str(), funcType}); - rewriter.replaceOp(op, call.getResults()); - return success(); - } - -private: - LoweringState &state; -}; - -template -class LowerPredicateStoreOpPattern final : public OpConversionPattern { -public: - explicit LowerPredicateStoreOpPattern(TypeConverter &typeConverter, - MLIRContext *context, - LoweringState &state) - : OpConversionPattern(typeConverter, context), state(state) {} - - LogicalResult - matchAndRewrite(StoreOp op, typename StoreOp::Adaptor adaptor, - ConversionPatternRewriter &rewriter) const override { - auto llvmDestType = - dyn_cast(adaptor.getDestination().getType()); - Type valueType = this->getTypeConverter()->convertType(op.getValue().getType()); - if (!llvmDestType || !valueType) - return rewriter.notifyMatchFailure( - op, "expected converted predicate-store operand types"); - - auto dist = parsePredicateStoreDistImmediate(op.getDist()); - if (!dist) - return rewriter.notifyMatchFailure( - op, "unsupported predicate-store dist immediate"); - - bool usePostIntrinsic = op.getUpdatedBase() != nullptr; - auto loweredOffset = lowerVPTOPredicateOffsetForIntrinsic( - op, adaptor.getDestination(), adaptor.getOffset(), usePostIntrinsic, - rewriter); - if (failed(loweredOffset)) { - return rewriter.notifyMatchFailure( - op, "failed to preserve predicate-store index offset"); - } - - SmallVector resultTypes; - if (failed(this->getTypeConverter()->convertTypes(op->getResultTypes(), - resultTypes)) || - resultTypes.size() != (usePostIntrinsic ? 1U : 0U)) - return rewriter.notifyMatchFailure( - op, "failed to convert predicate-store result types"); - - StringRef calleeName = - getPredicateStoreCallee(op.getContext(), usePostIntrinsic); - SmallVector args; - args.push_back(adaptor.getValue()); - args.push_back(loweredOffset->base); - args.push_back(loweredOffset->intrinsicOffset); - args.push_back(rewriter.create( - op.getLoc(), rewriter.getI32IntegerAttr(*dist))); - args.push_back(rewriter.create( - op.getLoc(), - rewriter.getI32IntegerAttr(usePostIntrinsic ? 1 : 0))); - auto funcType = rewriter.getFunctionType( - TypeRange{valueType, llvmDestType, rewriter.getI32Type(), - rewriter.getI32Type(), rewriter.getI32Type()}, - resultTypes); - auto call = rewriter.create(op.getLoc(), calleeName, - resultTypes, args); - state.plannedDecls.push_back(PlannedDecl{calleeName.str(), funcType}); - if (usePostIntrinsic) { - if (loweredOffset->updatedBase) { - rewriter.replaceOp(op, loweredOffset->updatedBase); - } else { - rewriter.replaceOp(op, call.getResults()); - } - } else { - rewriter.eraseOp(op); - } - return success(); - } - -private: - LoweringState &state; -}; - -template -class LowerPredicateLoadOpPattern final : public OpConversionPattern { -public: - explicit LowerPredicateLoadOpPattern(TypeConverter &typeConverter, - MLIRContext *context, - LoweringState &state) - : OpConversionPattern(typeConverter, context), state(state) {} - - LogicalResult - matchAndRewrite(LoadOp op, typename LoadOp::Adaptor adaptor, - ConversionPatternRewriter &rewriter) const override { - auto llvmSourceType = - dyn_cast(adaptor.getSource().getType()); - bool usePostIntrinsic = op.getUpdatedBase() != nullptr; - SmallVector resultTypes; - if (failed(this->getTypeConverter()->convertTypes(op->getResultTypes(), - resultTypes)) || - resultTypes.size() != (usePostIntrinsic ? 2U : 1U)) - return rewriter.notifyMatchFailure( - op, "failed to convert predicate-load result types"); - if (!llvmSourceType) - return rewriter.notifyMatchFailure( - op, "expected converted predicate-load operand/result types"); - - auto dist = parsePredicateLoadDistImmediate(op.getDist()); - if (!dist) - return rewriter.notifyMatchFailure( - op, "unsupported predicate-load dist immediate"); - - auto loweredOffset = lowerVPTOPredicateOffsetForIntrinsic( - op, adaptor.getSource(), adaptor.getOffset(), usePostIntrinsic, - rewriter); - if (failed(loweredOffset)) { - return rewriter.notifyMatchFailure( - op, "failed to preserve predicate-load index offset"); - } - - StringRef calleeName = - getPredicateLoadCallee(op.getContext(), usePostIntrinsic); - SmallVector args; - args.push_back(loweredOffset->base); - args.push_back(loweredOffset->intrinsicOffset); - args.push_back(rewriter.create( - op.getLoc(), rewriter.getI32IntegerAttr(*dist))); - args.push_back(rewriter.create( - op.getLoc(), - rewriter.getI32IntegerAttr(usePostIntrinsic ? 1 : 0))); - auto funcType = rewriter.getFunctionType( - TypeRange{llvmSourceType, rewriter.getI32Type(), rewriter.getI32Type(), - rewriter.getI32Type()}, - resultTypes); - auto call = rewriter.create(op.getLoc(), calleeName, - resultTypes, args); - state.plannedDecls.push_back(PlannedDecl{calleeName.str(), funcType}); - if (loweredOffset->updatedBase) { - rewriter.replaceOp( - op, ValueRange{call.getResult(0), loweredOffset->updatedBase}); - } else { - rewriter.replaceOp(op, call.getResults()); - } - return success(); - } - -private: - LoweringState &state; -}; - -template -class LowerSetLoopConfigOpPattern final : public OpConversionPattern { -public: - explicit LowerSetLoopConfigOpPattern(TypeConverter &typeConverter, - MLIRContext *context, - LoweringState &state) - : OpConversionPattern(typeConverter, context), state(state) {} - - LogicalResult - matchAndRewrite(LoopOp op, typename LoopOp::Adaptor adaptor, - ConversionPatternRewriter &rewriter) const override { - FailureOr packed = failure(); - if constexpr (std::is_same_v || - std::is_same_v) { - packed = packLoopSize(op, adaptor.getFirst(), adaptor.getSecond()); - } else { - packed = packLoopPair(op, adaptor.getFirst(), adaptor.getSecond()); - } - if (failed(packed)) - return rewriter.notifyMatchFailure(op, - "failed to pack loop configuration"); - - StringRef calleeName = buildSetLoopCallee(op.getContext()); - auto funcType = - rewriter.getFunctionType(TypeRange{rewriter.getI64Type()}, TypeRange{}); - rewriter.create(op.getLoc(), calleeName, TypeRange{}, - ValueRange{*packed}); - state.plannedDecls.push_back(PlannedDecl{calleeName.str(), funcType}); - rewriter.eraseOp(op); - return success(); - } - -private: - LoweringState &state; -}; - -template -class LowerUnaryConfigOpPattern final : public OpConversionPattern { -public: - explicit LowerUnaryConfigOpPattern(TypeConverter &typeConverter, - MLIRContext *context, - LoweringState &state) - : OpConversionPattern(typeConverter, context), state(state) {} - - LogicalResult - matchAndRewrite(ConfigOp op, typename ConfigOp::Adaptor adaptor, - ConversionPatternRewriter &rewriter) const override { - FailureOr encoded = - encodeMovPadValue(op.getLoc(), adaptor.getValue(), rewriter); - if (failed(encoded)) - return rewriter.notifyMatchFailure( - op, "expected 8/16/32-bit integer or float mov-pad payload"); - - StringRef calleeName = buildUnaryConfigCallee(op.getContext()); - auto funcType = - rewriter.getFunctionType(TypeRange{rewriter.getI64Type()}, TypeRange{}); - rewriter.create(op.getLoc(), calleeName, TypeRange{}, - ValueRange{*encoded}); - state.plannedDecls.push_back(PlannedDecl{calleeName.str(), funcType}); - rewriter.eraseOp(op); - return success(); - } - -private: - LoweringState &state; -}; - -template -class LowerUnaryI64ConfigOpPattern final : public OpConversionPattern { -public: - explicit LowerUnaryI64ConfigOpPattern(TypeConverter &typeConverter, - MLIRContext *context, - LoweringState &state) - : OpConversionPattern(typeConverter, context), state(state) {} - - LogicalResult - matchAndRewrite(ConfigOp op, typename ConfigOp::Adaptor adaptor, - ConversionPatternRewriter &rewriter) const override { - StringRef calleeName = buildUnaryConfigCallee(op.getContext()); - auto funcType = - rewriter.getFunctionType(TypeRange{adaptor.getValue().getType()}, - TypeRange{}); - rewriter.create(op.getLoc(), calleeName, TypeRange{}, - ValueRange{adaptor.getValue()}); - state.plannedDecls.push_back(PlannedDecl{calleeName.str(), funcType}); - rewriter.eraseOp(op); - return success(); - } - -private: - LoweringState &state; -}; - -class LowerStoreVfSimtInfoOpPattern final - : public OpConversionPattern { -public: - explicit LowerStoreVfSimtInfoOpPattern(TypeConverter &typeConverter, - MLIRContext *context, - LoweringState &state) - : OpConversionPattern(typeConverter, context), - state(state) {} - - LogicalResult - matchAndRewrite(pto::StoreVfSimtInfoOp op, - pto::StoreVfSimtInfoOp::Adaptor adaptor, - ConversionPatternRewriter &rewriter) const override { - Location loc = op.getLoc(); - Value dimZ = adaptor.getDimZ(); - Value dimY = adaptor.getDimY(); - Value dimX = adaptor.getDimX(); - if (!dimZ || !dimY || !dimX) - return rewriter.notifyMatchFailure(op, "missing converted SIMT dims"); - - auto i64Type = rewriter.getI64Type(); - auto castToI64 = [&](Value value) -> Value { - if (value.getType().isInteger(64)) - return value; - return rewriter.create(loc, i64Type, value).getResult(); - }; - - Value dimZI64 = castToI64(dimZ); - Value dimYI64 = castToI64(dimY); - Value dimXI64 = castToI64(dimX); - Value dimYShift = rewriter.create( - loc, i64Type, rewriter.getI64IntegerAttr(16)); - Value dimZShift = rewriter.create( - loc, i64Type, rewriter.getI64IntegerAttr(32)); - Value packedDimY = - rewriter.create(loc, dimYI64, dimYShift).getResult(); - Value packedDimZ = - rewriter.create(loc, dimZI64, dimZShift).getResult(); - Value payload = - rewriter.create(loc, dimXI64, packedDimY).getResult(); - payload = - rewriter.create(loc, payload, packedDimZ).getResult(); - - StringRef calleeName = buildStoreVfSimtInfoCallee(op.getContext()); - auto funcType = rewriter.getFunctionType(TypeRange{i64Type}, TypeRange{}); - rewriter.create(loc, calleeName, TypeRange{}, - ValueRange{payload}); - state.plannedDecls.push_back(PlannedDecl{calleeName.str(), funcType}); - rewriter.eraseOp(op); - return success(); - } - -private: - LoweringState &state; -}; - -template -static StringRef buildSimtFenceCallee(MLIRContext *context); - -template <> -StringRef buildSimtFenceCallee(MLIRContext *context) { - return buildSyncthreadsCallee(context); -} - -template <> -StringRef buildSimtFenceCallee(MLIRContext *context) { - return buildThreadfenceCallee(context); -} - -template <> -StringRef buildSimtFenceCallee(MLIRContext *context) { - return buildThreadfenceBlockCallee(context); -} - -template -class LowerSimtFenceOpPattern final : public OpConversionPattern { -public: - explicit LowerSimtFenceOpPattern(TypeConverter &typeConverter, - MLIRContext *context, - LoweringState &state) - : OpConversionPattern(typeConverter, context), - state(state) {} - - LogicalResult - matchAndRewrite(FenceOp op, typename FenceOp::Adaptor adaptor, - ConversionPatternRewriter &rewriter) const override { - (void)adaptor; - FunctionType funcType = rewriter.getFunctionType({}, {}); - StringRef calleeName = buildSimtFenceCallee(op.getContext()); - rewriter.create(op.getLoc(), calleeName, TypeRange{}, - ValueRange{}); - state.plannedDecls.push_back(PlannedDecl{calleeName.str(), funcType}); - rewriter.eraseOp(op); - return success(); - } - -private: - LoweringState &state; -}; - -struct SimtKeepResumePhysicalRegister { - int64_t baseRegister; - unsigned registerCount; -}; - -// TPERn names one 32-bit register, while TPERLn names the 64-bit pair whose -// base register is R(2n). Keep uses tied inputs so the compiler models the -// value captured by each fixed output without inline assembly instructions. -static std::string buildSimtKeepResumeConstraints( - ArrayRef physicalRegs, bool tieInputs) { - std::string result; - llvm::raw_string_ostream os(result); - for (auto [index, physicalReg] : llvm::enumerate(physicalRegs)) { - if (index != 0) - os << ","; - if (physicalReg.registerCount == 2) - os << "={TPERL" << physicalReg.baseRegister / 2 << "}"; - else - os << "={TPER" << physicalReg.baseRegister << "}"; - } - if (tieInputs) { - for (size_t index = 0; index < physicalRegs.size(); ++index) - os << "," << index; - } - return os.str(); -} - -template -static SmallVector collectConsecutiveOps(OpT first) { - SmallVector ops; - for (Operation *cur = first.getOperation(); cur; cur = cur->getNextNode()) { - auto typed = dyn_cast(cur); - if (!typed) - break; - ops.push_back(typed); - } - return ops; -} - -static bool hasPreviousSameOp(Operation *op) { - Operation *prev = op->getPrevNode(); - return prev && prev->getName() == op->getName(); -} - -static std::optional getSimtKeepResumeBitWidth(Type type) { - if (auto intType = dyn_cast(type)) { - if (intType.getWidth() <= 64) - return intType.getWidth(); - return std::nullopt; - } - if (type.isF16() || type.isBF16()) - return 16; - if (type.isF32()) - return 32; - return std::nullopt; -} - -static Value packSimtKeepResumePayload(Location loc, Value value, - ConversionPatternRewriter &rewriter) { - Type type = value.getType(); - std::optional width = getSimtKeepResumeBitWidth(type); - if (!width) - return {}; - - Type intType = rewriter.getIntegerType(*width); - Value bits = value; - if (!isa(type)) - bits = rewriter.create(loc, intType, value); - else if (bits.getType() != intType) - bits = rewriter.create(loc, intType, bits); - if (*width < 32) - return rewriter.create(loc, rewriter.getI32Type(), bits); - if (*width == 32 && bits.getType() != rewriter.getI32Type()) - return rewriter.create(loc, rewriter.getI32Type(), bits); - return bits; -} - -static Value unpackSimtKeepResumePayload(Location loc, Value value, - Type resultType, - ConversionPatternRewriter &rewriter) { - std::optional width = getSimtKeepResumeBitWidth(resultType); - if (!width) - return {}; - - Type intType = rewriter.getIntegerType(*width); - Value bits = value; - if (*width < 32) - bits = rewriter.create(loc, intType, bits); - else if (bits.getType() != intType) - bits = rewriter.create(loc, intType, bits); - - if (isa(resultType)) { - if (bits.getType() == resultType) - return bits; - return rewriter.create(loc, resultType, bits); - } - return rewriter.create(loc, resultType, bits); -} - -static unsigned getSimtKeepResumeRegisterCount(Type type) { - std::optional width = getSimtKeepResumeBitWidth(type); - return width && *width > 32 ? 2 : 1; -} - -static FailureOr> -computeSimtKeepResumePhysicalRegs( - ArrayRef> logicalSlots) { - SmallVector physicalRegs; - physicalRegs.reserve(logicalSlots.size()); - for (auto [slot, registerCount] : logicalSlots) { - if (slot < 0 || slot >= 123) - return failure(); - if (registerCount == 2 && ((slot % 2) != 0 || slot + 1 >= 123)) - return failure(); - // Slots are user-assigned storage words, not dense ordinals in the current - // keep/resume group. This keeps a consumer that resumes only a subset of - // slots from changing where the remaining slots are read from. - int64_t baseRegister = 4 + slot; - if (baseRegister + static_cast(registerCount) - 1 > 126) - return failure(); - physicalRegs.push_back({baseRegister, registerCount}); - } - return physicalRegs; -} - -static bool isValidSimtKeepResumeSlot(int64_t slot, unsigned registerCount) { - if (slot < 0 || slot >= 123) - return false; - if (registerCount == 2 && ((slot % 2) != 0 || slot + 1 >= 123)) - return false; - return true; -} - -class LowerKeepOpPattern final : public OpConversionPattern { -public: - explicit LowerKeepOpPattern(TypeConverter &typeConverter, - MLIRContext *context, LoweringState &) - : OpConversionPattern(typeConverter, context) {} - - LogicalResult - matchAndRewrite(pto::KeepOp op, pto::KeepOp::Adaptor adaptor, - ConversionPatternRewriter &rewriter) const override { - (void)adaptor; - if (hasPreviousSameOp(op.getOperation())) - return rewriter.notifyMatchFailure( - op, "only the first keep in a contiguous group is lowered"); - - SmallVector keepOps = collectConsecutiveOps(op); - SmallVector payloads; - SmallVector asmResultTypes; - SmallVector, 4> logicalSlots; - for (pto::KeepOp keep : keepOps) { - Value payload = rewriter.getRemappedValue(keep.getPayload()); - if (!payload) - return rewriter.notifyMatchFailure(keep, "payload is not remapped"); - payload = packSimtKeepResumePayload(keep.getLoc(), payload, rewriter); - if (!payload) - return rewriter.notifyMatchFailure( - keep, "expected integer scalar up to 64 bits or f16/bf16/f32"); - int64_t slot = keep.getSlot(); - unsigned registerCount = - getSimtKeepResumeRegisterCount(payload.getType()); - if (!isValidSimtKeepResumeSlot(slot, registerCount)) - return rewriter.notifyMatchFailure( - keep, - "slot must be in range [0, 122] and 64-bit slots must be even"); - logicalSlots.push_back({slot, registerCount}); - payloads.push_back(payload); - asmResultTypes.push_back(payload.getType()); - } - FailureOr> physicalRegs = - computeSimtKeepResumePhysicalRegs(logicalSlots); - if (failed(physicalRegs)) - return rewriter.notifyMatchFailure( - op, "keep slots must map to valid non-overlapping SIMT registers"); - - Type asmResultType = asmResultTypes.front(); - if (asmResultTypes.size() > 1) - asmResultType = - LLVM::LLVMStructType::getLiteral(op.getContext(), asmResultTypes); - rewriter.setInsertionPoint(op); - rewriter.create( - op.getLoc(), TypeRange{asmResultType}, payloads, "", - buildSimtKeepResumeConstraints(*physicalRegs, true), true, false, - LLVM::AsmDialectAttr::get(op.getContext(), LLVM::AsmDialect::AD_ATT), - ArrayAttr{}); - for (pto::KeepOp keep : llvm::reverse(keepOps)) - rewriter.eraseOp(keep); - return success(); - } -}; - -class LowerResumeOpPattern final : public OpConversionPattern { -public: - explicit LowerResumeOpPattern(TypeConverter &typeConverter, - MLIRContext *context, LoweringState &) - : OpConversionPattern(typeConverter, context) {} - - LogicalResult - matchAndRewrite(pto::ResumeOp op, pto::ResumeOp::Adaptor adaptor, - ConversionPatternRewriter &rewriter) const override { - (void)adaptor; - if (hasPreviousSameOp(op.getOperation())) - return rewriter.notifyMatchFailure( - op, "only the first resume in a contiguous group is lowered"); - - SmallVector resumeOps = collectConsecutiveOps(op); - SmallVector, 4> logicalSlots; - SmallVector asmResultTypes; - for (pto::ResumeOp resume : resumeOps) { - Type resultType = getTypeConverter()->convertType(resume.getType()); - if (!resultType || !getSimtKeepResumeBitWidth(resultType)) - return rewriter.notifyMatchFailure( - resume, "expected integer scalar up to 64 bits or f16/bf16/f32"); - int64_t slot = resume.getSlot(); - unsigned registerCount = getSimtKeepResumeRegisterCount(resultType); - if (!isValidSimtKeepResumeSlot(slot, registerCount)) - return rewriter.notifyMatchFailure( - resume, - "slot must be in range [0, 122] and 64-bit slots must be even"); - logicalSlots.push_back({slot, registerCount}); - asmResultTypes.push_back(rewriter.getIntegerType( - *getSimtKeepResumeBitWidth(resultType) > 32 ? 64 : 32)); - } - FailureOr> physicalRegs = - computeSimtKeepResumePhysicalRegs(logicalSlots); - if (failed(physicalRegs)) - return rewriter.notifyMatchFailure( - op, "resume slots must map to valid non-overlapping SIMT registers"); - - Type asmResultType = asmResultTypes.front(); - if (asmResultTypes.size() > 1) { - asmResultType = - LLVM::LLVMStructType::getLiteral(op.getContext(), asmResultTypes); - } - rewriter.setInsertionPoint(op); - auto asmOp = rewriter.create( - op.getLoc(), TypeRange{asmResultType}, ValueRange{}, "", - buildSimtKeepResumeConstraints(*physicalRegs, false), true, false, - LLVM::AsmDialectAttr::get(op.getContext(), LLVM::AsmDialect::AD_ATT), - ArrayAttr{}); - - if (resumeOps.size() == 1) { - Type resultType = getTypeConverter()->convertType(op.getType()); - Value result = unpackSimtKeepResumePayload(op.getLoc(), asmOp.getRes(), - resultType, rewriter); - if (!result) - return rewriter.notifyMatchFailure(op, "failed to unpack result"); - rewriter.replaceOp(op, result); - return success(); - } - - rewriter.setInsertionPointAfter(asmOp); - SmallVector results; - for (auto [index, resume] : llvm::enumerate(resumeOps)) { - auto extract = rewriter.create( - resume.getLoc(), asmOp.getRes(), - ArrayRef{static_cast(index)}); - Type resultType = getTypeConverter()->convertType(resume.getType()); - Value result = unpackSimtKeepResumePayload( - resume.getLoc(), extract.getRes(), resultType, rewriter); - if (!result) - return rewriter.notifyMatchFailure(resume, "failed to unpack result"); - results.push_back(result); - } - for (auto [resume, result] : llvm::zip(resumeOps, results)) - rewriter.replaceOp(resume, result); - return success(); - } -}; - -template -class LowerNullaryConfigOpPattern final : public OpConversionPattern { -public: - explicit LowerNullaryConfigOpPattern(TypeConverter &typeConverter, - MLIRContext *context, - LoweringState &state) - : OpConversionPattern(typeConverter, context), state(state) {} - - LogicalResult - matchAndRewrite(ConfigOp op, typename ConfigOp::Adaptor adaptor, - ConversionPatternRewriter &rewriter) const override { - (void)adaptor; - StringRef calleeName = buildNullaryConfigCallee(op.getContext()); - auto funcType = rewriter.getFunctionType(TypeRange{}, TypeRange{}); - rewriter.create(op.getLoc(), calleeName, TypeRange{}, - ValueRange{}); - state.plannedDecls.push_back(PlannedDecl{calleeName.str(), funcType}); - rewriter.eraseOp(op); - return success(); - } - -private: - LoweringState &state; -}; - -template -class LowerPipeEventSyncOpPattern final : public OpConversionPattern { -public: - explicit LowerPipeEventSyncOpPattern(TypeConverter &typeConverter, - MLIRContext *context, - LoweringState &state) - : OpConversionPattern(typeConverter, context), state(state) {} - - LogicalResult - matchAndRewrite(SyncOp op, typename SyncOp::Adaptor adaptor, - ConversionPatternRewriter &rewriter) const override { - (void)adaptor; - auto src = parsePipeImmediate(stringifyPIPE(op.getSrcPipe().getPipe())); - auto dst = parsePipeImmediate(stringifyPIPE(op.getDstPipe().getPipe())); - auto event = parseEventImmediate(stringifyEVENT(op.getEventId().getEvent())); - if (!src || !dst || !event) - return rewriter.notifyMatchFailure(op, "unsupported sync immediate"); - - StringRef calleeName = buildSyncCallee(op.getContext()); - Value srcValue = getI64Constant(rewriter, op.getLoc(), *src); - Value dstValue = getI64Constant(rewriter, op.getLoc(), *dst); - Value eventValue = getI64Constant(rewriter, op.getLoc(), *event); - auto funcType = rewriter.getFunctionType( - TypeRange{rewriter.getI64Type(), rewriter.getI64Type(), - rewriter.getI64Type()}, - TypeRange{}); - rewriter.create(op.getLoc(), calleeName, TypeRange{}, - ValueRange{srcValue, dstValue, eventValue}); - state.plannedDecls.push_back(PlannedDecl{calleeName.str(), funcType}); - rewriter.eraseOp(op); - return success(); - } - -private: - LoweringState &state; -}; - -template -class LowerPipeEventDynSyncOpPattern final : public OpConversionPattern { -public: - explicit LowerPipeEventDynSyncOpPattern(TypeConverter &typeConverter, - MLIRContext *context, - LoweringState &state) - : OpConversionPattern(typeConverter, context), state(state) {} - - LogicalResult - matchAndRewrite(SyncOp op, typename SyncOp::Adaptor adaptor, - ConversionPatternRewriter &rewriter) const override { - auto src = parsePipeImmediate(stringifyPIPE(op.getSrcPipe().getPipe())); - auto dst = parsePipeImmediate(stringifyPIPE(op.getDstPipe().getPipe())); - if (!src || !dst) - return rewriter.notifyMatchFailure(op, "unsupported sync pipe"); - - StringRef calleeName = buildSyncCallee(op.getContext()); - Value srcValue = getI64Constant(rewriter, op.getLoc(), *src); - Value dstValue = getI64Constant(rewriter, op.getLoc(), *dst); - - Value eventIdValue = adaptor.getEventId(); - if (!eventIdValue) - return rewriter.notifyMatchFailure(op, "missing event_id operand"); - - Value eventValue = eventIdValue; - - while (eventValue.getDefiningOp()) { - auto unrealizedCast = dyn_cast(eventValue.getDefiningOp()); - if (!unrealizedCast || unrealizedCast.getInputs().size() != 1) - break; - eventValue = unrealizedCast.getInputs()[0]; - } - - if (eventValue.getType().isIndex()) { - eventValue = rewriter.create(op.getLoc(), - rewriter.getI64Type(), - eventValue); - } else if (auto intType = dyn_cast(eventValue.getType())) { - if (intType.getWidth() < 64) { - eventValue = rewriter.create(op.getLoc(), - rewriter.getI64Type(), - eventValue); - } - } else { - return rewriter.notifyMatchFailure(op, "unexpected event_id type"); - } - - auto funcType = rewriter.getFunctionType( - TypeRange{rewriter.getI64Type(), rewriter.getI64Type(), - rewriter.getI64Type()}, - TypeRange{}); - rewriter.create(op.getLoc(), calleeName, TypeRange{}, - ValueRange{srcValue, dstValue, eventValue}); - state.plannedDecls.push_back(PlannedDecl{calleeName.str(), funcType}); - rewriter.eraseOp(op); - return success(); - } - -private: - LoweringState &state; -}; - -template -class LowerInterCoreSyncOpPattern final : public OpConversionPattern { -public: - explicit LowerInterCoreSyncOpPattern(TypeConverter &typeConverter, - MLIRContext *context, - LoweringState &state) - : OpConversionPattern(typeConverter, context), state(state) {} - - LogicalResult - matchAndRewrite(SyncOp op, typename SyncOp::Adaptor adaptor, - ConversionPatternRewriter &rewriter) const override { - auto pipe = parsePipeImmediate(stringifyPIPE(op.getPipe().getPipe())); - if (!pipe) - return rewriter.notifyMatchFailure(op, "unsupported inter-core sync pipe"); - - Value pipeValue = getI64Constant(rewriter, op.getLoc(), *pipe); - Value eventValue; - if (IntegerAttr eventIdAttr = op.getEventIdAttr()) { - eventValue = getI64Constant(rewriter, op.getLoc(), eventIdAttr.getInt()); - } else { - Value eventIdDyn = adaptor.getEventIdDyn(); - if (!eventIdDyn) - return rewriter.notifyMatchFailure( - op, "expected static or dynamic event-id operand"); - - eventValue = castIntegerLikeTo(op, eventIdDyn, rewriter.getI64Type()); - if (!eventValue) { - return rewriter.notifyMatchFailure( - op, "failed to cast dynamic event-id to i64"); - } - } - - StringRef calleeName = buildSyncCallee(op.getContext()); - SmallVector args{pipeValue, eventValue}; - if constexpr (std::is_same_v) { - int64_t mode = 2; - if (IntegerAttr attr = op.getFftsModeAttr()) { - mode = attr.getInt(); - } - Value modeValue = getI64Constant(rewriter, op.getLoc(), mode); - Value one = getI64Constant(rewriter, op.getLoc(), 1); - Value modeMask = getI64Constant(rewriter, op.getLoc(), 0x3); - Value eventMask = getI64Constant(rewriter, op.getLoc(), 0xf); - modeValue = rewriter.create(op.getLoc(), modeValue, - modeMask); - eventValue = rewriter.create(op.getLoc(), eventValue, - eventMask); - Value modeShift = rewriter.create(op.getLoc(), modeValue, - getI64Constant(rewriter, op.getLoc(), 4)); - Value eventShift = rewriter.create(op.getLoc(), eventValue, - getI64Constant(rewriter, op.getLoc(), 8)); - Value msg = rewriter.create(op.getLoc(), one, modeShift); - msg = rewriter.create(op.getLoc(), msg, eventShift); - args = {pipeValue, msg}; - } else if constexpr (std::is_same_v) { - calleeName = op.getEventIdAttr() - ? StringAttr::get(op.getContext(), - "llvm.hivm.WAIT.FLAG.DEV.PIPE.IMM") - .getValue() - : StringAttr::get(op.getContext(), - "llvm.hivm.WAIT.FLAG.DEV.PIPE.REG") - .getValue(); - } - auto funcType = rewriter.getFunctionType( - TypeRange{args[0].getType(), args[1].getType()}, TypeRange{}); - rewriter.create(op.getLoc(), calleeName, TypeRange{}, args); - state.plannedDecls.push_back(PlannedDecl{calleeName.str(), funcType}); - rewriter.eraseOp(op); - return success(); - } - -private: - LoweringState &state; -}; - -template -class LowerNamedSyncOpPattern final : public OpConversionPattern { -public: - explicit LowerNamedSyncOpPattern(TypeConverter &tc, MLIRContext *ctx, - LoweringState &state) - : OpConversionPattern(tc, ctx), state(state) {} - LogicalResult matchAndRewrite( - SyncOp op, typename SyncOp::Adaptor adaptor, - ConversionPatternRewriter &rewriter) const override { - auto pipe = parsePipeImmediate(stringifyPIPE(op.getPipe().getPipe())); - if (!pipe) { - return rewriter.notifyMatchFailure(op, "unsupported sync pipe"); - } - Value pipeValue = getI64Constant(rewriter, op.getLoc(), *pipe); - Value eventValue; - if (IntegerAttr attr = op.getEventIdAttr()) { - eventValue = getI64Constant(rewriter, op.getLoc(), attr.getInt()); - } else { - eventValue = castIntegerLikeTo(op, adaptor.getEventIdDyn(), - rewriter.getI64Type()); - if (!eventValue) { - return rewriter.notifyMatchFailure(op, "missing event-id operand"); - } - } - StringRef callee = buildSyncCallee(op.getContext()); - auto fnTy = rewriter.getFunctionType( - TypeRange{rewriter.getI64Type(), rewriter.getI64Type()}, TypeRange{}); - rewriter.create(op.getLoc(), callee, TypeRange{}, - ValueRange{pipeValue, eventValue}); - state.plannedDecls.push_back(PlannedDecl{callee.str(), fnTy}); - rewriter.eraseOp(op); - return success(); - } -private: - LoweringState &state; -}; - -class LowerBarrierOpPattern final : public OpConversionPattern { -public: - explicit LowerBarrierOpPattern(TypeConverter &typeConverter, - MLIRContext *context, LoweringState &state) - : OpConversionPattern(typeConverter, context), - state(state) {} - - LogicalResult - matchAndRewrite(pto::BarrierOp op, OpAdaptor adaptor, - ConversionPatternRewriter &rewriter) const override { - (void)adaptor; - if (isTargetArchA5(op.getOperation()) && - op.getPipe().getPipe() == PIPE::PIPE_V) { - op.emitError("internal error: A5 PIPE_V barrier should be erased before " - "VPTO LLVM lowering"); - return failure(); - } - - auto pipe = parsePipeImmediate(stringifyPIPE(op.getPipe().getPipe())); - if (!pipe) - return rewriter.notifyMatchFailure(op, "unsupported barrier pipe"); - - StringRef calleeName = buildSyncCallee(op.getContext()); - Value pipeValue = getI64Constant(rewriter, op.getLoc(), *pipe); - auto funcType = - rewriter.getFunctionType(TypeRange{rewriter.getI64Type()}, TypeRange{}); - rewriter.create(op.getLoc(), calleeName, TypeRange{}, - ValueRange{pipeValue}); - state.plannedDecls.push_back(PlannedDecl{calleeName.str(), funcType}); - rewriter.eraseOp(op); - return success(); - } - -private: - LoweringState &state; -}; - -class LowerMemBarOpPattern final : public OpConversionPattern { -public: - explicit LowerMemBarOpPattern(TypeConverter &typeConverter, - MLIRContext *context, LoweringState &state) - : OpConversionPattern(typeConverter, context), - state(state) {} - - LogicalResult - matchAndRewrite(pto::MemBarOp op, OpAdaptor adaptor, - ConversionPatternRewriter &rewriter) const override { - (void)adaptor; - StringRef calleeName = buildMemBarCallee(op.getKind().getKind(), op.getContext()); - auto funcType = rewriter.getFunctionType(TypeRange{}, TypeRange{}); - rewriter.create(op.getLoc(), calleeName, TypeRange{}, ValueRange{}); - state.plannedDecls.push_back(PlannedDecl{calleeName.str(), funcType}); - rewriter.eraseOp(op); - return success(); - } - -private: - LoweringState &state; -}; - -template -class LowerUnsupportedMemoryConsistencyOpPattern final - : public OpConversionPattern { -public: - explicit LowerUnsupportedMemoryConsistencyOpPattern( - TypeConverter &typeConverter, MLIRContext *context, - LoweringState &state) - : OpConversionPattern(typeConverter, context) { - (void)state; - } - - LogicalResult - matchAndRewrite(MemoryConsistencyOp op, - typename MemoryConsistencyOp::Adaptor adaptor, - ConversionPatternRewriter &rewriter) const override { - (void)adaptor; - (void)rewriter; - op.emitOpError() - << "is not supported by the VPTO backend yet; PTOAS validates the " - "memory-consistency contract, but high-level CMO/fence ops must be " - "lowered to `pto.dcci` or `pto.dsb` before VPTO LLVM lowering"; - return failure(); - } -}; - -class LowerDsbOpPattern final : public OpConversionPattern { -public: - explicit LowerDsbOpPattern(TypeConverter &typeConverter, MLIRContext *context, - LoweringState &state) - : OpConversionPattern(typeConverter, context), state(state) {} - - LogicalResult - matchAndRewrite(pto::DsbOp op, OpAdaptor adaptor, - ConversionPatternRewriter &rewriter) const override { - (void)adaptor; - StringRef calleeName = - StringAttr::get(op.getContext(), "llvm.hivm.DSB").getValue(); - Type i64Ty = rewriter.getI64Type(); - auto funcType = rewriter.getFunctionType(TypeRange{i64Ty}, TypeRange{}); - Value mem = - getI64Constant(rewriter, op.getLoc(), - getDsbMemImmediate(op.getMem().getKind())); - rewriter.create(op.getLoc(), calleeName, TypeRange{}, - ValueRange{mem}); - state.plannedDecls.push_back(PlannedDecl{calleeName.str(), funcType}); - rewriter.eraseOp(op); - return success(); - } - -private: - LoweringState &state; -}; - -class LowerDcciOpPattern final : public OpConversionPattern { -public: - explicit LowerDcciOpPattern(TypeConverter &typeConverter, MLIRContext *context, - LoweringState &state) - : OpConversionPattern(typeConverter, context), state(state) {} - - LogicalResult - matchAndRewrite(pto::DcciOp op, OpAdaptor adaptor, - ConversionPatternRewriter &rewriter) const override { - auto ptrType = dyn_cast(adaptor.getPtr().getType()); - if (!ptrType) - return rewriter.notifyMatchFailure(op, "expected LLVM pointer operand"); - - bool hasDst = static_cast(op.getDstAttr()); - StringRef calleeName = - buildDcciCallee(ptrType.getAddressSpace(), hasDst, op.getContext()); - - Type i64Ty = rewriter.getI64Type(); - SmallVector argTypes{ptrType, i64Ty}; - SmallVector args{ - adaptor.getPtr(), - getI64Constant(rewriter, op.getLoc(), - getDcciCacheLineImmediate(op.getCache().getKind()))}; - if (auto dst = op.getDstAttr()) { - argTypes.push_back(i64Ty); - args.push_back(getI64Constant(rewriter, op.getLoc(), - getDcciDstImmediate(dst.getKind()))); - } - - auto funcType = rewriter.getFunctionType(argTypes, TypeRange{}); - rewriter.create(op.getLoc(), calleeName, TypeRange{}, args); - state.plannedDecls.push_back(PlannedDecl{calleeName.str(), funcType}); - rewriter.eraseOp(op); - return success(); - } - -private: - LoweringState &state; -}; - -template -class LowerBufSyncOpPattern final : public OpConversionPattern { -public: - explicit LowerBufSyncOpPattern(TypeConverter &typeConverter, - MLIRContext *context, LoweringState &state) - : OpConversionPattern(typeConverter, context), state(state) {} - - LogicalResult - matchAndRewrite(BufSyncOp op, typename BufSyncOp::Adaptor adaptor, - ConversionPatternRewriter &rewriter) const override { - (void)adaptor; - PIPE pipe = PIPE::PIPE_UNASSIGNED; - if (auto pipeAttr = dyn_cast(op.getOpTypeAttr())) { - pipe = pipeAttr.getPipe(); - } else { - auto opTypeOr = parseSyncOpTypeLikeAttr(op.getOpTypeAttr()); - if (failed(opTypeOr)) - return rewriter.notifyMatchFailure( - op, "buffer sync expects pipe/sync_op_type/pipe_event_type attr"); - pipe = mapSyncOpTypeToPipe(*opTypeOr); - } - if (!isConcreteSyncPipe(pipe)) - return rewriter.notifyMatchFailure(op, - "buffer sync op_type cannot map to concrete pipe"); - - auto pipeImm = parsePipeImmediate(stringifyPIPE(pipe)); - if (!pipeImm) - return rewriter.notifyMatchFailure(op, "unsupported buffer sync pipe"); - - StringRef calleeName = buildSyncCallee(op.getContext()); - Value pipeValue = getI64Constant(rewriter, op.getLoc(), *pipeImm); - Value bufIdValue = - getI64Constant(rewriter, op.getLoc(), op.getBufIdAttr().getInt()); - Value modeValue = - getI64Constant(rewriter, op.getLoc(), op.getModeAttr().getInt()); - auto funcType = rewriter.getFunctionType( - TypeRange{rewriter.getI64Type(), rewriter.getI64Type(), - rewriter.getI64Type()}, - TypeRange{}); - rewriter.create(op.getLoc(), calleeName, TypeRange{}, - ValueRange{pipeValue, bufIdValue, modeValue}); - state.plannedDecls.push_back(PlannedDecl{calleeName.str(), funcType}); - rewriter.eraseOp(op); - return success(); - } - -private: - LoweringState &state; -}; - -template -class LowerBufDynSyncOpPattern final - : public OpConversionPattern { -public: - explicit LowerBufDynSyncOpPattern(TypeConverter &typeConverter, - MLIRContext *context, LoweringState &state) - : OpConversionPattern(typeConverter, context), - state(state) {} - - LogicalResult - matchAndRewrite(BufDynSyncOp op, typename BufDynSyncOp::Adaptor adaptor, - ConversionPatternRewriter &rewriter) const override { - PIPE pipe = PIPE::PIPE_UNASSIGNED; - if (auto pipeAttr = dyn_cast(op.getOpTypeAttr())) { - pipe = pipeAttr.getPipe(); - } else { - auto opTypeOr = parseSyncOpTypeLikeAttr(op.getOpTypeAttr()); - if (failed(opTypeOr)) - return rewriter.notifyMatchFailure( - op, "buffer sync expects pipe/sync_op_type/pipe_event_type attr"); - pipe = mapSyncOpTypeToPipe(*opTypeOr); - } - if (!isConcreteSyncPipe(pipe)) - return rewriter.notifyMatchFailure( - op, "buffer sync op_type cannot map to concrete pipe"); - - auto pipeImm = parsePipeImmediate(stringifyPIPE(pipe)); - if (!pipeImm) - return rewriter.notifyMatchFailure(op, "unsupported buffer sync pipe"); - - Value pipeValue = getI64Constant(rewriter, op.getLoc(), *pipeImm); - Value bufIdDyn = adaptor.getBufId(); - if (!bufIdDyn) - return rewriter.notifyMatchFailure( - op, "expected dynamic buf-id operand"); - Value bufIdValue = castIntegerLikeTo(op, bufIdDyn, rewriter.getI64Type()); - if (!bufIdValue) - return rewriter.notifyMatchFailure( - op, "failed to cast dynamic buf-id to i64"); - - bool isGetBuf = - std::is_same_v; - StringRef calleeName = - buildBufDynSyncCallee(op.getContext(), isGetBuf); - Value modeValue = - getI64Constant(rewriter, op.getLoc(), op.getModeAttr().getInt()); - auto funcType = rewriter.getFunctionType( - TypeRange{rewriter.getI64Type(), rewriter.getI64Type(), - rewriter.getI64Type()}, - TypeRange{}); - rewriter.create(op.getLoc(), calleeName, TypeRange{}, - ValueRange{pipeValue, bufIdValue, modeValue}); - state.plannedDecls.push_back(PlannedDecl{calleeName.str(), funcType}); - rewriter.eraseOp(op); - return success(); - } - -private: - LoweringState &state; -}; - -template -class LowerRuntimeQueryOpPattern final : public OpConversionPattern { -public: - explicit LowerRuntimeQueryOpPattern(TypeConverter &typeConverter, - MLIRContext *context, - LoweringState &state) - : OpConversionPattern(typeConverter, context), state(state) {} - - LogicalResult - matchAndRewrite(QueryOp op, typename QueryOp::Adaptor adaptor, - ConversionPatternRewriter &rewriter) const override { - (void)adaptor; - Type resultType = this->getTypeConverter()->convertType(op.getResult().getType()); - if (!resultType) - return rewriter.notifyMatchFailure(op, - "failed to convert runtime-query result type"); - - StringRef calleeName = buildRuntimeQueryCallee(op.getContext()); - auto funcType = rewriter.getFunctionType(TypeRange{}, TypeRange{resultType}); - auto call = rewriter.create(op.getLoc(), calleeName, - TypeRange{resultType}, ValueRange{}); - state.plannedDecls.push_back(PlannedDecl{calleeName.str(), funcType}); - rewriter.replaceOp(op, call.getResults()); - return success(); - } - -private: - LoweringState &state; -}; - -template -class LowerBlockRuntimeQueryOpPattern final - : public OpConversionPattern { -public: - explicit LowerBlockRuntimeQueryOpPattern(TypeConverter &typeConverter, - MLIRContext *context, - LoweringState &state) - : OpConversionPattern(typeConverter, context), state(state) {} - - LogicalResult - matchAndRewrite(QueryOp op, typename QueryOp::Adaptor adaptor, - ConversionPatternRewriter &rewriter) const override { - (void)adaptor; - Type resultType = - this->getTypeConverter()->convertType(op.getResult().getType()); - if (!resultType) - return rewriter.notifyMatchFailure( - op, "failed to convert block runtime-query result type"); - - auto funcOp = op->template getParentOfType(); - bool isSimtEntry = - funcOp && funcOp->hasAttr(pto::kPTOSimtEntryAttrName); - if (isSimtEntry && !resultType.isInteger(64)) - return rewriter.notifyMatchFailure( - op, "SIMT block runtime-query expects an i64 PTO result"); - - StringRef calleeName = - isSimtEntry ? buildSimtBlockQueryCallee(op.getContext()) - : buildRuntimeQueryCallee(op.getContext()); - Type callResultType = isSimtEntry ? rewriter.getI32Type() : resultType; - auto funcType = - rewriter.getFunctionType(TypeRange{}, TypeRange{callResultType}); - auto call = rewriter.create( - op.getLoc(), calleeName, TypeRange{callResultType}, ValueRange{}); - state.plannedDecls.push_back(PlannedDecl{calleeName.str(), funcType}); - - Value result = call.getResult(0); - if (isSimtEntry) - result = rewriter.create(op.getLoc(), resultType, result); - rewriter.replaceOp(op, result); - return success(); - } - -private: - LoweringState &state; -}; - -template -class LowerVoteOpPattern final : public OpConversionPattern { -public: - explicit LowerVoteOpPattern(TypeConverter &typeConverter, - MLIRContext *context, LoweringState &state) - : OpConversionPattern(typeConverter, context), state(state) {} - - LogicalResult - matchAndRewrite(VoteOp op, typename VoteOp::Adaptor adaptor, - ConversionPatternRewriter &rewriter) const override { - Type resultType = this->getTypeConverter()->convertType(op.getResult().getType()); - if (!resultType) - return rewriter.notifyMatchFailure(op, "failed to convert vote result type"); - - Type predType = this->getTypeConverter()->convertType(op.getPred().getType()); - if (!predType || predType != rewriter.getI1Type()) - return rewriter.notifyMatchFailure(op, "failed to convert vote predicate type"); - - StringRef calleeName = buildVoteCallee(op.getContext()); - auto funcType = rewriter.getFunctionType(TypeRange{predType}, TypeRange{resultType}); - auto call = rewriter.create(op.getLoc(), calleeName, - TypeRange{resultType}, - ValueRange{adaptor.getPred()}); - state.plannedDecls.push_back(PlannedDecl{calleeName.str(), funcType}); - rewriter.replaceOp(op, call.getResults()); - return success(); - } - -private: - LoweringState &state; -}; - -template -class LowerShuffleOpPattern final : public OpConversionPattern { -public: - explicit LowerShuffleOpPattern(TypeConverter &typeConverter, - MLIRContext *context, LoweringState &state) - : OpConversionPattern(typeConverter, context), state(state) {} - - LogicalResult - matchAndRewrite(ShuffleOp op, typename ShuffleOp::Adaptor adaptor, - ConversionPatternRewriter &rewriter) const override { - Type resultType = this->getTypeConverter()->convertType(op.getResult().getType()); - if (!resultType) - return rewriter.notifyMatchFailure(op, "failed to convert shuffle result type"); - - Type valueType = this->getTypeConverter()->convertType(op.getValue().getType()); - if (!valueType || valueType != resultType) - return rewriter.notifyMatchFailure(op, "unexpected converted shuffle operand type"); - - FailureOr calleeName = - buildShuffleCallee(op.getContext(), op.getValue().getType()); - if (failed(calleeName)) - return rewriter.notifyMatchFailure(op, "unsupported shuffle VPTO signature"); - - IntegerAttr widthAttr = op.getWidthAttr(); - Value controlValue; - unsigned controlMask = 0; - if constexpr (std::is_same_v) { - controlValue = adaptor.getIndex(); - controlMask = 0x1f; - } else if constexpr (std::is_same_v) { - controlValue = adaptor.getOffset(); - controlMask = 0; - } else if constexpr (std::is_same_v) { - controlValue = adaptor.getOffset(); - controlMask = 0x1f; - } else if constexpr (std::is_same_v) { - controlValue = adaptor.getMask(); - controlMask = 0x1f; - } - if (!controlValue) - return rewriter.notifyMatchFailure(op, "missing shuffle control operand"); - - Value control = buildShuffleControlValue( - rewriter, op.getLoc(), controlValue, widthAttr.getInt(), controlMask); - - Type i32Type = rewriter.getI32Type(); - auto funcType = rewriter.getFunctionType(TypeRange{resultType, i32Type}, - TypeRange{resultType}); - auto call = rewriter.create( - op.getLoc(), *calleeName, TypeRange{resultType}, - ValueRange{adaptor.getValue(), control}); - state.plannedDecls.push_back(PlannedDecl{calleeName->str(), funcType}); - rewriter.replaceOp(op, call.getResults()); - return success(); - } - -private: - LoweringState &state; -}; - -template -class LowerReduxOpPattern final : public OpConversionPattern { -public: - explicit LowerReduxOpPattern(TypeConverter &typeConverter, - MLIRContext *context, LoweringState &state) - : OpConversionPattern(typeConverter, context), state(state) {} - - LogicalResult - matchAndRewrite(ReduxOp op, typename ReduxOp::Adaptor adaptor, - ConversionPatternRewriter &rewriter) const override { - Type resultType = this->getTypeConverter()->convertType(op.getResult().getType()); - if (!resultType) - return rewriter.notifyMatchFailure(op, "failed to convert redux result type"); - - Type valueType = this->getTypeConverter()->convertType(op.getValue().getType()); - if (!valueType || valueType != resultType) - return rewriter.notifyMatchFailure(op, "unexpected converted redux operand type"); - - FailureOr calleeName = buildReduxCallee( - op.getContext(), op.getValue().getType(), op.getSignednessAttr()); - if (failed(calleeName)) - return rewriter.notifyMatchFailure(op, "unsupported redux VPTO signature"); - - auto funcType = rewriter.getFunctionType(TypeRange{resultType}, - TypeRange{resultType}); - auto call = rewriter.create(op.getLoc(), *calleeName, - TypeRange{resultType}, - ValueRange{adaptor.getValue()}); - state.plannedDecls.push_back(PlannedDecl{calleeName->str(), funcType}); - rewriter.replaceOp(op, call.getResults()); - return success(); - } - -private: - LoweringState &state; -}; - -template -class LowerAtomicBinaryOpPattern final : public OpConversionPattern { -public: - explicit LowerAtomicBinaryOpPattern(TypeConverter &typeConverter, - MLIRContext *context, - LoweringState &state) - : OpConversionPattern(typeConverter, context), state(state) {} - - LogicalResult - matchAndRewrite(AtomicOp op, typename AtomicOp::Adaptor adaptor, - ConversionPatternRewriter &rewriter) const override { - Type resultType = this->getTypeConverter()->convertType(op.getOld().getType()); - Type valueType = this->getTypeConverter()->convertType(op.getValue().getType()); - if (!resultType || !valueType || resultType != valueType) - return rewriter.notifyMatchFailure(op, - "unexpected atomic operand/result type"); - - Type ptrType = this->getTypeConverter()->convertType(op.getPtr().getType()); - if (!ptrType) - return rewriter.notifyMatchFailure(op, "failed to convert atomic pointer type"); - - FailureOr calleeName = buildAtomicCallee( - op.getContext(), op.getPtr().getType(), op.getValue().getType(), - op.getSignednessAttr()); - if (failed(calleeName)) - return rewriter.notifyMatchFailure(op, "unsupported atomic VPTO signature"); - - auto funcType = rewriter.getFunctionType( - TypeRange{ptrType, valueType, rewriter.getI32Type()}, - TypeRange{resultType}); - Value modeValue = getI32Constant( - rewriter, op.getLoc(), - static_cast(op.getL2cacheAttr() - ? op.getL2cacheAttr().getValue() - : pto::StL2Cache::NMFV)); - auto call = rewriter.create( - op.getLoc(), *calleeName, TypeRange{resultType}, - ValueRange{adaptor.getPtr(), adaptor.getValue(), modeValue}); - state.plannedDecls.push_back(PlannedDecl{calleeName->str(), funcType}); - rewriter.replaceOp(op, call.getResults()); - return success(); - } - -private: - LoweringState &state; -}; - -class LowerAtomicCasOpPattern final - : public OpConversionPattern { -public: - explicit LowerAtomicCasOpPattern(TypeConverter &typeConverter, - MLIRContext *context, - LoweringState &state) - : OpConversionPattern(typeConverter, context), - state(state) {} - - LogicalResult - matchAndRewrite(pto::AtomicCasOp op, pto::AtomicCasOp::Adaptor adaptor, - ConversionPatternRewriter &rewriter) const override { - Type resultType = this->getTypeConverter()->convertType(op.getOld().getType()); - Type compareType = - this->getTypeConverter()->convertType(op.getCompare().getType()); - Type valueType = this->getTypeConverter()->convertType(op.getValue().getType()); - if (!resultType || !compareType || !valueType || resultType != compareType || - resultType != valueType) - return rewriter.notifyMatchFailure(op, "unexpected atomic CAS type"); - - Type ptrType = this->getTypeConverter()->convertType(op.getPtr().getType()); - if (!ptrType) - return rewriter.notifyMatchFailure(op, "failed to convert atomic pointer type"); - - FailureOr calleeName = buildAtomicCallee( - op.getContext(), op.getPtr().getType(), op.getValue().getType(), - op.getSignednessAttr()); - if (failed(calleeName)) - return rewriter.notifyMatchFailure(op, "unsupported atomic CAS signature"); - - auto funcType = rewriter.getFunctionType( - TypeRange{ptrType, compareType, valueType, rewriter.getI32Type()}, - TypeRange{resultType}); - Value modeValue = getI32Constant( - rewriter, op.getLoc(), - static_cast(op.getL2cacheAttr() - ? op.getL2cacheAttr().getValue() - : pto::StL2Cache::NMFV)); - auto call = rewriter.create( - op.getLoc(), *calleeName, TypeRange{resultType}, - ValueRange{adaptor.getPtr(), adaptor.getCompare(), adaptor.getValue(), - modeValue}); - state.plannedDecls.push_back(PlannedDecl{calleeName->str(), funcType}); - rewriter.replaceOp(op, call.getResults()); - return success(); - } - -private: - LoweringState &state; -}; - -template -class LowerScalarIntrinsicOpPattern final : public OpConversionPattern { -public: - explicit LowerScalarIntrinsicOpPattern(TypeConverter &typeConverter, - MLIRContext *context, - LoweringState &state) - : OpConversionPattern(typeConverter, context), state(state) {} - - LogicalResult - matchAndRewrite(ScalarOp op, typename ScalarOp::Adaptor adaptor, - ConversionPatternRewriter &rewriter) const override { - SmallVector resultTypes; - if (failed(this->getTypeConverter()->convertTypes(op->getResultTypes(), - resultTypes))) - return rewriter.notifyMatchFailure(op, "failed to convert scalar result types"); - - SmallVector operandTypes; - if (failed(this->getTypeConverter()->convertTypes(op->getOperandTypes(), - operandTypes))) - return rewriter.notifyMatchFailure(op, "failed to convert scalar operand types"); - - StringRef calleeName = buildScalarIntrinsicCallee(op.getContext()); - auto funcType = rewriter.getFunctionType(operandTypes, resultTypes); - auto call = rewriter.create(op.getLoc(), calleeName, - resultTypes, adaptor.getOperands()); - state.plannedDecls.push_back(PlannedDecl{calleeName.str(), funcType}); - rewriter.replaceOp(op, call.getResults()); - return success(); - } - -private: - LoweringState &state; -}; - -class LowerMulhiOpPattern final : public OpConversionPattern { -public: - explicit LowerMulhiOpPattern(TypeConverter &typeConverter, - MLIRContext *context, LoweringState &state) - : OpConversionPattern(typeConverter, context), - state(state) {} - - LogicalResult - matchAndRewrite(pto::MulhiOp op, OpAdaptor adaptor, - ConversionPatternRewriter &rewriter) const override { - Type resultType = getTypeConverter()->convertType(op.getResult().getType()); - Type lhsType = getTypeConverter()->convertType(op.getLhs().getType()); - Type rhsType = getTypeConverter()->convertType(op.getRhs().getType()); - if (!resultType || !lhsType || !rhsType || lhsType != resultType || - rhsType != resultType) - return rewriter.notifyMatchFailure(op, "unexpected mulhi type"); - - pto::Signedness signedness = op.getSignednessAttr().getValue(); - FailureOr calleeName = - buildMulhiCallee(op.getContext(), op.getResult().getType(), signedness); - if (succeeded(calleeName)) { - auto funcType = - rewriter.getFunctionType(TypeRange{lhsType, rhsType}, TypeRange{resultType}); - auto call = rewriter.create( - op.getLoc(), *calleeName, TypeRange{resultType}, - ValueRange{adaptor.getLhs(), adaptor.getRhs()}); - state.plannedDecls.push_back(PlannedDecl{calleeName->str(), funcType}); - rewriter.replaceOp(op, call.getResults()); - return success(); - } - - if (!op.getResult().getType().isInteger(64) || - signedness != pto::Signedness::Signed) - return rewriter.notifyMatchFailure(op, "unsupported mulhi signature"); - - FailureOr unsignedCalleeName = - buildMulhiCallee(op.getContext(), op.getResult().getType(), - pto::Signedness::Unsigned); - if (failed(unsignedCalleeName)) - return rewriter.notifyMatchFailure(op, "unsupported mul64hi signature"); - - auto funcType = - rewriter.getFunctionType(TypeRange{lhsType, rhsType}, TypeRange{resultType}); - auto unsignedCall = rewriter.create( - op.getLoc(), *unsignedCalleeName, TypeRange{resultType}, - ValueRange{adaptor.getLhs(), adaptor.getRhs()}); - state.plannedDecls.push_back(PlannedDecl{unsignedCalleeName->str(), funcType}); - - Value zero = getI64Constant(rewriter, op.getLoc(), 0); - Value lhsNeg = rewriter.create( - op.getLoc(), LLVM::ICmpPredicate::slt, adaptor.getLhs(), zero); - Value rhsNeg = rewriter.create( - op.getLoc(), LLVM::ICmpPredicate::slt, adaptor.getRhs(), zero); - Value subRhs = rewriter.create( - op.getLoc(), unsignedCall.getResult(0), adaptor.getRhs()); - Value correctedLhs = rewriter.create( - op.getLoc(), resultType, lhsNeg, subRhs, unsignedCall.getResult(0)); - Value subLhs = rewriter.create( - op.getLoc(), correctedLhs, adaptor.getLhs()); - Value corrected = rewriter.create( - op.getLoc(), resultType, rhsNeg, subLhs, correctedLhs); - rewriter.replaceOp(op, corrected); - return success(); - } - -private: - LoweringState &state; -}; - -class LowerMulI32ToI64OpPattern final - : public OpConversionPattern { -public: - explicit LowerMulI32ToI64OpPattern(TypeConverter &typeConverter, - MLIRContext *context, - LoweringState &state) - : OpConversionPattern(typeConverter, context), - state(state) {} - - LogicalResult - matchAndRewrite(pto::MulI32ToI64Op op, OpAdaptor adaptor, - ConversionPatternRewriter &rewriter) const override { - Type resultType = getTypeConverter()->convertType(op.getResult().getType()); - Type lhsType = getTypeConverter()->convertType(op.getLhs().getType()); - Type rhsType = getTypeConverter()->convertType(op.getRhs().getType()); - if (!resultType || !lhsType || !rhsType) - return rewriter.notifyMatchFailure(op, "unexpected mul_i32toi64 type"); - - FailureOr calleeName = - buildMulI32ToI64Callee(op.getContext(), - op.getSignednessAttr().getValue()); - if (failed(calleeName)) - return rewriter.notifyMatchFailure(op, - "unsupported mul_i32toi64 signature"); - - auto funcType = - rewriter.getFunctionType(TypeRange{lhsType, rhsType}, TypeRange{resultType}); - auto call = rewriter.create( - op.getLoc(), *calleeName, TypeRange{resultType}, - ValueRange{adaptor.getLhs(), adaptor.getRhs()}); - state.plannedDecls.push_back(PlannedDecl{calleeName->str(), funcType}); - rewriter.replaceOp(op, call.getResults()); - return success(); - } - -private: - LoweringState &state; -}; - -class LowerSqrtOpPattern final : public OpConversionPattern { -public: - explicit LowerSqrtOpPattern(TypeConverter &typeConverter, - MLIRContext *context, LoweringState &state) - : OpConversionPattern(typeConverter, context), state(state) {} - - LogicalResult - matchAndRewrite(pto::SqrtOp op, pto::SqrtOp::Adaptor adaptor, - ConversionPatternRewriter &rewriter) const override { - Type resultType = this->getTypeConverter()->convertType(op.getResult().getType()); - Type valueType = this->getTypeConverter()->convertType(op.getValue().getType()); - if (!resultType || !valueType || valueType != resultType) - return rewriter.notifyMatchFailure(op, "unexpected sqrt operand/result type"); - - FailureOr calleeName = - buildSqrtCallee(op.getContext(), op.getValue().getType()); - if (failed(calleeName)) - return rewriter.notifyMatchFailure(op, "unsupported sqrt VPTO signature"); - - auto funcType = rewriter.getFunctionType(TypeRange{valueType}, - TypeRange{resultType}); - auto call = rewriter.create(op.getLoc(), *calleeName, - TypeRange{resultType}, - ValueRange{adaptor.getValue()}); - state.plannedDecls.push_back(PlannedDecl{calleeName->str(), funcType}); - rewriter.replaceOp(op, call.getResults()); - return success(); - } - -private: - LoweringState &state; -}; - -template -class LowerUnaryScalarMathOpPattern final : public OpConversionPattern { -public: - explicit LowerUnaryScalarMathOpPattern(TypeConverter &typeConverter, - MLIRContext *context, - LoweringState &state) - : OpConversionPattern(typeConverter, context), state(state) {} - - LogicalResult - matchAndRewrite(UnaryOp op, typename UnaryOp::Adaptor adaptor, - ConversionPatternRewriter &rewriter) const override { - Type resultType = this->getTypeConverter()->convertType(op.getResult().getType()); - Type valueType = this->getTypeConverter()->convertType(op.getValue().getType()); - if (!resultType || !valueType || valueType != resultType) - return rewriter.notifyMatchFailure(op, "unexpected unary scalar math type"); - - FailureOr calleeName = - buildUnaryScalarMathCallee(op.getContext(), op.getValue().getType()); - if (failed(calleeName)) - return rewriter.notifyMatchFailure(op, "unsupported unary scalar math signature"); - - auto funcType = rewriter.getFunctionType(TypeRange{valueType}, - TypeRange{resultType}); - auto call = rewriter.create(op.getLoc(), *calleeName, - TypeRange{resultType}, - ValueRange{adaptor.getValue()}); - state.plannedDecls.push_back(PlannedDecl{calleeName->str(), funcType}); - rewriter.replaceOp(op, call.getResults()); - return success(); - } - -private: - LoweringState &state; -}; - -template -class LowerBinaryScalarMathOpPattern final : public OpConversionPattern { -public: - explicit LowerBinaryScalarMathOpPattern(TypeConverter &typeConverter, - MLIRContext *context, - LoweringState &state) - : OpConversionPattern(typeConverter, context), state(state) {} - - LogicalResult - matchAndRewrite(BinaryOp op, typename BinaryOp::Adaptor adaptor, - ConversionPatternRewriter &rewriter) const override { - Type resultType = this->getTypeConverter()->convertType(op.getResult().getType()); - Type lhsType = this->getTypeConverter()->convertType(op.getLhs().getType()); - Type rhsType = this->getTypeConverter()->convertType(op.getRhs().getType()); - if (!resultType || !lhsType || !rhsType || - lhsType != rhsType || lhsType != resultType) - return rewriter.notifyMatchFailure(op, "unexpected binary scalar math type"); - - FailureOr calleeName = - buildBinaryScalarMathCallee(op.getContext(), op.getLhs().getType()); - if (failed(calleeName)) - return rewriter.notifyMatchFailure(op, "unsupported binary scalar math signature"); - - auto funcType = rewriter.getFunctionType(TypeRange{lhsType, rhsType}, - TypeRange{resultType}); - auto call = rewriter.create( - op.getLoc(), *calleeName, TypeRange{resultType}, - ValueRange{adaptor.getLhs(), adaptor.getRhs()}); - state.plannedDecls.push_back(PlannedDecl{calleeName->str(), funcType}); - rewriter.replaceOp(op, call.getResults()); - return success(); - } - -private: - LoweringState &state; -}; - -class LowerFmaOpPattern final : public OpConversionPattern { -public: - explicit LowerFmaOpPattern(TypeConverter &typeConverter, MLIRContext *context, - LoweringState &state) - : OpConversionPattern(typeConverter, context), state(state) {} - - LogicalResult - matchAndRewrite(pto::FmaOp op, pto::FmaOp::Adaptor adaptor, - ConversionPatternRewriter &rewriter) const override { - Type resultType = this->getTypeConverter()->convertType(op.getResult().getType()); - Type lhsType = this->getTypeConverter()->convertType(op.getLhs().getType()); - Type rhsType = this->getTypeConverter()->convertType(op.getRhs().getType()); - Type accType = this->getTypeConverter()->convertType(op.getAcc().getType()); - if (!resultType || !lhsType || !rhsType || !accType || - lhsType != rhsType || lhsType != accType || lhsType != resultType) - return rewriter.notifyMatchFailure(op, "unexpected fma scalar math type"); - - FailureOr calleeName = buildFmaCallee(op.getContext(), - op.getLhs().getType()); - if (failed(calleeName)) - return rewriter.notifyMatchFailure(op, "unsupported fma scalar signature"); - - auto funcType = rewriter.getFunctionType(TypeRange{lhsType, rhsType, accType}, - TypeRange{resultType}); - auto call = rewriter.create( - op.getLoc(), *calleeName, TypeRange{resultType}, - ValueRange{adaptor.getLhs(), adaptor.getRhs(), adaptor.getAcc()}); - state.plannedDecls.push_back(PlannedDecl{calleeName->str(), funcType}); - rewriter.replaceOp(op, call.getResults()); - return success(); - } - -private: - LoweringState &state; -}; - -class LowerConvertOpPattern final : public OpConversionPattern { -public: - explicit LowerConvertOpPattern(TypeConverter &typeConverter, - MLIRContext *context, LoweringState &state) - : OpConversionPattern(typeConverter, context), - state(state) {} - - LogicalResult - matchAndRewrite(pto::ConvertOp op, pto::ConvertOp::Adaptor adaptor, - ConversionPatternRewriter &rewriter) const override { - Type resultType = getTypeConverter()->convertType(op.getDst().getType()); - if (!resultType) - return rewriter.notifyMatchFailure(op, - "failed to convert result type"); - - FailureOr calleeName = - buildConvertCallee(op.getContext(), op.getSrc().getType(), - op.getDst().getType(), op.getSignednessAttr()); - if (failed(calleeName)) - return rewriter.notifyMatchFailure(op, "unsupported convert signature"); - - Value rounding = getI32Constant( - rewriter, op.getLoc(), static_cast(op.getRounding())); - Value saturation = getI32Constant( - rewriter, op.getLoc(), static_cast(op.getSaturation())); - - auto funcType = rewriter.getFunctionType( - TypeRange{adaptor.getSrc().getType(), rewriter.getI32Type(), - rewriter.getI32Type()}, - TypeRange{resultType}); - auto call = rewriter.create( - op.getLoc(), *calleeName, TypeRange{resultType}, - ValueRange{adaptor.getSrc(), rounding, saturation}); - state.plannedDecls.push_back(PlannedDecl{calleeName->str(), funcType}); - rewriter.replaceOp(op, call.getResults()); - return success(); - } - -private: - LoweringState &state; -}; - -class LowerGetVms4SrOpPattern final - : public OpConversionPattern { -public: - explicit LowerGetVms4SrOpPattern(TypeConverter &typeConverter, - MLIRContext *context, - LoweringState &state) - : OpConversionPattern(typeConverter, context), - state(state) {} - - LogicalResult - matchAndRewrite(pto::GetVms4SrOp op, pto::GetVms4SrOp::Adaptor adaptor, - ConversionPatternRewriter &rewriter) const override { - (void)adaptor; - SmallVector resultTypes; - if (failed(this->getTypeConverter()->convertTypes(op->getResultTypes(), - resultTypes)) || - resultTypes.size() != 4) - return rewriter.notifyMatchFailure( - op, "failed to convert get_vms4_sr result types"); - - StringRef calleeName = buildRuntimeQueryCallee( - op.getContext()); - auto funcType = - rewriter.getFunctionType(TypeRange{}, TypeRange{rewriter.getI64Type()}); - auto call = rewriter.create( - op.getLoc(), calleeName, TypeRange{rewriter.getI64Type()}, - ValueRange{}); - state.plannedDecls.push_back(PlannedDecl{calleeName.str(), funcType}); - - SmallVector counts; - counts.reserve(4); - Value raw = call.getResult(0); - for (unsigned i = 0; i < 4; ++i) { - Value shifted = raw; - if (i != 0) - shifted = rewriter.create( - op.getLoc(), raw, getI64Constant(rewriter, op.getLoc(), i * 16)); - counts.push_back(rewriter.create( - op.getLoc(), resultTypes[i], shifted)); - } - rewriter.replaceOp(op, counts); - return success(); - } - -private: - LoweringState &state; -}; - -template -class LowerBinaryI64PureOpPattern final : public OpConversionPattern { -public: - explicit LowerBinaryI64PureOpPattern(TypeConverter &typeConverter, - MLIRContext *context, - LoweringState &state) - : OpConversionPattern(typeConverter, context), state(state) {} - - LogicalResult - matchAndRewrite(BinaryOp op, typename BinaryOp::Adaptor adaptor, - ConversionPatternRewriter &rewriter) const override { - Type resultType = this->getTypeConverter()->convertType(op.getResult().getType()); - if (!resultType) - return rewriter.notifyMatchFailure(op, "failed to convert result type"); - - StringRef calleeName = buildBinaryI64PureCallee(op.getContext()); - auto funcType = - rewriter.getFunctionType(TypeRange{adaptor.getFirst().getType(), - adaptor.getSecond().getType()}, - TypeRange{resultType}); - auto call = rewriter.create( - op.getLoc(), calleeName, TypeRange{resultType}, - ValueRange{adaptor.getFirst(), adaptor.getSecond()}); - state.plannedDecls.push_back(PlannedDecl{calleeName.str(), funcType}); - rewriter.replaceOp(op, call.getResults()); - return success(); - } - -private: - LoweringState &state; -}; - -class ConvertVPTOUnrealizedCastOp final - : public OpConversionPattern { -public: - using OpConversionPattern::OpConversionPattern; - - LogicalResult - matchAndRewrite(UnrealizedConversionCastOp op, OpAdaptor adaptor, - ConversionPatternRewriter &rewriter) const override { - if (op->getNumOperands() != 1 || op->getNumResults() != 1) - return failure(); - if (!hasVPTOConvertibleType(op->getOperandTypes()) && - !hasVPTOConvertibleType(op->getResultTypes())) - return failure(); - - Type convertedResultType = - getTypeConverter()->convertType(op.getResult(0).getType()); - if (!convertedResultType) - return failure(); - - Value input = adaptor.getOperands().front(); - if (input.getType() != convertedResultType) - return failure(); - - rewriter.replaceOp(op, input); - return success(); - } -}; - -class ConvertPtoDeclareStructOp final - : public OpConversionPattern { -public: - using OpConversionPattern::OpConversionPattern; - - LogicalResult - matchAndRewrite(pto::DeclareStructOp op, OpAdaptor adaptor, - ConversionPatternRewriter &rewriter) const override { - (void)adaptor; - auto resultType = dyn_cast( - getTypeConverter()->convertType(op.getS().getType())); - if (!resultType) - return rewriter.notifyMatchFailure(op, - "expected LLVM pointer result type"); - auto structType = cast(op.getS().getType()); - Type storageType = getVPTOStructStorageType(structType, rewriter); - auto parentFunc = op->getParentOfType(); - if (!parentFunc) - return rewriter.notifyMatchFailure( - op, "expected struct declaration inside a function"); - - // A non-entry alloca is a dynamic stack allocation. Keep one stack slot per - // declaration per function invocation even when the declaration is nested - // in a loop or a region. - Value storage; - { - OpBuilder::InsertionGuard guard(rewriter); - Block &entryBlock = parentFunc.getBody().front(); - rewriter.setInsertionPointToStart(&entryBlock); - Value one = rewriter.create( - op.getLoc(), rewriter.getI64Type(), rewriter.getIndexAttr(1)); - storage = rewriter.create( - op.getLoc(), resultType, storageType, one, /*alignment=*/0); - } - rewriter.replaceOp(op, storage); - return success(); - } -}; - -class ConvertPtoStructGetOp final - : public OpConversionPattern { -public: - using OpConversionPattern::OpConversionPattern; - - LogicalResult - matchAndRewrite(pto::StructGetOp op, OpAdaptor adaptor, - ConversionPatternRewriter &rewriter) const override { - Type resultType = getTypeConverter()->convertType(op.getValue().getType()); - if (!resultType) - return rewriter.notifyMatchFailure(op, "could not convert result type"); - FailureOr address = getVPTOStructFieldAddress( - rewriter, op.getLoc(), adaptor.getS(), - cast(op.getS().getType()), op.getPath()); - if (failed(address)) - return rewriter.notifyMatchFailure(op, "invalid struct field path"); - rewriter.replaceOpWithNewOp( - op, resultType, *address, getNaturalByteAlignment(resultType)); - return success(); - } -}; - -class ConvertPtoStructSetOp final - : public OpConversionPattern { -public: - using OpConversionPattern::OpConversionPattern; - - LogicalResult - matchAndRewrite(pto::StructSetOp op, OpAdaptor adaptor, - ConversionPatternRewriter &rewriter) const override { - FailureOr address = getVPTOStructFieldAddress( - rewriter, op.getLoc(), adaptor.getS(), - cast(op.getS().getType()), op.getPath()); - if (failed(address)) - return rewriter.notifyMatchFailure(op, "invalid struct field path"); - rewriter.replaceOpWithNewOp( - op, adaptor.getValue(), *address, - getNaturalByteAlignment(adaptor.getValue().getType())); - return success(); - } -}; - -class ConvertArithSelectOp final : public OpConversionPattern { -public: - ConvertArithSelectOp(TypeConverter &typeConverter, MLIRContext *context) - : OpConversionPattern(typeConverter, context, - PatternBenefit(2)) {} - - LogicalResult - matchAndRewrite(arith::SelectOp op, OpAdaptor adaptor, - ConversionPatternRewriter &rewriter) const override { - if (!op.getCondition().getType().isInteger(1)) - return rewriter.notifyMatchFailure( - op, "only scalar i1 conditions supported for VPTO arith.select"); - - Type convertedResultType = - getTypeConverter()->convertType(op.getResult().getType()); - if (!convertedResultType) - return rewriter.notifyMatchFailure(op, "failed to convert result type"); - - Value trueValue = adaptor.getTrueValue(); - Value falseValue = adaptor.getFalseValue(); - if (trueValue.getType() != convertedResultType || - falseValue.getType() != convertedResultType) - return rewriter.notifyMatchFailure( - op, "converted true/false values must match result type"); - - rewriter.replaceOpWithNewOp( - op, convertedResultType, adaptor.getCondition(), trueValue, - falseValue); - return success(); - } -}; - -class ConvertPtoAddPtrOp final : public OpConversionPattern { -public: - using OpConversionPattern::OpConversionPattern; - - LogicalResult - matchAndRewrite(pto::AddPtrOp op, OpAdaptor adaptor, - ConversionPatternRewriter &rewriter) const override { - Type convertedResultType = getTypeConverter()->convertType(op.getResult().getType()); - auto llvmPtrType = dyn_cast(convertedResultType); - if (!llvmPtrType) - return rewriter.notifyMatchFailure(op, "expected LLVM pointer result type"); - - Value offset = adaptor.getOffset(); - if (offset.getType().isIndex()) - offset = rewriter.create(op.getLoc(), - rewriter.getI64Type(), offset); - - auto gep = rewriter.create( - op.getLoc(), llvmPtrType, - normalizeGEPElementTypeForLLVMLowering( - cast(op.getPtr().getType()).getElementType(), - rewriter), - adaptor.getPtr(), ValueRange{offset}); - rewriter.replaceOp(op, gep.getResult()); - return success(); - } -}; - -class ConvertPtoCastPtrOp final : public OpConversionPattern { -public: - using OpConversionPattern::OpConversionPattern; - - LogicalResult - matchAndRewrite(pto::CastPtrOp op, OpAdaptor adaptor, - ConversionPatternRewriter &rewriter) const override { - Type convertedResultType = - getTypeConverter()->convertType(op.getResult().getType()); - if (!convertedResultType) - return rewriter.notifyMatchFailure(op, - "could not convert castptr result type"); - - Value input = adaptor.getInput(); - Type inputType = input.getType(); - if (inputType == convertedResultType) { - rewriter.replaceOp(op, input); - return success(); - } - - if (auto llvmPtrType = dyn_cast(convertedResultType)) { - if (isa(inputType)) { - rewriter.replaceOpWithNewOp(op, llvmPtrType, input); - return success(); - } - auto sourcePtrType = dyn_cast(inputType); - if (!sourcePtrType) - return rewriter.notifyMatchFailure(op, - "expected integer or LLVM pointer input"); - if (sourcePtrType.getAddressSpace() == llvmPtrType.getAddressSpace()) { - rewriter.replaceOpWithNewOp(op, llvmPtrType, input); - return success(); - } - return rewriter.notifyMatchFailure( - op, "cross-address-space ptr casts are unsupported"); - } - - if (auto resultIntType = dyn_cast(convertedResultType)) { - if (isa(inputType)) { - rewriter.replaceOpWithNewOp(op, resultIntType, input); - return success(); - } - } - - return rewriter.notifyMatchFailure(op, "unsupported castptr conversion"); - } -}; - -class ConvertPtoLoadScalarOp final - : public OpConversionPattern { -public: - using OpConversionPattern::OpConversionPattern; - - LogicalResult - matchAndRewrite(pto::LoadScalarOp op, OpAdaptor adaptor, - ConversionPatternRewriter &rewriter) const override { - auto llvmPtrType = dyn_cast(adaptor.getPtr().getType()); - if (!llvmPtrType) - return rewriter.notifyMatchFailure(op, "expected LLVM pointer operand"); - - Type convertedValueType = - getTypeConverter()->convertType(op.getValue().getType()); - if (!convertedValueType) - return rewriter.notifyMatchFailure(op, - "could not convert load_scalar result type"); - - Value offset = adaptor.getOffset(); - if (offset.getType().isIndex()) - offset = rewriter.create(op.getLoc(), - rewriter.getI64Type(), offset); - - Value elemPtr = adaptor.getPtr(); - if (!matchPattern(offset, m_Zero())) { - elemPtr = rewriter.create(op.getLoc(), llvmPtrType, - normalizeGEPElementTypeForLLVMLowering( - convertedValueType, rewriter), - adaptor.getPtr(), - ValueRange{offset}); - } - - rewriter.replaceOpWithNewOp( - op, convertedValueType, elemPtr, - getNaturalByteAlignment(convertedValueType)); - return success(); - } -}; - -class ConvertPtoStoreScalarOp final - : public OpConversionPattern { -public: - using OpConversionPattern::OpConversionPattern; - - LogicalResult - matchAndRewrite(pto::StoreScalarOp op, OpAdaptor adaptor, - ConversionPatternRewriter &rewriter) const override { - auto llvmPtrType = dyn_cast(adaptor.getPtr().getType()); - if (!llvmPtrType) - return rewriter.notifyMatchFailure(op, "expected LLVM pointer operand"); - - Value offset = adaptor.getOffset(); - if (offset.getType().isIndex()) - offset = rewriter.create(op.getLoc(), - rewriter.getI64Type(), offset); - - Value elemPtr = adaptor.getPtr(); - if (!matchPattern(offset, m_Zero())) { - elemPtr = rewriter.create(op.getLoc(), llvmPtrType, - normalizeGEPElementTypeForLLVMLowering( - adaptor.getValue().getType(), - rewriter), - adaptor.getPtr(), ValueRange{offset}); - } - - rewriter.create(op.getLoc(), adaptor.getValue(), elemPtr, - getNaturalByteAlignment(adaptor.getValue().getType())); - rewriter.eraseOp(op); - return success(); - } -}; - -class ConvertPtoLoadOp final : public OpConversionPattern { -public: - ConvertPtoLoadOp(TypeConverter &typeConverter, MLIRContext *context, - LoweringState &) - : OpConversionPattern(typeConverter, context) {} - - LogicalResult - matchAndRewrite(pto::PTOLoadOp op, OpAdaptor adaptor, - ConversionPatternRewriter &rewriter) const override { - auto llvmPtrType = dyn_cast(adaptor.getPtr().getType()); - if (!llvmPtrType) - return rewriter.notifyMatchFailure(op, "expected LLVM pointer operand"); - - Type convertedValueType = - getTypeConverter()->convertType(op.getValue().getType()); - if (!convertedValueType) - return rewriter.notifyMatchFailure(op, "could not convert load result type"); - - Value offset = adaptor.getOffset(); - if (offset.getType().isIndex()) - offset = rewriter.create(op.getLoc(), - rewriter.getI64Type(), offset); - - Value elemPtr = adaptor.getPtr(); - if (!matchPattern(offset, m_Zero())) { - elemPtr = rewriter.create(op.getLoc(), llvmPtrType, - convertedValueType, - adaptor.getPtr(), - ValueRange{offset}); - } - - rewriter.replaceOpWithNewOp( - op, convertedValueType, elemPtr, - getNaturalByteAlignment(convertedValueType)); - return success(); - } -}; - -static Type getLdgCallResultType(Type valueType, Type convertedValueType, - ConversionPatternRewriter &rewriter) { - if (auto intType = dyn_cast(valueType)) { - unsigned width = intType.getWidth(); - if (width == 8 || width == 16) - return rewriter.getI32Type(); - return convertedValueType; - } - if (valueType.isF16() || valueType.isBF16() || valueType.isF32()) - return rewriter.getI32Type(); - if (valueType.isF64()) - return rewriter.getI64Type(); - if (pto::isPTOFloat8Type(valueType) || pto::isPTOHiFloat8Type(valueType)) - return rewriter.getI32Type(); - if (pto::isPTOPackedLdgStgVectorType(valueType)) { - unsigned totalBits = pto::getPTOPackedLdgStgTotalBits(valueType); - if (totalBits == 16) - return rewriter.getI32Type(); - if (totalBits == 32) - return rewriter.getI32Type(); - if (totalBits == 64) - return rewriter.getI64Type(); - } - return convertedValueType; -} - -static Value convertLdgCallResult(Location loc, Type valueType, - Type convertedValueType, Value callResult, - ConversionPatternRewriter &rewriter) { - if (auto intType = dyn_cast(valueType)) { - unsigned width = intType.getWidth(); - if (width == 8 || width == 16) - return rewriter.create( - loc, rewriter.getIntegerType(width), callResult); - return callResult; - } - - if (valueType.isF16() || valueType.isBF16()) { - Value payload = - rewriter.create(loc, rewriter.getI16Type(), callResult); - return rewriter.create(loc, convertedValueType, payload); - } - if (valueType.isF32() || valueType.isF64()) - return rewriter.create(loc, convertedValueType, - callResult); - if (pto::isPTOFloat8Type(valueType) || pto::isPTOHiFloat8Type(valueType)) { - Value payload = - rewriter.create(loc, rewriter.getI8Type(), callResult); - return rewriter.create(loc, convertedValueType, payload); - } - if (pto::isPTOPackedLdgStgVectorType(valueType)) { - unsigned totalBits = pto::getPTOPackedLdgStgTotalBits(valueType); - if (totalBits == 16) { - Value trunc = rewriter.create( - loc, rewriter.getI16Type(), callResult); - return rewriter.create(loc, convertedValueType, trunc); - } - return rewriter.create(loc, convertedValueType, - callResult); - } - return callResult; -} - -class ConvertPtoLdgOp final : public OpConversionPattern { -public: - ConvertPtoLdgOp(TypeConverter &typeConverter, MLIRContext *context, - LoweringState &state) - : OpConversionPattern(typeConverter, context), - state(state) {} - - LogicalResult - matchAndRewrite(pto::PTOLdgOp op, OpAdaptor adaptor, - ConversionPatternRewriter &rewriter) const override { - auto llvmPtrType = dyn_cast(adaptor.getPtr().getType()); - if (!llvmPtrType) - return rewriter.notifyMatchFailure(op, "expected LLVM pointer operand"); - - Type convertedValueType = - getTypeConverter()->convertType(op.getValue().getType()); - if (!convertedValueType) - return rewriter.notifyMatchFailure(op, "could not convert ldg result type"); - - Value offset = adaptor.getOffset(); - if (offset.getType().isIndex()) - offset = rewriter.create(op.getLoc(), - rewriter.getI64Type(), offset); - - Value elemPtr = adaptor.getPtr(); - if (!matchPattern(offset, m_Zero())) { - elemPtr = rewriter.create(op.getLoc(), llvmPtrType, - normalizeGEPElementTypeForLLVMLowering( - convertedValueType, rewriter), - adaptor.getPtr(), - ValueRange{offset}); - } - - auto ptrTy = cast(op.getPtr().getType()); - FailureOr ptr = reinterpretPointerToAddrSpace( - op, elemPtr, - static_cast(ptrTy.getMemorySpace().getAddressSpace())); - if (failed(ptr)) - return rewriter.notifyMatchFailure(op, "failed to map ldg pointer"); - - pto::L1Cache l1cache = op.getL1cacheAttr() - ? op.getL1cacheAttr().getValue() - : pto::L1Cache::Cache; - FailureOr calleeName = buildL1CacheLoadCallee( - op.getContext(), op.getValue().getType(), l1cache); - if (failed(calleeName)) - return rewriter.notifyMatchFailure(op, "unsupported ldg signature"); - - pto::LdL2Cache mode = op.getL2cacheAttr() - ? op.getL2cacheAttr().getValue() - : pto::LdL2Cache::NMFV; - Value modeValue = - getI32Constant(rewriter, op.getLoc(), static_cast(mode)); - Type callResultType = getLdgCallResultType(op.getValue().getType(), - convertedValueType, rewriter); - auto funcType = - rewriter.getFunctionType(TypeRange{ptr->getType(), rewriter.getI32Type()}, - TypeRange{callResultType}); - auto call = rewriter.create( - op.getLoc(), *calleeName, TypeRange{callResultType}, - ValueRange{*ptr, modeValue}); - state.plannedDecls.push_back(PlannedDecl{calleeName->str(), funcType}); - Value result = convertLdgCallResult(op.getLoc(), op.getValue().getType(), - convertedValueType, call.getResult(0), - rewriter); - rewriter.replaceOp(op, result); - return success(); - } - -private: - LoweringState &state; -}; - -class ConvertPtoStoreOp final : public OpConversionPattern { -public: - ConvertPtoStoreOp(TypeConverter &typeConverter, MLIRContext *context, - LoweringState &) - : OpConversionPattern(typeConverter, context) {} - - LogicalResult - matchAndRewrite(pto::PTOStoreOp op, OpAdaptor adaptor, - ConversionPatternRewriter &rewriter) const override { - auto llvmPtrType = dyn_cast(adaptor.getPtr().getType()); - if (!llvmPtrType) - return rewriter.notifyMatchFailure(op, "expected LLVM pointer operand"); - - Value offset = adaptor.getOffset(); - if (offset.getType().isIndex()) - offset = rewriter.create(op.getLoc(), - rewriter.getI64Type(), offset); - - Value elemPtr = adaptor.getPtr(); - if (!matchPattern(offset, m_Zero())) { - elemPtr = rewriter.create(op.getLoc(), llvmPtrType, - adaptor.getValue().getType(), - adaptor.getPtr(), ValueRange{offset}); - } - - rewriter.replaceOpWithNewOp( - op, adaptor.getValue(), elemPtr, - getNaturalByteAlignment(adaptor.getValue().getType())); - return success(); - } -}; - -static Value convertStgValue(Location loc, Type valueType, Value value, - ConversionPatternRewriter &rewriter) { - if (auto intType = dyn_cast(valueType)) { - unsigned width = intType.getWidth(); - if (width == 8) - return rewriter.create(loc, rewriter.getI32Type(), value); - if (width == 16) - return rewriter.create(loc, rewriter.getF16Type(), value); - return value; - } - - if (pto::isPTOFloat8Type(valueType) || pto::isPTOHiFloat8Type(valueType)) { - Value payload = - rewriter.create(loc, rewriter.getI8Type(), value); - return rewriter.create(loc, rewriter.getI32Type(), payload); - } - if (valueType.isBF16()) - return rewriter.create(loc, rewriter.getF16Type(), value); - if (valueType.isF32()) - return rewriter.create(loc, rewriter.getI32Type(), value); - if (valueType.isF64()) - return rewriter.create(loc, rewriter.getI64Type(), value); - if (pto::isPTOPackedLdgStgVectorType(valueType)) { - unsigned totalBits = pto::getPTOPackedLdgStgTotalBits(valueType); - if (totalBits == 16) - return rewriter.create(loc, rewriter.getF16Type(), - value); - if (totalBits == 32) - return rewriter.create(loc, rewriter.getI32Type(), - value); - if (totalBits == 64) - return rewriter.create(loc, rewriter.getI64Type(), - value); - } - return value; -} - -class ConvertPtoStgOp final : public OpConversionPattern { -public: - ConvertPtoStgOp(TypeConverter &typeConverter, MLIRContext *context, - LoweringState &state) - : OpConversionPattern(typeConverter, context), - state(state) {} - - LogicalResult - matchAndRewrite(pto::PTOStgOp op, OpAdaptor adaptor, - ConversionPatternRewriter &rewriter) const override { - auto llvmPtrType = dyn_cast(adaptor.getPtr().getType()); - if (!llvmPtrType) - return rewriter.notifyMatchFailure(op, "expected LLVM pointer operand"); - - Value offset = adaptor.getOffset(); - if (offset.getType().isIndex()) - offset = rewriter.create(op.getLoc(), - rewriter.getI64Type(), offset); - - Value elemPtr = adaptor.getPtr(); - if (!matchPattern(offset, m_Zero())) { - elemPtr = rewriter.create(op.getLoc(), llvmPtrType, - normalizeGEPElementTypeForLLVMLowering( - adaptor.getValue().getType(), - rewriter), - adaptor.getPtr(), ValueRange{offset}); - } - - auto ptrTy = cast(op.getPtr().getType()); - FailureOr ptr = reinterpretPointerToAddrSpace( - op, elemPtr, - static_cast(ptrTy.getMemorySpace().getAddressSpace())); - if (failed(ptr)) - return rewriter.notifyMatchFailure(op, "failed to map stg pointer"); - - pto::L1Cache l1cache = op.getL1cacheAttr() - ? op.getL1cacheAttr().getValue() - : pto::L1Cache::Cache; - FailureOr calleeName = buildL1CacheStoreCallee( - op.getContext(), op.getValue().getType(), l1cache); - if (failed(calleeName)) - return rewriter.notifyMatchFailure(op, "unsupported stg signature"); - - pto::StL2Cache mode = op.getL2cacheAttr() - ? op.getL2cacheAttr().getValue() - : pto::StL2Cache::NMFV; - Value modeValue = - getI32Constant(rewriter, op.getLoc(), static_cast(mode)); - Value storedValue = convertStgValue(op.getLoc(), op.getValue().getType(), - adaptor.getValue(), rewriter); - auto funcType = - rewriter.getFunctionType(TypeRange{ptr->getType(), storedValue.getType(), - rewriter.getI32Type()}, - TypeRange{}); - rewriter.create( - op.getLoc(), *calleeName, TypeRange{}, - ValueRange{*ptr, storedValue, modeValue}); - state.plannedDecls.push_back(PlannedDecl{calleeName->str(), funcType}); - rewriter.eraseOp(op); - return success(); - } - -private: - LoweringState &state; -}; - -static std::string buildLdDevCalleeName(unsigned width) { - return "llvm.hivm.LD.DEV.u" + std::to_string(width) + ".GM"; -} - -static std::string buildStDevCalleeName(unsigned width) { - return "llvm.hivm.ST.DEV.u" + std::to_string(width); -} - -class ConvertPtoLdDevOp final : public OpConversionPattern { -public: - ConvertPtoLdDevOp(TypeConverter &typeConverter, MLIRContext *context, - LoweringState &state) - : OpConversionPattern(typeConverter, context), - state(state) {} - - LogicalResult - matchAndRewrite(pto::PTOLdDevOp op, OpAdaptor adaptor, - ConversionPatternRewriter &rewriter) const override { - auto llvmPtrType = dyn_cast(adaptor.getPtr().getType()); - if (!llvmPtrType) - return rewriter.notifyMatchFailure(op, "expected LLVM pointer operand"); - - auto valueType = dyn_cast(op.getValue().getType()); - if (!valueType) - return rewriter.notifyMatchFailure(op, "expected integer result type"); - - Value offset = adaptor.getOffset(); - if (offset.getType().isIndex()) - offset = rewriter.create(op.getLoc(), - rewriter.getI64Type(), offset); - - Type convertedValueType = - getTypeConverter()->convertType(op.getValue().getType()); - if (!convertedValueType) - return rewriter.notifyMatchFailure(op, - "could not convert ld_dev result type"); - - Value elemPtr = adaptor.getPtr(); - if (!matchPattern(offset, m_Zero())) { - elemPtr = rewriter.create( - op.getLoc(), llvmPtrType, - normalizeGEPElementTypeForLLVMLowering(convertedValueType, rewriter), - adaptor.getPtr(), ValueRange{offset}); - } - - FailureOr gmPtr = reinterpretPointerToAddrSpace( - op, elemPtr, static_cast(pto::AddressSpace::GM)); - if (failed(gmPtr)) - return rewriter.notifyMatchFailure(op, "failed to map ld_dev GM pointer"); - - std::string calleeName = buildLdDevCalleeName(valueType.getWidth()); - Value intrinsicOffset = getI64Constant(rewriter, op.getLoc(), 0); - auto funcType = rewriter.getFunctionType( - TypeRange{gmPtr->getType(), rewriter.getI64Type()}, - TypeRange{rewriter.getI64Type()}); - auto call = rewriter.create( - op.getLoc(), calleeName, TypeRange{rewriter.getI64Type()}, - ValueRange{*gmPtr, intrinsicOffset}); - state.plannedDecls.push_back(PlannedDecl{calleeName, funcType}); - - Value result = call.getResult(0); - if (valueType.getWidth() < 64) - result = rewriter.create(op.getLoc(), convertedValueType, - result); - rewriter.replaceOp(op, result); - return success(); - } - -private: - LoweringState &state; -}; - -class ConvertPtoStDevOp final : public OpConversionPattern { -public: - ConvertPtoStDevOp(TypeConverter &typeConverter, MLIRContext *context, - LoweringState &state) - : OpConversionPattern(typeConverter, context), - state(state) {} - - LogicalResult - matchAndRewrite(pto::PTOStDevOp op, OpAdaptor adaptor, - ConversionPatternRewriter &rewriter) const override { - auto llvmPtrType = dyn_cast(adaptor.getPtr().getType()); - if (!llvmPtrType) - return rewriter.notifyMatchFailure(op, "expected LLVM pointer operand"); - - auto valueType = dyn_cast(op.getValue().getType()); - if (!valueType) - return rewriter.notifyMatchFailure(op, "expected integer value type"); - - Value offset = adaptor.getOffset(); - if (offset.getType().isIndex()) - offset = rewriter.create(op.getLoc(), - rewriter.getI64Type(), offset); - - Value elemPtr = adaptor.getPtr(); - if (!matchPattern(offset, m_Zero())) { - elemPtr = rewriter.create( - op.getLoc(), llvmPtrType, - normalizeGEPElementTypeForLLVMLowering(adaptor.getValue().getType(), - rewriter), - adaptor.getPtr(), ValueRange{offset}); - } - - FailureOr gmPtr = reinterpretPointerToAddrSpace( - op, elemPtr, static_cast(pto::AddressSpace::GM)); - if (failed(gmPtr)) - return rewriter.notifyMatchFailure(op, "failed to map st_dev GM pointer"); - - Value payload = adaptor.getValue(); - if (valueType.getWidth() < 64) - payload = rewriter.create(op.getLoc(), - rewriter.getI64Type(), payload); - - std::string calleeName = buildStDevCalleeName(valueType.getWidth()); - Value intrinsicOffset = getI64Constant(rewriter, op.getLoc(), 0); - auto funcType = rewriter.getFunctionType( - TypeRange{rewriter.getI64Type(), gmPtr->getType(), - rewriter.getI64Type()}, - TypeRange{}); - rewriter.create(op.getLoc(), calleeName, TypeRange{}, - ValueRange{payload, *gmPtr, intrinsicOffset}); - state.plannedDecls.push_back(PlannedDecl{calleeName, funcType}); - rewriter.eraseOp(op); - return success(); - } - -private: - LoweringState &state; -}; - -class ConvertVPTOTypedCarrierOp final : public ConversionPattern { -public: - ConvertVPTOTypedCarrierOp(TypeConverter &typeConverter, MLIRContext *context) - : ConversionPattern(typeConverter, MatchAnyOpTypeTag(), 1, context) {} - - LogicalResult - matchAndRewrite(Operation *op, ArrayRef operands, - ConversionPatternRewriter &rewriter) const override { - if (isa(op)) - return failure(); - Type propertyType; - if (auto allocaOp = dyn_cast(op)) - propertyType = allocaOp.getElemType(); - else if (auto gepOp = dyn_cast(op)) - propertyType = gepOp.getElemType(); - if (!hasVPTOConvertibleType(op->getOperandTypes()) && - !hasVPTOConvertibleType(op->getResultTypes()) && - !hasVPTOConvertibleType(propertyType)) - return failure(); - if (op->getNumRegions() != 0) - return rewriter.notifyMatchFailure( - op, "region ops with VPTO types are handled structurally"); - - SmallVector convertedResultTypes; - if (failed(typeConverter->convertTypes(op->getResultTypes(), - convertedResultTypes))) - return rewriter.notifyMatchFailure(op, "failed to convert result types"); - OperationState state(op->getLoc(), op->getName()); - state.addOperands(operands); - state.addTypes(convertedResultTypes); - state.addAttributes(op->getAttrs()); - state.addSuccessors(op->getSuccessors()); - state.propertiesAttr = op->getPropertiesAsAttribute(); - Operation *converted = rewriter.create(state); - if (propertyType) { - Type convertedPropertyType = typeConverter->convertType(propertyType); - if (!convertedPropertyType) - return rewriter.notifyMatchFailure( - op, "failed to convert LLVM element type"); - if (auto allocaOp = dyn_cast(converted)) - allocaOp.setElemType(convertedPropertyType); - else - cast(converted).setElemType(convertedPropertyType); - } - rewriter.replaceOp(op, converted->getResults()); - return success(); - } -}; - -static void populateVPTOOpLoweringPatterns(VPTOTypeConverter &typeConverter, - RewritePatternSet &patterns, - LoweringState &state) { - patterns.add, - LowerUnaryMaskedOpPattern, - LowerUnaryMaskedOpPattern, - LowerUnaryMaskedOpPattern, - LowerUnaryMaskedOpPattern, - LowerUnaryMaskedOpPattern, - LowerUnaryMaskedOpPattern, - LowerVsqzOpPattern, LowerVusqzOpPattern, - LowerVmulaOpPattern, LowerVmullOpPattern, - LowerBinaryMaskedOpPattern, - LowerBinaryMaskedOpPattern, - LowerBinaryMaskedOpPattern, - LowerBinaryMaskedOpPattern, - LowerBinaryMaskedOpPattern, - LowerBinaryMaskedOpPattern, - LowerBinaryMaskedOpPattern, - LowerBinaryMaskedOpPattern, - LowerBinaryMaskedOpPattern, - LowerTernaryMaskedOpPattern, - LowerBinaryMaskedOpPattern, - LowerCarryBinaryOpPattern, - LowerCarryBinaryOpPattern, - LowerCarryBinaryOpPattern, - LowerCarryBinaryOpPattern, - LowerBinaryMaskedOpPattern, - LowerBinaryMaskedOpPattern, - LowerVecScalarMaskedOpPattern, - LowerVecScalarMaskedOpPattern, - LowerVecScalarMaskedOpPattern, - LowerVecScalarMaskedOpPattern, - LowerVecScalarMaskedOpPattern, - LowerVecScalarMaskedOpPattern, - LowerVecScalarMaskedOpPattern, - LowerWideningReductionUnaryOpPattern, - LowerReductionUnaryOpPattern, - LowerReductionUnaryOpPattern, - LowerReductionUnaryOpPattern, - LowerReductionUnaryOpPattern, - LowerReductionUnaryOpPattern, - LowerReductionUnaryOpPattern, - LowerHistogramOpPattern, - LowerHistogramOpPattern, - LowerExtremaPredicateOpPattern, - LowerExtremaPredicateOpPattern, - LowerVdupOpPattern, - LowerVbrOpPattern, - LowerPredicatePackOpPattern, - LowerPredicatePackOpPattern, - LowerVselOpPattern, LowerVselrOpPattern, LowerPnotOpPattern, - LowerPredicateMaskBinaryOpPattern, - LowerPredicateMaskBinaryOpPattern, - LowerPredicateMaskBinaryOpPattern, - LowerPredicateMaskBinaryOpPattern, - LowerPredicatePairReorderOpPattern, - LowerPredicatePairReorderOpPattern, - LowerPredicatePairReorderOpPattern, - LowerPredicatePairReorderOpPattern, - LowerPredicatePairReorderOpPattern, - LowerPredicatePairReorderOpPattern, - LowerUnpackOpPattern, - LowerUnpackOpPattern, - LowerVpackOpPattern, - LowerInterleaveOpPattern, - LowerInterleaveOpPattern, - LowerCmpOpPattern, - LowerCmpOpPattern, - LowerPltOpPattern, - LowerPltOpPattern, - LowerPltOpPattern, - LowerPltmOpPattern, - LowerPltmOpPattern, - LowerPltmOpPattern, - LowerPsetOpPattern, - LowerPsetOpPattern, - LowerPsetOpPattern, - LowerPgeOpPattern, - LowerPgeOpPattern, - LowerPgeOpPattern, - LowerRuntimeQueryOpPattern, - LowerGetVms4SrOpPattern, - LowerRuntimeQueryOpPattern, - LowerRuntimeQueryOpPattern, - LowerRuntimeQueryOpPattern, - LowerRuntimeQueryOpPattern, - LowerRuntimeQueryOpPattern, - LowerRuntimeQueryOpPattern, - LowerRuntimeQueryOpPattern, - LowerRuntimeQueryOpPattern, - LowerRuntimeQueryOpPattern, - LowerRuntimeQueryOpPattern, - LowerRuntimeQueryOpPattern, - LowerRuntimeQueryOpPattern, - LowerRuntimeQueryOpPattern, - LowerRuntimeQueryOpPattern, - LowerRuntimeQueryOpPattern, - LowerRuntimeQueryOpPattern, - LowerRuntimeQueryOpPattern, - LowerRuntimeQueryOpPattern, - LowerRuntimeQueryOpPattern, - LowerRuntimeQueryOpPattern, - LowerRuntimeQueryOpPattern, - LowerVoteOpPattern, - LowerVoteOpPattern, - LowerVoteOpPattern, - LowerVoteOpPattern, - LowerShuffleOpPattern, - LowerShuffleOpPattern, - LowerShuffleOpPattern, - LowerShuffleOpPattern, - LowerReduxOpPattern, - LowerReduxOpPattern, - LowerReduxOpPattern, - LowerAtomicCasOpPattern, - LowerAtomicBinaryOpPattern, - LowerAtomicBinaryOpPattern, - LowerAtomicBinaryOpPattern, - LowerAtomicBinaryOpPattern, - LowerAtomicBinaryOpPattern, - LowerAtomicBinaryOpPattern, - LowerAtomicBinaryOpPattern, - LowerAtomicBinaryOpPattern, - LowerTrapOpPattern, - LowerScalarIntrinsicOpPattern, - LowerMulhiOpPattern, - LowerMulI32ToI64OpPattern, - LowerSqrtOpPattern, - LowerUnaryScalarMathOpPattern, - LowerUnaryScalarMathOpPattern, - LowerUnaryScalarMathOpPattern, - LowerUnaryScalarMathOpPattern, - LowerUnaryScalarMathOpPattern, - LowerUnaryScalarMathOpPattern, - LowerUnaryScalarMathOpPattern, - LowerBinaryScalarMathOpPattern, - LowerBinaryScalarMathOpPattern, - LowerBinaryScalarMathOpPattern, - LowerFmaOpPattern, - LowerConvertOpPattern, - LowerSimtFenceOpPattern, - LowerSimtFenceOpPattern, - LowerSimtFenceOpPattern, - LowerKeepOpPattern, - LowerResumeOpPattern, - LowerBinaryI64PureOpPattern, - LowerBinaryI64PureOpPattern, - LowerSetLoopConfigOpPattern, - LowerSetLoopConfigOpPattern, - LowerSetLoopConfigOpPattern, - LowerSetLoopConfigOpPattern, - LowerSetLoopConfigOpPattern, - LowerSetLoopConfigOpPattern, - LowerSetLoopConfigOpPattern, - LowerSetLoopConfigOpPattern, - LowerUnaryI64ConfigOpPattern, - LowerStoreVfSimtInfoOpPattern, - LowerUnaryConfigOpPattern, - LowerUnaryI64ConfigOpPattern, - LowerUnaryI64ConfigOpPattern, - LowerUnaryI64ConfigOpPattern, - LowerUnaryI64ConfigOpPattern, - LowerUnaryI64ConfigOpPattern, - LowerUnaryI64ConfigOpPattern, - LowerUnaryI64ConfigOpPattern, - LowerUnaryI64ConfigOpPattern, - LowerUnaryI64ConfigOpPattern, - LowerUnaryI64ConfigOpPattern, - LowerNullaryConfigOpPattern, - LowerNullaryConfigOpPattern, - LowerPipeEventSyncOpPattern, - LowerPipeEventSyncOpPattern, - LowerPipeEventDynSyncOpPattern, - LowerPipeEventDynSyncOpPattern, - LowerBarrierOpPattern, LowerMemBarOpPattern, - LowerUnsupportedMemoryConsistencyOpPattern, - LowerUnsupportedMemoryConsistencyOpPattern, - LowerDsbOpPattern, - LowerDcciOpPattern, - LowerBufSyncOpPattern, - LowerBufSyncOpPattern, - LowerBufDynSyncOpPattern, - LowerBufDynSyncOpPattern, - LowerBlockRuntimeQueryOpPattern, - LowerRuntimeQueryOpPattern, - LowerBlockRuntimeQueryOpPattern, - LowerRuntimeQueryOpPattern, - LowerVldsOpPattern, LowerVldsx2OpPattern, LowerVsldbOpPattern, - LowerVldasOpPattern, LowerInitAlignOpPattern, - LowerVldusOpPattern, LowerSprclrOpPattern, - LowerSprStoreOpPattern, - LowerSprStoreOpPattern, - LowerVstsOpPattern, LowerVsstbOpPattern, - LowerVstsx2OpPattern, - LowerVstarOpPattern, LowerVstasOpPattern, - LowerVgather2OpPattern, LowerVgather2BcOpPattern, - LowerVgatherbOpPattern, LowerVscatterOpPattern, - LowerVaxpyOpPattern, LowerVmulscvtOpPattern, - LowerVciOpPattern, LowerVexpdifOpPattern, - LowerVbitsortOpPattern, LowerVmrgsort4OpPattern, - LowerVtrcOpPattern, LowerVcvtOpPattern, - LowerVbitcastOpPattern, LowerPbitcastOpPattern, - LowerPredicateLoadOpPattern, - LowerPredicateLoadOpPattern, - LowerPredicateStoreOpPattern, - LowerPredicateStoreOpPattern, - LowerPstuOpPattern, LowerVstusOpPattern, LowerVsturOpPattern, - LowerInterCoreSyncOpPattern, - LowerInterCoreSyncOpPattern, - LowerNamedSyncOpPattern, - LowerNamedSyncOpPattern, - LowerCopyGmToCbufOpPattern, LowerLoadCbufToCaOpPattern, - LowerLoadCbufToCbOpPattern, - LowerLoadCbufToS4OpPattern, - LowerLoadCbufToS4OpPattern, - LowerLoadCbufToCaMxOpPattern, - LowerLoadCbufToCbMxOpPattern, LowerCopyMatrixCcToGmOpPattern, - LowerCopyMatrixCcToBufOpPattern, - LowerCopyMatrixCcToBufOpPattern, - LowerCopyCbufToBtOpPattern, LowerCopyCbufToFbufOpPattern, - LowerCopyGmToCbufMultiOpPattern, - LowerCopyGmToCbufMultiOpPattern, - LowerMadRawPattern, - LowerMadRawPattern, - LowerMadRawPattern, - LowerMadRawPattern, - LowerCopyOpPattern, - LowerCopyOpPattern, - LowerCopyUbufToUbufOpPattern, - LowerCopyCbufToUbufOpPattern, - LowerCopyUbufToCbufOpPattern, - LowerCreateCbufMatrixOpPattern>( - typeConverter, patterns.getContext(), state); -} - -static void configureVPTOOpLoweringTarget(ConversionTarget &target, - VPTOTypeConverter &typeConverter) { - (void)typeConverter; - target.addLegalOp(); - target.addLegalDialect(); - target.addLegalOp(); - target.addIllegalOp(); - target.addIllegalOp(); - target.addIllegalOp(); - target.addIllegalOp(); - target.addIllegalOp(); - target.addIllegalOp(); - target.addIllegalOp(); - target.markUnknownOpDynamicallyLegal([](Operation *op) { - return !isa(op); - }); -} - -static void populateVPTOStructuralTypePatterns( - VPTOTypeConverter &typeConverter, RewritePatternSet &patterns, - ConversionTarget &target) { - scf::populateSCFStructuralTypeConversionsAndLegality(typeConverter, patterns, - target); - populateFunctionOpInterfaceTypeConversionPattern(patterns, - typeConverter); - populateCallOpTypeConversionPattern(patterns, typeConverter); - populateBranchOpInterfaceTypeConversionPattern(patterns, typeConverter); - populateReturnOpTypeConversionPattern(patterns, typeConverter); -} - -static void foldVPTOTypeCasts(ModuleOp module, TypeConverter &typeConverter) { - SmallVector castsToFold; - module.walk([&](UnrealizedConversionCastOp castOp) { - if (castOp->getNumOperands() != 1 || castOp->getNumResults() != 1) - return; - if (!hasVPTOConvertibleType(castOp->getOperandTypes()) && - !hasVPTOConvertibleType(castOp->getResultTypes())) - return; - Type convertedResultType = - typeConverter.convertType(castOp.getResult(0).getType()); - if (convertedResultType && - convertedResultType == castOp.getOperand(0).getType()) - castsToFold.push_back(castOp); - }); - for (UnrealizedConversionCastOp castOp : castsToFold) { - castOp.getResult(0).replaceAllUsesWith(castOp.getOperand(0)); - castOp.erase(); - } -} - -static LogicalResult lowerVPTOOps(ModuleOp module, llvm::raw_ostream &diagOS) { - MLIRContext *context = module.getContext(); - VPTOTypeConverter typeConverter(context); - ConversionTarget target(*context); - RewritePatternSet patterns(context); - LoweringState state; - - configureVPTOOpLoweringTarget(target, typeConverter); - populateVPTOOpLoweringPatterns(typeConverter, patterns, state); - - if (failed(applyPartialConversion(module, target, std::move(patterns)))) { - diagOS << "VPTO LLVM emission failed: VPTO op lowering failed\n"; - return failure(); - } - if (failed(materializeDecls(module, state.plannedDecls, diagOS))) - return failure(); - return success(); -} - -static LogicalResult lowerVPTOTypes(ModuleOp module, llvm::raw_ostream &diagOS) { - MLIRContext *context = module.getContext(); - VPTOTypeConverter typeConverter(context); - ConversionTarget target(*context); - RewritePatternSet patterns(context); - LoweringState state; - - target.addLegalOp(); - target.addDynamicallyLegalOp([&](func::FuncOp op) { - return typeConverter.isSignatureLegal(op.getFunctionType()) && - typeConverter.isLegal(&op.getBody()); - }); - target.addDynamicallyLegalOp( - [&](func::CallOp op) { return typeConverter.isLegal(op); }); - target.addDynamicallyLegalOp( - [&](func::ReturnOp op) { return typeConverter.isLegal(op); }); - target.addDynamicallyLegalOp( - [&](Operation *op) { - return isLegalForBranchOpInterfaceTypeConversionPattern(op, - typeConverter); - }); - target.addDynamicallyLegalOp([&](arith::SelectOp op) { - return typeConverter.isLegal(op->getOperandTypes()) && - typeConverter.isLegal(op->getResultTypes()); - }); - target.addIllegalOp(); - target.addDynamicallyLegalOp( - [&](UnrealizedConversionCastOp op) { - return !hasVPTOConvertibleType(op->getOperandTypes()) && - !hasVPTOConvertibleType(op->getResultTypes()); - }); - target.addDynamicallyLegalOp([&](LLVM::AllocaOp op) { - return typeConverter.isLegal(op->getOperandTypes()) && - typeConverter.isLegal(op->getResultTypes()) && - typeConverter.isLegal(op.getElemType()); - }); - target.addDynamicallyLegalOp([&](LLVM::GEPOp op) { - return typeConverter.isLegal(op->getOperandTypes()) && - typeConverter.isLegal(op->getResultTypes()) && - typeConverter.isLegal(op.getElemType()); - }); - target.markUnknownOpDynamicallyLegal([&](Operation *op) { - return typeConverter.isLegal(op->getOperandTypes()) && - typeConverter.isLegal(op->getResultTypes()); - }); - - populateVPTOStructuralTypePatterns(typeConverter, patterns, target); - patterns.add(typeConverter, context); - patterns.add( - typeConverter, context, state); - patterns.add(typeConverter, context); - patterns.add(typeConverter, context); - patterns.add(typeConverter, context); - - if (failed(applyPartialConversion(module, target, std::move(patterns)))) { - diagOS << "VPTO LLVM emission failed: VPTO type lowering failed\n"; - return failure(); - } - if (failed(materializeDecls(module, state.plannedDecls, diagOS))) - return failure(); - foldVPTOTypeCasts(module, typeConverter); - return success(); -} - -static Type normalizeTypeForOfficialLLVMLowering(Type type, Builder &builder) { - type = convertVPTOType(type, builder); - return type; -} - -static void normalizeFuncSignaturesForOfficialLLVMLowering(ModuleOp module) { - Builder builder(module.getContext()); - - for (func::FuncOp funcOp : module.getOps()) { - FunctionType oldType = funcOp.getFunctionType(); - SmallVector newInputs; - SmallVector newResults; - bool changed = false; - - for (Type input : oldType.getInputs()) { - Type normalized = normalizeTypeForOfficialLLVMLowering(input, builder); - changed |= (normalized != input); - newInputs.push_back(normalized); - } - for (Type result : oldType.getResults()) { - Type normalized = normalizeTypeForOfficialLLVMLowering(result, builder); - changed |= (normalized != result); - newResults.push_back(normalized); - } - - if (!changed) - continue; - - auto newType = builder.getFunctionType(newInputs, newResults); - funcOp.setFunctionTypeAttr(TypeAttr::get(newType)); - - if (funcOp.isExternal()) - continue; - Block &entry = funcOp.getBody().front(); - for (auto [arg, newType] : llvm::zip(entry.getArguments(), newInputs)) - if (arg.getType() != newType) - arg.setType(newType); - } -} - -static void forceV300CtrlModeForVPTOFuncs(ModuleOp module) { - OpBuilder builder(module.getContext()); - - for (func::FuncOp funcOp : module.getOps()) { - if (!needsV300CtrlModeForVPTOFunc(funcOp)) - continue; - - Block &entry = funcOp.getBody().front(); - builder.setInsertionPointToStart(&entry); - auto i64Type = builder.getI64Type(); - auto bit60 = builder.create( - funcOp.getLoc(), i64Type, builder.getI64IntegerAttr(60)); - Value ctrl = - builder.create(funcOp.getLoc(), i64Type).getResult(); - Value ctrlV300 = builder - .create(funcOp.getLoc(), i64Type, - ctrl, bit60.getResult()) - .getResult(); - builder.create(funcOp.getLoc(), ctrlV300); - } -} - -static std::optional getKernelKind(ModuleOp module) { - auto kernelKind = module->getAttrOfType( - FunctionKernelKindAttr::name); - if (!kernelKind) - return std::nullopt; - return kernelKind.getKernelKind(); -} - -static VPTOEmissionOptions -makeDeviceEmissionOptions(const VPTOEmissionOptions &baseOptions, - FunctionKernelKind kind) { - VPTOEmissionOptions options = baseOptions; - constexpr llvm::StringLiteral kVecTargetFeatures = - "+ATOMIC,+ArchV130,+AregRedefinable,+ArithmeticBf16,+AtomicForB8 ," - "+F8e4m3,+F8e5m2,+F8e8m0,+FFTSBlk,+Fp4e1m2x2,+Fp4e2m1x2,+LDExtRefine," - "+MOVX8,+SPR7bits,+SyncV,+dav-c310-vec"; - constexpr llvm::StringLiteral kCubeTargetFeatures = - "+ATOMIC,+ArchV130,+AregRedefinable,+ArithmeticBf16,+AtomicForB8 ," - "+F8e4m3,+F8e5m2,+F8e8m0,+FFTSBlk,+Fp4e1m2x2,+Fp4e2m1x2,+LDExtRefine," - "+MOVX8,+SPR7bits,+SyncV,+dav-c310-cube"; - if (kind == FunctionKernelKind::Vector) { - options.march = "dav-c310-vec"; - options.aicoreArch = "dav-c310-vec"; - options.defaultTargetCPU = "dav-c310-vec"; - options.defaultTargetFeatures = kVecTargetFeatures.str(); - } else if (kind == FunctionKernelKind::Cube) { - options.march = "dav-c310-cube"; - options.aicoreArch = "dav-c310-cube"; - options.defaultTargetCPU = "dav-c310-cube"; - options.defaultTargetFeatures = kCubeTargetFeatures.str(); - } - return options; -} - -static FailureOr -getUniqueDeviceModuleByKernelKind(ModuleOp module, FunctionKernelKind kind, - llvm::raw_ostream &diagOS) { - ModuleOp matched; - for (ModuleOp child : module.getOps()) { - auto kernelKind = getKernelKind(child); - if (!kernelKind) - continue; - if (*kernelKind != kind) - continue; - if (matched) { - diagOS << "VPTO LLVM emission failed: duplicate device module with " - << FunctionKernelKindAttr::name << "\n"; - return failure(); - } - matched = child; - } - return matched; -} - -static LogicalResult renameKernelFunctionsForKernelKind(ModuleOp module, - llvm::raw_ostream &diagOS) { - auto kernelKind = getKernelKind(module); - if (!kernelKind) { - diagOS << "VPTO LLVM emission failed: device module missing " - << FunctionKernelKindAttr::name << "\n"; - return failure(); - } - - StringRef suffix; - if (*kernelKind == FunctionKernelKind::Vector) - suffix = kVectorSuffix; - else if (*kernelKind == FunctionKernelKind::Cube) - suffix = kCubeSuffix; - else { - diagOS << "VPTO LLVM emission failed: unsupported " - << FunctionKernelKindAttr::name << "\n"; - return failure(); - } - - for (func::FuncOp funcOp : module.getOps()) { - if (!pto::hasExplicitPTOEntryAttr(funcOp)) - continue; - if (funcOp.getSymName().ends_with(suffix)) - continue; - funcOp.setSymName((funcOp.getSymName() + suffix).str()); - } - return success(); -} - -struct LowerVPTOOpsPass final - : public PassWrapper> { - MLIR_DEFINE_EXPLICIT_INTERNAL_INLINE_TYPE_ID(LowerVPTOOpsPass) - - void runOnOperation() override { - materializeVecScopeCarrierLoops(getOperation()); - if (failed(lowerVPTOOps(getOperation(), llvm::errs()))) - signalPassFailure(); - } -}; - -struct LowerVPTOTypesPass final - : public PassWrapper> { - MLIR_DEFINE_EXPLICIT_INTERNAL_INLINE_TYPE_ID(LowerVPTOTypesPass) - - void runOnOperation() override { - if (failed(lowerVPTOTypes(getOperation(), llvm::errs()))) - signalPassFailure(); - } -}; - -struct NormalizeFuncSignaturesForLLVMLoweringPass final - : public PassWrapper> { - MLIR_DEFINE_EXPLICIT_INTERNAL_INLINE_TYPE_ID( - NormalizeFuncSignaturesForLLVMLoweringPass) - - void runOnOperation() override { - normalizeFuncSignaturesForOfficialLLVMLowering(getOperation()); - } -}; - -struct PrepareVPTOLLVMLoweringPass final - : public PassWrapper> { - MLIR_DEFINE_EXPLICIT_INTERNAL_INLINE_TYPE_ID(PrepareVPTOLLVMLoweringPass) - - void runOnOperation() override { - ModuleOp module = getOperation(); - pto::annotatePTOEntryFunctions(module); - forceV300CtrlModeForVPTOFuncs(module); - if (failed(renameKernelFunctionsForKernelKind(module, llvm::errs()))) - signalPassFailure(); - } -}; - -static llvm::StringSet -collectSimtEntryFunctionNames(ModuleOp module) { - llvm::StringSet simtEntries; - module.walk([&](func::FuncOp funcOp) { - if (funcOp->hasAttr(pto::kPTOSimtEntryAttrName)) - simtEntries.insert(funcOp.getSymName()); - }); - return simtEntries; -} - -static void applyArtifactVisibilityLinkage(ModuleOp sourceModule, - llvm::Module &llvmModule) { - llvm::StringMap externalByName; - sourceModule.walk([&](func::FuncOp funcOp) { - if (funcOp.isDeclaration()) - return; - externalByName[funcOp.getSymName()] = - pto::hasExternalArtifactVisibility(funcOp); - }); - - for (llvm::Function &function : llvmModule) { - auto it = externalByName.find(function.getName()); - if (it == externalByName.end()) - continue; - if (it->second) { - function.setLinkage(llvm::GlobalValue::ExternalLinkage); - continue; - } - function.setLinkage(llvm::GlobalValue::InternalLinkage); - } -} - -static void applySimtEntryCallingConvention( - llvm::Module &llvmModule, - const llvm::StringSet &simtEntryNames) { - for (llvm::Function &function : llvmModule) { - if (simtEntryNames.contains(function.getName())) { - function.setCallingConv(llvm::CallingConv::SimtEntry); - function.addFnAttr(llvm::Attribute::NoInline); - // Match Bisheng's C++ frontend shape for SIMT outlined bodies. The - // exported wrapper owns the real kernel metadata, while the SIMT body is - // an ODR helper called with the SIMT calling convention. In CANN beta.1, - // leaving the SIMT body as a strong GLOBAL FUNC makes the runtime count it - // as an extra kernel without matching .ascend.meta, which can corrupt the - // selected kernel metadata. linkonce_odr lowers to a weak helper symbol - // and avoids that beta.1 metadata mismatch. - function.setLinkage(llvm::GlobalValue::LinkOnceODRLinkage); - } - } - - for (llvm::Function &function : llvmModule) { - for (llvm::BasicBlock &block : function) { - for (llvm::Instruction &inst : block) { - auto *call = llvm::dyn_cast(&inst); - if (!call) - continue; - auto *callee = call->getCalledFunction(); - if (!callee || !simtEntryNames.contains(callee->getName())) - continue; - call->setCallingConv(llvm::CallingConv::SimtEntry); - } - } - } -} - -static FailureOr -emitDeviceLLVMModule(ModuleOp deviceModule, StringRef kernelKind, - const VPTOEmissionOptions &options, - const llvm::StringSet &simtEntryNames, - llvm::raw_ostream &diagOS) { - if (!deviceModule) - return EmittedLLVMModule{}; - if (failed(applyQueriedTargetAttrs(deviceModule, options, diagOS))) - return failure(); - - auto llvmContext = std::make_unique(); - registerBuiltinDialectTranslation(*deviceModule.getContext()); - registerLLVMDialectTranslation(*deviceModule.getContext()); - std::unique_ptr llvmModule = - translateModuleToLLVMIR(deviceModule.getOperation(), *llvmContext); - if (!llvmModule) { - diagOS << "VPTO LLVM emission failed: LLVM IR export failed for " - << kernelKind << " module\n"; - return failure(); - } - - applyArtifactVisibilityLinkage(deviceModule, *llvmModule); - for (llvm::Function &func : *llvmModule) { - if (!func.getName().starts_with("llvm.hivm.vscatter.")) - continue; - // Work around a bug in older Bisheng releases: vscatter was not modeled - // as writing through its destination pointer, so EarlyCSE could eliminate - // a load after vscatter as redundant. - func.setOnlyAccessesArgMemory(); - func.addFnAttr(llvm::Attribute::NoUnwind); - func.addFnAttr(llvm::Attribute::WriteOnly); - } - applySimtEntryCallingConvention(*llvmModule, simtEntryNames); - if (failed(attachAIVectorScopeMetadata(*llvmModule, diagOS))) - return failure(); - attachHIVMKernelAnnotations(*llvmModule, deviceModule); - llvmModule->setModuleIdentifier(("ptoas.hivm.official." + kernelKind).str()); - llvmModule->setSourceFileName(("ptoas.hivm.official." + kernelKind).str()); - return EmittedLLVMModule{std::move(llvmContext), std::move(llvmModule)}; -} - -template -static LogicalResult runPipeline(ModuleOp module, llvm::raw_ostream &diagOS, - EmitFn &&emit) { - OwningOpRef clonedOp(module->clone()); - ModuleOp clonedModule = cast(*clonedOp); - - if (failed(validateVPTOAuthoringIR(clonedModule, &diagOS))) { - diagOS << "VPTO LLVM emission failed: authoring-stage VPTO legality " - "validation failed\n"; - return failure(); - } - - PassManager pm(clonedModule.getContext()); - pm.enableVerifier(); - auto &kernelModulePM = pm.nest(); - kernelModulePM.addPass(std::make_unique()); - kernelModulePM.addPass(std::make_unique()); - kernelModulePM.addPass(std::make_unique()); - kernelModulePM.addPass( - std::make_unique()); - kernelModulePM.addPass(arith::createArithExpandOpsPass()); - // pto-convert-scf-to-cf-with-loop-hints performs the SCF-to-CF conversion for this pipeline: - // it runs the upstream conversion patterns plus a higher-benefit lowering - // for {pto.unroll = "enable"} loops that attaches llvm.loop_annotation to - // the latch, so the !llvm.loop.unroll.enable metadata survives into the - // emitted LLVM IR. It replaces createConvertSCFToCFPass here; running both - // would be redundant. - kernelModulePM.addNestedPass(pto::createPTOConvertSCFToCFWithLoopHintsPass()); - kernelModulePM.addPass(createArithToLLVMConversionPass()); - kernelModulePM.addPass(createConvertIndexToLLVMPass()); - kernelModulePM.addPass(createFinalizeMemRefToLLVMConversionPass()); - kernelModulePM.addPass(createConvertFuncToLLVMPass()); - kernelModulePM.addPass(createConvertControlFlowToLLVMPass()); - kernelModulePM.addPass(createReconcileUnrealizedCastsPass()); - if (failed(mlir::applyPassManagerCLOptions(pm))) { - diagOS << "VPTO LLVM emission failed: unable to apply MLIR pass manager " - "command-line options\n"; - return failure(); - } - if (failed(pm.run(clonedModule))) { - diagOS << "VPTO LLVM emission failed: official lowering pipeline failed\n"; - return failure(); - } - return emit(clonedModule); -} - -} // namespace +namespace mlir::pto { LogicalResult lowerVPTOModuleToLLVMModulesCANN900( ModuleOp module, const VPTOEmissionOptions &options, EmittedLLVMModule &cubeModule, EmittedLLVMModule &vectorModule, llvm::raw_ostream &diagOS) { - llvm::StringSet simtEntryNames = - collectSimtEntryFunctionNames(module); - cubeModule.context.reset(); - cubeModule.module.reset(); - vectorModule.context.reset(); - vectorModule.module.reset(); - return runPipeline(module, diagOS, - [&](ModuleOp loweredModule) { - auto vectorDeviceModule = - getUniqueDeviceModuleByKernelKind( - loweredModule, FunctionKernelKind::Vector, diagOS); - if (failed(vectorDeviceModule)) - return failure(); - auto cubeDeviceModule = - getUniqueDeviceModuleByKernelKind( - loweredModule, FunctionKernelKind::Cube, diagOS); - if (failed(cubeDeviceModule)) - return failure(); - - if (*vectorDeviceModule) { - auto vectorOptions = - makeDeviceEmissionOptions(options, FunctionKernelKind::Vector); - auto emitted = - emitDeviceLLVMModule(*vectorDeviceModule, "vector", vectorOptions, - simtEntryNames, diagOS); - if (failed(emitted)) - return failure(); - vectorModule.context = std::move(emitted->context); - vectorModule.module = std::move(emitted->module); - } - if (*cubeDeviceModule) { - auto cubeOptions = - makeDeviceEmissionOptions(options, FunctionKernelKind::Cube); - auto emitted = - emitDeviceLLVMModule(*cubeDeviceModule, "cube", cubeOptions, - simtEntryNames, diagOS); - if (failed(emitted)) - return failure(); - cubeModule.context = std::move(emitted->context); - cubeModule.module = std::move(emitted->module); - } - return success(); - }); + return detail::lowerCANN900Module(module, options, cubeModule, vectorModule, + diagOS); } } // namespace mlir::pto diff --git a/lib/PTO/Transforms/VPTOCANN900LLVMEmitterArithmeticPatterns.cpp b/lib/PTO/Transforms/VPTOCANN900LLVMEmitterArithmeticPatterns.cpp new file mode 100644 index 0000000000..c0440bc8e5 --- /dev/null +++ b/lib/PTO/Transforms/VPTOCANN900LLVMEmitterArithmeticPatterns.cpp @@ -0,0 +1,1840 @@ +// Copyright (c) 2026 Huawei Technologies Co., Ltd. +// This program is free software, you can redistribute it and/or modify it under the terms and conditions of +// CANN Open Software License Agreement Version 2.0 (the "License"). +// Please refer to the License for details. You may not use this file except in compliance with the License. +// THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED, +// INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE. +// See LICENSE in the root of the software repository for the full text of the License. + +#include "VPTOCANN900LLVMEmitterTemplates.h" + +namespace mlir::pto::detail { + +template class LowerUnaryMaskedOpPattern final : public OpConversionPattern { +public: + explicit LowerUnaryMaskedOpPattern(TypeConverter &typeConverter, MLIRContext *context, LoweringState &state) + : OpConversionPattern(typeConverter, context), state(state) {} + + LogicalResult matchAndRewrite(UnaryOp op, typename UnaryOp::Adaptor adaptor, + ConversionPatternRewriter &rewriter) const override { + FailureOr calleeName = buildUnaryMaskedCallee(op.getContext(), op.getResult().getType()); + if (failed(calleeName)) { + return rewriter.notifyMatchFailure(op, "unsupported unary VPTO signature"); + } + + Type resultType = this->getTypeConverter()->convertType(op.getResult().getType()); + if (!resultType) { + return rewriter.notifyMatchFailure(op, "failed to convert unary result type"); + } + + Value input = adaptor.getOperands()[0]; + Value mask = adaptor.getOperands()[1]; + Type expectedMaskType = this->getTypeConverter()->convertType(op->getOperand(1).getType()); + if (!input || !mask || input.getType() != resultType || mask.getType() != expectedMaskType) { + return rewriter.notifyMatchFailure(op, "unexpected converted unary VPTO operand types"); + } + + auto call = rewriter.create(op.getLoc(), *calleeName, TypeRange{resultType}, ValueRange{input, mask}); + state.plannedDecls.push_back(PlannedDecl{calleeName->str(), call.getCalleeType()}); + rewriter.replaceOp(op, call.getResults()); + return success(); + } + +private: + LoweringState &state; +}; + +class LowerVsqzOpPattern final : public OpConversionPattern { +public: + explicit LowerVsqzOpPattern(TypeConverter &typeConverter, MLIRContext *context, LoweringState &state) + : OpConversionPattern(typeConverter, context), state(state) {} + + LogicalResult matchAndRewrite(pto::VsqzOp op, pto::VsqzOp::Adaptor adaptor, + ConversionPatternRewriter &rewriter) const override { + FailureOr calleeName = buildVsqzCallee(op.getContext(), op.getResult().getType()); + if (failed(calleeName)) { + return rewriter.notifyMatchFailure(op, "unsupported vsqz VPTO signature"); + } + + Type resultType = this->getTypeConverter()->convertType(op.getResult().getType()); + Type maskType = this->getTypeConverter()->convertType(op.getMask().getType()); + if (!resultType || !maskType) { + return rewriter.notifyMatchFailure(op, "failed to convert vsqz types"); + } + + Value input = adaptor.getInput(); + Value mask = adaptor.getMask(); + if (!input || !mask || input.getType() != resultType || mask.getType() != maskType) { + return rewriter.notifyMatchFailure(op, "unexpected converted vsqz operand types"); + } + + Value storeHint = getI32Constant(rewriter, op.getLoc(), determineVsqzStoreHint(op)); + auto funcType = + rewriter.getFunctionType(TypeRange{resultType, maskType, storeHint.getType()}, TypeRange{resultType}); + auto call = rewriter.create(op.getLoc(), *calleeName, TypeRange{resultType}, + ValueRange{input, mask, storeHint}); + state.plannedDecls.push_back(PlannedDecl{calleeName->str(), funcType}); + rewriter.replaceOp(op, call.getResults()); + return success(); + } + +private: + LoweringState &state; +}; + +class LowerVusqzOpPattern final : public OpConversionPattern { +public: + explicit LowerVusqzOpPattern(TypeConverter &typeConverter, MLIRContext *context, LoweringState &state) + : OpConversionPattern(typeConverter, context), state(state) {} + + LogicalResult matchAndRewrite(pto::VusqzOp op, pto::VusqzOp::Adaptor adaptor, + ConversionPatternRewriter &rewriter) const override { + FailureOr calleeName = buildVusqzCallee(op.getContext(), op.getResult().getType()); + if (failed(calleeName)) { + return rewriter.notifyMatchFailure(op, "unsupported vusqz VPTO signature"); + } + + Type resultType = this->getTypeConverter()->convertType(op.getResult().getType()); + Type maskType = this->getTypeConverter()->convertType(op.getMask().getType()); + if (!resultType || !maskType) { + return rewriter.notifyMatchFailure(op, "failed to convert vusqz types"); + } + + Value src = adaptor.getSrc(); + Value mask = adaptor.getMask(); + if (!src || !mask || src.getType() != resultType || mask.getType() != maskType) { + return rewriter.notifyMatchFailure(op, "unexpected converted vusqz operand types"); + } + + auto funcType = rewriter.getFunctionType(TypeRange{resultType, maskType}, TypeRange{resultType}); + auto call = rewriter.create(op.getLoc(), *calleeName, TypeRange{resultType}, ValueRange{src, mask}); + state.plannedDecls.push_back(PlannedDecl{calleeName->str(), funcType}); + rewriter.replaceOp(op, call.getResults()); + return success(); + } + +private: + LoweringState &state; +}; + +class LowerVmulaOpPattern final : public OpConversionPattern { +public: + explicit LowerVmulaOpPattern(TypeConverter &typeConverter, MLIRContext *context, LoweringState &state) + : OpConversionPattern(typeConverter, context), state(state) {} + + LogicalResult matchAndRewrite(pto::VmulaOp op, pto::VmulaOp::Adaptor adaptor, + ConversionPatternRewriter &rewriter) const override { + FailureOr calleeName = buildVmulaCallee(op.getContext(), op.getResult().getType()); + if (failed(calleeName)) { + return rewriter.notifyMatchFailure(op, "unsupported vmula VPTO signature"); + } + + Type resultType = this->getTypeConverter()->convertType(op.getResult().getType()); + Type maskType = this->getTypeConverter()->convertType(op.getMask().getType()); + if (!resultType || !maskType) { + return rewriter.notifyMatchFailure(op, "failed to convert vmula types"); + } + + Value acc = adaptor.getAcc(); + Value lhs = adaptor.getLhs(); + Value rhs = adaptor.getRhs(); + Value mask = adaptor.getMask(); + if (!acc || !lhs || !rhs || !mask || acc.getType() != resultType || lhs.getType() != resultType || + rhs.getType() != resultType || mask.getType() != maskType) { + return rewriter.notifyMatchFailure(op, "unexpected converted vmula operand types"); + } + + auto funcType = + rewriter.getFunctionType(TypeRange{resultType, resultType, resultType, maskType}, TypeRange{resultType}); + auto call = + rewriter.create(op.getLoc(), *calleeName, TypeRange{resultType}, ValueRange{acc, lhs, rhs, mask}); + state.plannedDecls.push_back(PlannedDecl{calleeName->str(), funcType}); + rewriter.replaceOp(op, call.getResults()); + return success(); + } + +private: + LoweringState &state; +}; + +class LowerVmullOpPattern final : public OpConversionPattern { +public: + explicit LowerVmullOpPattern(TypeConverter &typeConverter, MLIRContext *context, LoweringState &state) + : OpConversionPattern(typeConverter, context), state(state) {} + + LogicalResult matchAndRewrite(pto::VmullOp op, pto::VmullOp::Adaptor adaptor, + ConversionPatternRewriter &rewriter) const override { + FailureOr calleeName = buildVmullCallee(op.getContext(), op.getLow().getType()); + if (failed(calleeName)) { + return rewriter.notifyMatchFailure(op, "unsupported vmull VPTO signature"); + } + + Type inputType = this->getTypeConverter()->convertType(op.getLhs().getType()); + Type maskType = this->getTypeConverter()->convertType(op.getMask().getType()); + SmallVector resultTypes; + if (!inputType || !maskType || failed(this->getTypeConverter()->convertTypes(op->getResultTypes(), resultTypes))) { + return rewriter.notifyMatchFailure(op, "failed to convert vmull types"); + } + if (resultTypes.size() != 2 || resultTypes[0] != resultTypes[1]) { + return rewriter.notifyMatchFailure(op, "unexpected converted vmull results"); + } + + Value lhs = adaptor.getLhs(); + Value rhs = adaptor.getRhs(); + Value mask = adaptor.getMask(); + if (!lhs || !rhs || !mask || lhs.getType() != inputType || rhs.getType() != inputType || + mask.getType() != maskType) { + return rewriter.notifyMatchFailure(op, "unexpected converted vmull operand types"); + } + + auto funcType = rewriter.getFunctionType(TypeRange{inputType, inputType, maskType}, resultTypes); + auto call = rewriter.create(op.getLoc(), *calleeName, resultTypes, ValueRange{lhs, rhs, mask}); + state.plannedDecls.push_back(PlannedDecl{calleeName->str(), funcType}); + rewriter.replaceOp(op, call.getResults()); + return success(); + } + +private: + LoweringState &state; +}; + +template struct BinaryMaskedCall { + Type resultType; + Value lhs; + Value rhs; + FailureOr calleeName; +}; + +template +static BinaryMaskedCall prepareBinaryMaskedCall(BinaryOp op, Value lhs, Value rhs, Type resultType, + ConversionPatternRewriter &rewriter) { + StringRef stem = getBinaryMaskedStem(); + FailureOr calleeName = + usesSignedBinaryCANN900Callee() + ? buildCANN900SignedModeTypedCallee(op.getContext(), op.getResult().getType(), stem, "x") + : buildCANN900ModeTypedCallee(op.getContext(), op.getResult().getType(), stem, "x"); + Type elementType = getElementTypeFromVectorLike(op.getResult().getType()); + if constexpr (std::is_same_v || std::is_same_v || + std::is_same_v) { + if (elementType && pto::isPTOLowPrecisionType(elementType)) { + calleeName = buildDirectLowpVLogicCallee(op.getContext(), op.getResult().getType(), stem, "x"); + if (failed(calleeName)) { + resultType = getLowpPayloadCarrierType(op.getResult().getType(), rewriter.getContext()); + if (resultType) { + lhs = castToPayloadABI(op.getLoc(), lhs, op.getResult().getType(), rewriter); + rhs = castToPayloadABI(op.getLoc(), rhs, op.getResult().getType(), rewriter); + calleeName = buildLowpPayloadVLogicCallee(op.getContext(), op.getResult().getType(), stem, "x"); + } + } + } + } + return BinaryMaskedCall{resultType, lhs, rhs, calleeName}; +} + +template class LowerBinaryMaskedOpPattern final : public OpConversionPattern { +public: + explicit LowerBinaryMaskedOpPattern(TypeConverter &typeConverter, MLIRContext *context, LoweringState &state) + : OpConversionPattern(typeConverter, context), state(state) {} + + LogicalResult matchAndRewrite(BinaryOp op, typename BinaryOp::Adaptor adaptor, + ConversionPatternRewriter &rewriter) const override { + Type resultType = this->getTypeConverter()->convertType(op.getResult().getType()); + if (!resultType) { + return rewriter.notifyMatchFailure(op, "failed to convert binary result type"); + } + + Value lhs = adaptor.getOperands()[0]; + Value rhs = adaptor.getOperands()[1]; + Value mask = adaptor.getOperands()[2]; + Type expectedMaskType = this->getTypeConverter()->convertType(op->getOperand(2).getType()); + if (!lhs || !rhs || !mask || lhs.getType() != resultType || rhs.getType() != resultType || + mask.getType() != expectedMaskType) { + return rewriter.notifyMatchFailure(op, "unexpected converted binary VPTO operand types"); + } + + BinaryMaskedCall callInfo = prepareBinaryMaskedCall(op, lhs, rhs, resultType, rewriter); + if (!callInfo.resultType || failed(callInfo.calleeName)) { + return rewriter.notifyMatchFailure(op, "unsupported binary VPTO signature"); + } + + auto call = rewriter.create(op.getLoc(), *callInfo.calleeName, TypeRange{callInfo.resultType}, + ValueRange{callInfo.lhs, callInfo.rhs, mask}); + state.plannedDecls.push_back(PlannedDecl{callInfo.calleeName->str(), call.getCalleeType()}); + Value result = castFromPayloadABI(op.getLoc(), call.getResult(0), op.getResult().getType(), resultType, rewriter); + rewriter.replaceOp(op, result); + return success(); + } + +private: + LoweringState &state; +}; + +template class LowerTernaryMaskedOpPattern final : public OpConversionPattern { +public: + explicit LowerTernaryMaskedOpPattern(TypeConverter &typeConverter, MLIRContext *context, LoweringState &state) + : OpConversionPattern(typeConverter, context), state(state) {} + + LogicalResult matchAndRewrite(TernaryOp op, typename TernaryOp::Adaptor adaptor, + ConversionPatternRewriter &rewriter) const override { + StringRef stem = getTernaryMaskedStem(); + FailureOr calleeName = + usesSignedTernaryCANN900Callee() + ? buildCANN900SignedModeTypedCallee(op.getContext(), op.getResult().getType(), stem, "m") + : buildCANN900ModeTypedCallee(op.getContext(), op.getResult().getType(), stem, "m"); + if (failed(calleeName)) { + return rewriter.notifyMatchFailure(op, "unsupported ternary VPTO signature"); + } + + Type resultType = this->getTypeConverter()->convertType(op.getResult().getType()); + Type expectedMaskType = this->getTypeConverter()->convertType(op.getMask().getType()); + if (!resultType || !expectedMaskType) { + return rewriter.notifyMatchFailure(op, "failed to convert ternary VPTO types"); + } + + Value acc = adaptor.getAcc(); + Value lhs = adaptor.getLhs(); + Value rhs = adaptor.getRhs(); + Value mask = adaptor.getMask(); + if (!acc || !lhs || !rhs || !mask || acc.getType() != resultType || lhs.getType() != resultType || + rhs.getType() != resultType || mask.getType() != expectedMaskType) { + return rewriter.notifyMatchFailure(op, "unexpected converted ternary VPTO operand types"); + } + + auto call = + rewriter.create(op.getLoc(), *calleeName, TypeRange{resultType}, ValueRange{acc, lhs, rhs, mask}); + state.plannedDecls.push_back(PlannedDecl{calleeName->str(), call.getCalleeType()}); + rewriter.replaceOp(op, call.getResults()); + return success(); + } + +private: + LoweringState &state; +}; + +template class LowerCarryBinaryOpPattern final : public OpConversionPattern { +public: + explicit LowerCarryBinaryOpPattern(TypeConverter &typeConverter, MLIRContext *context, LoweringState &state) + : OpConversionPattern(typeConverter, context), state(state) {} + + LogicalResult matchAndRewrite(CarryOp op, typename CarryOp::Adaptor adaptor, + ConversionPatternRewriter &rewriter) const override { + StringRef stem = getCarryBinaryStem(); + FailureOr calleeName = buildCarryBinaryCallee(op.getContext(), op.getResult().getType(), stem); + if (failed(calleeName)) { + return rewriter.notifyMatchFailure(op, "unsupported carry VPTO signature"); + } + + Type resultType = this->getTypeConverter()->convertType(op.getResult().getType()); + Type carryType = this->getTypeConverter()->convertType(op->getResult(1).getType()); + if (!resultType || !carryType) { + return rewriter.notifyMatchFailure(op, "failed to convert carry result types"); + } + + SmallVector callArgs; + callArgs.append(adaptor.getOperands().begin(), adaptor.getOperands().end()); + const size_t expectedArgCount = hasCarryInput() ? 4 : 3; + if (callArgs.size() != expectedArgCount || callArgs[0].getType() != resultType || + callArgs[1].getType() != resultType || callArgs.back().getType() != carryType) { + return rewriter.notifyMatchFailure(op, "unexpected converted carry operand types"); + } + if constexpr (hasCarryInput()) { + if (callArgs[2].getType() != carryType) { + return rewriter.notifyMatchFailure(op, "unexpected converted carry input operand type"); + } + } + + auto call = rewriter.create(op.getLoc(), *calleeName, TypeRange{resultType, carryType}, callArgs); + state.plannedDecls.push_back(PlannedDecl{calleeName->str(), call.getCalleeType()}); + rewriter.replaceOp(op, call.getResults()); + return success(); + } + +private: + LoweringState &state; +}; + +template class LowerCopyOpPattern final : public OpConversionPattern { +public: + explicit LowerCopyOpPattern(TypeConverter &typeConverter, MLIRContext *context, LoweringState &state) + : OpConversionPattern(typeConverter, context), state(state) {} + + LogicalResult matchAndRewrite(CopyOp op, typename CopyOp::Adaptor adaptor, + ConversionPatternRewriter &rewriter) const override { + FailureOr calleeName = failure(); + if constexpr (std::is_same_v) { + calleeName = buildCopyGmToUbCallee(op.getContext(), op.getSource().getType()); + } else { + calleeName = buildCopyUbToGmCallee(op.getContext()); + } + if (failed(calleeName)) { + return rewriter.notifyMatchFailure(op, "unsupported copy VPTO signature"); + } + + auto llvmSourceType = dyn_cast(adaptor.getOperands()[0].getType()); + auto llvmDestType = dyn_cast(adaptor.getOperands()[1].getType()); + if (!llvmSourceType || !llvmDestType) { + return rewriter.notifyMatchFailure(op, "expected LLVM pointer copy operands"); + } + + FailureOr config0 = failure(); + FailureOr config1 = failure(); + if constexpr (std::is_same_v) { + config0 = packCopyGmToUbConfig0(op, adaptor.getOperands()); + config1 = packCopyGmToUbConfig1(op, adaptor.getOperands()); + } else { + config0 = packCopyUbToGmConfig0(op, adaptor.getOperands()); + config1 = packCopyUbToGmConfig1(op, adaptor.getOperands()); + } + if (failed(config0) || failed(config1)) { + return rewriter.notifyMatchFailure(op, "failed to materialize copy config"); + } + + SmallVector args{adaptor.getOperands()[1], adaptor.getOperands()[0], *config0, *config1}; + auto funcType = rewriter.getFunctionType( + TypeRange{llvmDestType, llvmSourceType, rewriter.getI64Type(), rewriter.getI64Type()}, TypeRange{}); + auto call = rewriter.create(op.getLoc(), *calleeName, TypeRange{}, args); + state.plannedDecls.push_back(PlannedDecl{calleeName->str(), funcType}); + rewriter.eraseOp(op); + (void)call; + return success(); + } + +private: + LoweringState &state; +}; + +class LowerCopyUbufToUbufOpPattern final : public OpConversionPattern { +public: + explicit LowerCopyUbufToUbufOpPattern(TypeConverter &typeConverter, MLIRContext *context, LoweringState &state) + : OpConversionPattern(typeConverter, context), state(state) {} + + LogicalResult matchAndRewrite(pto::CopyUbufToUbufOp op, pto::CopyUbufToUbufOp::Adaptor adaptor, + ConversionPatternRewriter &rewriter) const override { + auto llvmSourceType = dyn_cast(adaptor.getOperands()[0].getType()); + auto llvmDestType = dyn_cast(adaptor.getOperands()[1].getType()); + if (!llvmSourceType || !llvmDestType) { + return rewriter.notifyMatchFailure(op, "expected LLVM pointer copy operands"); + } + + FailureOr config = packCopyUbToUbConfig(op, adaptor.getOperands()); + if (failed(config)) { + return rewriter.notifyMatchFailure(op, "failed to materialize copy config"); + } + + StringRef calleeName = buildCopyUbToUbCallee(op.getContext()); + SmallVector args{adaptor.getOperands()[1], adaptor.getOperands()[0], *config}; + auto funcType = + rewriter.getFunctionType(TypeRange{llvmDestType, llvmSourceType, rewriter.getI64Type()}, TypeRange{}); + auto call = rewriter.create(op.getLoc(), calleeName, TypeRange{}, args); + state.plannedDecls.push_back(PlannedDecl{calleeName.str(), funcType}); + rewriter.eraseOp(op); + (void)call; + return success(); + } + +private: + LoweringState &state; +}; + +class LowerCopyCbufToUbufOpPattern final : public OpConversionPattern { +public: + explicit LowerCopyCbufToUbufOpPattern(TypeConverter &typeConverter, MLIRContext *context, LoweringState &state) + : OpConversionPattern(typeConverter, context), state(state) {} + + LogicalResult matchAndRewrite(pto::CopyCbufToUbufOp op, pto::CopyCbufToUbufOp::Adaptor adaptor, + ConversionPatternRewriter &rewriter) const override { + Value sourceRaw = adaptor.getSource(); + Value destinationRaw = adaptor.getDestination(); + if (!sourceRaw || !destinationRaw) { + return rewriter.notifyMatchFailure(op, "expected converted operands"); + } + if (!isa(sourceRaw.getType()) || !isa(destinationRaw.getType())) { + return rewriter.notifyMatchFailure(op, "expected LLVM pointer src/dst"); + } + + constexpr unsigned cbufAddressSpace = static_cast(pto::AddressSpace::MAT); + constexpr unsigned ubufAddressSpace = static_cast(pto::AddressSpace::VEC); + FailureOr source = reinterpretPointerToAddrSpace(op, sourceRaw, cbufAddressSpace); + FailureOr destination = reinterpretPointerToAddrSpace(op, destinationRaw, ubufAddressSpace); + if (failed(source) || failed(destination)) { + return rewriter.notifyMatchFailure(op, "failed to map cbuf/ubuf pointer spaces"); + } + + FailureOr config = packCopyCbufToUbConfig(op, adaptor.getOperands()); + if (failed(config)) { + return rewriter.notifyMatchFailure(op, "failed to materialize copy config"); + } + + StringRef calleeName = buildCopyCbufToUbCallee(op.getContext()); + auto funcType = rewriter.getFunctionType( + TypeRange{destination->getType(), source->getType(), rewriter.getI64Type()}, TypeRange{}); + rewriter.create(op.getLoc(), calleeName, TypeRange{}, ValueRange{*destination, *source, *config}); + state.plannedDecls.push_back(PlannedDecl{calleeName.str(), funcType}); + rewriter.eraseOp(op); + return success(); + } + +private: + LoweringState &state; +}; + +class LowerCopyUbufToCbufOpPattern final : public OpConversionPattern { +public: + explicit LowerCopyUbufToCbufOpPattern(TypeConverter &typeConverter, MLIRContext *context, LoweringState &state) + : OpConversionPattern(typeConverter, context), state(state) {} + + LogicalResult matchAndRewrite(pto::CopyUbufToCbufOp op, pto::CopyUbufToCbufOp::Adaptor adaptor, + ConversionPatternRewriter &rewriter) const override { + Value sourceRaw = adaptor.getSource(); + Value destinationRaw = adaptor.getDestination(); + if (!sourceRaw || !destinationRaw) { + return rewriter.notifyMatchFailure(op, "expected converted operands"); + } + if (!isa(sourceRaw.getType()) || !isa(destinationRaw.getType())) { + return rewriter.notifyMatchFailure(op, "expected LLVM pointer src/dst"); + } + + constexpr unsigned ubufAddressSpace = static_cast(pto::AddressSpace::VEC); + constexpr unsigned cbufAddressSpace = static_cast(pto::AddressSpace::MAT); + FailureOr source = reinterpretPointerToAddrSpace(op, sourceRaw, ubufAddressSpace); + FailureOr destination = reinterpretPointerToAddrSpace(op, destinationRaw, cbufAddressSpace); + if (failed(source) || failed(destination)) { + return rewriter.notifyMatchFailure(op, "failed to map ubuf/cbuf pointer spaces"); + } + + FailureOr config = packCopyUbToCbufConfig(op, adaptor.getOperands()); + if (failed(config)) { + return rewriter.notifyMatchFailure(op, "failed to materialize copy config"); + } + + StringRef calleeName = buildCopyUbToCbufCallee(op.getContext()); + auto funcType = rewriter.getFunctionType( + TypeRange{destination->getType(), source->getType(), rewriter.getI64Type()}, TypeRange{}); + rewriter.create(op.getLoc(), calleeName, TypeRange{}, ValueRange{*destination, *source, *config}); + state.plannedDecls.push_back(PlannedDecl{calleeName.str(), funcType}); + rewriter.eraseOp(op); + return success(); + } + +private: + LoweringState &state; +}; + +struct CreateCbufMatrixFill { + StringRef calleeName; + Value pattern; +}; + +static FailureOr buildCreateCbufMatrixFill(pto::CreateCbufMatrixOp op, Value rawValue, + ConversionPatternRewriter &rewriter) { + Location loc = op.getLoc(); + Type i64Ty = rewriter.getI64Type(); + const uint64_t fillWordWidth = static_cast(op.getFillWordBits()); + if (fillWordWidth == 16) { + Value wordMask = getI32Constant(rewriter, loc, 0xFFFFU); + Value lowWord = rewriter.create(loc, rawValue, wordMask); + Value wordBits = rewriter.create(loc, rewriter.getI16Type(), lowWord); + Value pattern = rewriter.create(loc, rewriter.getF16Type(), wordBits); + return CreateCbufMatrixFill{"llvm.hivm.CREATE.CBUF.MATRIX.v3.u16.h", pattern}; + } + if (fillWordWidth == 32) { + Value pattern = rewriter.create(loc, i64Ty, rawValue); + return CreateCbufMatrixFill{"llvm.hivm.CREATE.CBUF.MATRIX.v3.u32", pattern}; + } + return failure(); +} + +static Value packCreateCbufMatrixConfig(Operation *anchor, Value repeatTimes, Value blockNum32b, Value dstGap32b, + ConversionPatternRewriter &rewriter) { + Location loc = anchor->getLoc(); + Value fieldMask = getI64Constant(rewriter, loc, 0x7FFFU); + auto maskField = [&](Value value) -> Value { return rewriter.create(loc, value, fieldMask); }; + auto shiftField = [&](Value value, uint64_t amount) -> Value { + return rewriter.create(loc, value, getI64Constant(rewriter, loc, amount)); + }; + Value config = maskField(repeatTimes); + config = rewriter.create(loc, config, shiftField(maskField(blockNum32b), 16)); + return rewriter.create(loc, config, shiftField(maskField(dstGap32b), 32)); +} + +class LowerCreateCbufMatrixOpPattern final : public OpConversionPattern { +public: + explicit LowerCreateCbufMatrixOpPattern(TypeConverter &typeConverter, MLIRContext *context, LoweringState &state) + : OpConversionPattern(typeConverter, context), state(state) {} + + LogicalResult matchAndRewrite(pto::CreateCbufMatrixOp op, pto::CreateCbufMatrixOp::Adaptor adaptor, + ConversionPatternRewriter &rewriter) const override { + Value destinationRaw = adaptor.getDst(); + Value rawValue = adaptor.getRawValue(); + Value repeatTimes = adaptor.getRepeatTimes(); + Value blockNum32b = adaptor.getBlockNum_32b(); + Value dstGap32b = adaptor.getDstGap_32b(); + if (!destinationRaw || !rawValue || !repeatTimes || !blockNum32b || !dstGap32b) { + return rewriter.notifyMatchFailure(op, "expected converted operands"); + } + if (!isa(destinationRaw.getType())) { + return rewriter.notifyMatchFailure(op, "expected LLVM pointer destination"); + } + + Type i32Ty = rewriter.getI32Type(); + Type i64Ty = rewriter.getI64Type(); + const bool validControlTypes = rawValue.getType() == i32Ty && repeatTimes.getType() == i64Ty && + blockNum32b.getType() == i64Ty && dstGap32b.getType() == i64Ty; + if (!validControlTypes) { + return rewriter.notifyMatchFailure(op, "expected i32 value and i64 controls"); + } + + constexpr unsigned cbufAddressSpace = static_cast(pto::AddressSpace::MAT); + FailureOr destination = reinterpretPointerToAddrSpace(op, destinationRaw, cbufAddressSpace); + if (failed(destination)) { + return rewriter.notifyMatchFailure(op, "failed to map destination to mat/l1"); + } + + FailureOr fill = buildCreateCbufMatrixFill(op, rawValue, rewriter); + if (failed(fill)) { + return rewriter.notifyMatchFailure(op, "expected a 16-bit or 32-bit fill word"); + } + + Location loc = op.getLoc(); + Value config = packCreateCbufMatrixConfig(op, repeatTimes, blockNum32b, dstGap32b, rewriter); + + auto funcType = + rewriter.getFunctionType(TypeRange{destination->getType(), i64Ty, fill->pattern.getType()}, TypeRange{}); + rewriter.create(loc, fill->calleeName, TypeRange{}, ValueRange{*destination, config, fill->pattern}); + state.plannedDecls.push_back(PlannedDecl{fill->calleeName.str(), funcType}); + rewriter.eraseOp(op); + return success(); + } + +private: + LoweringState &state; +}; + +static LogicalResult validateMadRawOperands(pto::MadRawOpInterface op, ValueRange operands, Value &bias, Value &xt, + ConversionPatternRewriter &rewriter) { + unsigned required = op.hasBiasOperand() ? 5 : 4; + if (operands.size() < required) { + return rewriter.notifyMatchFailure(op, "expected converted mad raw operands"); + } + bias = op.hasBiasOperand() ? operands[3] : Value(); + xt = operands[op.hasBiasOperand() ? 4 : 3]; + for (unsigned index : {0U, 1U, 2U}) { + if (!isa(operands[index].getType())) { + return rewriter.notifyMatchFailure(op, "expected LLVM pointer lhs/rhs/dst operands"); + } + } + if (bias && !isa(bias.getType())) { + return rewriter.notifyMatchFailure(op, "expected LLVM pointer bias operand"); + } + return success(); +} + +struct MadRawLoweringValues { + Value lhs; + Value rhs; + Value dst; + Value bias; + Value xt; +}; + +static FailureOr mapMadRawOperands(pto::MadRawOpInterface op, ValueRange convertedOperands, + ConversionPatternRewriter &rewriter) { + Value biasRaw; + Value xt; + if (failed(validateMadRawOperands(op, convertedOperands, biasRaw, xt, rewriter))) { + return failure(); + } + + constexpr unsigned caAddressSpace = static_cast(pto::AddressSpace::LEFT); + constexpr unsigned cbAddressSpace = static_cast(pto::AddressSpace::RIGHT); + constexpr unsigned ccAddressSpace = static_cast(pto::AddressSpace::ACC); + constexpr unsigned btAddressSpace = static_cast(pto::AddressSpace::BIAS); + FailureOr lhs = reinterpretPointerToAddrSpace(op, convertedOperands[0], caAddressSpace); + FailureOr rhs = reinterpretPointerToAddrSpace(op, convertedOperands[1], cbAddressSpace); + FailureOr dst = reinterpretPointerToAddrSpace(op, convertedOperands[2], ccAddressSpace); + FailureOr bias; + if (biasRaw) { + bias = reinterpretPointerToAddrSpace(op, biasRaw, btAddressSpace); + } + if (failed(lhs) || failed(rhs) || failed(dst) || (biasRaw && failed(bias))) { + return failure(); + } + return MadRawLoweringValues{*lhs, *rhs, *dst, biasRaw ? *bias : Value(), xt}; +} + +static LogicalResult lowerMadRawOp(pto::MadRawOpInterface op, ValueRange convertedOperands, + ConversionPatternRewriter &rewriter, LoweringState &state) { + FailureOr values = mapMadRawOperands(op, convertedOperands, rewriter); + if (failed(values)) { + return failure(); + } + Type i64Ty = rewriter.getI64Type(); + FailureOr calleeName = + op.isMadMxFamily() ? buildMxMadCallee(op.getContext(), op) : buildOrdinaryMadCallee(op.getContext(), op); + if (failed(calleeName)) { + return rewriter.notifyMatchFailure(op, "unsupported mad element types for raw dispatch"); + } + + Value callDst = values->dst; + if (values->bias) { + callDst = buildMadBiasDestination(op, rewriter, values->dst, values->bias); + } + auto funcType = rewriter.getFunctionType( + TypeRange{values->dst.getType(), values->lhs.getType(), values->rhs.getType(), i64Ty}, TypeRange{}); + auto call = rewriter.create(op->getLoc(), *calleeName, TypeRange{}, + ValueRange{callDst, values->lhs, values->rhs, values->xt}); + state.plannedDecls.push_back(PlannedDecl{calleeName->str(), funcType}); + rewriter.replaceOp(op, call.getResults()); + return success(); +} + +template class LowerMadRawPattern final : public OpConversionPattern { +public: + explicit LowerMadRawPattern(TypeConverter &typeConverter, MLIRContext *context, LoweringState &state) + : OpConversionPattern(typeConverter, context), state(state) {} + + LogicalResult matchAndRewrite(RawOp op, typename RawOp::Adaptor adaptor, + ConversionPatternRewriter &rewriter) const override { + auto raw = dyn_cast(op.getOperation()); + if (!raw) { + return failure(); + } + return lowerMadRawOp(raw, adaptor.getOperands(), rewriter, state); + } + +private: + LoweringState &state; +}; + +class LowerCopyGmToCbufOpPattern final : public OpConversionPattern { +public: + explicit LowerCopyGmToCbufOpPattern(TypeConverter &typeConverter, MLIRContext *context, LoweringState &state) + : OpConversionPattern(typeConverter, context), state(state) {} + + LogicalResult matchAndRewrite(pto::CopyGmToCbufOp op, pto::CopyGmToCbufOp::Adaptor adaptor, + ConversionPatternRewriter &rewriter) const override { + Value sourceRaw = adaptor.getSource(); + Value destinationRaw = adaptor.getDestination(); + Value nBurst = adaptor.getNBurst(); + Value lenBurst = adaptor.getLenBurst(); + Value srcStride = adaptor.getSrcStride(); + Value dstStride = adaptor.getDstStride(); + if (!sourceRaw || !destinationRaw || !nBurst || !lenBurst || !srcStride || !dstStride) { + return rewriter.notifyMatchFailure(op, "expected converted operands"); + } + + if (!isa(sourceRaw.getType()) || !isa(destinationRaw.getType())) { + return rewriter.notifyMatchFailure(op, "expected LLVM pointer src/dst"); + } + + Type i64Ty = rewriter.getI64Type(); + if (nBurst.getType() != i64Ty || lenBurst.getType() != i64Ty || srcStride.getType() != i64Ty || + dstStride.getType() != i64Ty) { + return rewriter.notifyMatchFailure(op, "expected i64 config operands"); + } + + constexpr unsigned gmAddressSpace = static_cast(pto::AddressSpace::GM); + constexpr unsigned cbufAddressSpace = static_cast(pto::AddressSpace::MAT); + FailureOr source = reinterpretPointerToAddrSpace(op, sourceRaw, gmAddressSpace); + FailureOr destination = reinterpretPointerToAddrSpace(op, destinationRaw, cbufAddressSpace); + if (failed(source) || failed(destination)) { + return rewriter.notifyMatchFailure(op, "failed to map cbuf/gm pointer spaces"); + } + + FailureOr calleeName = buildCopyGmToCbufCallee(op.getContext(), op.getSource().getType()); + if (failed(calleeName)) { + return rewriter.notifyMatchFailure(op, "unsupported copy_gm_to_cbuf element type"); + } + FailureOr config0 = packCopyGmToCbufConfig0(op, nBurst, lenBurst); + FailureOr config1 = packCopyGmToCbufConfig1(op, srcStride, dstStride); + if (failed(config0) || failed(config1)) { + return rewriter.notifyMatchFailure(op, "failed to pack copy_gm_to_cbuf config"); + } + + auto funcType = + rewriter.getFunctionType(TypeRange{destination->getType(), source->getType(), i64Ty, i64Ty}, TypeRange{}); + rewriter.create(op.getLoc(), *calleeName, TypeRange{}, + ValueRange{*destination, *source, *config0, *config1}); + state.plannedDecls.push_back(PlannedDecl{calleeName->str(), funcType}); + rewriter.eraseOp(op); + return success(); + } + +private: + LoweringState &state; +}; + +template class LowerCopyGmToCbufMultiOpPattern final : public OpConversionPattern { +public: + explicit LowerCopyGmToCbufMultiOpPattern(TypeConverter &typeConverter, MLIRContext *context, LoweringState &state) + : OpConversionPattern(typeConverter, context), state(state) {} + + LogicalResult matchAndRewrite(CopyOp op, typename CopyOp::Adaptor adaptor, + ConversionPatternRewriter &rewriter) const override { + Value sourceRaw = adaptor.getSource(); + Value destinationRaw = adaptor.getDestination(); + if (!sourceRaw || !destinationRaw) { + return rewriter.notifyMatchFailure(op, "expected converted operands"); + } + if (!isa(sourceRaw.getType()) || !isa(destinationRaw.getType())) { + return rewriter.notifyMatchFailure(op, "expected LLVM pointer src/dst"); + } + + constexpr unsigned gmAddressSpace = static_cast(pto::AddressSpace::GM); + constexpr unsigned cbufAddressSpace = static_cast(pto::AddressSpace::MAT); + FailureOr source = reinterpretPointerToAddrSpace(op, sourceRaw, gmAddressSpace); + FailureOr destination = reinterpretPointerToAddrSpace(op, destinationRaw, cbufAddressSpace); + if (failed(source) || failed(destination)) { + return rewriter.notifyMatchFailure(op, "failed to map cbuf/gm pointer spaces"); + } + + FailureOr config0 = packCopyGmToCbufMultiConfig0(op, adaptor.getSid(), adaptor.getLoop1SrcStride(), + adaptor.getL2CacheCtrl(), adaptor.getNValue()); + FailureOr config1 = + packCopyGmToCbufMultiConfig1(op, adaptor.getDValue(), adaptor.getLoop4SrcStride(), adaptor.getSmallc0En()); + if (failed(config0) || failed(config1)) { + return rewriter.notifyMatchFailure(op, "failed to pack multi copy config"); + } + + FailureOr calleeName = [&](MLIRContext *ctx, Type sourceType) -> FailureOr { + if constexpr (std::is_same_v) { + return buildCopyGmToCbufMultiNd2NzCallee(ctx, op.getSource().getType()); + } + return buildCopyGmToCbufMultiDn2NzCallee(ctx, sourceType); + }(op.getContext(), op.getSource().getType()); + if (failed(calleeName)) { + return rewriter.notifyMatchFailure(op, "unsupported copy_gm_to_cbuf_multi element type"); + } + + Type i64Ty = rewriter.getI64Type(); + auto funcType = + rewriter.getFunctionType(TypeRange{destination->getType(), source->getType(), i64Ty, i64Ty}, TypeRange{}); + rewriter.create(op.getLoc(), *calleeName, TypeRange{}, + ValueRange{*destination, *source, *config0, *config1}); + state.plannedDecls.push_back(PlannedDecl{calleeName->str(), funcType}); + rewriter.eraseOp(op); + return success(); + } + +private: + LoweringState &state; +}; + +class LowerCopyCbufToBtOpPattern final : public OpConversionPattern { +public: + explicit LowerCopyCbufToBtOpPattern(TypeConverter &typeConverter, MLIRContext *context, LoweringState &state) + : OpConversionPattern(typeConverter, context), state(state) {} + + LogicalResult matchAndRewrite(pto::CopyCbufToBtOp op, pto::CopyCbufToBtOp::Adaptor adaptor, + ConversionPatternRewriter &rewriter) const override { + Value sourceRaw = adaptor.getSource(); + Value destinationRaw = adaptor.getDestination(); + if (!sourceRaw || !destinationRaw) { + return rewriter.notifyMatchFailure(op, "expected converted operands"); + } + if (!isa(sourceRaw.getType()) || !isa(destinationRaw.getType())) { + return rewriter.notifyMatchFailure(op, "expected LLVM pointer src/dst"); + } + + constexpr unsigned cbufAddressSpace = static_cast(pto::AddressSpace::MAT); + constexpr unsigned btAddressSpace = static_cast(pto::AddressSpace::BIAS); + FailureOr source = reinterpretPointerToAddrSpace(op, sourceRaw, cbufAddressSpace); + FailureOr destinationPtr = reinterpretPointerToAddrSpace(op, destinationRaw, btAddressSpace); + if (failed(source) || failed(destinationPtr)) { + return rewriter.notifyMatchFailure(op, "failed to map cbuf/bt pointer spaces"); + } + + FailureOr config = + packCopyCbufToBtConfig(op, adaptor.getConvControl(), adaptor.getNBurst(), adaptor.getLenBurst(), + adaptor.getSourceGap(), adaptor.getDstGap()); + if (failed(config)) { + return rewriter.notifyMatchFailure(op, "failed to pack copy_cbuf_to_bt config"); + } + + Type i64Ty = rewriter.getI64Type(); + Value destination = rewriter.create(op.getLoc(), i64Ty, *destinationPtr); + FailureOr calleeName = buildCopyCbufToBtCallee(op); + if (failed(calleeName)) { + return rewriter.notifyMatchFailure(op, "unsupported copy_cbuf_to_bt source element type"); + } + auto funcType = rewriter.getFunctionType(TypeRange{i64Ty, source->getType(), i64Ty}, TypeRange{}); + rewriter.create(op.getLoc(), *calleeName, TypeRange{}, ValueRange{destination, *source, *config}); + state.plannedDecls.push_back(PlannedDecl{calleeName->str(), funcType}); + rewriter.eraseOp(op); + return success(); + } + +private: + LoweringState &state; +}; + +class LowerCopyCbufToFbufOpPattern final : public OpConversionPattern { +public: + explicit LowerCopyCbufToFbufOpPattern(TypeConverter &typeConverter, MLIRContext *context, LoweringState &state) + : OpConversionPattern(typeConverter, context), state(state) {} + + LogicalResult matchAndRewrite(pto::CopyCbufToFbufOp op, pto::CopyCbufToFbufOp::Adaptor adaptor, + ConversionPatternRewriter &rewriter) const override { + Value sourceRaw = adaptor.getSource(); + Value destinationRaw = adaptor.getDestination(); + if (!sourceRaw || !destinationRaw) { + return rewriter.notifyMatchFailure(op, "expected converted operands"); + } + if (!isa(sourceRaw.getType()) || !isa(destinationRaw.getType())) { + return rewriter.notifyMatchFailure(op, "expected LLVM pointer src/dst"); + } + + constexpr unsigned cbufAddressSpace = static_cast(pto::AddressSpace::MAT); + constexpr unsigned fbufAddressSpace = 7; + FailureOr source = reinterpretPointerToAddrSpace(op, sourceRaw, cbufAddressSpace); + FailureOr destination = reinterpretPointerToAddrSpace(op, destinationRaw, fbufAddressSpace); + if (failed(source) || failed(destination)) { + return rewriter.notifyMatchFailure(op, "failed to map cbuf/fbuf pointer spaces"); + } + + FailureOr config = packCopyCbufToFbufConfig(op, adaptor.getNBurst(), adaptor.getLenBurst(), + adaptor.getSourceGap(), adaptor.getDstGap()); + if (failed(config)) { + return rewriter.notifyMatchFailure(op, "failed to pack copy_cbuf_to_fbuf config"); + } + + Type i64Ty = rewriter.getI64Type(); + StringRef calleeName = buildCopyCbufToFbufCallee(op.getContext()); + auto funcType = rewriter.getFunctionType(TypeRange{destination->getType(), source->getType(), i64Ty}, TypeRange{}); + rewriter.create(op.getLoc(), calleeName, TypeRange{}, ValueRange{*destination, *source, *config}); + state.plannedDecls.push_back(PlannedDecl{calleeName.str(), funcType}); + rewriter.eraseOp(op); + return success(); + } + +private: + LoweringState &state; +}; + +class LowerLoadCbufToCaOpPattern final : public OpConversionPattern { +public: + explicit LowerLoadCbufToCaOpPattern(TypeConverter &typeConverter, MLIRContext *context, LoweringState &state) + : OpConversionPattern(typeConverter, context), state(state) {} + + LogicalResult matchAndRewrite(pto::LoadCbufToCaOp op, pto::LoadCbufToCaOp::Adaptor adaptor, + ConversionPatternRewriter &rewriter) const override { + Value sourceRaw = adaptor.getSource(); + Value destinationRaw = adaptor.getDestination(); + Value mStart = adaptor.getMStart(); + Value kStart = adaptor.getKStart(); + Value mStep = adaptor.getMStep(); + Value kStep = adaptor.getKStep(); + Value srcStride = adaptor.getSrcStride(); + Value dstStride = adaptor.getDstStride(); + if (!sourceRaw || !destinationRaw || !mStart || !kStart || !mStep || !kStep || !srcStride || !dstStride) { + return rewriter.notifyMatchFailure(op, "expected converted operands"); + } + + if (!isa(sourceRaw.getType()) || !isa(destinationRaw.getType())) { + return rewriter.notifyMatchFailure(op, "expected LLVM pointer src/dst"); + } + + Type i64Ty = rewriter.getI64Type(); + + constexpr unsigned cbufAddressSpace = static_cast(pto::AddressSpace::MAT); + constexpr unsigned caAddressSpace = static_cast(pto::AddressSpace::LEFT); + FailureOr source = reinterpretPointerToAddrSpace(op, sourceRaw, cbufAddressSpace); + FailureOr destination = reinterpretPointerToAddrSpace(op, destinationRaw, caAddressSpace); + if (failed(source) || failed(destination)) { + return rewriter.notifyMatchFailure(op, "failed to map cbuf/ca pointer spaces"); + } + + FailureOr config0 = packLoadCbufToCaConfig0(op, mStart, kStart, mStep, kStep); + FailureOr config1 = packLoadCbufToCaConfig1(op, srcStride, dstStride); + if (failed(config0) || failed(config1)) { + return rewriter.notifyMatchFailure(op, "failed to pack load_cbuf_to_ca config"); + } + Value transpose = getI64Constant(rewriter, op.getLoc(), op.getTranspose() ? 1 : 0); + + FailureOr calleeName = buildLoadCbufToCaCallee(op.getContext(), op.getSource().getType()); + if (failed(calleeName)) { + return rewriter.notifyMatchFailure(op, "unsupported load_cbuf_to_ca element type"); + } + auto funcType = rewriter.getFunctionType(TypeRange{destination->getType(), source->getType(), i64Ty, i64Ty, i64Ty}, + TypeRange{}); + rewriter.create(op.getLoc(), *calleeName, TypeRange{}, + ValueRange{*destination, *source, *config0, *config1, transpose}); + state.plannedDecls.push_back(PlannedDecl{calleeName->str(), funcType}); + rewriter.eraseOp(op); + return success(); + } + +private: + LoweringState &state; +}; + +template +static FailureOr> mapLoadCbufToS4Pointers(LoadOp op, typename LoadOp::Adaptor adaptor) { + Value sourceRaw = adaptor.getSource(); + Value destinationRaw = adaptor.getDestination(); + if (!sourceRaw || !destinationRaw || !isa(sourceRaw.getType()) || + !isa(destinationRaw.getType())) { + return failure(); + } + constexpr unsigned cbufAddressSpace = static_cast(pto::AddressSpace::MAT); + constexpr unsigned targetAddressSpace = std::is_same_v + ? static_cast(pto::AddressSpace::LEFT) + : static_cast(pto::AddressSpace::RIGHT); + FailureOr source = reinterpretPointerToAddrSpace(op, sourceRaw, cbufAddressSpace); + FailureOr destination = reinterpretPointerToAddrSpace(op, destinationRaw, targetAddressSpace); + if (failed(source) || failed(destination)) { + return failure(); + } + return std::pair{*source, *destination}; +} + +template static FailureOr buildLoadCbufToS4Callee(LoadOp op) { + if constexpr (std::is_same_v) { + return buildLoadCbufToCaS4Callee(op.getContext(), op.getSource().getType()); + } + return buildLoadCbufToCbS4Callee(op.getContext(), op.getSource().getType()); +} + +template class LowerLoadCbufToS4OpPattern final : public OpConversionPattern { +public: + explicit LowerLoadCbufToS4OpPattern(TypeConverter &typeConverter, MLIRContext *context, LoweringState &state) + : OpConversionPattern(typeConverter, context), state(state) {} + + LogicalResult matchAndRewrite(LoadOp op, typename LoadOp::Adaptor adaptor, + ConversionPatternRewriter &rewriter) const override { + FailureOr> pointers = mapLoadCbufToS4Pointers(op, adaptor); + if (failed(pointers)) { + return rewriter.notifyMatchFailure(op, "failed to map cbuf/cube pointer spaces"); + } + + FailureOr config0 = + packLoadCbufToS4Config0(op, adaptor.getMStart(), adaptor.getKStart(), adaptor.getMStep(), adaptor.getKStep()); + FailureOr config1 = packLoadCbufToS4Config1(op, adaptor.getSrcStride(), adaptor.getDstStride()); + if (failed(config0) || failed(config1)) { + return rewriter.notifyMatchFailure(op, "failed to pack load_cbuf_to_*_s4 config"); + } + + Value transpose = castIntegerLikeTo(op, adaptor.getTranspose(), rewriter.getI64Type()); + if (!transpose) { + return rewriter.notifyMatchFailure(op, "failed to cast transpose to i64"); + } + + FailureOr calleeName = buildLoadCbufToS4Callee(op); + if (failed(calleeName)) { + return rewriter.notifyMatchFailure(op, "unsupported load_cbuf_to_*_s4 element type"); + } + Type i64Ty = rewriter.getI64Type(); + auto funcType = rewriter.getFunctionType( + TypeRange{pointers->second.getType(), pointers->first.getType(), i64Ty, i64Ty, i64Ty}, TypeRange{}); + rewriter.create(op.getLoc(), *calleeName, TypeRange{}, + ValueRange{pointers->second, pointers->first, *config0, *config1, transpose}); + state.plannedDecls.push_back(PlannedDecl{calleeName->str(), funcType}); + rewriter.eraseOp(op); + return success(); + } + +private: + LoweringState &state; +}; + +class LowerLoadCbufToCbOpPattern final : public OpConversionPattern { +public: + explicit LowerLoadCbufToCbOpPattern(TypeConverter &typeConverter, MLIRContext *context, LoweringState &state) + : OpConversionPattern(typeConverter, context), state(state) {} + + LogicalResult matchAndRewrite(pto::LoadCbufToCbOp op, pto::LoadCbufToCbOp::Adaptor adaptor, + ConversionPatternRewriter &rewriter) const override { + Value sourceRaw = adaptor.getSource(); + Value destinationRaw = adaptor.getDestination(); + Value mStart = adaptor.getMStart(); + Value kStart = adaptor.getKStart(); + Value mStep = adaptor.getMStep(); + Value kStep = adaptor.getKStep(); + Value srcStride = adaptor.getSrcStride(); + Value dstStride = adaptor.getDstStride(); + if (!sourceRaw || !destinationRaw || !mStart || !kStart || !mStep || !kStep || !srcStride || !dstStride) { + return rewriter.notifyMatchFailure(op, "expected converted operands"); + } + + if (!isa(sourceRaw.getType()) || !isa(destinationRaw.getType())) { + return rewriter.notifyMatchFailure(op, "expected LLVM pointer src/dst"); + } + + Type i64Ty = rewriter.getI64Type(); + + constexpr unsigned cbufAddressSpace = static_cast(pto::AddressSpace::MAT); + constexpr unsigned cbAddressSpace = static_cast(pto::AddressSpace::RIGHT); + FailureOr source = reinterpretPointerToAddrSpace(op, sourceRaw, cbufAddressSpace); + FailureOr destination = reinterpretPointerToAddrSpace(op, destinationRaw, cbAddressSpace); + if (failed(source) || failed(destination)) { + return rewriter.notifyMatchFailure(op, "failed to map cbuf/cb pointer spaces"); + } + + bool transpose = op.getTranspose(); + FailureOr config0 = packLoadCbufToCbConfig0(op, mStart, kStart, mStep, kStep); + FailureOr config1 = packLoadCbufToCbConfig1(op, srcStride, dstStride); + if (failed(config0) || failed(config1)) { + return rewriter.notifyMatchFailure(op, "failed to pack load_cbuf_to_cb config"); + } + Value transposeValue = getI64Constant(rewriter, op.getLoc(), transpose ? 1 : 0); + + FailureOr calleeName = buildLoadCbufToCbCallee(op.getContext(), op.getSource().getType()); + if (failed(calleeName)) { + return rewriter.notifyMatchFailure(op, "unsupported load_cbuf_to_cb element type"); + } + auto funcType = rewriter.getFunctionType(TypeRange{destination->getType(), source->getType(), i64Ty, i64Ty, i64Ty}, + TypeRange{}); + rewriter.create(op.getLoc(), *calleeName, TypeRange{}, + ValueRange{*destination, *source, *config0, *config1, transposeValue}); + state.plannedDecls.push_back(PlannedDecl{calleeName->str(), funcType}); + rewriter.eraseOp(op); + return success(); + } + +private: + LoweringState &state; +}; + +class LowerLoadCbufToCaMxOpPattern final : public OpConversionPattern { +public: + explicit LowerLoadCbufToCaMxOpPattern(TypeConverter &typeConverter, MLIRContext *context, LoweringState &state) + : OpConversionPattern(typeConverter, context), state(state) {} + + LogicalResult matchAndRewrite(pto::LoadCbufToCaMxOp op, OpAdaptor adaptor, + ConversionPatternRewriter &rewriter) const override { + Value srcRaw = adaptor.getSource(); + Value dstRaw = adaptor.getDestination(); + if (!srcRaw || !dstRaw || !adaptor.getXStartPosition() || !adaptor.getYStartPosition() || !adaptor.getXStep() || + !adaptor.getYStep() || !adaptor.getSrcStride() || !adaptor.getDstStride()) { + return rewriter.notifyMatchFailure(op, "expected converted operands"); + } + if (!isa(srcRaw.getType()) || !isa(dstRaw.getType())) { + return rewriter.notifyMatchFailure(op, "expected LLVM pointer src/dst"); + } + + constexpr unsigned cbufAddressSpace = static_cast(pto::AddressSpace::MAT); + constexpr unsigned caAddressSpace = static_cast(pto::AddressSpace::LEFT); + FailureOr src = reinterpretPointerToAddrSpace(op, srcRaw, cbufAddressSpace); + FailureOr dst = reinterpretPointerToAddrSpace(op, dstRaw, caAddressSpace); + if (failed(src) || failed(dst)) { + return rewriter.notifyMatchFailure(op, "failed to map cbuf/ca pointer spaces"); + } + + Type sourceElemType = cast(op.getSource().getType()).getElementType(); + unsigned elemBitWidth = pto::getPTOStorageElemBitWidth(sourceElemType); + if (elemBitWidth == 0 || (elemBitWidth % 8) != 0) { + return rewriter.notifyMatchFailure(op, "unsupported load_cbuf_to_ca_mx element type"); + } + FailureOr config0 = packLoadCbufToCaConfig0(op, adaptor.getXStartPosition(), adaptor.getYStartPosition(), + adaptor.getXStep(), adaptor.getYStep()); + FailureOr config1 = packLoadCbufToCaConfig1(op, adaptor.getSrcStride(), adaptor.getDstStride()); + if (failed(config0) || failed(config1)) { + return rewriter.notifyMatchFailure(op, "failed to pack load_cbuf_to_ca_mx config"); + } + auto i64Ty = rewriter.getI64Type(); + Value dstAddr = rewriter.create(op.getLoc(), i64Ty, *dst); + + StringRef calleeName = buildLoadCbufToCaMxCallee(op.getContext()); + auto funcType = rewriter.getFunctionType(TypeRange{i64Ty, src->getType(), i64Ty, i64Ty}, TypeRange{}); + rewriter.create(op.getLoc(), calleeName, TypeRange{}, ValueRange{dstAddr, *src, *config0, *config1}); + state.plannedDecls.push_back(PlannedDecl{calleeName.str(), funcType}); + rewriter.eraseOp(op); + return success(); + } + +private: + LoweringState &state; +}; + +class LowerLoadCbufToCbMxOpPattern final : public OpConversionPattern { +public: + explicit LowerLoadCbufToCbMxOpPattern(TypeConverter &typeConverter, MLIRContext *context, LoweringState &state) + : OpConversionPattern(typeConverter, context), state(state) {} + + LogicalResult matchAndRewrite(pto::LoadCbufToCbMxOp op, OpAdaptor adaptor, + ConversionPatternRewriter &rewriter) const override { + Value srcRaw = adaptor.getSource(); + Value dstRaw = adaptor.getDestination(); + if (!srcRaw || !dstRaw || !adaptor.getXStartPosition() || !adaptor.getYStartPosition() || !adaptor.getXStep() || + !adaptor.getYStep() || !adaptor.getSrcStride() || !adaptor.getDstStride()) { + return rewriter.notifyMatchFailure(op, "expected converted operands"); + } + if (!isa(srcRaw.getType()) || !isa(dstRaw.getType())) { + return rewriter.notifyMatchFailure(op, "expected LLVM pointer src/dst"); + } + + constexpr unsigned cbufAddressSpace = static_cast(pto::AddressSpace::MAT); + constexpr unsigned cbAddressSpace = static_cast(pto::AddressSpace::RIGHT); + FailureOr src = reinterpretPointerToAddrSpace(op, srcRaw, cbufAddressSpace); + FailureOr dst = reinterpretPointerToAddrSpace(op, dstRaw, cbAddressSpace); + if (failed(src) || failed(dst)) { + return rewriter.notifyMatchFailure(op, "failed to map cbuf/cb pointer spaces"); + } + + Type sourceElemType = cast(op.getSource().getType()).getElementType(); + unsigned elemBitWidth = pto::getPTOStorageElemBitWidth(sourceElemType); + if (elemBitWidth == 0 || (elemBitWidth % 8) != 0) { + return rewriter.notifyMatchFailure(op, "unsupported load_cbuf_to_cb_mx element type"); + } + FailureOr config0 = packLoadCbufToCbConfig0(op, adaptor.getXStartPosition(), adaptor.getYStartPosition(), + adaptor.getXStep(), adaptor.getYStep()); + FailureOr config1 = packLoadCbufToCbConfig1(op, adaptor.getSrcStride(), adaptor.getDstStride()); + if (failed(config0) || failed(config1)) { + return rewriter.notifyMatchFailure(op, "failed to pack load_cbuf_to_cb_mx config"); + } + auto i64Ty = rewriter.getI64Type(); + Value dstAddr = rewriter.create(op.getLoc(), i64Ty, *dst); + + StringRef calleeName = buildLoadCbufToCbMxCallee(op.getContext()); + auto funcType = rewriter.getFunctionType(TypeRange{i64Ty, src->getType(), i64Ty, i64Ty}, TypeRange{}); + rewriter.create(op.getLoc(), calleeName, TypeRange{}, ValueRange{dstAddr, *src, *config0, *config1}); + state.plannedDecls.push_back(PlannedDecl{calleeName.str(), funcType}); + rewriter.eraseOp(op); + return success(); + } + +private: + LoweringState &state; +}; + +class LowerCopyMatrixCcToGmOpPattern final : public OpConversionPattern { +public: + explicit LowerCopyMatrixCcToGmOpPattern(TypeConverter &typeConverter, MLIRContext *context, LoweringState &state) + : OpConversionPattern(typeConverter, context), state(state) {} + + LogicalResult matchAndRewrite(pto::CopyMatrixCcToGmOp op, pto::CopyMatrixCcToGmOp::Adaptor adaptor, + ConversionPatternRewriter &rewriter) const override { + Value sourceRaw = adaptor.getSource(); + Value destinationRaw = adaptor.getDestination(); + Value xm = adaptor.getXm(); + Value xt = adaptor.getXt(); + if (!sourceRaw || !destinationRaw || !xm || !xt) { + return rewriter.notifyMatchFailure(op, "expected converted operands"); + } + + if (!isa(sourceRaw.getType()) || !isa(destinationRaw.getType())) { + return rewriter.notifyMatchFailure(op, "expected LLVM pointer src/dst"); + } + + Type i64Ty = rewriter.getI64Type(); + if (xm.getType() != i64Ty || xt.getType() != i64Ty) { + return rewriter.notifyMatchFailure(op, "expected i64 xm/xt operands"); + } + + constexpr unsigned gmAddressSpace = static_cast(pto::AddressSpace::GM); + constexpr unsigned ccAddressSpace = static_cast(pto::AddressSpace::ACC); + FailureOr source = reinterpretPointerToAddrSpace(op, sourceRaw, ccAddressSpace); + FailureOr destination = reinterpretPointerToAddrSpace(op, destinationRaw, gmAddressSpace); + if (failed(source) || failed(destination)) { + return rewriter.notifyMatchFailure(op, "failed to map cc/gm pointer spaces"); + } + + StringRef calleeName = buildCopyMatrixCcToGmCallee(op.getContext()); + auto funcType = + rewriter.getFunctionType(TypeRange{destination->getType(), source->getType(), i64Ty, i64Ty}, TypeRange{}); + rewriter.create(op.getLoc(), calleeName, TypeRange{}, ValueRange{*destination, *source, xm, xt}); + state.plannedDecls.push_back(PlannedDecl{calleeName.str(), funcType}); + rewriter.eraseOp(op); + return success(); + } + +private: + LoweringState &state; +}; + +template class LowerCopyMatrixCcToBufOpPattern final : public OpConversionPattern { +public: + explicit LowerCopyMatrixCcToBufOpPattern(TypeConverter &typeConverter, MLIRContext *context, LoweringState &state) + : OpConversionPattern(typeConverter, context), state(state) {} + + LogicalResult matchAndRewrite(CopyOp op, typename CopyOp::Adaptor adaptor, + ConversionPatternRewriter &rewriter) const override { + Value sourceRaw = adaptor.getSource(); + Value destinationRaw = adaptor.getDestination(); + if (!sourceRaw || !destinationRaw) { + return rewriter.notifyMatchFailure(op, "expected converted operands"); + } + if (!isa(sourceRaw.getType()) || !isa(destinationRaw.getType())) { + return rewriter.notifyMatchFailure(op, "expected LLVM pointer src/dst"); + } + + constexpr unsigned ccAddressSpace = static_cast(pto::AddressSpace::ACC); + constexpr unsigned targetAddressSpace = std::is_same_v + ? static_cast(pto::AddressSpace::MAT) + : static_cast(pto::AddressSpace::VEC); + FailureOr source = reinterpretPointerToAddrSpace(op, sourceRaw, ccAddressSpace); + FailureOr destination = reinterpretPointerToAddrSpace(op, destinationRaw, targetAddressSpace); + if (failed(source) || failed(destination)) { + return rewriter.notifyMatchFailure(op, "failed to map cc->buf pointer spaces"); + } + + Type i64Ty = rewriter.getI64Type(); + Value config0 = castIntegerLikeTo(op, adaptor.getConfig0(), i64Ty); + Value config1 = castIntegerLikeTo(op, adaptor.getConfig1(), i64Ty); + if (!config0 || !config1) { + return rewriter.notifyMatchFailure(op, "failed to cast config operands to i64"); + } + + FailureOr calleeName = std::is_same_v + ? FailureOr(buildCopyMatrixCcToCbufCallee(op.getContext())) + : buildCopyMatrixCcToUbCallee(op.getContext(), op.getDestination().getType()); + if (failed(calleeName)) { + return rewriter.notifyMatchFailure(op, "unsupported copy_matrix_cc_to_{cbuf,ub} element type"); + } + auto funcType = + rewriter.getFunctionType(TypeRange{destination->getType(), source->getType(), i64Ty, i64Ty}, TypeRange{}); + rewriter.create(op.getLoc(), *calleeName, TypeRange{}, + ValueRange{*destination, *source, config0, config1}); + state.plannedDecls.push_back(PlannedDecl{calleeName->str(), funcType}); + rewriter.eraseOp(op); + return success(); + } + +private: + LoweringState &state; +}; + +template class LowerVecScalarMaskedOpPattern final : public OpConversionPattern { +public: + explicit LowerVecScalarMaskedOpPattern(TypeConverter &typeConverter, MLIRContext *context, LoweringState &state) + : OpConversionPattern(typeConverter, context), state(state) {} + + LogicalResult matchAndRewrite(VecScalarOp op, typename VecScalarOp::Adaptor adaptor, + ConversionPatternRewriter &rewriter) const override { + StringRef stem = getVecScalarMaskedStem(); + FailureOr calleeName = + usesSignedVecScalarCANN900Callee() + ? buildCANN900SignedModeTypedCallee(op.getContext(), op.getResult().getType(), stem, "x") + : buildCANN900ModeTypedCallee(op.getContext(), op.getResult().getType(), stem, "x"); + if (failed(calleeName)) { + return rewriter.notifyMatchFailure(op, "unsupported vec-scalar VPTO signature"); + } + + Type resultType = this->getTypeConverter()->convertType(op.getResult().getType()); + if (!resultType) { + return rewriter.notifyMatchFailure(op, "failed to convert vec-scalar result type"); + } + + Value input = adaptor.getOperands()[0]; + Value scalar = adaptor.getOperands()[1]; + Value mask = adaptor.getOperands()[2]; + Type expectedMaskType = this->getTypeConverter()->convertType(op->getOperand(2).getType()); + if (!input || !scalar || !mask || input.getType() != resultType || mask.getType() != expectedMaskType) { + return rewriter.notifyMatchFailure(op, "unexpected converted vec-scalar VPTO operand types"); + } + + auto call = + rewriter.create(op.getLoc(), *calleeName, TypeRange{resultType}, ValueRange{input, scalar, mask}); + state.plannedDecls.push_back(PlannedDecl{calleeName->str(), call.getCalleeType()}); + rewriter.replaceOp(op, call.getResults()); + return success(); + } + +private: + LoweringState &state; +}; + +template class LowerReductionUnaryOpPattern final : public OpConversionPattern { +public: + explicit LowerReductionUnaryOpPattern(TypeConverter &typeConverter, MLIRContext *context, LoweringState &state) + : OpConversionPattern(typeConverter, context), state(state) {} + + LogicalResult matchAndRewrite(ReductionOp op, typename ReductionOp::Adaptor adaptor, + ConversionPatternRewriter &rewriter) const override { + StringRef stem = getReductionUnaryStem(); + FailureOr calleeName = + usesSignedReductionCANN900Callee() + ? buildCANN900SignedModeTypedCallee(op.getContext(), op.getResult().getType(), stem, "x") + : buildCANN900ModeTypedCallee(op.getContext(), op.getResult().getType(), stem, "x"); + if (failed(calleeName)) { + return rewriter.notifyMatchFailure(op, "unsupported reduction VPTO signature"); + } + + Type resultType = this->getTypeConverter()->convertType(op.getResult().getType()); + Type maskType = this->getTypeConverter()->convertType(op.getMask().getType()); + if (!resultType || !maskType) { + return rewriter.notifyMatchFailure(op, "failed to convert reduction result type"); + } + + Value input = adaptor.getInput(); + Value mask = adaptor.getMask(); + if (!input || !mask || input.getType() != resultType || mask.getType() != maskType) { + return rewriter.notifyMatchFailure(op, "unexpected converted reduction operand types"); + } + + auto call = rewriter.create(op.getLoc(), *calleeName, TypeRange{resultType}, ValueRange{input, mask}); + state.plannedDecls.push_back(PlannedDecl{calleeName->str(), call.getCalleeType()}); + rewriter.replaceOp(op, call.getResults()); + return success(); + } + +private: + LoweringState &state; +}; + +template class LowerHistogramOpPattern final : public OpConversionPattern { +public: + explicit LowerHistogramOpPattern(TypeConverter &typeConverter, MLIRContext *context, LoweringState &state) + : OpConversionPattern(typeConverter, context), state(state) {} + + LogicalResult matchAndRewrite(HistOp op, typename HistOp::Adaptor adaptor, + ConversionPatternRewriter &rewriter) const override { + StringRef calleeName = getHistogramCallee(op.getContext()); + if (calleeName.empty()) { + return rewriter.notifyMatchFailure(op, "unsupported histogram op"); + } + + Type resultType = this->getTypeConverter()->convertType(op.getResult().getType()); + Type sourceType = this->getTypeConverter()->convertType(op.getSource().getType()); + Type maskType = this->getTypeConverter()->convertType(op.getMask().getType()); + if (!resultType || !sourceType || !maskType) { + return rewriter.notifyMatchFailure(op, "failed to convert histogram types"); + } + + Value acc = adaptor.getAcc(); + Value source = adaptor.getSource(); + Value mask = adaptor.getMask(); + Value bin = adaptor.getBin(); + if (!acc || !source || !mask || !bin || acc.getType() != resultType || source.getType() != sourceType || + mask.getType() != maskType || !bin.getType().isInteger(32)) { + return rewriter.notifyMatchFailure(op, "unexpected converted histogram operand types"); + } + + auto funcType = rewriter.getFunctionType(TypeRange{resultType, sourceType, maskType, rewriter.getI32Type()}, + TypeRange{resultType}); + auto call = rewriter.create(op.getLoc(), calleeName, TypeRange{resultType}, + ValueRange{acc, source, mask, bin}); + state.plannedDecls.push_back(PlannedDecl{calleeName.str(), funcType}); + rewriter.replaceOp(op, call.getResults()); + return success(); + } + +private: + LoweringState &state; +}; + +template class LowerExtremaPredicateOpPattern final : public OpConversionPattern { +public: + explicit LowerExtremaPredicateOpPattern(TypeConverter &typeConverter, MLIRContext *context, LoweringState &state) + : OpConversionPattern(typeConverter, context), state(state) {} + + LogicalResult matchAndRewrite(ExtremaOp op, typename ExtremaOp::Adaptor adaptor, + ConversionPatternRewriter &rewriter) const override { + FailureOr calleeName = buildExtremaPredicateCallee(op.getContext(), op.getValue().getType()); + if (failed(calleeName)) { + return rewriter.notifyMatchFailure(op, "unsupported extrema-predicate VPTO signature"); + } + + Type valueType = this->getTypeConverter()->convertType(op.getValue().getType()); + Type predicateType = this->getTypeConverter()->convertType(op.getPredicate().getType()); + if (!valueType || !predicateType) { + return rewriter.notifyMatchFailure(op, "failed to convert extrema-predicate result types"); + } + + Value input = adaptor.getInput(); + Value mask = adaptor.getMask(); + if (!input || !mask || input.getType() != valueType || mask.getType() != predicateType) { + return rewriter.notifyMatchFailure(op, "unexpected converted extrema-predicate operand types"); + } + + auto call = rewriter.create(op.getLoc(), *calleeName, TypeRange{valueType, predicateType}, + ValueRange{input, mask}); + state.plannedDecls.push_back(PlannedDecl{calleeName->str(), call.getCalleeType()}); + rewriter.replaceOp(op, call.getResults()); + return success(); + } + +private: + LoweringState &state; +}; + +template +class LowerWideningReductionUnaryOpPattern final : public OpConversionPattern { +public: + explicit LowerWideningReductionUnaryOpPattern(TypeConverter &typeConverter, MLIRContext *context, + LoweringState &state) + : OpConversionPattern(typeConverter, context), state(state) {} + + LogicalResult matchAndRewrite(ReductionOp op, typename ReductionOp::Adaptor adaptor, + ConversionPatternRewriter &rewriter) const override { + StringRef stem = getReductionUnaryStem(); + FailureOr calleeName = buildCANN900WideningReductionCallee(op.getContext(), op.getInput().getType(), + op.getResult().getType(), stem, "x"); + if (failed(calleeName)) { + return rewriter.notifyMatchFailure(op, "unsupported widening reduction VPTO signature"); + } + + Type inputType = this->getTypeConverter()->convertType(op.getInput().getType()); + Type resultType = this->getTypeConverter()->convertType(op.getResult().getType()); + Type maskType = this->getTypeConverter()->convertType(op.getMask().getType()); + if (!inputType || !resultType || !maskType) { + return rewriter.notifyMatchFailure(op, "failed to convert widening reduction types"); + } + + Value input = adaptor.getInput(); + Value mask = adaptor.getMask(); + if (!input || !mask || input.getType() != inputType || mask.getType() != maskType) { + return rewriter.notifyMatchFailure(op, "unexpected converted widening reduction operand types"); + } + + auto call = rewriter.create(op.getLoc(), *calleeName, TypeRange{resultType}, ValueRange{input, mask}); + state.plannedDecls.push_back(PlannedDecl{calleeName->str(), call.getCalleeType()}); + rewriter.replaceOp(op, call.getResults()); + return success(); + } + +private: + LoweringState &state; +}; + +class LowerVselOpPattern final : public OpConversionPattern { +public: + explicit LowerVselOpPattern(TypeConverter &typeConverter, MLIRContext *context, LoweringState &state) + : OpConversionPattern(typeConverter, context), state(state) {} + + LogicalResult matchAndRewrite(pto::VselOp op, pto::VselOp::Adaptor adaptor, + ConversionPatternRewriter &rewriter) const override { + FailureOr calleeName = buildVselCallee(op.getContext(), op.getResult().getType()); + if (failed(calleeName)) { + return rewriter.notifyMatchFailure(op, "unsupported vsel VPTO signature"); + } + + Type resultType = this->getTypeConverter()->convertType(op.getResult().getType()); + Type maskType = this->getTypeConverter()->convertType(op.getMask().getType()); + if (!resultType || !maskType) { + return rewriter.notifyMatchFailure(op, "failed to convert vsel result type"); + } + + Value src0 = adaptor.getSrc0(); + Value src1 = adaptor.getSrc1(); + Value mask = adaptor.getMask(); + if (!src0 || !src1 || !mask || src0.getType() != resultType || src1.getType() != resultType || + mask.getType() != maskType) { + return rewriter.notifyMatchFailure(op, "unexpected converted vsel operand types"); + } + + auto call = + rewriter.create(op.getLoc(), *calleeName, TypeRange{resultType}, ValueRange{src0, src1, mask}); + state.plannedDecls.push_back(PlannedDecl{calleeName->str(), call.getCalleeType()}); + rewriter.replaceOp(op, call.getResults()); + return success(); + } + +private: + LoweringState &state; +}; + +class LowerVdupOpPattern final : public OpConversionPattern { +public: + explicit LowerVdupOpPattern(TypeConverter &typeConverter, MLIRContext *context, LoweringState &state) + : OpConversionPattern(typeConverter, context), state(state) {} + + LogicalResult matchAndRewrite(pto::VdupOp op, pto::VdupOp::Adaptor adaptor, + ConversionPatternRewriter &rewriter) const override { + FailureOr calleeName = buildVdupCallee(op.getContext(), op); + if (failed(calleeName)) { + return rewriter.notifyMatchFailure(op, "unsupported vdup VPTO signature"); + } + + Type resultType = this->getTypeConverter()->convertType(op.getResult().getType()); + Type maskType = this->getTypeConverter()->convertType(op.getMask().getType()); + if (!resultType || !maskType) { + return rewriter.notifyMatchFailure(op, "failed to convert vdup result type"); + } + + Value mask = adaptor.getMask(); + if (!mask || mask.getType() != maskType) { + return rewriter.notifyMatchFailure(op, "unexpected converted vdup mask type"); + } + + SmallVector callArgs; + bool vectorInput = isa(op.getInput().getType()); + if (vectorInput) { + Value input = adaptor.getInput(); + if (!input || input.getType() != resultType) { + return rewriter.notifyMatchFailure(op, "vector-input vdup requires matching result type"); + } + callArgs.push_back(input); + } else { + Type scalarType = getElementTypeFromVectorLike(op.getResult().getType()); + if (!scalarType || (op.getInput().getType() != scalarType && + !isCompatibleScalarForSemanticType(scalarType, op.getInput().getType()))) { + return rewriter.notifyMatchFailure(op, "unexpected scalar-input vdup type"); + } + FailureOr normalizedScalar = + normalizeVdupScalarOperand(rewriter, op.getLoc(), adaptor.getInput(), op.getResult().getType()); + if (failed(normalizedScalar)) { + return rewriter.notifyMatchFailure(op, "failed to normalize scalar vdup input"); + } + Value scalarForCall = + normalizeByteScalarOperandForCANN900VectorCall(rewriter, op.getLoc(), *normalizedScalar, scalarType); + callArgs.push_back(scalarForCall); + } + + callArgs.push_back(mask); + callArgs.push_back(getI32Constant(rewriter, op.getLoc(), 1)); + + auto call = rewriter.create(op.getLoc(), *calleeName, TypeRange{resultType}, callArgs); + state.plannedDecls.push_back(PlannedDecl{calleeName->str(), call.getCalleeType()}); + rewriter.replaceOp(op, call.getResults()); + return success(); + } + +private: + LoweringState &state; +}; + +class LowerVbrOpPattern final : public OpConversionPattern { +public: + explicit LowerVbrOpPattern(TypeConverter &typeConverter, MLIRContext *context, LoweringState &state) + : OpConversionPattern(typeConverter, context), state(state) {} + + LogicalResult matchAndRewrite(pto::VbrOp op, pto::VbrOp::Adaptor adaptor, + ConversionPatternRewriter &rewriter) const override { + FailureOr calleeName = buildVbrCallee(op.getContext(), op.getResult().getType()); + if (failed(calleeName)) { + return rewriter.notifyMatchFailure(op, "unsupported vbr VPTO signature"); + } + + Type resultType = this->getTypeConverter()->convertType(op.getResult().getType()); + if (!resultType) { + return rewriter.notifyMatchFailure(op, "failed to convert vbr result type"); + } + + Value scalar = adaptor.getValue(); + Type expectedScalarType = this->getTypeConverter()->convertType(op.getValue().getType()); + if (!scalar || !expectedScalarType || scalar.getType() != expectedScalarType) { + return rewriter.notifyMatchFailure(op, "unexpected converted vbr operand type"); + } + + scalar = normalizeByteScalarOperandForCANN900VectorCall( + rewriter, op.getLoc(), scalar, cast(op.getResult().getType()).getElementType()); + + auto funcType = rewriter.getFunctionType(TypeRange{scalar.getType()}, TypeRange{resultType}); + auto call = rewriter.create(op.getLoc(), *calleeName, TypeRange{resultType}, ValueRange{scalar}); + state.plannedDecls.push_back(PlannedDecl{calleeName->str(), funcType}); + rewriter.replaceOp(op, call.getResults()); + return success(); + } + +private: + LoweringState &state; +}; + +static FailureOr getVselrIntrinsicResultType(pto::VselrOp op, Type resultType, PatternRewriter &rewriter) { + auto resultVectorType = dyn_cast(resultType); + if (!resultVectorType) { + return failure(); + } + Type intrinsicResultType = resultType; + if (auto floatType = dyn_cast(resultVectorType.getElementType()); floatType && floatType.isF32()) { + intrinsicResultType = + VectorType::get(resultVectorType.getShape(), rewriter.getI32Type(), resultVectorType.getScalableDims()); + } + if (Type carrierType = getLowpPayloadCarrierType(op.getResult().getType(), rewriter.getContext())) { + intrinsicResultType = carrierType; + } + return intrinsicResultType; +} + +class LowerVselrOpPattern final : public OpConversionPattern { +public: + explicit LowerVselrOpPattern(TypeConverter &typeConverter, MLIRContext *context, LoweringState &state) + : OpConversionPattern(typeConverter, context), state(state) {} + + LogicalResult matchAndRewrite(pto::VselrOp op, pto::VselrOp::Adaptor adaptor, + ConversionPatternRewriter &rewriter) const override { + FailureOr calleeName = buildVselrCallee(op.getContext(), op.getResult().getType()); + if (failed(calleeName)) { + return rewriter.notifyMatchFailure(op, "unsupported vselr VPTO signature"); + } + + Type resultType = this->getTypeConverter()->convertType(op.getResult().getType()); + if (!resultType) { + return rewriter.notifyMatchFailure(op, "unexpected converted vselr result type"); + } + FailureOr intrinsicResultType = getVselrIntrinsicResultType(op, resultType, rewriter); + if (failed(intrinsicResultType)) { + return rewriter.notifyMatchFailure(op, "unexpected converted vselr result type"); + } + + Type indexType = this->getTypeConverter()->convertType(op.getSrc1().getType()); + if (!indexType) { + return rewriter.notifyMatchFailure(op, "failed to convert vselr index type"); + } + + Value src0 = adaptor.getSrc0(); + Value src1 = adaptor.getSrc1(); + if (!src0 || !src1 || src1.getType() != indexType) { + return rewriter.notifyMatchFailure(op, "unexpected converted vselr operand types"); + } + + if (src0.getType() != *intrinsicResultType) { + if (src0.getType() != resultType) { + return rewriter.notifyMatchFailure(op, "unexpected converted vselr source type"); + } + src0 = rewriter.create(op.getLoc(), *intrinsicResultType, src0); + } + + auto funcType = + rewriter.getFunctionType(TypeRange{*intrinsicResultType, indexType}, TypeRange{*intrinsicResultType}); + auto call = rewriter.create(op.getLoc(), *calleeName, TypeRange{*intrinsicResultType}, + ValueRange{src0, src1}); + state.plannedDecls.push_back(PlannedDecl{calleeName->str(), funcType}); + + Value result = call.getResult(0); + if (*intrinsicResultType != resultType) { + result = rewriter.create(op.getLoc(), resultType, result); + } + rewriter.replaceOp(op, ValueRange{result}); + return success(); + } + +private: + LoweringState &state; +}; + +class LowerPnotOpPattern final : public OpConversionPattern { +public: + explicit LowerPnotOpPattern(TypeConverter &typeConverter, MLIRContext *context, LoweringState &state) + : OpConversionPattern(typeConverter, context), state(state) {} + + LogicalResult matchAndRewrite(pto::PnotOp op, pto::PnotOp::Adaptor adaptor, + ConversionPatternRewriter &rewriter) const override { + Type resultType = this->getTypeConverter()->convertType(op.getResult().getType()); + if (!resultType) { + return rewriter.notifyMatchFailure(op, "failed to convert pnot result type"); + } + + Value input = adaptor.getInput(); + Value mask = adaptor.getMask(); + if (!input || !mask || input.getType() != resultType || mask.getType() != resultType) { + return rewriter.notifyMatchFailure(op, "unexpected converted pnot operand types"); + } + + StringRef calleeName = getPredicateMaskCallee(op.getContext()); + auto call = rewriter.create(op.getLoc(), calleeName, TypeRange{resultType}, ValueRange{input, mask}); + state.plannedDecls.push_back(PlannedDecl{calleeName.str(), call.getCalleeType()}); + rewriter.replaceOp(op, call.getResults()); + return success(); + } + +private: + LoweringState &state; +}; + +template class LowerInterleaveOpPattern final : public OpConversionPattern { +public: + explicit LowerInterleaveOpPattern(TypeConverter &typeConverter, MLIRContext *context, LoweringState &state) + : OpConversionPattern(typeConverter, context), state(state) {} + + LogicalResult matchAndRewrite(InterleaveOp op, typename InterleaveOp::Adaptor adaptor, + ConversionPatternRewriter &rewriter) const override { + StringRef stem = std::is_same_v ? "vintlv" : "vdintlv"; + FailureOr calleeName = buildInterleaveCallee(op.getContext(), op.getLow().getType(), stem); + if (failed(calleeName)) { + return rewriter.notifyMatchFailure(op, "unsupported interleave VPTO signature"); + } + + Type lowType = this->getTypeConverter()->convertType(op.getLow().getType()); + Type highType = this->getTypeConverter()->convertType(op.getHigh().getType()); + if (!lowType || !highType || lowType != highType) { + return rewriter.notifyMatchFailure(op, "failed to convert interleave result types"); + } + + Value lhs = adaptor.getLhs(); + Value rhs = adaptor.getRhs(); + if (!lhs || !rhs || lhs.getType() != lowType || rhs.getType() != lowType) { + return rewriter.notifyMatchFailure(op, "unexpected converted interleave operand types"); + } + + auto funcType = rewriter.getFunctionType(TypeRange{lowType, lowType}, TypeRange{lowType, highType}); + auto call = + rewriter.create(op.getLoc(), *calleeName, TypeRange{lowType, highType}, ValueRange{lhs, rhs}); + state.plannedDecls.push_back(PlannedDecl{calleeName->str(), funcType}); + rewriter.replaceOp(op, call.getResults()); + return success(); + } + +private: + LoweringState &state; +}; + +template class LowerPredicatePackOpPattern final : public OpConversionPattern { +public: + explicit LowerPredicatePackOpPattern(TypeConverter &typeConverter, MLIRContext *context, LoweringState &state) + : OpConversionPattern(typeConverter, context), state(state) {} + + LogicalResult matchAndRewrite(PackOp op, typename PackOp::Adaptor adaptor, + ConversionPatternRewriter &rewriter) const override { + Type resultType = this->getTypeConverter()->convertType(op.getResult().getType()); + if (!resultType) { + return rewriter.notifyMatchFailure(op, "failed to convert predicate-pack result type"); + } + + auto part = parseHiLoPartImmediate(op.getPart()); + if (!part) { + return rewriter.notifyMatchFailure(op, "unsupported predicate-pack part immediate"); + } + + Value input = adaptor.getInput(); + if (!input || input.getType() != resultType) { + return rewriter.notifyMatchFailure(op, "unexpected converted predicate-pack operand type"); + } + + Value partValue = rewriter.create(op.getLoc(), rewriter.getI32IntegerAttr(*part)); + StringRef calleeName = getPredicatePackCallee(op.getContext()); + auto funcType = rewriter.getFunctionType(TypeRange{resultType, rewriter.getI32Type()}, TypeRange{resultType}); + auto call = + rewriter.create(op.getLoc(), calleeName, TypeRange{resultType}, ValueRange{input, partValue}); + state.plannedDecls.push_back(PlannedDecl{calleeName.str(), funcType}); + rewriter.replaceOp(op, call.getResults()); + return success(); + } + +private: + LoweringState &state; +}; + +void populateVPTOArithmeticPatterns(VPTOTypeConverter &typeConverter, RewritePatternSet &patterns, + LoweringState &state) { + patterns.add, LowerUnaryMaskedOpPattern, + LowerUnaryMaskedOpPattern, LowerUnaryMaskedOpPattern, + LowerUnaryMaskedOpPattern, LowerUnaryMaskedOpPattern, + LowerUnaryMaskedOpPattern, LowerVsqzOpPattern, LowerVusqzOpPattern, LowerVmulaOpPattern, + LowerVmullOpPattern, LowerBinaryMaskedOpPattern, LowerBinaryMaskedOpPattern, + LowerBinaryMaskedOpPattern, LowerBinaryMaskedOpPattern, + LowerBinaryMaskedOpPattern, LowerBinaryMaskedOpPattern, + LowerBinaryMaskedOpPattern, LowerBinaryMaskedOpPattern, + LowerBinaryMaskedOpPattern, LowerTernaryMaskedOpPattern, + LowerBinaryMaskedOpPattern, LowerCarryBinaryOpPattern, + LowerCarryBinaryOpPattern, LowerCarryBinaryOpPattern, + LowerCarryBinaryOpPattern, LowerBinaryMaskedOpPattern, + LowerBinaryMaskedOpPattern, LowerVecScalarMaskedOpPattern, + LowerVecScalarMaskedOpPattern, LowerVecScalarMaskedOpPattern, + LowerVecScalarMaskedOpPattern, LowerVecScalarMaskedOpPattern, + LowerVecScalarMaskedOpPattern, LowerVecScalarMaskedOpPattern, + LowerWideningReductionUnaryOpPattern, LowerReductionUnaryOpPattern, + LowerReductionUnaryOpPattern, LowerReductionUnaryOpPattern, + LowerReductionUnaryOpPattern, LowerReductionUnaryOpPattern, + LowerReductionUnaryOpPattern, LowerHistogramOpPattern, + LowerHistogramOpPattern, LowerExtremaPredicateOpPattern, + LowerExtremaPredicateOpPattern, LowerVdupOpPattern, LowerVbrOpPattern, + LowerPredicatePackOpPattern, LowerPredicatePackOpPattern, + LowerVselOpPattern, LowerVselrOpPattern, LowerPnotOpPattern, LowerInterleaveOpPattern, + LowerInterleaveOpPattern, LowerCopyGmToCbufOpPattern, LowerLoadCbufToCaOpPattern, + LowerLoadCbufToCbOpPattern, LowerLoadCbufToS4OpPattern, + LowerLoadCbufToS4OpPattern, LowerLoadCbufToCaMxOpPattern, + LowerLoadCbufToCbMxOpPattern, LowerCopyMatrixCcToGmOpPattern, + LowerCopyMatrixCcToBufOpPattern, + LowerCopyMatrixCcToBufOpPattern, LowerCopyCbufToBtOpPattern, + LowerCopyCbufToFbufOpPattern, LowerCopyGmToCbufMultiOpPattern, + LowerCopyGmToCbufMultiOpPattern, LowerMadRawPattern, + LowerMadRawPattern, LowerMadRawPattern, + LowerMadRawPattern, LowerCopyOpPattern, + LowerCopyOpPattern, LowerCopyUbufToUbufOpPattern, LowerCopyCbufToUbufOpPattern, + LowerCopyUbufToCbufOpPattern, LowerCreateCbufMatrixOpPattern>(typeConverter, patterns.getContext(), + state); +} + +} // namespace mlir::pto::detail diff --git a/lib/PTO/Transforms/VPTOCANN900LLVMEmitterCalleeCore.cpp b/lib/PTO/Transforms/VPTOCANN900LLVMEmitterCalleeCore.cpp new file mode 100644 index 0000000000..19a5016f26 --- /dev/null +++ b/lib/PTO/Transforms/VPTOCANN900LLVMEmitterCalleeCore.cpp @@ -0,0 +1,450 @@ +// Copyright (c) 2026 Huawei Technologies Co., Ltd. +// This program is free software, you can redistribute it and/or modify it under the terms and conditions of +// CANN Open Software License Agreement Version 2.0 (the "License"). +// Please refer to the License for details. You may not use this file except in compliance with the License. +// THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED, +// INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE. +// See LICENSE in the root of the software repository for the full text of the License. + +#include "VPTOCANN900LLVMEmitterInternal.h" + +namespace mlir::pto::detail { + +FailureOr buildCarryBinaryCallee(MLIRContext *context, Type resultType, StringRef stem) { + std::string vec = getElementTypeFragment(cast(resultType).getElementType()); + auto lanes = getElementCountFromVectorLike(resultType); + if (vec.empty() || !lanes) { + return failure(); + } + return StringAttr::get(context, "llvm.hivm." + stem.str() + ".v" + std::to_string(*lanes) + vec).getValue(); +} + +FailureOr buildVselCallee(MLIRContext *context, Type resultType) { + std::string vec = getCANN900VectorTypeFragment(resultType); + if (vec.empty()) { + return failure(); + } + return StringAttr::get(context, "llvm.hivm.vsel." + vec).getValue(); +} + +FailureOr buildVselrCallee(MLIRContext *context, Type resultType) { + Type elementType = getElementTypeFromVectorLike(resultType); + auto lanes = getElementCountFromVectorLike(resultType); + if (!elementType || !lanes) { + return failure(); + } + + std::optional abi = getLowpPayloadABI(elementType, context); + std::string vec = abi ? "v" + std::to_string(*lanes) + abi->intrinsicElementFragment.str() + : getCANN900VectorTypeFragment(resultType); + if (vec.empty()) { + return failure(); + } + return StringAttr::get(context, "llvm.hivm.vselr." + vec).getValue(); +} + +FailureOr buildVdupCallee(MLIRContext *context, pto::VdupOp op) { + Type inputType = op.getInput().getType(); + Type resultType = op.getResult().getType(); + std::string vec = getCANN900VectorTypeFragment(resultType); + if (vec.empty()) { + return failure(); + } + + if (isa(inputType)) { + StringRef position = op.getPosition().value_or("LOWEST"); + StringRef family = position == "HIGHEST" ? "vdupm" : "vdup"; + return StringAttr::get(context, "llvm.hivm." + family.str() + ".z." + vec).getValue(); + } + + return StringAttr::get(context, "llvm.hivm.vdups.z." + vec).getValue(); +} + +FailureOr buildVbrCallee(MLIRContext *context, Type resultType) { + std::string vec = getCANN900VectorTypeFragment(resultType); + if (vec.empty()) { + return failure(); + } + return StringAttr::get(context, "llvm.hivm.vbr." + vec).getValue(); +} + +FailureOr buildPstuCallee(MLIRContext *context, pto::PstuOp op) { + if (auto maskType = dyn_cast(op.getValue().getType())) { + if (maskType.isB16()) { + return StringAttr::get(context, "llvm.hivm.pstu.b16").getValue(); + } + if (maskType.isB32()) { + return StringAttr::get(context, "llvm.hivm.pstu.b32").getValue(); + } + } + return failure(); +} + +FailureOr buildVstusCallee(MLIRContext *context, Type valueType) { + std::string vec = getMemoryElementTypeFragment(getElementTypeFromVectorLike(valueType)); + auto lanes = getElementCountFromVectorLike(valueType); + if (vec.empty() || !lanes) { + return failure(); + } + return StringAttr::get(context, "llvm.hivm.vstus.v" + std::to_string(*lanes) + vec).getValue(); +} + +FailureOr buildVstusPostCallee(MLIRContext *context, Type valueType) { + std::string vec = getMemoryElementTypeFragment(getElementTypeFromVectorLike(valueType)); + auto lanes = getElementCountFromVectorLike(valueType); + if (vec.empty() || !lanes) { + return failure(); + } + return StringAttr::get(context, "llvm.hivm.vstus.post.v" + std::to_string(*lanes) + vec).getValue(); +} + +StringRef buildVsturCallee(MLIRContext *context) { return StringAttr::get(context, "llvm.hivm.vstur").getValue(); } + +StringRef buildInitAlignCallee(MLIRContext *context) { + return StringAttr::get(context, "llvm.hivm.init.vector.align.data").getValue(); +} + +StringRef buildSprclrCallee(MLIRContext *context) { return StringAttr::get(context, "llvm.hivm.sprclr").getValue(); } + +StringRef buildSprstiCallee(MLIRContext *context, bool post) { + return StringAttr::get(context, post ? "llvm.hivm.sprsti.post" : "llvm.hivm.sprsti").getValue(); +} + +StringRef buildSprstsCallee(MLIRContext *context, bool post) { + return StringAttr::get(context, post ? "llvm.hivm.sprsts.post" : "llvm.hivm.sprsts").getValue(); +} + +StringRef buildStoreVfSimtInfoCallee(MLIRContext *context) { + return StringAttr::get(context, "llvm.hivm.store.vfsimt.info").getValue(); +} + +StringRef buildSyncthreadsCallee(MLIRContext *context) { + return StringAttr::get(context, "llvm.hivm.sync.workitems").getValue(); +} + +StringRef buildThreadfenceCallee(MLIRContext *context) { + return StringAttr::get(context, "llvm.hivm.fence.workitems").getValue(); +} + +StringRef buildThreadfenceBlockCallee(MLIRContext *context) { + return StringAttr::get(context, "llvm.hivm.fenceblock.workitems").getValue(); +} + +StringRef buildVstarCallee(MLIRContext *context) { return StringAttr::get(context, "llvm.hivm.vstar").getValue(); } + +StringRef buildVstasCallee(MLIRContext *context, bool post) { + return StringAttr::get(context, post ? "llvm.hivm.vstas.post" : "llvm.hivm.vstas").getValue(); +} + +Value buildShuffleControlValue(OpBuilder &builder, Location loc, Value controlValue, int64_t widthValue, + unsigned controlMask) { + Value lowBits = builder.create(loc, controlValue, getI32Constant(builder, loc, 0x1f)); + Value encodedWidth = getI32Constant(builder, loc, static_cast(32 - widthValue) << 16); + Value encodedMask = getI32Constant(builder, loc, static_cast(controlMask) << 8); + Value highBits = builder.create(loc, encodedWidth, encodedMask); + return builder.create(loc, highBits, lowBits); +} + +FailureOr buildAtomicCalleeName(MLIRContext *context, Type ptrType, Type valueType, Attribute signednessAttr, + StringRef opName) { + std::string elem = getAtomicElementTypeFragment(valueType, signednessAttr); + if (elem.empty()) { + return failure(); + } + auto ptrTy = dyn_cast(ptrType); + if (!ptrTy) { + return failure(); + } + + StringRef space; + switch (ptrTy.getMemorySpace().getAddressSpace()) { + case pto::AddressSpace::GM: + space = "G"; + break; + case pto::AddressSpace::VEC: + if (valueType.isInteger(64)) { + return failure(); + } + space = "S"; + break; + default: + return failure(); + } + + return StringAttr::get(context, "llvm.hivm.atom." + opName.str() + "." + space.str() + "." + elem).getValue(); +} + +FailureOr buildL1CacheLoadCallee(MLIRContext *context, Type resultType, pto::L1Cache l1cache) { + std::string elem; + if (auto intType = dyn_cast(resultType)) { + if (intType.getWidth() == 8) { + elem = "s8"; + } else if (intType.getWidth() == 16) { + elem = "s16"; + } else if (intType.getWidth() == 32) { + elem = "s32"; + } else if (intType.getWidth() == 64) { + elem = "s64"; + } + } else if (resultType.isF16() || resultType.isBF16()) { + elem = "s16"; + } else if (resultType.isF32()) { + elem = "s32"; + } else if (resultType.isF64()) { + elem = "s64"; + } else if (pto::isPTOFloat8Type(resultType) || pto::isPTOHiFloat8Type(resultType)) { + elem = "s8"; + } else if (pto::isPTOPackedLdgStgVectorType(resultType)) { + unsigned totalBits = pto::getPTOPackedLdgStgTotalBits(resultType); + if (totalBits == 16) { + elem = "s16"; + } else if (totalBits == 32) { + elem = "s32"; + } else if (totalBits == 64) { + elem = "s64"; + } + } + if (elem.empty()) { + return failure(); + } + StringRef l1cacheName = l1cache == pto::L1Cache::Cache ? "cache" : "uncache"; + return StringAttr::get(context, "llvm.hivm.ldg." + l1cacheName.str() + "." + elem).getValue(); +} + +FailureOr buildL1CacheStoreCallee(MLIRContext *context, Type valueType, pto::L1Cache l1cache) { + std::string elem; + if (auto intType = dyn_cast(valueType)) { + if (intType.getWidth() == 8) { + elem = "b8"; + } else if (intType.getWidth() == 16) { + elem = "b16"; + } else if (intType.getWidth() == 32) { + elem = "b32"; + } else if (intType.getWidth() == 64) { + elem = "b64"; + } + } else if (valueType.isF16() || valueType.isBF16()) { + elem = "b16"; + } else if (valueType.isF32()) { + elem = "b32"; + } else if (valueType.isF64()) { + elem = "b64"; + } else if (pto::isPTOFloat8Type(valueType) || pto::isPTOHiFloat8Type(valueType)) { + elem = "b8"; + } else if (pto::isPTOPackedLdgStgVectorType(valueType)) { + unsigned totalBits = pto::getPTOPackedLdgStgTotalBits(valueType); + if (totalBits == 16) { + elem = "b16"; + } else if (totalBits == 32) { + elem = "b32"; + } else if (totalBits == 64) { + elem = "b64"; + } + } + if (elem.empty()) { + return failure(); + } + StringRef l1cacheName = l1cache == pto::L1Cache::Cache ? "cache" : "uncache"; + return StringAttr::get(context, "llvm.hivm.stg." + l1cacheName.str() + "." + elem).getValue(); +} + +FailureOr buildMulhiCallee(MLIRContext *context, Type resultType, pto::Signedness signedness) { + if (resultType.isInteger(32)) { + return StringAttr::get(context, + signedness == pto::Signedness::Unsigned ? "llvm.hivm.mulhi.ui" : "llvm.hivm.mulhi.i") + .getValue(); + } + if (resultType.isInteger(64) && signedness == pto::Signedness::Unsigned) { + return StringAttr::get(context, "llvm.hivm.mul64hi.ui").getValue(); + } + return failure(); +} + +FailureOr buildMulI32ToI64Callee(MLIRContext *context, pto::Signedness signedness) { + return StringAttr::get(context, signedness == pto::Signedness::Unsigned ? "llvm.hivm.mul.i32toi64.ui" + : "llvm.hivm.mul.i32toi64.i") + .getValue(); +} + +std::string getScalarFloatBuiltinFragment(Type type) { + if (type.isF32()) { + return "f32"; + } + if (type.isF16()) { + return "f16"; + } + if (type.isBF16()) { + return "bf16"; + } + return {}; +} + +std::string getLLVMFloatBuiltinFragment(Type type) { + std::string scalar = getScalarFloatBuiltinFragment(type); + if (!scalar.empty()) { + return scalar; + } + + auto vecType = dyn_cast(type); + if (!vecType || vecType.getRank() != 1 || vecType.getDimSize(0) != 2) { + return {}; + } + Type elementType = vecType.getElementType(); + if (elementType.isF16()) { + return "v2f16"; + } + if (elementType.isBF16()) { + return "v2bf16"; + } + return {}; +} + +std::string getHIVMFloatBuiltinFragment(Type type) { + std::string scalar = getScalarFloatBuiltinFragment(type); + if (!scalar.empty()) { + return scalar; + } + + auto vecType = dyn_cast(type); + if (!vecType || vecType.getRank() != 1 || vecType.getDimSize(0) != 2) { + return {}; + } + Type elementType = vecType.getElementType(); + if (elementType.isF16()) { + return "f16x2"; + } + if (elementType.isBF16()) { + return "bf16x2"; + } + return {}; +} + +FailureOr buildSqrtCallee(MLIRContext *context, Type valueType) { + std::string elem = getLLVMFloatBuiltinFragment(valueType); + if (elem != "f32" && elem != "f16" && elem != "v2f16") { + return failure(); + } + return StringAttr::get(context, "llvm.sqrt." + elem).getValue(); +} + +std::string getScalarHIVMFloatShortFragment(Type type) { + if (type.isF32()) { + return "f"; + } + if (type.isF16()) { + return "h"; + } + if (type.isBF16()) { + return "y"; + } + return {}; +} + +FailureOr buildFmaCallee(MLIRContext *context, Type valueType) { + std::string elem = getHIVMFloatBuiltinFragment(valueType); + if (elem.empty()) { + return failure(); + } + return StringAttr::get(context, "llvm.hivm.ffma." + elem + ".rrr").getValue(); +} + +std::string getConvertScalarFragment(Type type, Attribute signednessAttr) { + if (auto vecType = dyn_cast(type)) { + if (vecType.getRank() != 1 || vecType.getDimSize(0) != 2) { + return {}; + } + Type elementType = vecType.getElementType(); + if (std::string elem = getLowPrecisionElementFragment(elementType); + !elem.empty() && !pto::isPTOFloat4PackedType(elementType)) { + return elem + "x2"; + } + if (elementType.isF32()) { + return "f32x2"; + } + if (elementType.isF16()) { + return "f16x2"; + } + if (elementType.isBF16()) { + return "bf16x2"; + } + return {}; + } + if (type.isF32()) { + return "fp32"; + } + if (type.isF16()) { + return "fp16"; + } + if (type.isBF16()) { + return "bf16"; + } + if (std::string elem = getLowPrecisionElementFragment(type); !elem.empty()) { + return elem; + } + auto intType = dyn_cast(type); + if (!intType || (intType.getWidth() != 32 && intType.getWidth() != 64) || !signednessAttr) { + return {}; + } + auto signedness = cast(signednessAttr).getValue(); + return std::string(signedness == pto::Signedness::Unsigned ? "u" : "s") + std::to_string(intType.getWidth()); +} + +FailureOr buildConvertCallee(MLIRContext *context, Type srcType, Type dstType, Attribute signednessAttr) { + std::string src = getConvertScalarFragment(srcType, signednessAttr); + std::string dst = getConvertScalarFragment(dstType, signednessAttr); + if (src.empty() || dst.empty()) { + return failure(); + } + return StringAttr::get(context, "llvm.hivm." + src + ".to." + dst).getValue(); +} + +FailureOr buildVldsPostCallee(MLIRContext *context, Type resultType) { + std::string vec = getMemoryElementTypeFragment(getElementTypeFromVectorLike(resultType)); + auto lanes = getElementCountFromVectorLike(resultType); + if (vec.empty() || !lanes) { + return failure(); + } + return StringAttr::get(context, "llvm.hivm.vldsx1.post.v" + std::to_string(*lanes) + vec).getValue(); +} + +FailureOr buildVstsPostCallee(MLIRContext *context, Type valueType) { + std::string vec = getMemoryElementTypeFragment(getElementTypeFromVectorLike(valueType)); + auto lanes = getElementCountFromVectorLike(valueType); + if (vec.empty() || !lanes) { + return failure(); + } + return StringAttr::get(context, "llvm.hivm.vstsx1.post.v" + std::to_string(*lanes) + vec).getValue(); +} + +StringRef buildVldasCallee(MLIRContext *context) { return StringAttr::get(context, "llvm.hivm.vldas").getValue(); } + +FailureOr buildVldusCallee(MLIRContext *context, Type resultType) { + std::string vec = getMemoryElementTypeFragment(getElementTypeFromVectorLike(resultType)); + auto lanes = getElementCountFromVectorLike(resultType); + if (vec.empty() || !lanes) { + return failure(); + } + return StringAttr::get(context, "llvm.hivm.vldus.v" + std::to_string(*lanes) + vec).getValue(); +} + +FailureOr buildVldusPostCallee(MLIRContext *context, Type resultType) { + std::string vec = getMemoryElementTypeFragment(getElementTypeFromVectorLike(resultType)); + auto lanes = getElementCountFromVectorLike(resultType); + if (vec.empty() || !lanes) { + return failure(); + } + return StringAttr::get(context, "llvm.hivm.vldus.post.v" + std::to_string(*lanes) + vec).getValue(); +} + +FailureOr buildVcmpCallee(MLIRContext *context, Type inputType, StringRef cmpMode, bool isScalarCompare) { + std::string vec = getCANN900VectorTypeFragment(inputType); + std::string signedness = getCANN900SignednessFragment(getElementTypeFromVectorLike(inputType)); + if (vec.empty() || signedness.empty()) { + return failure(); + } + StringRef stem = isScalarCompare ? "vcmps" : "vcmp"; + return StringAttr::get(context, "llvm.hivm." + stem.str() + "." + cmpMode.str() + "." + signedness + ".z." + vec) + .getValue(); +} + +} // namespace mlir::pto::detail diff --git a/lib/PTO/Transforms/VPTOCANN900LLVMEmitterCalleeMemory.cpp b/lib/PTO/Transforms/VPTOCANN900LLVMEmitterCalleeMemory.cpp new file mode 100644 index 0000000000..06aefa2a58 --- /dev/null +++ b/lib/PTO/Transforms/VPTOCANN900LLVMEmitterCalleeMemory.cpp @@ -0,0 +1,736 @@ +// Copyright (c) 2026 Huawei Technologies Co., Ltd. +// This program is free software, you can redistribute it and/or modify it under the terms and conditions of +// CANN Open Software License Agreement Version 2.0 (the "License"). +// Please refer to the License for details. You may not use this file except in compliance with the License. +// THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED, +// INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE. +// See LICENSE in the root of the software repository for the full text of the License. + +#include "VPTOCANN900LLVMEmitterInternal.h" + +namespace mlir::pto::detail { + +FailureOr buildCopyGmToUbCallee(MLIRContext *context, Type sourceType) { + auto ptrType = dyn_cast(sourceType); + if (!ptrType) { + return failure(); + } + Type elementType = ptrType.getElementType(); + if ((isa(elementType) && cast(elementType).getWidth() == 64) || elementType.isF64()) { + return StringAttr::get(context, "llvm.hivm.MOV.OUT.TO.UB.ALIGN.V2.s32.DV").getValue(); + } + std::string elem = getCopyElementFragment(elementType); + if (elem.empty()) { + return failure(); + } + return StringAttr::get(context, "llvm.hivm.MOV.OUT.TO.UB.ALIGN.V2." + elem + ".DV").getValue(); +} + +StringRef buildCopyUbToGmCallee(MLIRContext *context) { + return StringAttr::get(context, "llvm.hivm.MOV.UB.TO.OUT.ALIGN.V2.DV").getValue(); +} + +StringRef buildCopyUbToUbCallee(MLIRContext *context) { + return StringAttr::get(context, "llvm.hivm.MOV.UB.TO.UB.v310").getValue(); +} + +StringRef buildCopyCbufToUbCallee(MLIRContext *context) { + return StringAttr::get(context, "llvm.hivm.MOV.L1.TO.UB.v310").getValue(); +} + +StringRef buildCopyUbToCbufCallee(MLIRContext *context) { + return StringAttr::get(context, "llvm.hivm.MOV.UB.TO.L1.v310").getValue(); +} + +FailureOr buildOrdinaryMadCallee(MLIRContext *context, pto::MadRawOpInterface op) { + auto lhsType = dyn_cast(op.getLhs().getType()); + auto rhsType = dyn_cast(op.getRhs().getType()); + auto dstType = dyn_cast(op.getDst().getType()); + if (!lhsType || !rhsType || !dstType) { + return failure(); + } + + return buildMadTypedCalleeName(context, lhsType.getElementType(), rhsType.getElementType(), dstType.getElementType()); +} + +FailureOr buildMxMadCallee(MLIRContext *context, pto::MadRawOpInterface op) { + auto lhsType = dyn_cast(op.getLhs().getType()); + auto rhsType = dyn_cast(op.getRhs().getType()); + if (!lhsType || !rhsType) { + return failure(); + } + if (isMxElementType(lhsType.getElementType()) && isMxElementType(rhsType.getElementType())) { + return buildMadMxCalleeName(context, lhsType.getElementType(), rhsType.getElementType()); + } + return failure(); +} + +FailureOr buildCopyGmToCbufCallee(MLIRContext *context, Type sourceType) { + auto ptrType = dyn_cast(sourceType); + if (!ptrType) { + return failure(); + } + std::string elem = getCopyElementFragment(ptrType.getElementType()); + if (elem.empty()) { + return failure(); + } + return StringAttr::get(context, "llvm.hivm.MOV.OUT.TO.L1.ALIGN.V2." + elem + ".DV").getValue(); +} + +FailureOr buildCopyGmToCbufMultiNd2NzCallee(MLIRContext *context, Type sourceType) { + auto ptrType = dyn_cast(sourceType); + if (!ptrType) { + return failure(); + } + std::string elem = getNd2NzCopyElementFragment(ptrType.getElementType()); + if (elem.empty()) { + return failure(); + } + return StringAttr::get(context, "llvm.hivm.MOV.OUT.TO.L1.MULTI.ND2NZ." + elem + ".V310").getValue(); +} + +std::string getDn2NzCopyElementFragment(Type type) { + auto ptrType = dyn_cast(type); + if (!ptrType) { + return {}; + } + + Type elementType = ptrType.getElementType(); + std::string typeText; + llvm::raw_string_ostream os(typeText); + elementType.print(os); + os.flush(); + std::string lower = StringRef(typeText).lower(); + if (StringRef(lower).contains("e4m3") || StringRef(lower).contains("e5m2") || StringRef(lower).contains("e8m0") || + StringRef(lower).contains("hif8")) { + return "u8"; + } + + if (elementType.isF16() || elementType.isBF16()) { + return "u16"; + } + if (elementType.isF32()) { + return "u32"; + } + + if (auto intType = dyn_cast(elementType)) { + switch (intType.getWidth()) { + case 8: + return "u8"; + case 16: + return "u16"; + case 32: + return "u32"; + default: + return {}; + } + } + return {}; +} + +FailureOr buildCopyGmToCbufMultiDn2NzCallee(MLIRContext *context, Type sourceType) { + auto ptrType = dyn_cast(sourceType); + if (!ptrType) { + return failure(); + } + std::string elem = getDn2NzCopyElementFragment(sourceType); + if (elem.empty()) { + return failure(); + } + return StringAttr::get(context, "llvm.hivm.MOV.OUT.TO.L1.MULTI.DN2NZ." + elem).getValue(); +} + +FailureOr buildLoadCbufToCaCallee(MLIRContext *context, Type sourceType) { + auto ptrType = dyn_cast(sourceType); + if (!ptrType) { + return failure(); + } + std::string elem = getL0LoadElementFragment(ptrType.getElementType()); + if (elem.empty()) { + return failure(); + } + return StringAttr::get(context, "llvm.hivm.LOAD.L1.TO.L0A.2Dv2." + elem).getValue(); +} + +FailureOr buildLoadCbufToCbCallee(MLIRContext *context, Type sourceType) { + auto ptrType = dyn_cast(sourceType); + if (!ptrType) { + return failure(); + } + std::string elem = getL0LoadElementFragment(ptrType.getElementType()); + if (elem.empty()) { + return failure(); + } + return StringAttr::get(context, "llvm.hivm.LOAD.L1.TO.L0B.2Dv2." + elem).getValue(); +} + +FailureOr buildLoadCbufToCaS4Callee(MLIRContext *context, Type sourceType) { + auto ptrType = dyn_cast(sourceType); + if (!ptrType) { + return failure(); + } + Type elementType = ptrType.getElementType(); + if (!isa(elementType)) { + return failure(); + } + return StringAttr::get(context, "llvm.hivm.LOAD.L1.TO.L0A.2Dv2.s4").getValue(); +} + +FailureOr buildLoadCbufToCbS4Callee(MLIRContext *context, Type sourceType) { + auto ptrType = dyn_cast(sourceType); + if (!ptrType) { + return failure(); + } + Type elementType = ptrType.getElementType(); + if (!isa(elementType)) { + return failure(); + } + return StringAttr::get(context, "llvm.hivm.LOAD.L1.TO.L0B.2Dv2.s4").getValue(); +} + +StringRef buildLoadCbufToCaMxCallee(MLIRContext *context) { + return StringAttr::get(context, "llvm.hivm.LOAD.L1.TO.L0A.MX.2Dv2.v").getValue(); +} + +[[maybe_unused]] StringRef buildLoadCbufToCbMxCallee(MLIRContext *context) { + return StringAttr::get(context, "llvm.hivm.LOAD.L1.TO.L0B.MX.2Dv2.v").getValue(); +} + +StringRef buildCopyMatrixCcToGmCallee(MLIRContext *context) { + return StringAttr::get(context, "llvm.hivm.FIX.L0C.TO.OUT.f32.EXT").getValue(); +} + +StringRef buildCopyMatrixCcToCbufCallee(MLIRContext *context) { + return StringAttr::get(context, "llvm.hivm.FIX.L0C.TO.L1.f32.EXT").getValue(); +} + +FailureOr buildCopyMatrixCcToUbCallee(MLIRContext *context, Type destinationType) { + auto ptrType = dyn_cast(destinationType); + if (!ptrType) { + return failure(); + } + Type dstElem = ptrType.getElementType(); + if (dstElem.isF16()) { + return StringAttr::get(context, "llvm.hivm.FIX.L0C.TO.UB.f322f16.EXT").getValue(); + } + if (dstElem.isF32()) { + return StringAttr::get(context, "llvm.hivm.FIX.L0C.TO.UB.f32.EXT").getValue(); + } + return failure(); +} + +FailureOr buildCopyCbufToBtCallee(pto::CopyCbufToBtOp op) { + auto ptrType = dyn_cast(op.getSource().getType()); + if (!ptrType) { + return failure(); + } + Type srcElem = ptrType.getElementType(); + if (srcElem.isF16()) { + return StringAttr::get(op.getContext(), "llvm.hivm.MOV.L1.TO.BT.f16").getValue(); + } + if (srcElem.isBF16()) { + return StringAttr::get(op.getContext(), "llvm.hivm.MOV.L1.TO.BT.bf16").getValue(); + } + if (srcElem.isF32()) { + return StringAttr::get(op.getContext(), "llvm.hivm.MOV.L1.TO.BT.f32").getValue(); + } + if (auto intType = dyn_cast(srcElem); intType && intType.getWidth() == 32) { + return StringAttr::get(op.getContext(), "llvm.hivm.MOV.L1.TO.BT.s32").getValue(); + } + return failure(); +} + +StringRef buildCopyCbufToFbufCallee(MLIRContext *context) { + return StringAttr::get(context, "llvm.hivm.MOV.L1.TO.FB.v220").getValue(); +} + +StringRef buildPstiCallee(MLIRContext *context, bool post) { + return StringAttr::get(context, post ? "llvm.hivm.psti.post.b8" : "llvm.hivm.psti.b8").getValue(); +} + +StringRef buildPstsCallee(MLIRContext *context, bool post) { + return StringAttr::get(context, post ? "llvm.hivm.psts.post.b8" : "llvm.hivm.psts.b8").getValue(); +} + +StringRef buildPldiCallee(MLIRContext *context, bool post) { + return StringAttr::get(context, post ? "llvm.hivm.pldi.post.b8" : "llvm.hivm.pldi.b8").getValue(); +} + +StringRef buildPldsCallee(MLIRContext *context, bool post) { + return StringAttr::get(context, post ? "llvm.hivm.plds.post.b8" : "llvm.hivm.plds.b8").getValue(); +} + +StringRef buildPnotCallee(MLIRContext *context) { return StringAttr::get(context, "llvm.hivm.pnot.z").getValue(); } + +StringRef buildPselCallee(MLIRContext *context) { return StringAttr::get(context, "llvm.hivm.psel").getValue(); } + +StringRef buildPandCallee(MLIRContext *context) { return StringAttr::get(context, "llvm.hivm.pand.z").getValue(); } + +StringRef buildPorCallee(MLIRContext *context) { return StringAttr::get(context, "llvm.hivm.por.z").getValue(); } + +StringRef buildPxorCallee(MLIRContext *context) { return StringAttr::get(context, "llvm.hivm.pxor.z").getValue(); } + +StringRef buildPpackCallee(MLIRContext *context) { return StringAttr::get(context, "llvm.hivm.ppack.z").getValue(); } + +StringRef buildPunpackCallee(MLIRContext *context) { return StringAttr::get(context, "llvm.hivm.punpack").getValue(); } + +FailureOr buildInterleaveCallee(MLIRContext *context, Type resultType, StringRef stem) { + // bf16x2 has no dedicated vintlv/vdintlv intrinsic. It is a 32-bit packed + // pair lowered to i32 at the LLVM ABI, and (de)interleave is a bit-level + // lane shuffle, so the intrinsic serves the type. + if (pto::isPTOBF16x2Type(getElementTypeFromVectorLike(resultType))) { + auto lanes = getElementCountFromVectorLike(resultType); + if (lanes) { + return StringAttr::get(context, "llvm.hivm." + stem.str() + ".v" + std::to_string(*lanes) + "i32").getValue(); + } + } + std::string vec = getCANN900VectorTypeFragment(resultType); + if (vec.empty()) { + return failure(); + } + return StringAttr::get(context, "llvm.hivm." + stem.str() + "." + vec).getValue(); +} + +FailureOr buildUnpackCallee(MLIRContext *context, Type inputType, Type resultType, StringRef stem) { + (void)inputType; + std::string vec = getCANN900VectorTypeFragment(resultType); + if (vec.empty()) { + return failure(); + } + return StringAttr::get(context, "llvm.hivm." + stem.str() + "." + vec).getValue(); +} + +FailureOr buildVpackCallee(MLIRContext *context, Type inputType, Type resultType) { + (void)resultType; + std::string vec = getCANN900VectorTypeFragment(inputType); + if (vec.empty()) { + return failure(); + } + + return StringAttr::get(context, "llvm.hivm.vpack.x." + vec).getValue(); +} + +FailureOr buildVsqzCallee(MLIRContext *context, Type resultType) { + return buildCANN900ModeTypedCallee(context, resultType, "vsqz", "x"); +} + +FailureOr buildVusqzCallee(MLIRContext *context, Type resultType) { + return buildCANN900ModeTypedCallee(context, resultType, "vusqz", "m"); +} + +FailureOr buildVmulaCallee(MLIRContext *context, Type resultType) { + return buildCANN900SignedModeTypedCallee(context, resultType, "vmula", "m"); +} + +FailureOr buildVmullCallee(MLIRContext *context, Type resultType) { + return buildLaneTypedCallee(context, resultType, "vmull", ""); +} + +FailureOr buildVldsCallee(MLIRContext *context, Type resultType) { + std::string vec = getMemoryElementTypeFragment(getElementTypeFromVectorLike(resultType)); + auto lanes = getElementCountFromVectorLike(resultType); + if (vec.empty() || !lanes) { + return failure(); + } + return StringAttr::get(context, "llvm.hivm.vldsx1.v" + std::to_string(*lanes) + vec).getValue(); +} + +FailureOr buildVldsx2Callee(MLIRContext *context, Type resultType, bool post) { + std::string vec = getMemoryElementTypeFragment(getElementTypeFromVectorLike(resultType)); + auto lanes = getElementCountFromVectorLike(resultType); + if (vec.empty() || !lanes) { + return failure(); + } + return StringAttr::get(context, + "llvm.hivm.vldsx2" + std::string(post ? ".post" : "") + ".v" + std::to_string(*lanes) + vec) + .getValue(); +} + +FailureOr buildBlockStridedMemoryCallee(MLIRContext *context, Type vectorType, StringRef stem, bool post) { + Type elementType = getElementTypeFromVectorLike(vectorType); + auto lanes = getElementCountFromVectorLike(vectorType); + if (!elementType || !lanes) { + return failure(); + } + + std::string element; + if (auto intType = dyn_cast(elementType)) { + element = "i" + std::to_string(intType.getWidth()); + } else if (isLowpPayloadElementType(elementType)) { + element = "i8"; + } else { + element = getMemoryElementTypeFragment(elementType); + } + if (element.empty()) { + return failure(); + } + + return StringAttr::get(context, "llvm.hivm." + stem.str() + std::string(post ? ".post" : "") + ".v" + + std::to_string(*lanes) + element) + .getValue(); +} + +FailureOr buildVsldbCallee(MLIRContext *context, Type resultType, bool post) { + return buildBlockStridedMemoryCallee(context, resultType, "vsldb", post); +} + +FailureOr buildVstsCallee(MLIRContext *context, Type valueType) { + std::string vec = getMemoryElementTypeFragment(getElementTypeFromVectorLike(valueType)); + auto lanes = getElementCountFromVectorLike(valueType); + if (vec.empty() || !lanes) { + return failure(); + } + return StringAttr::get(context, "llvm.hivm.vstsx1.v" + std::to_string(*lanes) + vec).getValue(); +} + +FailureOr buildVstsx2Callee(MLIRContext *context, Type valueType) { + Type elementType = getElementTypeFromVectorLike(valueType); + auto lanes = getElementCountFromVectorLike(valueType); + if (!elementType || !lanes) { + return failure(); + } + + std::string element = getMemoryElementTypeFragment(elementType); + if (element.empty()) { + return failure(); + } + + return StringAttr::get(context, "llvm.hivm.vstsx2.v" + std::to_string(*lanes) + element).getValue(); +} + +FailureOr buildVsstbCallee(MLIRContext *context, Type valueType, bool post) { + return buildBlockStridedMemoryCallee(context, valueType, "vsstb", post); +} + +Type getVgather2SourceElementType(Type sourceType) { + if (auto ptrType = dyn_cast(sourceType)) { + return ptrType.getElementType(); + } + if (auto memrefType = dyn_cast(sourceType)) { + return memrefType.getElementType(); + } + return {}; +} + +FailureOr buildVgather2Callee(MLIRContext *context, Type sourceType, Type resultType) { + Type sourceElemType = getVgather2SourceElementType(sourceType); + Type resultElemType = getElementTypeFromVectorLike(resultType); + auto lanes = getElementCountFromVectorLike(resultType); + if (!sourceElemType || !resultElemType || !lanes) { + return failure(); + } + + std::string vec; + int64_t intrinsicLanes = *lanes; + if (pto::getPTOStorageElemBitWidth(sourceElemType) == 8) { + vec = getElementTypeFragment(sourceElemType); + intrinsicLanes *= 2; + } else { + vec = getElementTypeFragment(resultElemType); + } + if (vec.empty()) { + return failure(); + } + + return StringAttr::get(context, "llvm.hivm.vgather2.v300.v" + std::to_string(intrinsicLanes) + vec).getValue(); +} + +std::optional getFixedVectorBitWidth(Type type) { + auto vectorType = dyn_cast(type); + if (!vectorType || vectorType.getRank() != 1 || vectorType.isScalable()) { + return std::nullopt; + } + int64_t lanes = vectorType.getDimSize(0); + if (lanes <= 0) { + return std::nullopt; + } + auto elementType = dyn_cast(vectorType.getElementType()); + if (!elementType) { + return std::nullopt; + } + return static_cast(lanes) * elementType.getWidth(); +} + +FailureOr getVgather2OffsetsCarrierType(PatternRewriter &rewriter, Type sourceType, Type resultType, + Type offsetsType) { + Type sourceElemType = getVgather2SourceElementType(sourceType); + Type elementType = getElementTypeFromVectorLike(resultType); + auto lanes = getElementCountFromVectorLike(resultType); + if (!sourceElemType || !elementType || !lanes || *lanes <= 0) { + return failure(); + } + + Type carrierType = offsetsType; + if (pto::getPTOStorageElemBitWidth(elementType) == 16) { + if (*lanes % 2 != 0) { + return failure(); + } + carrierType = VectorType::get({*lanes / 2}, rewriter.getI32Type()); + } + + std::optional offsetsBits = getFixedVectorBitWidth(offsetsType); + std::optional carrierBits = getFixedVectorBitWidth(carrierType); + if (!offsetsBits || !carrierBits || *offsetsBits != *carrierBits) { + return failure(); + } + return carrierType; +} + +FailureOr buildVgather2BcCallee(MLIRContext *context, Type resultType) { + return buildLaneTypedCallee(context, resultType, "vgather2.bc", ""); +} + +FailureOr buildVgatherbCallee(MLIRContext *context, Type resultType) { + return buildLaneTypedCallee(context, resultType, "vgatherb.v310", ""); +} + +FailureOr buildVscatterCallee(MLIRContext *context, Type valueType) { + return buildLaneTypedCallee(context, valueType, "vscatter", ".v300"); +} + +FailureOr getVscatterOffsetsCarrierType(Type offsetsType) { return offsetsType; } + +FailureOr buildVaxpyCallee(MLIRContext *context, Type resultType) { + return buildCANN900ModeTypedCallee(context, resultType, "vaxpy", "m"); +} + +FailureOr buildVmulscvtCallee(MLIRContext *context, Type inputType, Type resultType) { + auto inputElemType = getElementTypeFromVectorLike(inputType); + auto resultElemType = getElementTypeFromVectorLike(resultType); + auto inputLanes = getElementCountFromVectorLike(inputType); + auto resultLanes = getElementCountFromVectorLike(resultType); + if (!inputElemType || !resultElemType || !inputLanes || !resultLanes) { + return failure(); + } + if (!inputElemType.isF32() || !resultElemType.isF16() || *inputLanes != 64 || *resultLanes != 128) { + return failure(); + } + return StringAttr::get(context, "llvm.hivm.vmulscvt.v128f16").getValue(); +} + +FailureOr buildVciCallee(MLIRContext *context, Type resultType) { + std::string vec = getCANN900VectorTypeFragment(resultType); + if (vec.empty()) { + return failure(); + } + return StringAttr::get(context, "llvm.hivm.vci." + vec).getValue(); +} + +FailureOr buildVtrcCallee(MLIRContext *context, Type resultType) { + std::string vec = getElementTypeFragment(getElementTypeFromVectorLike(resultType)); + auto lanes = getElementCountFromVectorLike(resultType); + if (vec.empty() || !lanes) { + return failure(); + } + return StringAttr::get(context, "llvm.hivm.vtrc." + vec + ".x").getValue(); +} + +FailureOr buildVexpdifCallee(MLIRContext *context, Type inputType, Type resultType) { + Type inputElem = getElementTypeFromVectorLike(inputType); + Type resultElem = getElementTypeFromVectorLike(resultType); + auto srcLanes = getElementCountFromVectorLike(inputType); + if (!srcLanes) { + return failure(); + } + if (inputElem.isF16() && resultElem.isF32() && *srcLanes == 128) { + return StringAttr::get(context, "llvm.hivm.vexpdif.interleave.v128f16").getValue(); + } + if (inputElem.isF32() && resultElem.isF32() && *srcLanes == 64) { + return StringAttr::get(context, "llvm.hivm.vexpdif.v64f32").getValue(); + } + return failure(); +} + +FailureOr buildVbitsortCallee(MLIRContext *context, pto::VbitsortOp op) { + Type sourceElemType = cast(op.getSource().getType()).getElementType(); + if (sourceElemType.isF16()) { + return StringAttr::get(context, "llvm.hivm.VBS32.V300.f16").getValue(); + } + if (sourceElemType.isF32()) { + return StringAttr::get(context, "llvm.hivm.VBS32.V300.f32").getValue(); + } + return failure(); +} + +FailureOr buildVmrgsort4Callee(MLIRContext *context, pto::Vmrgsort4Op op) { + Type elemType = cast(op.getDestination().getType()).getElementType(); + if (elemType.isF16()) { + return StringAttr::get(context, "llvm.hivm.VMRGSORT.f16.V300").getValue(); + } + if (elemType.isF32()) { + return StringAttr::get(context, "llvm.hivm.VMRGSORT.f32.V300").getValue(); + } + return failure(); +} + +FailureOr packVmrgsort4SourceAddr(Operation *anchor, Value source0, Value source1, Value source2, Value source3, + Type elemType) { + OpBuilder builder(anchor); + builder.setInsertionPoint(anchor); + Location loc = anchor->getLoc(); + unsigned addrShift = 0; + if (elemType.isF16()) { + addrShift = 3; + } else if (elemType.isF32()) { + addrShift = 3; + } else { + return failure(); + } + + auto packOne = [&](Value source, uint64_t laneShift) -> FailureOr { + FailureOr ubPtr = reinterpretPointerToAddrSpace(anchor, source, 6); + if (failed(ubPtr)) { + return failure(); + } + Value asInt = builder.create(loc, builder.getI64Type(), *ubPtr); + Value shifted = builder.create(loc, asInt, getI64Constant(builder, loc, addrShift)); + Value masked = builder.create(loc, shifted, getI64Constant(builder, loc, 0xFFFFULL)); + if (laneShift == 0) { + return masked; + } + return builder.create(loc, masked, getI64Constant(builder, loc, laneShift)).getResult(); + }; + + FailureOr low0 = packOne(source0, 0); + FailureOr low1 = packOne(source1, 16); + FailureOr low2 = packOne(source2, 32); + FailureOr low3 = packOne(source3, 48); + if (failed(low0) || failed(low1) || failed(low2) || failed(low3)) { + return failure(); + } + + Value packed01 = builder.create(loc, *low0, *low1); + Value packed23 = builder.create(loc, *low2, *low3); + Value packed = builder.create(loc, packed01, packed23); + Type ubPtrTy = LLVM::LLVMPointerType::get(anchor->getContext(), 6); + return builder.create(loc, ubPtrTy, packed).getResult(); +} + +FailureOr buildVcvtContract(pto::VcvtOp op) { + Type inputElemType = getElementTypeFromVectorLike(op.getInput().getType()); + Type resultElemType = getElementTypeFromVectorLike(op.getResult().getType()); + if (!inputElemType || !resultElemType) { + return failure(); + } + auto contract = lookupVcvtContract(classifyVcvtElemType(inputElemType), classifyVcvtElemType(resultElemType)); + if (!contract) { + return failure(); + } + return *contract; +} + +bool needsV300CtrlModeForVPTOFunc(func::FuncOp funcOp) { + if (!pto::isPTOEntryFunction(funcOp) || funcOp.getBlocks().empty()) { + return false; + } + + bool needsCtrlSetup = false; + funcOp.walk([&](pto::VcvtOp vcvtOp) { + FailureOr contract = buildVcvtContract(vcvtOp); + if (succeeded(contract) && (*contract).requiresSat) { + needsCtrlSetup = true; + return WalkResult::interrupt(); + } + return WalkResult::advance(); + }); + return needsCtrlSetup; +} + +FailureOr encodeMovPadValue(Location loc, Value value, ConversionPatternRewriter &rewriter) { + Type type = value.getType(); + Value payload = value; + unsigned bitWidth = 0; + + if (auto intType = dyn_cast(type)) { + bitWidth = intType.getWidth(); + } else if (auto floatType = dyn_cast(type)) { + bitWidth = floatType.getWidth(); + auto intType = rewriter.getIntegerType(bitWidth); + payload = rewriter.create(loc, intType, value); + } else { + return failure(); + } + + if (bitWidth != 8 && bitWidth != 16 && bitWidth != 32) { + return failure(); + } + + return rewriter.create(loc, rewriter.getI64Type(), payload).getResult(); +} + +StringRef buildMemBarCallee(MemBarKind kind, MLIRContext *context) { + switch (kind) { + case MemBarKind::VV_ALL: + return StringAttr::get(context, "llvm.hivm.mem.bar.vv.all").getValue(); + case MemBarKind::VST_VLD: + return StringAttr::get(context, "llvm.hivm.mem.bar.vst.vld").getValue(); + case MemBarKind::VLD_VST: + return StringAttr::get(context, "llvm.hivm.mem.bar.vld.vst").getValue(); + case MemBarKind::VST_VST: + return StringAttr::get(context, "llvm.hivm.mem.bar.vst.vst").getValue(); + case MemBarKind::VS_ALL: + return StringAttr::get(context, "llvm.hivm.mem.bar.vs.all").getValue(); + case MemBarKind::VST_LD: + return StringAttr::get(context, "llvm.hivm.mem.bar.vst.ld").getValue(); + case MemBarKind::VLD_ST: + return StringAttr::get(context, "llvm.hivm.mem.bar.vld.st").getValue(); + case MemBarKind::VST_ST: + return StringAttr::get(context, "llvm.hivm.mem.bar.vst.st").getValue(); + case MemBarKind::SV_ALL: + return StringAttr::get(context, "llvm.hivm.mem.bar.sv.all").getValue(); + case MemBarKind::ST_VLD: + return StringAttr::get(context, "llvm.hivm.mem.bar.st.vld").getValue(); + case MemBarKind::LD_VST: + return StringAttr::get(context, "llvm.hivm.mem.bar.ld.vst").getValue(); + case MemBarKind::ST_VST: + return StringAttr::get(context, "llvm.hivm.mem.bar.st.vst").getValue(); + case MemBarKind::SS_ALL: + return StringAttr::get(context, "llvm.hivm.mem.bar.ss.all").getValue(); + case MemBarKind::ST_LD: + return StringAttr::get(context, "llvm.hivm.mem.bar.st.ld").getValue(); + case MemBarKind::LD_ST: + return StringAttr::get(context, "llvm.hivm.mem.bar.ld.st").getValue(); + case MemBarKind::ST_ST: + return StringAttr::get(context, "llvm.hivm.mem.bar.st.st").getValue(); + } + llvm_unreachable("unexpected membar kind"); +} + +uint64_t getDsbMemImmediate(DsbMem kind) { return static_cast(kind); } + +uint64_t getDcciCacheLineImmediate(DcciCacheLine kind) { return static_cast(kind); } + +uint64_t getDcciDstImmediate(DcciDst kind) { return static_cast(kind); } + +StringRef buildDcciCallee(unsigned addressSpace, bool hasDst, MLIRContext *context) { + if (addressSpace == static_cast(pto::AddressSpace::GM)) { + return StringAttr::get(context, hasDst ? "llvm.hivm.DCCI.DST" : "llvm.hivm.DCCI").getValue(); + } + if (addressSpace == static_cast(pto::AddressSpace::VEC)) { + return StringAttr::get(context, hasDst ? "llvm.hivm.DCCI.DST.UB" : "llvm.hivm.DCCI.UB").getValue(); + } + llvm_unreachable("unexpected dcci address space"); +} + +StringRef buildBufDynSyncCallee(MLIRContext *context, bool isGetBuf) { + return StringAttr::get(context, isGetBuf ? "llvm.hivm.GET.BUF.mode" : "llvm.hivm.RLS.BUF.mode").getValue(); +} + +LogicalResult materializeDecls(ModuleOp module, ArrayRef plannedDecls, llvm::raw_ostream &diagOS) { + OpBuilder builder(module.getBodyRegion()); + builder.setInsertionPointToStart(&module.getBodyRegion().front()); + for (const PlannedDecl &decl : plannedDecls) { + if (func::FuncOp existing = module.lookupSymbol(decl.name)) { + if (existing.getFunctionType() != decl.type) { + diagOS << "VPTO LLVM emission failed: conflicting declaration for " << decl.name << "\n"; + return failure(); + } + continue; + } + auto func = builder.create(module.getLoc(), decl.name, decl.type); + func.setPrivate(); + } + return success(); +} + +} // namespace mlir::pto::detail diff --git a/lib/PTO/Transforms/VPTOCANN900LLVMEmitterInternal.h b/lib/PTO/Transforms/VPTOCANN900LLVMEmitterInternal.h new file mode 100644 index 0000000000..0ef3173f8c --- /dev/null +++ b/lib/PTO/Transforms/VPTOCANN900LLVMEmitterInternal.h @@ -0,0 +1,377 @@ +// Copyright (c) 2026 Huawei Technologies Co., Ltd. +// This program is free software, you can redistribute it and/or modify it under the terms and conditions of +// CANN Open Software License Agreement Version 2.0 (the "License"). +// Please refer to the License for details. You may not use this file except in compliance with the License. +// THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED, +// INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE. +// See LICENSE in the root of the software repository for the full text of the License. + +#pragma once + +#include "PTO/IR/PTO.h" +#include "PTO/IR/PTOSyncUtils.h" +#include "PTO/IR/PTOTypeUtils.h" +#include "PTO/Transforms/Passes.h" +#include "PTO/Transforms/VPTOLLVMEmitter.h" +#include "PTO/Transforms/VPTOLLVMEmitterHelper.h" + +#include "mlir/Conversion/Passes.h" +#include "mlir/Conversion/ReconcileUnrealizedCasts/ReconcileUnrealizedCasts.h" +#include "mlir/Conversion/SCFToControlFlow/SCFToControlFlow.h" +#include "mlir/Dialect/Arith/IR/Arith.h" +#include "mlir/Dialect/Arith/Transforms/Passes.h" +#include "mlir/Dialect/ControlFlow/IR/ControlFlowOps.h" +#include "mlir/Dialect/Func/IR/FuncOps.h" +#include "mlir/Dialect/Func/Transforms/FuncConversions.h" +#include "mlir/Dialect/LLVMIR/LLVMDialect.h" +#include "mlir/Dialect/SCF/Transforms/Patterns.h" +#include "mlir/IR/Builders.h" +#include "mlir/IR/BuiltinOps.h" +#include "mlir/IR/PatternMatch.h" +#include "mlir/Pass/Pass.h" +#include "mlir/Pass/PassManager.h" +#include "mlir/Target/LLVMIR/Dialect/Builtin/BuiltinToLLVMIRTranslation.h" +#include "mlir/Target/LLVMIR/Dialect/LLVMIR/LLVMToLLVMIRTranslation.h" +#include "mlir/Target/LLVMIR/Export.h" +#include "mlir/Transforms/DialectConversion.h" +#include "llvm/ADT/DenseMap.h" +#include "llvm/ADT/STLExtras.h" +#include "llvm/ADT/StringSet.h" +#include "llvm/Bitcode/BitcodeWriter.h" +#include "llvm/IR/Constants.h" +#include "llvm/IR/Function.h" +#include "llvm/IR/GlobalVariable.h" +#include "llvm/IR/Instructions.h" +#include "llvm/IR/LLVMContext.h" +#include "llvm/Support/raw_ostream.h" +#include "llvm/Transforms/Utils/ModuleUtils.h" + +namespace mlir::pto { + +void materializeVecScopeCarrierLoops(ModuleOp module); +LogicalResult applyQueriedTargetAttrs(ModuleOp module, const VPTOEmissionOptions &options, llvm::raw_ostream &diagOS); +LogicalResult attachAIVectorScopeMetadata(llvm::Module &llvmModule, llvm::raw_ostream &diagOS); +void attachHIVMKernelAnnotations(llvm::Module &llvmModule, ModuleOp sourceModule); + +namespace detail { + +inline constexpr llvm::StringLiteral kVectorSuffix = "_mix_aiv"; +inline constexpr llvm::StringLiteral kCubeSuffix = "_mix_aic"; + +struct PlannedDecl { + std::string name; + FunctionType type; +}; + +struct LoweringState { + SmallVector plannedDecls; +}; + +enum class VcvtElemKind { + Invalid, + F16, + BF16, + F32, + F8E4M3, + F8E5M2, + HiF8, + F4E1M2x2, + F4E2M1x2, + S8, + U8, + S16, + U16, + S32, + U32, + S64, +}; + +struct VcvtContract { + const char *intrinsic; + bool requiresRnd; + bool requiresSat; + bool requiresPart; + unsigned maskBitWidth; + bool satBeforeRnd = false; +}; + +struct MadCalleeContract { + StringRef lhs; + StringRef rhs; + StringRef dst; + StringRef callee; +}; + +struct LowpPayloadABI { + Type llvmElementType; + StringRef intrinsicElementFragment; +}; + +Type convertVPTOType(Type type, Builder &builder); +Value materializeVPTOCast(OpBuilder &builder, Type resultType, ValueRange inputs, Location loc); + +class VPTOTypeConverter final : public TypeConverter { +public: + explicit VPTOTypeConverter(MLIRContext *context) { + addConversion([](Type type) { return type; }); + addConversion([](Type type) -> Type { + Builder builder(type.getContext()); + return convertVPTOType(type, builder); + }); + addSourceMaterialization(materializeVPTOCast); + addTargetMaterialization(materializeVPTOCast); + } +}; + +Type getLowPrecisionLLVMType(Type type, MLIRContext *context); +bool isLLVMExtensionVectorElementType(Type type); +Type getLLVMCompatibleVectorType(ArrayRef shape, Type elementType, ArrayRef scalableDims); +Type normalizePayloadTypeForLLVMLowering(Type type, Builder &builder); +Type normalizeGEPElementTypeForLLVMLowering(Type type, Builder &builder); +Type convertVPTOType(Type type, Builder &builder); +unsigned getNaturalByteAlignment(Type type); +bool hasVPTOConvertibleType(Type type); +bool hasVPTOConvertibleType(TypeRange types); +Value materializeVPTOCast(OpBuilder &builder, Type resultType, ValueRange inputs, Location loc); +LLVM::LLVMStructType getVPTOStructStorageType(pto::StructType structType, Builder &builder); +FailureOr getVPTOStructFieldAddress(ConversionPatternRewriter &rewriter, Location loc, Value root, + pto::StructType rootType, ArrayRef path); +Value getI64Constant(OpBuilder &builder, Location loc, uint64_t value); +Value getI32Constant(OpBuilder &builder, Location loc, uint64_t value); +Value getI1Constant(OpBuilder &builder, Location loc, bool value); +bool isMxElementType(Type ty); +std::string getMadMxElementFragment(Type type); +FailureOr buildMadMxCalleeName(MLIRContext *context, Type lhsElem, Type rhsElem); +bool isSignedOrSignlessInteger(IntegerType intType, unsigned width); +std::string getMadRhsFragment(Type type); +bool isMadE4M3ElementType(Type type); +bool isMadE5M2ElementType(Type type); +std::string getMadDstFragment(Type type); +ArrayRef getMadCalleeContracts(); +std::string getMadLhsFragment(Type type); +FailureOr buildMadTypedCalleeName(MLIRContext *context, Type lhsElem, Type rhsElem, Type dstElem); +FailureOr buildLaneTypedCallee(MLIRContext *context, Type resultType, StringRef stem, StringRef suffix); +std::string getCANN900VectorElementFragment(Type type); +std::string getCANN900VectorTypeFragment(Type vectorType); +std::string getCANN900SignednessFragment(Type elemType); +FailureOr buildCANN900ModeTypedCallee(MLIRContext *context, Type vectorType, StringRef stem, StringRef mode); +FailureOr buildCANN900SignedModeTypedCallee(MLIRContext *context, Type vectorType, StringRef stem, + StringRef mode); +FailureOr buildCANN900WideningReductionCallee(MLIRContext *context, Type inputType, Type resultType, + StringRef stem, StringRef mode); +std::string getElementTypeFragment(Type type); +std::string getLowPrecisionElementFragment(Type type); +std::string getMemoryElementTypeFragment(Type type); +bool isLowpPayloadElementType(Type type); +std::optional getLowpPayloadABI(Type elementType, MLIRContext *context); +std::string getDirectLowpVLogicElementFragment(Type type); +FailureOr buildDirectLowpVLogicCallee(MLIRContext *context, Type vectorType, StringRef stem, StringRef mode); +FailureOr buildLowpPayloadVLogicCallee(MLIRContext *context, Type vectorType, StringRef stem, + StringRef mode); +Type getLowpPayloadCarrierType(Type vectorLikeType, MLIRContext *context); +Type getPayloadABIType(Type semanticType, Type convertedType, MLIRContext *context); +Value castToPayloadABI(Location loc, Value value, Type semanticType, ConversionPatternRewriter &rewriter); +Value castFromPayloadABI(Location loc, Value value, Type semanticType, Type convertedType, + ConversionPatternRewriter &rewriter); +std::string getAtomicElementTypeFragment(Type type, Attribute signednessAttr); +std::string getL0LoadElementFragment(Type type); +std::string getShuffleIntrinsicTypeFragment(Type type); +std::string getReduxIntrinsicTypeFragment(Type type, Attribute signednessAttr); +Type getElementTypeFromVectorLike(Type type); +std::optional getElementCountFromVectorLike(Type type); +Value castIntegerLikeTo(Operation *anchor, Value value, Type targetType); +FailureOr reinterpretPointerToAddrSpace(Operation *anchor, Value value, unsigned targetAddressSpace); +FailureOr normalizeVdupScalarOperand(OpBuilder &builder, Location loc, Value input, Type resultType); +Value normalizeByteScalarOperandForCANN900VectorCall(OpBuilder &builder, Location loc, Value input, + Type semanticElementType); +bool isCompatibleScalarForSemanticType(Type semanticType, Type scalarType); +std::string getCopyElementFragment(Type elementType); +std::string getNd2NzCopyElementFragment(Type elementType); +std::optional parsePredicatePatternImmediate(StringRef pattern); +std::optional parseHiLoPartImmediate(StringRef part); +std::optional parseRoundModeImmediate(StringRef roundMode); +std::optional parseSaturationImmediate(StringRef sat); +std::optional parsePartImmediate(StringRef part); +std::optional parseVcvtPartImmediate(StringRef part); +std::optional parsePredicateStoreDistImmediate(StringRef dist); +std::optional parsePredicateLoadDistImmediate(StringRef dist); +std::optional parsePostModeImmediate(StringRef mode); +std::optional parsePipeImmediate(StringRef pipe); +std::optional parseEventImmediate(StringRef event); +std::optional parseSprImmediate(StringRef spr); +std::optional getDistElementWidth(Type type); +VcvtElemKind classifyVcvtElemType(Type type); +std::optional lookupVcvtContract(VcvtElemKind src, VcvtElemKind dst); +uint64_t determineVsqzStoreHint(pto::VsqzOp vsqz); +std::optional parseLoadDistImmediate(StringRef dist, Type elementType); +FailureOr packShiftedFields(Operation *anchor, Value base, ArrayRef> fields); +std::optional parseLoadX2DistImmediate(StringRef dist, Type elementType); +std::optional parseStoreDistImmediate(StringRef dist, Type elementType); +bool isOnePointStoreDist(StringRef dist); +bool isMaskOnlyUsedByOnePointStores(Value mask); +std::optional parseStoreX2DistImmediate(StringRef dist, Type elementType); +Value packBlockRepeatStride(Operation *anchor, Value blockStride, Value repeatStride); +std::optional parseOrderImmediate(StringRef order); +FailureOr packLoopPair(Operation *anchor, Value low, Value high); +FailureOr packLoopSize(Operation *anchor, Value loop2, Value loop1); +FailureOr packCopyGmToUbConfig0(Operation *anchor, ValueRange operands); +FailureOr packCopyGmToUbConfig1(Operation *anchor, ValueRange operands); +FailureOr packCopyGmToUbConfig0(Operation *anchor, Value sid, Value nBurst, Value lenBurst, Value leftPadding, + Value rightPadding, Value dataSelect, Value cacheCtl); +FailureOr packCopyUbToGmConfig0(Operation *anchor, ValueRange operands); +FailureOr packCopyUbToGmConfig1(Operation *anchor, ValueRange operands); +FailureOr packCopyUbToGmConfig0(Operation *anchor, Value sid, Value nBurst, Value lenBurst, Value l2CacheCtl); +FailureOr packCopyUbToUbConfig(Operation *anchor, ValueRange operands); +FailureOr packCopyCbufToUbConfig(Operation *anchor, ValueRange operands); +FailureOr packCopyUbToCbufConfig(Operation *anchor, ValueRange operands); +FailureOr packCopyGmToCbufConfig0(Operation *anchor, Value nBurst, Value lenBurst); +FailureOr packCopyGmToCbufConfig1(Operation *anchor, Value srcStride, Value dstStride); +FailureOr packCopyGmToCbufMultiConfig0(Operation *anchor, Value sid, Value loop1SrcStride, Value l2CacheCtl, + Value nValue); +FailureOr packCopyGmToCbufMultiConfig1(Operation *anchor, Value dValue, Value loop4SrcStride, Value smallC0En); +FailureOr packCopyCbufToBtConfig(Operation *anchor, Value convControl, Value nBurst, Value lenBurst, + Value sourceGap, Value dstGap); +FailureOr packCopyCbufToFbufConfig(Operation *anchor, Value nBurst, Value lenBurst, Value sourceGap, + Value dstGap); +FailureOr packLoadCbufToS4Config0(Operation *anchor, Value mStart, Value kStart, Value mStep, Value kStep); +FailureOr packLoadCbufToS4Config1(Operation *anchor, Value srcStride, Value dstStride); +FailureOr packLoadCbufToCaConfig0(Operation *anchor, Value mStart, Value kStart, Value mStep, Value kStep); +FailureOr packLoadCbufToCaConfig1(Operation *anchor, Value srcStride, Value dstStride); +FailureOr packLoadCbufToCbConfig0(Operation *anchor, Value mStart, Value kStart, Value mStep, Value kStep); +FailureOr packLoadCbufToCbConfig1(Operation *anchor, Value srcStride, Value dstStride); +Value buildMadBiasDestination(Operation *anchor, ConversionPatternRewriter &rewriter, Value dst, Value bias); +FailureOr packVbitsortConfig(Operation *anchor, Value repeatTimes); +FailureOr materializeDynamicPltMask(ConversionPatternRewriter &rewriter, LoweringState &state, Location loc, + Value laneCount, Type vectorElemType); +FailureOr buildCarryBinaryCallee(MLIRContext *context, Type resultType, StringRef stem); +FailureOr buildVselCallee(MLIRContext *context, Type resultType); +FailureOr buildVselrCallee(MLIRContext *context, Type resultType); +FailureOr buildVdupCallee(MLIRContext *context, pto::VdupOp op); +FailureOr buildVbrCallee(MLIRContext *context, Type resultType); +FailureOr buildPstuCallee(MLIRContext *context, pto::PstuOp op); +FailureOr buildVstusCallee(MLIRContext *context, Type valueType); +FailureOr buildVstusPostCallee(MLIRContext *context, Type valueType); +StringRef buildVsturCallee(MLIRContext *context); +StringRef buildInitAlignCallee(MLIRContext *context); +StringRef buildSprclrCallee(MLIRContext *context); +StringRef buildSprstiCallee(MLIRContext *context, bool post); +StringRef buildSprstsCallee(MLIRContext *context, bool post); +StringRef buildStoreVfSimtInfoCallee(MLIRContext *context); +StringRef buildSyncthreadsCallee(MLIRContext *context); +StringRef buildThreadfenceCallee(MLIRContext *context); +StringRef buildThreadfenceBlockCallee(MLIRContext *context); +StringRef buildVstarCallee(MLIRContext *context); +StringRef buildVstasCallee(MLIRContext *context, bool post); +Value buildShuffleControlValue(OpBuilder &builder, Location loc, Value controlValue, int64_t widthValue, + unsigned controlMask); +FailureOr buildAtomicCalleeName(MLIRContext *context, Type ptrType, Type valueType, Attribute signednessAttr, + StringRef opName); +FailureOr buildL1CacheLoadCallee(MLIRContext *context, Type resultType, pto::L1Cache l1cache); +FailureOr buildL1CacheStoreCallee(MLIRContext *context, Type valueType, pto::L1Cache l1cache); +FailureOr buildMulhiCallee(MLIRContext *context, Type resultType, pto::Signedness signedness); +FailureOr buildMulI32ToI64Callee(MLIRContext *context, pto::Signedness signedness); +std::string getScalarFloatBuiltinFragment(Type type); +std::string getLLVMFloatBuiltinFragment(Type type); +std::string getHIVMFloatBuiltinFragment(Type type); +FailureOr buildSqrtCallee(MLIRContext *context, Type valueType); +std::string getScalarHIVMFloatShortFragment(Type type); +FailureOr buildFmaCallee(MLIRContext *context, Type valueType); +std::string getConvertScalarFragment(Type type, Attribute signednessAttr); +FailureOr buildConvertCallee(MLIRContext *context, Type srcType, Type dstType, Attribute signednessAttr); +FailureOr buildVldsPostCallee(MLIRContext *context, Type resultType); +FailureOr buildVstsPostCallee(MLIRContext *context, Type valueType); +StringRef buildVldasCallee(MLIRContext *context); +FailureOr buildVldusCallee(MLIRContext *context, Type resultType); +FailureOr buildVldusPostCallee(MLIRContext *context, Type resultType); +FailureOr buildVcmpCallee(MLIRContext *context, Type inputType, StringRef cmpMode, bool isScalarCompare); +FailureOr buildCopyGmToUbCallee(MLIRContext *context, Type sourceType); +StringRef buildCopyUbToGmCallee(MLIRContext *context); +StringRef buildCopyUbToUbCallee(MLIRContext *context); +StringRef buildCopyCbufToUbCallee(MLIRContext *context); +StringRef buildCopyUbToCbufCallee(MLIRContext *context); +FailureOr buildOrdinaryMadCallee(MLIRContext *context, pto::MadRawOpInterface op); +FailureOr buildMxMadCallee(MLIRContext *context, pto::MadRawOpInterface op); +FailureOr buildCopyGmToCbufCallee(MLIRContext *context, Type sourceType); +FailureOr buildCopyGmToCbufMultiNd2NzCallee(MLIRContext *context, Type sourceType); +std::string getDn2NzCopyElementFragment(Type type); +FailureOr buildCopyGmToCbufMultiDn2NzCallee(MLIRContext *context, Type sourceType); +FailureOr buildLoadCbufToCaCallee(MLIRContext *context, Type sourceType); +FailureOr buildLoadCbufToCbCallee(MLIRContext *context, Type sourceType); +FailureOr buildLoadCbufToCaS4Callee(MLIRContext *context, Type sourceType); +FailureOr buildLoadCbufToCbS4Callee(MLIRContext *context, Type sourceType); +StringRef buildLoadCbufToCaMxCallee(MLIRContext *context); +StringRef buildLoadCbufToCbMxCallee(MLIRContext *context); +StringRef buildCopyMatrixCcToGmCallee(MLIRContext *context); +StringRef buildCopyMatrixCcToCbufCallee(MLIRContext *context); +FailureOr buildCopyMatrixCcToUbCallee(MLIRContext *context, Type destinationType); +FailureOr buildCopyCbufToBtCallee(pto::CopyCbufToBtOp op); +StringRef buildCopyCbufToFbufCallee(MLIRContext *context); +StringRef buildPstiCallee(MLIRContext *context, bool post); +StringRef buildPstsCallee(MLIRContext *context, bool post); +StringRef buildPldiCallee(MLIRContext *context, bool post); +StringRef buildPldsCallee(MLIRContext *context, bool post); +StringRef buildPnotCallee(MLIRContext *context); +StringRef buildPselCallee(MLIRContext *context); +StringRef buildPandCallee(MLIRContext *context); +StringRef buildPorCallee(MLIRContext *context); +StringRef buildPxorCallee(MLIRContext *context); +StringRef buildPpackCallee(MLIRContext *context); +StringRef buildPunpackCallee(MLIRContext *context); +FailureOr buildInterleaveCallee(MLIRContext *context, Type resultType, StringRef stem); +FailureOr buildUnpackCallee(MLIRContext *context, Type inputType, Type resultType, StringRef stem); +FailureOr buildVpackCallee(MLIRContext *context, Type inputType, Type resultType); +FailureOr buildVsqzCallee(MLIRContext *context, Type resultType); +FailureOr buildVusqzCallee(MLIRContext *context, Type resultType); +FailureOr buildVmulaCallee(MLIRContext *context, Type resultType); +FailureOr buildVmullCallee(MLIRContext *context, Type resultType); +FailureOr buildVldsCallee(MLIRContext *context, Type resultType); +FailureOr buildVldsx2Callee(MLIRContext *context, Type resultType, bool post); +FailureOr buildBlockStridedMemoryCallee(MLIRContext *context, Type vectorType, StringRef stem, bool post); +FailureOr buildVsldbCallee(MLIRContext *context, Type resultType, bool post); +FailureOr buildVstsCallee(MLIRContext *context, Type valueType); +FailureOr buildVstsx2Callee(MLIRContext *context, Type valueType); +FailureOr buildVsstbCallee(MLIRContext *context, Type valueType, bool post); +Type getVgather2SourceElementType(Type sourceType); +FailureOr buildVgather2Callee(MLIRContext *context, Type sourceType, Type resultType); +std::optional getFixedVectorBitWidth(Type type); +FailureOr getVgather2OffsetsCarrierType(PatternRewriter &rewriter, Type sourceType, Type resultType, + Type offsetsType); +FailureOr buildVgather2BcCallee(MLIRContext *context, Type resultType); +FailureOr buildVgatherbCallee(MLIRContext *context, Type resultType); +FailureOr buildVscatterCallee(MLIRContext *context, Type valueType); +FailureOr getVscatterOffsetsCarrierType(Type offsetsType); +FailureOr buildVaxpyCallee(MLIRContext *context, Type resultType); +FailureOr buildVmulscvtCallee(MLIRContext *context, Type inputType, Type resultType); +FailureOr buildVciCallee(MLIRContext *context, Type resultType); +FailureOr buildVtrcCallee(MLIRContext *context, Type resultType); +FailureOr buildVexpdifCallee(MLIRContext *context, Type inputType, Type resultType); +FailureOr buildVbitsortCallee(MLIRContext *context, pto::VbitsortOp op); +FailureOr buildVmrgsort4Callee(MLIRContext *context, pto::Vmrgsort4Op op); +FailureOr packVmrgsort4SourceAddr(Operation *anchor, Value source0, Value source1, Value source2, Value source3, + Type elemType); +FailureOr buildVcvtContract(pto::VcvtOp op); +bool needsV300CtrlModeForVPTOFunc(func::FuncOp funcOp); +FailureOr encodeMovPadValue(Location loc, Value value, ConversionPatternRewriter &rewriter); +StringRef buildMemBarCallee(MemBarKind kind, MLIRContext *context); +uint64_t getDsbMemImmediate(DsbMem kind); +uint64_t getDcciCacheLineImmediate(DcciCacheLine kind); +uint64_t getDcciDstImmediate(DcciDst kind); +StringRef buildDcciCallee(unsigned addressSpace, bool hasDst, MLIRContext *context); +StringRef buildBufDynSyncCallee(MLIRContext *context, bool isGetBuf); +LogicalResult materializeDecls(ModuleOp module, ArrayRef plannedDecls, llvm::raw_ostream &diagOS); + +void populateVPTOArithmeticPatterns(VPTOTypeConverter &typeConverter, RewritePatternSet &patterns, + LoweringState &state); +void populateVPTOMemoryPatterns(VPTOTypeConverter &typeConverter, RewritePatternSet &patterns, LoweringState &state); +void populateVPTOVectorMemoryPatterns(VPTOTypeConverter &typeConverter, RewritePatternSet &patterns, + LoweringState &state); +void populateVPTOScalarPatterns(VPTOTypeConverter &typeConverter, RewritePatternSet &patterns, LoweringState &state); +void populateVPTOTypePatterns(VPTOTypeConverter &typeConverter, RewritePatternSet &patterns, ConversionTarget &target, + LoweringState &state); +void populateVPTOStructuralTypePatterns(VPTOTypeConverter &typeConverter, RewritePatternSet &patterns, + ConversionTarget &target); +LogicalResult lowerCANN900Module(ModuleOp module, const VPTOEmissionOptions &options, EmittedLLVMModule &cubeModule, + EmittedLLVMModule &vectorModule, llvm::raw_ostream &diagOS); + +} // namespace detail +} // namespace mlir::pto diff --git a/lib/PTO/Transforms/VPTOCANN900LLVMEmitterMemoryPatterns.cpp b/lib/PTO/Transforms/VPTOCANN900LLVMEmitterMemoryPatterns.cpp new file mode 100644 index 0000000000..69721677b6 --- /dev/null +++ b/lib/PTO/Transforms/VPTOCANN900LLVMEmitterMemoryPatterns.cpp @@ -0,0 +1,1553 @@ +// Copyright (c) 2026 Huawei Technologies Co., Ltd. +// This program is free software, you can redistribute it and/or modify it under the terms and conditions of +// CANN Open Software License Agreement Version 2.0 (the "License"). +// Please refer to the License for details. You may not use this file except in compliance with the License. +// THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED, +// INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE. +// See LICENSE in the root of the software repository for the full text of the License. + +#include "VPTOCANN900LLVMEmitterTemplates.h" + +namespace mlir::pto::detail { + +template class LowerUnpackOpPattern final : public OpConversionPattern { +public: + explicit LowerUnpackOpPattern(TypeConverter &typeConverter, MLIRContext *context, LoweringState &state) + : OpConversionPattern(typeConverter, context), state(state) {} + + LogicalResult matchAndRewrite(UnpackOp op, typename UnpackOp::Adaptor adaptor, + ConversionPatternRewriter &rewriter) const override { + StringRef stem = std::is_same_v ? "vsunpack" : "vzunpack"; + FailureOr calleeName = + buildUnpackCallee(op.getContext(), op.getSrc().getType(), op.getResult().getType(), stem); + if (failed(calleeName)) { + return rewriter.notifyMatchFailure(op, "unsupported unpack VPTO signature"); + } + + Type srcType = this->getTypeConverter()->convertType(op.getSrc().getType()); + Type resultType = this->getTypeConverter()->convertType(op.getResult().getType()); + if (!srcType || !resultType) { + return rewriter.notifyMatchFailure(op, "failed to convert unpack types"); + } + + Value src = adaptor.getSrc(); + if (!src || src.getType() != srcType) { + return rewriter.notifyMatchFailure(op, "unexpected converted unpack source type"); + } + + Value part = castIntegerLikeTo(op, adaptor.getPart(), rewriter.getI32Type()); + if (!part) { + return rewriter.notifyMatchFailure(op, "failed to materialize unpack part"); + } + + auto funcType = rewriter.getFunctionType(TypeRange{srcType, part.getType()}, TypeRange{resultType}); + auto call = rewriter.create(op.getLoc(), *calleeName, TypeRange{resultType}, ValueRange{src, part}); + state.plannedDecls.push_back(PlannedDecl{calleeName->str(), funcType}); + rewriter.replaceOp(op, call.getResults()); + return success(); + } + +private: + LoweringState &state; +}; + +class LowerVpackOpPattern final : public OpConversionPattern { +public: + explicit LowerVpackOpPattern(TypeConverter &typeConverter, MLIRContext *context, LoweringState &state) + : OpConversionPattern(typeConverter, context), state(state) {} + + LogicalResult matchAndRewrite(pto::VpackOp op, pto::VpackOp::Adaptor adaptor, + ConversionPatternRewriter &rewriter) const override { + FailureOr calleeName = + buildVpackCallee(op.getContext(), op.getSrc().getType(), op.getResult().getType()); + if (failed(calleeName)) { + return rewriter.notifyMatchFailure(op, "unsupported vpack VPTO signature"); + } + + Type srcType = this->getTypeConverter()->convertType(op.getSrc().getType()); + Type resultType = this->getTypeConverter()->convertType(op.getResult().getType()); + if (!srcType || !resultType) { + return rewriter.notifyMatchFailure(op, "failed to convert vpack types"); + } + + auto partImm = parseHiLoPartImmediate(op.getPart()); + if (!partImm) { + return rewriter.notifyMatchFailure(op, "unsupported vpack part immediate"); + } + + Value src = adaptor.getSrc(); + if (!src || src.getType() != srcType) { + return rewriter.notifyMatchFailure(op, "unexpected converted vpack source type"); + } + + Value part = getI32Constant(rewriter, op.getLoc(), *partImm); + auto funcType = rewriter.getFunctionType(TypeRange{srcType, part.getType()}, TypeRange{resultType}); + auto call = rewriter.create(op.getLoc(), *calleeName, TypeRange{resultType}, ValueRange{src, part}); + state.plannedDecls.push_back(PlannedDecl{calleeName->str(), funcType}); + rewriter.replaceOp(op, call.getResults()); + return success(); + } + +private: + LoweringState &state; +}; + +template +class LowerPredicateMaskBinaryOpPattern final : public OpConversionPattern { +public: + explicit LowerPredicateMaskBinaryOpPattern(TypeConverter &typeConverter, MLIRContext *context, LoweringState &state) + : OpConversionPattern(typeConverter, context), state(state) {} + + LogicalResult matchAndRewrite(PredicateMaskOp op, typename PredicateMaskOp::Adaptor adaptor, + ConversionPatternRewriter &rewriter) const override { + Type resultType = this->getTypeConverter()->convertType(op.getResult().getType()); + if (!resultType) { + return rewriter.notifyMatchFailure(op, "failed to convert predicate-mask result type"); + } + + Value src0 = adaptor.getSrc0(); + Value src1 = adaptor.getSrc1(); + Value mask = adaptor.getMask(); + if (!src0 || !src1 || !mask || src0.getType() != resultType || src1.getType() != resultType || + mask.getType() != resultType) { + return rewriter.notifyMatchFailure(op, "unexpected converted predicate-mask operand types"); + } + + StringRef calleeName = getPredicateMaskCallee(op.getContext()); + auto call = + rewriter.create(op.getLoc(), calleeName, TypeRange{resultType}, ValueRange{src0, src1, mask}); + state.plannedDecls.push_back(PlannedDecl{calleeName.str(), call.getCalleeType()}); + rewriter.replaceOp(op, call.getResults()); + return success(); + } + +private: + LoweringState &state; +}; + +template class LowerPredicatePairReorderOpPattern final : public OpConversionPattern { +public: + explicit LowerPredicatePairReorderOpPattern(TypeConverter &typeConverter, MLIRContext *context, LoweringState &state) + : OpConversionPattern(typeConverter, context), state(state) {} + + LogicalResult matchAndRewrite(ReorderOp op, typename ReorderOp::Adaptor adaptor, + ConversionPatternRewriter &rewriter) const override { + SmallVector resultTypes; + if (failed(this->getTypeConverter()->convertTypes(op->getResultTypes(), resultTypes))) { + return rewriter.notifyMatchFailure(op, "failed to convert predicate-pair-reorder result types"); + } + if (resultTypes.size() != 2 || resultTypes[0] != resultTypes[1]) { + return rewriter.notifyMatchFailure(op, "unexpected predicate-pair-reorder converted result types"); + } + + Value lhs = adaptor.getLhs(); + Value rhs = adaptor.getRhs(); + if (!lhs || !rhs || lhs.getType() != resultTypes[0] || rhs.getType() != resultTypes[0]) { + return rewriter.notifyMatchFailure(op, "unexpected converted predicate-pair-reorder operand types"); + } + + StringRef calleeName = buildPredicatePairReorderCallee(op.getContext()); + auto call = rewriter.create(op.getLoc(), calleeName, resultTypes, ValueRange{lhs, rhs}); + state.plannedDecls.push_back(PlannedDecl{calleeName.str(), call.getCalleeType()}); + rewriter.replaceOp(op, call.getResults()); + return success(); + } + +private: + LoweringState &state; +}; + +template class LowerCmpOpPattern final : public OpConversionPattern { +public: + explicit LowerCmpOpPattern(TypeConverter &typeConverter, MLIRContext *context, LoweringState &state) + : OpConversionPattern(typeConverter, context), state(state) {} + + LogicalResult matchAndRewrite(CmpOp op, typename CmpOp::Adaptor adaptor, + ConversionPatternRewriter &rewriter) const override { + constexpr bool isScalarCompare = std::is_same_v; + Type inputType = Type(); + if constexpr (isScalarCompare) { + inputType = op.getSrc().getType(); + } else { + inputType = op.getSrc0().getType(); + } + FailureOr calleeName = buildVcmpCallee(op.getContext(), inputType, op.getCmpMode(), isScalarCompare); + if (failed(calleeName)) { + return rewriter.notifyMatchFailure(op, "unsupported compare VPTO signature"); + } + + Type resultType = this->getTypeConverter()->convertType(op.getResult().getType()); + Type maskType = this->getTypeConverter()->convertType(op.getMask().getType()); + if (!resultType || !maskType) { + return rewriter.notifyMatchFailure(op, "failed to convert compare result type"); + } + if (resultType != maskType) { + return rewriter.notifyMatchFailure(op, "unexpected compare mask conversion"); + } + + SmallVector callArgs; + callArgs.append(adaptor.getOperands().begin(), adaptor.getOperands().end()); + if constexpr (isScalarCompare) { + if (callArgs.size() != 3 || !callArgs[0] || !callArgs[1] || !callArgs[2] || callArgs[2].getType() != maskType) { + return rewriter.notifyMatchFailure(op, "unexpected converted scalar-compare operand types"); + } + callArgs[1] = normalizeByteScalarOperandForCANN900VectorCall( + rewriter, op.getLoc(), callArgs[1], cast(op.getSrc().getType()).getElementType()); + } else { + if (callArgs.size() != 3 || !callArgs[0] || !callArgs[1] || !callArgs[2] || + callArgs[0].getType() != callArgs[1].getType() || callArgs[2].getType() != maskType) { + return rewriter.notifyMatchFailure(op, "unexpected converted compare operand types"); + } + } + + auto call = rewriter.create(op.getLoc(), *calleeName, TypeRange{resultType}, callArgs); + state.plannedDecls.push_back(PlannedDecl{calleeName->str(), call.getCalleeType()}); + rewriter.replaceOp(op, call.getResults()); + return success(); + } + +private: + LoweringState &state; +}; + +template class LowerPltOpPattern final : public OpConversionPattern { +public: + explicit LowerPltOpPattern(TypeConverter &typeConverter, MLIRContext *context, LoweringState &state) + : OpConversionPattern(typeConverter, context), state(state) {} + + LogicalResult matchAndRewrite(PltOp op, typename PltOp::Adaptor adaptor, + ConversionPatternRewriter &rewriter) const override { + Value laneCount = castIntegerLikeTo(op, adaptor.getScalar(), rewriter.getI32Type()); + if (!laneCount) { + return rewriter.notifyMatchFailure(op, "failed to materialize plt lane count"); + } + + SmallVector resultTypes; + if (failed(this->getTypeConverter()->convertTypes(op->getResultTypes(), resultTypes))) { + return rewriter.notifyMatchFailure(op, "failed to convert plt result types"); + } + + StringRef calleeName = buildPltCallee(op.getContext()); + auto funcType = rewriter.getFunctionType(TypeRange{rewriter.getI32Type()}, resultTypes); + auto call = rewriter.create(op.getLoc(), calleeName, resultTypes, ValueRange{laneCount}); + state.plannedDecls.push_back(PlannedDecl{calleeName.str(), funcType}); + rewriter.replaceOp(op, call.getResults()); + return success(); + } + +private: + LoweringState &state; +}; + +template class LowerPltmOpPattern final : public OpConversionPattern { +public: + explicit LowerPltmOpPattern(TypeConverter &typeConverter, MLIRContext *context, LoweringState &state) + : OpConversionPattern(typeConverter, context), state(state) {} + + LogicalResult matchAndRewrite(PltmOp op, typename PltmOp::Adaptor adaptor, + ConversionPatternRewriter &rewriter) const override { + SmallVector resultTypes; + if (failed(this->getTypeConverter()->convertTypes(op->getResultTypes(), resultTypes))) { + return rewriter.notifyMatchFailure(op, "failed to convert pltm result type"); + } + + Value loop = adaptor.getLoop(); + Value bound = adaptor.getBound(); + if (!loop || !bound || !loop.getType().isInteger(16) || !bound.getType().isInteger(32)) { + return rewriter.notifyMatchFailure(op, "unexpected converted pltm operand types"); + } + + StringRef calleeName = buildPltmCallee(op.getContext()); + auto funcType = rewriter.getFunctionType(TypeRange{rewriter.getI16Type(), rewriter.getI32Type()}, resultTypes); + auto call = rewriter.create(op.getLoc(), calleeName, resultTypes, ValueRange{loop, bound}); + state.plannedDecls.push_back(PlannedDecl{calleeName.str(), funcType}); + rewriter.replaceOp(op, call.getResults()); + return success(); + } + +private: + LoweringState &state; +}; + +template class LowerPsetOpPattern final : public OpConversionPattern { +public: + explicit LowerPsetOpPattern(TypeConverter &typeConverter, MLIRContext *context, LoweringState &state) + : OpConversionPattern(typeConverter, context), state(state) {} + + LogicalResult matchAndRewrite(PsetOp op, typename PsetOp::Adaptor adaptor, + ConversionPatternRewriter &rewriter) const override { + (void)adaptor; + auto pattern = parsePredicatePatternImmediate(op.getPattern()); + if (!pattern) { + return rewriter.notifyMatchFailure(op, "unsupported pset pattern"); + } + + SmallVector resultTypes; + if (failed(this->getTypeConverter()->convertTypes(op->getResultTypes(), resultTypes))) { + return rewriter.notifyMatchFailure(op, "failed to convert pset result types"); + } + + if (isMaskOnlyUsedByOnePointStores(op.getResult())) { + auto undef = rewriter.create(op.getLoc(), resultTypes.front()); + rewriter.replaceOp(op, undef.getResult()); + return success(); + } + + StringRef calleeName = buildPsetCallee(op.getContext()); + Value patternValue = rewriter.create(op.getLoc(), rewriter.getI32IntegerAttr(*pattern)); + auto funcType = rewriter.getFunctionType(TypeRange{rewriter.getI32Type()}, resultTypes); + auto call = rewriter.create(op.getLoc(), calleeName, resultTypes, ValueRange{patternValue}); + state.plannedDecls.push_back(PlannedDecl{calleeName.str(), funcType}); + rewriter.replaceOp(op, call.getResults()); + return success(); + } + +private: + LoweringState &state; +}; + +template class LowerPgeOpPattern final : public OpConversionPattern { +public: + explicit LowerPgeOpPattern(TypeConverter &typeConverter, MLIRContext *context, LoweringState &state) + : OpConversionPattern(typeConverter, context), state(state) {} + + LogicalResult matchAndRewrite(PgeOp op, typename PgeOp::Adaptor adaptor, + ConversionPatternRewriter &rewriter) const override { + (void)adaptor; + auto pattern = parsePredicatePatternImmediate(op.getPattern()); + if (!pattern) { + return rewriter.notifyMatchFailure(op, "unsupported pge pattern"); + } + + SmallVector resultTypes; + if (failed(this->getTypeConverter()->convertTypes(op->getResultTypes(), resultTypes))) { + return rewriter.notifyMatchFailure(op, "failed to convert pge result types"); + } + + if (isMaskOnlyUsedByOnePointStores(op.getResult())) { + auto undef = rewriter.create(op.getLoc(), resultTypes.front()); + rewriter.replaceOp(op, undef.getResult()); + return success(); + } + + StringRef calleeName = buildPgeCallee(op.getContext()); + Value patternValue = rewriter.create(op.getLoc(), rewriter.getI32IntegerAttr(*pattern)); + Value zero = rewriter.create(op.getLoc(), rewriter.getI32IntegerAttr(0)); + auto funcType = rewriter.getFunctionType(TypeRange{rewriter.getI32Type(), rewriter.getI32Type()}, resultTypes); + auto call = rewriter.create(op.getLoc(), calleeName, resultTypes, ValueRange{patternValue, zero}); + state.plannedDecls.push_back(PlannedDecl{calleeName.str(), funcType}); + rewriter.replaceOp(op, call.getResults()); + return success(); + } + +private: + LoweringState &state; +}; + +static SmallVector getVldsCallResultTypes(Type ptoResultType, ArrayRef resultTypes, bool usePostIntrinsic, + MLIRContext *context) { + SmallVector callResultTypes{getPayloadABIType(ptoResultType, resultTypes[0], context)}; + if (usePostIntrinsic) { + callResultTypes.push_back(resultTypes[1]); + } + return callResultTypes; +} + +static SmallVector getVldsReplacements(pto::VldsOp op, const VPTOLoweredAddressOffset &offset, func::CallOp call, + ArrayRef resultTypes, ConversionPatternRewriter &rewriter) { + Value loaded = castFromPayloadABI(op.getLoc(), call.getResult(0), op.getResult().getType(), resultTypes[0], rewriter); + SmallVector replacements{loaded}; + if (op.getUpdatedBase()) { + replacements.push_back(offset.updatedBase ? offset.updatedBase : call.getResult(1)); + } + return replacements; +} + +static SmallVector getVldsx2Replacements(pto::Vldsx2Op op, const VPTOLoweredAddressOffset &offset, + func::CallOp call, ArrayRef resultTypes, + ConversionPatternRewriter &rewriter) { + Value low = castFromPayloadABI(op.getLoc(), call.getResult(0), op.getLow().getType(), resultTypes[0], rewriter); + Value high = castFromPayloadABI(op.getLoc(), call.getResult(1), op.getHigh().getType(), resultTypes[1], rewriter); + SmallVector replacements{low, high}; + if (op.getUpdatedBase()) { + replacements.push_back(offset.updatedBase ? offset.updatedBase : call.getResult(2)); + } + return replacements; +} + +class LowerVldsOpPattern final : public OpConversionPattern { +public: + explicit LowerVldsOpPattern(TypeConverter &typeConverter, MLIRContext *context, LoweringState &state) + : OpConversionPattern(typeConverter, context), state(state) {} + + LogicalResult matchAndRewrite(pto::VldsOp op, OpAdaptor adaptor, ConversionPatternRewriter &rewriter) const override { + Type ptoResultType = op.getResult().getType(); + Type elementType = getElementTypeFromVectorLike(ptoResultType); + if (!elementType) { + return rewriter.notifyMatchFailure(op, "unsupported vlds element type"); + } + bool usePostIntrinsic = static_cast(op.getUpdatedBase()); + auto loweredOffset = lowerVPTOElementOffsetForIntrinsic(op, adaptor.getSource(), adaptor.getOffset(), elementType, + usePostIntrinsic, rewriter); + auto dist = parseLoadDistImmediate(op.getDist().value_or("NORM"), elementType); + bool invalidAddress = failed(loweredOffset) || !dist; + if (invalidAddress) { + return rewriter.notifyMatchFailure(op, "failed to materialize vlds operands"); + } + + SmallVector resultTypes; + if (failed(this->getTypeConverter()->convertTypes(op->getResultTypes(), resultTypes))) { + return rewriter.notifyMatchFailure(op, "failed to convert vlds result types"); + } + + if (usePostIntrinsic) { + if (resultTypes.size() != 2 || resultTypes[1] != adaptor.getSource().getType()) { + return rewriter.notifyMatchFailure(op, "unsupported vlds post-update results"); + } + } else if (resultTypes.size() != 1) { + return rewriter.notifyMatchFailure(op, "unsupported vlds result count"); + } + SmallVector callResultTypes = + getVldsCallResultTypes(ptoResultType, resultTypes, usePostIntrinsic, rewriter.getContext()); + + FailureOr calleeName = usePostIntrinsic ? buildVldsPostCallee(op.getContext(), ptoResultType) + : buildVldsCallee(op.getContext(), ptoResultType); + if (failed(calleeName)) { + return rewriter.notifyMatchFailure(op, "unsupported vlds signature"); + } + + Value distValue = getI32Constant(rewriter, op.getLoc(), *dist); + Value postValue = getI32Constant(rewriter, op.getLoc(), usePostIntrinsic ? 1 : 0); + SmallVector args{loweredOffset->base, loweredOffset->intrinsicOffset, distValue, postValue}; + auto funcType = + rewriter.getFunctionType(TypeRange{loweredOffset->base.getType(), loweredOffset->intrinsicOffset.getType(), + distValue.getType(), postValue.getType()}, + callResultTypes); + auto call = rewriter.create(op.getLoc(), *calleeName, callResultTypes, args); + state.plannedDecls.push_back(PlannedDecl{calleeName->str(), funcType}); + rewriter.replaceOp(op, getVldsReplacements(op, *loweredOffset, call, resultTypes, rewriter)); + return success(); + } + +private: + LoweringState &state; +}; + +class LowerVldsx2OpPattern final : public OpConversionPattern { +public: + explicit LowerVldsx2OpPattern(TypeConverter &typeConverter, MLIRContext *context, LoweringState &state) + : OpConversionPattern(typeConverter, context), state(state) {} + + LogicalResult matchAndRewrite(pto::Vldsx2Op op, pto::Vldsx2Op::Adaptor adaptor, + ConversionPatternRewriter &rewriter) const override { + Type elementType = getElementTypeFromVectorLike(op.getLow().getType()); + if (!elementType) { + return rewriter.notifyMatchFailure(op, "unsupported vldsx2 element type"); + } + + bool usePostIntrinsic = op.getUpdatedBase() != nullptr; + auto loweredOffset = lowerVPTOElementOffsetForIntrinsic(op, adaptor.getSource(), adaptor.getOffset(), elementType, + usePostIntrinsic, rewriter); + auto dist = parseLoadX2DistImmediate(op.getDist(), elementType); + bool invalidAddress = failed(loweredOffset) || !dist; + if (invalidAddress) { + return rewriter.notifyMatchFailure(op, "failed to materialize vldsx2 operands"); + } + + SmallVector resultTypes; + if (failed(this->getTypeConverter()->convertTypes(op->getResultTypes(), resultTypes)) || + resultTypes.size() != (usePostIntrinsic ? 3U : 2U)) { + return rewriter.notifyMatchFailure(op, "failed to convert vldsx2 result types"); + } + Type lowCallType = getPayloadABIType(op.getLow().getType(), resultTypes[0], rewriter.getContext()); + Type highCallType = getPayloadABIType(op.getHigh().getType(), resultTypes[1], rewriter.getContext()); + SmallVector callResultTypes{lowCallType, highCallType}; + if (usePostIntrinsic) { + callResultTypes.push_back(resultTypes[2]); + } + + FailureOr calleeName = buildVldsx2Callee(op.getContext(), op.getLow().getType(), usePostIntrinsic); + if (failed(calleeName)) { + return rewriter.notifyMatchFailure(op, "unsupported vldsx2 signature"); + } + + Value distValue = getI32Constant(rewriter, op.getLoc(), *dist); + Value postValue = getI32Constant(rewriter, op.getLoc(), usePostIntrinsic ? 1 : 0); + SmallVector args{loweredOffset->base, loweredOffset->intrinsicOffset, distValue, postValue}; + auto funcType = + rewriter.getFunctionType(TypeRange{loweredOffset->base.getType(), loweredOffset->intrinsicOffset.getType(), + distValue.getType(), postValue.getType()}, + callResultTypes); + auto call = rewriter.create(op.getLoc(), *calleeName, callResultTypes, args); + state.plannedDecls.push_back(PlannedDecl{calleeName->str(), funcType}); + rewriter.replaceOp(op, getVldsx2Replacements(op, *loweredOffset, call, resultTypes, rewriter)); + return success(); + } + +private: + LoweringState &state; +}; + +class LowerVsldbOpPattern final : public OpConversionPattern { +public: + explicit LowerVsldbOpPattern(TypeConverter &typeConverter, MLIRContext *context, LoweringState &state) + : OpConversionPattern(typeConverter, context), state(state) {} + + LogicalResult matchAndRewrite(pto::VsldbOp op, pto::VsldbOp::Adaptor adaptor, + ConversionPatternRewriter &rewriter) const override { + auto basePtr = dyn_cast(adaptor.getSource().getType()); + Value packedStride = packBlockRepeatStride(op, adaptor.getBlockStride(), adaptor.getRepeatStride()); + if (!basePtr || !packedStride) { + return rewriter.notifyMatchFailure(op, "failed to materialize vsldb operands"); + } + + bool usePostIntrinsic = op.getUpdatedBase() != nullptr; + SmallVector resultTypes; + if (failed(this->getTypeConverter()->convertTypes(op->getResultTypes(), resultTypes)) || + resultTypes.size() != (usePostIntrinsic ? 2U : 1U)) { + return rewriter.notifyMatchFailure(op, "failed to convert vsldb result type"); + } + + Type callResultType = getPayloadABIType(op.getResult().getType(), resultTypes[0], rewriter.getContext()); + SmallVector callResultTypes{callResultType}; + if (usePostIntrinsic) { + callResultTypes.push_back(resultTypes[1]); + } + + FailureOr calleeName = buildVsldbCallee(op.getContext(), op.getResult().getType(), usePostIntrinsic); + if (failed(calleeName)) { + return rewriter.notifyMatchFailure(op, "unsupported vsldb signature"); + } + Value postValue = getI32Constant(rewriter, op.getLoc(), usePostIntrinsic ? 1 : 0); + SmallVector args{adaptor.getSource(), packedStride, postValue, adaptor.getMask()}; + auto funcType = rewriter.getFunctionType(TypeRange{adaptor.getSource().getType(), packedStride.getType(), + postValue.getType(), adaptor.getMask().getType()}, + callResultTypes); + auto call = rewriter.create(op.getLoc(), *calleeName, callResultTypes, args); + state.plannedDecls.push_back(PlannedDecl{calleeName->str(), funcType}); + Value result = + castFromPayloadABI(op.getLoc(), call.getResult(0), op.getResult().getType(), resultTypes[0], rewriter); + if (usePostIntrinsic) { + rewriter.replaceOp(op, ValueRange{result, call.getResult(1)}); + } else { + rewriter.replaceOp(op, ValueRange{result}); + } + return success(); + } + +private: + LoweringState &state; +}; + +class LowerInitAlignOpPattern final : public OpConversionPattern { +public: + explicit LowerInitAlignOpPattern(TypeConverter &typeConverter, MLIRContext *context, LoweringState &state) + : OpConversionPattern(typeConverter, context), state(state) {} + + LogicalResult matchAndRewrite(pto::InitAlignOp op, pto::InitAlignOp::Adaptor adaptor, + ConversionPatternRewriter &rewriter) const override { + (void)adaptor; + Type resultType = this->getTypeConverter()->convertType(op.getResult().getType()); + if (!resultType) { + return rewriter.notifyMatchFailure(op, "failed to convert init_align result type"); + } + + StringRef calleeName = buildInitAlignCallee(op.getContext()); + auto funcType = rewriter.getFunctionType(TypeRange{}, TypeRange{resultType}); + auto call = rewriter.create(op.getLoc(), calleeName, TypeRange{resultType}); + state.plannedDecls.push_back(PlannedDecl{calleeName.str(), funcType}); + rewriter.replaceOp(op, call.getResults()); + return success(); + } + +private: + LoweringState &state; +}; + +class LowerVldasOpPattern final : public OpConversionPattern { +public: + explicit LowerVldasOpPattern(TypeConverter &typeConverter, MLIRContext *context, LoweringState &state) + : OpConversionPattern(typeConverter, context), state(state) {} + + LogicalResult matchAndRewrite(pto::VldasOp op, pto::VldasOp::Adaptor adaptor, + ConversionPatternRewriter &rewriter) const override { + auto sourceType = dyn_cast(adaptor.getSource().getType()); + Type resultType = this->getTypeConverter()->convertType(op.getResult().getType()); + if (!sourceType || !resultType) { + return rewriter.notifyMatchFailure(op, "expected converted vldas operand/result types"); + } + + StringRef calleeName = buildVldasCallee(op.getContext()); + auto funcType = rewriter.getFunctionType(TypeRange{adaptor.getSource().getType()}, TypeRange{resultType}); + auto call = + rewriter.create(op.getLoc(), calleeName, TypeRange{resultType}, ValueRange{adaptor.getSource()}); + state.plannedDecls.push_back(PlannedDecl{calleeName.str(), funcType}); + rewriter.replaceOp(op, call.getResults()); + return success(); + } + +private: + LoweringState &state; +}; + +struct VldusCallOperands { + SmallVector args; + Value explicitUpdatedBase; +}; + +static FailureOr buildVldusCallOperands(pto::VldusOp op, pto::VldusOp::Adaptor adaptor, + ConversionPatternRewriter &rewriter) { + SmallVector args{adaptor.getSource(), adaptor.getAlign()}; + if (!op.getUpdatedBase()) { + return VldusCallOperands{std::move(args), Value()}; + } + Type elementType = getElementTypeFromVectorLike(op.getResult().getType()); + auto loweredIncrement = + lowerVPTOElementOffsetForIntrinsic(op, adaptor.getSource(), adaptor.getIncrement(), elementType, + /*isPostUpdate=*/true, rewriter); + if (failed(loweredIncrement)) { + return failure(); + } + args.front() = loweredIncrement->base; + args.push_back(loweredIncrement->intrinsicOffset); + return VldusCallOperands{std::move(args), loweredIncrement->updatedBase}; +} + +class LowerVldusOpPattern final : public OpConversionPattern { +public: + explicit LowerVldusOpPattern(TypeConverter &typeConverter, MLIRContext *context, LoweringState &state) + : OpConversionPattern(typeConverter, context), state(state) {} + + LogicalResult matchAndRewrite(pto::VldusOp op, pto::VldusOp::Adaptor adaptor, + ConversionPatternRewriter &rewriter) const override { + auto sourceType = dyn_cast(adaptor.getSource().getType()); + SmallVector resultTypes; + bool usePostIntrinsic = static_cast(op.getUpdatedBase()); + if (!sourceType || failed(this->getTypeConverter()->convertTypes(op->getResultTypes(), resultTypes)) || + resultTypes.size() != (usePostIntrinsic ? 3U : 2U) || adaptor.getAlign().getType() != resultTypes[1] || + (usePostIntrinsic && resultTypes[2] != adaptor.getSource().getType())) { + return rewriter.notifyMatchFailure(op, "expected converted vldus operand/result types"); + } + + FailureOr calleeName = usePostIntrinsic ? buildVldusPostCallee(op.getContext(), op.getResult().getType()) + : buildVldusCallee(op.getContext(), op.getResult().getType()); + if (failed(calleeName)) { + return rewriter.notifyMatchFailure(op, "unsupported vldus signature"); + } + + Type callValueType = getPayloadABIType(op.getResult().getType(), resultTypes[0], rewriter.getContext()); + SmallVector intrinsicResultTypes{callValueType, resultTypes[1]}; + // The installed no-post A5 vldus intrinsic returns an extra hidden base ptr. + intrinsicResultTypes.push_back(adaptor.getSource().getType()); + + FailureOr callOperands = buildVldusCallOperands(op, adaptor, rewriter); + if (failed(callOperands)) { + return rewriter.notifyMatchFailure(op, "failed to convert vldus increment"); + } + SmallVector argTypes; + for (Value arg : callOperands->args) { + argTypes.push_back(arg.getType()); + } + auto funcType = rewriter.getFunctionType(argTypes, intrinsicResultTypes); + auto call = rewriter.create(op.getLoc(), *calleeName, intrinsicResultTypes, callOperands->args); + state.plannedDecls.push_back(PlannedDecl{calleeName->str(), funcType}); + Value loaded = + castFromPayloadABI(op.getLoc(), call.getResult(0), op.getResult().getType(), resultTypes[0], rewriter); + SmallVector replacements{loaded, call.getResult(1)}; + if (usePostIntrinsic) { + replacements.push_back(callOperands->explicitUpdatedBase ? callOperands->explicitUpdatedBase : call.getResult(2)); + } + rewriter.replaceOp(op, replacements); + return success(); + } + +private: + LoweringState &state; +}; + +class LowerSprclrOpPattern final : public OpConversionPattern { +public: + explicit LowerSprclrOpPattern(TypeConverter &typeConverter, MLIRContext *context, LoweringState &state) + : OpConversionPattern(typeConverter, context), state(state) {} + + LogicalResult matchAndRewrite(pto::SprclrOp op, pto::SprclrOp::Adaptor adaptor, + ConversionPatternRewriter &rewriter) const override { + (void)adaptor; + auto spr = parseSprImmediate(op.getSpr()); + if (!spr) { + return rewriter.notifyMatchFailure(op, "unsupported sprclr target"); + } + + StringRef calleeName = buildSprclrCallee(op.getContext()); + Value sprValue = rewriter.create(op.getLoc(), rewriter.getI16IntegerAttr(*spr)); + auto funcType = rewriter.getFunctionType(TypeRange{sprValue.getType()}, TypeRange{}); + rewriter.create(op.getLoc(), calleeName, TypeRange{}, ValueRange{sprValue}); + state.plannedDecls.push_back(PlannedDecl{calleeName.str(), funcType}); + rewriter.eraseOp(op); + return success(); + } + +private: + LoweringState &state; +}; + +template class LowerSprStoreOpPattern final : public OpConversionPattern { +public: + explicit LowerSprStoreOpPattern(TypeConverter &typeConverter, MLIRContext *context, LoweringState &state) + : OpConversionPattern(typeConverter, context), state(state) {} + + LogicalResult matchAndRewrite(SprStoreOp op, typename SprStoreOp::Adaptor adaptor, + ConversionPatternRewriter &rewriter) const override { + auto spr = parseSprImmediate(op.getSpr()); + if (!spr) { + return rewriter.notifyMatchFailure(op, "unsupported spr store target"); + } + auto destType = dyn_cast(adaptor.getDestination().getType()); + if (!destType || !adaptor.getOffset().getType().isInteger(32)) { + return rewriter.notifyMatchFailure(op, "expected converted spr store operands"); + } + + bool usePostIntrinsic = op.getUpdatedBase() != nullptr; + SmallVector resultTypes; + if (failed(this->getTypeConverter()->convertTypes(op->getResultTypes(), resultTypes)) || + resultTypes.size() != (usePostIntrinsic ? 1U : 0U)) { + return rewriter.notifyMatchFailure(op, "failed to convert spr store result types"); + } + + StringRef calleeName = buildSprStoreCallee(op.getContext(), usePostIntrinsic); + Value sprValue = rewriter.create(op.getLoc(), rewriter.getI16IntegerAttr(*spr)); + Value postValue = + rewriter.create(op.getLoc(), rewriter.getI32IntegerAttr(usePostIntrinsic ? 1 : 0)); + SmallVector args{sprValue, adaptor.getDestination(), adaptor.getOffset(), postValue}; + auto funcType = rewriter.getFunctionType(TypeRange{sprValue.getType(), adaptor.getDestination().getType(), + adaptor.getOffset().getType(), postValue.getType()}, + resultTypes); + auto call = rewriter.create(op.getLoc(), calleeName, resultTypes, args); + state.plannedDecls.push_back(PlannedDecl{calleeName.str(), funcType}); + if (usePostIntrinsic) { + rewriter.replaceOp(op, call.getResults()); + } else { + rewriter.eraseOp(op); + } + return success(); + } + +private: + LoweringState &state; +}; + +static Type getVPTOAddressElementType(Type addressType, Type fallbackType) { + if (auto ptrType = dyn_cast(addressType)) { + return ptrType.getElementType(); + } + if (auto memrefType = dyn_cast(addressType)) { + return memrefType.getElementType(); + } + return fallbackType; +} + +static SmallVector getVstsCallArgs(pto::VstsOp op, pto::VstsOp::Adaptor adaptor, + const VPTOLoweredAddressOffset &offset, uint64_t dist, bool usePostIntrinsic, + ConversionPatternRewriter &rewriter) { + Value value = castToPayloadABI(op.getLoc(), adaptor.getValue(), op.getValue().getType(), rewriter); + Value mask = adaptor.getMask(); + if (isOnePointStoreDist(op.getDist().value_or(""))) { + mask = rewriter.create(op.getLoc(), mask.getType()); + } + Value distValue = getI32Constant(rewriter, op.getLoc(), dist); + Value postValue = getI32Constant(rewriter, op.getLoc(), usePostIntrinsic ? 1 : 0); + return {value, offset.base, offset.intrinsicOffset, distValue, postValue, mask}; +} + +static LogicalResult replaceVstsOp(pto::VstsOp op, bool usePostIntrinsic, const VPTOLoweredAddressOffset &offset, + func::CallOp call, ConversionPatternRewriter &rewriter) { + if (!usePostIntrinsic) { + rewriter.eraseOp(op); + return success(); + } + rewriter.replaceOp(op, offset.updatedBase ? offset.updatedBase : call.getResult(0)); + return success(); +} + +class LowerVstsOpPattern final : public OpConversionPattern { +public: + explicit LowerVstsOpPattern(TypeConverter &typeConverter, MLIRContext *context, LoweringState &state) + : OpConversionPattern(typeConverter, context), state(state) {} + + LogicalResult matchAndRewrite(pto::VstsOp op, OpAdaptor adaptor, ConversionPatternRewriter &rewriter) const override { + Type elementType = getElementTypeFromVectorLike(op.getValue().getType()); + if (!elementType) { + return rewriter.notifyMatchFailure(op, "unsupported vsts element type"); + } + Type offsetElementType = getVPTOAddressElementType(op.getDestination().getType(), elementType); + bool usePostIntrinsic = static_cast(op.getUpdatedBase()); + auto loweredOffset = lowerVPTOElementOffsetForIntrinsic(op, adaptor.getDestination(), adaptor.getOffset(), + offsetElementType, usePostIntrinsic, rewriter); + auto dist = parseStoreDistImmediate(op.getDist().value_or(""), elementType); + bool invalidAddress = failed(loweredOffset) || !dist; + if (invalidAddress) { + return rewriter.notifyMatchFailure(op, "failed to materialize vsts operands"); + } + + FailureOr calleeName = op.getUpdatedBase() + ? buildVstsPostCallee(op.getContext(), op.getValue().getType()) + : buildVstsCallee(op.getContext(), op.getValue().getType()); + if (failed(calleeName)) { + return rewriter.notifyMatchFailure(op, "unsupported vsts signature"); + } + + SmallVector resultTypes; + if (failed(this->getTypeConverter()->convertTypes(op->getResultTypes(), resultTypes))) { + return rewriter.notifyMatchFailure(op, "failed to convert vsts result types"); + } + if (usePostIntrinsic) { + if (resultTypes.size() != 1 || resultTypes[0] != adaptor.getDestination().getType()) { + return rewriter.notifyMatchFailure(op, "unsupported vsts post-update result"); + } + } else if (!resultTypes.empty()) { + return rewriter.notifyMatchFailure(op, "unsupported vsts result count"); + } + + SmallVector args = getVstsCallArgs(op, adaptor, *loweredOffset, *dist, usePostIntrinsic, rewriter); + auto funcType = + rewriter.getFunctionType(TypeRange{args[0].getType(), loweredOffset->base.getType(), rewriter.getI32Type(), + rewriter.getI32Type(), rewriter.getI32Type(), args[5].getType()}, + resultTypes); + auto call = rewriter.create(op.getLoc(), *calleeName, resultTypes, args); + state.plannedDecls.push_back(PlannedDecl{calleeName->str(), funcType}); + return replaceVstsOp(op, usePostIntrinsic, *loweredOffset, call, rewriter); + } + +private: + LoweringState &state; +}; + +class LowerVsstbOpPattern final : public OpConversionPattern { +public: + explicit LowerVsstbOpPattern(TypeConverter &typeConverter, MLIRContext *context, LoweringState &state) + : OpConversionPattern(typeConverter, context), state(state) {} + + LogicalResult matchAndRewrite(pto::VsstbOp op, pto::VsstbOp::Adaptor adaptor, + ConversionPatternRewriter &rewriter) const override { + auto basePtr = dyn_cast(adaptor.getDestination().getType()); + Value packedStride = packBlockRepeatStride(op, adaptor.getBlockStride(), adaptor.getRepeatStride()); + if (!basePtr || !packedStride) { + return rewriter.notifyMatchFailure(op, "failed to materialize vsstb operands"); + } + + SmallVector resultTypes; + if (failed(this->getTypeConverter()->convertTypes(op->getResultTypes(), resultTypes))) { + return rewriter.notifyMatchFailure(op, "failed to convert vsstb result types"); + } + bool usePostIntrinsic = static_cast(op.getUpdatedBase()); + if (usePostIntrinsic) { + if (resultTypes.size() != 1 || resultTypes[0] != adaptor.getDestination().getType()) { + return rewriter.notifyMatchFailure(op, "unsupported vsstb post-update result"); + } + } else if (!resultTypes.empty()) { + return rewriter.notifyMatchFailure(op, "unsupported vsstb result count"); + } + + FailureOr calleeName = buildVsstbCallee(op.getContext(), op.getValue().getType(), usePostIntrinsic); + if (failed(calleeName)) { + return rewriter.notifyMatchFailure(op, "unsupported vsstb signature"); + } + Value zeroValue = getI32Constant(rewriter, op.getLoc(), usePostIntrinsic ? 1 : 0); + Value value = castToPayloadABI(op.getLoc(), adaptor.getValue(), op.getValue().getType(), rewriter); + SmallVector args{value, adaptor.getDestination(), packedStride, zeroValue, adaptor.getMask()}; + auto funcType = + rewriter.getFunctionType(TypeRange{value.getType(), adaptor.getDestination().getType(), packedStride.getType(), + zeroValue.getType(), adaptor.getMask().getType()}, + resultTypes); + auto call = rewriter.create(op.getLoc(), *calleeName, resultTypes, args); + state.plannedDecls.push_back(PlannedDecl{calleeName->str(), funcType}); + if (usePostIntrinsic) { + rewriter.replaceOp(op, call.getResults()); + } else { + rewriter.eraseOp(op); + } + return success(); + } + +private: + LoweringState &state; +}; + +class LowerVstsx2OpPattern final : public OpConversionPattern { +public: + explicit LowerVstsx2OpPattern(TypeConverter &typeConverter, MLIRContext *context, LoweringState &state) + : OpConversionPattern(typeConverter, context), state(state) {} + + LogicalResult matchAndRewrite(pto::Vstsx2Op op, pto::Vstsx2Op::Adaptor adaptor, + ConversionPatternRewriter &rewriter) const override { + Type elementType = getElementTypeFromVectorLike(op.getLow().getType()); + if (!elementType) { + return rewriter.notifyMatchFailure(op, "unsupported vstsx2 element type"); + } + + auto loweredOffset = + lowerVPTOElementOffsetForIntrinsic(op, adaptor.getDestination(), adaptor.getOffset(), elementType, + /*isPostUpdate=*/false, rewriter); + auto dist = parseStoreX2DistImmediate(op.getDist(), elementType); + bool invalidAddress = failed(loweredOffset) || !dist; + if (invalidAddress) { + return rewriter.notifyMatchFailure(op, "failed to materialize vstsx2 operands"); + } + + FailureOr calleeName = buildVstsx2Callee(op.getContext(), op.getLow().getType()); + if (failed(calleeName)) { + return rewriter.notifyMatchFailure(op, "unsupported vstsx2 signature"); + } + + Value distValue = getI32Constant(rewriter, op.getLoc(), *dist); + Value zeroValue = getI32Constant(rewriter, op.getLoc(), 0); + Value low = castToPayloadABI(op.getLoc(), adaptor.getLow(), op.getLow().getType(), rewriter); + Value high = castToPayloadABI(op.getLoc(), adaptor.getHigh(), op.getHigh().getType(), rewriter); + SmallVector args{low, high, loweredOffset->base, loweredOffset->intrinsicOffset, + distValue, zeroValue, adaptor.getMask()}; + auto funcType = rewriter.getFunctionType(TypeRange{low.getType(), high.getType(), loweredOffset->base.getType(), + loweredOffset->intrinsicOffset.getType(), distValue.getType(), + zeroValue.getType(), adaptor.getMask().getType()}, + TypeRange{}); + rewriter.create(op.getLoc(), *calleeName, TypeRange{}, args); + state.plannedDecls.push_back(PlannedDecl{calleeName->str(), funcType}); + rewriter.eraseOp(op); + return success(); + } + +private: + LoweringState &state; +}; + +class LowerPstuOpPattern final : public OpConversionPattern { +public: + explicit LowerPstuOpPattern(TypeConverter &typeConverter, MLIRContext *context, LoweringState &state) + : OpConversionPattern(typeConverter, context), state(state) {} + + LogicalResult matchAndRewrite(pto::PstuOp op, pto::PstuOp::Adaptor adaptor, + ConversionPatternRewriter &rewriter) const override { + FailureOr calleeName = buildPstuCallee(op.getContext(), op); + if (failed(calleeName)) { + return rewriter.notifyMatchFailure(op, "unsupported pstu signature"); + } + + SmallVector resultTypes; + if (failed(this->getTypeConverter()->convertTypes(op->getResultTypes(), resultTypes))) { + return rewriter.notifyMatchFailure(op, "failed to convert pstu result types"); + } + if (resultTypes.size() != 2) { + return rewriter.notifyMatchFailure(op, "unexpected converted pstu result arity"); + } + + auto baseType = dyn_cast(adaptor.getBase().getType()); + if (!baseType || adaptor.getAlignIn().getType() != resultTypes[0] || + adaptor.getBase().getType() != resultTypes[1]) { + return rewriter.notifyMatchFailure(op, "unexpected converted pstu operand/result types"); + } + + SmallVector args{adaptor.getValue(), adaptor.getBase(), adaptor.getAlignIn()}; + auto funcType = rewriter.getFunctionType( + TypeRange{adaptor.getValue().getType(), adaptor.getBase().getType(), adaptor.getAlignIn().getType()}, + resultTypes); + auto call = rewriter.create(op.getLoc(), *calleeName, resultTypes, args); + state.plannedDecls.push_back(PlannedDecl{calleeName->str(), funcType}); + rewriter.replaceOp(op, call.getResults()); + return success(); + } + +private: + LoweringState &state; +}; + +class LowerVstusOpPattern final : public OpConversionPattern { +public: + explicit LowerVstusOpPattern(TypeConverter &typeConverter, MLIRContext *context, LoweringState &state) + : OpConversionPattern(typeConverter, context), state(state) {} + + LogicalResult matchAndRewrite(pto::VstusOp op, pto::VstusOp::Adaptor adaptor, + ConversionPatternRewriter &rewriter) const override { + Type elementType = getElementTypeFromVectorLike(op.getValue().getType()); + if (!elementType) { + return rewriter.notifyMatchFailure(op, "unsupported vstus element type"); + } + + bool usePostIntrinsic = static_cast(op.getBaseOut()); + auto loweredOffset = lowerVPTOElementOffsetForIntrinsic(op, adaptor.getBase(), adaptor.getOffset(), elementType, + usePostIntrinsic, rewriter); + if (failed(loweredOffset)) { + return rewriter.notifyMatchFailure(op, "failed to convert vstus offset"); + } + + SmallVector resultTypes; + if (failed(this->getTypeConverter()->convertTypes(op->getResultTypes(), resultTypes))) { + return rewriter.notifyMatchFailure(op, "failed to convert vstus result types"); + } + auto baseType = dyn_cast(adaptor.getBase().getType()); + if (!baseType || resultTypes.size() != (usePostIntrinsic ? 2U : 1U) || + adaptor.getAlignIn().getType() != resultTypes[0] || + (usePostIntrinsic && resultTypes[1] != adaptor.getBase().getType())) { + return rewriter.notifyMatchFailure(op, "unexpected converted vstus operand/result types"); + } + + FailureOr calleeName = buildVstusCallee(op.getContext(), op.getValue().getType()); + if (usePostIntrinsic) { + calleeName = buildVstusPostCallee(op.getContext(), op.getValue().getType()); + } + if (failed(calleeName)) { + return rewriter.notifyMatchFailure(op, "unsupported vstus signature"); + } + Value value = castToPayloadABI(op.getLoc(), adaptor.getValue(), op.getValue().getType(), rewriter); + SmallVector args{value, loweredOffset->base, loweredOffset->intrinsicOffset, adaptor.getAlignIn()}; + auto funcType = + rewriter.getFunctionType(TypeRange{value.getType(), loweredOffset->base.getType(), + loweredOffset->intrinsicOffset.getType(), adaptor.getAlignIn().getType()}, + resultTypes); + auto call = rewriter.create(op.getLoc(), *calleeName, resultTypes, args); + state.plannedDecls.push_back(PlannedDecl{calleeName->str(), funcType}); + if (usePostIntrinsic && loweredOffset->updatedBase) { + rewriter.replaceOp(op, ValueRange{call.getResult(0), loweredOffset->updatedBase}); + } else { + rewriter.replaceOp(op, call.getResults()); + } + return success(); + } + +private: + LoweringState &state; +}; + +class LowerVsturOpPattern final : public OpConversionPattern { +public: + explicit LowerVsturOpPattern(TypeConverter &typeConverter, MLIRContext *context, LoweringState &state) + : OpConversionPattern(typeConverter, context), state(state) {} + + LogicalResult matchAndRewrite(pto::VsturOp op, pto::VsturOp::Adaptor adaptor, + ConversionPatternRewriter &rewriter) const override { + auto postMode = parsePostModeImmediate(op.getMode()); + if (!postMode) { + return rewriter.notifyMatchFailure(op, "unsupported vstur mode immediate"); + } + + Type resultType = this->getTypeConverter()->convertType(op.getAlignOut().getType()); + auto baseType = dyn_cast(adaptor.getBase().getType()); + if (!resultType || !baseType || adaptor.getAlignIn().getType() != resultType) { + return rewriter.notifyMatchFailure(op, "unexpected converted vstur operand/result types"); + } + + StringRef calleeName = buildVsturCallee(op.getContext()); + Value modeValue = getI32Constant(rewriter, op.getLoc(), *postMode); + Value zeroValue = getI32Constant(rewriter, op.getLoc(), 0); + Value value = castToPayloadABI(op.getLoc(), adaptor.getValue(), op.getValue().getType(), rewriter); + SmallVector args{value, adaptor.getBase(), adaptor.getAlignIn(), modeValue, zeroValue}; + auto funcType = + rewriter.getFunctionType(TypeRange{value.getType(), adaptor.getBase().getType(), adaptor.getAlignIn().getType(), + modeValue.getType(), zeroValue.getType()}, + TypeRange{resultType}); + auto call = rewriter.create(op.getLoc(), calleeName, TypeRange{resultType}, args); + state.plannedDecls.push_back(PlannedDecl{calleeName.str(), funcType}); + rewriter.replaceOp(op, call.getResults()); + return success(); + } + +private: + LoweringState &state; +}; + +class LowerVstarOpPattern final : public OpConversionPattern { +public: + explicit LowerVstarOpPattern(TypeConverter &typeConverter, MLIRContext *context, LoweringState &state) + : OpConversionPattern(typeConverter, context), state(state) {} + + LogicalResult matchAndRewrite(pto::VstarOp op, pto::VstarOp::Adaptor adaptor, + ConversionPatternRewriter &rewriter) const override { + auto baseType = dyn_cast(adaptor.getDestination().getType()); + Type alignType = this->getTypeConverter()->convertType(op.getValue().getType()); + if (!baseType || !alignType || adaptor.getValue().getType() != alignType) { + return rewriter.notifyMatchFailure(op, "unexpected converted vstar operand types"); + } + + StringRef calleeName = buildVstarCallee(op.getContext()); + Value zeroValue = getI32Constant(rewriter, op.getLoc(), 0); + SmallVector args{adaptor.getValue(), adaptor.getDestination(), zeroValue}; + auto funcType = rewriter.getFunctionType( + TypeRange{adaptor.getValue().getType(), adaptor.getDestination().getType(), zeroValue.getType()}, TypeRange{}); + rewriter.create(op.getLoc(), calleeName, TypeRange{}, args); + state.plannedDecls.push_back(PlannedDecl{calleeName.str(), funcType}); + rewriter.eraseOp(op); + return success(); + } + +private: + LoweringState &state; +}; + +class LowerVstasOpPattern final : public OpConversionPattern { +public: + explicit LowerVstasOpPattern(TypeConverter &typeConverter, MLIRContext *context, LoweringState &state) + : OpConversionPattern(typeConverter, context), state(state) {} + + LogicalResult matchAndRewrite(pto::VstasOp op, pto::VstasOp::Adaptor adaptor, + ConversionPatternRewriter &rewriter) const override { + auto baseType = dyn_cast(adaptor.getDestination().getType()); + Type alignType = this->getTypeConverter()->convertType(op.getValue().getType()); + auto dstType = dyn_cast(op.getDestination().getType()); + if (!baseType || !alignType || adaptor.getValue().getType() != alignType || !dstType) { + return rewriter.notifyMatchFailure(op, "unexpected converted vstas operand types"); + } + + bool usePostIntrinsic = op.getUpdatedBase() != nullptr; + auto loweredOffset = lowerVPTOElementOffsetForIntrinsic(op, adaptor.getDestination(), adaptor.getOffset(), + dstType.getElementType(), usePostIntrinsic, rewriter); + if (failed(loweredOffset)) { + return rewriter.notifyMatchFailure(op, "failed to convert vstas offset"); + } + + SmallVector resultTypes; + if (failed(this->getTypeConverter()->convertTypes(op->getResultTypes(), resultTypes)) || + resultTypes.size() != (usePostIntrinsic ? 1U : 0U)) { + return rewriter.notifyMatchFailure(op, "failed to convert vstas result types"); + } + + StringRef calleeName = buildVstasCallee(op.getContext(), usePostIntrinsic); + Value postValue = getI32Constant(rewriter, op.getLoc(), usePostIntrinsic ? 1 : 0); + SmallVector args{adaptor.getValue(), loweredOffset->base, loweredOffset->intrinsicOffset, postValue}; + auto funcType = rewriter.getFunctionType(TypeRange{adaptor.getValue().getType(), loweredOffset->base.getType(), + loweredOffset->intrinsicOffset.getType(), postValue.getType()}, + resultTypes); + auto call = rewriter.create(op.getLoc(), calleeName, resultTypes, args); + state.plannedDecls.push_back(PlannedDecl{calleeName.str(), funcType}); + if (usePostIntrinsic) { + Value updatedBase = loweredOffset->updatedBase ? loweredOffset->updatedBase : call.getResult(0); + rewriter.replaceOp(op, updatedBase); + } else { + rewriter.eraseOp(op); + } + return success(); + } + +private: + LoweringState &state; +}; + +class LowerVgather2OpPattern final : public OpConversionPattern { +public: + explicit LowerVgather2OpPattern(TypeConverter &typeConverter, MLIRContext *context, LoweringState &state) + : OpConversionPattern(typeConverter, context), state(state) {} + + LogicalResult matchAndRewrite(pto::Vgather2Op op, pto::Vgather2Op::Adaptor adaptor, + ConversionPatternRewriter &rewriter) const override { + Type elemType = getElementTypeFromVectorLike(op.getResult().getType()); + auto basePtr = dyn_cast(adaptor.getSource().getType()); + if (!elemType || !basePtr) { + return rewriter.notifyMatchFailure(op, "unexpected converted vgather2 operand types"); + } + + Type resultType = this->getTypeConverter()->convertType(op.getResult().getType()); + if (!resultType) { + return rewriter.notifyMatchFailure(op, "failed to convert vgather2 result type"); + } + + FailureOr calleeName = + buildVgather2Callee(op.getContext(), op.getSource().getType(), op.getResult().getType()); + if (failed(calleeName)) { + return rewriter.notifyMatchFailure(op, "unsupported vgather2 signature"); + } + + Value offsets = adaptor.getOffsets(); + FailureOr offsetsCarrierType = + getVgather2OffsetsCarrierType(rewriter, op.getSource().getType(), op.getResult().getType(), offsets.getType()); + if (failed(offsetsCarrierType)) { + return rewriter.notifyMatchFailure(op, "unsupported vgather2 offsets carrier"); + } + if (offsets.getType() != *offsetsCarrierType) { + offsets = rewriter.create(op.getLoc(), *offsetsCarrierType, offsets); + } + + auto funcType = rewriter.getFunctionType( + TypeRange{adaptor.getSource().getType(), *offsetsCarrierType, adaptor.getMask().getType()}, + TypeRange{resultType}); + auto call = rewriter.create(op.getLoc(), *calleeName, TypeRange{resultType}, + ValueRange{adaptor.getSource(), offsets, adaptor.getMask()}); + state.plannedDecls.push_back(PlannedDecl{calleeName->str(), funcType}); + rewriter.replaceOp(op, call.getResults()); + return success(); + } + +private: + LoweringState &state; +}; + +class LowerVgather2BcOpPattern final : public OpConversionPattern { +public: + explicit LowerVgather2BcOpPattern(TypeConverter &typeConverter, MLIRContext *context, LoweringState &state) + : OpConversionPattern(typeConverter, context), state(state) {} + + LogicalResult matchAndRewrite(pto::Vgather2BcOp op, pto::Vgather2BcOp::Adaptor adaptor, + ConversionPatternRewriter &rewriter) const override { + auto basePtr = dyn_cast(adaptor.getSource().getType()); + Type resultType = this->getTypeConverter()->convertType(op.getResult().getType()); + if (!basePtr || !resultType) { + return rewriter.notifyMatchFailure(op, "unexpected converted vgather2_bc operand/result types"); + } + + FailureOr calleeName = buildVgather2BcCallee(op.getContext(), op.getResult().getType()); + if (failed(calleeName)) { + return rewriter.notifyMatchFailure(op, "unsupported vgather2_bc signature"); + } + + auto funcType = rewriter.getFunctionType( + TypeRange{adaptor.getSource().getType(), adaptor.getOffsets().getType(), adaptor.getMask().getType()}, + TypeRange{resultType}); + auto call = rewriter.create(op.getLoc(), *calleeName, TypeRange{resultType}, + ValueRange{adaptor.getSource(), adaptor.getOffsets(), adaptor.getMask()}); + state.plannedDecls.push_back(PlannedDecl{calleeName->str(), funcType}); + rewriter.replaceOp(op, call.getResults()); + return success(); + } + +private: + LoweringState &state; +}; + +class LowerVgatherbOpPattern final : public OpConversionPattern { +public: + explicit LowerVgatherbOpPattern(TypeConverter &typeConverter, MLIRContext *context, LoweringState &state) + : OpConversionPattern(typeConverter, context), state(state) {} + + LogicalResult matchAndRewrite(pto::VgatherbOp op, pto::VgatherbOp::Adaptor adaptor, + ConversionPatternRewriter &rewriter) const override { + auto basePtr = dyn_cast(adaptor.getSource().getType()); + Type resultType = this->getTypeConverter()->convertType(op.getResult().getType()); + if (!basePtr || !resultType) { + return rewriter.notifyMatchFailure(op, "unexpected converted vgatherb operand/result types"); + } + + FailureOr calleeName = buildVgatherbCallee(op.getContext(), op.getResult().getType()); + if (failed(calleeName)) { + return rewriter.notifyMatchFailure(op, "unsupported vgatherb signature"); + } + + auto funcType = rewriter.getFunctionType( + TypeRange{adaptor.getSource().getType(), adaptor.getOffsets().getType(), adaptor.getMask().getType()}, + TypeRange{resultType}); + auto call = rewriter.create(op.getLoc(), *calleeName, TypeRange{resultType}, + ValueRange{adaptor.getSource(), adaptor.getOffsets(), adaptor.getMask()}); + state.plannedDecls.push_back(PlannedDecl{calleeName->str(), funcType}); + rewriter.replaceOp(op, call.getResults()); + return success(); + } + +private: + LoweringState &state; +}; + +class LowerVscatterOpPattern final : public OpConversionPattern { +public: + explicit LowerVscatterOpPattern(TypeConverter &typeConverter, MLIRContext *context, LoweringState &state) + : OpConversionPattern(typeConverter, context), state(state) {} + + LogicalResult matchAndRewrite(pto::VscatterOp op, pto::VscatterOp::Adaptor adaptor, + ConversionPatternRewriter &rewriter) const override { + Type elemType = getElementTypeFromVectorLike(op.getValue().getType()); + auto basePtr = dyn_cast(adaptor.getDestination().getType()); + if (!elemType || !basePtr) { + return rewriter.notifyMatchFailure(op, "unexpected converted vscatter operand types"); + } + + FailureOr calleeName = buildVscatterCallee(op.getContext(), op.getValue().getType()); + if (failed(calleeName)) { + return rewriter.notifyMatchFailure(op, "unsupported vscatter signature"); + } + + FailureOr offsetsCarrierType = getVscatterOffsetsCarrierType(adaptor.getOffsets().getType()); + if (failed(offsetsCarrierType)) { + return rewriter.notifyMatchFailure(op, "unsupported vscatter offsets carrier"); + } + + auto funcType = rewriter.getFunctionType(TypeRange{adaptor.getValue().getType(), adaptor.getDestination().getType(), + *offsetsCarrierType, adaptor.getMask().getType()}, + TypeRange{}); + rewriter.create( + op.getLoc(), *calleeName, TypeRange{}, + ValueRange{adaptor.getValue(), adaptor.getDestination(), adaptor.getOffsets(), adaptor.getMask()}); + state.plannedDecls.push_back(PlannedDecl{calleeName->str(), funcType}); + rewriter.eraseOp(op); + return success(); + } + +private: + LoweringState &state; +}; + +class LowerVaxpyOpPattern final : public OpConversionPattern { +public: + explicit LowerVaxpyOpPattern(TypeConverter &typeConverter, MLIRContext *context, LoweringState &state) + : OpConversionPattern(typeConverter, context), state(state) {} + + LogicalResult matchAndRewrite(pto::VaxpyOp op, pto::VaxpyOp::Adaptor adaptor, + ConversionPatternRewriter &rewriter) const override { + Type elemType = getElementTypeFromVectorLike(op.getResult().getType()); + if (!elemType) { + return rewriter.notifyMatchFailure(op, "unsupported vaxpy signature"); + } + + Type resultType = this->getTypeConverter()->convertType(op.getResult().getType()); + if (!resultType) { + return rewriter.notifyMatchFailure(op, "failed to convert vaxpy result type"); + } + + FailureOr calleeName = buildVaxpyCallee(op.getContext(), op.getResult().getType()); + if (failed(calleeName)) { + return rewriter.notifyMatchFailure(op, "unsupported vaxpy callee"); + } + + auto funcType = rewriter.getFunctionType(TypeRange{adaptor.getSrc1().getType(), adaptor.getSrc0().getType(), + adaptor.getAlpha().getType(), adaptor.getMask().getType()}, + TypeRange{resultType}); + auto call = rewriter.create( + op.getLoc(), *calleeName, TypeRange{resultType}, + ValueRange{adaptor.getSrc1(), adaptor.getSrc0(), adaptor.getAlpha(), adaptor.getMask()}); + state.plannedDecls.push_back(PlannedDecl{calleeName->str(), funcType}); + rewriter.replaceOp(op, call.getResults()); + return success(); + } + +private: + LoweringState &state; +}; + +class LowerVmulscvtOpPattern final : public OpConversionPattern { +public: + explicit LowerVmulscvtOpPattern(TypeConverter &typeConverter, MLIRContext *context, LoweringState &state) + : OpConversionPattern(typeConverter, context), state(state) {} + + LogicalResult matchAndRewrite(pto::VmulscvtOp op, pto::VmulscvtOp::Adaptor adaptor, + ConversionPatternRewriter &rewriter) const override { + auto roundMode = parseRoundModeImmediate(op.getRnd()); + if (!roundMode) { + return rewriter.notifyMatchFailure(op, "vmulscvt requires valid rnd attr"); + } + if (*roundMode != 1) { + return rewriter.notifyMatchFailure(op, "current vmulscvt lowering only supports rnd A"); + } + + auto part = parsePartImmediate(op.getPart()); + if (!part) { + return rewriter.notifyMatchFailure(op, "unsupported vmulscvt part"); + } + + Type resultType = this->getTypeConverter()->convertType(op.getResult().getType()); + if (!resultType) { + return rewriter.notifyMatchFailure(op, "failed to convert vmulscvt result type"); + } + + FailureOr calleeName = + buildVmulscvtCallee(op.getContext(), op.getInput().getType(), op.getResult().getType()); + if (failed(calleeName)) { + return rewriter.notifyMatchFailure(op, "unsupported vmulscvt signature"); + } + + Value partValue = getI32Constant(rewriter, op.getLoc(), *part); + auto funcType = rewriter.getFunctionType(TypeRange{adaptor.getInput().getType(), adaptor.getScalar().getType(), + adaptor.getMask().getType(), partValue.getType()}, + TypeRange{resultType}); + auto call = rewriter.create( + op.getLoc(), *calleeName, TypeRange{resultType}, + ValueRange{adaptor.getInput(), adaptor.getScalar(), adaptor.getMask(), partValue}); + state.plannedDecls.push_back(PlannedDecl{calleeName->str(), funcType}); + rewriter.replaceOp(op, call.getResults()); + return success(); + } + +private: + LoweringState &state; +}; + +class LowerVciOpPattern final : public OpConversionPattern { +public: + explicit LowerVciOpPattern(TypeConverter &typeConverter, MLIRContext *context, LoweringState &state) + : OpConversionPattern(typeConverter, context), state(state) {} + + LogicalResult matchAndRewrite(pto::VciOp op, pto::VciOp::Adaptor adaptor, + ConversionPatternRewriter &rewriter) const override { + auto order = parseOrderImmediate(op.getOrder().value_or("ASC")); + if (!order) { + return rewriter.notifyMatchFailure(op, "unsupported vci order"); + } + + Type resultType = this->getTypeConverter()->convertType(op.getResult().getType()); + if (!resultType) { + return rewriter.notifyMatchFailure(op, "failed to convert vci result type"); + } + + FailureOr calleeName = buildVciCallee(op.getContext(), op.getResult().getType()); + if (failed(calleeName)) { + return rewriter.notifyMatchFailure(op, "unsupported vci callee"); + } + + Value indexValue = adaptor.getIndex(); + + Value orderValue = getI32Constant(rewriter, op.getLoc(), *order); + auto funcType = + rewriter.getFunctionType(TypeRange{indexValue.getType(), orderValue.getType()}, TypeRange{resultType}); + auto call = rewriter.create(op.getLoc(), *calleeName, TypeRange{resultType}, + ValueRange{indexValue, orderValue}); + state.plannedDecls.push_back(PlannedDecl{calleeName->str(), funcType}); + rewriter.replaceOp(op, call.getResults()); + return success(); + } + +private: + LoweringState &state; +}; + +class LowerVexpdifOpPattern final : public OpConversionPattern { +public: + explicit LowerVexpdifOpPattern(TypeConverter &typeConverter, MLIRContext *context, LoweringState &state) + : OpConversionPattern(typeConverter, context), state(state) {} + + LogicalResult matchAndRewrite(pto::VexpdifOp op, pto::VexpdifOp::Adaptor adaptor, + ConversionPatternRewriter &rewriter) const override { + auto part = parsePartImmediate(op.getPart()); + if (!part) { + return rewriter.notifyMatchFailure(op, "unsupported vexpdif signature"); + } + + Type resultType = this->getTypeConverter()->convertType(op.getResult().getType()); + if (!resultType) { + return rewriter.notifyMatchFailure(op, "failed to convert vexpdif result type"); + } + + FailureOr calleeName = + buildVexpdifCallee(op.getContext(), op.getInput().getType(), op.getResult().getType()); + if (failed(calleeName)) { + return rewriter.notifyMatchFailure(op, "unsupported vexpdif callee"); + } + + Value partValue = getI32Constant(rewriter, op.getLoc(), *part); + auto funcType = rewriter.getFunctionType(TypeRange{adaptor.getInput().getType(), adaptor.getMax().getType(), + adaptor.getMask().getType(), partValue.getType()}, + TypeRange{resultType}); + auto call = + rewriter.create(op.getLoc(), *calleeName, TypeRange{resultType}, + ValueRange{adaptor.getInput(), adaptor.getMax(), adaptor.getMask(), partValue}); + state.plannedDecls.push_back(PlannedDecl{calleeName->str(), funcType}); + rewriter.replaceOp(op, call.getResults()); + return success(); + } + +private: + LoweringState &state; +}; + +class LowerVbitsortOpPattern final : public OpConversionPattern { +public: + explicit LowerVbitsortOpPattern(TypeConverter &typeConverter, MLIRContext *context, LoweringState &state) + : OpConversionPattern(typeConverter, context), state(state) {} + + LogicalResult matchAndRewrite(pto::VbitsortOp op, pto::VbitsortOp::Adaptor adaptor, + ConversionPatternRewriter &rewriter) const override { + auto dstType = dyn_cast(adaptor.getDestination().getType()); + auto srcType = dyn_cast(adaptor.getSource().getType()); + auto idxType = dyn_cast(adaptor.getIndices().getType()); + if (!dstType || !srcType || !idxType) { + return rewriter.notifyMatchFailure(op, "unexpected converted vbitsort operand types"); + } + + FailureOr config = packVbitsortConfig(op, adaptor.getRepeatTimes()); + if (failed(config)) { + return rewriter.notifyMatchFailure(op, "failed to pack vbitsort config"); + } + + FailureOr calleeName = buildVbitsortCallee(op.getContext(), op); + if (failed(calleeName)) { + return rewriter.notifyMatchFailure(op, "unsupported vbitsort signature"); + } + + auto funcType = + rewriter.getFunctionType(TypeRange{adaptor.getDestination().getType(), adaptor.getSource().getType(), + adaptor.getIndices().getType(), (*config).getType()}, + TypeRange{}); + rewriter.create( + op.getLoc(), *calleeName, TypeRange{}, + ValueRange{adaptor.getDestination(), adaptor.getSource(), adaptor.getIndices(), *config}); + state.plannedDecls.push_back(PlannedDecl{calleeName->str(), funcType}); + rewriter.eraseOp(op); + return success(); + } + +private: + LoweringState &state; +}; + +class LowerVmrgsort4OpPattern final : public OpConversionPattern { +public: + explicit LowerVmrgsort4OpPattern(TypeConverter &typeConverter, MLIRContext *context, LoweringState &state) + : OpConversionPattern(typeConverter, context), state(state) {} + + LogicalResult matchAndRewrite(pto::Vmrgsort4Op op, pto::Vmrgsort4Op::Adaptor adaptor, + ConversionPatternRewriter &rewriter) const override { + auto dstType = dyn_cast(adaptor.getDestination().getType()); + auto src0Type = dyn_cast(adaptor.getSource0().getType()); + auto src1Type = dyn_cast(adaptor.getSource1().getType()); + auto src2Type = dyn_cast(adaptor.getSource2().getType()); + auto src3Type = dyn_cast(adaptor.getSource3().getType()); + if (!dstType || !src0Type || !src1Type || !src2Type || !src3Type) { + return rewriter.notifyMatchFailure(op, "unexpected converted vmrgsort4 operand types"); + } + + Type elemType = cast(op.getDestination().getType()).getElementType(); + FailureOr packedSrc = packVmrgsort4SourceAddr(op, adaptor.getSource0(), adaptor.getSource1(), + adaptor.getSource2(), adaptor.getSource3(), elemType); + if (failed(packedSrc)) { + return rewriter.notifyMatchFailure(op, "failed to pack vmrgsort4 source addresses"); + } + + FailureOr dst = reinterpretPointerToAddrSpace(op, adaptor.getDestination(), 6); + if (failed(dst)) { + return rewriter.notifyMatchFailure(op, "failed to normalize vmrgsort4 destination"); + } + + FailureOr calleeName = buildVmrgsort4Callee(op.getContext(), op); + if (failed(calleeName)) { + return rewriter.notifyMatchFailure(op, "unsupported vmrgsort4 signature"); + } + + auto funcType = rewriter.getFunctionType(TypeRange{(*dst).getType(), (*packedSrc).getType(), + adaptor.getCount().getType(), adaptor.getConfig().getType()}, + TypeRange{}); + rewriter.create(op.getLoc(), *calleeName, TypeRange{}, + ValueRange{*dst, *packedSrc, adaptor.getCount(), adaptor.getConfig()}); + state.plannedDecls.push_back(PlannedDecl{calleeName->str(), funcType}); + rewriter.eraseOp(op); + return success(); + } + +private: + LoweringState &state; +}; + +void populateVPTOVectorMemoryPatterns(VPTOTypeConverter &typeConverter, RewritePatternSet &patterns, + LoweringState &state) { + patterns + .add, LowerPredicateMaskBinaryOpPattern, + LowerPredicateMaskBinaryOpPattern, LowerPredicateMaskBinaryOpPattern, + LowerPredicatePairReorderOpPattern, LowerPredicatePairReorderOpPattern, + LowerPredicatePairReorderOpPattern, LowerPredicatePairReorderOpPattern, + LowerPredicatePairReorderOpPattern, LowerPredicatePairReorderOpPattern, + LowerUnpackOpPattern, LowerUnpackOpPattern, LowerVpackOpPattern, + LowerCmpOpPattern, LowerCmpOpPattern, LowerPltOpPattern, + LowerPltOpPattern, LowerPltOpPattern, LowerPltmOpPattern, + LowerPltmOpPattern, LowerPltmOpPattern, LowerPsetOpPattern, + LowerPsetOpPattern, LowerPsetOpPattern, LowerPgeOpPattern, + LowerPgeOpPattern, LowerPgeOpPattern, LowerVldsOpPattern, LowerVldsx2OpPattern, + LowerVsldbOpPattern, LowerVldasOpPattern, LowerInitAlignOpPattern, LowerVldusOpPattern, LowerSprclrOpPattern, + LowerSprStoreOpPattern, LowerSprStoreOpPattern, LowerVstsOpPattern, + LowerVsstbOpPattern, LowerVstsx2OpPattern, LowerVstarOpPattern, LowerVstasOpPattern, LowerVgather2OpPattern, + LowerVgather2BcOpPattern, LowerVgatherbOpPattern, LowerVscatterOpPattern, LowerVaxpyOpPattern, + LowerVmulscvtOpPattern, LowerVciOpPattern, LowerVexpdifOpPattern, LowerVbitsortOpPattern, + LowerVmrgsort4OpPattern, LowerPstuOpPattern, LowerVstusOpPattern, LowerVsturOpPattern>( + typeConverter, patterns.getContext(), state); +} + +} // namespace mlir::pto::detail diff --git a/lib/PTO/Transforms/VPTOCANN900LLVMEmitterPacking.cpp b/lib/PTO/Transforms/VPTOCANN900LLVMEmitterPacking.cpp new file mode 100644 index 0000000000..ba78b3f516 --- /dev/null +++ b/lib/PTO/Transforms/VPTOCANN900LLVMEmitterPacking.cpp @@ -0,0 +1,658 @@ +// Copyright (c) 2026 Huawei Technologies Co., Ltd. +// This program is free software, you can redistribute it and/or modify it under the terms and conditions of +// CANN Open Software License Agreement Version 2.0 (the "License"). +// Please refer to the License for details. You may not use this file except in compliance with the License. +// THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED, +// INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE. +// See LICENSE in the root of the software repository for the full text of the License. + +#include "VPTOCANN900LLVMEmitterInternal.h" + +namespace mlir::pto::detail { + +FailureOr packShiftedFields(Operation *anchor, Value base, ArrayRef> fields) { + OpBuilder builder(anchor); + builder.setInsertionPoint(anchor); + Value result = castIntegerLikeTo(anchor, base, builder.getI64Type()); + if (!result) { + return failure(); + } + for (const auto &[field, shift] : fields) { + Value value = castIntegerLikeTo(anchor, field, builder.getI64Type()); + if (!value) { + return failure(); + } + Value shifted = + builder.create(anchor->getLoc(), value, getI64Constant(builder, anchor->getLoc(), shift)); + result = builder.create(anchor->getLoc(), result, shifted); + } + return result; +} + +std::optional parseLoadX2DistImmediate(StringRef dist, Type elementType) { + auto width = getDistElementWidth(elementType); + if (dist == "BDINTLV") { + return 10; + } + if (!width) { + return std::nullopt; + } + if (dist == "DINTLV_B8") { + return std::optional(11); + } + if (dist == "DINTLV_B16") { + return std::optional(12); + } + if (dist == "DINTLV_B32") { + return std::optional(19); + } + return std::nullopt; +} + +std::optional parseStoreDistImmediate(StringRef dist, Type elementType) { + auto width = getDistElementWidth(elementType); + if (dist.empty()) { + if (!width) { + return std::nullopt; + } + if (*width == 8) { + return 0; + } + if (*width == 16) { + return 1; + } + if (*width == 32) { + return 2; + } + return std::nullopt; + } + static constexpr std::pair encodings[] = { + {"NORM_B8", 0}, {"NORM_B16", 1}, {"NORM_B32", 2}, {"1PT_B8", 3}, {"1PT_B16", 4}, + {"1PT_B32", 5}, {"PK_B16", 6}, {"PK_B32", 7}, {"PK_B64", 10}, {"PK4_B32", 12}, + {"MRG4CHN_B8", 13}, {"MRG2CHN_B8", 14}, {"MRG2CHN_B16", 15}, + }; + for (const auto &[name, value] : encodings) { + if (dist == name) { + return value; + } + } + return std::nullopt; +} + +bool isOnePointStoreDist(StringRef dist) { return dist == "1PT_B8" || dist == "1PT_B16" || dist == "1PT_B32"; } + +bool isMaskOnlyUsedByOnePointStores(Value mask) { + return !mask.use_empty() && llvm::all_of(mask.getUsers(), [](Operation *user) { + auto store = dyn_cast(user); + return store && store.getDist() && isOnePointStoreDist(*store.getDist()); + }); +} + +std::optional parseStoreX2DistImmediate(StringRef dist, Type elementType) { + auto width = getDistElementWidth(elementType); + if (!width) { + return std::nullopt; + } + if (dist == "INTLV_B8") { + return std::optional(8); + } + if (dist == "INTLV_B16") { + return std::optional(9); + } + if (dist == "INTLV_B32") { + return std::optional(11); + } + return std::nullopt; +} + +Value packBlockRepeatStride(Operation *anchor, Value blockStride, Value repeatStride) { + OpBuilder builder(anchor); + builder.setInsertionPoint(anchor); + + Value blockI32 = castIntegerLikeTo(anchor, blockStride, builder.getI32Type()); + Value repeatI32 = castIntegerLikeTo(anchor, repeatStride, builder.getI32Type()); + if (!blockI32 || !repeatI32) { + return {}; + } + + auto c16 = builder.create(anchor->getLoc(), 16, 32); + auto blockShifted = builder.create(anchor->getLoc(), blockI32, c16); + return builder.create(anchor->getLoc(), blockShifted, repeatI32).getResult(); +} + +std::optional parseOrderImmediate(StringRef order) { + if (order.empty() || order == "ASC") { + return 0; + } + if (order == "DESC") { + return 1; + } + return std::nullopt; +} + +FailureOr packLoopPair(Operation *anchor, Value low, Value high) { + OpBuilder builder(anchor); + builder.setInsertionPoint(anchor); + + Value lowI64 = castIntegerLikeTo(anchor, low, builder.getI64Type()); + Value highI64 = castIntegerLikeTo(anchor, high, builder.getI64Type()); + if (!lowI64 || !highI64) { + return failure(); + } + + Value shift = getI64Constant(builder, anchor->getLoc(), 40); + Value highShifted = builder.create(anchor->getLoc(), highI64, shift).getResult(); + return builder.create(anchor->getLoc(), highShifted, lowI64).getResult(); +} + +FailureOr packLoopSize(Operation *anchor, Value loop2, Value loop1) { + OpBuilder builder(anchor); + builder.setInsertionPoint(anchor); + + Value loop2I64 = castIntegerLikeTo(anchor, loop2, builder.getI64Type()); + Value loop1I64 = castIntegerLikeTo(anchor, loop1, builder.getI64Type()); + if (!loop2I64 || !loop1I64) { + return failure(); + } + + Value shift = getI64Constant(builder, anchor->getLoc(), 21); + Value loop2Shifted = builder.create(anchor->getLoc(), loop2I64, shift).getResult(); + return builder.create(anchor->getLoc(), loop2Shifted, loop1I64).getResult(); +} + +FailureOr packCopyGmToUbConfig0(Operation *anchor, ValueRange operands) { + if (operands.size() != 11) { + return failure(); + } + + SmallVector, 6> fields = {{operands[3], 4}, {operands[4], 25}, {operands[5], 46}, + {operands[6], 52}, {operands[7], 58}, {operands[8], 60}}; + return packShiftedFields(anchor, operands[2], fields); +} + +FailureOr packCopyGmToUbConfig1(Operation *anchor, ValueRange operands) { + if (operands.size() != 11) { + return failure(); + } + return packLoopPair(anchor, operands[9], operands[10]); +} + +[[maybe_unused]] FailureOr packCopyGmToUbConfig0(Operation *anchor, Value sid, Value nBurst, Value lenBurst, + Value leftPadding, Value rightPadding, Value dataSelect, + Value cacheCtl) { + SmallVector operands(11); + operands[2] = sid; + operands[3] = nBurst; + operands[4] = lenBurst; + operands[5] = leftPadding; + operands[6] = rightPadding; + operands[7] = dataSelect; + operands[8] = cacheCtl; + return packCopyGmToUbConfig0(anchor, operands); +} + +FailureOr packCopyUbToGmConfig0(Operation *anchor, ValueRange operands) { + if (operands.size() != 8) { + return failure(); + } + + OpBuilder builder(anchor); + builder.setInsertionPoint(anchor); + Location loc = anchor->getLoc(); + + auto getI64Operand = [&](unsigned idx) -> Value { + return castIntegerLikeTo(anchor, operands[idx], builder.getI64Type()); + }; + + Value sid = getI64Operand(2); + Value nBurst = getI64Operand(3); + Value lenBurst = getI64Operand(4); + Value l2CacheCtl = getI64Operand(5); + if (!sid || !nBurst || !lenBurst || !l2CacheCtl) { + return failure(); + } + + auto shl = [&](Value value, uint64_t amount) -> Value { + return builder.create(loc, value, getI64Constant(builder, loc, amount)); + }; + auto bitOr = [&](Value lhs, Value rhs) -> Value { return builder.create(loc, lhs, rhs); }; + + Value config = sid; + config = bitOr(config, shl(nBurst, 4)); + config = bitOr(config, shl(lenBurst, 25)); + config = bitOr(config, shl(l2CacheCtl, 60)); + return config; +} + +FailureOr packCopyUbToGmConfig1(Operation *anchor, ValueRange operands) { + if (operands.size() != 8) { + return failure(); + } + return packLoopPair(anchor, operands[6], operands[7]); +} + +[[maybe_unused]] FailureOr packCopyUbToGmConfig0(Operation *anchor, Value sid, Value nBurst, Value lenBurst, + Value l2CacheCtl) { + SmallVector operands(8); + operands[2] = sid; + operands[3] = nBurst; + operands[4] = lenBurst; + operands[5] = l2CacheCtl; + return packCopyUbToGmConfig0(anchor, operands); +} + +FailureOr packCopyUbToUbConfig(Operation *anchor, ValueRange operands) { + if (operands.size() != 7) { + return failure(); + } + OpBuilder builder(anchor); + builder.setInsertionPoint(anchor); + Location loc = anchor->getLoc(); + + auto getI64Operand = [&](unsigned idx) -> Value { + return castIntegerLikeTo(anchor, operands[idx], builder.getI64Type()); + }; + + Value nBurst = getI64Operand(3); + Value lenBurst = getI64Operand(4); + Value srcStride = getI64Operand(5); + Value dstStride = getI64Operand(6); + if (!nBurst || !lenBurst || !srcStride || !dstStride) { + return failure(); + } + + auto shl = [&](Value value, uint64_t amount) -> Value { + return builder.create(loc, value, getI64Constant(builder, loc, amount)); + }; + auto bitOr = [&](Value lhs, Value rhs) -> Value { return builder.create(loc, lhs, rhs); }; + + Value config = nBurst; + config = bitOr(config, shl(lenBurst, 16)); + config = bitOr(config, shl(srcStride, 32)); + config = bitOr(config, shl(dstStride, 48)); + return config; +} + +FailureOr packCopyCbufToUbConfig(Operation *anchor, ValueRange operands) { + if (operands.size() != 7) { + return failure(); + } + OpBuilder builder(anchor); + builder.setInsertionPoint(anchor); + Location loc = anchor->getLoc(); + + auto getI64Operand = [&](unsigned idx) -> Value { + return castIntegerLikeTo(anchor, operands[idx], builder.getI64Type()); + }; + + Value sid = getI64Operand(2); + Value nBurst = getI64Operand(3); + Value lenBurst = getI64Operand(4); + Value srcStride = getI64Operand(5); + Value dstStride = getI64Operand(6); + if (!sid || !nBurst || !lenBurst || !srcStride || !dstStride) { + return failure(); + } + + auto shl = [&](Value value, uint64_t amount) -> Value { + return builder.create(loc, value, getI64Constant(builder, loc, amount)); + }; + auto bitOr = [&](Value lhs, Value rhs) -> Value { return builder.create(loc, lhs, rhs); }; + + Value config = sid; + config = bitOr(config, shl(nBurst, 4)); + config = bitOr(config, shl(lenBurst, 16)); + config = bitOr(config, shl(srcStride, 32)); + config = bitOr(config, shl(dstStride, 48)); + return config; +} + +FailureOr packCopyUbToCbufConfig(Operation *anchor, ValueRange operands) { + if (operands.size() != 7) { + return failure(); + } + OpBuilder builder(anchor); + builder.setInsertionPoint(anchor); + Location loc = anchor->getLoc(); + + auto getI64Operand = [&](unsigned idx) -> Value { + return castIntegerLikeTo(anchor, operands[idx], builder.getI64Type()); + }; + + Value sid = getI64Operand(2); + Value nBurst = getI64Operand(3); + Value lenBurst = getI64Operand(4); + Value srcStride = getI64Operand(5); + Value dstStride = getI64Operand(6); + if (!sid || !nBurst || !lenBurst || !srcStride || !dstStride) { + return failure(); + } + + auto shl = [&](Value value, uint64_t amount) -> Value { + return builder.create(loc, value, getI64Constant(builder, loc, amount)); + }; + auto bitOr = [&](Value lhs, Value rhs) -> Value { return builder.create(loc, lhs, rhs); }; + + Value config = sid; + config = bitOr(config, shl(nBurst, 4)); + config = bitOr(config, shl(lenBurst, 16)); + config = bitOr(config, shl(srcStride, 32)); + config = bitOr(config, shl(dstStride, 48)); + return config; +} + +FailureOr packCopyGmToCbufConfig0(Operation *anchor, Value nBurst, Value lenBurst) { + OpBuilder builder(anchor); + builder.setInsertionPoint(anchor); + Location loc = anchor->getLoc(); + + Value nBurstI64 = castIntegerLikeTo(anchor, nBurst, builder.getI64Type()); + Value lenBurstI64 = castIntegerLikeTo(anchor, lenBurst, builder.getI64Type()); + if (!nBurstI64 || !lenBurstI64) { + return failure(); + } + + auto shl = [&](Value value, uint64_t amount) -> Value { + return builder.create(loc, value, getI64Constant(builder, loc, amount)); + }; + auto bitOr = [&](Value lhs, Value rhs) -> Value { return builder.create(loc, lhs, rhs); }; + + Value config0 = getI64Constant(builder, loc, 0); // sid + config0 = bitOr(config0, shl(nBurstI64, 4)); // burst_num[24:4] + config0 = bitOr(config0, shl(lenBurstI64, 25)); // burst_len[45:25] + return config0; +} + +FailureOr packCopyGmToCbufConfig1(Operation *anchor, Value srcStride, Value dstStride) { + OpBuilder builder(anchor); + builder.setInsertionPoint(anchor); + Location loc = anchor->getLoc(); + + Value srcStrideI64 = castIntegerLikeTo(anchor, srcStride, builder.getI64Type()); + Value dstStrideI64 = castIntegerLikeTo(anchor, dstStride, builder.getI64Type()); + if (!srcStrideI64 || !dstStrideI64) { + return failure(); + } + + auto shl = [&](Value value, uint64_t amount) -> Value { + return builder.create(loc, value, getI64Constant(builder, loc, amount)); + }; + auto bitOr = [&](Value lhs, Value rhs) -> Value { return builder.create(loc, lhs, rhs); }; + + // config1 packs burst_src_stride[39:0] and burst_dst_stride[60:40]. + return bitOr(srcStrideI64, shl(dstStrideI64, 40)); +} + +FailureOr packCopyGmToCbufMultiConfig0(Operation *anchor, Value sid, Value loop1SrcStride, Value l2CacheCtl, + Value nValue) { + OpBuilder builder(anchor); + builder.setInsertionPoint(anchor); + Location loc = anchor->getLoc(); + + Value sidI64 = castIntegerLikeTo(anchor, sid, builder.getI64Type()); + Value loop1SrcStrideI64 = castIntegerLikeTo(anchor, loop1SrcStride, builder.getI64Type()); + Value l2CacheCtlI64 = castIntegerLikeTo(anchor, l2CacheCtl, builder.getI64Type()); + Value nValueI64 = castIntegerLikeTo(anchor, nValue, builder.getI64Type()); + if (!sidI64 || !loop1SrcStrideI64 || !l2CacheCtlI64 || !nValueI64) { + return failure(); + } + + auto shl = [&](Value value, uint64_t amount) -> Value { + return builder.create(loc, value, getI64Constant(builder, loc, amount)); + }; + auto bitOr = [&](Value lhs, Value rhs) -> Value { return builder.create(loc, lhs, rhs); }; + + Value config0 = sidI64; + config0 = bitOr(config0, shl(loop1SrcStrideI64, 4)); + config0 = bitOr(config0, shl(l2CacheCtlI64, 44)); + config0 = bitOr(config0, shl(nValueI64, 48)); + return config0; +} + +FailureOr packCopyGmToCbufMultiConfig1(Operation *anchor, Value dValue, Value loop4SrcStride, Value smallC0En) { + OpBuilder builder(anchor); + builder.setInsertionPoint(anchor); + Location loc = anchor->getLoc(); + + Value dValueI64 = castIntegerLikeTo(anchor, dValue, builder.getI64Type()); + Value loop4SrcStrideI64 = castIntegerLikeTo(anchor, loop4SrcStride, builder.getI64Type()); + Value smallC0EnI64 = castIntegerLikeTo(anchor, smallC0En, builder.getI64Type()); + if (!dValueI64 || !loop4SrcStrideI64 || !smallC0EnI64) { + return failure(); + } + + auto shl = [&](Value value, uint64_t amount) -> Value { + return builder.create(loc, value, getI64Constant(builder, loc, amount)); + }; + auto bitOr = [&](Value lhs, Value rhs) -> Value { return builder.create(loc, lhs, rhs); }; + + Value config1 = dValueI64; + config1 = bitOr(config1, shl(loop4SrcStrideI64, 21)); + config1 = bitOr(config1, shl(smallC0EnI64, 61)); + return config1; +} + +FailureOr packCopyCbufToBtConfig(Operation *anchor, Value convControl, Value nBurst, Value lenBurst, + Value sourceGap, Value dstGap) { + OpBuilder builder(anchor); + builder.setInsertionPoint(anchor); + Value zero = getI64Constant(builder, anchor->getLoc(), 0); + SmallVector, 5> fields = { + {convControl, 3}, {nBurst, 4}, {lenBurst, 16}, {sourceGap, 32}, {dstGap, 48}}; + return packShiftedFields(anchor, zero, fields); +} + +FailureOr packCopyCbufToFbufConfig(Operation *anchor, Value nBurst, Value lenBurst, Value sourceGap, + Value dstGap) { + OpBuilder builder(anchor); + builder.setInsertionPoint(anchor); + Location loc = anchor->getLoc(); + + Value nBurstI64 = castIntegerLikeTo(anchor, nBurst, builder.getI64Type()); + Value lenBurstI64 = castIntegerLikeTo(anchor, lenBurst, builder.getI64Type()); + Value sourceGapI64 = castIntegerLikeTo(anchor, sourceGap, builder.getI64Type()); + Value dstGapI64 = castIntegerLikeTo(anchor, dstGap, builder.getI64Type()); + if (!nBurstI64 || !lenBurstI64 || !sourceGapI64 || !dstGapI64) { + return failure(); + } + + auto shl = [&](Value value, uint64_t amount) -> Value { + return builder.create(loc, value, getI64Constant(builder, loc, amount)); + }; + auto bitOr = [&](Value lhs, Value rhs) -> Value { return builder.create(loc, lhs, rhs); }; + + Value config = shl(nBurstI64, 4); + config = bitOr(config, shl(lenBurstI64, 16)); + config = bitOr(config, shl(sourceGapI64, 32)); + config = bitOr(config, shl(dstGapI64, 48)); + return config; +} + +FailureOr packLoadCbufToS4Config0(Operation *anchor, Value mStart, Value kStart, Value mStep, Value kStep) { + OpBuilder builder(anchor); + builder.setInsertionPoint(anchor); + Location loc = anchor->getLoc(); + + Value mStartI64 = castIntegerLikeTo(anchor, mStart, builder.getI64Type()); + Value kStartI64 = castIntegerLikeTo(anchor, kStart, builder.getI64Type()); + Value mStepI64 = castIntegerLikeTo(anchor, mStep, builder.getI64Type()); + Value kStepI64 = castIntegerLikeTo(anchor, kStep, builder.getI64Type()); + if (!mStartI64 || !kStartI64 || !mStepI64 || !kStepI64) { + return failure(); + } + + auto shl = [&](Value value, uint64_t amount) -> Value { + return builder.create(loc, value, getI64Constant(builder, loc, amount)); + }; + auto bitOr = [&](Value lhs, Value rhs) -> Value { return builder.create(loc, lhs, rhs); }; + + Value config0 = mStartI64; + config0 = bitOr(config0, shl(kStartI64, 16)); + config0 = bitOr(config0, shl(mStepI64, 32)); + config0 = bitOr(config0, shl(kStepI64, 40)); + return config0; +} + +FailureOr packLoadCbufToS4Config1(Operation *anchor, Value srcStride, Value dstStride) { + OpBuilder builder(anchor); + builder.setInsertionPoint(anchor); + Location loc = anchor->getLoc(); + + Value srcStrideI64 = castIntegerLikeTo(anchor, srcStride, builder.getI64Type()); + Value dstStrideI64 = castIntegerLikeTo(anchor, dstStride, builder.getI64Type()); + if (!srcStrideI64 || !dstStrideI64) { + return failure(); + } + + auto shl = [&](Value value, uint64_t amount) -> Value { + return builder.create(loc, value, getI64Constant(builder, loc, amount)); + }; + return builder.create(loc, srcStrideI64, shl(dstStrideI64, 16)).getResult(); +} + +FailureOr packLoadCbufToCaConfig0(Operation *anchor, Value mStart, Value kStart, Value mStep, Value kStep) { + OpBuilder builder(anchor); + builder.setInsertionPoint(anchor); + Location loc = anchor->getLoc(); + + Value mStartI64 = castIntegerLikeTo(anchor, mStart, builder.getI64Type()); + Value kStartI64 = castIntegerLikeTo(anchor, kStart, builder.getI64Type()); + Value mStepI64 = castIntegerLikeTo(anchor, mStep, builder.getI64Type()); + Value kStepI64 = castIntegerLikeTo(anchor, kStep, builder.getI64Type()); + if (!mStartI64 || !kStartI64 || !mStepI64 || !kStepI64) { + return failure(); + } + + auto shl = [&](Value value, uint64_t amount) -> Value { + return builder.create(loc, value, getI64Constant(builder, loc, amount)); + }; + auto bitOr = [&](Value lhs, Value rhs) -> Value { return builder.create(loc, lhs, rhs); }; + + Value config0 = mStartI64; + config0 = bitOr(config0, shl(kStartI64, 16)); + config0 = bitOr(config0, shl(mStepI64, 32)); + config0 = bitOr(config0, shl(kStepI64, 40)); + return config0; +} + +FailureOr packLoadCbufToCaConfig1(Operation *anchor, Value srcStride, Value dstStride) { + OpBuilder builder(anchor); + builder.setInsertionPoint(anchor); + Location loc = anchor->getLoc(); + + Value srcStrideI64 = castIntegerLikeTo(anchor, srcStride, builder.getI64Type()); + Value dstStrideI64 = castIntegerLikeTo(anchor, dstStride, builder.getI64Type()); + if (!srcStrideI64 || !dstStrideI64) { + return failure(); + } + + auto shl = [&](Value value, uint64_t amount) -> Value { + return builder.create(loc, value, getI64Constant(builder, loc, amount)); + }; + return builder.create(loc, srcStrideI64, shl(dstStrideI64, 16)).getResult(); +} + +FailureOr packLoadCbufToCbConfig0(Operation *anchor, Value mStart, Value kStart, Value mStep, Value kStep) { + OpBuilder builder(anchor); + builder.setInsertionPoint(anchor); + Location loc = anchor->getLoc(); + + Value mStartI64 = castIntegerLikeTo(anchor, mStart, builder.getI64Type()); + Value kStartI64 = castIntegerLikeTo(anchor, kStart, builder.getI64Type()); + Value mStepI64 = castIntegerLikeTo(anchor, mStep, builder.getI64Type()); + Value kStepI64 = castIntegerLikeTo(anchor, kStep, builder.getI64Type()); + if (!mStartI64 || !kStartI64 || !mStepI64 || !kStepI64) { + return failure(); + } + + auto shl = [&](Value value, uint64_t amount) -> Value { + return builder.create(loc, value, getI64Constant(builder, loc, amount)); + }; + auto bitOr = [&](Value lhs, Value rhs) -> Value { return builder.create(loc, lhs, rhs); }; + + Value config0 = mStartI64; + config0 = bitOr(config0, shl(kStartI64, 16)); + config0 = bitOr(config0, shl(mStepI64, 32)); + config0 = bitOr(config0, shl(kStepI64, 40)); + return config0; +} + +FailureOr packLoadCbufToCbConfig1(Operation *anchor, Value srcStride, Value dstStride) { + OpBuilder builder(anchor); + builder.setInsertionPoint(anchor); + Location loc = anchor->getLoc(); + + Value srcStrideI64 = castIntegerLikeTo(anchor, srcStride, builder.getI64Type()); + Value dstStrideI64 = castIntegerLikeTo(anchor, dstStride, builder.getI64Type()); + if (!srcStrideI64 || !dstStrideI64) { + return failure(); + } + + auto shl = [&](Value value, uint64_t amount) -> Value { + return builder.create(loc, value, getI64Constant(builder, loc, amount)); + }; + return builder.create(loc, srcStrideI64, shl(dstStrideI64, 16)).getResult(); +} + +Value buildMadBiasDestination(Operation *anchor, ConversionPatternRewriter &rewriter, Value dst, Value bias) { + Type i64Ty = rewriter.getI64Type(); + Value dstAddr = rewriter.create(anchor->getLoc(), i64Ty, dst); + Value biasAddr = rewriter.create(anchor->getLoc(), i64Ty, bias); + Value lowMask = getI64Constant(rewriter, anchor->getLoc(), 0xffffffffULL); + Value dstLow = rewriter.create(anchor->getLoc(), dstAddr, lowMask); + Value biasLow = rewriter.create(anchor->getLoc(), biasAddr, lowMask); + Value biasHigh = + rewriter.create(anchor->getLoc(), biasLow, getI64Constant(rewriter, anchor->getLoc(), 32)); + Value packed = rewriter.create(anchor->getLoc(), dstLow, biasHigh); + return rewriter.create(anchor->getLoc(), dst.getType(), packed); +} + +FailureOr packVbitsortConfig(Operation *anchor, Value repeatTimes) { + OpBuilder builder(anchor); + builder.setInsertionPoint(anchor); + Location loc = anchor->getLoc(); + + Value repeatI64 = castIntegerLikeTo(anchor, repeatTimes, builder.getI64Type()); + if (!repeatI64) { + return failure(); + } + return builder.create(loc, repeatI64, getI64Constant(builder, loc, 56)).getResult(); +} + +[[maybe_unused]] FailureOr materializeDynamicPltMask(ConversionPatternRewriter &rewriter, LoweringState &state, + Location loc, Value laneCount, Type vectorElemType) { + Type i32Type = rewriter.getI32Type(); + Value laneCountI32 = laneCount; + if (laneCountI32.getType() != i32Type) { + laneCountI32 = castIntegerLikeTo(rewriter.getInsertionBlock()->getParentOp(), laneCountI32, i32Type); + if (!laneCountI32) { + return failure(); + } + } + + StringRef calleeName; + if (vectorElemType.isF32()) { + calleeName = StringRef("llvm.hivm.plt.b32.v300"); + } else if (vectorElemType.isF16() || vectorElemType.isBF16()) { + calleeName = StringRef("llvm.hivm.plt.b16.v300"); + } else if (auto intType = dyn_cast(vectorElemType)) { + if (intType.getWidth() == 32) { + calleeName = StringRef("llvm.hivm.plt.b32.v300"); + } else if (intType.getWidth() == 16) { + calleeName = StringRef("llvm.hivm.plt.b16.v300"); + } else if (intType.getWidth() == 8) { + calleeName = StringRef("llvm.hivm.plt.b8.v300"); + } + } + if (calleeName.empty()) { + return failure(); + } + + Type maskType = VectorType::get({256}, rewriter.getI1Type()); + auto funcType = rewriter.getFunctionType(TypeRange{i32Type}, TypeRange{maskType, i32Type}); + auto call = rewriter.create(loc, calleeName, funcType.getResults(), ValueRange{laneCountI32}); + state.plannedDecls.push_back(PlannedDecl{calleeName.str(), funcType}); + return call.getResult(0); +} + +} // namespace mlir::pto::detail diff --git a/lib/PTO/Transforms/VPTOCANN900LLVMEmitterPipeline.cpp b/lib/PTO/Transforms/VPTOCANN900LLVMEmitterPipeline.cpp new file mode 100644 index 0000000000..2f7787a557 --- /dev/null +++ b/lib/PTO/Transforms/VPTOCANN900LLVMEmitterPipeline.cpp @@ -0,0 +1,580 @@ +// Copyright (c) 2026 Huawei Technologies Co., Ltd. +// This program is free software, you can redistribute it and/or modify it under the terms and conditions of +// CANN Open Software License Agreement Version 2.0 (the "License"). +// Please refer to the License for details. You may not use this file except in compliance with the License. +// THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED, +// INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE. +// See LICENSE in the root of the software repository for the full text of the License. + +#include "VPTOCANN900LLVMEmitterInternal.h" + +namespace mlir::pto::detail { + +void populateVPTOOpLoweringPatterns(VPTOTypeConverter &typeConverter, RewritePatternSet &patterns, + LoweringState &state) { + populateVPTOArithmeticPatterns(typeConverter, patterns, state); + populateVPTOVectorMemoryPatterns(typeConverter, patterns, state); + populateVPTOScalarPatterns(typeConverter, patterns, state); +} + +void markIllegalVPTOSyncOps(ConversionTarget &target) { + target.addIllegalOp(); +} + +void markIllegalVPTOSimtOps(ConversionTarget &target) { + target.addIllegalOp< + pto::GetBlockIdxOp, pto::GetSubBlockIdxOp, pto::GetBlockNumOp, pto::GetSubBlockNumOp, pto::GetCtrlOp, + pto::GetVms4SrOp, pto::GetTidXOp, pto::GetTidYOp, pto::GetTidZOp, pto::GetBlockDimXOp, pto::GetBlockDimYOp, + pto::GetBlockDimZOp, pto::GetGridDimXOp, pto::GetGridDimYOp, pto::GetGridDimZOp, pto::GetBlockIdxXOp, + pto::GetBlockIdxYOp, pto::GetBlockIdxZOp, pto::GetVecCoreIdOp, pto::GetLaneIdOp, pto::GetClock32Op, + pto::GetClock64Op, pto::GetLaneMaskEqOp, pto::GetLaneMaskLeOp, pto::GetLaneMaskLtOp, pto::GetLaneMaskGeOp, + pto::GetLaneMaskGtOp, pto::VoteAllOp, pto::VoteAnyOp, pto::VoteUniOp, pto::VoteBallotOp, pto::ShuffleIdxOp, + pto::ShuffleUpOp, pto::ShuffleDownOp, pto::ShuffleBflyOp, pto::ReduxAddOp, pto::ReduxMaxOp, pto::ReduxMinOp, + pto::AtomicCasOp, pto::AtomicExchOp, pto::AtomicAddOp, pto::AtomicSubOp, pto::AtomicMinOp, pto::AtomicMaxOp, + pto::AtomicAndOp, pto::AtomicOrOp, pto::AtomicXorOp, pto::TrapOp, pto::PrmtOp, pto::MulhiOp, pto::MulI32ToI64Op, + pto::SqrtOp, pto::AbsFOp, pto::ExpOp, pto::LogOp, pto::CeilOp, pto::FloorOp, pto::RintOp, pto::RoundOp, + pto::FMinOp, pto::FMaxOp, pto::PowOp, pto::FmaOp, pto::ConvertOp, pto::SyncthreadsOp, pto::ThreadfenceOp, + pto::ThreadfenceBlockOp, pto::KeepOp, pto::ResumeOp>(); +} + +void markIllegalVPTOConfigOps(ConversionTarget &target) { + target.addIllegalOp(); + target.addIllegalOp(); +} + +void markIllegalVPTOMemoryOps(ConversionTarget &target) { + target.addIllegalOp(); +} + +void markIllegalVPTOPredicateOps(ConversionTarget &target) { + target.addIllegalOp(); +} + +void markIllegalVPTOArithmeticAndCopyOps(ConversionTarget &target) { + target.addIllegalOp< + pto::VabsOp, pto::VexpOp, pto::VlnOp, pto::VnegOp, pto::VsqrtOp, pto::VreluOp, pto::VnotOp, pto::VsqzOp, + pto::VusqzOp, pto::VmulaOp, pto::VmullOp, pto::VaddOp, pto::VsubOp, pto::VmulOp, pto::VdivOp, pto::VmaxOp, + pto::VminOp, pto::VandOp, pto::VorOp, pto::VxorOp, pto::VmaddOp, pto::VaddcOp, pto::VsubcOp, pto::VaddcsOp, + pto::VsubcsOp, pto::VshlOp, pto::VshrOp, pto::VmulsOp, pto::VaddsOp, pto::VmaxsOp, pto::VminsOp, pto::VlreluOp, + pto::VshlsOp, pto::VshrsOp, pto::VcaddOp, pto::VcmaxOp, pto::VcminOp, pto::VcgaddOp, pto::VcgmaxOp, pto::VcgminOp, + pto::VcpaddOp, pto::Chistv2Op, pto::Dhistv2Op, pto::VcbmaxOp, pto::VcbminOp, pto::VdupOp, pto::VbrOp, + pto::PpackOp, pto::PunpackOp, pto::PbitcastOp, pto::VselOp, pto::VselrOp, pto::PnotOp, pto::PselOp, pto::PandOp, + pto::PorOp, pto::PxorOp, pto::PdintlvB8Op, pto::PdintlvB16Op, pto::PdintlvB32Op, pto::PintlvB8Op, + pto::PintlvB16Op, pto::PintlvB32Op, pto::VsunpackOp, pto::VzunpackOp, pto::VpackOp, pto::VintlvOp, pto::VdintlvOp, + pto::VpreluOp, pto::VaxpyOp, pto::VmulscvtOp, pto::VciOp, pto::VexpdifOp, pto::VbitsortOp, pto::Vmrgsort4Op, + pto::VtrcOp, pto::VcvtOp, pto::VbitcastOp, pto::VcmpOp, pto::VcmpsOp, pto::CopyGmToUbufOp, pto::CopyUbufToGmOp, + pto::CopyUbufToUbufOp, pto::CopyCbufToUbufOp, pto::CopyUbufToCbufOp, pto::CopyGmToCbufOp, pto::CreateCbufMatrixOp, + pto::LoadCbufToCaOp, pto::LoadCbufToCbOp, pto::LoadCbufToCaS4Op, pto::LoadCbufToCbS4Op, pto::LoadCbufToCaMxOp, + pto::LoadCbufToCbMxOp, pto::CopyMatrixCcToGmOp, pto::CopyMatrixCcToCbufOp, pto::CopyMatrixCcToUbOp, + pto::CopyCbufToBtOp, pto::CopyCbufToFbufOp, pto::CopyGmToCbufMultiNd2NzOp, pto::CopyGmToCbufMultiDn2NzOp, + pto::MadOp, pto::MadAccOp, pto::MadBiasOp, pto::MadMxOp, pto::MadMxAccOp, pto::MadMxBiasOp, pto::MadRawOp, + pto::MadBiasRawOp, pto::MadMxRawOp, pto::MadMxBiasRawOp>(); +} + +void configureVPTOOpLoweringTarget(ConversionTarget &target, VPTOTypeConverter &typeConverter) { + (void)typeConverter; + target.addLegalOp(); + target.addLegalDialect(); + target.addLegalOp(); + markIllegalVPTOSyncOps(target); + markIllegalVPTOSimtOps(target); + markIllegalVPTOConfigOps(target); + markIllegalVPTOMemoryOps(target); + markIllegalVPTOPredicateOps(target); + markIllegalVPTOArithmeticAndCopyOps(target); + target.markUnknownOpDynamicallyLegal([](Operation *op) { return !isa(op); }); +} + +void populateVPTOStructuralTypePatterns(VPTOTypeConverter &typeConverter, RewritePatternSet &patterns, + ConversionTarget &target) { + scf::populateSCFStructuralTypeConversionsAndLegality(typeConverter, patterns, target); + populateFunctionOpInterfaceTypeConversionPattern(patterns, typeConverter); + populateCallOpTypeConversionPattern(patterns, typeConverter); + populateBranchOpInterfaceTypeConversionPattern(patterns, typeConverter); + populateReturnOpTypeConversionPattern(patterns, typeConverter); +} + +void configureVPTOTypeLoweringTarget(ConversionTarget &target, VPTOTypeConverter &typeConverter) { + target.addLegalOp(); + target.addDynamicallyLegalOp([&](func::FuncOp op) { + return typeConverter.isSignatureLegal(op.getFunctionType()) && typeConverter.isLegal(&op.getBody()); + }); + target.addDynamicallyLegalOp([&](Operation *op) { return typeConverter.isLegal(op); }); + target.addDynamicallyLegalOp( + [&](Operation *op) { return isLegalForBranchOpInterfaceTypeConversionPattern(op, typeConverter); }); + target.addDynamicallyLegalOp([&](arith::SelectOp op) { + return typeConverter.isLegal(op->getOperandTypes()) && typeConverter.isLegal(op->getResultTypes()); + }); + target.addIllegalOp(); +} + +void configureVPTOCarrierTypeLegality(ConversionTarget &target, VPTOTypeConverter &typeConverter) { + target.addDynamicallyLegalOp([&](UnrealizedConversionCastOp op) { + return !hasVPTOConvertibleType(op->getOperandTypes()) && !hasVPTOConvertibleType(op->getResultTypes()); + }); + target.addDynamicallyLegalOp([&](LLVM::AllocaOp op) { + return typeConverter.isLegal(op->getOperandTypes()) && typeConverter.isLegal(op->getResultTypes()) && + typeConverter.isLegal(op.getElemType()); + }); + target.addDynamicallyLegalOp([&](LLVM::GEPOp op) { + return typeConverter.isLegal(op->getOperandTypes()) && typeConverter.isLegal(op->getResultTypes()) && + typeConverter.isLegal(op.getElemType()); + }); + target.markUnknownOpDynamicallyLegal([&](Operation *op) { + return typeConverter.isLegal(op->getOperandTypes()) && typeConverter.isLegal(op->getResultTypes()); + }); +} +void foldVPTOTypeCasts(ModuleOp module, TypeConverter &typeConverter) { + SmallVector castsToFold; + module.walk([&](UnrealizedConversionCastOp castOp) { + if (castOp->getNumOperands() != 1 || castOp->getNumResults() != 1) { + return; + } + if (!hasVPTOConvertibleType(castOp->getOperandTypes()) && !hasVPTOConvertibleType(castOp->getResultTypes())) { + return; + } + Type convertedResultType = typeConverter.convertType(castOp.getResult(0).getType()); + if (convertedResultType && convertedResultType == castOp.getOperand(0).getType()) { + castsToFold.push_back(castOp); + } + }); + for (UnrealizedConversionCastOp castOp : castsToFold) { + castOp.getResult(0).replaceAllUsesWith(castOp.getOperand(0)); + castOp.erase(); + } +} + +LogicalResult lowerVPTOOps(ModuleOp module, llvm::raw_ostream &diagOS) { + MLIRContext *context = module.getContext(); + VPTOTypeConverter typeConverter(context); + ConversionTarget target(*context); + RewritePatternSet patterns(context); + LoweringState state; + + configureVPTOOpLoweringTarget(target, typeConverter); + populateVPTOOpLoweringPatterns(typeConverter, patterns, state); + + if (failed(applyPartialConversion(module, target, std::move(patterns)))) { + diagOS << "VPTO LLVM emission failed: VPTO op lowering failed\n"; + return failure(); + } + if (failed(materializeDecls(module, state.plannedDecls, diagOS))) { + return failure(); + } + return success(); +} + +LogicalResult lowerVPTOTypes(ModuleOp module, llvm::raw_ostream &diagOS) { + MLIRContext *context = module.getContext(); + VPTOTypeConverter typeConverter(context); + ConversionTarget target(*context); + RewritePatternSet patterns(context); + LoweringState state; + + configureVPTOTypeLoweringTarget(target, typeConverter); + configureVPTOCarrierTypeLegality(target, typeConverter); + populateVPTOStructuralTypePatterns(typeConverter, patterns, target); + populateVPTOTypePatterns(typeConverter, patterns, target, state); + + if (failed(applyPartialConversion(module, target, std::move(patterns)))) { + diagOS << "VPTO LLVM emission failed: VPTO type lowering failed\n"; + return failure(); + } + if (failed(materializeDecls(module, state.plannedDecls, diagOS))) { + return failure(); + } + foldVPTOTypeCasts(module, typeConverter); + return success(); +} + +Type normalizeTypeForOfficialLLVMLowering(Type type, Builder &builder) { + type = convertVPTOType(type, builder); + return type; +} + +void normalizeFuncSignaturesForOfficialLLVMLowering(ModuleOp module) { + Builder builder(module.getContext()); + + for (func::FuncOp funcOp : module.getOps()) { + FunctionType oldType = funcOp.getFunctionType(); + SmallVector newInputs; + SmallVector newResults; + bool changed = false; + + for (Type input : oldType.getInputs()) { + Type normalized = normalizeTypeForOfficialLLVMLowering(input, builder); + changed |= (normalized != input); + newInputs.push_back(normalized); + } + for (Type result : oldType.getResults()) { + Type normalized = normalizeTypeForOfficialLLVMLowering(result, builder); + changed |= (normalized != result); + newResults.push_back(normalized); + } + + if (!changed) { + continue; + } + + auto newType = builder.getFunctionType(newInputs, newResults); + funcOp.setFunctionTypeAttr(TypeAttr::get(newType)); + + if (funcOp.isExternal()) { + continue; + } + Block &entry = funcOp.getBody().front(); + for (auto [arg, newType] : llvm::zip(entry.getArguments(), newInputs)) { + if (arg.getType() != newType) { + arg.setType(newType); + } + } + } +} + +void forceV300CtrlModeForVPTOFuncs(ModuleOp module) { + OpBuilder builder(module.getContext()); + + for (func::FuncOp funcOp : module.getOps()) { + if (!needsV300CtrlModeForVPTOFunc(funcOp)) { + continue; + } + + Block &entry = funcOp.getBody().front(); + builder.setInsertionPointToStart(&entry); + auto i64Type = builder.getI64Type(); + auto bit60 = builder.create(funcOp.getLoc(), i64Type, builder.getI64IntegerAttr(60)); + Value ctrl = builder.create(funcOp.getLoc(), i64Type).getResult(); + Value ctrlV300 = builder.create(funcOp.getLoc(), i64Type, ctrl, bit60.getResult()).getResult(); + builder.create(funcOp.getLoc(), ctrlV300); + } +} + +std::optional getKernelKind(ModuleOp module) { + auto kernelKind = module->getAttrOfType(FunctionKernelKindAttr::name); + if (!kernelKind) { + return std::nullopt; + } + return kernelKind.getKernelKind(); +} + +VPTOEmissionOptions makeDeviceEmissionOptions(const VPTOEmissionOptions &baseOptions, FunctionKernelKind kind) { + VPTOEmissionOptions options = baseOptions; + constexpr llvm::StringLiteral kVecTargetFeatures = + "+ATOMIC,+ArchV130,+AregRedefinable,+ArithmeticBf16,+AtomicForB8 ," + "+F8e4m3,+F8e5m2,+F8e8m0,+FFTSBlk,+Fp4e1m2x2,+Fp4e2m1x2,+LDExtRefine," + "+MOVX8,+SPR7bits,+SyncV,+dav-c310-vec"; + constexpr llvm::StringLiteral kCubeTargetFeatures = + "+ATOMIC,+ArchV130,+AregRedefinable,+ArithmeticBf16,+AtomicForB8 ," + "+F8e4m3,+F8e5m2,+F8e8m0,+FFTSBlk,+Fp4e1m2x2,+Fp4e2m1x2,+LDExtRefine," + "+MOVX8,+SPR7bits,+SyncV,+dav-c310-cube"; + if (kind == FunctionKernelKind::Vector) { + options.march = "dav-c310-vec"; + options.aicoreArch = "dav-c310-vec"; + options.defaultTargetCPU = "dav-c310-vec"; + options.defaultTargetFeatures = kVecTargetFeatures.str(); + } else if (kind == FunctionKernelKind::Cube) { + options.march = "dav-c310-cube"; + options.aicoreArch = "dav-c310-cube"; + options.defaultTargetCPU = "dav-c310-cube"; + options.defaultTargetFeatures = kCubeTargetFeatures.str(); + } + return options; +} + +FailureOr getUniqueDeviceModuleByKernelKind(ModuleOp module, FunctionKernelKind kind, + llvm::raw_ostream &diagOS) { + ModuleOp matched; + for (ModuleOp child : module.getOps()) { + auto kernelKind = getKernelKind(child); + if (!kernelKind) { + continue; + } + if (*kernelKind != kind) { + continue; + } + if (matched) { + diagOS << "VPTO LLVM emission failed: duplicate device module with " << FunctionKernelKindAttr::name << "\n"; + return failure(); + } + matched = child; + } + return matched; +} + +LogicalResult renameKernelFunctionsForKernelKind(ModuleOp module, llvm::raw_ostream &diagOS) { + auto kernelKind = getKernelKind(module); + if (!kernelKind) { + diagOS << "VPTO LLVM emission failed: device module missing " << FunctionKernelKindAttr::name << "\n"; + return failure(); + } + + StringRef suffix; + if (*kernelKind == FunctionKernelKind::Vector) { + suffix = kVectorSuffix; + } else if (*kernelKind == FunctionKernelKind::Cube) { + suffix = kCubeSuffix; + } else { + diagOS << "VPTO LLVM emission failed: unsupported " << FunctionKernelKindAttr::name << "\n"; + return failure(); + } + + for (func::FuncOp funcOp : module.getOps()) { + if (!pto::hasExplicitPTOEntryAttr(funcOp)) { + continue; + } + if (funcOp.getSymName().ends_with(suffix)) { + continue; + } + funcOp.setSymName((funcOp.getSymName() + suffix).str()); + } + return success(); +} + +struct LowerVPTOOpsPass final : public PassWrapper> { + MLIR_DEFINE_EXPLICIT_INTERNAL_INLINE_TYPE_ID(LowerVPTOOpsPass) + + void runOnOperation() override { + materializeVecScopeCarrierLoops(getOperation()); + if (failed(lowerVPTOOps(getOperation(), llvm::errs()))) { + signalPassFailure(); + } + } +}; + +struct LowerVPTOTypesPass final : public PassWrapper> { + MLIR_DEFINE_EXPLICIT_INTERNAL_INLINE_TYPE_ID(LowerVPTOTypesPass) + + void runOnOperation() override { + if (failed(lowerVPTOTypes(getOperation(), llvm::errs()))) { + signalPassFailure(); + } + } +}; + +struct NormalizeFuncSignaturesForLLVMLoweringPass final + : public PassWrapper> { + MLIR_DEFINE_EXPLICIT_INTERNAL_INLINE_TYPE_ID(NormalizeFuncSignaturesForLLVMLoweringPass) + + void runOnOperation() override { normalizeFuncSignaturesForOfficialLLVMLowering(getOperation()); } +}; + +struct PrepareVPTOLLVMLoweringPass final : public PassWrapper> { + MLIR_DEFINE_EXPLICIT_INTERNAL_INLINE_TYPE_ID(PrepareVPTOLLVMLoweringPass) + + void runOnOperation() override { + ModuleOp module = getOperation(); + pto::annotatePTOEntryFunctions(module); + forceV300CtrlModeForVPTOFuncs(module); + if (failed(renameKernelFunctionsForKernelKind(module, llvm::errs()))) { + signalPassFailure(); + } + } +}; + +llvm::StringSet collectSimtEntryFunctionNames(ModuleOp module) { + llvm::StringSet simtEntries; + module.walk([&](func::FuncOp funcOp) { + if (funcOp->hasAttr(pto::kPTOSimtEntryAttrName)) { + simtEntries.insert(funcOp.getSymName()); + } + }); + return simtEntries; +} + +void applyArtifactVisibilityLinkage(ModuleOp sourceModule, llvm::Module &llvmModule) { + llvm::StringMap externalByName; + sourceModule.walk([&](func::FuncOp funcOp) { + if (funcOp.isDeclaration()) { + return; + } + externalByName[funcOp.getSymName()] = pto::hasExternalArtifactVisibility(funcOp); + }); + + for (llvm::Function &function : llvmModule) { + auto it = externalByName.find(function.getName()); + if (it == externalByName.end()) { + continue; + } + if (it->second) { + function.setLinkage(llvm::GlobalValue::ExternalLinkage); + continue; + } + function.setLinkage(llvm::GlobalValue::InternalLinkage); + } +} + +void applySimtEntryCallingConvention(llvm::Module &llvmModule, + const llvm::StringSet &simtEntryNames) { + for (llvm::Function &function : llvmModule) { + if (simtEntryNames.contains(function.getName())) { + function.setCallingConv(llvm::CallingConv::SimtEntry); + function.addFnAttr(llvm::Attribute::NoInline); + // Match Bisheng's C++ frontend shape for SIMT outlined bodies. The + // exported wrapper owns the real kernel metadata, while the SIMT body is + // an ODR helper called with the SIMT calling convention. In CANN beta.1, + // leaving the SIMT body as a strong GLOBAL FUNC makes the runtime count it + // as an extra kernel without matching .ascend.meta, which can corrupt the + // selected kernel metadata. linkonce_odr lowers to a weak helper symbol + // and avoids that beta.1 metadata mismatch. + function.setLinkage(llvm::GlobalValue::LinkOnceODRLinkage); + } + } + + for (llvm::Function &function : llvmModule) { + for (llvm::BasicBlock &block : function) { + for (llvm::Instruction &inst : block) { + auto *call = llvm::dyn_cast(&inst); + if (!call) { + continue; + } + auto *callee = call->getCalledFunction(); + if (!callee || !simtEntryNames.contains(callee->getName())) { + continue; + } + call->setCallingConv(llvm::CallingConv::SimtEntry); + } + } + } +} + +FailureOr emitDeviceLLVMModule(ModuleOp deviceModule, StringRef kernelKind, + const VPTOEmissionOptions &options, + const llvm::StringSet &simtEntryNames, + llvm::raw_ostream &diagOS) { + if (!deviceModule) { + return EmittedLLVMModule{}; + } + if (failed(applyQueriedTargetAttrs(deviceModule, options, diagOS))) { + return failure(); + } + + auto llvmContext = std::make_unique(); + registerBuiltinDialectTranslation(*deviceModule.getContext()); + registerLLVMDialectTranslation(*deviceModule.getContext()); + std::unique_ptr llvmModule = translateModuleToLLVMIR(deviceModule.getOperation(), *llvmContext); + if (!llvmModule) { + diagOS << "VPTO LLVM emission failed: LLVM IR export failed for " << kernelKind << " module\n"; + return failure(); + } + + applyArtifactVisibilityLinkage(deviceModule, *llvmModule); + for (llvm::Function &func : *llvmModule) { + if (!func.getName().starts_with("llvm.hivm.vscatter.")) { + continue; + } + // Work around a bug in older Bisheng releases: vscatter was not modeled + // as writing through its destination pointer, so EarlyCSE could eliminate + // a load after vscatter as redundant. + func.setOnlyAccessesArgMemory(); + func.addFnAttr(llvm::Attribute::NoUnwind); + func.addFnAttr(llvm::Attribute::WriteOnly); + } + applySimtEntryCallingConvention(*llvmModule, simtEntryNames); + if (failed(attachAIVectorScopeMetadata(*llvmModule, diagOS))) { + return failure(); + } + attachHIVMKernelAnnotations(*llvmModule, deviceModule); + llvmModule->setModuleIdentifier(("ptoas.hivm.official." + kernelKind).str()); + llvmModule->setSourceFileName(("ptoas.hivm.official." + kernelKind).str()); + return EmittedLLVMModule{std::move(llvmContext), std::move(llvmModule)}; +} + +template LogicalResult runPipeline(ModuleOp module, llvm::raw_ostream &diagOS, EmitFn &&emit) { + OwningOpRef clonedOp(module->clone()); + ModuleOp clonedModule = cast(*clonedOp); + + if (failed(validateVPTOAuthoringIR(clonedModule, &diagOS))) { + diagOS << "VPTO LLVM emission failed: authoring-stage VPTO legality " + "validation failed\n"; + return failure(); + } + + PassManager pm(clonedModule.getContext()); + pm.enableVerifier(); + auto &kernelModulePM = pm.nest(); + kernelModulePM.addPass(std::make_unique()); + kernelModulePM.addPass(std::make_unique()); + kernelModulePM.addPass(std::make_unique()); + kernelModulePM.addPass(std::make_unique()); + kernelModulePM.addPass(arith::createArithExpandOpsPass()); + // pto-convert-scf-to-cf-with-loop-hints performs the SCF-to-CF conversion for this pipeline: + // it runs the upstream conversion patterns plus a higher-benefit lowering + // for {pto.unroll = "enable"} loops that attaches llvm.loop_annotation to + // the latch, so the !llvm.loop.unroll.enable metadata survives into the + // emitted LLVM IR. It replaces createConvertSCFToCFPass here; running both + // would be redundant. + kernelModulePM.addNestedPass(pto::createPTOConvertSCFToCFWithLoopHintsPass()); + kernelModulePM.addPass(createArithToLLVMConversionPass()); + kernelModulePM.addPass(createConvertIndexToLLVMPass()); + kernelModulePM.addPass(createFinalizeMemRefToLLVMConversionPass()); + kernelModulePM.addPass(createConvertFuncToLLVMPass()); + kernelModulePM.addPass(createConvertControlFlowToLLVMPass()); + kernelModulePM.addPass(createReconcileUnrealizedCastsPass()); + if (failed(mlir::applyPassManagerCLOptions(pm))) { + diagOS << "VPTO LLVM emission failed: unable to apply MLIR pass manager " + "command-line options\n"; + return failure(); + } + if (failed(pm.run(clonedModule))) { + diagOS << "VPTO LLVM emission failed: official lowering pipeline failed\n"; + return failure(); + } + return emit(clonedModule); +} + +LogicalResult lowerCANN900Module(ModuleOp module, const VPTOEmissionOptions &options, EmittedLLVMModule &cubeModule, + EmittedLLVMModule &vectorModule, llvm::raw_ostream &diagOS) { + llvm::StringSet simtEntryNames = collectSimtEntryFunctionNames(module); + cubeModule.context.reset(); + cubeModule.module.reset(); + vectorModule.context.reset(); + vectorModule.module.reset(); + return runPipeline(module, diagOS, [&](ModuleOp loweredModule) { + auto vectorDeviceModule = getUniqueDeviceModuleByKernelKind(loweredModule, FunctionKernelKind::Vector, diagOS); + if (failed(vectorDeviceModule)) { + return failure(); + } + auto cubeDeviceModule = getUniqueDeviceModuleByKernelKind(loweredModule, FunctionKernelKind::Cube, diagOS); + if (failed(cubeDeviceModule)) { + return failure(); + } + + if (*vectorDeviceModule) { + auto vectorOptions = makeDeviceEmissionOptions(options, FunctionKernelKind::Vector); + auto emitted = emitDeviceLLVMModule(*vectorDeviceModule, "vector", vectorOptions, simtEntryNames, diagOS); + if (failed(emitted)) { + return failure(); + } + vectorModule.context = std::move(emitted->context); + vectorModule.module = std::move(emitted->module); + } + if (*cubeDeviceModule) { + auto cubeOptions = makeDeviceEmissionOptions(options, FunctionKernelKind::Cube); + auto emitted = emitDeviceLLVMModule(*cubeDeviceModule, "cube", cubeOptions, simtEntryNames, diagOS); + if (failed(emitted)) { + return failure(); + } + cubeModule.context = std::move(emitted->context); + cubeModule.module = std::move(emitted->module); + } + return success(); + }); +} + +} // namespace mlir::pto::detail diff --git a/lib/PTO/Transforms/VPTOCANN900LLVMEmitterScalarPatterns.cpp b/lib/PTO/Transforms/VPTOCANN900LLVMEmitterScalarPatterns.cpp new file mode 100644 index 0000000000..69eb3fc6e8 --- /dev/null +++ b/lib/PTO/Transforms/VPTOCANN900LLVMEmitterScalarPatterns.cpp @@ -0,0 +1,1875 @@ +// Copyright (c) 2026 Huawei Technologies Co., Ltd. +// This program is free software, you can redistribute it and/or modify it under the terms and conditions of +// CANN Open Software License Agreement Version 2.0 (the "License"). +// Please refer to the License for details. You may not use this file except in compliance with the License. +// THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED, +// INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE. +// See LICENSE in the root of the software repository for the full text of the License. + +#include "VPTOCANN900LLVMEmitterTemplates.h" + +namespace mlir::pto::detail { + +class LowerTrapOpPattern final : public OpConversionPattern { +public: + explicit LowerTrapOpPattern(TypeConverter &typeConverter, MLIRContext *context, LoweringState &state) + : OpConversionPattern(typeConverter, context), state(state) {} + + LogicalResult matchAndRewrite(pto::TrapOp op, pto::TrapOp::Adaptor, + ConversionPatternRewriter &rewriter) const override { + constexpr StringLiteral calleeName = "llvm.hivm.TRAP"; + auto funcType = rewriter.getFunctionType({}, {}); + rewriter.create(op.getLoc(), calleeName, TypeRange{}, ValueRange{}); + state.plannedDecls.push_back(PlannedDecl{calleeName.str(), funcType}); + rewriter.eraseOp(op); + return success(); + } + +private: + LoweringState &state; +}; + +static LogicalResult appendVcvtImmediate(pto::VcvtOp op, StringRef name, std::optional immediate, + SmallVectorImpl &callArgs, SmallVectorImpl &argTypes, + ConversionPatternRewriter &rewriter) { + if (!immediate) { + StringRef message = name == "rnd" ? "vcvt requires valid rnd attr" + : name == "sat" ? "vcvt requires valid sat attr" + : "vcvt requires valid part attr"; + return rewriter.notifyMatchFailure(op, message); + } + Value value = getI32Constant(rewriter, op.getLoc(), *immediate); + callArgs.push_back(value); + argTypes.push_back(value.getType()); + return success(); +} + +static LogicalResult appendVcvtOptionalArguments(pto::VcvtOp op, const VcvtContract &contract, + SmallVectorImpl &callArgs, SmallVectorImpl &argTypes, + ConversionPatternRewriter &rewriter) { + auto appendRound = [&]() { + return appendVcvtImmediate(op, "rnd", op.getRndAttr() ? parseRoundModeImmediate(*op.getRnd()) : std::nullopt, + callArgs, argTypes, rewriter); + }; + auto appendSaturation = [&]() { + return appendVcvtImmediate(op, "sat", op.getSatAttr() ? parseSaturationImmediate(*op.getSat()) : std::nullopt, + callArgs, argTypes, rewriter); + }; + if (contract.satBeforeRnd) { + if (contract.requiresSat && failed(appendSaturation())) { + return failure(); + } + if (contract.requiresRnd && failed(appendRound())) { + return failure(); + } + } else { + if (contract.requiresRnd && failed(appendRound())) { + return failure(); + } + if (contract.requiresSat && failed(appendSaturation())) { + return failure(); + } + } + if (!contract.requiresPart) { + return success(); + } + return appendVcvtImmediate(op, "part", op.getPartAttr() ? parseVcvtPartImmediate(*op.getPart()) : std::nullopt, + callArgs, argTypes, rewriter); +} + +class LowerVcvtOpPattern final : public OpConversionPattern { +public: + explicit LowerVcvtOpPattern(TypeConverter &typeConverter, MLIRContext *context, LoweringState &state) + : OpConversionPattern(typeConverter, context), state(state) {} + + LogicalResult matchAndRewrite(pto::VcvtOp op, pto::VcvtOp::Adaptor adaptor, + ConversionPatternRewriter &rewriter) const override { + FailureOr contract = buildVcvtContract(op); + if (failed(contract)) { + return rewriter.notifyMatchFailure(op, "unsupported vcvt type pair"); + } + + Type resultType = this->getTypeConverter()->convertType(op.getResult().getType()); + if (!resultType) { + return rewriter.notifyMatchFailure(op, "failed to convert vcvt result type"); + } + + SmallVector callArgs; + SmallVector argTypes; + callArgs.push_back(adaptor.getInput()); + argTypes.push_back(adaptor.getInput().getType()); + callArgs.push_back(adaptor.getMask()); + argTypes.push_back(adaptor.getMask().getType()); + if (failed(appendVcvtOptionalArguments(op, *contract, callArgs, argTypes, rewriter))) { + return failure(); + } + + auto funcType = rewriter.getFunctionType(argTypes, TypeRange{resultType}); + auto call = + rewriter.create(op.getLoc(), StringRef((*contract).intrinsic), TypeRange{resultType}, callArgs); + state.plannedDecls.push_back(PlannedDecl{std::string((*contract).intrinsic), funcType}); + rewriter.replaceOp(op, call.getResults()); + return success(); + } + +private: + LoweringState &state; +}; + +class LowerVbitcastOpPattern final : public OpConversionPattern { +public: + explicit LowerVbitcastOpPattern(TypeConverter &typeConverter, MLIRContext *context, LoweringState &state) + : OpConversionPattern(typeConverter, context) {} + + LogicalResult matchAndRewrite(pto::VbitcastOp op, pto::VbitcastOp::Adaptor adaptor, + ConversionPatternRewriter &rewriter) const override { + // A vbitcast whose result has no users is a dead noop (Pure). Erase it + // instead of emitting an LLVM bitcast the device compiler may not lower + // (e.g. bf16x2 <-> bf16 physical views). + if (op->use_empty()) { + rewriter.eraseOp(op); + return success(); + } + Type resultType = this->getTypeConverter()->convertType(op.getResult().getType()); + if (!resultType) { + return rewriter.notifyMatchFailure(op, "failed to convert vbitcast result type"); + } + rewriter.replaceOpWithNewOp(op, resultType, adaptor.getInput()); + return success(); + } +}; + +class LowerPbitcastOpPattern final : public OpConversionPattern { +public: + explicit LowerPbitcastOpPattern(TypeConverter &typeConverter, MLIRContext *context, LoweringState &state) + : OpConversionPattern(typeConverter, context) {} + + LogicalResult matchAndRewrite(pto::PbitcastOp op, pto::PbitcastOp::Adaptor adaptor, + ConversionPatternRewriter &rewriter) const override { + Type resultType = this->getTypeConverter()->convertType(op.getResult().getType()); + if (!resultType) { + return rewriter.notifyMatchFailure(op, "failed to convert pbitcast result type"); + } + if (adaptor.getInput().getType() != resultType) { + return rewriter.notifyMatchFailure(op, "pbitcast expects identical lowered input/result types"); + } + rewriter.replaceOp(op, adaptor.getInput()); + return success(); + } +}; + +class LowerVtrcOpPattern final : public OpConversionPattern { +public: + explicit LowerVtrcOpPattern(TypeConverter &typeConverter, MLIRContext *context, LoweringState &state) + : OpConversionPattern(typeConverter, context), state(state) {} + + LogicalResult matchAndRewrite(pto::VtrcOp op, pto::VtrcOp::Adaptor adaptor, + ConversionPatternRewriter &rewriter) const override { + auto roundMode = parseRoundModeImmediate(op.getRoundMode()); + if (!roundMode) { + return rewriter.notifyMatchFailure(op, "unsupported vtrc signature"); + } + + Type resultType = this->getTypeConverter()->convertType(op.getResult().getType()); + if (!resultType) { + return rewriter.notifyMatchFailure(op, "failed to convert vtrc result type"); + } + + FailureOr calleeName = buildVtrcCallee(op.getContext(), op.getResult().getType()); + if (failed(calleeName)) { + return rewriter.notifyMatchFailure(op, "unsupported vtrc callee"); + } + + Value roundValue = getI32Constant(rewriter, op.getLoc(), *roundMode); + auto funcType = rewriter.getFunctionType( + TypeRange{adaptor.getInput().getType(), roundValue.getType(), adaptor.getMask().getType()}, + TypeRange{resultType}); + auto call = rewriter.create(op.getLoc(), *calleeName, TypeRange{resultType}, + ValueRange{adaptor.getInput(), roundValue, adaptor.getMask()}); + state.plannedDecls.push_back(PlannedDecl{calleeName->str(), funcType}); + rewriter.replaceOp(op, call.getResults()); + return success(); + } + +private: + LoweringState &state; +}; + +template +static SmallVector buildPredicateStoreCallArgs(StoreOp op, typename StoreOp::Adaptor adaptor, + const VPTOLoweredAddressOffset &offset, uint64_t dist, + bool usePostIntrinsic, ConversionPatternRewriter &rewriter) { + Value distValue = getI32Constant(rewriter, op.getLoc(), dist); + Value postValue = getI32Constant(rewriter, op.getLoc(), usePostIntrinsic ? 1 : 0); + return {adaptor.getValue(), offset.base, offset.intrinsicOffset, distValue, postValue}; +} + +template +static void replacePredicateStoreOp(StoreOp op, bool usePostIntrinsic, const VPTOLoweredAddressOffset &offset, + func::CallOp call, ConversionPatternRewriter &rewriter) { + if (!usePostIntrinsic) { + rewriter.eraseOp(op); + return; + } + if (offset.updatedBase) { + rewriter.replaceOp(op, offset.updatedBase); + return; + } + rewriter.replaceOp(op, call.getResults()); +} + +template class LowerPredicateStoreOpPattern final : public OpConversionPattern { +public: + explicit LowerPredicateStoreOpPattern(TypeConverter &typeConverter, MLIRContext *context, LoweringState &state) + : OpConversionPattern(typeConverter, context), state(state) {} + + LogicalResult matchAndRewrite(StoreOp op, typename StoreOp::Adaptor adaptor, + ConversionPatternRewriter &rewriter) const override { + auto llvmDestType = dyn_cast(adaptor.getDestination().getType()); + Type valueType = this->getTypeConverter()->convertType(op.getValue().getType()); + if (!llvmDestType || !valueType) { + return rewriter.notifyMatchFailure(op, "expected converted predicate-store operand types"); + } + + auto dist = parsePredicateStoreDistImmediate(op.getDist()); + if (!dist) { + return rewriter.notifyMatchFailure(op, "unsupported predicate-store dist immediate"); + } + + bool usePostIntrinsic = op.getUpdatedBase() != nullptr; + auto loweredOffset = lowerVPTOPredicateOffsetForIntrinsic(op, adaptor.getDestination(), adaptor.getOffset(), + usePostIntrinsic, rewriter); + if (failed(loweredOffset)) { + return rewriter.notifyMatchFailure(op, "failed to preserve predicate-store index offset"); + } + + SmallVector resultTypes; + if (failed(this->getTypeConverter()->convertTypes(op->getResultTypes(), resultTypes)) || + resultTypes.size() != (usePostIntrinsic ? 1U : 0U)) { + return rewriter.notifyMatchFailure(op, "failed to convert predicate-store result types"); + } + + StringRef calleeName = getPredicateStoreCallee(op.getContext(), usePostIntrinsic); + SmallVector args = + buildPredicateStoreCallArgs(op, adaptor, *loweredOffset, *dist, usePostIntrinsic, rewriter); + auto funcType = rewriter.getFunctionType( + TypeRange{valueType, llvmDestType, rewriter.getI32Type(), rewriter.getI32Type(), rewriter.getI32Type()}, + resultTypes); + auto call = rewriter.create(op.getLoc(), calleeName, resultTypes, args); + state.plannedDecls.push_back(PlannedDecl{calleeName.str(), funcType}); + replacePredicateStoreOp(op, usePostIntrinsic, *loweredOffset, call, rewriter); + return success(); + } + +private: + LoweringState &state; +}; + +template class LowerPredicateLoadOpPattern final : public OpConversionPattern { +public: + explicit LowerPredicateLoadOpPattern(TypeConverter &typeConverter, MLIRContext *context, LoweringState &state) + : OpConversionPattern(typeConverter, context), state(state) {} + + LogicalResult matchAndRewrite(LoadOp op, typename LoadOp::Adaptor adaptor, + ConversionPatternRewriter &rewriter) const override { + auto llvmSourceType = dyn_cast(adaptor.getSource().getType()); + bool usePostIntrinsic = op.getUpdatedBase() != nullptr; + SmallVector resultTypes; + if (failed(this->getTypeConverter()->convertTypes(op->getResultTypes(), resultTypes)) || + resultTypes.size() != (usePostIntrinsic ? 2U : 1U)) { + return rewriter.notifyMatchFailure(op, "failed to convert predicate-load result types"); + } + if (!llvmSourceType) { + return rewriter.notifyMatchFailure(op, "expected converted predicate-load operand/result types"); + } + + auto dist = parsePredicateLoadDistImmediate(op.getDist()); + if (!dist) { + return rewriter.notifyMatchFailure(op, "unsupported predicate-load dist immediate"); + } + + auto loweredOffset = + lowerVPTOPredicateOffsetForIntrinsic(op, adaptor.getSource(), adaptor.getOffset(), usePostIntrinsic, rewriter); + if (failed(loweredOffset)) { + return rewriter.notifyMatchFailure(op, "failed to preserve predicate-load index offset"); + } + + StringRef calleeName = getPredicateLoadCallee(op.getContext(), usePostIntrinsic); + SmallVector args; + args.push_back(loweredOffset->base); + args.push_back(loweredOffset->intrinsicOffset); + args.push_back(rewriter.create(op.getLoc(), rewriter.getI32IntegerAttr(*dist))); + args.push_back( + rewriter.create(op.getLoc(), rewriter.getI32IntegerAttr(usePostIntrinsic ? 1 : 0))); + auto funcType = rewriter.getFunctionType( + TypeRange{llvmSourceType, rewriter.getI32Type(), rewriter.getI32Type(), rewriter.getI32Type()}, resultTypes); + auto call = rewriter.create(op.getLoc(), calleeName, resultTypes, args); + state.plannedDecls.push_back(PlannedDecl{calleeName.str(), funcType}); + if (loweredOffset->updatedBase) { + rewriter.replaceOp(op, ValueRange{call.getResult(0), loweredOffset->updatedBase}); + } else { + rewriter.replaceOp(op, call.getResults()); + } + return success(); + } + +private: + LoweringState &state; +}; + +template class LowerSetLoopConfigOpPattern final : public OpConversionPattern { +public: + explicit LowerSetLoopConfigOpPattern(TypeConverter &typeConverter, MLIRContext *context, LoweringState &state) + : OpConversionPattern(typeConverter, context), state(state) {} + + LogicalResult matchAndRewrite(LoopOp op, typename LoopOp::Adaptor adaptor, + ConversionPatternRewriter &rewriter) const override { + FailureOr packed = failure(); + if constexpr (std::is_same_v || + std::is_same_v) { + packed = packLoopSize(op, adaptor.getFirst(), adaptor.getSecond()); + } else { + packed = packLoopPair(op, adaptor.getFirst(), adaptor.getSecond()); + } + if (failed(packed)) { + return rewriter.notifyMatchFailure(op, "failed to pack loop configuration"); + } + + StringRef calleeName = buildSetLoopCallee(op.getContext()); + auto funcType = rewriter.getFunctionType(TypeRange{rewriter.getI64Type()}, TypeRange{}); + rewriter.create(op.getLoc(), calleeName, TypeRange{}, ValueRange{*packed}); + state.plannedDecls.push_back(PlannedDecl{calleeName.str(), funcType}); + rewriter.eraseOp(op); + return success(); + } + +private: + LoweringState &state; +}; + +template class LowerUnaryConfigOpPattern final : public OpConversionPattern { +public: + explicit LowerUnaryConfigOpPattern(TypeConverter &typeConverter, MLIRContext *context, LoweringState &state) + : OpConversionPattern(typeConverter, context), state(state) {} + + LogicalResult matchAndRewrite(ConfigOp op, typename ConfigOp::Adaptor adaptor, + ConversionPatternRewriter &rewriter) const override { + FailureOr encoded = encodeMovPadValue(op.getLoc(), adaptor.getValue(), rewriter); + if (failed(encoded)) { + return rewriter.notifyMatchFailure(op, "expected 8/16/32-bit integer or float mov-pad payload"); + } + + StringRef calleeName = buildUnaryConfigCallee(op.getContext()); + auto funcType = rewriter.getFunctionType(TypeRange{rewriter.getI64Type()}, TypeRange{}); + rewriter.create(op.getLoc(), calleeName, TypeRange{}, ValueRange{*encoded}); + state.plannedDecls.push_back(PlannedDecl{calleeName.str(), funcType}); + rewriter.eraseOp(op); + return success(); + } + +private: + LoweringState &state; +}; + +template class LowerUnaryI64ConfigOpPattern final : public OpConversionPattern { +public: + explicit LowerUnaryI64ConfigOpPattern(TypeConverter &typeConverter, MLIRContext *context, LoweringState &state) + : OpConversionPattern(typeConverter, context), state(state) {} + + LogicalResult matchAndRewrite(ConfigOp op, typename ConfigOp::Adaptor adaptor, + ConversionPatternRewriter &rewriter) const override { + StringRef calleeName = buildUnaryConfigCallee(op.getContext()); + auto funcType = rewriter.getFunctionType(TypeRange{adaptor.getValue().getType()}, TypeRange{}); + rewriter.create(op.getLoc(), calleeName, TypeRange{}, ValueRange{adaptor.getValue()}); + state.plannedDecls.push_back(PlannedDecl{calleeName.str(), funcType}); + rewriter.eraseOp(op); + return success(); + } + +private: + LoweringState &state; +}; + +class LowerStoreVfSimtInfoOpPattern final : public OpConversionPattern { +public: + explicit LowerStoreVfSimtInfoOpPattern(TypeConverter &typeConverter, MLIRContext *context, LoweringState &state) + : OpConversionPattern(typeConverter, context), state(state) {} + + LogicalResult matchAndRewrite(pto::StoreVfSimtInfoOp op, pto::StoreVfSimtInfoOp::Adaptor adaptor, + ConversionPatternRewriter &rewriter) const override { + Location loc = op.getLoc(); + Value dimZ = adaptor.getDimZ(); + Value dimY = adaptor.getDimY(); + Value dimX = adaptor.getDimX(); + if (!dimZ || !dimY || !dimX) { + return rewriter.notifyMatchFailure(op, "missing converted SIMT dims"); + } + + auto i64Type = rewriter.getI64Type(); + auto castToI64 = [&](Value value) -> Value { + if (value.getType().isInteger(64)) { + return value; + } + return rewriter.create(loc, i64Type, value).getResult(); + }; + + Value dimZI64 = castToI64(dimZ); + Value dimYI64 = castToI64(dimY); + Value dimXI64 = castToI64(dimX); + Value dimYShift = rewriter.create(loc, i64Type, rewriter.getI64IntegerAttr(16)); + Value dimZShift = rewriter.create(loc, i64Type, rewriter.getI64IntegerAttr(32)); + Value packedDimY = rewriter.create(loc, dimYI64, dimYShift).getResult(); + Value packedDimZ = rewriter.create(loc, dimZI64, dimZShift).getResult(); + Value payload = rewriter.create(loc, dimXI64, packedDimY).getResult(); + payload = rewriter.create(loc, payload, packedDimZ).getResult(); + + StringRef calleeName = buildStoreVfSimtInfoCallee(op.getContext()); + auto funcType = rewriter.getFunctionType(TypeRange{i64Type}, TypeRange{}); + rewriter.create(loc, calleeName, TypeRange{}, ValueRange{payload}); + state.plannedDecls.push_back(PlannedDecl{calleeName.str(), funcType}); + rewriter.eraseOp(op); + return success(); + } + +private: + LoweringState &state; +}; + +template static StringRef buildSimtFenceCallee(MLIRContext *context); + +template <> StringRef buildSimtFenceCallee(MLIRContext *context) { + return buildSyncthreadsCallee(context); +} + +template <> StringRef buildSimtFenceCallee(MLIRContext *context) { + return buildThreadfenceCallee(context); +} + +template <> StringRef buildSimtFenceCallee(MLIRContext *context) { + return buildThreadfenceBlockCallee(context); +} + +template class LowerSimtFenceOpPattern final : public OpConversionPattern { +public: + explicit LowerSimtFenceOpPattern(TypeConverter &typeConverter, MLIRContext *context, LoweringState &state) + : OpConversionPattern(typeConverter, context), state(state) {} + + LogicalResult matchAndRewrite(FenceOp op, typename FenceOp::Adaptor adaptor, + ConversionPatternRewriter &rewriter) const override { + (void)adaptor; + FunctionType funcType = rewriter.getFunctionType({}, {}); + StringRef calleeName = buildSimtFenceCallee(op.getContext()); + rewriter.create(op.getLoc(), calleeName, TypeRange{}, ValueRange{}); + state.plannedDecls.push_back(PlannedDecl{calleeName.str(), funcType}); + rewriter.eraseOp(op); + return success(); + } + +private: + LoweringState &state; +}; + +struct SimtKeepResumePhysicalRegister { + int64_t baseRegister; + unsigned registerCount; +}; + +// TPERn names one 32-bit register, while TPERLn names the 64-bit pair whose +// base register is R(2n). Keep uses tied inputs so the compiler models the +// value captured by each fixed output without inline assembly instructions. +static std::string buildSimtKeepResumeConstraints(ArrayRef physicalRegs, + bool tieInputs) { + std::string result; + llvm::raw_string_ostream os(result); + for (auto [index, physicalReg] : llvm::enumerate(physicalRegs)) { + if (index != 0) { + os << ","; + } + if (physicalReg.registerCount == 2) { + os << "={TPERL" << physicalReg.baseRegister / 2 << "}"; + } else { + os << "={TPER" << physicalReg.baseRegister << "}"; + } + } + if (tieInputs) { + for (size_t index = 0; index < physicalRegs.size(); ++index) { + os << "," << index; + } + } + return os.str(); +} + +template static SmallVector collectConsecutiveOps(OpT first) { + SmallVector ops; + for (Operation *cur = first.getOperation(); cur; cur = cur->getNextNode()) { + auto typed = dyn_cast(cur); + if (!typed) { + break; + } + ops.push_back(typed); + } + return ops; +} + +static bool hasPreviousSameOp(Operation *op) { + Operation *prev = op->getPrevNode(); + return prev && prev->getName() == op->getName(); +} + +static std::optional getSimtKeepResumeBitWidth(Type type) { + if (auto intType = dyn_cast(type)) { + if (intType.getWidth() <= 64) { + return intType.getWidth(); + } + return std::nullopt; + } + if (type.isF16() || type.isBF16()) { + return 16; + } + if (type.isF32()) { + return 32; + } + return std::nullopt; +} + +static Value packSimtKeepResumePayload(Location loc, Value value, ConversionPatternRewriter &rewriter) { + Type type = value.getType(); + std::optional width = getSimtKeepResumeBitWidth(type); + if (!width) { + return {}; + } + + Type intType = rewriter.getIntegerType(*width); + Value bits = value; + if (!isa(type)) { + bits = rewriter.create(loc, intType, value); + } else if (bits.getType() != intType) { + bits = rewriter.create(loc, intType, bits); + } + if (*width < 32) { + return rewriter.create(loc, rewriter.getI32Type(), bits); + } + if (*width == 32 && bits.getType() != rewriter.getI32Type()) { + return rewriter.create(loc, rewriter.getI32Type(), bits); + } + return bits; +} + +static Value unpackSimtKeepResumePayload(Location loc, Value value, Type resultType, + ConversionPatternRewriter &rewriter) { + std::optional width = getSimtKeepResumeBitWidth(resultType); + if (!width) { + return {}; + } + + Type intType = rewriter.getIntegerType(*width); + Value bits = value; + if (*width < 32) { + bits = rewriter.create(loc, intType, bits); + } else if (bits.getType() != intType) { + bits = rewriter.create(loc, intType, bits); + } + + if (isa(resultType)) { + if (bits.getType() == resultType) { + return bits; + } + return rewriter.create(loc, resultType, bits); + } + return rewriter.create(loc, resultType, bits); +} + +static unsigned getSimtKeepResumeRegisterCount(Type type) { + std::optional width = getSimtKeepResumeBitWidth(type); + return width && *width > 32 ? 2 : 1; +} + +static FailureOr> +computeSimtKeepResumePhysicalRegs(ArrayRef> logicalSlots) { + SmallVector physicalRegs; + physicalRegs.reserve(logicalSlots.size()); + for (auto [slot, registerCount] : logicalSlots) { + if (slot < 0 || slot >= 123) { + return failure(); + } + if (registerCount == 2 && ((slot % 2) != 0 || slot + 1 >= 123)) { + return failure(); + } + // Slots are user-assigned storage words, not dense ordinals in the current + // keep/resume group. This keeps a consumer that resumes only a subset of + // slots from changing where the remaining slots are read from. + int64_t baseRegister = 4 + slot; + if (baseRegister + static_cast(registerCount) - 1 > 126) { + return failure(); + } + physicalRegs.push_back({baseRegister, registerCount}); + } + return physicalRegs; +} + +static bool isValidSimtKeepResumeSlot(int64_t slot, unsigned registerCount) { + if (slot < 0 || slot >= 123) { + return false; + } + if (registerCount == 2 && ((slot % 2) != 0 || slot + 1 >= 123)) { + return false; + } + return true; +} + +struct ResumeGroupTypes { + SmallVector, 4> logicalSlots; + SmallVector asmResultTypes; +}; + +static FailureOr collectResumeGroupTypes(ArrayRef resumeOps, + const TypeConverter &typeConverter, + ConversionPatternRewriter &rewriter) { + ResumeGroupTypes types; + for (unsigned index = 0; index < resumeOps.size(); ++index) { + pto::ResumeOp resume = resumeOps[index]; + Type resultType = typeConverter.convertType(resume.getType()); + std::optional bitWidth = getSimtKeepResumeBitWidth(resultType); + if (!resultType || !bitWidth) { + return failure(); + } + unsigned registerCount = getSimtKeepResumeRegisterCount(resultType); + if (!isValidSimtKeepResumeSlot(resume.getSlot(), registerCount)) { + return failure(); + } + types.logicalSlots.push_back({resume.getSlot(), registerCount}); + types.asmResultTypes.push_back(rewriter.getIntegerType(*bitWidth > 32 ? 64 : 32)); + } + return types; +} + +static LogicalResult replaceResumeGroup(ArrayRef resumeOps, LLVM::InlineAsmOp asmOp, + const TypeConverter &typeConverter, ConversionPatternRewriter &rewriter) { + SmallVector results; + for (unsigned index = 0; index < resumeOps.size(); ++index) { + pto::ResumeOp resume = resumeOps[index]; + auto extract = rewriter.create(resume.getLoc(), asmOp.getRes(), + ArrayRef{static_cast(index)}); + Type resultType = typeConverter.convertType(resume.getType()); + Value result = unpackSimtKeepResumePayload(resume.getLoc(), extract.getRes(), resultType, rewriter); + if (!result) { + return failure(); + } + results.push_back(result); + } + for (auto [resume, result] : llvm::zip(resumeOps, results)) { + rewriter.replaceOp(resume, result); + } + return success(); +} + +class LowerKeepOpPattern final : public OpConversionPattern { +public: + explicit LowerKeepOpPattern(TypeConverter &typeConverter, MLIRContext *context, LoweringState &) + : OpConversionPattern(typeConverter, context) {} + + LogicalResult matchAndRewrite(pto::KeepOp op, pto::KeepOp::Adaptor adaptor, + ConversionPatternRewriter &rewriter) const override { + (void)adaptor; + if (hasPreviousSameOp(op.getOperation())) { + return rewriter.notifyMatchFailure(op, "only the first keep in a contiguous group is lowered"); + } + + SmallVector keepOps = collectConsecutiveOps(op); + SmallVector payloads; + SmallVector asmResultTypes; + SmallVector, 4> logicalSlots; + for (pto::KeepOp keep : keepOps) { + Value payload = rewriter.getRemappedValue(keep.getPayload()); + if (!payload) { + return rewriter.notifyMatchFailure(keep, "payload is not remapped"); + } + payload = packSimtKeepResumePayload(keep.getLoc(), payload, rewriter); + if (!payload) { + return rewriter.notifyMatchFailure(keep, "expected integer scalar up to 64 bits or f16/bf16/f32"); + } + int64_t slot = keep.getSlot(); + unsigned registerCount = getSimtKeepResumeRegisterCount(payload.getType()); + if (!isValidSimtKeepResumeSlot(slot, registerCount)) { + return rewriter.notifyMatchFailure(keep, "slot must be in range [0, 122] and 64-bit slots must be even"); + } + logicalSlots.push_back({slot, registerCount}); + payloads.push_back(payload); + asmResultTypes.push_back(payload.getType()); + } + FailureOr> physicalRegs = + computeSimtKeepResumePhysicalRegs(logicalSlots); + if (failed(physicalRegs)) { + return rewriter.notifyMatchFailure(op, "keep slots must map to valid non-overlapping SIMT registers"); + } + + Type asmResultType = asmResultTypes.front(); + if (asmResultTypes.size() > 1) { + asmResultType = LLVM::LLVMStructType::getLiteral(op.getContext(), asmResultTypes); + } + rewriter.setInsertionPoint(op); + rewriter.create( + op.getLoc(), TypeRange{asmResultType}, payloads, "", buildSimtKeepResumeConstraints(*physicalRegs, true), true, + false, LLVM::AsmDialectAttr::get(op.getContext(), LLVM::AsmDialect::AD_ATT), ArrayAttr{}); + for (pto::KeepOp keep : llvm::reverse(keepOps)) { + rewriter.eraseOp(keep); + } + return success(); + } +}; + +class LowerResumeOpPattern final : public OpConversionPattern { +public: + explicit LowerResumeOpPattern(TypeConverter &typeConverter, MLIRContext *context, LoweringState &) + : OpConversionPattern(typeConverter, context) {} + + LogicalResult matchAndRewrite(pto::ResumeOp op, pto::ResumeOp::Adaptor adaptor, + ConversionPatternRewriter &rewriter) const override { + (void)adaptor; + if (hasPreviousSameOp(op.getOperation())) { + return rewriter.notifyMatchFailure(op, "only the first resume in a contiguous group is lowered"); + } + + SmallVector resumeOps = collectConsecutiveOps(op); + FailureOr groupTypes = collectResumeGroupTypes(resumeOps, *getTypeConverter(), rewriter); + if (failed(groupTypes)) { + return rewriter.notifyMatchFailure(op, "resume slots or result types are unsupported"); + } + FailureOr> physicalRegs = + computeSimtKeepResumePhysicalRegs(groupTypes->logicalSlots); + if (failed(physicalRegs)) { + return rewriter.notifyMatchFailure(op, "resume slots must map to valid non-overlapping SIMT registers"); + } + + Type asmResultType = groupTypes->asmResultTypes.front(); + if (groupTypes->asmResultTypes.size() > 1) { + asmResultType = LLVM::LLVMStructType::getLiteral(op.getContext(), groupTypes->asmResultTypes); + } + rewriter.setInsertionPoint(op); + auto asmOp = rewriter.create( + op.getLoc(), TypeRange{asmResultType}, ValueRange{}, "", buildSimtKeepResumeConstraints(*physicalRegs, false), + true, false, LLVM::AsmDialectAttr::get(op.getContext(), LLVM::AsmDialect::AD_ATT), ArrayAttr{}); + + if (resumeOps.size() == 1) { + Type resultType = getTypeConverter()->convertType(op.getType()); + Value result = unpackSimtKeepResumePayload(op.getLoc(), asmOp.getRes(), resultType, rewriter); + if (!result) { + return rewriter.notifyMatchFailure(op, "failed to unpack result"); + } + rewriter.replaceOp(op, result); + return success(); + } + + rewriter.setInsertionPointAfter(asmOp); + if (failed(replaceResumeGroup(resumeOps, asmOp, *getTypeConverter(), rewriter))) { + return rewriter.notifyMatchFailure(op, "failed to unpack resume results"); + } + return success(); + } +}; + +template class LowerNullaryConfigOpPattern final : public OpConversionPattern { +public: + explicit LowerNullaryConfigOpPattern(TypeConverter &typeConverter, MLIRContext *context, LoweringState &state) + : OpConversionPattern(typeConverter, context), state(state) {} + + LogicalResult matchAndRewrite(ConfigOp op, typename ConfigOp::Adaptor adaptor, + ConversionPatternRewriter &rewriter) const override { + (void)adaptor; + StringRef calleeName = buildNullaryConfigCallee(op.getContext()); + auto funcType = rewriter.getFunctionType(TypeRange{}, TypeRange{}); + rewriter.create(op.getLoc(), calleeName, TypeRange{}, ValueRange{}); + state.plannedDecls.push_back(PlannedDecl{calleeName.str(), funcType}); + rewriter.eraseOp(op); + return success(); + } + +private: + LoweringState &state; +}; + +template class LowerPipeEventSyncOpPattern final : public OpConversionPattern { +public: + explicit LowerPipeEventSyncOpPattern(TypeConverter &typeConverter, MLIRContext *context, LoweringState &state) + : OpConversionPattern(typeConverter, context), state(state) {} + + LogicalResult matchAndRewrite(SyncOp op, typename SyncOp::Adaptor adaptor, + ConversionPatternRewriter &rewriter) const override { + (void)adaptor; + auto src = parsePipeImmediate(stringifyPIPE(op.getSrcPipe().getPipe())); + auto dst = parsePipeImmediate(stringifyPIPE(op.getDstPipe().getPipe())); + auto event = parseEventImmediate(stringifyEVENT(op.getEventId().getEvent())); + if (!src || !dst || !event) { + return rewriter.notifyMatchFailure(op, "unsupported sync immediate"); + } + + StringRef calleeName = buildSyncCallee(op.getContext()); + Value srcValue = getI64Constant(rewriter, op.getLoc(), *src); + Value dstValue = getI64Constant(rewriter, op.getLoc(), *dst); + Value eventValue = getI64Constant(rewriter, op.getLoc(), *event); + auto funcType = rewriter.getFunctionType( + TypeRange{rewriter.getI64Type(), rewriter.getI64Type(), rewriter.getI64Type()}, TypeRange{}); + rewriter.create(op.getLoc(), calleeName, TypeRange{}, ValueRange{srcValue, dstValue, eventValue}); + state.plannedDecls.push_back(PlannedDecl{calleeName.str(), funcType}); + rewriter.eraseOp(op); + return success(); + } + +private: + LoweringState &state; +}; + +template class LowerPipeEventDynSyncOpPattern final : public OpConversionPattern { +public: + explicit LowerPipeEventDynSyncOpPattern(TypeConverter &typeConverter, MLIRContext *context, LoweringState &state) + : OpConversionPattern(typeConverter, context), state(state) {} + + LogicalResult matchAndRewrite(SyncOp op, typename SyncOp::Adaptor adaptor, + ConversionPatternRewriter &rewriter) const override { + auto src = parsePipeImmediate(stringifyPIPE(op.getSrcPipe().getPipe())); + auto dst = parsePipeImmediate(stringifyPIPE(op.getDstPipe().getPipe())); + if (!src || !dst) { + return rewriter.notifyMatchFailure(op, "unsupported sync pipe"); + } + + StringRef calleeName = buildSyncCallee(op.getContext()); + Value srcValue = getI64Constant(rewriter, op.getLoc(), *src); + Value dstValue = getI64Constant(rewriter, op.getLoc(), *dst); + + Value eventIdValue = adaptor.getEventId(); + if (!eventIdValue) { + return rewriter.notifyMatchFailure(op, "missing event_id operand"); + } + + Value eventValue = eventIdValue; + + while (eventValue.getDefiningOp()) { + auto unrealizedCast = dyn_cast(eventValue.getDefiningOp()); + if (!unrealizedCast || unrealizedCast.getInputs().size() != 1) { + break; + } + eventValue = unrealizedCast.getInputs()[0]; + } + + if (eventValue.getType().isIndex()) { + eventValue = rewriter.create(op.getLoc(), rewriter.getI64Type(), eventValue); + } else if (auto intType = dyn_cast(eventValue.getType())) { + if (intType.getWidth() < 64) { + eventValue = rewriter.create(op.getLoc(), rewriter.getI64Type(), eventValue); + } + } else { + return rewriter.notifyMatchFailure(op, "unexpected event_id type"); + } + + auto funcType = rewriter.getFunctionType( + TypeRange{rewriter.getI64Type(), rewriter.getI64Type(), rewriter.getI64Type()}, TypeRange{}); + rewriter.create(op.getLoc(), calleeName, TypeRange{}, ValueRange{srcValue, dstValue, eventValue}); + state.plannedDecls.push_back(PlannedDecl{calleeName.str(), funcType}); + rewriter.eraseOp(op); + return success(); + } + +private: + LoweringState &state; +}; + +template +static FailureOr getInterCoreEventValue(SyncOp op, typename SyncOp::Adaptor adaptor, + ConversionPatternRewriter &rewriter) { + if (IntegerAttr eventIdAttr = op.getEventIdAttr()) { + return getI64Constant(rewriter, op.getLoc(), eventIdAttr.getInt()); + } + Value eventId = adaptor.getEventIdDyn(); + if (!eventId) { + return failure(); + } + Value eventValue = castIntegerLikeTo(op, eventId, rewriter.getI64Type()); + if (!eventValue) { + return failure(); + } + return eventValue; +} + +template +static SmallVector buildInterCoreSyncArgs(SyncOp op, Value pipeValue, Value eventValue, StringRef &calleeName, + ConversionPatternRewriter &rewriter) { + if constexpr (std::is_same_v) { + int64_t mode = op.getFftsModeAttr() ? op.getFftsModeAttr().getInt() : 2; + Value modeValue = getI64Constant(rewriter, op.getLoc(), mode); + modeValue = rewriter.create(op.getLoc(), modeValue, getI64Constant(rewriter, op.getLoc(), 0x3)); + eventValue = rewriter.create(op.getLoc(), eventValue, getI64Constant(rewriter, op.getLoc(), 0xf)); + Value modeShift = rewriter.create(op.getLoc(), modeValue, getI64Constant(rewriter, op.getLoc(), 4)); + Value eventShift = + rewriter.create(op.getLoc(), eventValue, getI64Constant(rewriter, op.getLoc(), 8)); + Value message = rewriter.create(op.getLoc(), getI64Constant(rewriter, op.getLoc(), 1), modeShift); + message = rewriter.create(op.getLoc(), message, eventShift); + return {pipeValue, message}; + } + if constexpr (std::is_same_v) { + calleeName = op.getEventIdAttr() ? StringAttr::get(op.getContext(), "llvm.hivm.WAIT.FLAG.DEV.PIPE.IMM").getValue() + : StringAttr::get(op.getContext(), "llvm.hivm.WAIT.FLAG.DEV.PIPE.REG").getValue(); + } + return {pipeValue, eventValue}; +} + +template class LowerInterCoreSyncOpPattern final : public OpConversionPattern { +public: + explicit LowerInterCoreSyncOpPattern(TypeConverter &typeConverter, MLIRContext *context, LoweringState &state) + : OpConversionPattern(typeConverter, context), state(state) {} + + LogicalResult matchAndRewrite(SyncOp op, typename SyncOp::Adaptor adaptor, + ConversionPatternRewriter &rewriter) const override { + auto pipe = parsePipeImmediate(stringifyPIPE(op.getPipe().getPipe())); + if (!pipe) { + return rewriter.notifyMatchFailure(op, "unsupported inter-core sync pipe"); + } + + Value pipeValue = getI64Constant(rewriter, op.getLoc(), *pipe); + FailureOr eventValue = getInterCoreEventValue(op, adaptor, rewriter); + if (failed(eventValue)) { + return rewriter.notifyMatchFailure(op, "expected a valid static or dynamic event-id operand"); + } + + StringRef calleeName = buildSyncCallee(op.getContext()); + SmallVector args = buildInterCoreSyncArgs(op, pipeValue, *eventValue, calleeName, rewriter); + auto funcType = rewriter.getFunctionType(TypeRange{args[0].getType(), args[1].getType()}, TypeRange{}); + rewriter.create(op.getLoc(), calleeName, TypeRange{}, args); + state.plannedDecls.push_back(PlannedDecl{calleeName.str(), funcType}); + rewriter.eraseOp(op); + return success(); + } + +private: + LoweringState &state; +}; + +template class LowerNamedSyncOpPattern final : public OpConversionPattern { +public: + explicit LowerNamedSyncOpPattern(TypeConverter &tc, MLIRContext *ctx, LoweringState &state) + : OpConversionPattern(tc, ctx), state(state) {} + LogicalResult matchAndRewrite(SyncOp op, typename SyncOp::Adaptor adaptor, + ConversionPatternRewriter &rewriter) const override { + auto pipe = parsePipeImmediate(stringifyPIPE(op.getPipe().getPipe())); + if (!pipe) { + return rewriter.notifyMatchFailure(op, "unsupported sync pipe"); + } + Value pipeValue = getI64Constant(rewriter, op.getLoc(), *pipe); + Value eventValue; + if (IntegerAttr attr = op.getEventIdAttr()) { + eventValue = getI64Constant(rewriter, op.getLoc(), attr.getInt()); + } else { + eventValue = castIntegerLikeTo(op, adaptor.getEventIdDyn(), rewriter.getI64Type()); + if (!eventValue) { + return rewriter.notifyMatchFailure(op, "missing event-id operand"); + } + } + StringRef callee = buildSyncCallee(op.getContext()); + auto fnTy = rewriter.getFunctionType(TypeRange{rewriter.getI64Type(), rewriter.getI64Type()}, TypeRange{}); + rewriter.create(op.getLoc(), callee, TypeRange{}, ValueRange{pipeValue, eventValue}); + state.plannedDecls.push_back(PlannedDecl{callee.str(), fnTy}); + rewriter.eraseOp(op); + return success(); + } + +private: + LoweringState &state; +}; + +class LowerBarrierOpPattern final : public OpConversionPattern { +public: + explicit LowerBarrierOpPattern(TypeConverter &typeConverter, MLIRContext *context, LoweringState &state) + : OpConversionPattern(typeConverter, context), state(state) {} + + LogicalResult matchAndRewrite(pto::BarrierOp op, OpAdaptor adaptor, + ConversionPatternRewriter &rewriter) const override { + (void)adaptor; + if (isTargetArchA5(op.getOperation()) && op.getPipe().getPipe() == PIPE::PIPE_V) { + op.emitError("internal error: A5 PIPE_V barrier should be erased before " + "VPTO LLVM lowering"); + return failure(); + } + + auto pipe = parsePipeImmediate(stringifyPIPE(op.getPipe().getPipe())); + if (!pipe) { + return rewriter.notifyMatchFailure(op, "unsupported barrier pipe"); + } + + StringRef calleeName = buildSyncCallee(op.getContext()); + Value pipeValue = getI64Constant(rewriter, op.getLoc(), *pipe); + auto funcType = rewriter.getFunctionType(TypeRange{rewriter.getI64Type()}, TypeRange{}); + rewriter.create(op.getLoc(), calleeName, TypeRange{}, ValueRange{pipeValue}); + state.plannedDecls.push_back(PlannedDecl{calleeName.str(), funcType}); + rewriter.eraseOp(op); + return success(); + } + +private: + LoweringState &state; +}; + +class LowerMemBarOpPattern final : public OpConversionPattern { +public: + explicit LowerMemBarOpPattern(TypeConverter &typeConverter, MLIRContext *context, LoweringState &state) + : OpConversionPattern(typeConverter, context), state(state) {} + + LogicalResult matchAndRewrite(pto::MemBarOp op, OpAdaptor adaptor, + ConversionPatternRewriter &rewriter) const override { + (void)adaptor; + StringRef calleeName = buildMemBarCallee(op.getKind().getKind(), op.getContext()); + auto funcType = rewriter.getFunctionType(TypeRange{}, TypeRange{}); + rewriter.create(op.getLoc(), calleeName, TypeRange{}, ValueRange{}); + state.plannedDecls.push_back(PlannedDecl{calleeName.str(), funcType}); + rewriter.eraseOp(op); + return success(); + } + +private: + LoweringState &state; +}; + +template +class LowerUnsupportedMemoryConsistencyOpPattern final : public OpConversionPattern { +public: + explicit LowerUnsupportedMemoryConsistencyOpPattern(TypeConverter &typeConverter, MLIRContext *context, + LoweringState &state) + : OpConversionPattern(typeConverter, context) { + (void)state; + } + + LogicalResult matchAndRewrite(MemoryConsistencyOp op, typename MemoryConsistencyOp::Adaptor adaptor, + ConversionPatternRewriter &rewriter) const override { + (void)adaptor; + (void)rewriter; + op.emitOpError() << "is not supported by the VPTO backend yet; PTOAS validates the " + "memory-consistency contract, but high-level CMO/fence ops must be " + "lowered to `pto.dcci` or `pto.dsb` before VPTO LLVM lowering"; + return failure(); + } +}; + +class LowerDsbOpPattern final : public OpConversionPattern { +public: + explicit LowerDsbOpPattern(TypeConverter &typeConverter, MLIRContext *context, LoweringState &state) + : OpConversionPattern(typeConverter, context), state(state) {} + + LogicalResult matchAndRewrite(pto::DsbOp op, OpAdaptor adaptor, ConversionPatternRewriter &rewriter) const override { + (void)adaptor; + StringRef calleeName = StringAttr::get(op.getContext(), "llvm.hivm.DSB").getValue(); + Type i64Ty = rewriter.getI64Type(); + auto funcType = rewriter.getFunctionType(TypeRange{i64Ty}, TypeRange{}); + Value mem = getI64Constant(rewriter, op.getLoc(), getDsbMemImmediate(op.getMem().getKind())); + rewriter.create(op.getLoc(), calleeName, TypeRange{}, ValueRange{mem}); + state.plannedDecls.push_back(PlannedDecl{calleeName.str(), funcType}); + rewriter.eraseOp(op); + return success(); + } + +private: + LoweringState &state; +}; + +class LowerDcciOpPattern final : public OpConversionPattern { +public: + explicit LowerDcciOpPattern(TypeConverter &typeConverter, MLIRContext *context, LoweringState &state) + : OpConversionPattern(typeConverter, context), state(state) {} + + LogicalResult matchAndRewrite(pto::DcciOp op, OpAdaptor adaptor, ConversionPatternRewriter &rewriter) const override { + auto ptrType = dyn_cast(adaptor.getPtr().getType()); + if (!ptrType) { + return rewriter.notifyMatchFailure(op, "expected LLVM pointer operand"); + } + + bool hasDst = static_cast(op.getDstAttr()); + StringRef calleeName = buildDcciCallee(ptrType.getAddressSpace(), hasDst, op.getContext()); + + Type i64Ty = rewriter.getI64Type(); + SmallVector argTypes{ptrType, i64Ty}; + SmallVector args{adaptor.getPtr(), + getI64Constant(rewriter, op.getLoc(), getDcciCacheLineImmediate(op.getCache().getKind()))}; + if (auto dst = op.getDstAttr()) { + argTypes.push_back(i64Ty); + args.push_back(getI64Constant(rewriter, op.getLoc(), getDcciDstImmediate(dst.getKind()))); + } + + auto funcType = rewriter.getFunctionType(argTypes, TypeRange{}); + rewriter.create(op.getLoc(), calleeName, TypeRange{}, args); + state.plannedDecls.push_back(PlannedDecl{calleeName.str(), funcType}); + rewriter.eraseOp(op); + return success(); + } + +private: + LoweringState &state; +}; + +template class LowerBufSyncOpPattern final : public OpConversionPattern { +public: + explicit LowerBufSyncOpPattern(TypeConverter &typeConverter, MLIRContext *context, LoweringState &state) + : OpConversionPattern(typeConverter, context), state(state) {} + + LogicalResult matchAndRewrite(BufSyncOp op, typename BufSyncOp::Adaptor adaptor, + ConversionPatternRewriter &rewriter) const override { + (void)adaptor; + PIPE pipe = PIPE::PIPE_UNASSIGNED; + if (auto pipeAttr = dyn_cast(op.getOpTypeAttr())) { + pipe = pipeAttr.getPipe(); + } else { + auto opTypeOr = parseSyncOpTypeLikeAttr(op.getOpTypeAttr()); + if (failed(opTypeOr)) { + return rewriter.notifyMatchFailure(op, "buffer sync expects pipe/sync_op_type/pipe_event_type attr"); + } + pipe = mapSyncOpTypeToPipe(*opTypeOr); + } + if (!isConcreteSyncPipe(pipe)) { + return rewriter.notifyMatchFailure(op, "buffer sync op_type cannot map to concrete pipe"); + } + + auto pipeImm = parsePipeImmediate(stringifyPIPE(pipe)); + if (!pipeImm) { + return rewriter.notifyMatchFailure(op, "unsupported buffer sync pipe"); + } + + StringRef calleeName = buildSyncCallee(op.getContext()); + Value pipeValue = getI64Constant(rewriter, op.getLoc(), *pipeImm); + Value bufIdValue = getI64Constant(rewriter, op.getLoc(), op.getBufIdAttr().getInt()); + Value modeValue = getI64Constant(rewriter, op.getLoc(), op.getModeAttr().getInt()); + auto funcType = rewriter.getFunctionType( + TypeRange{rewriter.getI64Type(), rewriter.getI64Type(), rewriter.getI64Type()}, TypeRange{}); + rewriter.create(op.getLoc(), calleeName, TypeRange{}, ValueRange{pipeValue, bufIdValue, modeValue}); + state.plannedDecls.push_back(PlannedDecl{calleeName.str(), funcType}); + rewriter.eraseOp(op); + return success(); + } + +private: + LoweringState &state; +}; + +template class LowerBufDynSyncOpPattern final : public OpConversionPattern { +public: + explicit LowerBufDynSyncOpPattern(TypeConverter &typeConverter, MLIRContext *context, LoweringState &state) + : OpConversionPattern(typeConverter, context), state(state) {} + + LogicalResult matchAndRewrite(BufDynSyncOp op, typename BufDynSyncOp::Adaptor adaptor, + ConversionPatternRewriter &rewriter) const override { + PIPE pipe = PIPE::PIPE_UNASSIGNED; + if (auto pipeAttr = dyn_cast(op.getOpTypeAttr())) { + pipe = pipeAttr.getPipe(); + } else { + auto opTypeOr = parseSyncOpTypeLikeAttr(op.getOpTypeAttr()); + if (failed(opTypeOr)) { + return rewriter.notifyMatchFailure(op, "buffer sync expects pipe/sync_op_type/pipe_event_type attr"); + } + pipe = mapSyncOpTypeToPipe(*opTypeOr); + } + if (!isConcreteSyncPipe(pipe)) { + return rewriter.notifyMatchFailure(op, "buffer sync op_type cannot map to concrete pipe"); + } + + auto pipeImm = parsePipeImmediate(stringifyPIPE(pipe)); + if (!pipeImm) { + return rewriter.notifyMatchFailure(op, "unsupported buffer sync pipe"); + } + + Value pipeValue = getI64Constant(rewriter, op.getLoc(), *pipeImm); + Value bufIdDyn = adaptor.getBufId(); + if (!bufIdDyn) { + return rewriter.notifyMatchFailure(op, "expected dynamic buf-id operand"); + } + Value bufIdValue = castIntegerLikeTo(op, bufIdDyn, rewriter.getI64Type()); + if (!bufIdValue) { + return rewriter.notifyMatchFailure(op, "failed to cast dynamic buf-id to i64"); + } + + bool isGetBuf = std::is_same_v; + StringRef calleeName = buildBufDynSyncCallee(op.getContext(), isGetBuf); + Value modeValue = getI64Constant(rewriter, op.getLoc(), op.getModeAttr().getInt()); + auto funcType = rewriter.getFunctionType( + TypeRange{rewriter.getI64Type(), rewriter.getI64Type(), rewriter.getI64Type()}, TypeRange{}); + rewriter.create(op.getLoc(), calleeName, TypeRange{}, ValueRange{pipeValue, bufIdValue, modeValue}); + state.plannedDecls.push_back(PlannedDecl{calleeName.str(), funcType}); + rewriter.eraseOp(op); + return success(); + } + +private: + LoweringState &state; +}; + +template class LowerRuntimeQueryOpPattern final : public OpConversionPattern { +public: + explicit LowerRuntimeQueryOpPattern(TypeConverter &typeConverter, MLIRContext *context, LoweringState &state) + : OpConversionPattern(typeConverter, context), state(state) {} + + LogicalResult matchAndRewrite(QueryOp op, typename QueryOp::Adaptor adaptor, + ConversionPatternRewriter &rewriter) const override { + (void)adaptor; + Type resultType = this->getTypeConverter()->convertType(op.getResult().getType()); + if (!resultType) { + return rewriter.notifyMatchFailure(op, "failed to convert runtime-query result type"); + } + + StringRef calleeName = buildRuntimeQueryCallee(op.getContext()); + auto funcType = rewriter.getFunctionType(TypeRange{}, TypeRange{resultType}); + auto call = rewriter.create(op.getLoc(), calleeName, TypeRange{resultType}, ValueRange{}); + state.plannedDecls.push_back(PlannedDecl{calleeName.str(), funcType}); + rewriter.replaceOp(op, call.getResults()); + return success(); + } + +private: + LoweringState &state; +}; + +template class LowerBlockRuntimeQueryOpPattern final : public OpConversionPattern { +public: + explicit LowerBlockRuntimeQueryOpPattern(TypeConverter &typeConverter, MLIRContext *context, LoweringState &state) + : OpConversionPattern(typeConverter, context), state(state) {} + + LogicalResult matchAndRewrite(QueryOp op, typename QueryOp::Adaptor adaptor, + ConversionPatternRewriter &rewriter) const override { + (void)adaptor; + Type resultType = this->getTypeConverter()->convertType(op.getResult().getType()); + if (!resultType) { + return rewriter.notifyMatchFailure(op, "failed to convert block runtime-query result type"); + } + + auto funcOp = op->template getParentOfType(); + bool isSimtEntry = funcOp && funcOp->hasAttr(pto::kPTOSimtEntryAttrName); + if (isSimtEntry && !resultType.isInteger(64)) { + return rewriter.notifyMatchFailure(op, "SIMT block runtime-query expects an i64 PTO result"); + } + + StringRef calleeName = isSimtEntry ? buildSimtBlockQueryCallee(op.getContext()) + : buildRuntimeQueryCallee(op.getContext()); + Type callResultType = isSimtEntry ? rewriter.getI32Type() : resultType; + auto funcType = rewriter.getFunctionType(TypeRange{}, TypeRange{callResultType}); + auto call = rewriter.create(op.getLoc(), calleeName, TypeRange{callResultType}, ValueRange{}); + state.plannedDecls.push_back(PlannedDecl{calleeName.str(), funcType}); + + Value result = call.getResult(0); + if (isSimtEntry) { + result = rewriter.create(op.getLoc(), resultType, result); + } + rewriter.replaceOp(op, result); + return success(); + } + +private: + LoweringState &state; +}; + +template class LowerVoteOpPattern final : public OpConversionPattern { +public: + explicit LowerVoteOpPattern(TypeConverter &typeConverter, MLIRContext *context, LoweringState &state) + : OpConversionPattern(typeConverter, context), state(state) {} + + LogicalResult matchAndRewrite(VoteOp op, typename VoteOp::Adaptor adaptor, + ConversionPatternRewriter &rewriter) const override { + Type resultType = this->getTypeConverter()->convertType(op.getResult().getType()); + if (!resultType) { + return rewriter.notifyMatchFailure(op, "failed to convert vote result type"); + } + + Type predType = this->getTypeConverter()->convertType(op.getPred().getType()); + if (!predType || predType != rewriter.getI1Type()) { + return rewriter.notifyMatchFailure(op, "failed to convert vote predicate type"); + } + + StringRef calleeName = buildVoteCallee(op.getContext()); + auto funcType = rewriter.getFunctionType(TypeRange{predType}, TypeRange{resultType}); + auto call = + rewriter.create(op.getLoc(), calleeName, TypeRange{resultType}, ValueRange{adaptor.getPred()}); + state.plannedDecls.push_back(PlannedDecl{calleeName.str(), funcType}); + rewriter.replaceOp(op, call.getResults()); + return success(); + } + +private: + LoweringState &state; +}; + +template class LowerShuffleOpPattern final : public OpConversionPattern { +public: + explicit LowerShuffleOpPattern(TypeConverter &typeConverter, MLIRContext *context, LoweringState &state) + : OpConversionPattern(typeConverter, context), state(state) {} + + LogicalResult matchAndRewrite(ShuffleOp op, typename ShuffleOp::Adaptor adaptor, + ConversionPatternRewriter &rewriter) const override { + Type resultType = this->getTypeConverter()->convertType(op.getResult().getType()); + if (!resultType) { + return rewriter.notifyMatchFailure(op, "failed to convert shuffle result type"); + } + + Type valueType = this->getTypeConverter()->convertType(op.getValue().getType()); + if (!valueType || valueType != resultType) { + return rewriter.notifyMatchFailure(op, "unexpected converted shuffle operand type"); + } + + FailureOr calleeName = buildShuffleCallee(op.getContext(), op.getValue().getType()); + if (failed(calleeName)) { + return rewriter.notifyMatchFailure(op, "unsupported shuffle VPTO signature"); + } + + IntegerAttr widthAttr = op.getWidthAttr(); + Value controlValue; + unsigned controlMask = 0; + if constexpr (std::is_same_v) { + controlValue = adaptor.getIndex(); + controlMask = 0x1f; + } else if constexpr (std::is_same_v) { + controlValue = adaptor.getOffset(); + controlMask = 0; + } else if constexpr (std::is_same_v) { + controlValue = adaptor.getOffset(); + controlMask = 0x1f; + } else if constexpr (std::is_same_v) { + controlValue = adaptor.getMask(); + controlMask = 0x1f; + } + if (!controlValue) { + return rewriter.notifyMatchFailure(op, "missing shuffle control operand"); + } + + Value control = buildShuffleControlValue(rewriter, op.getLoc(), controlValue, widthAttr.getInt(), controlMask); + + Type i32Type = rewriter.getI32Type(); + auto funcType = rewriter.getFunctionType(TypeRange{resultType, i32Type}, TypeRange{resultType}); + auto call = rewriter.create(op.getLoc(), *calleeName, TypeRange{resultType}, + ValueRange{adaptor.getValue(), control}); + state.plannedDecls.push_back(PlannedDecl{calleeName->str(), funcType}); + rewriter.replaceOp(op, call.getResults()); + return success(); + } + +private: + LoweringState &state; +}; + +template class LowerReduxOpPattern final : public OpConversionPattern { +public: + explicit LowerReduxOpPattern(TypeConverter &typeConverter, MLIRContext *context, LoweringState &state) + : OpConversionPattern(typeConverter, context), state(state) {} + + LogicalResult matchAndRewrite(ReduxOp op, typename ReduxOp::Adaptor adaptor, + ConversionPatternRewriter &rewriter) const override { + Type resultType = this->getTypeConverter()->convertType(op.getResult().getType()); + if (!resultType) { + return rewriter.notifyMatchFailure(op, "failed to convert redux result type"); + } + + Type valueType = this->getTypeConverter()->convertType(op.getValue().getType()); + if (!valueType || valueType != resultType) { + return rewriter.notifyMatchFailure(op, "unexpected converted redux operand type"); + } + + FailureOr calleeName = + buildReduxCallee(op.getContext(), op.getValue().getType(), op.getSignednessAttr()); + if (failed(calleeName)) { + return rewriter.notifyMatchFailure(op, "unsupported redux VPTO signature"); + } + + auto funcType = rewriter.getFunctionType(TypeRange{resultType}, TypeRange{resultType}); + auto call = + rewriter.create(op.getLoc(), *calleeName, TypeRange{resultType}, ValueRange{adaptor.getValue()}); + state.plannedDecls.push_back(PlannedDecl{calleeName->str(), funcType}); + rewriter.replaceOp(op, call.getResults()); + return success(); + } + +private: + LoweringState &state; +}; + +template class LowerAtomicBinaryOpPattern final : public OpConversionPattern { +public: + explicit LowerAtomicBinaryOpPattern(TypeConverter &typeConverter, MLIRContext *context, LoweringState &state) + : OpConversionPattern(typeConverter, context), state(state) {} + + LogicalResult matchAndRewrite(AtomicOp op, typename AtomicOp::Adaptor adaptor, + ConversionPatternRewriter &rewriter) const override { + Type resultType = this->getTypeConverter()->convertType(op.getOld().getType()); + Type valueType = this->getTypeConverter()->convertType(op.getValue().getType()); + if (!resultType || !valueType || resultType != valueType) { + return rewriter.notifyMatchFailure(op, "unexpected atomic operand/result type"); + } + + Type ptrType = this->getTypeConverter()->convertType(op.getPtr().getType()); + if (!ptrType) { + return rewriter.notifyMatchFailure(op, "failed to convert atomic pointer type"); + } + + FailureOr calleeName = buildAtomicCallee(op.getContext(), op.getPtr().getType(), + op.getValue().getType(), op.getSignednessAttr()); + if (failed(calleeName)) { + return rewriter.notifyMatchFailure(op, "unsupported atomic VPTO signature"); + } + + auto funcType = + rewriter.getFunctionType(TypeRange{ptrType, valueType, rewriter.getI32Type()}, TypeRange{resultType}); + Value modeValue = getI32Constant( + rewriter, op.getLoc(), + static_cast(op.getL2cacheAttr() ? op.getL2cacheAttr().getValue() : pto::StL2Cache::NMFV)); + auto call = rewriter.create(op.getLoc(), *calleeName, TypeRange{resultType}, + ValueRange{adaptor.getPtr(), adaptor.getValue(), modeValue}); + state.plannedDecls.push_back(PlannedDecl{calleeName->str(), funcType}); + rewriter.replaceOp(op, call.getResults()); + return success(); + } + +private: + LoweringState &state; +}; + +class LowerAtomicCasOpPattern final : public OpConversionPattern { +public: + explicit LowerAtomicCasOpPattern(TypeConverter &typeConverter, MLIRContext *context, LoweringState &state) + : OpConversionPattern(typeConverter, context), state(state) {} + + LogicalResult matchAndRewrite(pto::AtomicCasOp op, pto::AtomicCasOp::Adaptor adaptor, + ConversionPatternRewriter &rewriter) const override { + Type resultType = this->getTypeConverter()->convertType(op.getOld().getType()); + Type compareType = this->getTypeConverter()->convertType(op.getCompare().getType()); + Type valueType = this->getTypeConverter()->convertType(op.getValue().getType()); + if (!resultType || !compareType || !valueType || resultType != compareType || resultType != valueType) { + return rewriter.notifyMatchFailure(op, "unexpected atomic CAS type"); + } + + Type ptrType = this->getTypeConverter()->convertType(op.getPtr().getType()); + if (!ptrType) { + return rewriter.notifyMatchFailure(op, "failed to convert atomic pointer type"); + } + + FailureOr calleeName = buildAtomicCallee( + op.getContext(), op.getPtr().getType(), op.getValue().getType(), op.getSignednessAttr()); + if (failed(calleeName)) { + return rewriter.notifyMatchFailure(op, "unsupported atomic CAS signature"); + } + + auto funcType = rewriter.getFunctionType(TypeRange{ptrType, compareType, valueType, rewriter.getI32Type()}, + TypeRange{resultType}); + Value modeValue = getI32Constant( + rewriter, op.getLoc(), + static_cast(op.getL2cacheAttr() ? op.getL2cacheAttr().getValue() : pto::StL2Cache::NMFV)); + auto call = rewriter.create( + op.getLoc(), *calleeName, TypeRange{resultType}, + ValueRange{adaptor.getPtr(), adaptor.getCompare(), adaptor.getValue(), modeValue}); + state.plannedDecls.push_back(PlannedDecl{calleeName->str(), funcType}); + rewriter.replaceOp(op, call.getResults()); + return success(); + } + +private: + LoweringState &state; +}; + +template class LowerScalarIntrinsicOpPattern final : public OpConversionPattern { +public: + explicit LowerScalarIntrinsicOpPattern(TypeConverter &typeConverter, MLIRContext *context, LoweringState &state) + : OpConversionPattern(typeConverter, context), state(state) {} + + LogicalResult matchAndRewrite(ScalarOp op, typename ScalarOp::Adaptor adaptor, + ConversionPatternRewriter &rewriter) const override { + SmallVector resultTypes; + if (failed(this->getTypeConverter()->convertTypes(op->getResultTypes(), resultTypes))) { + return rewriter.notifyMatchFailure(op, "failed to convert scalar result types"); + } + + SmallVector operandTypes; + if (failed(this->getTypeConverter()->convertTypes(op->getOperandTypes(), operandTypes))) { + return rewriter.notifyMatchFailure(op, "failed to convert scalar operand types"); + } + + StringRef calleeName = buildScalarIntrinsicCallee(op.getContext()); + auto funcType = rewriter.getFunctionType(operandTypes, resultTypes); + auto call = rewriter.create(op.getLoc(), calleeName, resultTypes, adaptor.getOperands()); + state.plannedDecls.push_back(PlannedDecl{calleeName.str(), funcType}); + rewriter.replaceOp(op, call.getResults()); + return success(); + } + +private: + LoweringState &state; +}; + +class LowerMulhiOpPattern final : public OpConversionPattern { +public: + explicit LowerMulhiOpPattern(TypeConverter &typeConverter, MLIRContext *context, LoweringState &state) + : OpConversionPattern(typeConverter, context), state(state) {} + + LogicalResult matchAndRewrite(pto::MulhiOp op, OpAdaptor adaptor, + ConversionPatternRewriter &rewriter) const override { + Type resultType = getTypeConverter()->convertType(op.getResult().getType()); + Type lhsType = getTypeConverter()->convertType(op.getLhs().getType()); + Type rhsType = getTypeConverter()->convertType(op.getRhs().getType()); + if (!resultType || !lhsType || !rhsType || lhsType != resultType || rhsType != resultType) { + return rewriter.notifyMatchFailure(op, "unexpected mulhi type"); + } + + pto::Signedness signedness = op.getSignednessAttr().getValue(); + FailureOr calleeName = buildMulhiCallee(op.getContext(), op.getResult().getType(), signedness); + if (succeeded(calleeName)) { + auto funcType = rewriter.getFunctionType(TypeRange{lhsType, rhsType}, TypeRange{resultType}); + auto call = rewriter.create(op.getLoc(), *calleeName, TypeRange{resultType}, + ValueRange{adaptor.getLhs(), adaptor.getRhs()}); + state.plannedDecls.push_back(PlannedDecl{calleeName->str(), funcType}); + rewriter.replaceOp(op, call.getResults()); + return success(); + } + + if (!op.getResult().getType().isInteger(64) || signedness != pto::Signedness::Signed) { + return rewriter.notifyMatchFailure(op, "unsupported mulhi signature"); + } + + FailureOr unsignedCalleeName = + buildMulhiCallee(op.getContext(), op.getResult().getType(), pto::Signedness::Unsigned); + if (failed(unsignedCalleeName)) { + return rewriter.notifyMatchFailure(op, "unsupported mul64hi signature"); + } + + auto funcType = rewriter.getFunctionType(TypeRange{lhsType, rhsType}, TypeRange{resultType}); + auto unsignedCall = rewriter.create(op.getLoc(), *unsignedCalleeName, TypeRange{resultType}, + ValueRange{adaptor.getLhs(), adaptor.getRhs()}); + state.plannedDecls.push_back(PlannedDecl{unsignedCalleeName->str(), funcType}); + + Value zero = getI64Constant(rewriter, op.getLoc(), 0); + Value lhsNeg = rewriter.create(op.getLoc(), LLVM::ICmpPredicate::slt, adaptor.getLhs(), zero); + Value rhsNeg = rewriter.create(op.getLoc(), LLVM::ICmpPredicate::slt, adaptor.getRhs(), zero); + Value subRhs = rewriter.create(op.getLoc(), unsignedCall.getResult(0), adaptor.getRhs()); + Value correctedLhs = + rewriter.create(op.getLoc(), resultType, lhsNeg, subRhs, unsignedCall.getResult(0)); + Value subLhs = rewriter.create(op.getLoc(), correctedLhs, adaptor.getLhs()); + Value corrected = rewriter.create(op.getLoc(), resultType, rhsNeg, subLhs, correctedLhs); + rewriter.replaceOp(op, corrected); + return success(); + } + +private: + LoweringState &state; +}; + +class LowerMulI32ToI64OpPattern final : public OpConversionPattern { +public: + explicit LowerMulI32ToI64OpPattern(TypeConverter &typeConverter, MLIRContext *context, LoweringState &state) + : OpConversionPattern(typeConverter, context), state(state) {} + + LogicalResult matchAndRewrite(pto::MulI32ToI64Op op, OpAdaptor adaptor, + ConversionPatternRewriter &rewriter) const override { + Type resultType = getTypeConverter()->convertType(op.getResult().getType()); + Type lhsType = getTypeConverter()->convertType(op.getLhs().getType()); + Type rhsType = getTypeConverter()->convertType(op.getRhs().getType()); + if (!resultType || !lhsType || !rhsType) { + return rewriter.notifyMatchFailure(op, "unexpected mul_i32toi64 type"); + } + + FailureOr calleeName = buildMulI32ToI64Callee(op.getContext(), op.getSignednessAttr().getValue()); + if (failed(calleeName)) { + return rewriter.notifyMatchFailure(op, "unsupported mul_i32toi64 signature"); + } + + auto funcType = rewriter.getFunctionType(TypeRange{lhsType, rhsType}, TypeRange{resultType}); + auto call = rewriter.create(op.getLoc(), *calleeName, TypeRange{resultType}, + ValueRange{adaptor.getLhs(), adaptor.getRhs()}); + state.plannedDecls.push_back(PlannedDecl{calleeName->str(), funcType}); + rewriter.replaceOp(op, call.getResults()); + return success(); + } + +private: + LoweringState &state; +}; + +class LowerSqrtOpPattern final : public OpConversionPattern { +public: + explicit LowerSqrtOpPattern(TypeConverter &typeConverter, MLIRContext *context, LoweringState &state) + : OpConversionPattern(typeConverter, context), state(state) {} + + LogicalResult matchAndRewrite(pto::SqrtOp op, pto::SqrtOp::Adaptor adaptor, + ConversionPatternRewriter &rewriter) const override { + Type resultType = this->getTypeConverter()->convertType(op.getResult().getType()); + Type valueType = this->getTypeConverter()->convertType(op.getValue().getType()); + if (!resultType || !valueType || valueType != resultType) { + return rewriter.notifyMatchFailure(op, "unexpected sqrt operand/result type"); + } + + FailureOr calleeName = buildSqrtCallee(op.getContext(), op.getValue().getType()); + if (failed(calleeName)) { + return rewriter.notifyMatchFailure(op, "unsupported sqrt VPTO signature"); + } + + auto funcType = rewriter.getFunctionType(TypeRange{valueType}, TypeRange{resultType}); + auto call = + rewriter.create(op.getLoc(), *calleeName, TypeRange{resultType}, ValueRange{adaptor.getValue()}); + state.plannedDecls.push_back(PlannedDecl{calleeName->str(), funcType}); + rewriter.replaceOp(op, call.getResults()); + return success(); + } + +private: + LoweringState &state; +}; + +template class LowerUnaryScalarMathOpPattern final : public OpConversionPattern { +public: + explicit LowerUnaryScalarMathOpPattern(TypeConverter &typeConverter, MLIRContext *context, LoweringState &state) + : OpConversionPattern(typeConverter, context), state(state) {} + + LogicalResult matchAndRewrite(UnaryOp op, typename UnaryOp::Adaptor adaptor, + ConversionPatternRewriter &rewriter) const override { + Type resultType = this->getTypeConverter()->convertType(op.getResult().getType()); + Type valueType = this->getTypeConverter()->convertType(op.getValue().getType()); + if (!resultType || !valueType || valueType != resultType) { + return rewriter.notifyMatchFailure(op, "unexpected unary scalar math type"); + } + + FailureOr calleeName = buildUnaryScalarMathCallee(op.getContext(), op.getValue().getType()); + if (failed(calleeName)) { + return rewriter.notifyMatchFailure(op, "unsupported unary scalar math signature"); + } + + auto funcType = rewriter.getFunctionType(TypeRange{valueType}, TypeRange{resultType}); + auto call = + rewriter.create(op.getLoc(), *calleeName, TypeRange{resultType}, ValueRange{adaptor.getValue()}); + state.plannedDecls.push_back(PlannedDecl{calleeName->str(), funcType}); + rewriter.replaceOp(op, call.getResults()); + return success(); + } + +private: + LoweringState &state; +}; + +template class LowerBinaryScalarMathOpPattern final : public OpConversionPattern { +public: + explicit LowerBinaryScalarMathOpPattern(TypeConverter &typeConverter, MLIRContext *context, LoweringState &state) + : OpConversionPattern(typeConverter, context), state(state) {} + + LogicalResult matchAndRewrite(BinaryOp op, typename BinaryOp::Adaptor adaptor, + ConversionPatternRewriter &rewriter) const override { + Type resultType = this->getTypeConverter()->convertType(op.getResult().getType()); + Type lhsType = this->getTypeConverter()->convertType(op.getLhs().getType()); + Type rhsType = this->getTypeConverter()->convertType(op.getRhs().getType()); + if (!resultType || !lhsType || !rhsType || lhsType != rhsType || lhsType != resultType) { + return rewriter.notifyMatchFailure(op, "unexpected binary scalar math type"); + } + + FailureOr calleeName = buildBinaryScalarMathCallee(op.getContext(), op.getLhs().getType()); + if (failed(calleeName)) { + return rewriter.notifyMatchFailure(op, "unsupported binary scalar math signature"); + } + + auto funcType = rewriter.getFunctionType(TypeRange{lhsType, rhsType}, TypeRange{resultType}); + auto call = rewriter.create(op.getLoc(), *calleeName, TypeRange{resultType}, + ValueRange{adaptor.getLhs(), adaptor.getRhs()}); + state.plannedDecls.push_back(PlannedDecl{calleeName->str(), funcType}); + rewriter.replaceOp(op, call.getResults()); + return success(); + } + +private: + LoweringState &state; +}; + +class LowerFmaOpPattern final : public OpConversionPattern { +public: + explicit LowerFmaOpPattern(TypeConverter &typeConverter, MLIRContext *context, LoweringState &state) + : OpConversionPattern(typeConverter, context), state(state) {} + + LogicalResult matchAndRewrite(pto::FmaOp op, pto::FmaOp::Adaptor adaptor, + ConversionPatternRewriter &rewriter) const override { + Type resultType = this->getTypeConverter()->convertType(op.getResult().getType()); + Type lhsType = this->getTypeConverter()->convertType(op.getLhs().getType()); + Type rhsType = this->getTypeConverter()->convertType(op.getRhs().getType()); + Type accType = this->getTypeConverter()->convertType(op.getAcc().getType()); + if (!resultType || !lhsType || !rhsType || !accType || lhsType != rhsType || lhsType != accType || + lhsType != resultType) { + return rewriter.notifyMatchFailure(op, "unexpected fma scalar math type"); + } + + FailureOr calleeName = buildFmaCallee(op.getContext(), op.getLhs().getType()); + if (failed(calleeName)) { + return rewriter.notifyMatchFailure(op, "unsupported fma scalar signature"); + } + + auto funcType = rewriter.getFunctionType(TypeRange{lhsType, rhsType, accType}, TypeRange{resultType}); + auto call = rewriter.create(op.getLoc(), *calleeName, TypeRange{resultType}, + ValueRange{adaptor.getLhs(), adaptor.getRhs(), adaptor.getAcc()}); + state.plannedDecls.push_back(PlannedDecl{calleeName->str(), funcType}); + rewriter.replaceOp(op, call.getResults()); + return success(); + } + +private: + LoweringState &state; +}; + +class LowerConvertOpPattern final : public OpConversionPattern { +public: + explicit LowerConvertOpPattern(TypeConverter &typeConverter, MLIRContext *context, LoweringState &state) + : OpConversionPattern(typeConverter, context), state(state) {} + + LogicalResult matchAndRewrite(pto::ConvertOp op, pto::ConvertOp::Adaptor adaptor, + ConversionPatternRewriter &rewriter) const override { + Type resultType = getTypeConverter()->convertType(op.getDst().getType()); + if (!resultType) { + return rewriter.notifyMatchFailure(op, "failed to convert result type"); + } + + FailureOr calleeName = + buildConvertCallee(op.getContext(), op.getSrc().getType(), op.getDst().getType(), op.getSignednessAttr()); + if (failed(calleeName)) { + return rewriter.notifyMatchFailure(op, "unsupported convert signature"); + } + + Value rounding = getI32Constant(rewriter, op.getLoc(), static_cast(op.getRounding())); + Value saturation = getI32Constant(rewriter, op.getLoc(), static_cast(op.getSaturation())); + + auto funcType = rewriter.getFunctionType( + TypeRange{adaptor.getSrc().getType(), rewriter.getI32Type(), rewriter.getI32Type()}, TypeRange{resultType}); + auto call = rewriter.create(op.getLoc(), *calleeName, TypeRange{resultType}, + ValueRange{adaptor.getSrc(), rounding, saturation}); + state.plannedDecls.push_back(PlannedDecl{calleeName->str(), funcType}); + rewriter.replaceOp(op, call.getResults()); + return success(); + } + +private: + LoweringState &state; +}; + +class LowerGetVms4SrOpPattern final : public OpConversionPattern { +public: + explicit LowerGetVms4SrOpPattern(TypeConverter &typeConverter, MLIRContext *context, LoweringState &state) + : OpConversionPattern(typeConverter, context), state(state) {} + + LogicalResult matchAndRewrite(pto::GetVms4SrOp op, pto::GetVms4SrOp::Adaptor adaptor, + ConversionPatternRewriter &rewriter) const override { + (void)adaptor; + SmallVector resultTypes; + if (failed(this->getTypeConverter()->convertTypes(op->getResultTypes(), resultTypes)) || resultTypes.size() != 4) { + return rewriter.notifyMatchFailure(op, "failed to convert get_vms4_sr result types"); + } + + StringRef calleeName = buildRuntimeQueryCallee(op.getContext()); + auto funcType = rewriter.getFunctionType(TypeRange{}, TypeRange{rewriter.getI64Type()}); + auto call = rewriter.create(op.getLoc(), calleeName, TypeRange{rewriter.getI64Type()}, ValueRange{}); + state.plannedDecls.push_back(PlannedDecl{calleeName.str(), funcType}); + + SmallVector counts; + counts.reserve(4); + Value raw = call.getResult(0); + for (unsigned i = 0; i < 4; ++i) { + Value shifted = raw; + if (i != 0) { + shifted = rewriter.create(op.getLoc(), raw, getI64Constant(rewriter, op.getLoc(), i * 16)); + } + counts.push_back(rewriter.create(op.getLoc(), resultTypes[i], shifted)); + } + rewriter.replaceOp(op, counts); + return success(); + } + +private: + LoweringState &state; +}; + +template class LowerBinaryI64PureOpPattern final : public OpConversionPattern { +public: + explicit LowerBinaryI64PureOpPattern(TypeConverter &typeConverter, MLIRContext *context, LoweringState &state) + : OpConversionPattern(typeConverter, context), state(state) {} + + LogicalResult matchAndRewrite(BinaryOp op, typename BinaryOp::Adaptor adaptor, + ConversionPatternRewriter &rewriter) const override { + Type resultType = this->getTypeConverter()->convertType(op.getResult().getType()); + if (!resultType) { + return rewriter.notifyMatchFailure(op, "failed to convert result type"); + } + + StringRef calleeName = buildBinaryI64PureCallee(op.getContext()); + auto funcType = rewriter.getFunctionType(TypeRange{adaptor.getFirst().getType(), adaptor.getSecond().getType()}, + TypeRange{resultType}); + auto call = rewriter.create(op.getLoc(), calleeName, TypeRange{resultType}, + ValueRange{adaptor.getFirst(), adaptor.getSecond()}); + state.plannedDecls.push_back(PlannedDecl{calleeName.str(), funcType}); + rewriter.replaceOp(op, call.getResults()); + return success(); + } + +private: + LoweringState &state; +}; + +static void populateVPTOSIMTAndScalarPatterns(VPTOTypeConverter &typeConverter, RewritePatternSet &patterns, + LoweringState &state) { + patterns.add< + LowerRuntimeQueryOpPattern, LowerGetVms4SrOpPattern, LowerRuntimeQueryOpPattern, + LowerRuntimeQueryOpPattern, LowerRuntimeQueryOpPattern, + LowerRuntimeQueryOpPattern, LowerRuntimeQueryOpPattern, + LowerRuntimeQueryOpPattern, LowerRuntimeQueryOpPattern, + LowerRuntimeQueryOpPattern, LowerRuntimeQueryOpPattern, + LowerRuntimeQueryOpPattern, LowerRuntimeQueryOpPattern, + LowerRuntimeQueryOpPattern, LowerRuntimeQueryOpPattern, + LowerRuntimeQueryOpPattern, LowerRuntimeQueryOpPattern, + LowerRuntimeQueryOpPattern, LowerRuntimeQueryOpPattern, + LowerRuntimeQueryOpPattern, LowerRuntimeQueryOpPattern, + LowerRuntimeQueryOpPattern, LowerRuntimeQueryOpPattern, + LowerVoteOpPattern, LowerVoteOpPattern, LowerVoteOpPattern, + LowerVoteOpPattern, LowerShuffleOpPattern, + LowerShuffleOpPattern, LowerShuffleOpPattern, + LowerShuffleOpPattern, LowerReduxOpPattern, + LowerReduxOpPattern, LowerReduxOpPattern, LowerAtomicCasOpPattern, + LowerAtomicBinaryOpPattern, LowerAtomicBinaryOpPattern, + LowerAtomicBinaryOpPattern, LowerAtomicBinaryOpPattern, + LowerAtomicBinaryOpPattern, LowerAtomicBinaryOpPattern, + LowerAtomicBinaryOpPattern, LowerAtomicBinaryOpPattern, LowerTrapOpPattern, + LowerScalarIntrinsicOpPattern, LowerMulhiOpPattern, LowerMulI32ToI64OpPattern, LowerSqrtOpPattern, + LowerUnaryScalarMathOpPattern, LowerUnaryScalarMathOpPattern, + LowerUnaryScalarMathOpPattern, LowerUnaryScalarMathOpPattern, + LowerUnaryScalarMathOpPattern, LowerUnaryScalarMathOpPattern, + LowerUnaryScalarMathOpPattern, LowerBinaryScalarMathOpPattern, + LowerBinaryScalarMathOpPattern, LowerBinaryScalarMathOpPattern, LowerFmaOpPattern, + LowerConvertOpPattern, LowerSimtFenceOpPattern, LowerSimtFenceOpPattern, + LowerSimtFenceOpPattern, LowerKeepOpPattern, LowerResumeOpPattern, + LowerBinaryI64PureOpPattern, LowerBinaryI64PureOpPattern, + LowerSetLoopConfigOpPattern, + LowerSetLoopConfigOpPattern, LowerSetLoopConfigOpPattern, + LowerSetLoopConfigOpPattern, + LowerSetLoopConfigOpPattern, LowerSetLoopConfigOpPattern, + LowerSetLoopConfigOpPattern, LowerSetLoopConfigOpPattern, + LowerUnaryI64ConfigOpPattern, LowerStoreVfSimtInfoOpPattern>(typeConverter, patterns.getContext(), + state); +} + +static void populateVPTOConfigAndSyncPatterns(VPTOTypeConverter &typeConverter, RewritePatternSet &patterns, + LoweringState &state) { + patterns + .add, LowerUnaryI64ConfigOpPattern, + LowerUnaryI64ConfigOpPattern, LowerUnaryI64ConfigOpPattern, + LowerUnaryI64ConfigOpPattern, + LowerUnaryI64ConfigOpPattern, + LowerUnaryI64ConfigOpPattern, LowerUnaryI64ConfigOpPattern, + LowerUnaryI64ConfigOpPattern, LowerUnaryI64ConfigOpPattern, + LowerUnaryI64ConfigOpPattern, LowerNullaryConfigOpPattern, + LowerNullaryConfigOpPattern, LowerPipeEventSyncOpPattern, + LowerPipeEventSyncOpPattern, LowerPipeEventDynSyncOpPattern, + LowerPipeEventDynSyncOpPattern, LowerBarrierOpPattern, LowerMemBarOpPattern, + LowerUnsupportedMemoryConsistencyOpPattern, + LowerUnsupportedMemoryConsistencyOpPattern, LowerDsbOpPattern, LowerDcciOpPattern, + LowerBufSyncOpPattern, LowerBufSyncOpPattern, + LowerBufDynSyncOpPattern, LowerBufDynSyncOpPattern, + LowerBlockRuntimeQueryOpPattern, LowerRuntimeQueryOpPattern, + LowerBlockRuntimeQueryOpPattern, LowerRuntimeQueryOpPattern, + LowerVtrcOpPattern, LowerVcvtOpPattern, LowerVbitcastOpPattern, LowerPbitcastOpPattern, + LowerPredicateLoadOpPattern, LowerPredicateLoadOpPattern, + LowerPredicateStoreOpPattern, LowerPredicateStoreOpPattern, + LowerInterCoreSyncOpPattern, LowerInterCoreSyncOpPattern, + LowerNamedSyncOpPattern, LowerNamedSyncOpPattern>( + typeConverter, patterns.getContext(), state); +} + +void populateVPTOScalarPatterns(VPTOTypeConverter &typeConverter, RewritePatternSet &patterns, LoweringState &state) { + populateVPTOSIMTAndScalarPatterns(typeConverter, patterns, state); + populateVPTOConfigAndSyncPatterns(typeConverter, patterns, state); +} + +} // namespace mlir::pto::detail diff --git a/lib/PTO/Transforms/VPTOCANN900LLVMEmitterTemplates.h b/lib/PTO/Transforms/VPTOCANN900LLVMEmitterTemplates.h new file mode 100644 index 0000000000..54a3c40490 --- /dev/null +++ b/lib/PTO/Transforms/VPTOCANN900LLVMEmitterTemplates.h @@ -0,0 +1,819 @@ +// Copyright (c) 2026 Huawei Technologies Co., Ltd. +// This program is free software, you can redistribute it and/or modify it under the terms and conditions of +// CANN Open Software License Agreement Version 2.0 (the "License"). +// Please refer to the License for details. You may not use this file except in compliance with the License. +// THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED, +// INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE. +// See LICENSE in the root of the software repository for the full text of the License. + +#pragma once + +#include "VPTOCANN900LLVMEmitterInternal.h" + +namespace mlir::pto::detail { + +template StringRef getUnaryMaskedStem() { + if constexpr (std::is_same_v) { + return "vabs"; + } + if constexpr (std::is_same_v) { + return "vexp"; + } + if constexpr (std::is_same_v) { + return "vln"; + } + if constexpr (std::is_same_v) { + return "vneg"; + } + if constexpr (std::is_same_v) { + return "vsqrt"; + } + if constexpr (std::is_same_v) { + return "vrelu"; + } + if constexpr (std::is_same_v) { + return "vnot"; + } + return {}; +} + +template FailureOr buildUnaryMaskedCallee(MLIRContext *context, Type resultType) { + StringRef stem = getUnaryMaskedStem(); + if (stem.empty()) { + return failure(); + } + return buildCANN900ModeTypedCallee(context, resultType, stem, "x"); +} + +template StringRef getBinaryMaskedStem() { + if constexpr (std::is_same_v) { + return "vadd"; + } + if constexpr (std::is_same_v) { + return "vsub"; + } + if constexpr (std::is_same_v) { + return "vmul"; + } + if constexpr (std::is_same_v) { + return "vdiv"; + } + if constexpr (std::is_same_v) { + return "vmax"; + } + if constexpr (std::is_same_v) { + return "vmin"; + } + if constexpr (std::is_same_v) { + return "vand"; + } + if constexpr (std::is_same_v) { + return "vor"; + } + if constexpr (std::is_same_v) { + return "vxor"; + } + if constexpr (std::is_same_v) { + return "vshl"; + } + if constexpr (std::is_same_v) { + return "vshr"; + } + if constexpr (std::is_same_v) { + return "vprelu"; + } + return {}; +} + +template StringRef getTernaryMaskedStem() { + if constexpr (std::is_same_v) { + return "vmadd"; + } + return {}; +} + +template constexpr bool usesSignedBinaryCANN900Callee() { + return !std::is_same_v && !std::is_same_v && + !std::is_same_v && !std::is_same_v; +} + +template constexpr bool usesSignedTernaryCANN900Callee() { return false; } + +template StringRef getCarryBinaryStem() { + if constexpr (std::is_same_v) { + return "vaddc"; + } + if constexpr (std::is_same_v) { + return "vsubc"; + } + if constexpr (std::is_same_v) { + return "vaddcs"; + } + if constexpr (std::is_same_v) { + return "vsubcs"; + } + return {}; +} + +template constexpr bool hasCarryInput() { + return std::is_same_v || std::is_same_v; +} + +template StringRef buildRuntimeQueryCallee(MLIRContext *context); + +template <> inline StringRef buildRuntimeQueryCallee(MLIRContext *context) { + return StringAttr::get(context, "llvm.hivm.GET.CTRL").getValue(); +} + +template <> inline StringRef buildRuntimeQueryCallee(MLIRContext *context) { + return StringAttr::get(context, "llvm.hivm.GET.VMS4.SR").getValue(); +} + +template <> inline StringRef buildRuntimeQueryCallee(MLIRContext *context) { + return StringAttr::get(context, "llvm.hivm.get.TID.X").getValue(); +} + +template <> inline StringRef buildRuntimeQueryCallee(MLIRContext *context) { + return StringAttr::get(context, "llvm.hivm.get.TID.Y").getValue(); +} + +template <> inline StringRef buildRuntimeQueryCallee(MLIRContext *context) { + return StringAttr::get(context, "llvm.hivm.get.TID.Z").getValue(); +} + +template <> inline StringRef buildRuntimeQueryCallee(MLIRContext *context) { + return StringAttr::get(context, "llvm.hivm.get.BLOCK.DIM.X").getValue(); +} + +template <> inline StringRef buildRuntimeQueryCallee(MLIRContext *context) { + return StringAttr::get(context, "llvm.hivm.get.BLOCK.DIM.Y").getValue(); +} + +template <> inline StringRef buildRuntimeQueryCallee(MLIRContext *context) { + return StringAttr::get(context, "llvm.hivm.get.BLOCK.DIM.Z").getValue(); +} + +template <> inline StringRef buildRuntimeQueryCallee(MLIRContext *context) { + return StringAttr::get(context, "llvm.hivm.get.GRID.DIM.X").getValue(); +} + +template <> inline StringRef buildRuntimeQueryCallee(MLIRContext *context) { + return StringAttr::get(context, "llvm.hivm.get.GRID.DIM.Y").getValue(); +} + +template <> inline StringRef buildRuntimeQueryCallee(MLIRContext *context) { + return StringAttr::get(context, "llvm.hivm.get.GRID.DIM.Z").getValue(); +} + +template <> inline StringRef buildRuntimeQueryCallee(MLIRContext *context) { + return StringAttr::get(context, "llvm.hivm.get.BLOCK.IDX.X").getValue(); +} + +template <> inline StringRef buildRuntimeQueryCallee(MLIRContext *context) { + return StringAttr::get(context, "llvm.hivm.get.BLOCK.IDX.Y").getValue(); +} + +template <> inline StringRef buildRuntimeQueryCallee(MLIRContext *context) { + return StringAttr::get(context, "llvm.hivm.get.BLOCK.IDX.Z").getValue(); +} + +template <> inline StringRef buildRuntimeQueryCallee(MLIRContext *context) { + return StringAttr::get(context, "llvm.hivm.tpe.get.VECCOREID").getValue(); +} + +template <> inline StringRef buildRuntimeQueryCallee(MLIRContext *context) { + return StringAttr::get(context, "llvm.hivm.get.laneID").getValue(); +} + +template <> inline StringRef buildRuntimeQueryCallee(MLIRContext *context) { + return StringAttr::get(context, "llvm.hivm.get.CLOCK32").getValue(); +} + +template <> inline StringRef buildRuntimeQueryCallee(MLIRContext *context) { + return StringAttr::get(context, "llvm.hivm.get.CLOCK64").getValue(); +} + +template <> inline StringRef buildRuntimeQueryCallee(MLIRContext *context) { + return StringAttr::get(context, "llvm.hivm.get.LANEMASK.EQ").getValue(); +} + +template <> inline StringRef buildRuntimeQueryCallee(MLIRContext *context) { + return StringAttr::get(context, "llvm.hivm.get.LANEMASK.LE").getValue(); +} + +template <> inline StringRef buildRuntimeQueryCallee(MLIRContext *context) { + return StringAttr::get(context, "llvm.hivm.get.LANEMASK.LT").getValue(); +} + +template <> inline StringRef buildRuntimeQueryCallee(MLIRContext *context) { + return StringAttr::get(context, "llvm.hivm.get.LANEMASK.GE").getValue(); +} + +template <> inline StringRef buildRuntimeQueryCallee(MLIRContext *context) { + return StringAttr::get(context, "llvm.hivm.get.LANEMASK.GT").getValue(); +} + +template StringRef buildSprStoreCallee(MLIRContext *context, bool post); + +template <> inline StringRef buildSprStoreCallee(MLIRContext *context, bool post) { + return buildSprstiCallee(context, post); +} + +template <> inline StringRef buildSprStoreCallee(MLIRContext *context, bool post) { + return buildSprstsCallee(context, post); +} + +template StringRef buildUnaryConfigCallee(MLIRContext *context); + +template <> inline StringRef buildUnaryConfigCallee(MLIRContext *context) { + return StringAttr::get(context, "llvm.hivm.SET.CTRL").getValue(); +} + +template StringRef buildVoteCallee(MLIRContext *context); + +template <> inline StringRef buildVoteCallee(MLIRContext *context) { + return StringAttr::get(context, "llvm.hivm.vote.all").getValue(); +} + +template <> inline StringRef buildVoteCallee(MLIRContext *context) { + return StringAttr::get(context, "llvm.hivm.vote.any").getValue(); +} + +template <> inline StringRef buildVoteCallee(MLIRContext *context) { + return StringAttr::get(context, "llvm.hivm.vote.uni").getValue(); +} + +template <> inline StringRef buildVoteCallee(MLIRContext *context) { + return StringAttr::get(context, "llvm.hivm.vote.ballot").getValue(); +} + +template StringRef buildBinaryI64PureCallee(MLIRContext *context); + +template <> inline StringRef buildBinaryI64PureCallee(MLIRContext *context) { + return StringAttr::get(context, "llvm.hivm.SBITSET0").getValue(); +} + +template <> inline StringRef buildBinaryI64PureCallee(MLIRContext *context) { + return StringAttr::get(context, "llvm.hivm.SBITSET1").getValue(); +} + +template FailureOr buildShuffleCallee(MLIRContext *context, Type valueType); + +template <> inline FailureOr buildShuffleCallee(MLIRContext *context, Type valueType) { + std::string elem = getShuffleIntrinsicTypeFragment(valueType); + if (elem.empty()) { + return failure(); + } + return StringAttr::get(context, "llvm.hivm.shfl.idx." + elem).getValue(); +} + +template <> inline FailureOr buildShuffleCallee(MLIRContext *context, Type valueType) { + std::string elem = getShuffleIntrinsicTypeFragment(valueType); + if (elem.empty()) { + return failure(); + } + return StringAttr::get(context, "llvm.hivm.shfl.up." + elem).getValue(); +} + +template <> inline FailureOr buildShuffleCallee(MLIRContext *context, Type valueType) { + std::string elem = getShuffleIntrinsicTypeFragment(valueType); + if (elem.empty()) { + return failure(); + } + return StringAttr::get(context, "llvm.hivm.shfl.down." + elem).getValue(); +} + +template <> inline FailureOr buildShuffleCallee(MLIRContext *context, Type valueType) { + std::string elem = getShuffleIntrinsicTypeFragment(valueType); + if (elem.empty()) { + return failure(); + } + return StringAttr::get(context, "llvm.hivm.shfl.bfly." + elem).getValue(); +} + +template +FailureOr buildReduxCallee(MLIRContext *context, Type valueType, Attribute signednessAttr); + +template <> +inline FailureOr buildReduxCallee(MLIRContext *context, Type valueType, + Attribute signednessAttr) { + std::string elem = getReduxIntrinsicTypeFragment(valueType, signednessAttr); + if (elem.empty()) { + return failure(); + } + return StringAttr::get(context, "llvm.hivm.redux.add." + elem).getValue(); +} + +template <> +inline FailureOr buildReduxCallee(MLIRContext *context, Type valueType, + Attribute signednessAttr) { + std::string elem = getReduxIntrinsicTypeFragment(valueType, signednessAttr); + if (elem.empty()) { + return failure(); + } + return StringAttr::get(context, "llvm.hivm.redux.max." + elem).getValue(); +} + +template <> +inline FailureOr buildReduxCallee(MLIRContext *context, Type valueType, + Attribute signednessAttr) { + std::string elem = getReduxIntrinsicTypeFragment(valueType, signednessAttr); + if (elem.empty()) { + return failure(); + } + return StringAttr::get(context, "llvm.hivm.redux.min." + elem).getValue(); +} + +template StringRef buildScalarIntrinsicCallee(MLIRContext *context); + +template <> inline StringRef buildScalarIntrinsicCallee(MLIRContext *context) { + return StringAttr::get(context, "llvm.hivm.prmt").getValue(); +} + +template FailureOr buildUnaryScalarMathCallee(MLIRContext *context, Type valueType); + +template <> inline FailureOr buildUnaryScalarMathCallee(MLIRContext *context, Type valueType) { + std::string elem = getLLVMFloatBuiltinFragment(valueType); + if (elem != "f16" && elem != "f32" && elem != "v2f16" && elem != "v2bf16") { + return failure(); + } + return StringAttr::get(context, "llvm.fabs." + elem).getValue(); +} + +template <> inline FailureOr buildUnaryScalarMathCallee(MLIRContext *context, Type valueType) { + std::string elem = getLLVMFloatBuiltinFragment(valueType); + if (elem != "f32" && elem != "f16" && elem != "v2f16") { + return failure(); + } + return StringAttr::get(context, "llvm.exp." + elem).getValue(); +} + +template <> inline FailureOr buildUnaryScalarMathCallee(MLIRContext *context, Type valueType) { + std::string elem = getLLVMFloatBuiltinFragment(valueType); + if (elem != "f32" && elem != "f16" && elem != "v2f16") { + return failure(); + } + return StringAttr::get(context, "llvm.log." + elem).getValue(); +} + +template <> inline FailureOr buildUnaryScalarMathCallee(MLIRContext *context, Type valueType) { + std::string elem = getScalarHIVMFloatShortFragment(valueType); + if (elem.empty()) { + return failure(); + } + return StringAttr::get(context, "llvm.hivm.ceil." + elem).getValue(); +} + +template <> inline FailureOr buildUnaryScalarMathCallee(MLIRContext *context, Type valueType) { + std::string elem = getScalarHIVMFloatShortFragment(valueType); + if (elem.empty()) { + return failure(); + } + return StringAttr::get(context, "llvm.hivm.floor." + elem).getValue(); +} + +template <> inline FailureOr buildUnaryScalarMathCallee(MLIRContext *context, Type valueType) { + std::string elem = getScalarHIVMFloatShortFragment(valueType); + if (elem.empty()) { + return failure(); + } + return StringAttr::get(context, "llvm.hivm.rint." + elem).getValue(); +} + +template <> inline FailureOr buildUnaryScalarMathCallee(MLIRContext *context, Type valueType) { + std::string elem = getScalarHIVMFloatShortFragment(valueType); + if (elem.empty()) { + return failure(); + } + return StringAttr::get(context, "llvm.hivm.round." + elem).getValue(); +} + +template FailureOr buildBinaryScalarMathCallee(MLIRContext *context, Type valueType); + +template <> inline FailureOr buildBinaryScalarMathCallee(MLIRContext *context, Type valueType) { + std::string elem = getLLVMFloatBuiltinFragment(valueType); + if (elem != "f16" && elem != "f32" && elem != "bf16" && elem != "v2f16" && elem != "v2bf16") { + return failure(); + } + return StringAttr::get(context, "llvm.minnum." + elem).getValue(); +} + +template <> inline FailureOr buildBinaryScalarMathCallee(MLIRContext *context, Type valueType) { + std::string elem = getLLVMFloatBuiltinFragment(valueType); + if (elem != "f16" && elem != "f32" && elem != "bf16" && elem != "v2f16" && elem != "v2bf16") { + return failure(); + } + return StringAttr::get(context, "llvm.maxnum." + elem).getValue(); +} + +template <> inline FailureOr buildBinaryScalarMathCallee(MLIRContext *context, Type valueType) { + std::string elem = getLLVMFloatBuiltinFragment(valueType); + if (elem != "f32" && elem != "f16" && elem != "v2f16") { + return failure(); + } + return StringAttr::get(context, "llvm.pow." + elem).getValue(); +} + +template StringRef getVecScalarMaskedStem() { + if constexpr (std::is_same_v) { + return "vmuls"; + } + if constexpr (std::is_same_v) { + return "vadds"; + } + if constexpr (std::is_same_v) { + return "vmaxs"; + } + if constexpr (std::is_same_v) { + return "vmins"; + } + if constexpr (std::is_same_v) { + return "vlrelu"; + } + if constexpr (std::is_same_v) { + return "vshls"; + } + if constexpr (std::is_same_v) { + return "vshrs"; + } + return {}; +} + +template constexpr bool usesSignedVecScalarCANN900Callee() { + return !std::is_same_v; +} + +template StringRef getReductionUnaryStem() { + if constexpr (std::is_same_v) { + return "vcadd"; + } + if constexpr (std::is_same_v) { + return "vcmax"; + } + if constexpr (std::is_same_v) { + return "vcmin"; + } + if constexpr (std::is_same_v) { + return "vcgadd"; + } + if constexpr (std::is_same_v) { + return "vcgmax"; + } + if constexpr (std::is_same_v) { + return "vcgmin"; + } + if constexpr (std::is_same_v) { + return "vcpadd"; + } + return {}; +} + +template StringRef getHistogramCallee(MLIRContext *context) { + if constexpr (std::is_same_v) { + return StringAttr::get(context, "llvm.hivm.chistv2.m").getValue(); + } + if constexpr (std::is_same_v) { + return StringAttr::get(context, "llvm.hivm.dhistv2.m").getValue(); + } + return {}; +} + +template StringRef getExtremaPredicateStem() { + if constexpr (std::is_same_v) { + return "vcbmax"; + } + if constexpr (std::is_same_v) { + return "vcbmin"; + } + return {}; +} + +template FailureOr buildExtremaPredicateCallee(MLIRContext *context, Type resultType) { + return buildCANN900SignedModeTypedCallee(context, resultType, getExtremaPredicateStem(), "x"); +} + +template constexpr bool usesSignedReductionCANN900Callee() { + return !std::is_same_v; +} + +template StringRef buildPredicatePairReorderCallee(MLIRContext *context); + +template <> inline StringRef buildPredicatePairReorderCallee(MLIRContext *context) { + return StringAttr::get(context, "llvm.hivm.pdintlv.b8").getValue(); +} + +template <> inline StringRef buildPredicatePairReorderCallee(MLIRContext *context) { + return StringAttr::get(context, "llvm.hivm.pdintlv.b16").getValue(); +} + +template <> inline StringRef buildPredicatePairReorderCallee(MLIRContext *context) { + return StringAttr::get(context, "llvm.hivm.pdintlv.b32").getValue(); +} + +template <> inline StringRef buildPredicatePairReorderCallee(MLIRContext *context) { + return StringAttr::get(context, "llvm.hivm.pintlv.b8").getValue(); +} + +template <> inline StringRef buildPredicatePairReorderCallee(MLIRContext *context) { + return StringAttr::get(context, "llvm.hivm.pintlv.b16").getValue(); +} + +template <> inline StringRef buildPredicatePairReorderCallee(MLIRContext *context) { + return StringAttr::get(context, "llvm.hivm.pintlv.b32").getValue(); +} + +template StringRef getPredicateStoreCallee(MLIRContext *context, bool post); + +template <> inline StringRef getPredicateStoreCallee(MLIRContext *context, bool post) { + return buildPstiCallee(context, post); +} + +template <> inline StringRef getPredicateStoreCallee(MLIRContext *context, bool post) { + return buildPstsCallee(context, post); +} + +template StringRef getPredicateLoadCallee(MLIRContext *context, bool post); + +template <> inline StringRef getPredicateLoadCallee(MLIRContext *context, bool post) { + return buildPldiCallee(context, post); +} + +template <> inline StringRef getPredicateLoadCallee(MLIRContext *context, bool post) { + return buildPldsCallee(context, post); +} + +template StringRef getPredicateMaskCallee(MLIRContext *context); + +template <> inline StringRef getPredicateMaskCallee(MLIRContext *context) { + return buildPnotCallee(context); +} + +template <> inline StringRef getPredicateMaskCallee(MLIRContext *context) { + return buildPselCallee(context); +} + +template <> inline StringRef getPredicateMaskCallee(MLIRContext *context) { + return buildPandCallee(context); +} + +template <> inline StringRef getPredicateMaskCallee(MLIRContext *context) { + return buildPorCallee(context); +} + +template <> inline StringRef getPredicateMaskCallee(MLIRContext *context) { + return buildPxorCallee(context); +} + +template StringRef getPredicatePackCallee(MLIRContext *context); + +template <> inline StringRef getPredicatePackCallee(MLIRContext *context) { + return buildPpackCallee(context); +} + +template <> inline StringRef getPredicatePackCallee(MLIRContext *context) { + return buildPunpackCallee(context); +} + +template StringRef buildPltCallee(MLIRContext *context); + +template <> inline StringRef buildPltCallee(MLIRContext *context) { + return StringAttr::get(context, "llvm.hivm.plt.b8.v300").getValue(); +} + +template <> inline StringRef buildPltCallee(MLIRContext *context) { + return StringAttr::get(context, "llvm.hivm.plt.b16.v300").getValue(); +} + +template <> inline StringRef buildPltCallee(MLIRContext *context) { + return StringAttr::get(context, "llvm.hivm.plt.b32.v300").getValue(); +} + +template StringRef buildPltmCallee(MLIRContext *context); + +template <> inline StringRef buildPltmCallee(MLIRContext *context) { + return StringAttr::get(context, "llvm.hivm.pltm.b8.v300").getValue(); +} + +template <> inline StringRef buildPltmCallee(MLIRContext *context) { + return StringAttr::get(context, "llvm.hivm.pltm.b16.v300").getValue(); +} + +template <> inline StringRef buildPltmCallee(MLIRContext *context) { + return StringAttr::get(context, "llvm.hivm.pltm.b32.v300").getValue(); +} + +template StringRef buildPsetCallee(MLIRContext *context); + +template <> inline StringRef buildPsetCallee(MLIRContext *context) { + return StringAttr::get(context, "llvm.hivm.pset.b8").getValue(); +} + +template <> inline StringRef buildPsetCallee(MLIRContext *context) { + return StringAttr::get(context, "llvm.hivm.pset.b16").getValue(); +} + +template <> inline StringRef buildPsetCallee(MLIRContext *context) { + return StringAttr::get(context, "llvm.hivm.pset.b32").getValue(); +} + +template StringRef buildPgeCallee(MLIRContext *context); + +template <> inline StringRef buildPgeCallee(MLIRContext *context) { + return StringAttr::get(context, "llvm.hivm.pge.b8").getValue(); +} + +template <> inline StringRef buildPgeCallee(MLIRContext *context) { + return StringAttr::get(context, "llvm.hivm.pge.b16").getValue(); +} + +template <> inline StringRef buildPgeCallee(MLIRContext *context) { + return StringAttr::get(context, "llvm.hivm.pge.b32").getValue(); +} + +template StringRef buildSetLoopCallee(MLIRContext *context); + +template StringRef buildUnaryConfigCallee(MLIRContext *context); + +template StringRef buildNullaryConfigCallee(MLIRContext *context); + +template <> inline StringRef buildSetLoopCallee(MLIRContext *context) { + return StringAttr::get(context, "llvm.hivm.SET.LOOP2.STRIDE.OUTTOUB").getValue(); +} + +template <> inline StringRef buildSetLoopCallee(MLIRContext *context) { + return StringAttr::get(context, "llvm.hivm.SET.LOOP1.STRIDE.OUTTOUB").getValue(); +} + +template <> inline StringRef buildSetLoopCallee(MLIRContext *context) { + return StringAttr::get(context, "llvm.hivm.SET.LOOP.SIZE.OUTTOUB").getValue(); +} + +template <> inline StringRef buildSetLoopCallee(MLIRContext *context) { + return StringAttr::get(context, "llvm.hivm.SET.LOOP2.STRIDE.UBTOOUT").getValue(); +} + +template <> inline StringRef buildSetLoopCallee(MLIRContext *context) { + return StringAttr::get(context, "llvm.hivm.SET.LOOP1.STRIDE.UBTOOUT").getValue(); +} + +template <> inline StringRef buildSetLoopCallee(MLIRContext *context) { + return StringAttr::get(context, "llvm.hivm.SET.LOOP.SIZE.UBTOOUT").getValue(); +} + +template <> inline StringRef buildSetLoopCallee(MLIRContext *context) { + return StringAttr::get(context, "llvm.hivm.SET.LOOP3.PARA").getValue(); +} + +template <> inline StringRef buildSetLoopCallee(MLIRContext *context) { + return StringAttr::get(context, "llvm.hivm.SET.CHANNEL.PARA").getValue(); +} + +template <> inline StringRef buildUnaryConfigCallee(MLIRContext *context) { + return StringAttr::get(context, "llvm.hivm.SET.MOV.PAD.VAL").getValue(); +} + +template <> inline StringRef buildUnaryConfigCallee(MLIRContext *context) { + return StringAttr::get(context, "llvm.hivm.SET.QUANT.PRE.v300").getValue(); +} + +template <> inline StringRef buildUnaryConfigCallee(MLIRContext *context) { + return StringAttr::get(context, "llvm.hivm.SET.RELU.ALPHA").getValue(); +} + +template <> inline StringRef buildUnaryConfigCallee(MLIRContext *context) { + return StringAttr::get(context, "llvm.hivm.SET.FIX.CLIP.RELU").getValue(); +} + +template <> inline StringRef buildUnaryConfigCallee(MLIRContext *context) { + return StringAttr::get(context, "llvm.hivm.SET.LOOP2.STRIDE.OUTTOL1").getValue(); +} + +template <> inline StringRef buildUnaryConfigCallee(MLIRContext *context) { + return StringAttr::get(context, "llvm.hivm.SET.LOOP1.STRIDE.OUTTOL1").getValue(); +} + +template <> inline StringRef buildUnaryConfigCallee(MLIRContext *context) { + return StringAttr::get(context, "llvm.hivm.SET.LOOP.SIZE.OUTTOL1").getValue(); +} + +template <> inline StringRef buildUnaryConfigCallee(MLIRContext *context) { + return StringAttr::get(context, "llvm.hivm.SET.MTE2.NZ.PARA").getValue(); +} + +template <> inline StringRef buildUnaryConfigCallee(MLIRContext *context) { + return StringAttr::get(context, "llvm.hivm.SET.PAD.VAL.OUTTOL1").getValue(); +} + +template <> inline StringRef buildUnaryConfigCallee(MLIRContext *context) { + return StringAttr::get(context, "llvm.hivm.SET.FPC").getValue(); +} + +template <> inline StringRef buildUnaryConfigCallee(MLIRContext *context) { + return StringAttr::get(context, "llvm.hivm.SET.ST.ATOMIC.CFG").getValue(); +} + +template <> inline StringRef buildNullaryConfigCallee(MLIRContext *context) { + return StringAttr::get(context, "llvm.hivm.SET.ATOMIC.S32").getValue(); +} + +template <> inline StringRef buildNullaryConfigCallee(MLIRContext *context) { + return StringAttr::get(context, "llvm.hivm.SET.ATOMIC.S8").getValue(); +} + +template StringRef buildSyncCallee(MLIRContext *context); + +template <> inline StringRef buildSyncCallee(MLIRContext *context) { + return StringAttr::get(context, "llvm.hivm.SET.FLAG.IMM").getValue(); +} + +template <> inline StringRef buildSyncCallee(MLIRContext *context) { + return StringAttr::get(context, "llvm.hivm.WAIT.FLAG.IMM").getValue(); +} + +template <> inline StringRef buildSyncCallee(MLIRContext *context) { + return StringAttr::get(context, "llvm.hivm.SET.FLAG.REG").getValue(); +} + +template <> inline StringRef buildSyncCallee(MLIRContext *context) { + return StringAttr::get(context, "llvm.hivm.WAIT.FLAG.REG").getValue(); +} + +template <> inline StringRef buildSyncCallee(MLIRContext *context) { + return StringAttr::get(context, "llvm.hivm.BARRIER").getValue(); +} + +template <> inline StringRef buildSyncCallee(MLIRContext *context) { + return StringAttr::get(context, "llvm.hivm.SET.CROSS.CORE").getValue(); +} + +template <> inline StringRef buildSyncCallee(MLIRContext *context) { + return StringAttr::get(context, "llvm.hivm.WAIT.FLAG.DEV.REG").getValue(); +} + +template <> inline StringRef buildSyncCallee(MLIRContext *context) { + return StringAttr::get(context, "llvm.hivm.SET.INTRA.BLOCK.mode").getValue(); +} + +template <> inline StringRef buildSyncCallee(MLIRContext *context) { + return StringAttr::get(context, "llvm.hivm.WAIT.INTRA.BLOCK.mode").getValue(); +} + +template <> inline StringRef buildSyncCallee(MLIRContext *context) { + return StringAttr::get(context, "llvm.hivm.GET.BUFI.mode").getValue(); +} + +template <> inline StringRef buildSyncCallee(MLIRContext *context) { + return StringAttr::get(context, "llvm.hivm.RLS.BUFI.mode").getValue(); +} + +template StringRef buildRuntimeQueryCallee(MLIRContext *context); + +template <> inline StringRef buildRuntimeQueryCallee(MLIRContext *context) { + return StringAttr::get(context, "llvm.hivm.GET.BLOCK.IDX").getValue(); +} + +template <> inline StringRef buildRuntimeQueryCallee(MLIRContext *context) { + return StringAttr::get(context, "llvm.hivm.GET.SUBBLOCKID").getValue(); +} + +template <> inline StringRef buildRuntimeQueryCallee(MLIRContext *context) { + return StringAttr::get(context, "llvm.hivm.GET.BLOCK.NUM").getValue(); +} + +template <> inline StringRef buildRuntimeQueryCallee(MLIRContext *context) { + return StringAttr::get(context, "llvm.hivm.GET.SUBBLOCKDIM").getValue(); +} + +template StringRef buildSimtBlockQueryCallee(MLIRContext *context); + +template <> inline StringRef buildSimtBlockQueryCallee(MLIRContext *context) { + return StringAttr::get(context, "llvm.hivm.tpe.get.BLOCK.IDX").getValue(); +} + +template <> inline StringRef buildSimtBlockQueryCallee(MLIRContext *context) { + return StringAttr::get(context, "llvm.hivm.tpe.get.BLOCK.NUM").getValue(); +} + +template +FailureOr buildAtomicCallee(MLIRContext *context, Type ptrType, Type valueType, Attribute signednessAttr); + +#define PTO_DECLARE_ATOMIC_CALLEE(OP, NAME) \ + template <> \ + inline FailureOr buildAtomicCallee(MLIRContext * context, Type ptrType, Type valueType, \ + Attribute signednessAttr) { \ + return buildAtomicCalleeName(context, ptrType, valueType, signednessAttr, NAME); \ + } + +PTO_DECLARE_ATOMIC_CALLEE(AtomicCasOp, "CAS") +PTO_DECLARE_ATOMIC_CALLEE(AtomicExchOp, "EXCH") +PTO_DECLARE_ATOMIC_CALLEE(AtomicAddOp, "ADD") +PTO_DECLARE_ATOMIC_CALLEE(AtomicSubOp, "SUB") +PTO_DECLARE_ATOMIC_CALLEE(AtomicMinOp, "MIN") +PTO_DECLARE_ATOMIC_CALLEE(AtomicMaxOp, "MAX") +PTO_DECLARE_ATOMIC_CALLEE(AtomicAndOp, "AND") +PTO_DECLARE_ATOMIC_CALLEE(AtomicOrOp, "OR") +PTO_DECLARE_ATOMIC_CALLEE(AtomicXorOp, "XOR") + +#undef PTO_DECLARE_ATOMIC_CALLEE + +} // namespace mlir::pto::detail diff --git a/lib/PTO/Transforms/VPTOCANN900LLVMEmitterTypeHelpers.cpp b/lib/PTO/Transforms/VPTOCANN900LLVMEmitterTypeHelpers.cpp new file mode 100644 index 0000000000..109b33393a --- /dev/null +++ b/lib/PTO/Transforms/VPTOCANN900LLVMEmitterTypeHelpers.cpp @@ -0,0 +1,1372 @@ +// Copyright (c) 2026 Huawei Technologies Co., Ltd. +// This program is free software, you can redistribute it and/or modify it under the terms and conditions of +// CANN Open Software License Agreement Version 2.0 (the "License"). +// Please refer to the License for details. You may not use this file except in compliance with the License. +// THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED, +// INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE. +// See LICENSE in the root of the software repository for the full text of the License. + +#include "VPTOCANN900LLVMEmitterInternal.h" + +namespace mlir::pto::detail { + +Type getLowPrecisionLLVMType(Type type, MLIRContext *context) { + if (pto::isPTOHiFloat8Type(type)) { + return LLVM::LLVMHiFloat8Type::get(context); + } + if (isa(type)) { + return LLVM::LLVMFloat4E1M2x2Type::get(context); + } + if (isa(type)) { + return LLVM::LLVMFloat4E2M1x2Type::get(context); + } + if (pto::isPTOFloat8E4M3LikeType(type)) { + return LLVM::LLVMFloat8E4M3Type::get(context); + } + if (pto::isPTOFloat8E5M2LikeType(type)) { + return LLVM::LLVMFloat8E5M2Type::get(context); + } + return {}; +} + +bool isLLVMExtensionVectorElementType(Type type) { + return isa(type); +} + +Type getLLVMCompatibleVectorType(ArrayRef shape, Type elementType, ArrayRef scalableDims = {}) { + if (shape.size() == 1 && isLLVMExtensionVectorElementType(elementType)) { + return LLVM::LLVMFixedVectorType::get(elementType, shape.front()); + } + return VectorType::get(shape, elementType, scalableDims); +} + +Type normalizePayloadTypeForLLVMLowering(Type type, Builder &builder) { + if (pto::isPTOHiFloat8x2Type(type)) { + return getLLVMCompatibleVectorType({2}, LLVM::LLVMHiFloat8Type::get(builder.getContext())); + } + // bf16x2 is a 4-byte packed pair; lower it as an opaque i32 so vregs whose + // element type is bf16x2 get a valid LLVM type. + if (pto::isPTOBF16x2Type(type)) { + return builder.getI32Type(); + } + if (Type lowpType = getLowPrecisionLLVMType(type, builder.getContext())) { + return lowpType; + } + + if (auto intType = dyn_cast(type)) { + if (!intType.isSignless()) { + return builder.getIntegerType(intType.getWidth()); + } + return type; + } + + if (auto vecType = dyn_cast(type)) { + Type normalizedElement = normalizePayloadTypeForLLVMLowering(vecType.getElementType(), builder); + if (normalizedElement == vecType.getElementType()) { + return type; + } + return getLLVMCompatibleVectorType(vecType.getShape(), normalizedElement, vecType.getScalableDims()); + } + + return type; +} + +Type normalizeGEPElementTypeForLLVMLowering(Type type, Builder &builder) { + if (pto::isPTOHiFloat8x2Type(type)) { + return builder.getI16Type(); + } + // bf16x2 is 4 bytes, not an 8-bit low-precision type. + if (pto::isPTOBF16x2Type(type)) { + return builder.getI32Type(); + } + if (pto::isPTOLowPrecisionType(type)) { + return builder.getI8Type(); + } + if (isa(type)) { + return builder.getI8Type(); + } + + if (auto vecType = dyn_cast(type)) { + Type normalizedElement = normalizeGEPElementTypeForLLVMLowering(vecType.getElementType(), builder); + if (normalizedElement == vecType.getElementType()) { + return normalizePayloadTypeForLLVMLowering(type, builder); + } + return getLLVMCompatibleVectorType(vecType.getShape(), normalizedElement, vecType.getScalableDims()); + } + + if (auto vecType = dyn_cast(type)) { + Type normalizedElement = normalizeGEPElementTypeForLLVMLowering(vecType.getElementType(), builder); + if (normalizedElement == vecType.getElementType()) { + return normalizePayloadTypeForLLVMLowering(type, builder); + } + return getLLVMCompatibleVectorType({vecType.getNumElements()}, normalizedElement); + } + + return normalizePayloadTypeForLLVMLowering(type, builder); +} + +Type convertVPTOType(Type type, Builder &builder) { + if (auto vecType = dyn_cast(type)) { + Type elementType = normalizePayloadTypeForLLVMLowering(vecType.getElementType(), builder); + return getLLVMCompatibleVectorType({vecType.getElementCount()}, elementType); + } + if (isa(type)) { + return VectorType::get({256}, builder.getI1Type()); + } + if (isa(type)) { + return VectorType::get({32}, builder.getI8Type()); + } + if (isa(type)) { + return LLVM::LLVMPointerType::get(builder.getContext()); + } + if (auto ptrType = dyn_cast(type)) { + return LLVM::LLVMPointerType::get(builder.getContext(), + static_cast(ptrType.getMemorySpace().getAddressSpace())); + } + return normalizePayloadTypeForLLVMLowering(type, builder); +} + +unsigned getNaturalByteAlignment(Type type) { + if (auto vecType = dyn_cast(type)) { + unsigned elemAlign = getNaturalByteAlignment(vecType.getElementType()); + if (!elemAlign) { + return 0; + } + int64_t elems = 1; + for (int64_t dim : vecType.getShape()) { + elems *= dim; + } + return elemAlign * static_cast(elems); + } + if (auto vecType = dyn_cast(type)) { + unsigned elemAlign = getNaturalByteAlignment(vecType.getElementType()); + if (!elemAlign) { + return 0; + } + return elemAlign * vecType.getNumElements(); + } + if (auto intType = dyn_cast(type)) { + return llvm::divideCeil(static_cast(intType.getWidth()), 8U); + } + if (pto::isPTOHiFloat8x2Type(type)) { + return 2; + } + if (pto::isPTOBF16x2Type(type)) { + return 4; + } + if (pto::isPTOLowPrecisionType(type)) { + return 1; + } + if (type.isF16() || type.isBF16()) { + return 2; + } + if (type.isF32()) { + return 4; + } + if (type.isF64()) { + return 8; + } + return 0; +} + +bool hasVPTOConvertibleType(Type type) { + if (!type) { + return false; + } + if (isa(type) || + pto::isPTOLowPrecisionType(type)) { + return true; + } + if (auto vecType = dyn_cast(type)) { + return hasVPTOConvertibleType(vecType.getElementType()); + } + return false; +} + +bool hasVPTOConvertibleType(TypeRange types) { + return llvm::any_of(types, [](Type type) { return hasVPTOConvertibleType(type); }); +} + +Value materializeVPTOCast(OpBuilder &builder, Type resultType, ValueRange inputs, Location loc) { + if (inputs.size() != 1) { + return {}; + } + return builder.create(loc, TypeRange{resultType}, inputs).getResult(0); +} + +// Struct values carry the address of stack-local storage. Keep the pointee +// type local to struct access lowering so the public type conversion remains +// an opaque LLVM pointer, consistent with other pointer-like PTO handles. +LLVM::LLVMStructType getVPTOStructStorageType(pto::StructType structType, Builder &builder) { + struct Frame { + pto::StructType type; + bool materialize; + }; + + // PTO structs form an acyclic type tree. Build literal LLVM struct types in + // explicit post-order so deeply nested legal structs do not consume the C++ + // call stack during lowering. + SmallVector worklist{{structType, false}}; + llvm::DenseMap storageTypes; + while (!worklist.empty()) { + Frame frame = worklist.pop_back_val(); + if (!frame.materialize) { + worklist.push_back({frame.type, true}); + for (Type fieldType : frame.type.getFieldTypes()) { + if (auto nestedStruct = dyn_cast(fieldType)) { + worklist.push_back({nestedStruct, false}); + } + } + continue; + } + + SmallVector fieldTypes; + fieldTypes.reserve(frame.type.getNumFields()); + for (Type fieldType : frame.type.getFieldTypes()) { + if (auto nestedStruct = dyn_cast(fieldType)) { + fieldTypes.push_back(storageTypes.find(nestedStruct)->second); + continue; + } + fieldTypes.push_back(convertVPTOType(fieldType, builder)); + } + storageTypes[frame.type] = LLVM::LLVMStructType::getLiteral(builder.getContext(), fieldTypes); + } + return storageTypes.find(structType)->second; +} + +FailureOr getVPTOStructFieldAddress(ConversionPatternRewriter &rewriter, Location loc, Value root, + pto::StructType rootType, ArrayRef path) { + auto pointerType = LLVM::LLVMPointerType::get(rewriter.getContext()); + Value address = root; + pto::StructType currentType = rootType; + for (auto [depth, index] : llvm::enumerate(path)) { + if (index < 0 || index >= static_cast(currentType.getNumFields())) { + return failure(); + } + Type storageType = getVPTOStructStorageType(currentType, rewriter); + address = rewriter.create(loc, pointerType, storageType, address, + ArrayRef{0, static_cast(index)}); + Type fieldType = currentType.getFieldType(static_cast(index)); + if (depth + 1 == path.size()) { + continue; + } + auto nestedStruct = dyn_cast(fieldType); + if (!nestedStruct) { + return failure(); + } + currentType = nestedStruct; + } + return address; +} + +Value getI64Constant(OpBuilder &builder, Location loc, uint64_t value) { + return builder.create(loc, builder.getI64IntegerAttr(value)).getResult(); +} + +Value getI32Constant(OpBuilder &builder, Location loc, uint64_t value) { + return builder.create(loc, builder.getI32IntegerAttr(value)).getResult(); +} + +[[maybe_unused]] Value getI1Constant(OpBuilder &builder, Location loc, bool value) { + return builder.create(loc, builder.getIntegerAttr(builder.getI1Type(), value ? 1 : 0)).getResult(); +} + +bool isMxElementType(Type ty) { + if (auto floatType = dyn_cast(ty)) { + return floatType.getWidth() == 8; + } + if (isa(ty)) { + return true; + } + std::string typeText; + llvm::raw_string_ostream os(typeText); + ty.print(os); + os.flush(); + return StringRef(typeText).starts_with("f8"); +} + +std::string getMadMxElementFragment(Type type) { + if (type.isF16()) { + return "f16"; + } + if (type.isBF16()) { + return "bf16"; + } + + std::string typeText; + llvm::raw_string_ostream os(typeText); + type.print(os); + os.flush(); + + std::string lower = StringRef(typeText).lower(); + if (StringRef(lower).contains("e4m3")) { + return "e4m3"; + } + if (StringRef(lower).contains("e5m2")) { + return "e5m2"; + } + if (StringRef(lower).contains("hif4")) { + return "hif4"; + } + if (StringRef(lower).contains("e2m1x2")) { + return "e2m1x2"; + } + if (StringRef(lower).contains("e1m2x2")) { + return "e1m2x2"; + } + return {}; +} + +FailureOr buildMadMxCalleeName(MLIRContext *context, Type lhsElem, Type rhsElem) { + std::string lhs = getMadMxElementFragment(lhsElem); + std::string rhs = getMadMxElementFragment(rhsElem); + if (lhs.empty() || rhs.empty()) { + return failure(); + } + return StringAttr::get(context, "llvm.hivm.MMAD.MX." + lhs + rhs).getValue(); +} + +bool isSignedOrSignlessInteger(IntegerType intType, unsigned width) { + return intType && intType.getWidth() == width && (intType.isSigned() || intType.isSignless()); +} + +std::string getMadRhsFragment(Type type) { + if (type.isF16()) { + return "f16"; + } + if (type.isBF16()) { + return "bf16"; + } + if (type.isF32()) { + return "f32"; + } + if (isMadE4M3ElementType(type)) { + return "e4m3"; + } + if (isMadE5M2ElementType(type)) { + return "e5m2"; + } + if (pto::isPTOHiFloat8Type(type)) { + return "hif8"; + } + if (auto intType = dyn_cast(type)) { + if (isSignedOrSignlessInteger(intType, 4)) { + return "s4"; + } + if (isSignedOrSignlessInteger(intType, 8)) { + return "s8"; + } + if (intType.isUnsigned() && intType.getWidth() == 2) { + return "u2"; + } + } + + std::string typeText; + llvm::raw_string_ostream os(typeText); + type.print(os); + os.flush(); + std::string lower = StringRef(typeText).lower(); + if (StringRef(lower).contains("e8m0")) { + return "e8m0"; + } + return {}; +} + +bool isMadE4M3ElementType(Type type) { return pto::isPTOFloat8E4M3LikeType(type); } + +bool isMadE5M2ElementType(Type type) { return pto::isPTOFloat8E5M2LikeType(type); } + +std::string getMadDstFragment(Type type) { + if (type.isF16()) { + return "f16"; + } + if (type.isF32()) { + return "f32"; + } + if (auto intType = dyn_cast(type)) { + if (isSignedOrSignlessInteger(intType, 32)) { + return "s32"; + } + } + return {}; +} + +ArrayRef getMadCalleeContracts() { + static constexpr MadCalleeContract contracts[] = { + {"f16", "f16", "f32", "llvm.hivm.MAD.f162f32.c310"}, {"f16", "f16", "f16", "llvm.hivm.MAD.f162f16"}, + {"f16", "f16", "s32", "llvm.hivm.MAD.f162s32.1952"}, {"bf16", "bf16", "f32", "llvm.hivm.MAD.bf162f32.c310"}, + {"f32", "f32", "f32", "llvm.hivm.MAD.f322f32.c310"}, {"s8", "s8", "s32", "llvm.hivm.MAD.s8.c310"}, + {"e4m3", "e4m3", "f32", "llvm.hivm.MAD.e4m3e4m3.c310"}, {"e4m3", "e5m2", "f32", "llvm.hivm.MAD.e4m3e5m2.c310"}, + {"e5m2", "e4m3", "f32", "llvm.hivm.MAD.e5m2e4m3.c310"}, {"e5m2", "e5m2", "f32", "llvm.hivm.MAD.e5m2e5m2.c310"}, + {"hif8", "hif8", "f32", "llvm.hivm.MAD.e4m3e4m3.c310"}, {"f16", "s4", "", "llvm.hivm.MAD.f16s4.c310"}, + {"f16", "s8", "", "llvm.hivm.MAD.f16s8.c310"}, {"f16", "u2", "", "llvm.hivm.MAD.f16u2"}, + {"f16", "e8m0", "", "llvm.hivm.MAD.f16e8m0.c310"}, + }; + return contracts; +} + +std::string getMadLhsFragment(Type type) { + if (type.isF16()) { + return "f16"; + } + if (type.isBF16()) { + return "bf16"; + } + if (type.isF32()) { + return "f32"; + } + if (isSignedOrSignlessInteger(dyn_cast(type), 8)) { + return "s8"; + } + if (isMadE4M3ElementType(type)) { + return "e4m3"; + } + if (isMadE5M2ElementType(type)) { + return "e5m2"; + } + if (pto::isPTOHiFloat8Type(type)) { + return "hif8"; + } + return {}; +} + +FailureOr buildMadTypedCalleeName(MLIRContext *context, Type lhsElem, Type rhsElem, Type dstElem) { + if (pto::isPTOHiFloat8Type(lhsElem) && pto::isPTOHiFloat8Type(rhsElem) && dstElem.isF32()) { + return StringAttr::get(context, "llvm.hivm.MAD.e4m3e4m3.c310").getValue(); + } + std::string lhs = getMadLhsFragment(lhsElem); + std::string rhs = getMadRhsFragment(rhsElem); + std::string dst = getMadDstFragment(dstElem); + for (const MadCalleeContract &contract : getMadCalleeContracts()) { + if (contract.lhs == lhs && contract.rhs == rhs && (contract.dst.empty() || contract.dst == dst)) { + return StringAttr::get(context, contract.callee).getValue(); + } + } + return failure(); +} + +FailureOr buildLaneTypedCallee(MLIRContext *context, Type resultType, StringRef stem, StringRef suffix) { + std::string vec = getElementTypeFragment(getElementTypeFromVectorLike(resultType)); + auto lanes = getElementCountFromVectorLike(resultType); + if (vec.empty() || !lanes) { + return failure(); + } + + return StringAttr::get(context, "llvm.hivm." + stem.str() + ".v" + std::to_string(*lanes) + vec + suffix.str()) + .getValue(); +} + +std::string getLowPrecisionElementFragment(Type type); + +std::string getCANN900VectorElementFragment(Type type) { + if (type.isF16()) { + return "f16"; + } + if (type.isBF16()) { + return "bf16"; + } + if (type.isF32()) { + return "f32"; + } + if (std::string lowPrecision = getLowPrecisionElementFragment(type); !lowPrecision.empty()) { + return lowPrecision; + } + if (auto intType = dyn_cast(type)) { + return "i" + std::to_string(intType.getWidth()); + } + return {}; +} + +std::string getCANN900VectorTypeFragment(Type vectorType) { + std::string elem = getCANN900VectorElementFragment(getElementTypeFromVectorLike(vectorType)); + auto lanes = getElementCountFromVectorLike(vectorType); + if (elem.empty() || !lanes) { + return {}; + } + return "v" + std::to_string(*lanes) + elem; +} + +std::string getCANN900SignednessFragment(Type elemType) { + if (elemType.isF16() || elemType.isBF16() || elemType.isF32()) { + return "s"; + } + if (auto intType = dyn_cast(elemType)) { + return intType.isUnsigned() ? "u" : "s"; + } + return {}; +} + +FailureOr buildCANN900ModeTypedCallee(MLIRContext *context, Type vectorType, StringRef stem, + StringRef mode) { + std::string vec = getCANN900VectorTypeFragment(vectorType); + if (vec.empty()) { + return failure(); + } + return StringAttr::get(context, "llvm.hivm." + stem.str() + "." + mode.str() + "." + vec).getValue(); +} + +FailureOr buildCANN900SignedModeTypedCallee(MLIRContext *context, Type vectorType, StringRef stem, + StringRef mode) { + std::string vec = getCANN900VectorTypeFragment(vectorType); + std::string signedness = getCANN900SignednessFragment(getElementTypeFromVectorLike(vectorType)); + if (vec.empty() || signedness.empty()) { + return failure(); + } + return StringAttr::get(context, "llvm.hivm." + stem.str() + "." + signedness + "." + mode.str() + "." + vec) + .getValue(); +} + +FailureOr buildCANN900WideningReductionCallee(MLIRContext *context, Type inputType, Type resultType, + StringRef stem, StringRef mode) { + std::string inputVec = getCANN900VectorTypeFragment(inputType); + std::string resultVec = getCANN900VectorTypeFragment(resultType); + std::string signedness = getCANN900SignednessFragment(getElementTypeFromVectorLike(inputType)); + if (inputVec.empty() || resultVec.empty() || signedness.empty()) { + return failure(); + } + return StringAttr::get(context, "llvm.hivm." + stem.str() + "." + signedness + "." + mode.str() + "." + resultVec + + "." + inputVec) + .getValue(); +} + +std::string getElementTypeFragment(Type type) { + if (type.isF16()) { + return "f16"; + } + if (type.isBF16()) { + return "bf16"; + } + if (type.isF32()) { + return "f32"; + } + if (auto intType = dyn_cast(type)) { + return (intType.isUnsigned() ? "u" : "s") + std::to_string(intType.getWidth()); + } + return {}; +} + +std::string getLowPrecisionElementFragment(Type type) { + if (pto::isPTOHiFloat8x2Type(type)) { + return "hif8x2"; + } + if (pto::isPTOHiFloat8Type(type)) { + return "hif8"; + } + if (isa(type)) { + return "f4e1m2x2"; + } + if (isa(type)) { + return "f4e2m1x2"; + } + if (pto::isPTOBF16x2Type(type)) { + return "bf16x2"; + } + if (pto::isPTOFloat8E4M3LikeType(type)) { + return "f8e4m3"; + } + if (pto::isPTOFloat8E5M2LikeType(type)) { + return "f8e5m2"; + } + return {}; +} + +std::string getMemoryElementTypeFragment(Type type) { + if (auto intType = dyn_cast(type)) { + return "i" + std::to_string(intType.getWidth()); + } + if (pto::isPTOHiFloat8Type(type)) { + return "s8"; + } + if (std::string elem = getElementTypeFragment(type); !elem.empty()) { + return elem; + } + return getLowPrecisionElementFragment(type); +} + +bool isLowpPayloadElementType(Type type) { + return pto::isPTOFloat8Type(type) || pto::isPTOHiFloat8Type(type) || pto::isPTOFloat4PackedType(type); +} + +std::optional getLowpPayloadABI(Type elementType, MLIRContext *context) { + if (!isLowpPayloadElementType(elementType)) { + return std::nullopt; + } + return LowpPayloadABI{IntegerType::get(context, 8), "u8"}; +} + +std::string getDirectLowpVLogicElementFragment(Type type) { + if (pto::isPTOFloat8E4M3LikeType(type)) { + return "fp8e4m3"; + } + if (pto::isPTOFloat8E5M2LikeType(type)) { + return "fp8e5m2"; + } + return {}; +} + +FailureOr buildDirectLowpVLogicCallee(MLIRContext *context, Type vectorType, StringRef stem, + StringRef mode) { + Type elementType = getElementTypeFromVectorLike(vectorType); + auto lanes = getElementCountFromVectorLike(vectorType); + std::string elem = getDirectLowpVLogicElementFragment(elementType); + if (elem.empty() || !lanes) { + return failure(); + } + return StringAttr::get(context, "llvm.hivm." + stem.str() + "." + mode.str() + ".v" + std::to_string(*lanes) + elem) + .getValue(); +} + +FailureOr buildLowpPayloadVLogicCallee(MLIRContext *context, Type vectorType, StringRef stem, + StringRef mode) { + Type elementType = getElementTypeFromVectorLike(vectorType); + auto lanes = getElementCountFromVectorLike(vectorType); + std::optional abi = getLowpPayloadABI(elementType, context); + if (!abi || !lanes) { + return failure(); + } + return StringAttr::get(context, "llvm.hivm." + stem.str() + "." + mode.str() + ".v" + std::to_string(*lanes) + + abi->intrinsicElementFragment.str()) + .getValue(); +} + +Type getLowpPayloadCarrierType(Type vectorLikeType, MLIRContext *context) { + Type elementType = getElementTypeFromVectorLike(vectorLikeType); + std::optional abi = getLowpPayloadABI(elementType, context); + if (!abi) { + return {}; + } + auto lanes = getElementCountFromVectorLike(vectorLikeType); + if (!lanes) { + return {}; + } + return VectorType::get({*lanes}, abi->llvmElementType); +} + +Type getPayloadABIType(Type semanticType, Type convertedType, MLIRContext *context) { + if (Type carrierType = getLowpPayloadCarrierType(semanticType, context)) { + return carrierType; + } + return convertedType; +} + +Value castToPayloadABI(Location loc, Value value, Type semanticType, ConversionPatternRewriter &rewriter) { + Type carrierType = getLowpPayloadCarrierType(semanticType, rewriter.getContext()); + if (!carrierType || carrierType == value.getType()) { + return value; + } + return rewriter.create(loc, carrierType, value); +} + +Value castFromPayloadABI(Location loc, Value value, Type semanticType, Type convertedType, + ConversionPatternRewriter &rewriter) { + Type carrierType = getLowpPayloadCarrierType(semanticType, rewriter.getContext()); + if (!carrierType || carrierType == convertedType) { + return value; + } + return rewriter.create(loc, convertedType, value); +} + +std::string getAtomicElementTypeFragment(Type type, Attribute signednessAttr) { + if (auto vecType = dyn_cast(type)) { + if (vecType.getRank() != 1 || vecType.getDimSize(0) != 2) { + return {}; + } + if (vecType.getElementType().isF16()) { + return "f16x2"; + } + if (vecType.getElementType().isBF16()) { + return "bf16x2"; + } + return {}; + } + if (type.isF16()) { + return "fp16"; + } + if (type.isBF16()) { + return "bf16"; + } + if (type.isF32()) { + return "fp32"; + } + auto intType = dyn_cast(type); + if (!intType) { + return {}; + } + if (intType.getWidth() != 32 && intType.getWidth() != 64) { + return {}; + } + if (signednessAttr) { + auto signedness = cast(signednessAttr).getValue(); + return std::string(signedness == pto::Signedness::Unsigned ? "u" : "s") + std::to_string(intType.getWidth()); + } + return std::string(intType.isUnsigned() ? "u" : "s") + std::to_string(intType.getWidth()); +} + +std::string getL0LoadElementFragment(Type type) { + std::string elem = getElementTypeFragment(type); + if (!elem.empty()) { + return elem; + } + + std::string typeText; + llvm::raw_string_ostream os(typeText); + type.print(os); + os.flush(); + std::string lower = StringRef(typeText).lower(); + if (StringRef(lower).contains("e4m3") || StringRef(lower).contains("e5m2") || StringRef(lower).contains("e8m0") || + StringRef(lower).contains("hif8") || StringRef(lower).contains("e1m2x2") || StringRef(lower).contains("e2m1x2")) { + return "s8"; + } + return {}; +} + +std::string getShuffleIntrinsicTypeFragment(Type type) { + if (auto intType = dyn_cast(type)) { + switch (intType.getWidth()) { + case 32: + return "i32"; + case 64: + return "i64"; + default: + return {}; + } + } + if (type.isF16()) { + return "f16"; + } + if (type.isF32()) { + return "f32"; + } + if (auto vecType = dyn_cast(type)) { + if (vecType.getRank() == 1 && vecType.getDimSize(0) == 2 && vecType.getElementType().isF16()) { + return "v2f16"; + } + } + return {}; +} + +std::string getReduxIntrinsicTypeFragment(Type type, Attribute signednessAttr) { + if (auto intType = dyn_cast(type)) { + if (intType.getWidth() != 32) { + return {}; + } + bool isUnsigned = false; + if (signednessAttr) { + isUnsigned = cast(signednessAttr).getValue() == pto::Signedness::Unsigned; + } + return isUnsigned ? "u32" : "s32"; + } + if (type.isF16()) { + return "f16"; + } + if (type.isF32()) { + return "f32"; + } + return {}; +} + +Type getElementTypeFromVectorLike(Type type) { + if (auto vecType = dyn_cast(type)) { + return vecType.getElementType(); + } + if (auto vecType = dyn_cast(type)) { + return vecType.getElementType(); + } + if (auto vecType = dyn_cast(type)) { + return vecType.getElementType(); + } + return {}; +} + +std::optional getElementCountFromVectorLike(Type type) { + if (auto vecType = dyn_cast(type)) { + return vecType.getElementCount(); + } + if (auto vecType = dyn_cast(type)) { + if (vecType.getRank() != 1) { + return std::nullopt; + } + return vecType.getShape().front(); + } + if (auto vecType = dyn_cast(type)) { + return vecType.getNumElements(); + } + return std::nullopt; +} + +Value castIntegerLikeTo(Operation *anchor, Value value, Type targetType) { + OpBuilder builder(anchor); + builder.setInsertionPoint(anchor); + + if (value.getType() == targetType) { + return value; + } + + auto targetInt = dyn_cast(targetType); + if (value.getType().isIndex() && targetInt) { + return builder.create(anchor->getLoc(), targetType, value); + } + if (auto sourceInt = dyn_cast(value.getType())) { + if (targetInt) { + if (sourceInt.getWidth() < targetInt.getWidth()) { + return builder.create(anchor->getLoc(), targetType, value); + } + if (sourceInt.getWidth() > targetInt.getWidth()) { + return builder.create(anchor->getLoc(), targetType, value); + } + return value; + } + if (targetType.isIndex()) { + return builder.create(anchor->getLoc(), targetType, value); + } + } + + return {}; +} + +FailureOr reinterpretPointerToAddrSpace(Operation *anchor, Value value, unsigned targetAddressSpace) { + auto sourcePtrType = dyn_cast(value.getType()); + if (!sourcePtrType) { + return failure(); + } + if (sourcePtrType.getAddressSpace() == targetAddressSpace) { + return value; + } + + OpBuilder builder(anchor); + builder.setInsertionPoint(anchor); + Location loc = anchor->getLoc(); + Value asInt = builder.create(loc, builder.getI64Type(), value); + Type targetPtrType = LLVM::LLVMPointerType::get(anchor->getContext(), targetAddressSpace); + return builder.create(loc, targetPtrType, asInt).getResult(); +} + +FailureOr normalizeVdupScalarOperand(OpBuilder &builder, Location loc, Value input, Type resultType) { + auto intType = dyn_cast(input.getType()); + if (!intType || intType.getWidth() != 8) { + return input; + } + + Type resultElemType = getElementTypeFromVectorLike(resultType); + std::string resultElemFragment = getElementTypeFragment(resultElemType); + if (resultElemFragment != "s8" && resultElemFragment != "u8") { + return input; + } + + if (intType.isSignless()) { + return input; + } + + Type signlessType = builder.getIntegerType(intType.getWidth()); + return builder.create(loc, TypeRange{signlessType}, input).getResult(0); +} + +Value normalizeByteScalarOperandForCANN900VectorCall(OpBuilder &builder, Location loc, Value input, + Type semanticElementType) { + (void)semanticElementType; + auto intType = dyn_cast(input.getType()); + if (!intType || intType.getWidth() != 8 || intType.isSignless()) { + return input; + } + + Type signlessType = builder.getIntegerType(8); + return builder.create(loc, TypeRange{signlessType}, input).getResult(0); +} + +bool isCompatibleScalarForSemanticType(Type semanticType, Type scalarType) { + if (semanticType == scalarType) { + return true; + } + + auto semanticInt = dyn_cast(semanticType); + auto scalarInt = dyn_cast(scalarType); + if (!semanticInt || !scalarInt || semanticInt.getWidth() != scalarInt.getWidth()) { + return false; + } + + if (semanticInt.isSigned()) { + return scalarInt.isSigned() || scalarInt.isSignless(); + } + if (semanticInt.isUnsigned()) { + return scalarInt.isUnsigned() || scalarInt.isSignless(); + } + return scalarInt.isSignless(); +} + +std::string getCopyElementFragment(Type elementType) { + if (!elementType) { + return {}; + } + if (elementType.isF16()) { + return "f16"; + } + if (elementType.isBF16()) { + return "bf16"; + } + if (elementType.isF32()) { + return "f32"; + } + // Handle FP8 family (e4m3/e5m2/e8m0/hif8) used by cube-matmul/mad_mx. + std::string typeText; + llvm::raw_string_ostream os(typeText); + elementType.print(os); + os.flush(); + std::string lower = StringRef(typeText).lower(); + if (StringRef(lower).contains("e4m3")) { + return "e4m3"; + } + if (StringRef(lower).contains("e5m2")) { + return "e5m2"; + } + if (StringRef(lower).contains("e8m0")) { + return "e8m0"; + } + if (StringRef(lower).contains("hif8")) { + return "hif8"; + } + if (StringRef(lower).contains("e1m2x2") || StringRef(lower).contains("e2m1x2")) { + return "u8"; + } + if (auto intType = dyn_cast(elementType)) { + switch (intType.getWidth()) { + case 8: + return intType.isUnsigned() ? "u8" : "s8"; + case 16: + return intType.isUnsigned() ? "u16" : "s16"; + case 32: + return intType.isUnsigned() ? "u32" : "s32"; + default: + return {}; + } + } + return {}; +} + +std::string getNd2NzCopyElementFragment(Type elementType) { + if (!elementType) { + return {}; + } + std::string typeText; + llvm::raw_string_ostream os(typeText); + elementType.print(os); + os.flush(); + std::string lower = StringRef(typeText).lower(); + if (StringRef(lower).contains("e4m3") || StringRef(lower).contains("e5m2") || StringRef(lower).contains("e8m0") || + StringRef(lower).contains("hif8")) { + return "U8"; + } + if (StringRef(lower).contains("e1m2x2") || StringRef(lower).contains("e2m1x2")) { + return "U8"; + } + + if (elementType.isF16() || elementType.isBF16()) { + return "U16"; + } + if (elementType.isF32()) { + return "U32"; + } + if (auto intType = dyn_cast(elementType)) { + switch (intType.getWidth()) { + case 8: + return "U8"; + case 16: + return "U16"; + case 32: + return "U32"; + default: + return {}; + } + } + return {}; +} + +std::optional parsePredicatePatternImmediate(StringRef pattern) { + if (pattern == "PAT_ALL") { + return 0; + } + if (pattern == "PAT_VL1") { + return 1; + } + if (pattern == "PAT_VL2") { + return 2; + } + if (pattern == "PAT_VL3") { + return 3; + } + if (pattern == "PAT_VL4") { + return 4; + } + if (pattern == "PAT_VL8") { + return 5; + } + if (pattern == "PAT_VL16") { + return 6; + } + if (pattern == "PAT_VL32") { + return 7; + } + if (pattern == "PAT_VL64") { + return 8; + } + if (pattern == "PAT_VL128") { + return 9; + } + if (pattern == "PAT_M3") { + return 10; + } + if (pattern == "PAT_M4") { + return 11; + } + if (pattern == "PAT_H") { + return 12; + } + if (pattern == "PAT_Q") { + return 13; + } + if (pattern == "PAT_ALLF") { + return 15; + } + return std::nullopt; +} + +std::optional parseHiLoPartImmediate(StringRef part) { + if (part == "LOWER") { + return 0; + } + if (part == "HIGHER") { + return 1; + } + return std::nullopt; +} + +std::optional parseRoundModeImmediate(StringRef roundMode) { + if (roundMode == "R" || roundMode == "ROUND_R") { + return 0; + } + if (roundMode == "A" || roundMode == "ROUND_A") { + return 1; + } + if (roundMode == "F" || roundMode == "ROUND_F") { + return 2; + } + if (roundMode == "C" || roundMode == "ROUND_C") { + return 3; + } + if (roundMode == "Z" || roundMode == "ROUND_Z") { + return 4; + } + if (roundMode == "O" || roundMode == "ROUND_O") { + return 5; + } + if (roundMode == "H" || roundMode == "ROUND_H") { + return 6; + } + return std::nullopt; +} + +std::optional parseSaturationImmediate(StringRef sat) { + if (sat == "SAT") { + return 1; + } + if (sat == "NOSAT") { + return 0; + } + return std::nullopt; +} + +std::optional parsePartImmediate(StringRef part) { + if (part == "EVEN" || part == "PART_EVEN") { + return 0; + } + if (part == "ODD" || part == "PART_ODD") { + return 1; + } + return std::nullopt; +} + +std::optional parseVcvtPartImmediate(StringRef part) { + if (part == "EVEN" || part == "PART_EVEN" || part == "P0" || part == "PART_P0") { + return 0; + } + if (part == "ODD" || part == "PART_ODD" || part == "P1" || part == "PART_P1") { + return 1; + } + if (part == "P2" || part == "PART_P2") { + return 2; + } + if (part == "P3" || part == "PART_P3") { + return 3; + } + return std::nullopt; +} + +std::optional parsePredicateStoreDistImmediate(StringRef dist) { + if (dist == "NORM") { + return 0; + } + if (dist == "PK") { + return 1; + } + return std::nullopt; +} + +std::optional parsePredicateLoadDistImmediate(StringRef dist) { + if (dist.empty() || dist == "NORM") { + return 0; + } + if (dist == "US") { + return 1; + } + if (dist == "DS") { + return 2; + } + return std::nullopt; +} + +std::optional parsePostModeImmediate(StringRef mode) { + if (mode == "NO_POST_UPDATE") { + return 0; + } + if (mode == "POST_UPDATE") { + return 1; + } + return std::nullopt; +} + +std::optional parsePipeImmediate(StringRef pipe) { + if (pipe == "PIPE_S") { + return 0; + } + if (pipe == "PIPE_V") { + return 1; + } + if (pipe == "PIPE_M") { + return 2; + } + if (pipe == "PIPE_MTE1") { + return 3; + } + if (pipe == "PIPE_MTE2") { + return 4; + } + if (pipe == "PIPE_MTE3") { + return 5; + } + if (pipe == "PIPE_ALL") { + return 6; + } + if (pipe == "PIPE_MTE4") { + return 7; + } + if (pipe == "PIPE_MTE5") { + return 8; + } + if (pipe == "PIPE_V2") { + return 9; + } + if (pipe == "PIPE_FIX") { + return 10; + } + if (pipe == "VIRTUAL_PIPE_MTE2_L1A") { + return 11; + } + if (pipe == "VIRTUAL_PIPE_MTE2_L1B") { + return 12; + } + return std::nullopt; +} + +std::optional parseEventImmediate(StringRef event) { + if (!event.consume_front("EVENT_ID")) { + return std::nullopt; + } + uint64_t value = 0; + if (event.getAsInteger(10, value)) { + return std::nullopt; + } + return value; +} + +std::optional parseSprImmediate(StringRef spr) { + if (spr == "AR") { + return 74; + } + return std::nullopt; +} + +std::optional getDistElementWidth(Type type) { + if (auto intType = dyn_cast(type)) { + return intType.getWidth(); + } + if (isLowpPayloadElementType(type)) { + return 8; + } + if (type.isF16() || type.isBF16()) { + return 16; + } + if (type.isF32()) { + return 32; + } + if (type.isF64()) { + return 64; + } + // bf16x2 is a 32-bit packed pair; its dist width is 32 (i32/align4 ABI). + if (pto::isPTOBF16x2Type(type)) { + return 32; + } + return std::nullopt; +} + +VcvtElemKind classifyVcvtElemType(Type type) { + if (type.isF16()) { + return VcvtElemKind::F16; + } + if (type.isBF16()) { + return VcvtElemKind::BF16; + } + if (type.isF32()) { + return VcvtElemKind::F32; + } + if (pto::isPTOFloat8E4M3LikeType(type)) { + return VcvtElemKind::F8E4M3; + } + if (pto::isPTOFloat8E5M2LikeType(type)) { + return VcvtElemKind::F8E5M2; + } + if (pto::isPTOHiFloat8Type(type)) { + return VcvtElemKind::HiF8; + } + if (isa(type)) { + return VcvtElemKind::F4E1M2x2; + } + if (isa(type)) { + return VcvtElemKind::F4E2M1x2; + } + if (auto intType = dyn_cast(type)) { + switch (intType.getWidth()) { + case 8: + return intType.isUnsigned() ? VcvtElemKind::U8 : VcvtElemKind::S8; + case 16: + return intType.isUnsigned() ? VcvtElemKind::U16 : VcvtElemKind::S16; + case 32: + return intType.isUnsigned() ? VcvtElemKind::U32 : VcvtElemKind::S32; + case 64: + return intType.isUnsigned() ? VcvtElemKind::Invalid : VcvtElemKind::S64; + default: + return VcvtElemKind::Invalid; + } + } + return VcvtElemKind::Invalid; +} + +struct VcvtContractEntry { + VcvtElemKind src; + VcvtElemKind dst; + VcvtContract contract; +}; + +constexpr VcvtContractEntry kVcvtContractEntries[] = { + {VcvtElemKind::F32, VcvtElemKind::F8E4M3, {"llvm.hivm.vcvtff.f322f8e4m3.x", true, true, true, 32, false}}, + {VcvtElemKind::F32, VcvtElemKind::F8E5M2, {"llvm.hivm.vcvtff.f322f8e5m2.x", true, true, true, 32, false}}, + {VcvtElemKind::F32, VcvtElemKind::HiF8, {"llvm.hivm.vcvtff.f322hif8.x", true, true, true, 32, false}}, + {VcvtElemKind::F32, VcvtElemKind::F16, {"llvm.hivm.vcvtff.f322f16.x", true, true, true, 32, false}}, + {VcvtElemKind::F32, VcvtElemKind::BF16, {"llvm.hivm.vcvtff.f322bf16.x", true, true, true, 32, false}}, + {VcvtElemKind::F32, VcvtElemKind::S16, {"llvm.hivm.vcvtfi.f322s16.x", true, true, true, 32, false}}, + {VcvtElemKind::F32, VcvtElemKind::S32, {"llvm.hivm.vcvtfi.f322s32.x", true, true, false, 32, false}}, + {VcvtElemKind::F32, VcvtElemKind::S64, {"llvm.hivm.vcvtfi.f322s64.x", true, true, true, 32, false}}, + {VcvtElemKind::F16, VcvtElemKind::F8E4M3, {"llvm.hivm.vcvtff.f162f8e4m3.x", true, true, true, 16, false}}, + {VcvtElemKind::F16, VcvtElemKind::F8E5M2, {"llvm.hivm.vcvtff.f162f8e5m2.x", true, true, true, 16, false}}, + {VcvtElemKind::F16, VcvtElemKind::HiF8, {"llvm.hivm.vcvtff.f162hif8.x", true, true, true, 16, false}}, + {VcvtElemKind::F16, VcvtElemKind::F32, {"llvm.hivm.vcvtff.f162f32.x", false, false, true, 16, false}}, + {VcvtElemKind::F16, VcvtElemKind::S32, {"llvm.hivm.vcvtfi.f162s32.x", true, false, true, 16, false}}, + {VcvtElemKind::F16, VcvtElemKind::S16, {"llvm.hivm.vcvtfi.f162s16.x", true, true, false, 16, false}}, + {VcvtElemKind::F16, VcvtElemKind::S8, {"llvm.hivm.vcvtfi.f162s8.x", true, true, true, 16, false}}, + {VcvtElemKind::F16, VcvtElemKind::U8, {"llvm.hivm.vcvtfi.f162u8.x", true, true, true, 16, false}}, + {VcvtElemKind::BF16, VcvtElemKind::F8E4M3, {"llvm.hivm.vcvtff.bf162f8e4m3.x", true, true, true, 16, false}}, + {VcvtElemKind::BF16, VcvtElemKind::F8E5M2, {"llvm.hivm.vcvtff.bf162f8e5m2.x", true, true, true, 16, false}}, + {VcvtElemKind::BF16, VcvtElemKind::F4E1M2x2, {"llvm.hivm.vcvtff2.bf162f4e1m2x2.x", true, false, true, 16, false}}, + {VcvtElemKind::BF16, VcvtElemKind::F4E2M1x2, {"llvm.hivm.vcvtff2.bf162f4e2m1x2.x", true, false, true, 16, false}}, + {VcvtElemKind::BF16, VcvtElemKind::F16, {"llvm.hivm.vcvtff.bf162f16.x", true, true, false, 16, true}}, + {VcvtElemKind::BF16, VcvtElemKind::F32, {"llvm.hivm.vcvtff.bf162f32.x", false, false, true, 16, false}}, + {VcvtElemKind::BF16, VcvtElemKind::S32, {"llvm.hivm.vcvtfi.bf162s32.x", true, true, true, 16, false}}, + {VcvtElemKind::U8, VcvtElemKind::F16, {"llvm.hivm.vcvtif.u82f16.x", false, false, true, 8, false}}, + {VcvtElemKind::U8, VcvtElemKind::U16, {"llvm.hivm.vcvtii.u82u16.x", false, false, true, 8, false}}, + {VcvtElemKind::U8, VcvtElemKind::U32, {"llvm.hivm.vcvtii.u82u32.x", false, false, true, 8, false}}, + {VcvtElemKind::S8, VcvtElemKind::F16, {"llvm.hivm.vcvtif.s82f16.x", false, false, true, 8, false}}, + {VcvtElemKind::S8, VcvtElemKind::S16, {"llvm.hivm.vcvtii.s82s16.x", false, false, true, 8, false}}, + {VcvtElemKind::S8, VcvtElemKind::S32, {"llvm.hivm.vcvtii.s82s32.x", false, false, true, 8, false}}, + {VcvtElemKind::U16, VcvtElemKind::U8, {"llvm.hivm.vcvtii.u162u8.x", false, true, true, 16, false}}, + {VcvtElemKind::U16, VcvtElemKind::U32, {"llvm.hivm.vcvtii.u162u32.x", false, false, true, 16, false}}, + {VcvtElemKind::S16, VcvtElemKind::F16, {"llvm.hivm.vcvtif.s162f16.x", true, false, false, 16, false}}, + {VcvtElemKind::S16, VcvtElemKind::F32, {"llvm.hivm.vcvtif.s162f32.x", false, false, true, 16, false}}, + {VcvtElemKind::S16, VcvtElemKind::U8, {"llvm.hivm.vcvtii.s162u8.x", false, true, true, 16, false}}, + {VcvtElemKind::S16, VcvtElemKind::U32, {"llvm.hivm.vcvtii.s162u32.x", false, false, true, 16, false}}, + {VcvtElemKind::S16, VcvtElemKind::S32, {"llvm.hivm.vcvtii.s162s32.x", false, false, true, 16, false}}, + {VcvtElemKind::U32, VcvtElemKind::U8, {"llvm.hivm.vcvtii.u322u8.x", false, true, true, 32, false}}, + {VcvtElemKind::U32, VcvtElemKind::U16, {"llvm.hivm.vcvtii.u322u16.x", false, true, true, 32, false}}, + {VcvtElemKind::U32, VcvtElemKind::S16, {"llvm.hivm.vcvtii.u322s16.x", false, true, true, 32, false}}, + {VcvtElemKind::S32, VcvtElemKind::F32, {"llvm.hivm.vcvtif.s322f32.x", true, false, false, 32, false}}, + {VcvtElemKind::S32, VcvtElemKind::U8, {"llvm.hivm.vcvtii.s322u8.x", false, true, true, 32, false}}, + {VcvtElemKind::S32, VcvtElemKind::U16, {"llvm.hivm.vcvtii.s322u16.x", false, true, true, 32, false}}, + {VcvtElemKind::S32, VcvtElemKind::S16, {"llvm.hivm.vcvtii.s322s16.x", false, true, true, 32, false}}, + {VcvtElemKind::S32, VcvtElemKind::S64, {"llvm.hivm.vcvtii.s322s64.x", false, false, true, 32, false}}, + {VcvtElemKind::S64, VcvtElemKind::F32, {"llvm.hivm.vcvtif.s642f32.x", true, false, true, 32, false}}, + {VcvtElemKind::S64, VcvtElemKind::S32, {"llvm.hivm.vcvtii.s642s32.x", false, true, true, 32, false}}, + {VcvtElemKind::F8E4M3, VcvtElemKind::F32, {"llvm.hivm.vcvtff.f8e4m32f32.x", false, false, true, 8, false}}, + {VcvtElemKind::F8E5M2, VcvtElemKind::F32, {"llvm.hivm.vcvtff.f8e5m22f32.x", false, false, true, 8, false}}, + {VcvtElemKind::HiF8, VcvtElemKind::F32, {"llvm.hivm.vcvtff.hif82f32.x", false, false, true, 8, false}}, + {VcvtElemKind::F4E1M2x2, VcvtElemKind::BF16, {"llvm.hivm.vcvtff2.f4e1m2x22bf16.x", false, false, true, 8, false}}, + {VcvtElemKind::F4E2M1x2, VcvtElemKind::BF16, {"llvm.hivm.vcvtff2.f4e2m1x22bf16.x", false, false, true, 8, false}}, +}; + +std::optional lookupVcvtContract(VcvtElemKind src, VcvtElemKind dst) { + for (const VcvtContractEntry &entry : kVcvtContractEntries) { + if (entry.src == src && entry.dst == dst) { + return entry.contract; + } + } + return std::nullopt; +} +// VSQZ #st hint must only be set when the compacted vector feeds VSTUR. +// Emitting #st=1 without a matching VSTUR consumer can deadlock hardware queues. +uint64_t determineVsqzStoreHint(pto::VsqzOp vsqz) { + Value result = vsqz.getResult(); + for (Operation *user : result.getUsers()) { + auto vstur = dyn_cast(user); + if (!vstur) { + continue; + } + if (vstur.getValue() == result) { + return 1; + } + } + return 0; +} + +std::optional parseLoadDistImmediate(StringRef dist, Type elementType) { + if (dist.empty() || dist == "NORM") { + return 0; + } + static constexpr std::pair kModes[] = { + {"BRC_B8", 1}, {"BRC_B16", 2}, {"BRC_B32", 3}, {"US_B8", 6}, {"US_B16", 7}, + {"DS_B8", 8}, {"DS_B16", 9}, {"UNPK_B8", 13}, {"UNPK_B16", 14}, {"UNPK_B32", 18}, + {"BRC_BLK", 15}, {"E2B_B16", 16}, {"E2B_B32", 17}, {"SPLT2CHN_B8", 22}, {"SPLT2CHN_B16", 23}, + }; + for (const auto &[name, value] : kModes) { + if (dist == name) { + return value; + } + } + if (dist == "UNPK4" || dist == "SPLT4CHN") { + auto width = getDistElementWidth(elementType); + if (!width || *width != 8) { + return std::nullopt; + } + return dist == "UNPK4" ? 20 : 21; + } + return std::nullopt; +} + +} // namespace mlir::pto::detail diff --git a/lib/PTO/Transforms/VPTOCANN900LLVMEmitterTypePatterns.cpp b/lib/PTO/Transforms/VPTOCANN900LLVMEmitterTypePatterns.cpp new file mode 100644 index 0000000000..e4d3c46773 --- /dev/null +++ b/lib/PTO/Transforms/VPTOCANN900LLVMEmitterTypePatterns.cpp @@ -0,0 +1,741 @@ +// Copyright (c) 2026 Huawei Technologies Co., Ltd. +// This program is free software, you can redistribute it and/or modify it under the terms and conditions of +// CANN Open Software License Agreement Version 2.0 (the "License"). +// Please refer to the License for details. You may not use this file except in compliance with the License. +// THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED, +// INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE. +// See LICENSE in the root of the software repository for the full text of the License. + +#include "VPTOCANN900LLVMEmitterInternal.h" + +namespace mlir::pto::detail { + +class ConvertVPTOUnrealizedCastOp final : public OpConversionPattern { +public: + using OpConversionPattern::OpConversionPattern; + + LogicalResult matchAndRewrite(UnrealizedConversionCastOp op, OpAdaptor adaptor, + ConversionPatternRewriter &rewriter) const override { + if (op->getNumOperands() != 1 || op->getNumResults() != 1) { + return failure(); + } + if (!hasVPTOConvertibleType(op->getOperandTypes()) && !hasVPTOConvertibleType(op->getResultTypes())) { + return failure(); + } + + Type convertedResultType = getTypeConverter()->convertType(op.getResult(0).getType()); + if (!convertedResultType) { + return failure(); + } + + Value input = adaptor.getOperands().front(); + if (input.getType() != convertedResultType) { + return failure(); + } + + rewriter.replaceOp(op, input); + return success(); + } +}; + +class ConvertPtoDeclareStructOp final : public OpConversionPattern { +public: + using OpConversionPattern::OpConversionPattern; + + LogicalResult matchAndRewrite(pto::DeclareStructOp op, OpAdaptor adaptor, + ConversionPatternRewriter &rewriter) const override { + (void)adaptor; + Type declaredType = op.getResult().getType(); + auto resultType = dyn_cast(getTypeConverter()->convertType(declaredType)); + if (!resultType) { + return rewriter.notifyMatchFailure(op, "expected LLVM pointer result type"); + } + auto structType = cast(declaredType); + Type storageType = getVPTOStructStorageType(structType, rewriter); + auto parentFunc = op->getParentOfType(); + if (!parentFunc) { + return rewriter.notifyMatchFailure(op, "expected struct declaration inside a function"); + } + + // A non-entry alloca is a dynamic stack allocation. Keep one stack slot per + // declaration per function invocation even when the declaration is nested + // in a loop or a region. + Value storage; + { + OpBuilder::InsertionGuard guard(rewriter); + Block &entryBlock = parentFunc.getBody().front(); + rewriter.setInsertionPointToStart(&entryBlock); + Value one = rewriter.create(op.getLoc(), rewriter.getI64Type(), rewriter.getIndexAttr(1)); + storage = rewriter.create(op.getLoc(), resultType, storageType, one, /*alignment=*/0); + } + rewriter.replaceOp(op, storage); + return success(); + } +}; + +class ConvertPtoStructGetOp final : public OpConversionPattern { +public: + using OpConversionPattern::OpConversionPattern; + + LogicalResult matchAndRewrite(pto::StructGetOp op, OpAdaptor adaptor, + ConversionPatternRewriter &rewriter) const override { + Type resultType = getTypeConverter()->convertType(op.getValue().getType()); + if (!resultType) { + return rewriter.notifyMatchFailure(op, "could not convert result type"); + } + Value structValue = adaptor.getOperands().front(); + auto structType = cast(op->getOperand(0).getType()); + FailureOr address = getVPTOStructFieldAddress(rewriter, op.getLoc(), structValue, structType, op.getPath()); + if (failed(address)) { + return rewriter.notifyMatchFailure(op, "invalid struct field path"); + } + rewriter.replaceOpWithNewOp(op, resultType, *address, getNaturalByteAlignment(resultType)); + return success(); + } +}; + +class ConvertPtoStructSetOp final : public OpConversionPattern { +public: + using OpConversionPattern::OpConversionPattern; + + LogicalResult matchAndRewrite(pto::StructSetOp op, OpAdaptor adaptor, + ConversionPatternRewriter &rewriter) const override { + Value structValue = adaptor.getOperands().front(); + auto structType = cast(op->getOperand(0).getType()); + FailureOr address = getVPTOStructFieldAddress(rewriter, op.getLoc(), structValue, structType, op.getPath()); + if (failed(address)) { + return rewriter.notifyMatchFailure(op, "invalid struct field path"); + } + rewriter.replaceOpWithNewOp(op, adaptor.getValue(), *address, + getNaturalByteAlignment(adaptor.getValue().getType())); + return success(); + } +}; + +class ConvertArithSelectOp final : public OpConversionPattern { +public: + ConvertArithSelectOp(TypeConverter &typeConverter, MLIRContext *context) + : OpConversionPattern(typeConverter, context, PatternBenefit(2)) {} + + LogicalResult matchAndRewrite(arith::SelectOp op, OpAdaptor adaptor, + ConversionPatternRewriter &rewriter) const override { + if (!op.getCondition().getType().isInteger(1)) { + return rewriter.notifyMatchFailure(op, "only scalar i1 conditions supported for VPTO arith.select"); + } + + Type convertedResultType = getTypeConverter()->convertType(op.getResult().getType()); + if (!convertedResultType) { + return rewriter.notifyMatchFailure(op, "failed to convert result type"); + } + + Value trueValue = adaptor.getTrueValue(); + Value falseValue = adaptor.getFalseValue(); + if (trueValue.getType() != convertedResultType || falseValue.getType() != convertedResultType) { + return rewriter.notifyMatchFailure(op, "converted true/false values must match result type"); + } + + rewriter.replaceOpWithNewOp(op, convertedResultType, adaptor.getCondition(), trueValue, + falseValue); + return success(); + } +}; + +class ConvertPtoAddPtrOp final : public OpConversionPattern { +public: + using OpConversionPattern::OpConversionPattern; + + LogicalResult matchAndRewrite(pto::AddPtrOp op, OpAdaptor adaptor, + ConversionPatternRewriter &rewriter) const override { + Type convertedResultType = getTypeConverter()->convertType(op.getResult().getType()); + auto llvmPtrType = dyn_cast(convertedResultType); + if (!llvmPtrType) { + return rewriter.notifyMatchFailure(op, "expected LLVM pointer result type"); + } + + Value offset = adaptor.getOffset(); + if (offset.getType().isIndex()) { + offset = rewriter.create(op.getLoc(), rewriter.getI64Type(), offset); + } + + auto gep = rewriter.create( + op.getLoc(), llvmPtrType, + normalizeGEPElementTypeForLLVMLowering(cast(op.getPtr().getType()).getElementType(), rewriter), + adaptor.getPtr(), ValueRange{offset}); + rewriter.replaceOp(op, gep.getResult()); + return success(); + } +}; + +class ConvertPtoCastPtrOp final : public OpConversionPattern { +public: + using OpConversionPattern::OpConversionPattern; + + LogicalResult matchAndRewrite(pto::CastPtrOp op, OpAdaptor adaptor, + ConversionPatternRewriter &rewriter) const override { + Type convertedResultType = getTypeConverter()->convertType(op.getResult().getType()); + if (!convertedResultType) { + return rewriter.notifyMatchFailure(op, "could not convert castptr result type"); + } + + Value input = adaptor.getInput(); + Type inputType = input.getType(); + if (inputType == convertedResultType) { + rewriter.replaceOp(op, input); + return success(); + } + + if (auto llvmPtrType = dyn_cast(convertedResultType)) { + if (isa(inputType)) { + rewriter.replaceOpWithNewOp(op, llvmPtrType, input); + return success(); + } + auto sourcePtrType = dyn_cast(inputType); + if (!sourcePtrType) { + return rewriter.notifyMatchFailure(op, "expected integer or LLVM pointer input"); + } + if (sourcePtrType.getAddressSpace() == llvmPtrType.getAddressSpace()) { + rewriter.replaceOpWithNewOp(op, llvmPtrType, input); + return success(); + } + return rewriter.notifyMatchFailure(op, "cross-address-space ptr casts are unsupported"); + } + + if (auto resultIntType = dyn_cast(convertedResultType)) { + if (isa(inputType)) { + rewriter.replaceOpWithNewOp(op, resultIntType, input); + return success(); + } + } + + return rewriter.notifyMatchFailure(op, "unsupported castptr conversion"); + } +}; + +class ConvertPtoLoadScalarOp final : public OpConversionPattern { +public: + using OpConversionPattern::OpConversionPattern; + + LogicalResult matchAndRewrite(pto::LoadScalarOp op, OpAdaptor adaptor, + ConversionPatternRewriter &rewriter) const override { + auto llvmPtrType = dyn_cast(adaptor.getPtr().getType()); + if (!llvmPtrType) { + return rewriter.notifyMatchFailure(op, "expected LLVM pointer operand"); + } + + Type convertedValueType = getTypeConverter()->convertType(op.getValue().getType()); + if (!convertedValueType) { + return rewriter.notifyMatchFailure(op, "could not convert load_scalar result type"); + } + + Value offset = adaptor.getOffset(); + if (offset.getType().isIndex()) { + offset = rewriter.create(op.getLoc(), rewriter.getI64Type(), offset); + } + + Value elemPtr = adaptor.getPtr(); + if (!matchPattern(offset, m_Zero())) { + elemPtr = rewriter.create(op.getLoc(), llvmPtrType, + normalizeGEPElementTypeForLLVMLowering(convertedValueType, rewriter), + adaptor.getPtr(), ValueRange{offset}); + } + + rewriter.replaceOpWithNewOp(op, convertedValueType, elemPtr, + getNaturalByteAlignment(convertedValueType)); + return success(); + } +}; + +class ConvertPtoStoreScalarOp final : public OpConversionPattern { +public: + using OpConversionPattern::OpConversionPattern; + + LogicalResult matchAndRewrite(pto::StoreScalarOp op, OpAdaptor adaptor, + ConversionPatternRewriter &rewriter) const override { + auto llvmPtrType = dyn_cast(adaptor.getPtr().getType()); + if (!llvmPtrType) { + return rewriter.notifyMatchFailure(op, "expected LLVM pointer operand"); + } + + Value offset = adaptor.getOffset(); + if (offset.getType().isIndex()) { + offset = rewriter.create(op.getLoc(), rewriter.getI64Type(), offset); + } + + Value elemPtr = adaptor.getPtr(); + if (!matchPattern(offset, m_Zero())) { + elemPtr = rewriter.create( + op.getLoc(), llvmPtrType, normalizeGEPElementTypeForLLVMLowering(adaptor.getValue().getType(), rewriter), + adaptor.getPtr(), ValueRange{offset}); + } + + rewriter.create(op.getLoc(), adaptor.getValue(), elemPtr, + getNaturalByteAlignment(adaptor.getValue().getType())); + rewriter.eraseOp(op); + return success(); + } +}; + +class ConvertPtoLoadOp final : public OpConversionPattern { +public: + ConvertPtoLoadOp(TypeConverter &typeConverter, MLIRContext *context, LoweringState &) + : OpConversionPattern(typeConverter, context) {} + + LogicalResult matchAndRewrite(pto::PTOLoadOp op, OpAdaptor adaptor, + ConversionPatternRewriter &rewriter) const override { + auto llvmPtrType = dyn_cast(adaptor.getPtr().getType()); + if (!llvmPtrType) { + return rewriter.notifyMatchFailure(op, "expected LLVM pointer operand"); + } + + Type convertedValueType = getTypeConverter()->convertType(op.getValue().getType()); + if (!convertedValueType) { + return rewriter.notifyMatchFailure(op, "could not convert load result type"); + } + + Value offset = adaptor.getOffset(); + if (offset.getType().isIndex()) { + offset = rewriter.create(op.getLoc(), rewriter.getI64Type(), offset); + } + + Value elemPtr = adaptor.getPtr(); + if (!matchPattern(offset, m_Zero())) { + elemPtr = rewriter.create(op.getLoc(), llvmPtrType, convertedValueType, adaptor.getPtr(), + ValueRange{offset}); + } + + rewriter.replaceOpWithNewOp(op, convertedValueType, elemPtr, + getNaturalByteAlignment(convertedValueType)); + return success(); + } +}; + +static Type getLdgCallResultType(Type valueType, Type convertedValueType, ConversionPatternRewriter &rewriter) { + if (auto intType = dyn_cast(valueType)) { + unsigned width = intType.getWidth(); + if (width == 8 || width == 16) { + return rewriter.getI32Type(); + } + return convertedValueType; + } + if (valueType.isF16() || valueType.isBF16() || valueType.isF32()) { + return rewriter.getI32Type(); + } + if (valueType.isF64()) { + return rewriter.getI64Type(); + } + if (pto::isPTOFloat8Type(valueType) || pto::isPTOHiFloat8Type(valueType)) { + return rewriter.getI32Type(); + } + if (pto::isPTOPackedLdgStgVectorType(valueType)) { + unsigned totalBits = pto::getPTOPackedLdgStgTotalBits(valueType); + if (totalBits == 16) { + return rewriter.getI32Type(); + } + if (totalBits == 32) { + return rewriter.getI32Type(); + } + if (totalBits == 64) { + return rewriter.getI64Type(); + } + } + return convertedValueType; +} + +static Value convertLdgCallResult(Location loc, Type valueType, Type convertedValueType, Value callResult, + ConversionPatternRewriter &rewriter) { + if (auto intType = dyn_cast(valueType)) { + unsigned width = intType.getWidth(); + if (width == 8 || width == 16) { + return rewriter.create(loc, rewriter.getIntegerType(width), callResult); + } + return callResult; + } + + if (valueType.isF16() || valueType.isBF16()) { + Value payload = rewriter.create(loc, rewriter.getI16Type(), callResult); + return rewriter.create(loc, convertedValueType, payload); + } + if (valueType.isF32() || valueType.isF64()) { + return rewriter.create(loc, convertedValueType, callResult); + } + if (pto::isPTOFloat8Type(valueType) || pto::isPTOHiFloat8Type(valueType)) { + Value payload = rewriter.create(loc, rewriter.getI8Type(), callResult); + return rewriter.create(loc, convertedValueType, payload); + } + if (pto::isPTOPackedLdgStgVectorType(valueType)) { + unsigned totalBits = pto::getPTOPackedLdgStgTotalBits(valueType); + if (totalBits == 16) { + Value trunc = rewriter.create(loc, rewriter.getI16Type(), callResult); + return rewriter.create(loc, convertedValueType, trunc); + } + return rewriter.create(loc, convertedValueType, callResult); + } + return callResult; +} + +static FailureOr preparePtoLdgAddress(pto::PTOLdgOp op, pto::PTOLdgOp::Adaptor adaptor, Type convertedValueType, + ConversionPatternRewriter &rewriter) { + auto llvmPtrType = dyn_cast(adaptor.getPtr().getType()); + if (!llvmPtrType) { + return failure(); + } + + Value offset = adaptor.getOffset(); + if (offset.getType().isIndex()) { + offset = rewriter.create(op.getLoc(), rewriter.getI64Type(), offset); + } + + Value elemPtr = adaptor.getPtr(); + if (!matchPattern(offset, m_Zero())) { + Type elementType = normalizeGEPElementTypeForLLVMLowering(convertedValueType, rewriter); + elemPtr = rewriter.create(op.getLoc(), llvmPtrType, elementType, adaptor.getPtr(), ValueRange{offset}); + } + + auto ptrType = cast(op.getPtr().getType()); + return reinterpretPointerToAddrSpace(op, elemPtr, static_cast(ptrType.getMemorySpace().getAddressSpace())); +} + +class ConvertPtoLdgOp final : public OpConversionPattern { +public: + ConvertPtoLdgOp(TypeConverter &typeConverter, MLIRContext *context, LoweringState &state) + : OpConversionPattern(typeConverter, context), state(state) {} + + LogicalResult matchAndRewrite(pto::PTOLdgOp op, OpAdaptor adaptor, + ConversionPatternRewriter &rewriter) const override { + if (!isa(adaptor.getPtr().getType())) { + return rewriter.notifyMatchFailure(op, "expected LLVM pointer operand"); + } + + Type convertedValueType = getTypeConverter()->convertType(op.getValue().getType()); + if (!convertedValueType) { + return rewriter.notifyMatchFailure(op, "could not convert ldg result type"); + } + + FailureOr ptr = preparePtoLdgAddress(op, adaptor, convertedValueType, rewriter); + if (failed(ptr)) { + return rewriter.notifyMatchFailure(op, "failed to map ldg pointer"); + } + + pto::L1Cache l1cache = op.getL1cacheAttr() ? op.getL1cacheAttr().getValue() : pto::L1Cache::Cache; + FailureOr calleeName = buildL1CacheLoadCallee(op.getContext(), op.getValue().getType(), l1cache); + if (failed(calleeName)) { + return rewriter.notifyMatchFailure(op, "unsupported ldg signature"); + } + + pto::LdL2Cache mode = op.getL2cacheAttr() ? op.getL2cacheAttr().getValue() : pto::LdL2Cache::NMFV; + Value modeValue = getI32Constant(rewriter, op.getLoc(), static_cast(mode)); + Type callResultType = getLdgCallResultType(op.getValue().getType(), convertedValueType, rewriter); + auto funcType = + rewriter.getFunctionType(TypeRange{ptr->getType(), rewriter.getI32Type()}, TypeRange{callResultType}); + auto call = + rewriter.create(op.getLoc(), *calleeName, TypeRange{callResultType}, ValueRange{*ptr, modeValue}); + state.plannedDecls.push_back(PlannedDecl{calleeName->str(), funcType}); + Value result = + convertLdgCallResult(op.getLoc(), op.getValue().getType(), convertedValueType, call.getResult(0), rewriter); + rewriter.replaceOp(op, result); + return success(); + } + +private: + LoweringState &state; +}; + +class ConvertPtoStoreOp final : public OpConversionPattern { +public: + ConvertPtoStoreOp(TypeConverter &typeConverter, MLIRContext *context, LoweringState &) + : OpConversionPattern(typeConverter, context) {} + + LogicalResult matchAndRewrite(pto::PTOStoreOp op, OpAdaptor adaptor, + ConversionPatternRewriter &rewriter) const override { + auto llvmPtrType = dyn_cast(adaptor.getPtr().getType()); + if (!llvmPtrType) { + return rewriter.notifyMatchFailure(op, "expected LLVM pointer operand"); + } + + Value offset = adaptor.getOffset(); + if (offset.getType().isIndex()) { + offset = rewriter.create(op.getLoc(), rewriter.getI64Type(), offset); + } + + Value elemPtr = adaptor.getPtr(); + if (!matchPattern(offset, m_Zero())) { + elemPtr = rewriter.create(op.getLoc(), llvmPtrType, adaptor.getValue().getType(), adaptor.getPtr(), + ValueRange{offset}); + } + + rewriter.replaceOpWithNewOp(op, adaptor.getValue(), elemPtr, + getNaturalByteAlignment(adaptor.getValue().getType())); + return success(); + } +}; + +static Value convertStgValue(Location loc, Type valueType, Value value, ConversionPatternRewriter &rewriter) { + if (auto intType = dyn_cast(valueType)) { + unsigned width = intType.getWidth(); + if (width == 8) { + return rewriter.create(loc, rewriter.getI32Type(), value); + } + if (width == 16) { + return rewriter.create(loc, rewriter.getF16Type(), value); + } + return value; + } + + if (pto::isPTOFloat8Type(valueType) || pto::isPTOHiFloat8Type(valueType)) { + Value payload = rewriter.create(loc, rewriter.getI8Type(), value); + return rewriter.create(loc, rewriter.getI32Type(), payload); + } + if (valueType.isBF16()) { + return rewriter.create(loc, rewriter.getF16Type(), value); + } + if (valueType.isF32()) { + return rewriter.create(loc, rewriter.getI32Type(), value); + } + if (valueType.isF64()) { + return rewriter.create(loc, rewriter.getI64Type(), value); + } + if (pto::isPTOPackedLdgStgVectorType(valueType)) { + unsigned totalBits = pto::getPTOPackedLdgStgTotalBits(valueType); + if (totalBits == 16) { + return rewriter.create(loc, rewriter.getF16Type(), value); + } + if (totalBits == 32) { + return rewriter.create(loc, rewriter.getI32Type(), value); + } + if (totalBits == 64) { + return rewriter.create(loc, rewriter.getI64Type(), value); + } + } + return value; +} + +class ConvertPtoStgOp final : public OpConversionPattern { +public: + ConvertPtoStgOp(TypeConverter &typeConverter, MLIRContext *context, LoweringState &state) + : OpConversionPattern(typeConverter, context), state(state) {} + + LogicalResult matchAndRewrite(pto::PTOStgOp op, OpAdaptor adaptor, + ConversionPatternRewriter &rewriter) const override { + auto llvmPtrType = dyn_cast(adaptor.getPtr().getType()); + if (!llvmPtrType) { + return rewriter.notifyMatchFailure(op, "expected LLVM pointer operand"); + } + + Value offset = adaptor.getOffset(); + if (offset.getType().isIndex()) { + offset = rewriter.create(op.getLoc(), rewriter.getI64Type(), offset); + } + + Value elemPtr = adaptor.getPtr(); + if (!matchPattern(offset, m_Zero())) { + elemPtr = rewriter.create( + op.getLoc(), llvmPtrType, normalizeGEPElementTypeForLLVMLowering(adaptor.getValue().getType(), rewriter), + adaptor.getPtr(), ValueRange{offset}); + } + + auto ptrTy = cast(op.getPtr().getType()); + FailureOr ptr = + reinterpretPointerToAddrSpace(op, elemPtr, static_cast(ptrTy.getMemorySpace().getAddressSpace())); + if (failed(ptr)) { + return rewriter.notifyMatchFailure(op, "failed to map stg pointer"); + } + + pto::L1Cache l1cache = op.getL1cacheAttr() ? op.getL1cacheAttr().getValue() : pto::L1Cache::Cache; + FailureOr calleeName = buildL1CacheStoreCallee(op.getContext(), op.getValue().getType(), l1cache); + if (failed(calleeName)) { + return rewriter.notifyMatchFailure(op, "unsupported stg signature"); + } + + pto::StL2Cache mode = op.getL2cacheAttr() ? op.getL2cacheAttr().getValue() : pto::StL2Cache::NMFV; + Value modeValue = getI32Constant(rewriter, op.getLoc(), static_cast(mode)); + Value storedValue = convertStgValue(op.getLoc(), op.getValue().getType(), adaptor.getValue(), rewriter); + auto funcType = + rewriter.getFunctionType(TypeRange{ptr->getType(), storedValue.getType(), rewriter.getI32Type()}, TypeRange{}); + rewriter.create(op.getLoc(), *calleeName, TypeRange{}, ValueRange{*ptr, storedValue, modeValue}); + state.plannedDecls.push_back(PlannedDecl{calleeName->str(), funcType}); + rewriter.eraseOp(op); + return success(); + } + +private: + LoweringState &state; +}; + +static std::string buildLdDevCalleeName(unsigned width) { return "llvm.hivm.LD.DEV.u" + std::to_string(width) + ".GM"; } + +static std::string buildStDevCalleeName(unsigned width) { return "llvm.hivm.ST.DEV.u" + std::to_string(width); } + +class ConvertPtoLdDevOp final : public OpConversionPattern { +public: + ConvertPtoLdDevOp(TypeConverter &typeConverter, MLIRContext *context, LoweringState &state) + : OpConversionPattern(typeConverter, context), state(state) {} + + LogicalResult matchAndRewrite(pto::PTOLdDevOp op, OpAdaptor adaptor, + ConversionPatternRewriter &rewriter) const override { + auto llvmPtrType = dyn_cast(adaptor.getPtr().getType()); + if (!llvmPtrType) { + return rewriter.notifyMatchFailure(op, "expected LLVM pointer operand"); + } + + auto valueType = dyn_cast(op.getValue().getType()); + if (!valueType) { + return rewriter.notifyMatchFailure(op, "expected integer result type"); + } + + Value offset = adaptor.getOffset(); + if (offset.getType().isIndex()) { + offset = rewriter.create(op.getLoc(), rewriter.getI64Type(), offset); + } + + Type convertedValueType = getTypeConverter()->convertType(op.getValue().getType()); + if (!convertedValueType) { + return rewriter.notifyMatchFailure(op, "could not convert ld_dev result type"); + } + + Value elemPtr = adaptor.getPtr(); + if (!matchPattern(offset, m_Zero())) { + elemPtr = rewriter.create(op.getLoc(), llvmPtrType, + normalizeGEPElementTypeForLLVMLowering(convertedValueType, rewriter), + adaptor.getPtr(), ValueRange{offset}); + } + + FailureOr gmPtr = reinterpretPointerToAddrSpace(op, elemPtr, static_cast(pto::AddressSpace::GM)); + if (failed(gmPtr)) { + return rewriter.notifyMatchFailure(op, "failed to map ld_dev GM pointer"); + } + + std::string calleeName = buildLdDevCalleeName(valueType.getWidth()); + Value intrinsicOffset = getI64Constant(rewriter, op.getLoc(), 0); + auto funcType = + rewriter.getFunctionType(TypeRange{gmPtr->getType(), rewriter.getI64Type()}, TypeRange{rewriter.getI64Type()}); + auto call = rewriter.create(op.getLoc(), calleeName, TypeRange{rewriter.getI64Type()}, + ValueRange{*gmPtr, intrinsicOffset}); + state.plannedDecls.push_back(PlannedDecl{calleeName, funcType}); + + Value result = call.getResult(0); + if (valueType.getWidth() < 64) { + result = rewriter.create(op.getLoc(), convertedValueType, result); + } + rewriter.replaceOp(op, result); + return success(); + } + +private: + LoweringState &state; +}; + +class ConvertPtoStDevOp final : public OpConversionPattern { +public: + ConvertPtoStDevOp(TypeConverter &typeConverter, MLIRContext *context, LoweringState &state) + : OpConversionPattern(typeConverter, context), state(state) {} + + LogicalResult matchAndRewrite(pto::PTOStDevOp op, OpAdaptor adaptor, + ConversionPatternRewriter &rewriter) const override { + auto llvmPtrType = dyn_cast(adaptor.getPtr().getType()); + if (!llvmPtrType) { + return rewriter.notifyMatchFailure(op, "expected LLVM pointer operand"); + } + + auto valueType = dyn_cast(op.getValue().getType()); + if (!valueType) { + return rewriter.notifyMatchFailure(op, "expected integer value type"); + } + + Value offset = adaptor.getOffset(); + if (offset.getType().isIndex()) { + offset = rewriter.create(op.getLoc(), rewriter.getI64Type(), offset); + } + + Value elemPtr = adaptor.getPtr(); + if (!matchPattern(offset, m_Zero())) { + elemPtr = rewriter.create( + op.getLoc(), llvmPtrType, normalizeGEPElementTypeForLLVMLowering(adaptor.getValue().getType(), rewriter), + adaptor.getPtr(), ValueRange{offset}); + } + + FailureOr gmPtr = reinterpretPointerToAddrSpace(op, elemPtr, static_cast(pto::AddressSpace::GM)); + if (failed(gmPtr)) { + return rewriter.notifyMatchFailure(op, "failed to map st_dev GM pointer"); + } + + Value payload = adaptor.getValue(); + if (valueType.getWidth() < 64) { + payload = rewriter.create(op.getLoc(), rewriter.getI64Type(), payload); + } + + std::string calleeName = buildStDevCalleeName(valueType.getWidth()); + Value intrinsicOffset = getI64Constant(rewriter, op.getLoc(), 0); + auto funcType = rewriter.getFunctionType(TypeRange{rewriter.getI64Type(), gmPtr->getType(), rewriter.getI64Type()}, + TypeRange{}); + rewriter.create(op.getLoc(), calleeName, TypeRange{}, ValueRange{payload, *gmPtr, intrinsicOffset}); + state.plannedDecls.push_back(PlannedDecl{calleeName, funcType}); + rewriter.eraseOp(op); + return success(); + } + +private: + LoweringState &state; +}; + +class ConvertVPTOTypedCarrierOp final : public ConversionPattern { +public: + ConvertVPTOTypedCarrierOp(TypeConverter &typeConverter, MLIRContext *context) + : ConversionPattern(typeConverter, MatchAnyOpTypeTag(), 1, context) {} + + LogicalResult matchAndRewrite(Operation *op, ArrayRef operands, + ConversionPatternRewriter &rewriter) const override { + if (isa(op)) { + return failure(); + } + Type propertyType; + if (auto allocaOp = dyn_cast(op)) { + propertyType = allocaOp.getElemType(); + } else if (auto gepOp = dyn_cast(op)) { + propertyType = gepOp.getElemType(); + } + if (!hasVPTOConvertibleType(op->getOperandTypes()) && !hasVPTOConvertibleType(op->getResultTypes()) && + !hasVPTOConvertibleType(propertyType)) { + return failure(); + } + if (op->getNumRegions() != 0) { + return rewriter.notifyMatchFailure(op, "region ops with VPTO types are handled structurally"); + } + + SmallVector convertedResultTypes; + if (failed(typeConverter->convertTypes(op->getResultTypes(), convertedResultTypes))) { + return rewriter.notifyMatchFailure(op, "failed to convert result types"); + } + OperationState state(op->getLoc(), op->getName()); + state.addOperands(operands); + state.addTypes(convertedResultTypes); + state.addAttributes(op->getAttrs()); + state.addSuccessors(op->getSuccessors()); + state.propertiesAttr = op->getPropertiesAsAttribute(); + Operation *converted = rewriter.create(state); + if (propertyType) { + Type convertedPropertyType = typeConverter->convertType(propertyType); + if (!convertedPropertyType) { + return rewriter.notifyMatchFailure(op, "failed to convert LLVM element type"); + } + if (auto allocaOp = dyn_cast(converted)) { + allocaOp.setElemType(convertedPropertyType); + } else { + cast(converted).setElemType(convertedPropertyType); + } + } + rewriter.replaceOp(op, converted->getResults()); + return success(); + } +}; +void populateVPTOTypePatterns(VPTOTypeConverter &typeConverter, RewritePatternSet &patterns, ConversionTarget &target, + LoweringState &state) { + MLIRContext *context = patterns.getContext(); + patterns.add(typeConverter, context); + patterns + .add( + typeConverter, context, state); + patterns.add(typeConverter, context); +} + +} // namespace mlir::pto::detail