[clang] [llvm] [HLSL][Matrix] Use canonical column-major indexing in CodeGen (PR #227087)

Farzon Lotfi via llvm-commits llvm-commits at lists.llvm.org
Mon Sep 28 11:56:15 PDT 2026


https://github.com/farzonl created https://github.com/llvm/llvm-project/pull/227087

Make MatrixBuilder::CreateIndex always compute column-major indices and remove the row-major and column-major helper functions. Update matrix subscripting, row extraction, flattening, and elementwise casts to consistently construct and access matrices in canonical column-major register order.

Keep matrix layout as a memory property by converting canonical values when storing to row-major memory. Route matrix prvalue materialization through the matrix-aware store path so element access on expressions such as transpose(m)._m01 observes the correct temporary layout.

Update CodeGen tests to cover canonical indexing for both layout defaults, explicit row-major elementwise-cast destinations, and matrix swizzles on prvalue temporaries.

Assisted by Copilot via GPT 5.6-sol

>From 5d6c58592c0ec30e6df8df753fed796cae09d2e5 Mon Sep 17 00:00:00 2001
From: Farzon Lotfi <farzonlotfi at microsoft.com>
Date: Mon, 28 Sep 2026 14:39:10 -0400
Subject: [PATCH] [HLSL][Matrix] Use canonical column-major indexing in CodeGen

Make MatrixBuilder::CreateIndex always compute column-major indices and remove
the row-major and column-major helper functions. Update matrix subscripting,
row extraction, flattening, and elementwise casts to consistently construct
and access matrices in canonical column-major register order.

Keep matrix layout as a memory property by converting canonical values when
storing to row-major memory. Route matrix prvalue materialization through the
matrix-aware store path so element access on expressions such as
transpose(m)._m01 observes the correct temporary layout.

Update CodeGen tests to cover canonical indexing for both layout defaults,
explicit row-major elementwise-cast destinations, and matrix swizzles on
prvalue temporaries.
---
 clang/lib/CodeGen/CGExpr.cpp                  | 27 +++++------------
 clang/lib/CodeGen/CGExprScalar.cpp            | 11 ++-----
 clang/test/CodeGen/matrix-type-indexing.c     | 15 ++++------
 .../MatrixElementRowColFlags.hlsl             | 29 +++++++++++++++++++
 .../BasicFeatures/VectorElementwiseCast.hlsl  |  6 ++--
 .../BasicFeatures/matrix-type-indexing.hlsl   |  9 ++----
 .../matrix-layout-attr-overrides-default.hlsl | 19 ++++++++++--
 llvm/include/llvm/IR/MatrixBuilder.h          | 26 ++---------------
 8 files changed, 70 insertions(+), 72 deletions(-)

diff --git a/clang/lib/CodeGen/CGExpr.cpp b/clang/lib/CodeGen/CGExpr.cpp
index 25e775f8328b9..a826fa0089043 100644
--- a/clang/lib/CodeGen/CGExpr.cpp
+++ b/clang/lib/CodeGen/CGExpr.cpp
@@ -2364,11 +2364,8 @@ LValue CodeGenFunction::EmitMatrixElementExpr(const MatrixElementExpr *E) {
     llvm::Value *Mat = EmitScalarExpr(E->getBase());
     Address MatMem = CreateMemTemp(E->getBase()->getType());
     QualType Ty = E->getBase()->getType();
-    llvm::Type *LTy = convertTypeForLoadStore(Ty, Mat->getType());
-    if (LTy->getScalarSizeInBits() > Mat->getType()->getScalarSizeInBits())
-      Mat = Builder.CreateZExt(Mat, LTy);
-    Builder.CreateStore(Mat, MatMem);
     Base = MakeAddrLValue(MatMem, Ty, AlignmentSource::Decl);
+    EmitStoreOfScalar(Mat, Base, /*isInit=*/true);
   }
   QualType ResultType =
       E->getType().withCVRQualifiers(Base.getQuals().getCVRQualifiers());
@@ -2430,7 +2427,8 @@ static void EmitStoreOfMatrixScalar(llvm::Value *value, LValue lvalue,
     const auto *MatrixTy = lvalue.getType()->castAs<ConstantMatrixType>();
     llvm::MatrixBuilder MB(CGF.Builder);
     value = MB.CreateColumnMajorToRowMajorTransform(
-        value, MatrixTy->getNumRows(), MatrixTy->getNumColumns());
+        value, MatrixTy->getNumRows(), MatrixTy->getNumColumns(),
+        "TEMP_ROW_MAJOR");
   }
   Address Addr = MaybeConvertMatrixAddress(lvalue.getAddress(), CGF,
                                            value->getType()->isVectorTy());
@@ -2650,9 +2648,7 @@ RValue CodeGenFunction::EmitLoadOfLValue(LValue LV, SourceLocation Loc) {
         ColIdx = ColConstsIndices->getAggregateElement(Col);
       else
         ColIdx = llvm::ConstantInt::get(Row->getType(), Col);
-      bool IsMatrixRowMajor = isMatrixRowMajor(getLangOpts(), MatTy);
-      llvm::Value *EltIndex =
-          MB.CreateIndex(Row, ColIdx, NumRows, NumCols, IsMatrixRowMajor);
+      llvm::Value *EltIndex = MB.CreateIndex(Row, ColIdx, NumRows);
       llvm::Value *Elt = Builder.CreateExtractElement(MatrixVec, EltIndex);
       llvm::Value *Lane = llvm::ConstantInt::get(Builder.getInt32Ty(), Col);
       Result = Builder.CreateInsertElement(Result, Elt, Lane);
@@ -2969,9 +2965,7 @@ void CodeGenFunction::EmitStoreThroughLValue(RValue Src, LValue Dst,
           ColIdx = ColConstsIndices->getAggregateElement(Col);
         else
           ColIdx = llvm::ConstantInt::get(Row->getType(), Col);
-        bool IsMatrixRowMajor = isMatrixRowMajor(getLangOpts(), Dst.getType());
-        llvm::Value *EltIndex =
-            MB.CreateIndex(Row, ColIdx, NumRows, NumCols, IsMatrixRowMajor);
+        llvm::Value *EltIndex = MB.CreateIndex(Row, ColIdx, NumRows);
         llvm::Value *Lane = llvm::ConstantInt::get(Builder.getInt32Ty(), Col);
         llvm::Value *Zero = llvm::ConstantInt::get(Int32Ty, 0);
         llvm::Value *NewElt = Builder.CreateExtractElement(RowVal, Lane);
@@ -5432,12 +5426,8 @@ LValue CodeGenFunction::EmitMatrixSubscriptExpr(const MatrixSubscriptExpr *E) {
   llvm::Value *ColIdx = EmitMatrixIndexExpr(E->getColumnIdx());
   llvm::MatrixBuilder MB(Builder);
   const auto *MatrixTy = E->getBase()->getType()->castAs<ConstantMatrixType>();
-  unsigned NumCols = MatrixTy->getNumColumns();
   unsigned NumRows = MatrixTy->getNumRows();
-  bool IsMatrixRowMajor =
-      isMatrixRowMajor(getLangOpts(), E->getBase()->getType());
-  llvm::Value *FinalIdx =
-      MB.CreateIndex(RowIdx, ColIdx, NumRows, NumCols, IsMatrixRowMajor);
+  llvm::Value *FinalIdx = MB.CreateIndex(RowIdx, ColIdx, NumRows);
 
   return LValue::MakeMatrixElt(
       MaybeConvertMatrixAddress(Base.getAddress(), *this), FinalIdx,
@@ -7662,15 +7652,12 @@ void CodeGenFunction::FlattenAccessAndTypeLValue(
       LValue Base = MakeAddrLValue(GEP, T);
       Address MatAddr = MaybeConvertMatrixAddress(Base.getAddress(), *this);
       unsigned NumRows = MT->getNumRows();
-      unsigned NumCols = MT->getNumColumns();
-      bool IsMatrixRowMajor = isMatrixRowMajor(getLangOpts(), T);
       llvm::MatrixBuilder MB(Builder);
       for (unsigned Row = 0; Row < MT->getNumRows(); Row++) {
         for (unsigned Col = 0; Col < MT->getNumColumns(); Col++) {
           llvm::Value *RowIdx = llvm::ConstantInt::get(IdxTy, Row);
           llvm::Value *ColIdx = llvm::ConstantInt::get(IdxTy, Col);
-          llvm::Value *Idx = MB.CreateIndex(RowIdx, ColIdx, NumRows, NumCols,
-                                            IsMatrixRowMajor);
+          llvm::Value *Idx = MB.CreateIndex(RowIdx, ColIdx, NumRows);
           LValue LV =
               LValue::MakeMatrixElt(MatAddr, Idx, MT->getElementType(),
                                     Base.getBaseInfo(), TBAAAccessInfo());
diff --git a/clang/lib/CodeGen/CGExprScalar.cpp b/clang/lib/CodeGen/CGExprScalar.cpp
index 0c310816bc268..40306b1f20e28 100644
--- a/clang/lib/CodeGen/CGExprScalar.cpp
+++ b/clang/lib/CodeGen/CGExprScalar.cpp
@@ -2231,8 +2231,7 @@ Value *ScalarExprEmitter::VisitMatrixSingleSubscriptExpr(
 
   for (unsigned Col = 0; Col != NumColumns; ++Col) {
     Value *ColVal = llvm::ConstantInt::get(RowIdx->getType(), Col);
-    Value *EltIdx = MB.CreateIndex(RowIdx, ColVal, NumRows, NumColumns,
-                                   /*IsRowMajor=*/false, "matrix_row_idx");
+    Value *EltIdx = MB.CreateIndex(RowIdx, ColVal, NumRows, "matrix_row_idx");
     Value *Elt =
         Builder.CreateExtractElement(FlatMatrix, EltIdx, "matrix_elem");
     Value *Lane = llvm::ConstantInt::get(Builder.getInt32Ty(), Col);
@@ -2254,10 +2253,8 @@ Value *ScalarExprEmitter::VisitMatrixSubscriptExpr(MatrixSubscriptExpr *E) {
   llvm::MatrixBuilder MB(Builder);
 
   Value *Idx;
-  unsigned NumCols = MatrixTy->getNumColumns();
   unsigned NumRows = MatrixTy->getNumRows();
-  Idx = MB.CreateIndex(RowIdx, ColumnIdx, NumRows, NumCols,
-                       /*IsRowMajor=*/false);
+  Idx = MB.CreateIndex(RowIdx, ColumnIdx, NumRows);
 
   if (CGF.CGM.getCodeGenOpts().OptimizationLevel > 0)
     MB.CreateIndexAssumption(Idx, MatrixTy->getNumElementsFlattened());
@@ -2578,8 +2575,6 @@ static Value *EmitHLSLElementwiseCast(CodeGenFunction &CGF, LValue SrcVal,
            "Flattened type on RHS must have the same number or more elements "
            "than vector on LHS.");
 
-    bool IsRowMajor = isMatrixRowMajor(CGF.getLangOpts(), DestTy);
-
     llvm::Value *V = CGF.Builder.CreateLoad(
         CGF.CreateIRTempWithoutCast(DestTy, "flatcast.tmp"));
     // V is an allocated temporary for constructing the matrix.
@@ -2594,7 +2589,7 @@ static Value *EmitHLSLElementwiseCast(CodeGenFunction &CGF, LValue SrcVal,
         llvm::Value *Cast = CGF.EmitScalarConversion(
             RVal.getScalarVal(), LoadList[LoadIdx].getType(),
             MatTy->getElementType(), Loc);
-        unsigned MatrixIdx = MatTy->getFlattenedIndex(Row, Col, IsRowMajor);
+        unsigned MatrixIdx = MatTy->getColumnMajorFlattenedIndex(Row, Col);
         V = CGF.Builder.CreateInsertElement(V, Cast, MatrixIdx);
       }
     }
diff --git a/clang/test/CodeGen/matrix-type-indexing.c b/clang/test/CodeGen/matrix-type-indexing.c
index 20eece3d646d4..cc8117e043c55 100644
--- a/clang/test/CodeGen/matrix-type-indexing.c
+++ b/clang/test/CodeGen/matrix-type-indexing.c
@@ -1,6 +1,6 @@
-// RUN: %clang_cc1 -fenable-matrix -fmatrix-memory-layout=row-major -triple x86_64-apple-darwin %s -emit-llvm -disable-llvm-passes -o - | FileCheck %s --check-prefixes=CHECK,ROW-CHECK
-// RUN: %clang_cc1 -fenable-matrix -fmatrix-memory-layout=column-major -triple x86_64-apple-darwin %s -emit-llvm -disable-llvm-passes -o - | FileCheck %s --check-prefixes=CHECK,COL-CHECK
-// RUN: %clang_cc1 -fenable-matrix -triple x86_64-apple-darwin %s -emit-llvm -disable-llvm-passes -o - | FileCheck %s --check-prefixes=CHECK,COL-CHECK
+// RUN: %clang_cc1 -fenable-matrix -fmatrix-memory-layout=row-major -triple x86_64-apple-darwin %s -emit-llvm -disable-llvm-passes -o - | FileCheck %s
+// RUN: %clang_cc1 -fenable-matrix -fmatrix-memory-layout=column-major -triple x86_64-apple-darwin %s -emit-llvm -disable-llvm-passes -o - | FileCheck %s
+// RUN: %clang_cc1 -fenable-matrix -triple x86_64-apple-darwin %s -emit-llvm -disable-llvm-passes -o - | FileCheck %s
 
 typedef float fx2x3_t __attribute__((matrix_type(2, 3)));
  float Out[6];
@@ -42,13 +42,10 @@ float returnMatrixSubscriptExpr(int row, int col, fx2x3_t M) {
 void storeAtMatrixSubscriptExpr(int row, int col, float value) {
     // CHECK-LABEL: storeAtMatrixSubscriptExpr
     // CHECK: [[value_load:%.*]] = load float, ptr [[value_ptr:%.*]], align 4
-    // ROW-CHECK: [[row_offset:%.*]] = mul i64 [[row_load:%.*]], 3
-    // ROW-CHECK-NEXT: [[row_major_index:%.*]] = add i64 [[row_offset]], [[col_load:%.*]]
-    // COL-CHECK: [[col_offset:%.*]] = mul i64 [[col_load:%.*]], 2
-    // COL-CHECK-NEXT: [[col_major_index:%.*]] = add i64 [[col_offset]], [[row_load:%.*]]
+    // CHECK: [[col_offset:%.*]] = mul i64 [[col_load:%.*]], 2
+    // CHECK-NEXT: [[col_major_index:%.*]] = add i64 [[col_offset]], [[row_load:%.*]]
     // CHECK-NEXT: [[matrix_as_vec:%.*]] = load <6 x float>, ptr @gM, align 4
-    // ROW-CHECK-NEXT: [[matrix_after_insert:%.*]] = insertelement <6 x float> [[matrix_as_vec]], float [[value_load]], i64 [[row_major_index]]
-    // COL-CHECK-NEXT: [[matrix_after_insert:%.*]] = insertelement <6 x float> [[matrix_as_vec]], float [[value_load]], i64 [[col_major_index]]
+    // CHECK-NEXT: [[matrix_after_insert:%.*]] = insertelement <6 x float> [[matrix_as_vec]], float [[value_load]], i64 [[col_major_index]]
     // CHECK-NEXT: store <6 x float> [[matrix_after_insert]], ptr @gM, align 4
     gM[row][col] = value;
 }
diff --git a/clang/test/CodeGenHLSL/BasicFeatures/MatrixElementRowColFlags.hlsl b/clang/test/CodeGenHLSL/BasicFeatures/MatrixElementRowColFlags.hlsl
index 666a387a8450a..ce14a04fe80ff 100644
--- a/clang/test/CodeGenHLSL/BasicFeatures/MatrixElementRowColFlags.hlsl
+++ b/clang/test/CodeGenHLSL/BasicFeatures/MatrixElementRowColFlags.hlsl
@@ -13,6 +13,8 @@
 
 
 // CHECK-LABEL: define {{.*}} @_Z16getScalarElementu11matrix_typeILm3ELm2EfE
+// ROW: [[TEMP_ROW_MAJOR:%.*]] = call {{.*}} <6 x float> @llvm.matrix.transpose.v6f32(<6 x float> %{{.*}}, i32 3, i32 2)
+// ROW-NEXT: store <6 x float> [[TEMP_ROW_MAJOR]], ptr
 // CHECK: load <6 x float>, ptr
 // COL-NEXT: extractelement <6 x float> {{.*}}, i32 4
 // ROW-NEXT: extractelement <6 x float> {{.*}}, i32 3
@@ -21,6 +23,8 @@ export float getScalarElement(float3x2 M) {
 }
 
 // CHECK-LABEL: define {{.*}} @_Z18getSwizzleElementsu11matrix_typeILm3ELm2EfE
+// ROW: [[TEMP_ROW_MAJOR:%.*]] = call {{.*}} <6 x float> @llvm.matrix.transpose.v6f32(<6 x float> %{{.*}}, i32 3, i32 2)
+// ROW-NEXT: store <6 x float> [[TEMP_ROW_MAJOR]], ptr
 // CHECK: load <6 x float>, ptr
 // COL-NEXT: shufflevector <6 x float> {{.*}}, <6 x float> poison, <4 x i32> <i32 0, i32 3, i32 1, i32 4>
 // ROW-NEXT: shufflevector <6 x float> {{.*}}, <6 x float> poison, <4 x i32> <i32 0, i32 1, i32 2, i32 3>
@@ -29,9 +33,34 @@ export float4 getSwizzleElements(float3x2 M) {
 }
 
 // CHECK-LABEL: define {{.*}} @_Z22getZeroBasedSwizzleEltu11matrix_typeILm3ELm2EfE
+// ROW: [[TEMP_ROW_MAJOR:%.*]] = call {{.*}} <6 x float> @llvm.matrix.transpose.v6f32(<6 x float> %{{.*}}, i32 3, i32 2)
+// ROW-NEXT: store <6 x float> [[TEMP_ROW_MAJOR]], ptr
 // CHECK: load <6 x float>, ptr
 // COL-NEXT: shufflevector <6 x float> {{.*}}, <6 x float> poison, <2 x i32> <i32 1, i32 3>
 // ROW-NEXT: shufflevector <6 x float> {{.*}}, <6 x float> poison, <2 x i32> <i32 2, i32 1>
 export float2 getZeroBasedSwizzleElt(float3x2 M) {
   return M._m10_m01;
 }
+
+// transpose(m) produces a canonical column-major register value. Matrix
+// element access materializes that prvalue in a matrix-typed temporary, which
+// must use the selected memory layout.
+export float swizzle_prvalue(float2x3 m) {
+  return transpose(m)._m01;
+}
+
+// COL-LABEL: define {{.*}} float @_Z15swizzle_prvalue
+// COL: [[TEMP:%.*]] = alloca [2 x <3 x float>]
+// COL: [[RESULT_COL_MAJOR:%.*]] = call {{.*}} <6 x float> @llvm.matrix.transpose.v6f32(<6 x float> %{{.*}}, i32 2, i32 3)
+// COL: store <6 x float> [[RESULT_COL_MAJOR]], ptr [[TEMP]]
+// COL: [[FROM_TEMP:%.*]] = load <6 x float>, ptr [[TEMP]]
+// COL: extractelement <6 x float> [[FROM_TEMP]], i32 3
+
+// ROW-LABEL: define {{.*}} float @_Z15swizzle_prvalue
+// ROW: [[TEMP:%.*]] = alloca [3 x <2 x float>]
+// ROW: [[INPUT_COL_MAJOR:%.*]] = call {{.*}} <6 x float> @llvm.matrix.transpose.v6f32(<6 x float> %{{.*}}, i32 3, i32 2)
+// ROW: [[RESULT_COL_MAJOR:%.*]] = call {{.*}} <6 x float> @llvm.matrix.transpose.v6f32(<6 x float> [[INPUT_COL_MAJOR]], i32 2, i32 3)
+// ROW: [[TEMP_ROW_MAJOR:%.*]] = call {{.*}} <6 x float> @llvm.matrix.transpose.v6f32(<6 x float> [[RESULT_COL_MAJOR]], i32 3, i32 2)
+// ROW: store <6 x float> [[TEMP_ROW_MAJOR]], ptr [[TEMP]]
+// ROW: [[FROM_TEMP:%.*]] = load <6 x float>, ptr [[TEMP]]
+// ROW: extractelement <6 x float> [[FROM_TEMP]], i32 1
diff --git a/clang/test/CodeGenHLSL/BasicFeatures/VectorElementwiseCast.hlsl b/clang/test/CodeGenHLSL/BasicFeatures/VectorElementwiseCast.hlsl
index 5156dccb26eca..90a4fa61bd1e3 100644
--- a/clang/test/CodeGenHLSL/BasicFeatures/VectorElementwiseCast.hlsl
+++ b/clang/test/CodeGenHLSL/BasicFeatures/VectorElementwiseCast.hlsl
@@ -147,12 +147,10 @@ export void call6(Derived D) {
 // CHECK-NEXT:    [[MATRIXEXT:%.*]] = extractelement <4 x float> [[TMP2]], i32 0
 // CHECK-NEXT:    [[TMP3:%.*]] = insertelement <4 x float> [[TMP1]], float [[MATRIXEXT]], i64 0
 // CHECK-NEXT:    [[TMP4:%.*]] = load <4 x float>, ptr [[MATRIX_GEP]], align 4
-// COL-CHECK-NEXT:    [[MATRIXEXT1:%.*]] = extractelement <4 x float> [[TMP4]], i32 2
-// ROW-CHECK-NEXT:    [[MATRIXEXT1:%.*]] = extractelement <4 x float> [[TMP4]], i32 1
+// CHECK-NEXT:    [[MATRIXEXT1:%.*]] = extractelement <4 x float> [[TMP4]], i32 2
 // CHECK-NEXT:    [[TMP5:%.*]] = insertelement <4 x float> [[TMP3]], float [[MATRIXEXT1]], i64 1
 // CHECK-NEXT:    [[TMP6:%.*]] = load <4 x float>, ptr [[MATRIX_GEP]], align 4
-// COL-CHECK-NEXT:    [[MATRIXEXT2:%.*]] = extractelement <4 x float> [[TMP6]], i32 1
-// ROW-CHECK-NEXT:    [[MATRIXEXT2:%.*]] = extractelement <4 x float> [[TMP6]], i32 2
+// CHECK-NEXT:    [[MATRIXEXT2:%.*]] = extractelement <4 x float> [[TMP6]], i32 1
 // CHECK-NEXT:    [[TMP7:%.*]] = insertelement <4 x float> [[TMP5]], float [[MATRIXEXT2]], i64 2
 // CHECK-NEXT:    [[TMP8:%.*]] = load <4 x float>, ptr [[MATRIX_GEP]], align 4
 // CHECK-NEXT:    [[MATRIXEXT3:%.*]] = extractelement <4 x float> [[TMP8]], i32 3
diff --git a/clang/test/CodeGenHLSL/BasicFeatures/matrix-type-indexing.hlsl b/clang/test/CodeGenHLSL/BasicFeatures/matrix-type-indexing.hlsl
index 85c65d14616a5..12425033abea1 100644
--- a/clang/test/CodeGenHLSL/BasicFeatures/matrix-type-indexing.hlsl
+++ b/clang/test/CodeGenHLSL/BasicFeatures/matrix-type-indexing.hlsl
@@ -38,12 +38,9 @@ half returnMatrixSubscriptExpr(int row, int col, half2x3 M) {
 void storeAtMatrixSubscriptExpr(int row, int col, half value) {
     // CHECK-LABEL: storeAtMatrixSubscriptExpr
     // CHECK: [[value_load:%.*]] = load half, ptr [[value_ptr:%.*]], align 2
-    // ROW-CHECK: [[row_offset:%.*]] = mul i32 [[row_load:%.*]], 3
-    // ROW-CHECK-NEXT: [[row_major_index:%.*]] = add i32 [[row_offset]], [[col_load:%.*]]
-    // COL-CHECK: [[col_offset:%.*]] = mul i32 [[col_load:%.*]], 2
-    // COL-CHECK-NEXT: [[col_major_index:%.*]] = add i32 [[col_offset]], [[row_load:%.*]]
-    // ROW-CHECK-NEXT: [[matrix_gep:%.*]] = getelementptr <6 x half>, ptr addrspace(2) @gM, i32 0, i32 [[row_major_index]]
-    // COL-CHECK-NEXT: [[matrix_gep:%.*]] = getelementptr <6 x half>, ptr addrspace(2) @gM, i32 0, i32 [[col_major_index]]
+    // CHECK: [[col_offset:%.*]] = mul i32 [[col_load:%.*]], 2
+    // CHECK-NEXT: [[col_major_index:%.*]] = add i32 [[col_offset]], [[row_load:%.*]]
+    // CHECK-NEXT: [[matrix_gep:%.*]] = getelementptr <6 x half>, ptr addrspace(2) @gM, i32 0, i32 [[col_major_index]]
     // CHECK-NEXT: store half [[value_load]], ptr addrspace(2) [[matrix_gep]], align 2
     gM[row][col] = value;
 }
diff --git a/clang/test/CodeGenHLSL/matrix-layout-attr-overrides-default.hlsl b/clang/test/CodeGenHLSL/matrix-layout-attr-overrides-default.hlsl
index a04cf58b45b7c..cb73364bd6e0e 100644
--- a/clang/test/CodeGenHLSL/matrix-layout-attr-overrides-default.hlsl
+++ b/clang/test/CodeGenHLSL/matrix-layout-attr-overrides-default.hlsl
@@ -7,6 +7,7 @@
 //
 //   * `MatrixSubscriptExpr` index computation
 //   * `MatrixSingleSubscriptExpr` row extraction
+//   * `CK_HLSLElementwiseCast` matrix construction
 //   * `__builtin_hlsl_mul` matrix-multiply transpose insertion
 //   * `__builtin_hlsl_transpose` row/col dimension swap
 //   * `CK_HLSLMatrixTruncation` shuffle mask
@@ -40,8 +41,8 @@ export float subscript_cm(int row, int col, column_major float2x3 m) {
 // CHECK: extractelement <6 x float> %{{.*}}, i32 [[IDX]]
 
 // -----------------------------------------------------------------------------
-// MatrixSingleSubscriptExpr (row extraction): attribute selects the per-element
-// index formula even when the TU default disagrees.
+// MatrixSingleSubscriptExpr (row extraction) uses canonical column-major
+// indexing even when the destination storage layout is row-major.
 // -----------------------------------------------------------------------------
 
 // Row extraction also indexes the canonical column-major prvalue.
@@ -66,6 +67,20 @@ export float3 row_extract_cm(int row, column_major float2x3 m) {
 // CHECK: add i32 2, [[ROW]]
 // CHECK: add i32 4, [[ROW]]
 
+// -----------------------------------------------------------------------------
+// CK_HLSLElementwiseCast produces a canonical column-major register value.
+// An explicit row-major destination affects only the subsequent memory store.
+// -----------------------------------------------------------------------------
+typedef row_major float2x2 RowMajorMatrix;
+
+export float cast_row_major(float4 v) {
+  RowMajorMatrix m = (RowMajorMatrix)v;
+  return m[0][1];
+}
+// CHECK-LABEL: define {{.*}} float @_Z14cast_row_major
+// CHECK: [[SECOND:%.*]] = extractelement <4 x float> %{{.*}}, i32 1
+// CHECK: insertelement <4 x float> %{{.*}}, float [[SECOND]], i64 2
+
 // -----------------------------------------------------------------------------
 // __builtin_hlsl_mul (vector * matrix): row-major operand triggers a transpose
 // before the column-major matrix.multiply intrinsic.
diff --git a/llvm/include/llvm/IR/MatrixBuilder.h b/llvm/include/llvm/IR/MatrixBuilder.h
index 41cd5ea0efd93..8c0609401ebd6 100644
--- a/llvm/include/llvm/IR/MatrixBuilder.h
+++ b/llvm/include/llvm/IR/MatrixBuilder.h
@@ -260,39 +260,19 @@ class MatrixBuilder {
     else
       B.CreateAssumption(Cmp);
   }
-  /// Compute the index to access the element at (\p RowIdx, \p ColumnIdx) from
-  /// a matrix with \p NumRows or \p NumCols embedded in a vector depending
-  /// on matrix major ordering.
+  /// Compute the column-major index to access the element at
+  /// (\p RowIdx, \p ColumnIdx) from a matrix with \p NumRows embedded in a
+  /// vector.
   Value *CreateIndex(Value *RowIdx, Value *ColumnIdx, unsigned NumRows,
-                     unsigned NumCols, bool IsMatrixRowMajor = false,
                      Twine const &Name = "") {
     unsigned MaxWidth = std::max(RowIdx->getType()->getScalarSizeInBits(),
                                  ColumnIdx->getType()->getScalarSizeInBits());
     Type *IntTy = IntegerType::get(RowIdx->getType()->getContext(), MaxWidth);
     RowIdx = B.CreateZExt(RowIdx, IntTy);
     ColumnIdx = B.CreateZExt(ColumnIdx, IntTy);
-    if (IsMatrixRowMajor) {
-      Value *NumColsV = B.getIntN(MaxWidth, NumCols);
-      return CreateRowMajorIndex(RowIdx, ColumnIdx, NumColsV, Name);
-    }
     Value *NumRowsV = B.getIntN(MaxWidth, NumRows);
-    return CreateColumnMajorIndex(RowIdx, ColumnIdx, NumRowsV, Name);
-  }
-
-private:
-  /// Compute the index to access the element at (\p RowIdx, \p ColumnIdx) from
-  /// a matrix with \p NumRows embedded in a vector.
-  Value *CreateColumnMajorIndex(Value *RowIdx, Value *ColumnIdx,
-                                Value *NumRowsV, Twine const &Name) {
     return B.CreateAdd(B.CreateMul(ColumnIdx, NumRowsV), RowIdx);
   }
-
-  /// Compute the index to access the element at (\p RowIdx, \p ColumnIdx) from
-  /// a matrix with \p NumCols embedded in a vector.
-  Value *CreateRowMajorIndex(Value *RowIdx, Value *ColumnIdx, Value *NumColsV,
-                             Twine const &Name) {
-    return B.CreateAdd(B.CreateMul(RowIdx, NumColsV), ColumnIdx);
-  }
 };
 
 } // end namespace llvm



More information about the llvm-commits mailing list