Author: Erich Keane Date: 2026-10-09T05:53:18-07:00 New Revision: f8470a952862a7fd047a866a7f8168898b7008dd
URL: https://github.com/llvm/llvm-project/commit/f8470a952862a7fd047a866a7f8168898b7008dd DIFF: https://github.com/llvm/llvm-project/commit/f8470a952862a7fd047a866a7f8168898b7008dd.diff LOG: [CIR][Matrix] Implement vector splat/add for CIR lowering (#230185) This patch extends cir.add/cir.fadd to work on matrixes of int/float, and lower correctly to add/fadd in LLVMIR. It also extends cir.vec.splat to work with a matrix as well, so we could get mixed-matrix-int/float operations to work. I considered making this its own operation, however it is so nearly identical to cir.vec.splat (and will become more so as we extend matrix) that it didn't seem valuable to consider it separately, particularly as they lower to the same things in LLVM. One limitation: Splat is sometimes constant-folded during simplify. However, we don't yet have a constant attribute type for a matrix, so this is left for future work. Added: clang/test/CIR/CodeGen/matrix-add.c clang/test/CIR/Lowering/matrix-add.cir Modified: clang/include/clang/CIR/Dialect/IR/CIROps.td clang/include/clang/CIR/Dialect/IR/CIRTypeConstraints.td clang/lib/CIR/CodeGen/CIRGenBuilder.h clang/lib/CIR/CodeGen/CIRGenExprScalar.cpp clang/lib/CIR/Dialect/Transforms/CIRSimplify.cpp clang/lib/CIR/Lowering/DirectToLLVM/LowerToLLVM.cpp clang/test/CIR/IR/invalid-matrix.cir clang/test/CIR/IR/matrix.cir Removed: ################################################################################ diff --git a/clang/include/clang/CIR/Dialect/IR/CIROps.td b/clang/include/clang/CIR/Dialect/IR/CIROps.td index 1f264ce184c4a..c0de90183c255 100644 --- a/clang/include/clang/CIR/Dialect/IR/CIROps.td +++ b/clang/include/clang/CIR/Dialect/IR/CIROps.td @@ -2800,12 +2800,12 @@ class CIR_BinaryOpWithOverflowFlags<string mnemonic, Type type, //===----------------------------------------------------------------------===// def CIR_AddOp - : CIR_BinaryOpWithOverflowFlags<"add", CIR_AnyIntOrVecOfIntType> { + : CIR_BinaryOpWithOverflowFlags<"add", CIR_AnyIntOrVecOrMatrixOfIntType> { let summary = "Integer addition"; let description = [{ The `cir.add` operation performs addition on integer operands. Both - operands and the result must have the same integer or vector-of-integer - type. + operands and the result must have the same integer, vector-of-integer or + matrix-of-integer type. The optional `nsw` (no signed wrap) and `nuw` (no unsigned wrap) unit attributes indicate that the result is poison if signed or unsigned @@ -2821,6 +2821,7 @@ def CIR_AddOp %2 = cir.add nuw %a, %b : !u32i %3 = cir.add sat %a, %b : !s32i %4 = cir.add %va, %vb : !cir.vector<4 x !s32i> + %5 = cir.add %ma, %mb : !cir.matrix<3 x 3 x !s32i> ``` }]; } @@ -2943,8 +2944,9 @@ def CIR_RemOp : CIR_BinaryOp<"rem", CIR_AnyIntOrVecOfIntType> { // // The optional `fenv` attribute describes constraints on the floating-point // handling of the operation. -class CIR_FPBinaryOp<string mnemonic, list<Trait> traits = []> - : CIR_BinaryOp<mnemonic, CIR_AnyFloatOrVecOfFloatType, +class CIR_FPBinaryOp<string mnemonic, list<Trait> traits = [], + Type type = CIR_AnyFloatOrVecOfFloatType> + : CIR_BinaryOp<mnemonic, type, !listconcat(CIR_FenvOpTraits, traits), CIR_DynamicMemoryEffects> { let arguments = !con(commonArgs, (ins OptionalAttr<CIR_FenvAttr>:$fenv)); @@ -2965,12 +2967,13 @@ class CIR_FPBinaryOp<string mnemonic, list<Trait> traits = []> // FAddOp //===----------------------------------------------------------------------===// -def CIR_FAddOp : CIR_FPBinaryOp<"fadd"> { +def CIR_FAddOp + : CIR_FPBinaryOp<"fadd", [], CIR_AnyFloatOrVecOrMatrixOfFloatType> { let summary = "Floating-point addition"; let description = [{ The `cir.fadd` operation performs floating-point addition on its operands. - Both operands and the result must have the same floating-point scalar or - vector-of-float type. + Both operands and the result must have the same floating-point scalar, + vector-of-float or matrix-of-float type. Example: @@ -2978,6 +2981,7 @@ def CIR_FAddOp : CIR_FPBinaryOp<"fadd"> { %0 = cir.fadd %a, %b : !cir.float %1 = cir.fadd %a, %b : !cir.double %2 = cir.fadd %va, %vb : !cir.vector<4 x !cir.float> + %3 = cir.fadd %ma, %mb : !cir.matrix<3 x 3 x !cir.float> ``` }]; @@ -6553,12 +6557,16 @@ def CIR_VecTernaryOp : CIR_Op<"vec.ternary", [ def CIR_VecSplatOp : CIR_Op<"vec.splat", [ Pure, TypesMatchWith<"type of 'value' matches element type of 'result'", - "result", "value", "mlir::cast<cir::VectorType>($_self).getElementType()"> + "result", "value", + "mlir::isa<cir::VectorType>($_self) " + "? mlir::cast<cir::VectorType>($_self).getElementType() " + ": mlir::cast<cir::MatrixType>($_self).getElementType()"> ]> { - let summary = "Convert a scalar into a vector"; + let summary = "Convert a scalar into a vector or matrix"; let description = [{ - The `cir.vec.splat` operation creates a vector value from a scalar value. - All elements of the vector have the same value, that of the given scalar. + The `cir.vec.splat` operation creates a vector or matrix value from a + scalar value. All elements of the result have the same value, that of the + given scalar. It's a separate operation from `cir.vec.create` because more efficient LLVM IR can be generated for it, and because some optimization and @@ -6568,11 +6576,13 @@ def CIR_VecSplatOp : CIR_Op<"vec.splat", [ ``` %value = cir.const #cir.int<3> : !s32i %value_vec = cir.vec.splat %value : !s32i, !cir.vector<4 x !s32i> + %value_mat = cir.vec.splat %value : !s32i, !cir.matrix<2 x 2 x !s32i> ``` }]; let arguments = (ins CIR_VectorElementType:$value); - let results = (outs CIR_VectorType:$result); + let results = (outs AnyTypeOf<[CIR_VectorType, CIR_MatrixType], + "vector or matrix type">:$result); let assemblyFormat = [{ $value `:` type($value) `,` qualified(type($result)) attr-dict diff --git a/clang/include/clang/CIR/Dialect/IR/CIRTypeConstraints.td b/clang/include/clang/CIR/Dialect/IR/CIRTypeConstraints.td index a3de83f6cab63..97d7b30236b82 100644 --- a/clang/include/clang/CIR/Dialect/IR/CIRTypeConstraints.td +++ b/clang/include/clang/CIR/Dialect/IR/CIRTypeConstraints.td @@ -391,10 +391,38 @@ def CIR_AnyBitwiseType def CIR_MatrixElementType : AnyTypeOf<[CIR_AnyBoolType, CIR_AnyIntOrFloatType], - "any cir boolean, integer, floating point"> { + "any cir boolean, integer, floating point type"> { let cppFunctionName = "isValidMatrixTypeElementType"; } +def CIR_AnyMatrixType : CIR_TypeBase<"::cir::MatrixType", "matrix type">; + +class CIR_MatrixElementTypePred<Pred pred> : SubstLeaves<"$_self", + "::mlir::cast<::cir::MatrixType>($_self).getElementType()", pred>; + +class CIR_MatrixTypeOf<list<Type> types, string summary = ""> + : CIR_ConfinedType<CIR_AnyMatrixType, + [Or<!foreach(type, types, CIR_MatrixElementTypePred<type.predicate>)>], + !if(!empty(summary), + "matrix of " # CIR_TypeSummaries<types>.value, + summary)>; + +def CIR_MatrixOfIntType : CIR_MatrixTypeOf<[CIR_AnyIntType]>; +def CIR_MatrixOfFloatType : CIR_MatrixTypeOf<[CIR_AnyFloatType]>; + +def CIR_AnyIntOrVecOrMatrixOfIntType + : AnyTypeOf<[CIR_AnyIntType, CIR_VectorOfIntType, CIR_MatrixOfIntType], + "integer or vector/matrix of integer type"> { + let cppFunctionName = "isIntOrVectorOrMatrixOfIntType"; +} + +def CIR_AnyFloatOrVecOrMatrixOfFloatType + : AnyTypeOf<[CIR_AnyFloatType, CIR_VectorOfFloatType, + CIR_MatrixOfFloatType], + "floating point or vector/matrix of floating point type"> { + let cppFunctionName = "isFPOrVectorOrMatrixOfFPType"; +} + //===----------------------------------------------------------------------===// // Data member type predicates //===----------------------------------------------------------------------===// diff --git a/clang/lib/CIR/CodeGen/CIRGenBuilder.h b/clang/lib/CIR/CodeGen/CIRGenBuilder.h index b3b3805e8bb2b..12695704786bc 100644 --- a/clang/lib/CIR/CodeGen/CIRGenBuilder.h +++ b/clang/lib/CIR/CodeGen/CIRGenBuilder.h @@ -852,6 +852,20 @@ class CIRGenBuilderTy : public cir::CIRBaseBuilderTy { stride, isVolatile); } + std::pair<mlir::Value, mlir::Value> + splatMatrixOpOperandsIfNecessary(mlir::Location loc, mlir::Value lhs, + mlir::Value rhs) { + assert(mlir::isa<cir::MatrixType>(lhs.getType()) || + mlir::isa<cir::MatrixType>(rhs.getType())); + + if (!mlir::isa<cir::MatrixType>(lhs.getType())) + lhs = cir::VecSplatOp::create(*this, loc, rhs.getType(), lhs); + else if (!mlir::isa<cir::MatrixType>(rhs.getType())) + rhs = cir::VecSplatOp::create(*this, loc, lhs.getType(), rhs); + + return {lhs, rhs}; + } + template <typename... Operands> mlir::Value emitIntrinsicCallOp(mlir::Location loc, const llvm::StringRef str, const mlir::Type &resTy, Operands &&...op) { diff --git a/clang/lib/CIR/CodeGen/CIRGenExprScalar.cpp b/clang/lib/CIR/CodeGen/CIRGenExprScalar.cpp index 836f27cf5345e..8ae0be29d6d2b 100644 --- a/clang/lib/CIR/CodeGen/CIRGenExprScalar.cpp +++ b/clang/lib/CIR/CodeGen/CIRGenExprScalar.cpp @@ -2413,9 +2413,15 @@ mlir::Value ScalarExprEmitter::emitAdd(const BinOpInfo &ops) { } } if (ops.fullType->isConstantMatrixType()) { - assert(!cir::MissingFeatures::matrixType()); - cgf.cgm.errorNYI("ScalarExprEmitter::emitAdd: matrix types"); - return {}; + // Like llvm::MatrixBuilder::CreateAdd, splat a scalar operand to the matrix + // type before adding. + auto [lhs, rhs] = + builder.splatMatrixOpOperandsIfNecessary(loc, ops.lhs, ops.rhs); + + CIRGenFunction::CIRGenFPOptionsRAII fpOptsRAII(cgf, ops.fpFeatures); + if (cir::isFPOrVectorOrMatrixOfFPType(lhs.getType())) + return builder.createFAdd(loc, lhs, rhs); + return builder.createAdd(loc, lhs, rhs); } if (ops.compType->isUnsignedIntegerType() && diff --git a/clang/lib/CIR/Dialect/Transforms/CIRSimplify.cpp b/clang/lib/CIR/Dialect/Transforms/CIRSimplify.cpp index ad1e58e6ab5cd..9d569be7b80cb 100644 --- a/clang/lib/CIR/Dialect/Transforms/CIRSimplify.cpp +++ b/clang/lib/CIR/Dialect/Transforms/CIRSimplify.cpp @@ -371,7 +371,12 @@ struct SimplifyVecSplat : public OpRewritePattern<VecSplatOp> { !mlir::isa_and_nonnull<cir::FPAttr>(value)) return mlir::failure(); - cir::VectorType resultType = op.getResult().getType(); + // FIXME(CIR): We should consider making a matrix constant attribute so that + // we can simplify it here too. + assert(!MissingFeatures::matrixType()); + auto resultType = mlir::dyn_cast<cir::VectorType>(op.getResult().getType()); + if (!resultType) + return mlir::failure(); SmallVector<mlir::Attribute, 16> elements(resultType.getSize(), value); auto constVecAttr = cir::ConstVectorAttr::get( resultType, mlir::ArrayAttr::get(getContext(), elements)); diff --git a/clang/lib/CIR/Lowering/DirectToLLVM/LowerToLLVM.cpp b/clang/lib/CIR/Lowering/DirectToLLVM/LowerToLLVM.cpp index 7a1e61b526264..d87145b32a692 100644 --- a/clang/lib/CIR/Lowering/DirectToLLVM/LowerToLLVM.cpp +++ b/clang/lib/CIR/Lowering/DirectToLLVM/LowerToLLVM.cpp @@ -68,11 +68,11 @@ namespace direct { //===----------------------------------------------------------------------===// namespace { -/// If the given type is a vector type, return the vector's element type. +/// If the given type is a vector or matrix type, return its element type. /// Otherwise return the given type unchanged. -mlir::Type elementTypeIfVector(mlir::Type type) { +mlir::Type elementTypeIfVectorOrMatrix(mlir::Type type) { return llvm::TypeSwitch<mlir::Type, mlir::Type>(type) - .Case<cir::VectorType, mlir::VectorType>( + .Case<cir::VectorType, cir::MatrixType, mlir::VectorType>( [](auto p) { return p.getElementType(); }) .Default([](mlir::Type p) { return p; }); } @@ -1734,9 +1734,9 @@ mlir::LogicalResult CIRToLLVMCastOpLowering::matchAndRewrite( mlir::Value llvmSrcVal = adaptor.getSrc(); mlir::Type llvmDstType = getTypeConverter()->convertType(dstType); cir::IntType srcIntType = - mlir::cast<cir::IntType>(elementTypeIfVector(srcType)); + mlir::cast<cir::IntType>(elementTypeIfVectorOrMatrix(srcType)); cir::IntType dstIntType = - mlir::cast<cir::IntType>(elementTypeIfVector(dstType)); + mlir::cast<cir::IntType>(elementTypeIfVectorOrMatrix(dstType)); rewriter.replaceOp(castOp, getLLVMIntCast(rewriter, llvmSrcVal, llvmDstType, srcIntType.isUnsigned(), srcIntType.getWidth(), @@ -1747,8 +1747,8 @@ mlir::LogicalResult CIRToLLVMCastOpLowering::matchAndRewrite( mlir::Value llvmSrcVal = adaptor.getSrc(); mlir::Type llvmDstTy = getTypeConverter()->convertType(castOp.getType()); - mlir::Type srcTy = elementTypeIfVector(castOp.getSrc().getType()); - mlir::Type dstTy = elementTypeIfVector(castOp.getType()); + mlir::Type srcTy = elementTypeIfVectorOrMatrix(castOp.getSrc().getType()); + mlir::Type dstTy = elementTypeIfVectorOrMatrix(castOp.getType()); if (!mlir::isa<cir::FPTypeInterface>(dstTy) || !mlir::isa<cir::FPTypeInterface>(srcTy)) @@ -1813,8 +1813,9 @@ mlir::LogicalResult CIRToLLVMCastOpLowering::matchAndRewrite( mlir::Type llvmDstTy = getTypeConverter()->convertType(dstTy); // Compare element widths so this also handles vector bool -> int casts. auto srcElemTy = mlir::cast<mlir::IntegerType>( - elementTypeIfVector(llvmSrcVal.getType())); - auto dstElemTy = mlir::cast<cir::IntType>(elementTypeIfVector(dstTy)); + elementTypeIfVectorOrMatrix(llvmSrcVal.getType())); + auto dstElemTy = + mlir::cast<cir::IntType>(elementTypeIfVectorOrMatrix(dstTy)); if (srcElemTy.getWidth() == dstElemTy.getWidth()) rewriter.replaceOpWithNewOp<mlir::LLVM::BitcastOp>(castOp, llvmDstTy, @@ -1836,9 +1837,9 @@ mlir::LogicalResult CIRToLLVMCastOpLowering::matchAndRewrite( mlir::Type dstTy = castOp.getType(); mlir::Value llvmSrcVal = adaptor.getSrc(); mlir::Type llvmDstTy = getTypeConverter()->convertType(dstTy); - bool isSigned = - mlir::cast<cir::IntType>(elementTypeIfVector(castOp.getSrc().getType())) - .isSigned(); + bool isSigned = mlir::cast<cir::IntType>( + elementTypeIfVectorOrMatrix(castOp.getSrc().getType())) + .isSigned(); if (cir::FenvAttr fenv = castOp.getFenvAttr()) { return lowerToConstrainedFPIntrinsic( castOp, llvmSrcVal, fenv, llvmDstTy, rewriter, @@ -1857,7 +1858,7 @@ mlir::LogicalResult CIRToLLVMCastOpLowering::matchAndRewrite( mlir::Value llvmSrcVal = adaptor.getSrc(); mlir::Type llvmDstTy = getTypeConverter()->convertType(dstTy); bool isSigned = - mlir::cast<cir::IntType>(elementTypeIfVector(castOp.getType())) + mlir::cast<cir::IntType>(elementTypeIfVectorOrMatrix(castOp.getType())) .isSigned(); if (cir::FenvAttr fenv = castOp.getFenvAttr()) { return lowerToConstrainedFPIntrinsic( @@ -3404,7 +3405,7 @@ mlir::LogicalResult CIRToLLVMMinusOpLowering::matchAndRewrite( mlir::LogicalResult CIRToLLVMNotOpLowering::matchAndRewrite( cir::NotOp op, OpAdaptor adaptor, mlir::ConversionPatternRewriter &rewriter) const { - mlir::Type elementType = elementTypeIfVector(op.getType()); + mlir::Type elementType = elementTypeIfVectorOrMatrix(op.getType()); bool isVector = mlir::isa<cir::VectorType>(op.getType()); mlir::Type llvmType = adaptor.getInput().getType(); mlir::Location loc = op.getLoc(); @@ -3476,7 +3477,7 @@ template <typename UIntSatOp, typename SIntSatOp, typename IntOp, static mlir::LogicalResult lowerSaturatableArithOp(CIROp op, mlir::Value lhs, mlir::Value rhs, mlir::ConversionPatternRewriter &rewriter) { - const mlir::Type eltType = elementTypeIfVector(op.getRhs().getType()); + const mlir::Type eltType = elementTypeIfVectorOrMatrix(op.getRhs().getType()); assert(cir::isIntOrBoolType(eltType) && "saturatable arith op expects integer operand types"); if (op.getSaturated()) { @@ -3509,7 +3510,8 @@ mlir::LogicalResult CIRToLLVMSubOpLowering::matchAndRewrite( mlir::LogicalResult CIRToLLVMMulOpLowering::matchAndRewrite( cir::MulOp op, OpAdaptor adaptor, mlir::ConversionPatternRewriter &rewriter) const { - assert(cir::isIntOrBoolType(elementTypeIfVector(op.getRhs().getType())) && + assert(cir::isIntOrBoolType( + elementTypeIfVectorOrMatrix(op.getRhs().getType())) && "cir.mul expects integer operand types"); rewriter.replaceOpWithNewOp<mlir::LLVM::MulOp>( op, adaptor.getLhs(), adaptor.getRhs(), intOverflowFlag(op)); @@ -3521,7 +3523,7 @@ template <typename UIntOp, typename SIntOp, typename CIROp> static mlir::LogicalResult lowerIntBinaryOp(CIROp op, mlir::Value lhs, mlir::Value rhs, mlir::ConversionPatternRewriter &rewriter) { - const mlir::Type eltType = elementTypeIfVector(op.getRhs().getType()); + const mlir::Type eltType = elementTypeIfVectorOrMatrix(op.getRhs().getType()); assert(cir::isIntOrBoolType(eltType) && "integer binary op expects integer operand types"); if (isIntTypeUnsigned(eltType)) @@ -3551,7 +3553,7 @@ lowerMinMaxOp(CIROp op, typename CIROp::Adaptor adaptor, mlir::ConversionPatternRewriter &rewriter) { const mlir::Value lhs = adaptor.getLhs(); const mlir::Value rhs = adaptor.getRhs(); - if (isIntTypeUnsigned(elementTypeIfVector(op.getRhs().getType()))) + if (isIntTypeUnsigned(elementTypeIfVectorOrMatrix(op.getRhs().getType()))) rewriter.replaceOpWithNewOp<UIntOp>(op, lhs, rhs); else rewriter.replaceOpWithNewOp<SIntOp>(op, lhs, rhs); @@ -5211,8 +5213,7 @@ mlir::LogicalResult CIRToLLVMVecSplatOpLowering::matchAndRewrite( // element in the vector. Start with an undef vector. Insert the value into // the first element. Then use a `shufflevector` with a mask of all 0 to // fill out the entire vector with that value. - cir::VectorType vecTy = op.getType(); - mlir::Type llvmTy = typeConverter->convertType(vecTy); + mlir::Type llvmTy = typeConverter->convertType(op.getType()); mlir::Location loc = op.getLoc(); mlir::Value poison = mlir::LLVM::PoisonOp::create(rewriter, loc, llvmTy); @@ -5246,7 +5247,8 @@ mlir::LogicalResult CIRToLLVMVecSplatOpLowering::matchAndRewrite( mlir::LLVM::ConstantOp::create(rewriter, loc, rewriter.getI64Type(), 0); mlir::Value oneElement = mlir::LLVM::InsertElementOp::create( rewriter, loc, poison, elementValue, indexValue); - SmallVector<int32_t> zeroValues(vecTy.getSize(), 0); + SmallVector<int32_t> zeroValues( + mlir::cast<mlir::VectorType>(llvmTy).getNumElements(), 0); rewriter.replaceOpWithNewOp<mlir::LLVM::ShuffleVectorOp>(op, oneElement, poison, zeroValues); return mlir::success(); diff --git a/clang/test/CIR/CodeGen/matrix-add.c b/clang/test/CIR/CodeGen/matrix-add.c new file mode 100644 index 0000000000000..adfe1e671ad23 --- /dev/null +++ b/clang/test/CIR/CodeGen/matrix-add.c @@ -0,0 +1,207 @@ +// RUN: %clang_cc1 -triple x86_64-unknown-linux-gnu -fenable-matrix -fclangir -emit-cir %s -o %t.cir +// RUN: FileCheck --input-file=%t.cir %s -check-prefix=CIR +// RUN: %clang_cc1 -triple x86_64-unknown-linux-gnu -fenable-matrix -fclangir -emit-llvm %s -o %t-cir.ll +// RUN: FileCheck --input-file=%t-cir.ll %s -check-prefix=LLVM +// RUN: %clang_cc1 -triple x86_64-unknown-linux-gnu -fenable-matrix -emit-llvm %s -o %t.ll +// RUN: FileCheck --input-file=%t.ll %s -check-prefix=LLVM + +typedef double dx5x5_t __attribute__((matrix_type(5, 5))); +typedef float fx2x3_t __attribute__((matrix_type(2, 3))); +typedef int ix9x3_t __attribute__((matrix_type(9, 3))); +typedef unsigned long long ullx4x2_t __attribute__((matrix_type(4, 2))); + +void add_matrix_matrix_double() { + dx5x5_t a; + dx5x5_t b; + dx5x5_t c; + a = b + c; +} + +// CIR-LABEL: cir.func {{.*}}@add_matrix_matrix_double +// CIR: %[[B:.*]] = cir.load {{.*}} : !cir.ptr<!cir.matrix<5 x 5 x !cir.double>>, !cir.matrix<5 x 5 x !cir.double> +// CIR: %[[C:.*]] = cir.load {{.*}} : !cir.ptr<!cir.matrix<5 x 5 x !cir.double>>, !cir.matrix<5 x 5 x !cir.double> +// CIR: %[[RES:.*]] = cir.fadd %[[B]], %[[C]] : !cir.matrix<5 x 5 x !cir.double> +// CIR: cir.store {{.*}} %[[RES]], {{.*}} : !cir.matrix<5 x 5 x !cir.double>, !cir.ptr<!cir.matrix<5 x 5 x !cir.double>> + +// LLVM-LABEL: define {{.*}}void @add_matrix_matrix_double( +// LLVM: %[[B:.*]] = load <25 x double>, ptr {{.*}}, align 8 +// LLVM: %[[C:.*]] = load <25 x double>, ptr {{.*}}, align 8 +// LLVM: %[[RES:.*]] = fadd <25 x double> %[[B]], %[[C]] +// LLVM: store <25 x double> %[[RES]], ptr {{.*}}, align 8 + +void add_compound_assign_matrix_double() { + dx5x5_t a; + dx5x5_t b; + a += b; +} + +// CIR-LABEL: cir.func {{.*}}@add_compound_assign_matrix_double +// CIR: %[[B:.*]] = cir.load {{.*}} : !cir.ptr<!cir.matrix<5 x 5 x !cir.double>>, !cir.matrix<5 x 5 x !cir.double> +// CIR: %[[A:.*]] = cir.load {{.*}} : !cir.ptr<!cir.matrix<5 x 5 x !cir.double>>, !cir.matrix<5 x 5 x !cir.double> +// CIR: %[[RES:.*]] = cir.fadd %[[A]], %[[B]] : !cir.matrix<5 x 5 x !cir.double> +// CIR: cir.store {{.*}} %[[RES]], {{.*}} : !cir.matrix<5 x 5 x !cir.double>, !cir.ptr<!cir.matrix<5 x 5 x !cir.double>> + +// LLVM-LABEL: define {{.*}}void @add_compound_assign_matrix_double( +// LLVM: %[[B:.*]] = load <25 x double>, ptr {{.*}}, align 8 +// LLVM: %[[A:.*]] = load <25 x double>, ptr {{.*}}, align 8 +// LLVM: %[[RES:.*]] = fadd <25 x double> %[[A]], %[[B]] +// LLVM: store <25 x double> %[[RES]], ptr {{.*}}, align 8 + +void add_matrix_matrix_float() { + fx2x3_t a; + fx2x3_t b; + fx2x3_t c; + a = b + c; +} + +// CIR-LABEL: cir.func {{.*}}@add_matrix_matrix_float +// CIR: %[[B:.*]] = cir.load {{.*}} : !cir.ptr<!cir.matrix<2 x 3 x !cir.float>>, !cir.matrix<2 x 3 x !cir.float> +// CIR: %[[C:.*]] = cir.load {{.*}} : !cir.ptr<!cir.matrix<2 x 3 x !cir.float>>, !cir.matrix<2 x 3 x !cir.float> +// CIR: %[[RES:.*]] = cir.fadd %[[B]], %[[C]] : !cir.matrix<2 x 3 x !cir.float> + +// LLVM-LABEL: define {{.*}}void @add_matrix_matrix_float( +// LLVM: %[[B:.*]] = load <6 x float>, ptr {{.*}}, align 4 +// LLVM: %[[C:.*]] = load <6 x float>, ptr {{.*}}, align 4 +// LLVM: %[[RES:.*]] = fadd <6 x float> %[[B]], %[[C]] +// LLVM: store <6 x float> %[[RES]], ptr {{.*}}, align 4 + +void add_matrix_matrix_int() { + ix9x3_t a; + ix9x3_t b; + ix9x3_t c; + a = b + c; +} + +// CIR-LABEL: cir.func {{.*}}@add_matrix_matrix_int +// CIR: %[[B:.*]] = cir.load {{.*}} : !cir.ptr<!cir.matrix<9 x 3 x !s32i>>, !cir.matrix<9 x 3 x !s32i> +// CIR: %[[C:.*]] = cir.load {{.*}} : !cir.ptr<!cir.matrix<9 x 3 x !s32i>>, !cir.matrix<9 x 3 x !s32i> +// CIR: %[[RES:.*]] = cir.add %[[B]], %[[C]] : !cir.matrix<9 x 3 x !s32i> +// CIR: cir.store {{.*}} %[[RES]], {{.*}} : !cir.matrix<9 x 3 x !s32i>, !cir.ptr<!cir.matrix<9 x 3 x !s32i>> + +// LLVM-LABEL: define {{.*}}void @add_matrix_matrix_int( +// LLVM: %[[B:.*]] = load <27 x i32>, ptr {{.*}}, align 4 +// LLVM: %[[C:.*]] = load <27 x i32>, ptr {{.*}}, align 4 +// LLVM: %[[RES:.*]] = add <27 x i32> %[[B]], %[[C]] +// LLVM: store <27 x i32> %[[RES]], ptr {{.*}}, align 4 + +void add_compound_assign_matrix_int() { + ix9x3_t a; + ix9x3_t b; + a += b; +} + +// CIR-LABEL: cir.func {{.*}}@add_compound_assign_matrix_int +// CIR: %[[B:.*]] = cir.load {{.*}} : !cir.ptr<!cir.matrix<9 x 3 x !s32i>>, !cir.matrix<9 x 3 x !s32i> +// CIR: %[[A:.*]] = cir.load {{.*}} : !cir.ptr<!cir.matrix<9 x 3 x !s32i>>, !cir.matrix<9 x 3 x !s32i> +// CIR: %[[RES:.*]] = cir.add %[[A]], %[[B]] : !cir.matrix<9 x 3 x !s32i> + +// LLVM-LABEL: define {{.*}}void @add_compound_assign_matrix_int( +// LLVM: %[[B:.*]] = load <27 x i32>, ptr {{.*}}, align 4 +// LLVM: %[[A:.*]] = load <27 x i32>, ptr {{.*}}, align 4 +// LLVM: %[[RES:.*]] = add <27 x i32> %[[A]], %[[B]] +// LLVM: store <27 x i32> %[[RES]], ptr {{.*}}, align 4 + +void add_matrix_matrix_unsigned() { + ullx4x2_t a; + ullx4x2_t b; + ullx4x2_t c; + a = b + c; +} + +// CIR-LABEL: cir.func {{.*}}@add_matrix_matrix_unsigned +// CIR: %[[RES:.*]] = cir.add %{{.*}}, %{{.*}} : !cir.matrix<4 x 2 x !u64i> + +// LLVM-LABEL: define {{.*}}void @add_matrix_matrix_unsigned( +// LLVM: %[[RES:.*]] = add <8 x i64> %{{.*}}, %{{.*}} + +void add_matrix_scalar_double_double() { + dx5x5_t a; + double vd; + a = a + vd; +} + +// CIR-LABEL: cir.func {{.*}}@add_matrix_scalar_double_double +// CIR: %[[MAT:.*]] = cir.load {{.*}} : !cir.ptr<!cir.matrix<5 x 5 x !cir.double>>, !cir.matrix<5 x 5 x !cir.double> +// CIR: %[[SCALAR:.*]] = cir.load {{.*}} : !cir.ptr<!cir.double>, !cir.double +// CIR: %[[SPLAT:.*]] = cir.vec.splat %[[SCALAR]] : !cir.double, !cir.matrix<5 x 5 x !cir.double> +// CIR: %[[RES:.*]] = cir.fadd %[[MAT]], %[[SPLAT]] : !cir.matrix<5 x 5 x !cir.double> + +// LLVM-LABEL: define {{.*}}void @add_matrix_scalar_double_double( +// LLVM: %[[MAT:.*]] = load <25 x double>, ptr {{.*}}, align 8 +// LLVM: %[[SCALAR:.*]] = load double, ptr {{.*}}, align 8 +// LLVM: %[[EMBED:.*]] = insertelement <25 x double> poison, double %[[SCALAR]], i64 0 +// LLVM: %[[SPLAT:.*]] = shufflevector <25 x double> %[[EMBED]], <25 x double> poison, <25 x i32> zeroinitializer +// LLVM: %[[RES:.*]] = fadd <25 x double> %[[MAT]], %[[SPLAT]] +// LLVM: store <25 x double> %[[RES]], ptr {{.*}}, align 8 + +void add_scalar_matrix_double_double() { + dx5x5_t a; + double vd; + a = vd + a; +} + +// CIR-LABEL: cir.func {{.*}}@add_scalar_matrix_double_double +// CIR: %[[SCALAR:.*]] = cir.load {{.*}} : !cir.ptr<!cir.double>, !cir.double +// CIR: %[[MAT:.*]] = cir.load {{.*}} : !cir.ptr<!cir.matrix<5 x 5 x !cir.double>>, !cir.matrix<5 x 5 x !cir.double> +// CIR: %[[SPLAT:.*]] = cir.vec.splat %[[SCALAR]] : !cir.double, !cir.matrix<5 x 5 x !cir.double> +// CIR: %[[RES:.*]] = cir.fadd %[[SPLAT]], %[[MAT]] : !cir.matrix<5 x 5 x !cir.double> + +// LLVM-LABEL: define {{.*}}void @add_scalar_matrix_double_double( +// LLVM: %[[SCALAR:.*]] = load double, ptr {{.*}}, align 8 +// LLVM: %[[MAT:.*]] = load <25 x double>, ptr {{.*}}, align 8 +// LLVM: %[[EMBED:.*]] = insertelement <25 x double> poison, double %[[SCALAR]], i64 0 +// LLVM: %[[SPLAT:.*]] = shufflevector <25 x double> %[[EMBED]], <25 x double> poison, <25 x i32> zeroinitializer +// LLVM: %[[RES:.*]] = fadd <25 x double> %[[SPLAT]], %[[MAT]] + +void add_matrix_scalar_double_float() { + dx5x5_t a; + float vf; + a = a + vf; +} + +// CIR-LABEL: cir.func {{.*}}@add_matrix_scalar_double_float +// CIR: %[[SCALAR:.*]] = cir.load {{.*}} : !cir.ptr<!cir.float>, !cir.float +// CIR: %[[EXT:.*]] = cir.cast floating %[[SCALAR]] : !cir.float -> !cir.double +// CIR: %[[SPLAT:.*]] = cir.vec.splat %[[EXT]] : !cir.double, !cir.matrix<5 x 5 x !cir.double> +// CIR: cir.fadd %{{.*}}, %[[SPLAT]] : !cir.matrix<5 x 5 x !cir.double> + +// LLVM-LABEL: define {{.*}}void @add_matrix_scalar_double_float( +// LLVM: %[[SCALAR:.*]] = load float, ptr {{.*}}, align 4 +// LLVM: %[[EXT:.*]] = fpext float %[[SCALAR]] to double +// LLVM: %[[EMBED:.*]] = insertelement <25 x double> poison, double %[[EXT]], i64 0 +// LLVM: %[[SPLAT:.*]] = shufflevector <25 x double> %[[EMBED]], <25 x double> poison, <25 x i32> zeroinitializer +// LLVM: fadd <25 x double> %{{.*}}, %[[SPLAT]] + +void add_matrix_scalar_int_long() { + ix9x3_t a; + long vl; + a = a + vl; +} + +// CIR-LABEL: cir.func {{.*}}@add_matrix_scalar_int_long +// CIR: %[[SCALAR:.*]] = cir.load {{.*}} : !cir.ptr<!s64i>, !s64i +// CIR: %[[TRUNC:.*]] = cir.cast integral %[[SCALAR]] : !s64i -> !s32i +// CIR: %[[SPLAT:.*]] = cir.vec.splat %[[TRUNC]] : !s32i, !cir.matrix<9 x 3 x !s32i> +// CIR: cir.add %{{.*}}, %[[SPLAT]] : !cir.matrix<9 x 3 x !s32i> + +// LLVM-LABEL: define {{.*}}void @add_matrix_scalar_int_long( +// LLVM: %[[SCALAR:.*]] = load i64, ptr {{.*}}, align 8 +// LLVM: %[[TRUNC:.*]] = trunc i64 %[[SCALAR]] to i32 +// LLVM: %[[EMBED:.*]] = insertelement <27 x i32> poison, i32 %[[TRUNC]], i64 0 +// LLVM: %[[SPLAT:.*]] = shufflevector <27 x i32> %[[EMBED]], <27 x i32> poison, <27 x i32> zeroinitializer +// LLVM: add <27 x i32> %{{.*}}, %[[SPLAT]] + +void add_compound_matrix_scalar_float_float() { + fx2x3_t b; + float vf; + b += vf; +} + +// CIR-LABEL: cir.func {{.*}}@add_compound_matrix_scalar_float_float +// CIR: %[[SPLAT:.*]] = cir.vec.splat %{{.*}} : !cir.float, !cir.matrix<2 x 3 x !cir.float> +// CIR: cir.fadd %{{.*}}, %[[SPLAT]] : !cir.matrix<2 x 3 x !cir.float> + +// LLVM-LABEL: define {{.*}}void @add_compound_matrix_scalar_float_float( +// LLVM: %[[EMBED:.*]] = insertelement <6 x float> poison, float %{{.*}}, i64 0 +// LLVM: %[[SPLAT:.*]] = shufflevector <6 x float> %[[EMBED]], <6 x float> poison, <6 x i32> zeroinitializer +// LLVM: fadd <6 x float> %{{.*}}, %[[SPLAT]] diff --git a/clang/test/CIR/IR/invalid-matrix.cir b/clang/test/CIR/IR/invalid-matrix.cir index a68bdecd26012..ce48317aa87ef 100644 --- a/clang/test/CIR/IR/invalid-matrix.cir +++ b/clang/test/CIR/IR/invalid-matrix.cir @@ -67,3 +67,55 @@ cir.func @builtin_tranpose_ diff erent_sizes() { %3 = cir.matrix.transpose %2 : <3 x 2 x !cir.float>, !cir.matrix<3 x 3 x !cir.float> cir.return } + +// ----- + +!s32i = !cir.int<s, 32> + +cir.func @add_of_float_matrix(%arg0: !cir.matrix<2 x 2 x !cir.float>) { + // expected-error@+1 {{'cir.add' op operand #0 must be integer or vector/matrix of integer type}} + %0 = cir.add %arg0, %arg0 : !cir.matrix<2 x 2 x !cir.float> + cir.return +} + +// ----- + +!s32i = !cir.int<s, 32> + +cir.func @fadd_of_int_matrix(%arg0: !cir.matrix<2 x 2 x !s32i>) { + // expected-error@+1 {{'cir.fadd' op operand #0 must be floating point or vector/matrix of floating point type}} + %0 = cir.fadd %arg0, %arg0 : !cir.matrix<2 x 2 x !s32i> + cir.return +} + +// ----- + +!s32i = !cir.int<s, 32> + +// expected-note@+1 {{prior use here}} +cir.func @add_mismatched_matrix_types(%arg0: !cir.matrix<2 x 2 x !s32i>, %arg1: !cir.matrix<2 x 3 x !s32i>) { + // expected-error@+1 {{use of value '%arg1' expects diff erent type than prior uses}} + %0 = cir.add %arg0, %arg1 : !cir.matrix<2 x 2 x !s32i> + cir.return +} + +// ----- + +!s32i = !cir.int<s, 32> + +// Only element-wise ops accept matrices; mul on a matrix is not element-wise. +cir.func @mul_of_matrix(%arg0: !cir.matrix<2 x 2 x !s32i>) { + // expected-error@+1 {{'cir.mul' op operand #0 must be integer or vector of integer type}} + %0 = cir.mul %arg0, %arg0 : !cir.matrix<2 x 2 x !s32i> + cir.return +} + +// ----- + +!s32i = !cir.int<s, 32> + +cir.func @splat_element_type_mismatch(%arg0: !s32i) { + // expected-error@+1 {{type of 'value' matches element type of 'result'}} + %0 = cir.vec.splat %arg0 : !s32i, !cir.matrix<2 x 2 x !cir.float> + cir.return +} diff --git a/clang/test/CIR/IR/matrix.cir b/clang/test/CIR/IR/matrix.cir index e0262694c533a..ea4f239a13146 100644 --- a/clang/test/CIR/IR/matrix.cir +++ b/clang/test/CIR/IR/matrix.cir @@ -17,4 +17,32 @@ cir.func @valid_matrix_type() { // CHECK: cir.return // CHECK: } +cir.func @matrix_add(%arg0: !cir.matrix<2 x 3 x !s32i>, %arg1: !cir.matrix<2 x 3 x !s32i>, + %arg2: !cir.matrix<3 x 3 x !cir.float>, %arg3: !cir.matrix<3 x 3 x !cir.float>) { + %0 = cir.add %arg0, %arg1 : !cir.matrix<2 x 3 x !s32i> + %1 = cir.add nsw %arg0, %arg1 : !cir.matrix<2 x 3 x !s32i> + %2 = cir.fadd %arg2, %arg3 : !cir.matrix<3 x 3 x !cir.float> + cir.return +} + +// CHECK: cir.func @matrix_add(%[[A:.*]]: !cir.matrix<2 x 3 x !s32i>, %[[B:.*]]: !cir.matrix<2 x 3 x !s32i>, +// CHECK-SAME: %[[C:.*]]: !cir.matrix<3 x 3 x !cir.float>, %[[D:.*]]: !cir.matrix<3 x 3 x !cir.float>) { +// CHECK: %{{.*}} = cir.add %[[A]], %[[B]] : !cir.matrix<2 x 3 x !s32i> +// CHECK: %{{.*}} = cir.add nsw %[[A]], %[[B]] : !cir.matrix<2 x 3 x !s32i> +// CHECK: %{{.*}} = cir.fadd %[[C]], %[[D]] : !cir.matrix<3 x 3 x !cir.float> +// CHECK: cir.return +// CHECK: } + +cir.func @matrix_splat(%arg0: !s32i, %arg1: !cir.float) { + %0 = cir.vec.splat %arg0 : !s32i, !cir.matrix<2 x 3 x !s32i> + %1 = cir.vec.splat %arg1 : !cir.float, !cir.matrix<3 x 3 x !cir.float> + cir.return +} + +// CHECK: cir.func @matrix_splat(%[[I:.*]]: !s32i, %[[F:.*]]: !cir.float) { +// CHECK: %{{.*}} = cir.vec.splat %[[I]] : !s32i, !cir.matrix<2 x 3 x !s32i> +// CHECK: %{{.*}} = cir.vec.splat %[[F]] : !cir.float, !cir.matrix<3 x 3 x !cir.float> +// CHECK: cir.return +// CHECK: } + } diff --git a/clang/test/CIR/Lowering/matrix-add.cir b/clang/test/CIR/Lowering/matrix-add.cir new file mode 100644 index 0000000000000..fe2bfb99ab482 --- /dev/null +++ b/clang/test/CIR/Lowering/matrix-add.cir @@ -0,0 +1,70 @@ +// RUN: cir-opt %s -cir-to-llvm -o - | FileCheck %s -check-prefix=MLIR +// RUN: cir-translate %s -cir-to-llvmir --target x86_64-unknown-linux-gnu --disable-cc-lowering | FileCheck %s -check-prefix=LLVM + +!s32i = !cir.int<s, 32> + +module { + cir.func @int_matrix_add(%arg0: !cir.matrix<2 x 3 x !s32i>, + %arg1: !cir.matrix<2 x 3 x !s32i>) -> !cir.matrix<2 x 3 x !s32i> { + %0 = cir.add %arg0, %arg1 : !cir.matrix<2 x 3 x !s32i> + cir.return %0 : !cir.matrix<2 x 3 x !s32i> + } + + // MLIR-LABEL: llvm.func @int_matrix_add + // MLIR: llvm.add %{{.*}}, %{{.*}} : vector<6xi32> + // LLVM-LABEL: define {{.*}}@int_matrix_add( + // LLVM: add <6 x i32> %{{.*}}, %{{.*}} + + cir.func @int_matrix_add_nsw(%arg0: !cir.matrix<2 x 3 x !s32i>, + %arg1: !cir.matrix<2 x 3 x !s32i>) -> !cir.matrix<2 x 3 x !s32i> { + %0 = cir.add nsw %arg0, %arg1 : !cir.matrix<2 x 3 x !s32i> + cir.return %0 : !cir.matrix<2 x 3 x !s32i> + } + + // MLIR-LABEL: llvm.func @int_matrix_add_nsw + // MLIR: llvm.add %{{.*}}, %{{.*}} overflow<nsw> : vector<6xi32> + // LLVM-LABEL: define {{.*}}@int_matrix_add_nsw( + // LLVM: add nsw <6 x i32> %{{.*}}, %{{.*}} + + cir.func @fp_matrix_add(%arg0: !cir.matrix<3 x 3 x !cir.float>, + %arg1: !cir.matrix<3 x 3 x !cir.float>) -> !cir.matrix<3 x 3 x !cir.float> { + %0 = cir.fadd %arg0, %arg1 : !cir.matrix<3 x 3 x !cir.float> + cir.return %0 : !cir.matrix<3 x 3 x !cir.float> + } + + // MLIR-LABEL: llvm.func @fp_matrix_add + // MLIR: llvm.fadd %{{.*}}, %{{.*}} : vector<9xf32> + // LLVM-LABEL: define {{.*}}@fp_matrix_add( + // LLVM: fadd <9 x float> %{{.*}}, %{{.*}} + + cir.func @matrix_splat_add(%arg0: !cir.matrix<3 x 3 x !cir.float>, + %arg1: !cir.float) -> !cir.matrix<3 x 3 x !cir.float> { + %0 = cir.vec.splat %arg1 : !cir.float, !cir.matrix<3 x 3 x !cir.float> + %1 = cir.fadd %arg0, %0 : !cir.matrix<3 x 3 x !cir.float> + cir.return %1 : !cir.matrix<3 x 3 x !cir.float> + } + + // MLIR-LABEL: llvm.func @matrix_splat_add + // MLIR: %[[POISON:.*]] = llvm.mlir.poison : vector<9xf32> + // MLIR: %[[IDX:.*]] = llvm.mlir.constant(0 : i64) : i64 + // MLIR: %[[INS:.*]] = llvm.insertelement %{{.*}}, %[[POISON]][%[[IDX]] : i64] : vector<9xf32> + // MLIR: %[[SPLAT:.*]] = llvm.shufflevector %[[INS]], %[[POISON]] [0, 0, 0, 0, 0, 0, 0, 0, 0] : vector<9xf32> + // MLIR: llvm.fadd %{{.*}}, %[[SPLAT]] : vector<9xf32> + // LLVM-LABEL: define {{.*}}@matrix_splat_add( + // LLVM: %[[INS:.*]] = insertelement <9 x float> poison, float %{{.*}}, i64 0 + // LLVM: %[[SPLAT:.*]] = shufflevector <9 x float> %[[INS]], <9 x float> poison, <9 x i32> zeroinitializer + // LLVM: fadd <9 x float> %{{.*}}, %[[SPLAT]] + + // A splat of a constant lowers to a dense constant vector. + cir.func @matrix_const_splat() -> !cir.matrix<2 x 3 x !s32i> { + %0 = cir.const #cir.int<7> : !s32i + %1 = cir.vec.splat %0 : !s32i, !cir.matrix<2 x 3 x !s32i> + cir.return %1 : !cir.matrix<2 x 3 x !s32i> + } + + // MLIR-LABEL: llvm.func @matrix_const_splat + // MLIR: %[[C:.*]] = llvm.mlir.constant(dense<7> : vector<6xi32>) : vector<6xi32> + // MLIR: llvm.return %[[C]] + // LLVM-LABEL: define {{.*}}@matrix_const_splat( + // LLVM: ret <6 x i32> splat (i32 7) +} _______________________________________________ cfe-commits mailing list [email protected] https://lists.llvm.org/cgi-bin/mailman/listinfo/cfe-commits
