[clang] [HLSL] Modify the Matrix orientation to be AST based instead of Attribute (PR #214001)
Farzon Lotfi via cfe-commits
cfe-commits at lists.llvm.org
Tue Aug 11 11:52:34 PDT 2026
https://github.com/farzonl updated https://github.com/llvm/llvm-project/pull/214001
>From 695dfeb131f4e8d84e757ffae526dad639ff1443 Mon Sep 17 00:00:00 2001
From: Farzon Lotfi <farzonlotfi at microsoft.com>
Date: Mon, 3 Aug 2026 14:43:54 -0400
Subject: [PATCH 1/2] [HLSL] Modify the Matrix orientation to be AST based
instead of Attribute
fixes https://github.com/llvm/llvm-project/issues/213996
fixes https://github.com/llvm/llvm-project/issues/211977
fixes https://godbolt.org/z/rhTYx1KGf
Store explicit matrix orientation in ConstantMatrixType so layout
survives desugaring, arrays, resources, serialization, and template deduction.
Assisted by Copilot with GPT-5.6-Sol
---
clang/include/clang/AST/ASTContext.h | 7 +-
clang/include/clang/AST/MatrixUtils.h | 26 ++------
clang/include/clang/AST/PropertiesBase.td | 7 ++
clang/include/clang/AST/TypeBase.h | 25 ++++++--
clang/include/clang/AST/TypeProperties.td | 5 +-
clang/lib/AST/ASTContext.cpp | 44 +++++++++++--
clang/lib/AST/ASTImporter.cpp | 2 +-
clang/lib/AST/Type.cpp | 14 ++--
clang/lib/CodeGen/HLSLBufferLayoutBuilder.cpp | 3 +-
clang/lib/Sema/SemaHLSL.cpp | 12 +---
clang/lib/Sema/SemaTemplateDeduction.cpp | 11 ++++
clang/lib/Sema/SemaType.cpp | 11 +++-
clang/lib/Sema/TreeTransform.h | 11 ++++
.../test/CodeGenHLSL/matrix-array-layout.hlsl | 64 +++++++++++++++++++
14 files changed, 188 insertions(+), 54 deletions(-)
create mode 100644 clang/test/CodeGenHLSL/matrix-array-layout.hlsl
diff --git a/clang/include/clang/AST/ASTContext.h b/clang/include/clang/AST/ASTContext.h
index 763039e690dec..4b8adcceed285 100644
--- a/clang/include/clang/AST/ASTContext.h
+++ b/clang/include/clang/AST/ASTContext.h
@@ -1850,14 +1850,17 @@ class ASTContext : public RefCountedBase<ASTContext> {
///
/// \pre \p ElementType must be a valid matrix element type (see
/// MatrixType::isValidElementType).
- QualType getConstantMatrixType(QualType ElementType, unsigned NumRows,
- unsigned NumColumns) const;
+ QualType getConstantMatrixType(
+ QualType ElementType, unsigned NumRows, unsigned NumColumns,
+ std::optional<MatrixType::LayoutKind> Layout = std::nullopt) const;
/// Return the unique reference to the matrix type of the specified element
/// type and size
QualType getDependentSizedMatrixType(QualType ElementType, Expr *RowExpr,
Expr *ColumnExpr,
SourceLocation AttrLoc) const;
+ QualType getMatrixTypeWithLayout(QualType T,
+ MatrixType::LayoutKind Layout) const;
QualType getDependentAddressSpaceType(QualType PointeeType,
Expr *AddrSpaceExpr,
diff --git a/clang/include/clang/AST/MatrixUtils.h b/clang/include/clang/AST/MatrixUtils.h
index ef6cbba6ba7c0..1057bbd67a526 100644
--- a/clang/include/clang/AST/MatrixUtils.h
+++ b/clang/include/clang/AST/MatrixUtils.h
@@ -15,32 +15,18 @@
#define LLVM_CLANG_AST_MATRIXUTILS_H
#include "clang/AST/Type.h"
-#include "clang/Basic/AttrKinds.h"
#include "clang/Basic/LangOptions.h"
namespace clang {
/// Returns true if matrices of \p T should be laid out in row-major order.
///
-/// In HLSL mode, an `HLSLRowMajor` / `HLSLColumnMajor` AttributedType anywhere
-/// in the sugar chain of \p T (imprinted by Sema when a source decl carries
-/// `[[hlsl::row_major]]` / `[[hlsl::column_major]]`) takes precedence over the
-/// `-fmatrix-memory-layout=` default carried in \p LangOpts. Otherwise the
-/// LangOptions default is used.
+/// An explicit layout stored on the matrix type takes precedence over the
+/// `-fmatrix-memory-layout=` default carried in \p LangOpts.
inline bool isMatrixRowMajor(const LangOptions &LangOpts, QualType T) {
- if (LangOpts.HLSL && !T.isNull()) {
- QualType Cur = T;
- while (const auto *AT = Cur->getAs<AttributedType>()) {
- switch (AT->getAttrKind()) {
- case attr::HLSLRowMajor:
- return true;
- case attr::HLSLColumnMajor:
- return false;
- default:
- break;
- }
- Cur = AT->getModifiedType();
- }
- }
+ if (LangOpts.HLSL && !T.isNull())
+ if (const auto *MT = T->getAs<ConstantMatrixType>())
+ if (auto Layout = MT->getLayout())
+ return *Layout == MatrixType::LayoutKind::RowMajor;
return LangOpts.getDefaultMatrixMemoryLayout() ==
LangOptions::MatrixMemoryLayout::MatrixRowMajor;
}
diff --git a/clang/include/clang/AST/PropertiesBase.td b/clang/include/clang/AST/PropertiesBase.td
index 25ef4c26a9aa1..bab9787d0dc13 100644
--- a/clang/include/clang/AST/PropertiesBase.td
+++ b/clang/include/clang/AST/PropertiesBase.td
@@ -130,6 +130,13 @@ def LValuePathSerializationHelper :
PropertyType<"APValue::LValuePathSerializationHelper"> {
let BufferElementTypes = [ LValuePathEntry ];
}
+def MatrixLayoutKind : EnumPropertyType<"MatrixType::LayoutKind"> {
+ let PackOptional =
+ "value.value_or(static_cast<MatrixType::LayoutKind>(2))";
+ let UnpackOptional =
+ "value == static_cast<MatrixType::LayoutKind>(2) ? std::nullopt : "
+ "std::optional<MatrixType::LayoutKind>(value)";
+}
def NestedNameSpecifier : PropertyType<"NestedNameSpecifier">;
def NestedNameSpecifierKind : EnumPropertyType<"NestedNameSpecifier::Kind">;
def OverloadedOperatorKind : EnumPropertyType;
diff --git a/clang/include/clang/AST/TypeBase.h b/clang/include/clang/AST/TypeBase.h
index 4851c4e5185dd..b6129b1814d70 100644
--- a/clang/include/clang/AST/TypeBase.h
+++ b/clang/include/clang/AST/TypeBase.h
@@ -4443,9 +4443,16 @@ class MatrixType : public Type, public llvm::FoldingSetNode {
protected:
friend class ASTContext;
+public:
+ /// A matrix's explicit `row_major`/`column_major` orientation, if any was
+ /// spelled in source; absent means no explicit orientation was given.
+ enum class LayoutKind : uint8_t { RowMajor, ColumnMajor };
+
+private:
/// The element type of the matrix.
QualType ElementType;
+protected:
MatrixType(QualType ElementTy, QualType CanonElementTy);
MatrixType(TypeClass TypeClass, QualType ElementTy, QualType CanonElementTy,
@@ -4497,11 +4504,18 @@ class ConstantMatrixType final : public MatrixType {
unsigned NumRows;
unsigned NumColumns;
+ /// Only ConstantMatrixType has a layout; it isn't meaningful until the
+ /// matrix dimensions are concrete. Absent if no explicit orientation was
+ /// spelled in source.
+ std::optional<LayoutKind> Layout;
+
ConstantMatrixType(QualType MatrixElementType, unsigned NRows,
- unsigned NColumns, QualType CanonElementType);
+ unsigned NColumns, QualType CanonElementType,
+ std::optional<LayoutKind> Layout);
ConstantMatrixType(TypeClass typeClass, QualType MatrixType, unsigned NRows,
- unsigned NColumns, QualType CanonElementType);
+ unsigned NColumns, QualType CanonElementType,
+ std::optional<LayoutKind> Layout);
public:
/// Returns the number of rows in the matrix.
@@ -4510,6 +4524,8 @@ class ConstantMatrixType final : public MatrixType {
/// Returns the number of columns in the matrix.
unsigned getNumColumns() const { return NumColumns; }
+ std::optional<LayoutKind> getLayout() const { return Layout; }
+
/// Returns the number of elements required to embed the matrix into a vector.
unsigned getNumElementsFlattened() const {
return getNumRows() * getNumColumns();
@@ -4555,16 +4571,17 @@ class ConstantMatrixType final : public MatrixType {
}
void Profile(llvm::FoldingSetNodeID &ID) {
- Profile(ID, getElementType(), getNumRows(), getNumColumns(),
+ Profile(ID, getElementType(), getNumRows(), getNumColumns(), getLayout(),
getTypeClass());
}
static void Profile(llvm::FoldingSetNodeID &ID, QualType ElementType,
unsigned NumRows, unsigned NumColumns,
- TypeClass TypeClass) {
+ std::optional<LayoutKind> Layout, TypeClass TypeClass) {
ID.AddPointer(ElementType.getAsOpaquePtr());
ID.AddInteger(NumRows);
ID.AddInteger(NumColumns);
+ ID.AddInteger(Layout ? llvm::to_underlying(*Layout) + 1 : 0);
ID.AddInteger(TypeClass);
}
diff --git a/clang/include/clang/AST/TypeProperties.td b/clang/include/clang/AST/TypeProperties.td
index 1185a3b1dc670..f9066d07b0f9c 100644
--- a/clang/include/clang/AST/TypeProperties.td
+++ b/clang/include/clang/AST/TypeProperties.td
@@ -254,9 +254,12 @@ let Class = ConstantMatrixType in {
def : Property<"numColumns", UInt32> {
let Read = [{ node->getNumColumns() }];
}
+ def : Property<"layout", Optional<MatrixLayoutKind>> {
+ let Read = [{ node->getLayout() }];
+ }
def : Creator<[{
- return ctx.getConstantMatrixType(elementType, numRows, numColumns);
+ return ctx.getConstantMatrixType(elementType, numRows, numColumns, layout);
}]>;
}
diff --git a/clang/lib/AST/ASTContext.cpp b/clang/lib/AST/ASTContext.cpp
index 5f1e5b30ee50c..c21ad7b7e01b2 100644
--- a/clang/lib/AST/ASTContext.cpp
+++ b/clang/lib/AST/ASTContext.cpp
@@ -4855,10 +4855,11 @@ ASTContext::getDependentSizedExtVectorType(QualType vecType,
return QualType(New, 0);
}
-QualType ASTContext::getConstantMatrixType(QualType ElementTy, unsigned NumRows,
- unsigned NumColumns) const {
+QualType ASTContext::getConstantMatrixType(
+ QualType ElementTy, unsigned NumRows, unsigned NumColumns,
+ std::optional<MatrixType::LayoutKind> Layout) const {
llvm::FoldingSetNodeID ID;
- ConstantMatrixType::Profile(ID, ElementTy, NumRows, NumColumns,
+ ConstantMatrixType::Profile(ID, ElementTy, NumRows, NumColumns, Layout,
Type::ConstantMatrix);
assert(MatrixType::isValidElementType(ElementTy, getLangOpts()) &&
@@ -4872,8 +4873,8 @@ QualType ASTContext::getConstantMatrixType(QualType ElementTy, unsigned NumRows,
QualType Canonical;
if (!ElementTy.isCanonical()) {
- Canonical =
- getConstantMatrixType(getCanonicalType(ElementTy), NumRows, NumColumns);
+ Canonical = getConstantMatrixType(getCanonicalType(ElementTy), NumRows,
+ NumColumns, Layout);
ConstantMatrixType *NewIP = MatrixTypes.FindNodeOrInsertPos(ID, InsertPos);
assert(!NewIP && "Matrix type shouldn't already exist in the map");
@@ -4881,7 +4882,7 @@ QualType ASTContext::getConstantMatrixType(QualType ElementTy, unsigned NumRows,
}
auto *New = new (*this, alignof(ConstantMatrixType))
- ConstantMatrixType(ElementTy, NumRows, NumColumns, Canonical);
+ ConstantMatrixType(ElementTy, NumRows, NumColumns, Canonical, Layout);
MatrixTypes.InsertNode(New, InsertPos);
Types.push_back(New);
return QualType(New, 0);
@@ -4928,6 +4929,37 @@ QualType ASTContext::getDependentSizedMatrixType(QualType ElementTy,
return QualType(New, 0);
}
+QualType
+ASTContext::getMatrixTypeWithLayout(QualType T,
+ MatrixType::LayoutKind Layout) const {
+ Qualifiers Quals = T.getQualifiers();
+ const Type *Ty = T->getUnqualifiedDesugaredType();
+
+ if (const auto *MT = dyn_cast<ConstantMatrixType>(Ty))
+ return getQualifiedType(getConstantMatrixType(MT->getElementType(),
+ MT->getNumRows(),
+ MT->getNumColumns(), Layout),
+ Quals);
+
+ // `row_major`/`column_major` are HLSL-only and only ever applied to a
+ // non-dependent type, so a ConstantArrayType (or its HLSL parameter-decayed
+ // ArrayParameterType subclass) is the only array kind that can wrap a
+ // matrix here: HLSL has no VLAs or incomplete data arrays, and dependent
+ // array bounds imply a dependent type, which is rejected before this point.
+ const auto *CAT = dyn_cast<ConstantArrayType>(Ty);
+ if (!CAT)
+ return T;
+
+ QualType Result = getConstantArrayType(
+ getMatrixTypeWithLayout(CAT->getElementType(), Layout), CAT->getSize(),
+ CAT->getSizeExpr(), CAT->getSizeModifier(),
+ CAT->getIndexTypeCVRQualifiers());
+ if (isa<ArrayParameterType>(CAT))
+ Result = getArrayParameterType(Result);
+
+ return getQualifiedType(Result, Quals);
+}
+
QualType ASTContext::getDependentAddressSpaceType(QualType PointeeType,
Expr *AddrSpaceExpr,
SourceLocation AttrLoc) const {
diff --git a/clang/lib/AST/ASTImporter.cpp b/clang/lib/AST/ASTImporter.cpp
index 3ad71a223903c..f3345cedb6517 100644
--- a/clang/lib/AST/ASTImporter.cpp
+++ b/clang/lib/AST/ASTImporter.cpp
@@ -2095,7 +2095,7 @@ ExpectedType clang::ASTNodeImporter::VisitConstantMatrixType(
return ToElementTypeOrErr.takeError();
return Importer.getToContext().getConstantMatrixType(
- *ToElementTypeOrErr, T->getNumRows(), T->getNumColumns());
+ *ToElementTypeOrErr, T->getNumRows(), T->getNumColumns(), T->getLayout());
}
ExpectedType clang::ASTNodeImporter::VisitDependentAddressSpaceType(
diff --git a/clang/lib/AST/Type.cpp b/clang/lib/AST/Type.cpp
index e51e7de9f176a..3dae2684236b8 100644
--- a/clang/lib/AST/Type.cpp
+++ b/clang/lib/AST/Type.cpp
@@ -413,15 +413,17 @@ MatrixType::MatrixType(TypeClass tc, QualType matrixType, QualType canonType,
ElementType(matrixType) {}
ConstantMatrixType::ConstantMatrixType(QualType matrixType, unsigned nRows,
- unsigned nColumns, QualType canonType)
- : ConstantMatrixType(ConstantMatrix, matrixType, nRows, nColumns,
- canonType) {}
+ unsigned nColumns, QualType canonType,
+ std::optional<LayoutKind> Layout)
+ : ConstantMatrixType(ConstantMatrix, matrixType, nRows, nColumns, canonType,
+ Layout) {}
ConstantMatrixType::ConstantMatrixType(TypeClass tc, QualType matrixType,
unsigned nRows, unsigned nColumns,
- QualType canonType)
+ QualType canonType,
+ std::optional<LayoutKind> Layout)
: MatrixType(tc, matrixType, canonType), NumRows(nRows),
- NumColumns(nColumns) {}
+ NumColumns(nColumns), Layout(Layout) {}
DependentSizedMatrixType::DependentSizedMatrixType(QualType ElementType,
QualType CanonicalType,
@@ -1187,7 +1189,7 @@ struct SimpleTransformVisitor : public TypeVisitor<Derived, QualType> {
return QualType(T, 0);
return Ctx.getConstantMatrixType(elementType, T->getNumRows(),
- T->getNumColumns());
+ T->getNumColumns(), T->getLayout());
}
QualType VisitOverflowBehaviorType(const OverflowBehaviorType *T) {
diff --git a/clang/lib/CodeGen/HLSLBufferLayoutBuilder.cpp b/clang/lib/CodeGen/HLSLBufferLayoutBuilder.cpp
index b09fc4c91df2a..075eaa8d82c13 100644
--- a/clang/lib/CodeGen/HLSLBufferLayoutBuilder.cpp
+++ b/clang/lib/CodeGen/HLSLBufferLayoutBuilder.cpp
@@ -136,8 +136,7 @@ llvm::Type *HLSLBufferLayoutBuilder::layOutArray(const ConstantArrayType *AT) {
llvm::Type *HLSLBufferLayoutBuilder::layOutMatrix(QualType Ty) {
// ConvertTypeForMem already handles row/column-major layout and bool
// promotion, producing [Count x <VecLen x EltTy>]. We just need to add
- // cbuffer padding between the array elements. Pass the sugared QualType so
- // that the `row_major`/`column_major` orientation attribute is preserved.
+ // cbuffer padding between the array elements.
llvm::ArrayType *MemTy =
cast<llvm::ArrayType>(CGM.getTypes().ConvertTypeForMem(Ty));
return padArrayElements(MemTy->getElementType(), MemTy->getNumElements());
diff --git a/clang/lib/Sema/SemaHLSL.cpp b/clang/lib/Sema/SemaHLSL.cpp
index c353c3fec3f62..05ae000b6f18f 100644
--- a/clang/lib/Sema/SemaHLSL.cpp
+++ b/clang/lib/Sema/SemaHLSL.cpp
@@ -525,15 +525,8 @@ static const Type *createHostLayoutType(Sema &S, const Type *Ty) {
return Ty;
}
-// Returns the type to use for a host layout struct field. For most types this
-// is the unqualified desugared type. Matrix types, however, retain their sugar
-// so that the row_major/column_major orientation (carried as an AttributedType)
-// is preserved; the orientation determines the in-memory cbuffer layout.
static const Type *getHostLayoutFieldType(QualType QT) {
- const Type *Desugared = QT->getUnqualifiedDesugaredType();
- if (Desugared->isConstantMatrixType())
- return QT.getTypePtr();
- return Desugared;
+ return QT->getUnqualifiedDesugaredType();
}
// Creates a field declaration of given name and type for HLSL buffer layout
@@ -2845,8 +2838,7 @@ void SemaHLSL::propagateContextualMatrixLayout(Expr *E, QualType DestType) {
if (!CallMat || CallMat->getNumRows() != DestMat->getNumRows() ||
CallMat->getNumColumns() != DestMat->getNumColumns())
return;
- // Re-type the call with the destination sugar so CodeGen lowers into that
- // layout, not the TU default.
+ // Re-type the call with the destination layout.
Call->setType(DestType.getUnqualifiedType());
}
diff --git a/clang/lib/Sema/SemaTemplateDeduction.cpp b/clang/lib/Sema/SemaTemplateDeduction.cpp
index 3c45806c47a6e..f0db5218ffe1e 100644
--- a/clang/lib/Sema/SemaTemplateDeduction.cpp
+++ b/clang/lib/Sema/SemaTemplateDeduction.cpp
@@ -3747,6 +3747,17 @@ CheckOriginalCallArgDeduction(Sema &S, TemplateDeductionInfo &Info,
DeducedA.getAtomicUnqualifiedType()))
return TemplateDeductionResult::Success;
+ if (S.getLangOpts().HLSL) {
+ const auto *AMatrix = A->getAs<ConstantMatrixType>();
+ const auto *DeducedMatrix = DeducedA->getAs<ConstantMatrixType>();
+ if (AMatrix && DeducedMatrix &&
+ AMatrix->getNumRows() == DeducedMatrix->getNumRows() &&
+ AMatrix->getNumColumns() == DeducedMatrix->getNumColumns() &&
+ Context.hasSameType(AMatrix->getElementType(),
+ DeducedMatrix->getElementType()))
+ return TemplateDeductionResult::Success;
+ }
+
// Strip off references on the argument types; they aren't needed for
// the following checks.
if (const ReferenceType *DeducedARef = DeducedA->getAs<ReferenceType>())
diff --git a/clang/lib/Sema/SemaType.cpp b/clang/lib/Sema/SemaType.cpp
index dc3564c8b17fd..c1a0eec829885 100644
--- a/clang/lib/Sema/SemaType.cpp
+++ b/clang/lib/Sema/SemaType.cpp
@@ -9136,8 +9136,15 @@ static void processTypeAttrs(TypeProcessingState &state, QualType &type,
case ParsedAttr::AT_HLSLRowMajor:
case ParsedAttr::AT_HLSLColumnMajor:
if (Attr *A =
- state.getSema().HLSL().buildMatrixLayoutTypeAttr(type, attr))
- type = state.getAttributedType(A, type, type);
+ state.getSema().HLSL().buildMatrixLayoutTypeAttr(type, attr)) {
+ MatrixType::LayoutKind Layout =
+ attr.getKind() == ParsedAttr::AT_HLSLRowMajor
+ ? MatrixType::LayoutKind::RowMajor
+ : MatrixType::LayoutKind::ColumnMajor;
+ QualType Equivalent =
+ state.getSema().Context.getMatrixTypeWithLayout(type, Layout);
+ type = state.getAttributedType(A, type, Equivalent);
+ }
attr.setUsedAsTypeAttr();
break;
OBJC_POINTER_TYPE_ATTRS_CASELIST:
diff --git a/clang/lib/Sema/TreeTransform.h b/clang/lib/Sema/TreeTransform.h
index 2083fcd372e81..0b96b943e925f 100644
--- a/clang/lib/Sema/TreeTransform.h
+++ b/clang/lib/Sema/TreeTransform.h
@@ -7721,6 +7721,17 @@ QualType TreeTransform<Derived>::TransformAttributedType(TypeLocBuilder &TLB,
return QualType();
}
+ if (SemaRef.getLangOpts().HLSL &&
+ (oldType->getAttrKind() == attr::HLSLRowMajor ||
+ oldType->getAttrKind() == attr::HLSLColumnMajor)) {
+ MatrixType::LayoutKind Layout =
+ oldType->getAttrKind() == attr::HLSLRowMajor
+ ? MatrixType::LayoutKind::RowMajor
+ : MatrixType::LayoutKind::ColumnMajor;
+ equivalentType =
+ SemaRef.Context.getMatrixTypeWithLayout(equivalentType, Layout);
+ }
+
// Check whether we can add nullability; it is only represented as
// type sugar, and therefore cannot be diagnosed in any other way.
if (auto nullability = oldType->getImmediateNullability()) {
diff --git a/clang/test/CodeGenHLSL/matrix-array-layout.hlsl b/clang/test/CodeGenHLSL/matrix-array-layout.hlsl
new file mode 100644
index 0000000000000..b2faba5483e52
--- /dev/null
+++ b/clang/test/CodeGenHLSL/matrix-array-layout.hlsl
@@ -0,0 +1,64 @@
+// RUN: %clang_cc1 -triple dxil-pc-shadermodel6.4-compute -x hlsl -emit-llvm -finclude-default-header -disable-llvm-passes -o - %s | FileCheck %s
+
+// StructuredBuffer is read-only and stores one handle per resource.
+// CHECK: type { target("dx.RawBuffer", [3 x <2 x float>], 0, 0) }
+
+// CHECK: type { target("dx.RawBuffer", [2 x <3 x float>], 1, 0), target("dx.RawBuffer", [2 x <3 x float>], 1, 0) }
+// CHECK: type { target("dx.RawBuffer", [3 x <2 x float>], 1, 1), target("dx.RawBuffer", [3 x <2 x float>], 1, 1) }
+
+// The array element layout matches the bare matrix layout for each orientation.
+// CHECK: %rm_arr = alloca [2 x [2 x <3 x float>]], align 4
+// CHECK: %cm_arr = alloca [2 x [3 x <2 x float>]], align 4
+// CHECK: %rm_bare = alloca [2 x <3 x float>], align 4
+// CHECK: %cm_bare = alloca [3 x <2 x float>], align 4
+// CHECK: %[[RM_VALUE:.*]] = load <6 x float>, ptr %rm_bare, align 4
+// CHECK: %[[RM_ELEMENT:.*]] = getelementptr inbounds [2 x [2 x <3 x float>]], ptr %rm_arr, i32 0, i32 0
+// CHECK: store <6 x float> %[[RM_VALUE]], ptr %[[RM_ELEMENT]], align 4
+// CHECK: %[[CM_VALUE:.*]] = load <6 x float>, ptr %cm_bare, align 4
+// CHECK: %[[CM_ELEMENT:.*]] = getelementptr inbounds [2 x [3 x <2 x float>]], ptr %cm_arr, i32 0, i32 0
+// CHECK: store <6 x float> %[[CM_VALUE]], ptr %[[CM_ELEMENT]], align 4
+
+export void f() {
+ row_major float2x3 rm_arr[2];
+ column_major float2x3 cm_arr[2];
+ row_major float2x3 rm_bare;
+ column_major float2x3 cm_bare;
+ rm_arr[0] = rm_bare;
+ cm_arr[0] = cm_bare;
+}
+
+StructuredBuffer<column_major float2x3> ColumnSource : register(t0);
+RWStructuredBuffer<row_major float2x3> RowDestination : register(u0);
+RasterizerOrderedStructuredBuffer<column_major float2x3> ColumnDestination
+ : register(u1);
+
+[numthreads(1,1,1)]
+void main() {
+ RowDestination[0] = ColumnSource[0];
+ ColumnDestination[0] = RowDestination[0];
+}
+
+// CHECK-LABEL: define internal void @_Z4mainv()
+// CHECK: %[[CM_PTR:.*]] = call {{.*}} ptr {{.*}}StructuredBuffer{{.*}}ColumnSource
+// CHECK: %[[CM_LOAD:.*]] = load <6 x float>, ptr %[[CM_PTR]], align 4
+// CHECK: %[[RM_PTR:.*]] = call {{.*}} ptr {{.*}}RWStructuredBuffer{{.*}}RowDestination
+// CHECK: store <6 x float> %[[CM_LOAD]], ptr %[[RM_PTR]], align 4
+// CHECK: %[[RM_SRC_PTR:.*]] = call {{.*}} ptr {{.*}}RWStructuredBuffer{{.*}}RowDestination
+// CHECK: %[[RM_LOAD:.*]] = load <6 x float>, ptr %[[RM_SRC_PTR]], align 4
+// CHECK: %[[CM_DST_PTR:.*]] = call {{.*}} ptr {{.*}}RasterizerOrderedStructuredBuffer{{.*}}ColumnDestination
+// CHECK: store <6 x float> %[[RM_LOAD]], ptr %[[CM_DST_PTR]], align 4
+
+// CHECK-LABEL: define linkonce_odr hidden {{.*}} ptr @_ZNK4hlsl16StructuredBuffer
+// CHECK: %[[CM_HANDLE_PTR:.*]] = getelementptr {{.*}}%"class.hlsl::StructuredBuffer", ptr {{.*}}, i32 0, i32 0
+// CHECK: %[[CM_HANDLE:.*]] = load target("dx.RawBuffer", [3 x <2 x float>], 0, 0), ptr %[[CM_HANDLE_PTR]], align 4
+// CHECK: call ptr @llvm.dx.resource.getpointer{{.*}}(target("dx.RawBuffer", [3 x <2 x float>], 0, 0) %[[CM_HANDLE]], i32 {{.*}})
+
+// CHECK-LABEL: define linkonce_odr hidden {{.*}} ptr @_ZNK4hlsl18RWStructuredBuffer
+// CHECK: %[[RM_HANDLE_PTR:.*]] = getelementptr {{.*}}%"class.hlsl::RWStructuredBuffer", ptr {{.*}}, i32 0, i32 0
+// CHECK: %[[RM_HANDLE:.*]] = load target("dx.RawBuffer", [2 x <3 x float>], 1, 0), ptr %[[RM_HANDLE_PTR]], align 4
+// CHECK: call ptr @llvm.dx.resource.getpointer{{.*}}(target("dx.RawBuffer", [2 x <3 x float>], 1, 0) %[[RM_HANDLE]], i32 {{.*}})
+
+// CHECK-LABEL: define linkonce_odr hidden {{.*}} ptr @_ZNK4hlsl33RasterizerOrderedStructuredBuffer
+// CHECK: %[[CM_DST_HANDLE_PTR:.*]] = getelementptr {{.*}}%"class.hlsl::RasterizerOrderedStructuredBuffer", ptr {{.*}}, i32 0, i32 0
+// CHECK: %[[CM_DST_HANDLE:.*]] = load target("dx.RawBuffer", [3 x <2 x float>], 1, 1), ptr %[[CM_DST_HANDLE_PTR]], align 4
+// CHECK: call ptr @llvm.dx.resource.getpointer{{.*}}(target("dx.RawBuffer", [3 x <2 x float>], 1, 1) %[[CM_DST_HANDLE]], i32 {{.*}})
>From eb9b2fbbdb8e4423f177db9be7cea3d9c0271971 Mon Sep 17 00:00:00 2001
From: Farzon Lotfi <farzonlotfi at microsoft.com>
Date: Tue, 11 Aug 2026 14:52:17 -0400
Subject: [PATCH 2/2] Add ICK_HLSL_Matrix_Layout typr conversions and type name
mangling for layout
---
clang/include/clang/Sema/Overload.h | 3 +
clang/lib/AST/ASTContext.cpp | 4 +-
clang/lib/AST/ItaniumMangle.cpp | 5 ++
clang/lib/AST/MicrosoftMangle.cpp | 3 +
clang/lib/Sema/SemaExpr.cpp | 15 ++++
clang/lib/Sema/SemaExprCXX.cpp | 20 +++++
clang/lib/Sema/SemaOverload.cpp | 12 ++-
.../test/CodeGenHLSL/matrix-array-layout.hlsl | 26 +++++-
.../matrix-layout-attr-overrides-default.hlsl | 86 ++++++++++++-------
clang/test/SemaHLSL/matrix_layout_attr.hlsl | 27 ++++++
10 files changed, 161 insertions(+), 40 deletions(-)
diff --git a/clang/include/clang/Sema/Overload.h b/clang/include/clang/Sema/Overload.h
index 1e412ff6fc9e2..e4a108404eed6 100644
--- a/clang/include/clang/Sema/Overload.h
+++ b/clang/include/clang/Sema/Overload.h
@@ -210,6 +210,9 @@ class Sema;
/// HLSL matrix splat from scalar or boolean type.
ICK_HLSL_Matrix_Splat,
+ /// HLSL conversion between matrix memory layouts.
+ ICK_HLSL_Matrix_Layout,
+
/// The number of conversion kinds
ICK_Num_Conversion_Kinds,
};
diff --git a/clang/lib/AST/ASTContext.cpp b/clang/lib/AST/ASTContext.cpp
index c21ad7b7e01b2..6dec8becedafe 100644
--- a/clang/lib/AST/ASTContext.cpp
+++ b/clang/lib/AST/ASTContext.cpp
@@ -14617,8 +14617,10 @@ static QualType getCommonNonSugarTypeNode(const ASTContext &Ctx, const Type *X,
*MY = cast<ConstantMatrixType>(Y);
assert(MX->getNumRows() == MY->getNumRows());
assert(MX->getNumColumns() == MY->getNumColumns());
+ assert(MX->getLayout() == MY->getLayout());
return Ctx.getConstantMatrixType(getCommonElementType(Ctx, MX, MY),
- MX->getNumRows(), MX->getNumColumns());
+ MX->getNumRows(), MX->getNumColumns(),
+ MX->getLayout());
}
case Type::DependentSizedMatrix: {
const auto *MX = cast<DependentSizedMatrixType>(X),
diff --git a/clang/lib/AST/ItaniumMangle.cpp b/clang/lib/AST/ItaniumMangle.cpp
index f8e6b898be250..49dc7043849ee 100644
--- a/clang/lib/AST/ItaniumMangle.cpp
+++ b/clang/lib/AST/ItaniumMangle.cpp
@@ -4406,6 +4406,11 @@ void CXXNameMangler::mangleType(const ConstantMatrixType *T) {
llvm::APSInt Columns(BitWidth);
Columns = T->getNumColumns();
mangleIntegerLiteral(ASTCtx.getSizeType(), Columns);
+ if (std::optional<MatrixType::LayoutKind> Layout = T->getLayout()) {
+ llvm::APSInt LayoutValue(BitWidth);
+ LayoutValue = static_cast<unsigned>(*Layout) + 1;
+ mangleIntegerLiteral(ASTCtx.getSizeType(), LayoutValue);
+ }
mangleType(T->getElementType());
Out << "E";
}
diff --git a/clang/lib/AST/MicrosoftMangle.cpp b/clang/lib/AST/MicrosoftMangle.cpp
index a1f2b671d6216..4ad552b9e7445 100644
--- a/clang/lib/AST/MicrosoftMangle.cpp
+++ b/clang/lib/AST/MicrosoftMangle.cpp
@@ -3766,6 +3766,9 @@ void MicrosoftCXXNameMangler::mangleType(const ConstantMatrixType *T,
Extra.mangleIntegerLiteral(llvm::APSInt::getUnsigned(T->getNumRows()));
Extra.mangleIntegerLiteral(llvm::APSInt::getUnsigned(T->getNumColumns()));
+ if (std::optional<MatrixType::LayoutKind> Layout = T->getLayout())
+ Extra.mangleIntegerLiteral(
+ llvm::APSInt::getUnsigned(static_cast<unsigned>(*Layout) + 1));
mangleArtificialTagType(TagTypeKind::Struct, TemplateMangling, {"__clang"});
}
diff --git a/clang/lib/Sema/SemaExpr.cpp b/clang/lib/Sema/SemaExpr.cpp
index 59b8c9b60663c..73134751814d0 100644
--- a/clang/lib/Sema/SemaExpr.cpp
+++ b/clang/lib/Sema/SemaExpr.cpp
@@ -13778,6 +13778,21 @@ QualType Sema::CheckMatrixElementwiseOperands(ExprResult &LHS, ExprResult &RHS,
if (Context.hasSameType(LHSType, RHSType))
return Context.getCommonSugaredType(LHSType, RHSType);
+ if (const auto *LHSConstantMat = dyn_cast_or_null<ConstantMatrixType>(
+ LHSMatType)) {
+ const auto *RHSConstantMat =
+ dyn_cast_or_null<ConstantMatrixType>(RHSMatType);
+ if (RHSConstantMat &&
+ LHSConstantMat->getNumRows() == RHSConstantMat->getNumRows() &&
+ LHSConstantMat->getNumColumns() == RHSConstantMat->getNumColumns() &&
+ Context.hasSameUnqualifiedType(LHSConstantMat->getElementType(),
+ RHSConstantMat->getElementType())) {
+ RHS = tryConvertExprToType(RHS.get(), LHSType);
+ if (!RHS.isInvalid())
+ return LHSType;
+ }
+ }
+
// Type conversion may change LHS/RHS. Keep copies to the original results, in
// case we have to return InvalidOperands.
ExprResult OriginalLHS = LHS;
diff --git a/clang/lib/Sema/SemaExprCXX.cpp b/clang/lib/Sema/SemaExprCXX.cpp
index 538604aa2e64b..f570010fcc325 100644
--- a/clang/lib/Sema/SemaExprCXX.cpp
+++ b/clang/lib/Sema/SemaExprCXX.cpp
@@ -5306,6 +5306,7 @@ Sema::PerformImplicitConversion(Expr *From, QualType ToType,
case ICK_HLSL_Matrix_Truncation:
case ICK_HLSL_Vector_Splat:
case ICK_HLSL_Matrix_Splat:
+ case ICK_HLSL_Matrix_Layout:
llvm_unreachable("Improper second standard conversion");
}
@@ -5363,6 +5364,11 @@ Sema::PerformImplicitConversion(Expr *From, QualType ToType,
.get();
break;
}
+ case ICK_HLSL_Matrix_Layout:
+ From = ImpCastExprToType(From, ToType, CK_HLSLMatrixTruncation,
+ From->getValueKind())
+ .get();
+ break;
case ICK_Identity:
default:
llvm_unreachable("Improper element standard conversion");
@@ -6011,6 +6017,20 @@ QualType Sema::CXXCheckConditionalOperands(ExprResult &Cond, ExprResult &LHS,
return QualType();
}
+ if (getLangOpts().HLSL) {
+ const auto *LMat = LTy->getAs<ConstantMatrixType>();
+ const auto *RMat = RTy->getAs<ConstantMatrixType>();
+ if (LMat && RMat && LMat->getNumRows() == RMat->getNumRows() &&
+ LMat->getNumColumns() == RMat->getNumColumns() &&
+ Context.hasSameUnqualifiedType(LMat->getElementType(),
+ RMat->getElementType())) {
+ RHS = tryConvertExprToType(RHS.get(), LTy);
+ if (RHS.isInvalid())
+ return QualType();
+ RTy = RHS.get()->getType();
+ }
+ }
+
// C++11 [expr.cond]p3
// Otherwise, if the second and third operand have different types, and
// either has (cv) class type [...] an attempt is made to convert each of
diff --git a/clang/lib/Sema/SemaOverload.cpp b/clang/lib/Sema/SemaOverload.cpp
index d66dea0d918fa..740cb974da1f6 100644
--- a/clang/lib/Sema/SemaOverload.cpp
+++ b/clang/lib/Sema/SemaOverload.cpp
@@ -173,6 +173,7 @@ ImplicitConversionRank clang::GetConversionRank(ImplicitConversionKind Kind) {
ICR_Conversion,
ICR_HLSL_Scalar_Widening,
ICR_HLSL_Scalar_Widening,
+ ICR_Exact_Match,
};
static_assert(std::size(Rank) == (int)ICK_Num_Conversion_Kinds);
return Rank[(int)Kind];
@@ -237,6 +238,7 @@ static const char *GetImplicitConversionName(ImplicitConversionKind Kind) {
"Non-decaying array conversion",
"HLSL vector splat",
"HLSL matrix splat",
+ "HLSL matrix layout conversion",
};
static_assert(std::size(Name) == (int)ICK_Num_Conversion_Kinds);
return Name[Kind];
@@ -2165,10 +2167,13 @@ static bool IsMatrixConversion(Sema &S, QualType FromType, QualType ToType,
if (FromRows < ToRows)
return false;
- if (FromRows == ToRows && FromCols == ToCols)
- ElConv = ICK_Identity;
- else
+ if (FromRows == ToRows && FromCols == ToCols) {
+ ElConv = FromMatrixType->getLayout() == ToMatrixType->getLayout()
+ ? ICK_Identity
+ : ICK_HLSL_Matrix_Layout;
+ } else {
ElConv = ICK_HLSL_Matrix_Truncation;
+ }
QualType FromElTy = FromMatrixType->getElementType();
QualType ToElTy = ToMatrixType->getElementType();
@@ -6453,6 +6458,7 @@ static bool CheckConvertedConstantConversions(Sema &S,
case ICK_RVV_Vector_Conversion:
case ICK_HLSL_Vector_Splat:
case ICK_HLSL_Matrix_Splat:
+ case ICK_HLSL_Matrix_Layout:
case ICK_Vector_Splat:
case ICK_Complex_Real:
case ICK_Block_Pointer_Conversion:
diff --git a/clang/test/CodeGenHLSL/matrix-array-layout.hlsl b/clang/test/CodeGenHLSL/matrix-array-layout.hlsl
index b2faba5483e52..65ce6789798c0 100644
--- a/clang/test/CodeGenHLSL/matrix-array-layout.hlsl
+++ b/clang/test/CodeGenHLSL/matrix-array-layout.hlsl
@@ -2,6 +2,7 @@
// StructuredBuffer is read-only and stores one handle per resource.
// CHECK: type { target("dx.RawBuffer", [3 x <2 x float>], 0, 0) }
+// CHECK: type { target("dx.RawBuffer", [2 x <3 x float>], 0, 0) }
// CHECK: type { target("dx.RawBuffer", [2 x <3 x float>], 1, 0), target("dx.RawBuffer", [2 x <3 x float>], 1, 0) }
// CHECK: type { target("dx.RawBuffer", [3 x <2 x float>], 1, 1), target("dx.RawBuffer", [3 x <2 x float>], 1, 1) }
@@ -27,7 +28,18 @@ export void f() {
cm_arr[0] = cm_bare;
}
+float use_default_layout(float2x3 M) { return M[0][0]; }
+
+export float call_default_layout(row_major float2x3 M) {
+ return use_default_layout(M);
+}
+
+// CHECK-LABEL: define {{.*}} float @_Z19call_default_layoutu11matrix_typeILm2ELm3ELm1EfE
+// CHECK: %[[CALL_LAYOUT:.*]] = shufflevector <6 x float> %{{.*}}, <6 x float> poison, <6 x i32> <i32 0, i32 3, i32 1, i32 4, i32 2, i32 5>
+// CHECK: call {{.*}} float @_Z18use_default_layoutu11matrix_typeILm2ELm3EfE(<6 x float> {{.*}}%[[CALL_LAYOUT]])
+
StructuredBuffer<column_major float2x3> ColumnSource : register(t0);
+StructuredBuffer<row_major float2x3> RowSource : register(t1);
RWStructuredBuffer<row_major float2x3> RowDestination : register(u0);
RasterizerOrderedStructuredBuffer<column_major float2x3> ColumnDestination
: register(u1);
@@ -36,19 +48,22 @@ RasterizerOrderedStructuredBuffer<column_major float2x3> ColumnDestination
void main() {
RowDestination[0] = ColumnSource[0];
ColumnDestination[0] = RowDestination[0];
+ RowDestination[1] = RowSource[0];
}
// CHECK-LABEL: define internal void @_Z4mainv()
// CHECK: %[[CM_PTR:.*]] = call {{.*}} ptr {{.*}}StructuredBuffer{{.*}}ColumnSource
// CHECK: %[[CM_LOAD:.*]] = load <6 x float>, ptr %[[CM_PTR]], align 4
+// CHECK: %[[CM_TO_RM:.*]] = shufflevector <6 x float> %[[CM_LOAD]], <6 x float> poison, <6 x i32> <i32 0, i32 2, i32 4, i32 1, i32 3, i32 5>
// CHECK: %[[RM_PTR:.*]] = call {{.*}} ptr {{.*}}RWStructuredBuffer{{.*}}RowDestination
-// CHECK: store <6 x float> %[[CM_LOAD]], ptr %[[RM_PTR]], align 4
+// CHECK: store <6 x float> %[[CM_TO_RM]], ptr %[[RM_PTR]], align 4
// CHECK: %[[RM_SRC_PTR:.*]] = call {{.*}} ptr {{.*}}RWStructuredBuffer{{.*}}RowDestination
// CHECK: %[[RM_LOAD:.*]] = load <6 x float>, ptr %[[RM_SRC_PTR]], align 4
+// CHECK: %[[RM_TO_CM:.*]] = shufflevector <6 x float> %[[RM_LOAD]], <6 x float> poison, <6 x i32> <i32 0, i32 3, i32 1, i32 4, i32 2, i32 5>
// CHECK: %[[CM_DST_PTR:.*]] = call {{.*}} ptr {{.*}}RasterizerOrderedStructuredBuffer{{.*}}ColumnDestination
-// CHECK: store <6 x float> %[[RM_LOAD]], ptr %[[CM_DST_PTR]], align 4
+// CHECK: store <6 x float> %[[RM_TO_CM]], ptr %[[CM_DST_PTR]], align 4
-// CHECK-LABEL: define linkonce_odr hidden {{.*}} ptr @_ZNK4hlsl16StructuredBuffer
+// CHECK-LABEL: define linkonce_odr hidden {{.*}} ptr @_ZNK4hlsl16StructuredBufferIu11matrix_typeILm2ELm3ELm2EfEEixEj
// CHECK: %[[CM_HANDLE_PTR:.*]] = getelementptr {{.*}}%"class.hlsl::StructuredBuffer", ptr {{.*}}, i32 0, i32 0
// CHECK: %[[CM_HANDLE:.*]] = load target("dx.RawBuffer", [3 x <2 x float>], 0, 0), ptr %[[CM_HANDLE_PTR]], align 4
// CHECK: call ptr @llvm.dx.resource.getpointer{{.*}}(target("dx.RawBuffer", [3 x <2 x float>], 0, 0) %[[CM_HANDLE]], i32 {{.*}})
@@ -62,3 +77,8 @@ void main() {
// CHECK: %[[CM_DST_HANDLE_PTR:.*]] = getelementptr {{.*}}%"class.hlsl::RasterizerOrderedStructuredBuffer", ptr {{.*}}, i32 0, i32 0
// CHECK: %[[CM_DST_HANDLE:.*]] = load target("dx.RawBuffer", [3 x <2 x float>], 1, 1), ptr %[[CM_DST_HANDLE_PTR]], align 4
// CHECK: call ptr @llvm.dx.resource.getpointer{{.*}}(target("dx.RawBuffer", [3 x <2 x float>], 1, 1) %[[CM_DST_HANDLE]], i32 {{.*}})
+
+// CHECK-LABEL: define linkonce_odr hidden {{.*}} ptr @_ZNK4hlsl16StructuredBufferIu11matrix_typeILm2ELm3ELm1EfEEixEj
+// CHECK: %[[RM_SOURCE_HANDLE_PTR:.*]] = getelementptr {{.*}}%"class.hlsl::StructuredBuffer{{(\.0)?}}", ptr {{.*}}, i32 0, i32 0
+// CHECK: %[[RM_SOURCE_HANDLE:.*]] = load target("dx.RawBuffer", [2 x <3 x float>], 0, 0), ptr %[[RM_SOURCE_HANDLE_PTR]], align 4
+// CHECK: call ptr @llvm.dx.resource.getpointer{{.*}}(target("dx.RawBuffer", [2 x <3 x float>], 0, 0) %[[RM_SOURCE_HANDLE]], i32 {{.*}})
diff --git a/clang/test/CodeGenHLSL/matrix-layout-attr-overrides-default.hlsl b/clang/test/CodeGenHLSL/matrix-layout-attr-overrides-default.hlsl
index dfafa2b0b7e61..a259d915657ca 100644
--- a/clang/test/CodeGenHLSL/matrix-layout-attr-overrides-default.hlsl
+++ b/clang/test/CodeGenHLSL/matrix-layout-attr-overrides-default.hlsl
@@ -19,7 +19,7 @@
export float subscript_rm(int row, int col, row_major float2x3 m) {
return m[row][col];
}
-// CHECK-LABEL: define {{.*}} float @_Z12subscript_rmiiu11matrix_typeILm2ELm3EfE
+// CHECK-LABEL: define {{.*}} float @_Z12subscript_rmiiu11matrix_typeILm2ELm3ELm1EfE
// CHECK: [[ROW:%.*]] = load i32, ptr %row.addr
// CHECK: [[COL:%.*]] = load i32, ptr %col.addr
// CHECK: [[OFFSET:%.*]] = mul i32 [[ROW]], 3
@@ -32,7 +32,7 @@ export float subscript_rm(int row, int col, row_major float2x3 m) {
export float subscript_cm(int row, int col, column_major float2x3 m) {
return m[row][col];
}
-// CHECK-LABEL: define {{.*}} float @_Z12subscript_cmiiu11matrix_typeILm2ELm3EfE
+// CHECK-LABEL: define {{.*}} float @_Z12subscript_cmiiu11matrix_typeILm2ELm3ELm2EfE
// CHECK: [[ROW:%.*]] = load i32, ptr %row.addr
// CHECK: [[COL:%.*]] = load i32, ptr %col.addr
// CHECK: [[OFFSET:%.*]] = mul i32 [[COL]], 2
@@ -49,7 +49,7 @@ export float subscript_cm(int row, int col, column_major float2x3 m) {
export float3 row_extract_rm(int row, row_major float2x3 m) {
return m[row];
}
-// CHECK-LABEL: define {{.*}} <3 x float> @_Z14row_extract_rmiu11matrix_typeILm2ELm3EfE
+// CHECK-LABEL: define {{.*}} <3 x float> @_Z14row_extract_rmiu11matrix_typeILm2ELm3ELm1EfE
// CHECK: [[ROW:%.*]] = load i32, ptr %row.addr
// CHECK: [[ROW_OFFSET0:%.*]] = mul i32 [[ROW]], 3
// CHECK: add i32 [[ROW_OFFSET0]], 0
@@ -64,7 +64,7 @@ export float3 row_extract_rm(int row, row_major float2x3 m) {
export float3 row_extract_cm(int row, column_major float2x3 m) {
return m[row];
}
-// CHECK-LABEL: define {{.*}} <3 x float> @_Z14row_extract_cmiu11matrix_typeILm2ELm3EfE
+// CHECK-LABEL: define {{.*}} <3 x float> @_Z14row_extract_cmiu11matrix_typeILm2ELm3ELm2EfE
// CHECK: [[ROW:%.*]] = load i32, ptr %row.addr
// CHECK: add i32 0, [[ROW]]
// CHECK: add i32 2, [[ROW]]
@@ -75,15 +75,19 @@ export float3 row_extract_cm(int row, column_major float2x3 m) {
// before the column-major matrix.multiply intrinsic.
// -----------------------------------------------------------------------------
export float3 vec_mat_rm(float2 v, row_major float2x3 m) { return mul(v, m); }
-// CHECK-LABEL: define {{.*}} <3 x float> @_Z10vec_mat_rmDv2_fu11matrix_typeILm2ELm3EfE
-// CHECK: [[T:%.*]] = call {{.*}} <6 x float> @llvm.matrix.transpose.v6f32(<6 x float> %{{.*}}, i32 3, i32 2)
+// CHECK-LABEL: define {{.*}} <3 x float> @_Z10vec_mat_rmDv2_fu11matrix_typeILm2ELm3ELm1EfE
+// COLMAJOR: [[T:%.*]] = shufflevector <6 x float> %{{.*}}, <6 x float> poison, <6 x i32> <i32 0, i32 3, i32 1, i32 4, i32 2, i32 5>
+// ROWMAJOR: [[S:%.*]] = shufflevector <6 x float> %{{.*}}, <6 x float> poison, <6 x i32> <i32 0, i32 1, i32 2, i32 3, i32 4, i32 5>
+// ROWMAJOR: [[T:%.*]] = call {{.*}} <6 x float> @llvm.matrix.transpose.v6f32(<6 x float> [[S]], i32 3, i32 2)
// CHECK: call {{.*}} <3 x float> @llvm.matrix.multiply.v3f32.v2f32.v6f32(<2 x float> %{{.*}}, <6 x float> [[T]], i32 1, i32 2, i32 3)
// Column-major operand: no transpose is inserted before matrix.multiply.
export float3 vec_mat_cm(float2 v, column_major float2x3 m) { return mul(v, m); }
-// CHECK-LABEL: define {{.*}} <3 x float> @_Z10vec_mat_cmDv2_fu11matrix_typeILm2ELm3EfE
-// CHECK-NOT: @llvm.matrix.transpose
-// CHECK: call {{.*}} <3 x float> @llvm.matrix.multiply.v3f32.v2f32.v6f32(<2 x float> %{{.*}}, <6 x float> %{{.*}}, i32 1, i32 2, i32 3)
+// CHECK-LABEL: define {{.*}} <3 x float> @_Z10vec_mat_cmDv2_fu11matrix_typeILm2ELm3ELm2EfE
+// COLMAJOR: [[T:%.*]] = shufflevector <6 x float> %{{.*}}, <6 x float> poison, <6 x i32> <i32 0, i32 1, i32 2, i32 3, i32 4, i32 5>
+// ROWMAJOR: [[S:%.*]] = shufflevector <6 x float> %{{.*}}, <6 x float> poison, <6 x i32> <i32 0, i32 2, i32 4, i32 1, i32 3, i32 5>
+// ROWMAJOR: [[T:%.*]] = call {{.*}} <6 x float> @llvm.matrix.transpose.v6f32(<6 x float> [[S]], i32 3, i32 2)
+// CHECK: call {{.*}} <3 x float> @llvm.matrix.multiply.v3f32.v2f32.v6f32(<2 x float> %{{.*}}, <6 x float> [[T]], i32 1, i32 2, i32 3)
// -----------------------------------------------------------------------------
// __builtin_hlsl_mul (matrix * matrix): mixed per-decl layouts cause a
@@ -94,17 +98,29 @@ export float3 vec_mat_cm(float2 v, column_major float2x3 m) { return mul(v, m);
export float2x2 mat_mat_rm_cm(row_major float2x3 a, column_major float3x2 b) { return mul(a, b); }
// CHECK-LABEL: define {{.*}} <4 x float> @_Z13mat_mat_rm_cm
// CHECK: [[AMat:%.*]] = load <6 x float>, ptr %a.addr, align 4
+// COLMAJOR: [[A:%.*]] = shufflevector <6 x float> [[AMat]], {{.*}} <i32 0, i32 3, i32 1, i32 4, i32 2, i32 5>
+// ROWMAJOR: [[A:%.*]] = shufflevector <6 x float> [[AMat]], {{.*}} <i32 0, i32 1, i32 2, i32 3, i32 4, i32 5>
// CHECK: [[BMat:%.*]] = load <6 x float>, ptr %b.addr, align 4
-// CHECK: [[T:%.*]] = call {{.*}} <6 x float> @llvm.matrix.transpose.v6f32(<6 x float> [[AMat]], i32 3, i32 2)
-// CHECK: call {{.*}} <4 x float> @llvm.matrix.multiply.v4f32.v6f32.v6f32(<6 x float> [[T]], <6 x float> [[BMat]], i32 2, i32 3, i32 2)
+// COLMAJOR: [[B:%.*]] = shufflevector <6 x float> [[BMat]], {{.*}} <i32 0, i32 1, i32 2, i32 3, i32 4, i32 5>
+// COLMAJOR: call {{.*}} @llvm.matrix.multiply{{.*}}(<6 x float> [[A]], <6 x float> [[B]],
+// ROWMAJOR: [[B:%.*]] = shufflevector <6 x float> [[BMat]], {{.*}} <i32 0, i32 3, i32 1, i32 4, i32 2, i32 5>
+// ROWMAJOR: [[AT:%.*]] = call {{.*}} @llvm.matrix.transpose{{.*}}(<6 x float> [[A]], i32 3, i32 2)
+// ROWMAJOR: [[BT:%.*]] = call {{.*}} @llvm.matrix.transpose{{.*}}(<6 x float> [[B]], i32 2, i32 3)
+// ROWMAJOR: call {{.*}} @llvm.matrix.multiply{{.*}}(<6 x float> [[AT]], <6 x float> [[BT]],
// LHS column-major, RHS row-major: only RHS is transposed.
export float2x2 mat_mat_cm_rm(column_major float2x3 a, row_major float3x2 b) { return mul(a, b); }
// CHECK-LABEL: define {{.*}} <4 x float> @_Z13mat_mat_cm_rm
// CHECK: [[AMat:%.*]] = load <6 x float>, ptr %a.addr, align 4
+// COLMAJOR: [[A:%.*]] = shufflevector <6 x float> [[AMat]], {{.*}} <i32 0, i32 1, i32 2, i32 3, i32 4, i32 5>
+// ROWMAJOR: [[A:%.*]] = shufflevector <6 x float> [[AMat]], {{.*}} <i32 0, i32 2, i32 4, i32 1, i32 3, i32 5>
// CHECK: [[BMat:%.*]] = load <6 x float>, ptr %b.addr, align 4
-// CHECK: [[T:%.*]] = call {{.*}} <6 x float> @llvm.matrix.transpose.v6f32(<6 x float> [[BMat]], i32 2, i32 3)
-// CHECK: call {{.*}} <4 x float> @llvm.matrix.multiply.v4f32.v6f32.v6f32(<6 x float> [[AMat]], <6 x float> [[T]], i32 2, i32 3, i32 2)
+// COLMAJOR: [[B:%.*]] = shufflevector <6 x float> [[BMat]], {{.*}} <i32 0, i32 2, i32 4, i32 1, i32 3, i32 5>
+// COLMAJOR: call {{.*}} @llvm.matrix.multiply{{.*}}(<6 x float> [[A]], <6 x float> [[B]],
+// ROWMAJOR: [[B:%.*]] = shufflevector <6 x float> [[BMat]], {{.*}} <i32 0, i32 1, i32 2, i32 3, i32 4, i32 5>
+// ROWMAJOR: [[AT:%.*]] = call {{.*}} @llvm.matrix.transpose{{.*}}(<6 x float> [[A]], i32 3, i32 2)
+// ROWMAJOR: [[BT:%.*]] = call {{.*}} @llvm.matrix.transpose{{.*}}(<6 x float> [[B]], i32 2, i32 3)
+// ROWMAJOR: call {{.*}} @llvm.matrix.multiply{{.*}}(<6 x float> [[AT]], <6 x float> [[BT]],
// Destination layout: the result is column-major, so no transpose is needed.
export column_major float2x2 mat_mat_dst_cm(column_major float2x3 a, column_major float3x2 b) { return mul(a, b); }
@@ -121,42 +137,46 @@ export row_major float2x2 mat_mat_dst_rm(column_major float2x3 a, column_major f
// Row-major source -> column-major destination: bits already transposed, no-op.
export column_major float3x2 transpose_rm_to_cm(row_major float2x3 m) { return transpose(m); }
-// CHECK-LABEL: define {{.*}} <6 x float> @_Z18transpose_rm_to_cmu11matrix_typeILm2ELm3EfE
-// CHECK-NOT: @llvm.matrix.transpose
-// CHECK: ret <6 x float>
+// CHECK-LABEL: define {{.*}} <6 x float> @_Z18transpose_rm_to_cmu11matrix_typeILm2ELm3ELm1EfE
+// COLMAJOR: [[S:%.*]] = shufflevector <6 x float> %{{.*}}, {{.*}} <i32 0, i32 3, i32 1, i32 4, i32 2, i32 5>
+// COLMAJOR: call {{.*}} @llvm.matrix.transpose{{.*}}(<6 x float> [[S]], i32 2, i32 3)
+// ROWMAJOR-NOT: @llvm.matrix.transpose
+// ROWMAJOR: ret <6 x float>
// Column-major source -> row-major destination: bits already transposed, no-op.
export row_major float3x2 transpose_cm_to_rm(column_major float2x3 m) { return transpose(m); }
-// CHECK-LABEL: define {{.*}} <6 x float> @_Z18transpose_cm_to_rmu11matrix_typeILm2ELm3EfE
-// CHECK-NOT: @llvm.matrix.transpose
-// CHECK: ret <6 x float>
+// CHECK-LABEL: define {{.*}} <6 x float> @_Z18transpose_cm_to_rmu11matrix_typeILm2ELm3ELm2EfE
+// COLMAJOR-NOT: @llvm.matrix.transpose
+// COLMAJOR: ret <6 x float>
+// ROWMAJOR: [[S:%.*]] = shufflevector <6 x float> %{{.*}}, {{.*}} <i32 0, i32 2, i32 4, i32 1, i32 3, i32 5>
+// ROWMAJOR: call {{.*}} @llvm.matrix.transpose{{.*}}(<6 x float> [[S]], i32 3, i32 2)
// Row-major source -> row-major destination: real transpose, dims swapped.
export row_major float3x2 transpose_rm_to_rm(row_major float2x3 m) { return transpose(m); }
-// CHECK-LABEL: define {{.*}} <6 x float> @_Z18transpose_rm_to_rmu11matrix_typeILm2ELm3EfE
-// CHECK: call {{.*}} <6 x float> @llvm.matrix.transpose.v6f32(<6 x float> %{{.*}}, i32 3, i32 2)
+// CHECK-LABEL: define {{.*}} <6 x float> @_Z18transpose_rm_to_rmu11matrix_typeILm2ELm3ELm1EfE
+// COLMAJOR-NOT: @llvm.matrix.transpose
+// ROWMAJOR: call {{.*}} <6 x float> @llvm.matrix.transpose.v6f32(<6 x float> %{{.*}}, i32 3, i32 2)
// Column-major source -> column-major destination: real transpose, natural dims.
export column_major float3x2 transpose_cm_to_cm(column_major float2x3 m) { return transpose(m); }
-// CHECK-LABEL: define {{.*}} <6 x float> @_Z18transpose_cm_to_cmu11matrix_typeILm2ELm3EfE
-// CHECK: call {{.*}} <6 x float> @llvm.matrix.transpose.v6f32(<6 x float> %{{.*}}, i32 2, i32 3)
+// CHECK-LABEL: define {{.*}} <6 x float> @_Z18transpose_cm_to_cmu11matrix_typeILm2ELm3ELm2EfE
+// COLMAJOR: call {{.*}} <6 x float> @llvm.matrix.transpose.v6f32(<6 x float> %{{.*}}, i32 2, i32 3)
+// ROWMAJOR-NOT: @llvm.matrix.transpose
// Default-layout return type: the TU `-fmatrix-memory-layout=` default
// flips between a real transpose and a no-op depending on the default.
export float3x2 transpose_rm(row_major float2x3 m) { return transpose(m); }
-// CHECK-LABEL: define {{.*}} <6 x float> @_Z12transpose_rmu11matrix_typeILm2ELm3EfE
-// COLMAJOR-NOT: @llvm.matrix.transpose
-// COLMAJOR: ret <6 x float>
+// CHECK-LABEL: define {{.*}} <6 x float> @_Z12transpose_rmu11matrix_typeILm2ELm3ELm1EfE
+// COLMAJOR: call {{.*}} <6 x float> @llvm.matrix.transpose.v6f32(<6 x float> %{{.*}}, i32 2, i32 3)
// ROWMAJOR: call {{.*}} <6 x float> @llvm.matrix.transpose.v6f32(<6 x float> %{{.*}}, i32 3, i32 2)
// column-major default: src/dst match -> real transpose, natural dims.
// row-major default: src/dst differ -> bits already transposed, no-op.
export float3x2 transpose_cm(column_major float2x3 m) { return transpose(m); }
-// CHECK-LABEL: define {{.*}} <6 x float> @_Z12transpose_cmu11matrix_typeILm2ELm3EfE
+// CHECK-LABEL: define {{.*}} <6 x float> @_Z12transpose_cmu11matrix_typeILm2ELm3ELm2EfE
// COLMAJOR: call {{.*}} <6 x float> @llvm.matrix.transpose.v6f32(<6 x float> %{{.*}}, i32 2, i32 3)
-// ROWMAJOR-NOT: @llvm.matrix.transpose
-// ROWMAJOR: ret <6 x float>
+// ROWMAJOR: call {{.*}} <6 x float> @llvm.matrix.transpose.v6f32(<6 x float> %{{.*}}, i32 3, i32 2)
// -----------------------------------------------------------------------------
// CK_HLSLMatrixTruncation: the shuffle mask that picks elements from the
@@ -170,12 +190,12 @@ typedef column_major float3x3 CM33;
// Row-major source 3x2 -> row-major dest 2x2: flat row-major mask is {0,1,2,3}.
export row_major float2x2 truncate_rm(row_major float3x2 m) { return (RM22)m; }
-// CHECK-LABEL: define {{.*}} <4 x float> @_Z11truncate_rmu11matrix_typeILm3ELm2EfE
+// CHECK-LABEL: define {{.*}} <4 x float> @_Z11truncate_rmu11matrix_typeILm3ELm2ELm1EfE
// CHECK: shufflevector <6 x float> %{{.*}}, <6 x float> poison, <4 x i32> <i32 0, i32 1, i32 2, i32 3>
// Column-major source 3x2 -> column-major dest 2x2: flat column-major mask is {0,1,3,4}.
export column_major float2x2 truncate_cm(column_major float3x2 m) { return (CM22)m; }
-// CHECK-LABEL: define {{.*}} <4 x float> @_Z11truncate_cmu11matrix_typeILm3ELm2EfE
+// CHECK-LABEL: define {{.*}} <4 x float> @_Z11truncate_cmu11matrix_typeILm3ELm2ELm2EfE
// CHECK: shufflevector <6 x float> %{{.*}}, <6 x float> poison, <4 x i32> <i32 0, i32 1, i32 3, i32 4>
// -----------------------------------------------------------------------------
@@ -192,7 +212,7 @@ export column_major float2x2 truncate_cm(column_major float3x2 m) { return (CM22
// (1,0)->mask[1]=4 (1,1)->mask[4]=5 (1,2)->mask[7]=6
// (2,0)->mask[2]=8 (2,1)->mask[5]=9 (2,2)->mask[8]=10
export column_major float3x3 truncate_rm_to_cm(row_major float3x4 m) { return (CM33)m; }
-// CHECK-LABEL: define {{.*}} <9 x float> @_Z17truncate_rm_to_cmu11matrix_typeILm3ELm4EfE
+// CHECK-LABEL: define {{.*}} <9 x float> @_Z17truncate_rm_to_cmu11matrix_typeILm3ELm4ELm1EfE
// CHECK: shufflevector <12 x float> %{{.*}}, <12 x float> poison, <9 x i32> <i32 0, i32 4, i32 8, i32 1, i32 5, i32 9, i32 2, i32 6, i32 10>
// Column-major src 3x4 -> row-major dst 3x3.
@@ -201,7 +221,7 @@ export column_major float3x3 truncate_rm_to_cm(row_major float3x4 m) { return (C
// (1,0)->mask[3]=1 (1,1)->mask[4]=4 (1,2)->mask[5]=7
// (2,0)->mask[6]=2 (2,1)->mask[7]=5 (2,2)->mask[8]=8
export row_major float3x3 truncate_cm_to_rm(column_major float3x4 m) { return (RM33)m; }
-// CHECK-LABEL: define {{.*}} <9 x float> @_Z17truncate_cm_to_rmu11matrix_typeILm3ELm4EfE
+// CHECK-LABEL: define {{.*}} <9 x float> @_Z17truncate_cm_to_rmu11matrix_typeILm3ELm4ELm2EfE
// CHECK: shufflevector <12 x float> %{{.*}}, <12 x float> poison, <9 x i32> <i32 0, i32 3, i32 6, i32 1, i32 4, i32 7, i32 2, i32 5, i32 8>
// -----------------------------------------------------------------------------
diff --git a/clang/test/SemaHLSL/matrix_layout_attr.hlsl b/clang/test/SemaHLSL/matrix_layout_attr.hlsl
index 3e953f07557e6..a1a2176bc1dbf 100644
--- a/clang/test/SemaHLSL/matrix_layout_attr.hlsl
+++ b/clang/test/SemaHLSL/matrix_layout_attr.hlsl
@@ -44,6 +44,33 @@ column_major float4x4 Col2Row(row_major float4x4 M) {
void bar(row_major float4x4 M, column_major float4x4 M2) {}
+void takes_default_layout(float2x2 M) {}
+
+float2x2 layout_conversions(row_major float2x2 RM,
+ column_major float2x2 CM, bool SelectRM) {
+ takes_default_layout(RM);
+ takes_default_layout(CM);
+
+ float2x2 Result = RM;
+ Result = CM;
+ Result += RM;
+ Result -= CM;
+ Result *= RM;
+
+ float2x2 Sum = RM + CM;
+ float2x2 Difference = RM - CM;
+ float2x2 Product = RM * CM;
+ float2x2 MulProduct = mul(RM, CM);
+ return SelectRM ? RM : CM;
+}
+
+typedef float FLOAT;
+
+float2x2 common_matrix_element_sugar(row_major matrix<FLOAT, 2, 2> A,
+ row_major matrix<float, 2, 2> B) {
+ return A + B;
+}
+
//Invalid:
// expected-error at +1 {{'row_major' attribute can only be applied to a matrix type}}
void foo(column_major float4x4 mat, row_major int i) {}
More information about the cfe-commits
mailing list