[clang] 139c878 - [CIR] Add Matrix column major load op (#228229)

via cfe-commits cfe-commits at lists.llvm.org
Sat Oct 3 11:19:05 PDT 2026


Author: Amr Hesham
Date: 2026-10-03T18:18:53Z
New Revision: 139c8783301b458de3fe1270c9454d508ff13ae7

URL: https://github.com/llvm/llvm-project/commit/139c8783301b458de3fe1270c9454d508ff13ae7
DIFF: https://github.com/llvm/llvm-project/commit/139c8783301b458de3fe1270c9454d508ff13ae7.diff

LOG: [CIR] Add Matrix column major load op (#228229)

Add support for the matrix column-major load operation

Issue #221772

Added: 
    

Modified: 
    clang/include/clang/CIR/Dialect/IR/CIROps.td
    clang/lib/CIR/CodeGen/CIRGenBuilder.h
    clang/lib/CIR/CodeGen/CIRGenBuiltin.cpp
    clang/lib/CIR/Dialect/IR/CIRMemorySlot.cpp
    clang/lib/CIR/Lowering/DirectToLLVM/LowerToLLVM.cpp
    clang/test/CIR/CodeGen/matrix.cpp

Removed: 
    


################################################################################
diff  --git a/clang/include/clang/CIR/Dialect/IR/CIROps.td b/clang/include/clang/CIR/Dialect/IR/CIROps.td
index e450423c12ce2..3f02f619ced66 100644
--- a/clang/include/clang/CIR/Dialect/IR/CIROps.td
+++ b/clang/include/clang/CIR/Dialect/IR/CIROps.td
@@ -6375,6 +6375,54 @@ def CIR_VecSplatOp : CIR_Op<"vec.splat", [
   }];
 }
 
+//===----------------------------------------------------------------------===//
+// MatrixColumnMajorLoadOp
+//===----------------------------------------------------------------------===//
+
+def CIR_MatrixColumnMajorLoadOp : CIR_Op<"matrix.column_major_load", [
+  DeclareOpInterfaceMethods<PromotableMemOpInterface>,
+]> {
+  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 a load of any matrix type with a stride to compute
+  the start address of the 
diff erent 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>
+
+  %result = cir.matrix.column_major_load %ptr : <!cir.double>, %stride : !u64i
+     volatile, !cir.matrix<5 x 5 x !cir.double>
+  ```
+  }];
+
+  let arguments = (ins
+    Arg<CIR_PointerType, "the address to load from", [MemRead]>:$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 d224feb83b03d..c754c2442f098 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 9203cc0b9f722..c47165ec153a7 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/Dialect/IR/CIRMemorySlot.cpp b/clang/lib/CIR/Dialect/IR/CIRMemorySlot.cpp
index d6de6b6e80799..8a42f9caaa293 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 MatrixColumnMajorLoadOp
+//===----------------------------------------------------------------------===//
+
+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
 //===----------------------------------------------------------------------===//

diff  --git a/clang/lib/CIR/Lowering/DirectToLLVM/LowerToLLVM.cpp b/clang/lib/CIR/Lowering/DirectToLLVM/LowerToLLVM.cpp
index 8c6679c1c5c51..1871da4ec0607 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 9b8634f2e8247..98f0b99f7dd05 100644
--- a/clang/test/CIR/CodeGen/matrix.cpp
+++ b/clang/test/CIR/CodeGen/matrix.cpp
@@ -84,3 +84,39 @@ void builtin_matrix_transpose_
diff erent_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


        


More information about the cfe-commits mailing list