[clang] [CIR] Add Matrix column major load op (PR #228229)
Amr Hesham via cfe-commits
cfe-commits at lists.llvm.org
Sat Oct 3 09:43:57 PDT 2026
https://github.com/AmrDeveloper updated https://github.com/llvm/llvm-project/pull/228229
>From 3d7dc9ea1ef17a806cddb3258e30d06f743f1bf9 Mon Sep 17 00:00:00 2001
From: Amr Hesham <amr96 at programmer.net>
Date: Thu, 1 Oct 2026 20:49:41 +0200
Subject: [PATCH 1/4] [CIR] Add Matrix column major load op
---
clang/include/clang/CIR/Dialect/IR/CIROps.td | 41 +++++++++++++++++++
clang/lib/CIR/CodeGen/CIRGenBuilder.h | 9 ++++
clang/lib/CIR/CodeGen/CIRGenBuiltin.cpp | 18 +++++++-
.../CIR/Lowering/DirectToLLVM/LowerToLLVM.cpp | 22 +++++++---
clang/test/CIR/CodeGen/matrix.cpp | 36 ++++++++++++++++
5 files changed, 120 insertions(+), 6 deletions(-)
diff --git a/clang/include/clang/CIR/Dialect/IR/CIROps.td b/clang/include/clang/CIR/Dialect/IR/CIROps.td
index e450423c12ce24..d8d6305ee262e8 100644
--- a/clang/include/clang/CIR/Dialect/IR/CIROps.td
+++ b/clang/include/clang/CIR/Dialect/IR/CIROps.td
@@ -6375,6 +6375,47 @@ def CIR_VecSplatOp : CIR_Op<"vec.splat", [
}];
}
+//===----------------------------------------------------------------------===//
+// MatrixColumnMajorLoadOp
+//===----------------------------------------------------------------------===//
+
+def CIR_MatrixColumnMajorLoadOp : CIR_Op<"matrix.column_major_load", [
+ Pure,
+]> {
+ let summary = "Matrix column major load";
+ let description = [{
+ The `cir.matrix.column_major_load` operation provides a representation for
+ the `__builtin_matrix_column_major_load` builtin and corresponds to the
+ `llvm.matrix.column.major.load` intrinsic in LLVM IR.
+
+ This operation performs load any matrix type with a stride to compute
+ the start address of the different columns.
+
+ ```
+ %result = cir.matrix.column_major_load %ptr : <!cir.double>, %stride : !u64i,
+ !cir.matrix<5 x 5 x !cir.double>
+
+ %result = cir.matrix.column_major_load %ptr : <!cir.double>, %stride : !u64i
+ volatile, !cir.matrix<5 x 5 x !cir.double>
+ ```
+ }];
+
+ let arguments = (ins
+ CIR_PointerType:$value,
+ CIR_IntType:$stride,
+ UnitAttr:$is_volatile
+ );
+
+ let results = (outs CIR_MatrixType:$result);
+
+ let assemblyFormat = [{
+ $value `:` type($value) `,`
+ $stride `:` type($stride)
+ (`volatile` $is_volatile^)?
+ `,` qualified(type($result)) attr-dict
+ }];
+}
+
//===----------------------------------------------------------------------===//
// MatrixTransposeOp
//===----------------------------------------------------------------------===//
diff --git a/clang/lib/CIR/CodeGen/CIRGenBuilder.h b/clang/lib/CIR/CodeGen/CIRGenBuilder.h
index d224feb83b03df..c754c2442f098c 100644
--- a/clang/lib/CIR/CodeGen/CIRGenBuilder.h
+++ b/clang/lib/CIR/CodeGen/CIRGenBuilder.h
@@ -825,6 +825,15 @@ class CIRGenBuilderTy : public cir::CIRBaseBuilderTy {
return createVecShuffle(loc, vec1, poison, mask);
}
+ cir::MatrixColumnMajorLoadOp createMatrixColumnMajorLoad(mlir::Location loc,
+ mlir::Type resultTy,
+ mlir::Value value,
+ mlir::Value stride,
+ bool isVolatile) {
+ return cir::MatrixColumnMajorLoadOp::create(*this, loc, resultTy, value,
+ stride, isVolatile);
+ }
+
cir::MatrixTransposeOp createMatrixTranspose(mlir::Location loc,
mlir::Value matrix) {
auto inputTy = mlir::cast<cir::MatrixType>(matrix.getType());
diff --git a/clang/lib/CIR/CodeGen/CIRGenBuiltin.cpp b/clang/lib/CIR/CodeGen/CIRGenBuiltin.cpp
index 9203cc0b9f7220..c47165ec153a7e 100644
--- a/clang/lib/CIR/CodeGen/CIRGenBuiltin.cpp
+++ b/clang/lib/CIR/CodeGen/CIRGenBuiltin.cpp
@@ -2291,7 +2291,23 @@ RValue CIRGenFunction::emitBuiltinExpr(const GlobalDecl &gd, unsigned builtinID,
mlir::Value result = builder.createMatrixTranspose(loc, matrix);
return RValue::get(result);
}
- case Builtin::BI__builtin_matrix_column_major_load:
+ case Builtin::BI__builtin_matrix_column_major_load: {
+ // Emit everything that isn't dependent on the first parameter type
+ mlir::Value stride = emitScalarExpr(e->getArg(3));
+ const QualType resultTy = e->getType();
+ mlir::Type resultType = convertType(resultTy);
+ auto *ptrTy = e->getArg(0)->getType()->getAs<PointerType>();
+ assert(ptrTy && "arg0 must be of pointer type");
+ bool isVolatile = ptrTy->getPointeeType().isVolatileQualified();
+ Address src = emitPointerWithAlignment(e->getArg(0));
+ emitNonNullArgCheck(RValue::get(src.emitRawPointer()),
+ e->getArg(0)->getType(), e->getArg(0)->getExprLoc(), fd,
+ 0);
+ mlir::Value dataPtr = src.emitRawPointer();
+ mlir::Value result = builder.createMatrixColumnMajorLoad(
+ loc, resultType, dataPtr, stride, isVolatile);
+ return RValue::get(result);
+ }
case Builtin::BI__builtin_matrix_column_major_store:
case Builtin::BI__builtin_masked_load:
case Builtin::BI__builtin_masked_expand_load:
diff --git a/clang/lib/CIR/Lowering/DirectToLLVM/LowerToLLVM.cpp b/clang/lib/CIR/Lowering/DirectToLLVM/LowerToLLVM.cpp
index 8c6679c1c5c512..1871da4ec06073 100644
--- a/clang/lib/CIR/Lowering/DirectToLLVM/LowerToLLVM.cpp
+++ b/clang/lib/CIR/Lowering/DirectToLLVM/LowerToLLVM.cpp
@@ -5260,15 +5260,27 @@ mlir::LogicalResult CIRToLLVMVecTernaryOpLowering::matchAndRewrite(
return mlir::success();
}
+mlir::LogicalResult CIRToLLVMMatrixColumnMajorLoadOpLowering::matchAndRewrite(
+ cir::MatrixColumnMajorLoadOp op, OpAdaptor adaptor,
+ mlir::ConversionPatternRewriter &rewriter) const {
+ cir::MatrixType resultMatrixTy = op.getResult().getType();
+ mlir::Type resultTy = typeConverter->convertType(resultMatrixTy);
+ rewriter.replaceOpWithNewOp<mlir::LLVM::MatrixColumnMajorLoadOp>(
+ op, resultTy, adaptor.getValue(), adaptor.getStride(),
+ rewriter.getBoolAttr(op.getIsVolatile()),
+ rewriter.getI32IntegerAttr(resultMatrixTy.getRowNum()),
+ rewriter.getI32IntegerAttr(resultMatrixTy.getColumnNum()));
+ return mlir::success();
+}
+
mlir::LogicalResult CIRToLLVMMatrixTransposeOpLowering::matchAndRewrite(
cir::MatrixTransposeOp op, OpAdaptor adaptor,
mlir::ConversionPatternRewriter &rewriter) const {
- cir::MatrixType matrixTy = op.getValue().getType();
- mlir::Type resultTy =
- typeConverter->convertType(op->getResultTypes().front());
+ cir::MatrixType resultMatrixTy = op.getValue().getType();
+ mlir::Type resultTy = typeConverter->convertType(resultMatrixTy);
rewriter.replaceOpWithNewOp<mlir::LLVM::MatrixTransposeOp>(
- +op, resultTy, adaptor.getValue(), matrixTy.getRowNum(),
- matrixTy.getColumnNum());
+ op, resultTy, adaptor.getValue(), resultMatrixTy.getRowNum(),
+ resultMatrixTy.getColumnNum());
return mlir::success();
}
diff --git a/clang/test/CIR/CodeGen/matrix.cpp b/clang/test/CIR/CodeGen/matrix.cpp
index 9b8634f2e82473..98f0b99f7dd052 100644
--- a/clang/test/CIR/CodeGen/matrix.cpp
+++ b/clang/test/CIR/CodeGen/matrix.cpp
@@ -84,3 +84,39 @@ void builtin_matrix_transpose_different_sizes() {
// LLVM: %[[TMP_A:.*]] = load <6 x float>, ptr %[[A_ADDR]], align 4
// LLVM: %[[TRANSPOSE:.*]] = call <6 x float> @llvm.matrix.transpose.v6f32(<6 x float> %[[TMP_A]], i32 3, i32 2)
// LLVM: store <6 x float> %[[TRANSPOSE:.*]], ptr %[[B_ADDR]], align 4
+
+void column_major_load() {
+ float *ptr;
+ matrix3x3 matrix = __builtin_matrix_column_major_load(ptr, 3, 3, 3);
+}
+
+// CIR: %[[PTR_ADDR:.*]] = cir.alloca "ptr" {{.*}} : !cir.ptr<!cir.ptr<!cir.float>>
+// CIR: %[[MATRIX_ADDR:.*]] = cir.alloca "matrix" {{.*}} init : !cir.ptr<!cir.matrix<3 x 3 x !cir.float>>
+// CIR: %[[STRIDE:.*]] = cir.const #cir.int<3> : !u64i
+// CIR: %[[TMP_PTR:.*]] = cir.load {{.*}} %[[PTR_ADDR]] : !cir.ptr<!cir.ptr<!cir.float>>, !cir.ptr<!cir.float>
+// CIR: %[[RESULT:.*]] = cir.matrix.column_major_load %[[TMP_PTR]] : <!cir.float>, %[[STRIDE]] : !u64i, !cir.matrix<3 x 3 x !cir.float>
+// CIR: cir.store {{.*}} %[[RESULT]], %[[MATRIX_ADDR]] : !cir.matrix<3 x 3 x !cir.float>, !cir.ptr<!cir.matrix<3 x 3 x !cir.float>>
+
+// LLVM: %[[PTR_ADDR:.*]] = alloca ptr, align 8
+// LLVM: %[[MATRIX_ADDR:.*]] = alloca [9 x float], align 4
+// LLVM: %[[TMP_PTR:.*]] = load ptr, ptr %[[PTR_ADDR]], align 8
+// LLVM: %[[RESULT:.*]] = call <9 x float> @llvm.matrix.column.major.load.v9f32.i64(ptr align 4 %[[TMP_PTR]], i64 3, i1 false, i32 3, i32 3)
+// LLVM: store <9 x float> %[[RESULT]], ptr %[[MATRIX_ADDR]], align 4
+
+void column_major_volatile_load() {
+ volatile float *ptr;
+ matrix3x3 matrix = __builtin_matrix_column_major_load(ptr, 3, 3, 3);
+}
+
+// CIR: %[[PTR_ADDR:.*]] = cir.alloca "ptr" {{.*}} : !cir.ptr<!cir.ptr<!cir.float>>
+// CIR: %[[MATRIX_ADDR:.*]] = cir.alloca "matrix" {{.*}} init : !cir.ptr<!cir.matrix<3 x 3 x !cir.float>>
+// CIR: %[[STRIDE:.*]] = cir.const #cir.int<3> : !u64i
+// CIR: %[[TMP_PTR:.*]] = cir.load {{.*}} %[[PTR_ADDR]] : !cir.ptr<!cir.ptr<!cir.float>>, !cir.ptr<!cir.float>
+// CIR: %[[RESULT:.*]] = cir.matrix.column_major_load %[[TMP_PTR]] : <!cir.float>, %[[STRIDE]] : !u64i volatile, !cir.matrix<3 x 3 x !cir.float>
+// CIR: cir.store {{.*}} %[[RESULT]], %[[MATRIX_ADDR]] : !cir.matrix<3 x 3 x !cir.float>, !cir.ptr<!cir.matrix<3 x 3 x !cir.float>>
+
+// LLVM: %[[PTR_ADDR:.*]] = alloca ptr, align 8
+// LLVM: %[[MATRIX_ADDR:.*]] = alloca [9 x float], align 4
+// LLVM: %[[TMP_PTR:.*]] = load ptr, ptr %[[PTR_ADDR]], align 8
+// LLVM: %[[RESULT:.*]] = call <9 x float> @llvm.matrix.column.major.load.v9f32.i64(ptr align 4 %[[TMP_PTR]], i64 3, i1 true, i32 3, i32 3)
+// LLVM: store <9 x float> %[[RESULT]], ptr %[[MATRIX_ADDR]], align 4
>From 5bd9f22930e1aa938da3e4c045401fc374c39443 Mon Sep 17 00:00:00 2001
From: Amr Hesham <amr96 at programmer.net>
Date: Fri, 2 Oct 2026 18:27:38 +0200
Subject: [PATCH 2/4] Address code review comments
---
clang/include/clang/CIR/Dialect/IR/CIROps.td | 9 ++++++++-
1 file changed, 8 insertions(+), 1 deletion(-)
diff --git a/clang/include/clang/CIR/Dialect/IR/CIROps.td b/clang/include/clang/CIR/Dialect/IR/CIROps.td
index d8d6305ee262e8..7068b7717b11bc 100644
--- a/clang/include/clang/CIR/Dialect/IR/CIROps.td
+++ b/clang/include/clang/CIR/Dialect/IR/CIROps.td
@@ -6388,10 +6388,17 @@ def CIR_MatrixColumnMajorLoadOp : CIR_Op<"matrix.column_major_load", [
the `__builtin_matrix_column_major_load` builtin and corresponds to the
`llvm.matrix.column.major.load` intrinsic in LLVM IR.
- This operation performs load any matrix type with a stride to compute
+ This operation performs a load of any matrix type with a stride to compute
the start address of the different columns.
+ The `stride` argument is the column stride which much be greater than or equal
+ to `row`, giving the following 2x3 matrix `[[1, 2, 3], [4, 5, 6]]` with a
+ stride of 2, the result matrix will be `[[1, 3, 5], [2, 4, 6]]`.
+
```
+ %result = cir.matrix.column_major_load %ptr : <!cir.float>, %stride : !u64i,
+ !cir.matrix<2 x 3 x !cir.float>
+
%result = cir.matrix.column_major_load %ptr : <!cir.double>, %stride : !u64i,
!cir.matrix<5 x 5 x !cir.double>
>From 134d4aad71f85b6dce8ea7e45070b066b810e5cd Mon Sep 17 00:00:00 2001
From: Amr Hesham <amr96 at programmer.net>
Date: Sat, 3 Oct 2026 16:53:19 +0200
Subject: [PATCH 3/4] Address code review comments
---
clang/include/clang/CIR/Dialect/IR/CIROps.td | 4 +-
clang/lib/CIR/Dialect/IR/CIRMemorySlot.cpp | 43 ++++++++++++++++++++
2 files changed, 45 insertions(+), 2 deletions(-)
diff --git a/clang/include/clang/CIR/Dialect/IR/CIROps.td b/clang/include/clang/CIR/Dialect/IR/CIROps.td
index 7068b7717b11bc..26fd82dfdcc841 100644
--- a/clang/include/clang/CIR/Dialect/IR/CIROps.td
+++ b/clang/include/clang/CIR/Dialect/IR/CIROps.td
@@ -6380,7 +6380,7 @@ def CIR_VecSplatOp : CIR_Op<"vec.splat", [
//===----------------------------------------------------------------------===//
def CIR_MatrixColumnMajorLoadOp : CIR_Op<"matrix.column_major_load", [
- Pure,
+ DeclareOpInterfaceMethods<PromotableMemOpInterface>,
]> {
let summary = "Matrix column major load";
let description = [{
@@ -6408,7 +6408,7 @@ def CIR_MatrixColumnMajorLoadOp : CIR_Op<"matrix.column_major_load", [
}];
let arguments = (ins
- CIR_PointerType:$value,
+ Arg<CIR_PointerType, "the address to store the value", [MemRead]>:$value,
CIR_IntType:$stride,
UnitAttr:$is_volatile
);
diff --git a/clang/lib/CIR/Dialect/IR/CIRMemorySlot.cpp b/clang/lib/CIR/Dialect/IR/CIRMemorySlot.cpp
index d6de6b6e807994..414e285988a1a3 100644
--- a/clang/lib/CIR/Dialect/IR/CIRMemorySlot.cpp
+++ b/clang/lib/CIR/Dialect/IR/CIRMemorySlot.cpp
@@ -170,6 +170,49 @@ bool cir::CopyOp::canUsesBeRemoved(
dataLayout.getTypeSize(slot.elemType);
}
+//===----------------------------------------------------------------------===//
+// Interfaces for LoadOp
+//===----------------------------------------------------------------------===//
+
+bool cir::MatrixColumnMajorLoadOp::loadsFrom(const MemorySlot &slot) {
+ return getValue() == slot.ptr;
+}
+
+bool cir::MatrixColumnMajorLoadOp::storesTo(const MemorySlot &slot) {
+ return false;
+}
+
+Value cir::MatrixColumnMajorLoadOp::getStored(const MemorySlot &slot,
+ OpBuilder &builder,
+ Value reachingDef,
+ const DataLayout &dataLayout) {
+ llvm_unreachable("getStored should not be called on MatrixColumnMajorLoadOp");
+}
+
+bool cir::MatrixColumnMajorLoadOp::canUsesBeRemoved(
+ const MemorySlot &slot, const SmallPtrSetImpl<OpOperand *> &blockingUses,
+ SmallVectorImpl<OpOperand *> &newBlockingUses,
+ const DataLayout &dataLayout) {
+ if (blockingUses.size() != 1)
+ return false;
+
+ // Volatile load should not be removed.
+ if (getIsVolatile())
+ return false;
+
+ Value blockingUse = (*blockingUses.begin())->get();
+ return blockingUse == slot.ptr && getValue() == slot.ptr &&
+ getType() == slot.elemType;
+}
+
+DeletionKind cir::MatrixColumnMajorLoadOp::removeBlockingUses(
+ const MemorySlot &slot, const SmallPtrSetImpl<OpOperand *> &blockingUses,
+ OpBuilder &builder, Value reachingDefinition,
+ const DataLayout &dataLayout) {
+ getResult().replaceAllUsesWith(reachingDefinition);
+ return DeletionKind::Delete;
+}
+
//===----------------------------------------------------------------------===//
// Interfaces for CastOp
//===----------------------------------------------------------------------===//
>From 2ea8305da32a5eccd850d35c7b4d4e4082df046b Mon Sep 17 00:00:00 2001
From: Amr Hesham <amr96 at programmer.net>
Date: Sat, 3 Oct 2026 17:00:20 +0200
Subject: [PATCH 4/4] Fix parameter description
---
clang/include/clang/CIR/Dialect/IR/CIROps.td | 2 +-
1 file changed, 1 insertion(+), 1 deletion(-)
diff --git a/clang/include/clang/CIR/Dialect/IR/CIROps.td b/clang/include/clang/CIR/Dialect/IR/CIROps.td
index 26fd82dfdcc841..3f02f619ced663 100644
--- a/clang/include/clang/CIR/Dialect/IR/CIROps.td
+++ b/clang/include/clang/CIR/Dialect/IR/CIROps.td
@@ -6408,7 +6408,7 @@ def CIR_MatrixColumnMajorLoadOp : CIR_Op<"matrix.column_major_load", [
}];
let arguments = (ins
- Arg<CIR_PointerType, "the address to store the value", [MemRead]>:$value,
+ Arg<CIR_PointerType, "the address to load from", [MemRead]>:$value,
CIR_IntType:$stride,
UnitAttr:$is_volatile
);
More information about the cfe-commits
mailing list