[clang] [HLSL][Matrix] Use canonical column-major indexing in CodeGen (PR #227087)
Farzon Lotfi via cfe-commits
cfe-commits at lists.llvm.org
Tue Sep 29 07:46:09 PDT 2026
https://github.com/farzonl updated https://github.com/llvm/llvm-project/pull/227087
>From 61a1deb95dd2101e6f084f0ad46ec1581a9dc39f 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] [Clang][HLSL] Keep matrix SSA values in column-major order
Keep matrix register values in canonical column-major order while preserving
the declared matrix layout for memory accesses.
Construct elementwise matrix casts in column-major lane order and use
column-major indexing when extracting elements from SSA values. Continue to
use layout-aware indexing for LValues and other physical memory operations.
Materialize matrix prvalues through the matrix-aware scalar store path so
row-major temporaries are transposed at the register-to-memory boundary.
Add coverage for matrix indexing, casts, swizzles, transposed prvalues, and
explicit row-major out-parameter accesses.
---
clang/lib/CodeGen/CGExpr.cpp | 28 ++++++------
clang/lib/CodeGen/CGExprScalar.cpp | 9 ++--
.../MatrixElementRowColFlags.hlsl | 29 ++++++++++++
.../BasicFeatures/MatrixElementTypeCast.hlsl | 8 ++--
.../matrix-layout-attr-overrides-default.hlsl | 45 ++++++++++++++++++-
5 files changed, 92 insertions(+), 27 deletions(-)
diff --git a/clang/lib/CodeGen/CGExpr.cpp b/clang/lib/CodeGen/CGExpr.cpp
index 25e775f8328b9..8410233238a79 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());
@@ -2651,9 +2649,9 @@ RValue CodeGenFunction::EmitLoadOfLValue(LValue LV, SourceLocation Loc) {
else
ColIdx = llvm::ConstantInt::get(Row->getType(), Col);
bool IsMatrixRowMajor = isMatrixRowMajor(getLangOpts(), MatTy);
- llvm::Value *EltIndex =
+ llvm::Value *MemoryIndex =
MB.CreateIndex(Row, ColIdx, NumRows, NumCols, IsMatrixRowMajor);
- llvm::Value *Elt = Builder.CreateExtractElement(MatrixVec, EltIndex);
+ llvm::Value *Elt = Builder.CreateExtractElement(MatrixVec, MemoryIndex);
llvm::Value *Lane = llvm::ConstantInt::get(Builder.getInt32Ty(), Col);
Result = Builder.CreateInsertElement(Result, Elt, Lane);
}
@@ -2970,13 +2968,13 @@ void CodeGenFunction::EmitStoreThroughLValue(RValue Src, LValue Dst,
else
ColIdx = llvm::ConstantInt::get(Row->getType(), Col);
bool IsMatrixRowMajor = isMatrixRowMajor(getLangOpts(), Dst.getType());
- llvm::Value *EltIndex =
+ llvm::Value *MemoryIndex =
MB.CreateIndex(Row, ColIdx, NumRows, NumCols, IsMatrixRowMajor);
llvm::Value *Lane = llvm::ConstantInt::get(Builder.getInt32Ty(), Col);
llvm::Value *Zero = llvm::ConstantInt::get(Int32Ty, 0);
llvm::Value *NewElt = Builder.CreateExtractElement(RowVal, Lane);
- Address DstElemAddr =
- Builder.CreateGEP(DstAddr, {Zero, EltIndex}, DestAddrTy, ElemAlign);
+ Address DstElemAddr = Builder.CreateGEP(DstAddr, {Zero, MemoryIndex},
+ DestAddrTy, ElemAlign);
Builder.CreateStore(NewElt, DstElemAddr, Dst.isVolatileQualified());
}
@@ -5436,11 +5434,11 @@ LValue CodeGenFunction::EmitMatrixSubscriptExpr(const MatrixSubscriptExpr *E) {
unsigned NumRows = MatrixTy->getNumRows();
bool IsMatrixRowMajor =
isMatrixRowMajor(getLangOpts(), E->getBase()->getType());
- llvm::Value *FinalIdx =
+ llvm::Value *MemoryIndex =
MB.CreateIndex(RowIdx, ColIdx, NumRows, NumCols, IsMatrixRowMajor);
return LValue::MakeMatrixElt(
- MaybeConvertMatrixAddress(Base.getAddress(), *this), FinalIdx,
+ MaybeConvertMatrixAddress(Base.getAddress(), *this), MemoryIndex,
E->getBase()->getType(), Base.getBaseInfo(), TBAAAccessInfo());
}
@@ -7669,10 +7667,10 @@ void CodeGenFunction::FlattenAccessAndTypeLValue(
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 *MemoryIndex = MB.CreateIndex(RowIdx, ColIdx, NumRows,
+ NumCols, IsMatrixRowMajor);
LValue LV =
- LValue::MakeMatrixElt(MatAddr, Idx, MT->getElementType(),
+ LValue::MakeMatrixElt(MatAddr, MemoryIndex, MT->getElementType(),
Base.getBaseInfo(), TBAAAccessInfo());
AccessList.emplace_back(LV);
}
diff --git a/clang/lib/CodeGen/CGExprScalar.cpp b/clang/lib/CodeGen/CGExprScalar.cpp
index 308c45bcbee4f..67f6673234a27 100644
--- a/clang/lib/CodeGen/CGExprScalar.cpp
+++ b/clang/lib/CodeGen/CGExprScalar.cpp
@@ -2253,11 +2253,10 @@ Value *ScalarExprEmitter::VisitMatrixSubscriptExpr(MatrixSubscriptExpr *E) {
const auto *MatrixTy = E->getBase()->getType()->castAs<ConstantMatrixType>();
llvm::MatrixBuilder MB(Builder);
- Value *Idx;
unsigned NumCols = MatrixTy->getNumColumns();
unsigned NumRows = MatrixTy->getNumRows();
- Idx = MB.CreateIndex(RowIdx, ColumnIdx, NumRows, NumCols,
- /*IsRowMajor=*/false);
+ Value *Idx = MB.CreateIndex(RowIdx, ColumnIdx, NumRows, NumCols,
+ /*IsRowMajor=*/false);
if (CGF.CGM.getCodeGenOpts().OptimizationLevel > 0)
MB.CreateIndexAssumption(Idx, MatrixTy->getNumElementsFlattened());
@@ -2577,8 +2576,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 = llvm::PoisonValue::get(CGF.ConvertType(DestTy));
// V is an allocated temporary for constructing the matrix.
for (unsigned Row = 0, RE = MatTy->getNumRows(); Row < RE; Row++) {
@@ -2592,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/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/MatrixElementTypeCast.hlsl b/clang/test/CodeGenHLSL/BasicFeatures/MatrixElementTypeCast.hlsl
index 8a27efef75255..126b4d238ada1 100644
--- a/clang/test/CodeGenHLSL/BasicFeatures/MatrixElementTypeCast.hlsl
+++ b/clang/test/CodeGenHLSL/BasicFeatures/MatrixElementTypeCast.hlsl
@@ -301,10 +301,10 @@ struct Derived : BFields {
// ROW-CHECK-NEXT: [[BF_SHL:%.*]] = shl i24 [[BF_LOAD]], 9
// ROW-CHECK-NEXT: [[BF_ASHR:%.*]] = ashr i24 [[BF_SHL]], 9
// ROW-CHECK-NEXT: [[BF_CAST:%.*]] = sext i24 [[BF_ASHR]] to i32
-// ROW-CHECK-NEXT: [[TMP3:%.*]] = insertelement <4 x i32> [[TMP2]], i32 [[BF_CAST]], i64 1
+// ROW-CHECK-NEXT: [[TMP3:%.*]] = insertelement <4 x i32> [[TMP2]], i32 [[BF_CAST]], i64 2
// ROW-CHECK-NEXT: [[TMP4:%.*]] = load float, ptr [[GEP2]], align 4
// ROW-CHECK-NEXT: [[CONV4:%.*]] = fptosi float [[TMP4]] to i32
-// ROW-CHECK-NEXT: [[TMP5:%.*]] = insertelement <4 x i32> [[TMP3]], i32 [[CONV4]], i64 2
+// ROW-CHECK-NEXT: [[TMP5:%.*]] = insertelement <4 x i32> [[TMP3]], i32 [[CONV4]], i64 1
// ROW-CHECK-NEXT: [[TMP6:%.*]] = load i32, ptr [[GEP3]], align 4
// ROW-CHECK-NEXT: [[TMP7:%.*]] = insertelement <4 x i32> [[TMP5]], i32 [[TMP6]], i64 3
// ROW-CHECK-NEXT: [[TMP8:%.*]] = call <4 x i32> @llvm.matrix.transpose.v4i32(<4 x i32> [[TMP7]], i32 2, i32 2)
@@ -361,10 +361,10 @@ void call4(Derived D) {
// ROW-CHECK-NEXT: [[TMP3:%.*]] = insertelement <4 x float> poison, float [[VECEXT]], i64 0
// ROW-CHECK-NEXT: [[TMP4:%.*]] = load <4 x float>, ptr [[VECTOR_GEP]], align 4
// ROW-CHECK-NEXT: [[VECEXT1:%.*]] = extractelement <4 x float> [[TMP4]], i32 1
-// ROW-CHECK-NEXT: [[TMP5:%.*]] = insertelement <4 x float> [[TMP3]], float [[VECEXT1]], i64 1
+// ROW-CHECK-NEXT: [[TMP5:%.*]] = insertelement <4 x float> [[TMP3]], float [[VECEXT1]], i64 2
// ROW-CHECK-NEXT: [[TMP6:%.*]] = load <4 x float>, ptr [[VECTOR_GEP]], align 4
// ROW-CHECK-NEXT: [[VECEXT2:%.*]] = extractelement <4 x float> [[TMP6]], i32 2
-// ROW-CHECK-NEXT: [[TMP7:%.*]] = insertelement <4 x float> [[TMP5]], float [[VECEXT2]], i64 2
+// ROW-CHECK-NEXT: [[TMP7:%.*]] = insertelement <4 x float> [[TMP5]], float [[VECEXT2]], i64 1
// ROW-CHECK-NEXT: [[TMP8:%.*]] = load <4 x float>, ptr [[VECTOR_GEP]], align 4
// ROW-CHECK-NEXT: [[VECEXT3:%.*]] = extractelement <4 x float> [[TMP8]], i32 3
// ROW-CHECK-NEXT: [[TMP9:%.*]] = insertelement <4 x float> [[TMP7]], float [[VECEXT3]], i64 3
diff --git a/clang/test/CodeGenHLSL/matrix-layout-attr-overrides-default.hlsl b/clang/test/CodeGenHLSL/matrix-layout-attr-overrides-default.hlsl
index a04cf58b45b7c..25a788ca2991f 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
@@ -26,6 +27,32 @@ export float subscript_rm(int row, int col, row_major float2x3 m) {
// CHECK: [[IDX:%.*]] = add i32 [[OFFSET]], [[ROW]]
// CHECK: extractelement <6 x float> %{{.*}}, i32 [[IDX]]
+// Matrix out parameters point directly to storage, so their element indices
+// must use the declared memory layout.
+export void write_row_major_element(out row_major float2x3 m, float value) {
+ m[0][1] = value;
+}
+// CHECK-LABEL: define void @_Z23write_row_major_element{{.*}}(
+// CHECK: [[M_PTR:%.*]] = load ptr, ptr %m.addr
+// CHECK: [[ELEMENT_ADDR:%.*]] = getelementptr <6 x float>, ptr [[M_PTR]], i32 0, i32 1
+// CHECK: store float %{{.*}}, ptr [[ELEMENT_ADDR]]
+
+export float call_write_row_major_element(float value) {
+ column_major float2x3 source = {1, 2, 3, 4, 5, 6};
+ write_row_major_element(source, value);
+ return source._m01;
+}
+// CHECK-LABEL: define {{.*}} float @_Z28call_write_row_major_element
+// CHECK: [[SOURCE:%.*]] = alloca [3 x <2 x float>]
+// CHECK-NEXT: [[OUT_TEMP:%.*]] = alloca <6 x float>
+// CHECK: call void @_Z23write_row_major_element{{.*}}(ptr {{.*}} [[OUT_TEMP]], float {{.*}})
+// CHECK-NEXT: [[ROW_MAJOR:%.*]] = load <6 x float>, ptr [[OUT_TEMP]]
+// CHECK-NEXT: [[TO_COLUMN_MAJOR:%.*]] = call {{.*}} <6 x float> @llvm.matrix.transpose.v6f32(<6 x float> [[ROW_MAJOR]], i32 3, i32 2)
+// CHECK-NEXT: store <6 x float> [[TO_COLUMN_MAJOR]], ptr [[SOURCE]]
+// CHECK: [[CANONICAL:%.*]] = load <6 x float>, ptr [[SOURCE]]
+// CHECK-NEXT: [[ELEMENT:%.*]] = extractelement <6 x float> [[CANONICAL]], i32 2
+// CHECK-NEXT: ret float [[ELEMENT]]
+
// -----------------------------------------------------------------------------
// MatrixSubscriptExpr indexing: column-major attr -> Col*NumRows + Row
// -----------------------------------------------------------------------------
@@ -40,8 +67,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 +93,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.
More information about the cfe-commits
mailing list