[Mlir-commits] [mlir] [mlir][AMDGPU] Extend amdgpu.transpose_load for gfx1250 (PR #198354)
Krzysztof Drewniak
llvmlistbot at llvm.org
Mon May 18 10:48:37 PDT 2026
https://github.com/krzysz00 created https://github.com/llvm/llvm-project/pull/198354
This commit adds support for gfx1250's ds_load_tr* instructions to `amdgpu.transpose_load` since they're pretty close to the gfx950 ones.
>From bafbabf6b5dda7e1dc4f106e8e83a7e7e05fe3db Mon Sep 17 00:00:00 2001
From: Krzysztof Drewniak <Krzysztof.Drewniak at amd.com>
Date: Mon, 18 May 2026 17:44:53 +0000
Subject: [PATCH] [mlir][AMDGPU] Extend amdgpu.transpose_load for gfx1250
This commit adds support for gfx1250's ds_load_tr* instructions to
`amdgpu.transpose_load` since they're pretty close to the gfx950 ones.
Co-Authored-By: Codex <codex at openai.com>
---
.../mlir/Dialect/AMDGPU/IR/AMDGPUOps.td | 16 ++-
.../AMDGPUToROCDL/AMDGPUToROCDL.cpp | 132 +++++++++++++-----
mlir/lib/Dialect/AMDGPU/IR/AMDGPUOps.cpp | 33 +++--
.../AMDGPUToROCDL/transpose_load.mlir | 10 +-
.../AMDGPUToROCDL/transpose_load_gfx1250.mlir | 41 ++++++
.../transpose_load_gfx1250_invalid.mlir | 8 ++
.../transpose_load_gfx950_invalid.mlir | 8 ++
mlir/test/Dialect/AMDGPU/invalid.mlir | 2 +-
8 files changed, 190 insertions(+), 60 deletions(-)
create mode 100644 mlir/test/Conversion/AMDGPUToROCDL/transpose_load_gfx1250.mlir
create mode 100644 mlir/test/Conversion/AMDGPUToROCDL/transpose_load_gfx1250_invalid.mlir
create mode 100644 mlir/test/Conversion/AMDGPUToROCDL/transpose_load_gfx950_invalid.mlir
diff --git a/mlir/include/mlir/Dialect/AMDGPU/IR/AMDGPUOps.td b/mlir/include/mlir/Dialect/AMDGPU/IR/AMDGPUOps.td
index 7d33ca163fb2f..c41503d6dfc64 100644
--- a/mlir/include/mlir/Dialect/AMDGPU/IR/AMDGPUOps.td
+++ b/mlir/include/mlir/Dialect/AMDGPU/IR/AMDGPUOps.td
@@ -1110,7 +1110,7 @@ def AMDGPU_DotOp :
let summary = "MLIR wrapper for AMDGPU v_dot* intrinsics";
let description = [{
The `amdgpu.dot` op is an MLIR wrapper over the `v_dot*` family of intrinsics,
- which compute `D = sum_i A[i] * B[i] + C`.
+ which compute `D = sum_i A[i] * B[i] + C`.
Variants (source, dest, signedness, chipset -> intrinsic).
@@ -1420,9 +1420,11 @@ def AMDGPU_TransposeLoadOp :
AMDGPU_Op<"transpose_load", [SameVariadicOperandSize]>,
Arguments<(ins Arg<AnyMemRef, "buffer to transpose load from", [MemRead]>:$src, Variadic<Index>:$srcIndices)>,
Results<(outs AnyTypeOf<[AnyVectorOfNonZeroRank]>:$result)> {
- let summary = "MLIR wrapper for CDNA Transpose Load instructions";
+ let summary = "MLIR wrapper for CDNA transpose Load instructions";
let description = [{
- The `amdgpu.transpose_load` op is a wrapper around the `ds_read_tr` instructions.
+ The `amdgpu.transpose_load` op is a wrapper around the `ds_read_tr` instructions
+ on gfx9 and the `ds_load_tr` family of instructions on gfx1250.
+
The transpose load op represents a subgroup load from LDS memory,
where the subgroup of threads collectively reads a matrix from the source
memref, with each thread reading a vector of the matrix, and gets a transposed matrix
@@ -1430,8 +1432,9 @@ def AMDGPU_TransposeLoadOp :
indices, and the thread's read result is a vector of the corresponding row of the transposed
matrix.
- This op is a direct wrapper around the ROCDL `ds_read_tr` family intrinsics. Please refer
- to the CDNA4 ISA documentation for more details about its exact semantics.
+ This op is a direct wrapper around the ROCDL `ds_read_tr` family intrinsics on
+ CDNA4 and the `ds_load_tr` family of instructions on CDNA5. Please refer
+ to the respective ISA documentation for more details about its exact semantics.
Format example:
```
@@ -1442,7 +1445,8 @@ def AMDGPU_TransposeLoadOp :
* `$srcIndices`: indices into `$src` to read from for this thread.
* `$result`: target register this transpose load instruction will write to.
- Note: Lowering is only supported on gfx950 and up.
+ Note: Lowering is only supported on gfx950 and gfx1250, with different
+ permitted load types.
}];
let assemblyFormat = [{
$src `[` $srcIndices `]` attr-dict `:` type($src) `->` type($result)
diff --git a/mlir/lib/Conversion/AMDGPUToROCDL/AMDGPUToROCDL.cpp b/mlir/lib/Conversion/AMDGPUToROCDL/AMDGPUToROCDL.cpp
index 8b9b3e194adf9..541922e634957 100644
--- a/mlir/lib/Conversion/AMDGPUToROCDL/AMDGPUToROCDL.cpp
+++ b/mlir/lib/Conversion/AMDGPUToROCDL/AMDGPUToROCDL.cpp
@@ -2167,8 +2167,9 @@ struct TransposeLoadOpLowering
LogicalResult
matchAndRewrite(TransposeLoadOp op, TransposeLoadOpAdaptor adaptor,
ConversionPatternRewriter &rewriter) const override {
- if (chipset != kGfx950)
- return op.emitOpError("Non-gfx950 chipset not supported");
+ if (chipset != kGfx950 && chipset < kGfx1250)
+ return op.emitOpError(
+ "transpose_load is only supported on gfx950 and gfx1250+");
Location loc = op.getLoc();
auto srcMemRefType = cast<MemRefType>(op.getSrc().getType());
@@ -2191,43 +2192,104 @@ struct TransposeLoadOpLowering
size_t elementTypeSize =
resultType.getElementType().getIntOrFloatBitWidth();
- // ROCDL transpose load intrinsics return vectors of 32-bit integers, if
- // the element size is smaller than 16 bits.
- Type rocdlResultType = VectorType::get((numElements * elementTypeSize) / 32,
- rewriter.getIntegerType(32));
Type llvmResultType = typeConverter->convertType(resultType);
+ // ROCDL transpose load intrinsics return vectors of 32-bit integers for
+ // sub-16-bit element types, and otherwise return the converted result type.
+ Type rocdlResultType =
+ elementTypeSize < 16
+ ? VectorType::get((numElements * elementTypeSize) / 32,
+ rewriter.getIntegerType(32))
+ : llvmResultType;
- switch (elementTypeSize) {
- case 4: {
- assert(numElements == 16);
- auto rocdlOp = ROCDL::ds_read_tr4_b64::create(rewriter, loc,
- rocdlResultType, srcPtr);
- rewriter.replaceOpWithNewOp<LLVM::BitcastOp>(op, llvmResultType, rocdlOp);
- break;
- }
- case 6: {
- assert(numElements == 16);
- auto rocdlOp = ROCDL::ds_read_tr6_b96::create(rewriter, loc,
- rocdlResultType, srcPtr);
- rewriter.replaceOpWithNewOp<LLVM::BitcastOp>(op, llvmResultType, rocdlOp);
- break;
- }
- case 8: {
- assert(numElements == 8);
- auto rocdlOp = ROCDL::ds_read_tr8_b64::create(rewriter, loc,
- rocdlResultType, srcPtr);
- rewriter.replaceOpWithNewOp<LLVM::BitcastOp>(op, llvmResultType, rocdlOp);
- break;
- }
- case 16: {
- assert(numElements == 4);
- rewriter.replaceOpWithNewOp<ROCDL::ds_read_tr16_b64>(op, llvmResultType,
- srcPtr);
- break;
+ auto emitNumElementsError = [&](size_t expected, StringRef chipsetName) {
+ return op.emitOpError()
+ << elementTypeSize << "-bit transpose_load requires " << expected
+ << " elements on " << chipsetName;
+ };
+
+ Value intrinsic;
+ if (chipset >= kGfx1250) {
+ switch (elementTypeSize) {
+ case 4: {
+ if (numElements != 16)
+ return emitNumElementsError(16, "gfx1250+");
+ intrinsic =
+ ROCDL::DsLoadTr4_B64::create(rewriter, loc, rocdlResultType, srcPtr)
+ .getResult();
+ break;
+ }
+ case 6: {
+ if (numElements != 16)
+ return emitNumElementsError(16, "gfx1250+");
+ intrinsic =
+ ROCDL::DsLoadTr6_B96::create(rewriter, loc, rocdlResultType, srcPtr)
+ .getResult();
+ break;
+ }
+ case 8: {
+ if (numElements != 8)
+ return emitNumElementsError(8, "gfx1250+");
+ intrinsic =
+ ROCDL::DsLoadTr8_B64::create(rewriter, loc, rocdlResultType, srcPtr)
+ .getResult();
+ break;
+ }
+ case 16: {
+ if (numElements != 8)
+ return emitNumElementsError(8, "gfx1250+");
+ intrinsic = ROCDL::DsLoadTr16_B128::create(rewriter, loc,
+ rocdlResultType, srcPtr)
+ .getResult();
+ break;
+ }
+ default:
+ return op.emitOpError("Unsupported element size for transpose load");
+ }
+ } else {
+ switch (elementTypeSize) {
+ case 4: {
+ if (numElements != 16)
+ return emitNumElementsError(16, "gfx950");
+ intrinsic = ROCDL::ds_read_tr4_b64::create(rewriter, loc,
+ rocdlResultType, srcPtr)
+ .getResult();
+ break;
+ }
+ case 6: {
+ if (numElements != 16)
+ return emitNumElementsError(16, "gfx950");
+ intrinsic = ROCDL::ds_read_tr6_b96::create(rewriter, loc,
+ rocdlResultType, srcPtr)
+ .getResult();
+ break;
+ }
+ case 8: {
+ if (numElements != 8)
+ return emitNumElementsError(8, "gfx950");
+ intrinsic = ROCDL::ds_read_tr8_b64::create(rewriter, loc,
+ rocdlResultType, srcPtr)
+ .getResult();
+ break;
+ }
+ case 16: {
+ if (numElements != 4)
+ return emitNumElementsError(4, "gfx950");
+ intrinsic = ROCDL::ds_read_tr16_b64::create(rewriter, loc,
+ rocdlResultType, srcPtr)
+ .getResult();
+ break;
+ }
+ default:
+ return op.emitOpError("Unsupported element size for transpose load");
+ }
}
- default:
- return op.emitOpError("Unsupported element size for transpose load");
+
+ assert(intrinsic && "expected ROCDL transpose load intrinsic");
+ if (intrinsic.getType() == llvmResultType) {
+ rewriter.replaceOp(op, intrinsic);
+ return success();
}
+ rewriter.replaceOpWithNewOp<LLVM::BitcastOp>(op, llvmResultType, intrinsic);
return success();
}
};
diff --git a/mlir/lib/Dialect/AMDGPU/IR/AMDGPUOps.cpp b/mlir/lib/Dialect/AMDGPU/IR/AMDGPUOps.cpp
index 209d52ec7a1c8..652720a729d01 100644
--- a/mlir/lib/Dialect/AMDGPU/IR/AMDGPUOps.cpp
+++ b/mlir/lib/Dialect/AMDGPU/IR/AMDGPUOps.cpp
@@ -1109,23 +1109,30 @@ LogicalResult TransposeLoadOp::verify() {
size_t elementTypeSize =
transferType.getElementType().getIntOrFloatBitWidth();
- // ElementSize -> NumElements
- const llvm::SmallDenseMap<size_t, size_t> kValidLoadSizeMap = {
- {4, 16},
- {6, 16},
- {8, 8},
- {16, 4},
+ auto emitNumElementsError = [&](StringRef expected) {
+ return emitOpError(
+ "Transferring type size mismatch: expected num of elements: ")
+ << expected;
};
- auto validNumElems = kValidLoadSizeMap.find(elementTypeSize);
- if (validNumElems == kValidLoadSizeMap.end())
+ switch (elementTypeSize) {
+ case 4:
+ case 6:
+ if (numElements != 16)
+ return emitNumElementsError("16");
+ break;
+ case 8:
+ if (numElements != 8)
+ return emitNumElementsError("8");
+ break;
+ case 16:
+ if (numElements != 4 && numElements != 8)
+ return emitNumElementsError("4 or 8");
+ break;
+ default:
return emitOpError("Unsupported element type size for transpose load: ")
<< elementTypeSize << " bits";
-
- if (numElements != validNumElems->second)
- return emitOpError(
- "Transferring type size mismatch: expected num of elements: ")
- << validNumElems->second;
+ }
return success();
}
diff --git a/mlir/test/Conversion/AMDGPUToROCDL/transpose_load.mlir b/mlir/test/Conversion/AMDGPUToROCDL/transpose_load.mlir
index 73481e4bced20..dcc6624cdb37b 100644
--- a/mlir/test/Conversion/AMDGPUToROCDL/transpose_load.mlir
+++ b/mlir/test/Conversion/AMDGPUToROCDL/transpose_load.mlir
@@ -4,7 +4,7 @@
// CHECK-LABEL: func @transpose_load_to_rocdl_4xf16
func.func @transpose_load_to_rocdl_4xf16(%idx1 : index, %idx2 : index, %wgmem : memref<128x72xf16, 3>) -> vector<4xf16> {
// CHECK: rocdl.ds.read.tr16.b64
- // CHECK-OLD: error: 'amdgpu.transpose_load' op Non-gfx950 chipset not supported
+ // CHECK-OLD: error: 'amdgpu.transpose_load' op transpose_load is only supported on gfx950 and gfx1250+
%0 = amdgpu.transpose_load %wgmem[%idx1, %idx2] : memref<128x72xf16, 3> -> vector<4xf16>
return %0 : vector<4xf16>
}
@@ -16,7 +16,7 @@ func.func @transpose_load_to_rocdl_8xi8(%idx1 : index, %idx2 : index, %wgmem : m
// CHECK: %[[RES:.*]] = rocdl.ds.read.tr8.b64
// CHECK-SAME: -> vector<2xi32>
// CHECK-NEXT: llvm.bitcast %[[RES]] : vector<2xi32> to vector<8xi8>
- // CHECK-OLD: error: 'amdgpu.transpose_load' op Non-gfx950 chipset not supported
+ // CHECK-OLD: error: 'amdgpu.transpose_load' op transpose_load is only supported on gfx950 and gfx1250+
%0 = amdgpu.transpose_load %wgmem[%idx1, %idx2] : memref<128x128xi8, 3> -> vector<8xi8>
return %0 : vector<8xi8>
}
@@ -28,7 +28,7 @@ func.func @transpose_load_to_rocdl_i4_memrefxi8(%idx1 : index, %idx2 : index, %w
// CHECK: %[[RES:.*]] = rocdl.ds.read.tr4.b64
// CHECK-SAME: -> vector<2xi32>
// CHECK-NEXT: llvm.bitcast %[[RES]] : vector<2xi32> to vector<16xi4>
- // CHECK-OLD: error: 'amdgpu.transpose_load' op Non-gfx950 chipset not supported
+ // CHECK-OLD: error: 'amdgpu.transpose_load' op transpose_load is only supported on gfx950 and gfx1250+
%0 = amdgpu.transpose_load %wgmem[%idx1, %idx2] : memref<128x32xi8, 3> -> vector<16xi4>
return %0 : vector<16xi4>
}
@@ -40,7 +40,7 @@ func.func @transpose_load_to_rocdl_i6_memrefxi8(%idx1 : index, %idx2 : index, %w
// CHECK: %[[RES:.*]] = rocdl.ds.read.tr6.b96
// CHECK-SAME: -> vector<3xi32>
// CHECK-NEXT: llvm.bitcast %[[RES]] : vector<3xi32> to vector<16xi6>
- // CHECK-OLD: error: 'amdgpu.transpose_load' op Non-gfx950 chipset not supported
+ // CHECK-OLD: error: 'amdgpu.transpose_load' op transpose_load is only supported on gfx950 and gfx1250+
%0 = amdgpu.transpose_load %wgmem[%idx1, %idx2] : memref<128x32xi8, 3> -> vector<16xi6>
return %0 : vector<16xi6>
}
@@ -50,7 +50,7 @@ func.func @transpose_load_to_rocdl_i6_memrefxi8(%idx1 : index, %idx2 : index, %w
// CHECK-LABEL: func @transpose_load_to_rocdl_i16_memrefxi8
func.func @transpose_load_to_rocdl_i16_memrefxi8(%idx1 : index, %idx2 : index, %wgmem : memref<128x32xi8, 3>) -> vector<4xi16> {
// CHECK: rocdl.ds.read.tr16.b64
- // CHECK-OLD: error: 'amdgpu.transpose_load' op Non-gfx950 chipset not supported
+ // CHECK-OLD: error: 'amdgpu.transpose_load' op transpose_load is only supported on gfx950 and gfx1250+
%0 = amdgpu.transpose_load %wgmem[%idx1, %idx2] : memref<128x32xi8, 3> -> vector<4xi16>
return %0 : vector<4xi16>
}
diff --git a/mlir/test/Conversion/AMDGPUToROCDL/transpose_load_gfx1250.mlir b/mlir/test/Conversion/AMDGPUToROCDL/transpose_load_gfx1250.mlir
new file mode 100644
index 0000000000000..98ce6b7ea3001
--- /dev/null
+++ b/mlir/test/Conversion/AMDGPUToROCDL/transpose_load_gfx1250.mlir
@@ -0,0 +1,41 @@
+// RUN: mlir-opt %s --split-input-file -convert-amdgpu-to-rocdl=chipset=gfx1250 | FileCheck %s
+
+// CHECK-LABEL: func @transpose_load_to_rocdl_8xf16
+func.func @transpose_load_to_rocdl_8xf16(%idx1 : index, %idx2 : index, %wgmem : memref<128x72xf16, 3>) -> vector<8xf16> {
+ // CHECK: rocdl.ds.load.tr16.b128
+ %0 = amdgpu.transpose_load %wgmem[%idx1, %idx2] : memref<128x72xf16, 3> -> vector<8xf16>
+ return %0 : vector<8xf16>
+}
+
+// -----
+
+// CHECK-LABEL: func @transpose_load_to_rocdl_8xi8
+func.func @transpose_load_to_rocdl_8xi8(%idx1 : index, %idx2 : index, %wgmem : memref<128x128xi8, 3>) -> vector<8xi8> {
+ // CHECK: %[[RES:.*]] = rocdl.ds.load.tr8.b64
+ // CHECK-SAME: -> vector<2xi32>
+ // CHECK-NEXT: llvm.bitcast %[[RES]] : vector<2xi32> to vector<8xi8>
+ %0 = amdgpu.transpose_load %wgmem[%idx1, %idx2] : memref<128x128xi8, 3> -> vector<8xi8>
+ return %0 : vector<8xi8>
+}
+
+// -----
+
+// CHECK-LABEL: func @transpose_load_to_rocdl_16xi4
+func.func @transpose_load_to_rocdl_16xi4(%idx1 : index, %idx2 : index, %wgmem : memref<128x32xi8, 3>) -> vector<16xi4> {
+ // CHECK: %[[RES:.*]] = rocdl.ds.load.tr4.b64
+ // CHECK-SAME: -> vector<2xi32>
+ // CHECK-NEXT: llvm.bitcast %[[RES]] : vector<2xi32> to vector<16xi4>
+ %0 = amdgpu.transpose_load %wgmem[%idx1, %idx2] : memref<128x32xi8, 3> -> vector<16xi4>
+ return %0 : vector<16xi4>
+}
+
+// -----
+
+// CHECK-LABEL: func @transpose_load_to_rocdl_16xi6
+func.func @transpose_load_to_rocdl_16xi6(%idx1 : index, %idx2 : index, %wgmem : memref<128x32xi8, 3>) -> vector<16xi6> {
+ // CHECK: %[[RES:.*]] = rocdl.ds.load.tr6.b96
+ // CHECK-SAME: -> vector<3xi32>
+ // CHECK-NEXT: llvm.bitcast %[[RES]] : vector<3xi32> to vector<16xi6>
+ %0 = amdgpu.transpose_load %wgmem[%idx1, %idx2] : memref<128x32xi8, 3> -> vector<16xi6>
+ return %0 : vector<16xi6>
+}
diff --git a/mlir/test/Conversion/AMDGPUToROCDL/transpose_load_gfx1250_invalid.mlir b/mlir/test/Conversion/AMDGPUToROCDL/transpose_load_gfx1250_invalid.mlir
new file mode 100644
index 0000000000000..61acdfe245c7c
--- /dev/null
+++ b/mlir/test/Conversion/AMDGPUToROCDL/transpose_load_gfx1250_invalid.mlir
@@ -0,0 +1,8 @@
+// RUN: mlir-opt %s --split-input-file --verify-diagnostics -convert-amdgpu-to-rocdl=chipset=gfx1250
+
+func.func @transpose_load_to_rocdl_4xf16(%idx1 : index, %idx2 : index, %wgmem : memref<128x72xf16, 3>) -> vector<4xf16> {
+ // expected-error at +2 {{'amdgpu.transpose_load' op 16-bit transpose_load requires 8 elements on gfx1250+}}
+ // expected-error at +1 {{failed to legalize operation 'amdgpu.transpose_load'}}
+ %0 = amdgpu.transpose_load %wgmem[%idx1, %idx2] : memref<128x72xf16, 3> -> vector<4xf16>
+ return %0 : vector<4xf16>
+}
diff --git a/mlir/test/Conversion/AMDGPUToROCDL/transpose_load_gfx950_invalid.mlir b/mlir/test/Conversion/AMDGPUToROCDL/transpose_load_gfx950_invalid.mlir
new file mode 100644
index 0000000000000..682f989ede83b
--- /dev/null
+++ b/mlir/test/Conversion/AMDGPUToROCDL/transpose_load_gfx950_invalid.mlir
@@ -0,0 +1,8 @@
+// RUN: mlir-opt %s --split-input-file --verify-diagnostics -convert-amdgpu-to-rocdl=chipset=gfx950
+
+func.func @transpose_load_to_rocdl_8xf16(%idx1 : index, %idx2 : index, %wgmem : memref<128x72xf16, 3>) -> vector<8xf16> {
+ // expected-error at +2 {{'amdgpu.transpose_load' op 16-bit transpose_load requires 4 elements on gfx950}}
+ // expected-error at +1 {{failed to legalize operation 'amdgpu.transpose_load'}}
+ %0 = amdgpu.transpose_load %wgmem[%idx1, %idx2] : memref<128x72xf16, 3> -> vector<8xf16>
+ return %0 : vector<8xf16>
+}
diff --git a/mlir/test/Dialect/AMDGPU/invalid.mlir b/mlir/test/Dialect/AMDGPU/invalid.mlir
index 4e4cfe53298c7..08e943492cf16 100644
--- a/mlir/test/Dialect/AMDGPU/invalid.mlir
+++ b/mlir/test/Dialect/AMDGPU/invalid.mlir
@@ -294,7 +294,7 @@ func.func @transpose_load_elem_f32(%idx1 : index, %idx2 : index, %mem : memref<1
// -----
func.func @transpose_load_vector_size_f16(%idx1 : index, %idx2 : index, %mem : memref<128x32xf16, 3>) -> vector<2xf16> {
- // expected-error at +1 {{'amdgpu.transpose_load' op Transferring type size mismatch: expected num of elements: 4}}
+ // expected-error at +1 {{'amdgpu.transpose_load' op Transferring type size mismatch: expected num of elements: 4 or 8}}
%0 = amdgpu.transpose_load %mem[%idx1, %idx2] : memref<128x32xf16, 3> -> vector<2xf16>
func.return %0 : vector<2xf16>
}
More information about the Mlir-commits
mailing list