[clang] CIR] Add Matrix column major store op (PR #228917)

via cfe-commits cfe-commits at lists.llvm.org
Sun Oct 4 09:48:15 PDT 2026


llvmorg-github-actions[bot] wrote:


<!--LLVM PR SUMMARY COMMENT-->

@llvm/pr-subscribers-clang

Author: Amr Hesham (AmrDeveloper)

<details>
<summary>Changes</summary>

Add support for the matrix column-major store operation

Issue https://github.com/llvm/llvm-project/issues/221772

---
Full diff: https://github.com/llvm/llvm-project/pull/228917.diff


6 Files Affected:

- (modified) clang/include/clang/CIR/Dialect/IR/CIROps.td (+42) 
- (modified) clang/lib/CIR/CodeGen/CIRGenBuilder.h (+9) 
- (modified) clang/lib/CIR/CodeGen/CIRGenBuiltin.cpp (+16-1) 
- (modified) clang/lib/CIR/Dialect/IR/CIRMemorySlot.cpp (+42) 
- (modified) clang/lib/CIR/Lowering/DirectToLLVM/LowerToLLVM.cpp (+12) 
- (modified) clang/test/CIR/CodeGen/matrix.cpp (+38) 


``````````diff
diff --git a/clang/include/clang/CIR/Dialect/IR/CIROps.td b/clang/include/clang/CIR/Dialect/IR/CIROps.td
index 3f02f619ced66..75d3622510a59 100644
--- a/clang/include/clang/CIR/Dialect/IR/CIROps.td
+++ b/clang/include/clang/CIR/Dialect/IR/CIROps.td
@@ -6423,6 +6423,48 @@ def CIR_MatrixColumnMajorLoadOp : CIR_Op<"matrix.column_major_load", [
   }];
 }
 
+//===----------------------------------------------------------------------===//
+// MatrixColumnMajorStoreOp
+//===----------------------------------------------------------------------===//
+
+def CIR_MatrixColumnMajorStoreOp : CIR_Op<"matrix.column_major_store", [
+  DeclareOpInterfaceMethods<PromotableMemOpInterface>,
+]> {
+  let summary = "Matrix column major store";
+  let description = [{
+    The `cir.matrix.column_major_store` operation provides a representation for
+    the `__builtin_matrix_column_major_store` builtin and corresponds to the
+    `llvm.matrix.column.major.store` intrinsic in LLVM IR.
+
+    This operation performs a store of data from 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 stored data will be `[1, 3, 5, 2, 4, 6]`.
+
+    ```
+    cir.matrix.column_major_store %matrix : <3 x 3 x !cir.float>, 
+      %data : <!cir.float>, %stribe : !u64i
+    ```
+  }];
+
+  let arguments = (ins
+    CIR_MatrixType:$matrix,
+    Arg<CIR_PointerType, "the address to store the value", [MemWrite]>:$value,
+    CIR_IntType:$stride,
+    UnitAttr:$is_volatile
+  );
+
+  let assemblyFormat = [{
+    $matrix `:` type($matrix) `,`
+    $value `:` type($value) `,`
+    $stride `:` type($stride)
+    (`volatile` $is_volatile^)?
+    attr-dict
+  }];
+}
+
 //===----------------------------------------------------------------------===//
 // MatrixTransposeOp
 //===----------------------------------------------------------------------===//
diff --git a/clang/lib/CIR/CodeGen/CIRGenBuilder.h b/clang/lib/CIR/CodeGen/CIRGenBuilder.h
index 1d8dea9dc7f1d..b3b3805e8bb2b 100644
--- a/clang/lib/CIR/CodeGen/CIRGenBuilder.h
+++ b/clang/lib/CIR/CodeGen/CIRGenBuilder.h
@@ -843,6 +843,15 @@ class CIRGenBuilderTy : public cir::CIRBaseBuilderTy {
     return cir::MatrixTransposeOp::create(*this, loc, resultTy, matrix);
   }
 
+  cir::MatrixColumnMajorStoreOp createMatrixColumnMajorStore(mlir::Location loc,
+                                                             mlir::Value matrix,
+                                                             mlir::Value data,
+                                                             mlir::Value stride,
+                                                             bool isVolatile) {
+    return cir::MatrixColumnMajorStoreOp::create(*this, loc, matrix, data,
+                                                 stride, isVolatile);
+  }
+
   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/CIRGenBuiltin.cpp b/clang/lib/CIR/CodeGen/CIRGenBuiltin.cpp
index c47165ec153a7..97ce3fbcabcbc 100644
--- a/clang/lib/CIR/CodeGen/CIRGenBuiltin.cpp
+++ b/clang/lib/CIR/CodeGen/CIRGenBuiltin.cpp
@@ -2308,7 +2308,22 @@ RValue CIRGenFunction::emitBuiltinExpr(const GlobalDecl &gd, unsigned builtinID,
         loc, resultType, dataPtr, stride, isVolatile);
     return RValue::get(result);
   }
-  case Builtin::BI__builtin_matrix_column_major_store:
+  case Builtin::BI__builtin_matrix_column_major_store: {
+    mlir::Value matrix = emitScalarExpr(e->getArg(0));
+    Address dst = emitPointerWithAlignment(e->getArg(1));
+    mlir::Value stride = emitScalarExpr(e->getArg(2));
+
+    auto *ptrTy = e->getArg(1)->getType()->getAs<PointerType>();
+    assert(ptrTy && "arg1 must be of pointer type");
+    bool isVolatile = ptrTy->getPointeeType().isVolatileQualified();
+
+    emitNonNullArgCheck(RValue::get(dst.emitRawPointer()),
+                        e->getArg(1)->getType(), e->getArg(1)->getExprLoc(), fd,
+                        0);
+    builder.createMatrixColumnMajorStore(loc, matrix, dst.emitRawPointer(),
+                                         stride, isVolatile);
+    return RValue::get(nullptr);
+  }
   case Builtin::BI__builtin_masked_load:
   case Builtin::BI__builtin_masked_expand_load:
   case Builtin::BI__builtin_masked_gather:
diff --git a/clang/lib/CIR/Dialect/IR/CIRMemorySlot.cpp b/clang/lib/CIR/Dialect/IR/CIRMemorySlot.cpp
index 8a42f9caaa293..2ff8b9bd63603 100644
--- a/clang/lib/CIR/Dialect/IR/CIRMemorySlot.cpp
+++ b/clang/lib/CIR/Dialect/IR/CIRMemorySlot.cpp
@@ -213,6 +213,48 @@ DeletionKind cir::MatrixColumnMajorLoadOp::removeBlockingUses(
   return DeletionKind::Delete;
 }
 
+//===----------------------------------------------------------------------===//
+// Interfaces for MatrixColumnMajorStoreOp
+//===----------------------------------------------------------------------===//
+
+bool cir::MatrixColumnMajorStoreOp::loadsFrom(const MemorySlot &slot) {
+  return false;
+}
+
+bool cir::MatrixColumnMajorStoreOp::storesTo(const MemorySlot &slot) {
+  return getValue() == slot.ptr;
+}
+
+Value cir::MatrixColumnMajorStoreOp::getStored(const MemorySlot &slot,
+                                               OpBuilder &builder,
+                                               Value reachingDef,
+                                               const DataLayout &dataLayout) {
+  return getMatrix();
+}
+
+bool cir::MatrixColumnMajorStoreOp::canUsesBeRemoved(
+    const MemorySlot &slot, const SmallPtrSetImpl<OpOperand *> &blockingUses,
+    SmallVectorImpl<OpOperand *> &newBlockingUses,
+    const DataLayout &dataLayout) {
+  if (blockingUses.size() != 1)
+    return false;
+
+  // Volatile store should not be removed.
+  if (getIsVolatile())
+    return false;
+
+  Value blockingUse = (*blockingUses.begin())->get();
+  return blockingUse == slot.ptr && getValue() == slot.ptr &&
+         getValue() != slot.ptr && slot.elemType == getValue().getType();
+}
+
+DeletionKind cir::MatrixColumnMajorStoreOp::removeBlockingUses(
+    const MemorySlot &slot, const SmallPtrSetImpl<OpOperand *> &blockingUses,
+    OpBuilder &builder, Value reachingDefinition,
+    const DataLayout &dataLayout) {
+  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 1871da4ec0607..6426841e6237f 100644
--- a/clang/lib/CIR/Lowering/DirectToLLVM/LowerToLLVM.cpp
+++ b/clang/lib/CIR/Lowering/DirectToLLVM/LowerToLLVM.cpp
@@ -5273,6 +5273,18 @@ mlir::LogicalResult CIRToLLVMMatrixColumnMajorLoadOpLowering::matchAndRewrite(
   return mlir::success();
 }
 
+mlir::LogicalResult CIRToLLVMMatrixColumnMajorStoreOpLowering::matchAndRewrite(
+    cir::MatrixColumnMajorStoreOp op, OpAdaptor adaptor,
+    mlir::ConversionPatternRewriter &rewriter) const {
+  cir::MatrixType matrixTy = op.getMatrix().getType();
+  rewriter.replaceOpWithNewOp<mlir::LLVM::MatrixColumnMajorStoreOp>(
+      op, adaptor.getMatrix(), adaptor.getValue(), adaptor.getStride(),
+      rewriter.getBoolAttr(op.getIsVolatile()),
+      rewriter.getI32IntegerAttr(matrixTy.getRowNum()),
+      rewriter.getI32IntegerAttr(matrixTy.getColumnNum()));
+  return mlir::success();
+}
+
 mlir::LogicalResult CIRToLLVMMatrixTransposeOpLowering::matchAndRewrite(
     cir::MatrixTransposeOp op, OpAdaptor adaptor,
     mlir::ConversionPatternRewriter &rewriter) const {
diff --git a/clang/test/CIR/CodeGen/matrix.cpp b/clang/test/CIR/CodeGen/matrix.cpp
index 98f0b99f7dd05..b1aa26455a25e 100644
--- a/clang/test/CIR/CodeGen/matrix.cpp
+++ b/clang/test/CIR/CodeGen/matrix.cpp
@@ -120,3 +120,41 @@ void column_major_volatile_load() {
 // 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
+
+void column_major_store() {
+  matrix3x3 matrix;
+  float *ptr;
+  __builtin_matrix_column_major_store(matrix, ptr, 3);
+}
+
+// CIR: %[[MATRIX_ADDR:.*]] = cir.alloca "matrix" {{.*}} : !cir.ptr<!cir.matrix<3 x 3 x !cir.float>>
+// CIR: %[[PTR_ADDR:.*]] = cir.alloca "ptr" {{.*}} : !cir.ptr<!cir.ptr<!cir.float>>
+// CIR: %[[TMP_MATRIX:.*]] = cir.load {{.*}} %[[MATRIX_ADDR]] : !cir.ptr<!cir.matrix<3 x 3 x !cir.float>>, !cir.matrix<3 x 3 x !cir.float>
+// CIR: %[[TMP_PTR:.*]] = cir.load {{.*}} %[[PTR_ADDR]] : !cir.ptr<!cir.ptr<!cir.float>>, !cir.ptr<!cir.float>
+// CIR: %[[STRIDE:.*]] = cir.const #cir.int<3> : !u64i
+// CIR: cir.matrix.column_major_store %[[TMP_MATRIX]] : <3 x 3 x !cir.float>, %[[TMP_PTR]] : <!cir.float>, %[[STRIDE]] : !u64i
+
+// LLVM: %[[MATRIX_ADDR:.*]] = alloca [9 x float], align 4
+// LLVM: %[[PTR_ADDR:.*]] = alloca ptr, align 8
+// LLVM: %[[TMP_MATRIX:.*]] = load <9 x float>, ptr %[[MATRIX_ADDR]], align 4
+// LLVM: %[[TMP_PTR:.*]] = load ptr, ptr %[[PTR_ADDR]], align 8
+// LLVM: call void @llvm.matrix.column.major.store.v9f32.i64(<9 x float> %[[TMP_MATRIX]], ptr align 4 %[[TMP_PTR]], i64 3, i1 false, i32 3, i32 3)
+
+void column_major_volatile_store() {
+  matrix3x3 matrix;
+  volatile float *ptr;
+  __builtin_matrix_column_major_store(matrix, ptr, 3);
+}
+
+// CIR: %[[MATRIX_ADDR:.*]] = cir.alloca "matrix" {{.*}} : !cir.ptr<!cir.matrix<3 x 3 x !cir.float>>
+// CIR: %[[PTR_ADDR:.*]] = cir.alloca "ptr" {{.*}} : !cir.ptr<!cir.ptr<!cir.float>>
+// CIR: %[[TMP_MATRIX:.*]] = cir.load {{.*}} %[[MATRIX_ADDR]] : !cir.ptr<!cir.matrix<3 x 3 x !cir.float>>, !cir.matrix<3 x 3 x !cir.float>
+// CIR: %[[TMP_PTR:.*]] = cir.load {{.*}} %[[PTR_ADDR]] : !cir.ptr<!cir.ptr<!cir.float>>, !cir.ptr<!cir.float>
+// CIR: %[[STRIDE:.*]] = cir.const #cir.int<3> : !u64i
+// CIR: cir.matrix.column_major_store %[[TMP_MATRIX]] : <3 x 3 x !cir.float>, %[[TMP_PTR]] : <!cir.float>, %[[STRIDE]] : !u64i volatile
+
+// LLVM: %[[MATRIX_ADDR:.*]] = alloca [9 x float], align 4
+// LLVM: %[[PTR_ADDR:.*]] = alloca ptr, align 8
+// LLVM: %[[TMP_MATRIX:.*]] = load <9 x float>, ptr %[[MATRIX_ADDR]], align 4
+// LLVM: %[[TMP_PTR:.*]] = load ptr, ptr %[[PTR_ADDR]], align 8
+// LLVM: call void @llvm.matrix.column.major.store.v9f32.i64(<9 x float> %[[TMP_MATRIX]], ptr align 4 %[[TMP_PTR]], i64 3, i1 true, i32 3, i32 3)

``````````

</details>


https://github.com/llvm/llvm-project/pull/228917


More information about the cfe-commits mailing list