MLIR 24.0.0git
AMDGPUToROCDL.cpp
Go to the documentation of this file.
1//===- AMDGPUToROCDL.cpp - AMDGPU to ROCDL dialect conversion -------===//
2//
3// Part of the LLVM Project, under the Apache License v2.0 with LLVM Exceptions.
4// See https://llvm.org/LICENSE.txt for license information.
5// SPDX-License-Identifier: Apache-2.0 WITH LLVM-exception
6//
7//===----------------------------------------------------------------------===//
8
10
21#include "mlir/IR/Attributes.h"
24#include "mlir/IR/Matchers.h"
26#include "mlir/Pass/Pass.h"
27
29
30#include "llvm/ADT/STLExtras.h"
31#include "llvm/ADT/TypeSwitch.h"
32#include "llvm/Support/AMDGPUAddrSpace.h"
33#include "llvm/Support/Casting.h"
34#include "llvm/Support/ErrorHandling.h"
35#include <cstdint>
36#include <optional>
37
38namespace mlir {
39#define GEN_PASS_DEF_CONVERTAMDGPUTOROCDLPASS
40#include "mlir/Conversion/Passes.h.inc"
41} // namespace mlir
42
43using namespace mlir;
44using namespace mlir::amdgpu;
45
46// Define commonly used chipsets versions for convenience.
47constexpr Chipset kGfx908 = Chipset(9, 0, 8);
48constexpr Chipset kGfx90a = Chipset(9, 0, 0xa);
49constexpr Chipset kGfx942 = Chipset(9, 4, 2);
50constexpr Chipset kGfx950 = Chipset(9, 5, 0);
51constexpr Chipset kGfx1200 = Chipset(12, 0, 0);
52constexpr Chipset kGfx1250 = Chipset(12, 5, 0);
53
54// Predicates mirroring the LLVM AMDGPU `HasDot{N}Insts` features that gate
55// the `v_dot*` instructions consumed by the `amdgpu.dot` lowering.
56static bool hasDot1Insts(const Chipset &chipset) {
57 if (chipset.majorVersion == 9)
58 return chipset >= Chipset(9, 0, 6);
59 if (chipset.majorVersion == 10) {
60 if (chipset.minorVersion == 1)
61 return chipset.steppingVersion == 1u || chipset.steppingVersion == 2u;
62 return chipset.minorVersion >= 3u;
63 }
64 return false;
65}
66
67static bool hasDot2Insts(const Chipset &chipset) {
68 return hasDot1Insts(chipset);
69}
70
71static bool hasDot7Insts(const Chipset &chipset) {
72 return chipset.majorVersion >= 11 || hasDot1Insts(chipset);
73}
74
75static bool hasDot8Insts(const Chipset &chipset) {
76 return chipset.majorVersion >= 11;
77}
78
79static bool hasDot9Insts(const Chipset &chipset) {
80 if (chipset.majorVersion == 11)
81 return true;
82 return chipset.majorVersion == 12 && chipset.minorVersion == 0;
83}
84
85static bool hasDot10Insts(const Chipset &chipset) {
86 if (chipset.majorVersion == 11)
87 return true;
88 if (chipset.majorVersion == 12)
89 return chipset.minorVersion == 0;
90 return hasDot1Insts(chipset);
91}
92
93static bool hasDot11Insts(const Chipset &chipset) {
94 if (chipset.majorVersion == 11)
95 return chipset.minorVersion == 7u;
96 return chipset.majorVersion == 12 && chipset.minorVersion == 0;
97}
98
99static bool hasDot12Insts(const Chipset &chipset) {
100 if (chipset == Chipset(9, 5, 0))
101 return true;
102 if (chipset.majorVersion == 11)
103 return true;
104 return chipset.majorVersion == 12 && chipset.minorVersion == 0;
105}
106
107static bool has45BitNumRecordsBufferResource(const Chipset &chipset) {
108 return chipset.majorVersion > 12 ||
109 (chipset.majorVersion == 12 && chipset.minorVersion >= 5);
110}
111
112/// Convert an unsigned number `val` to i32.
113static Value convertUnsignedToI32(ConversionPatternRewriter &rewriter,
114 Location loc, Value val) {
115 IntegerType i32 = rewriter.getI32Type();
116 // Force check that `val` is of int type.
117 auto valTy = cast<IntegerType>(val.getType());
118 if (i32 == valTy)
119 return val;
120 return valTy.getWidth() > 32
121 ? Value(LLVM::TruncOp::create(rewriter, loc, i32, val))
122 : Value(LLVM::ZExtOp::create(rewriter, loc, i32, val));
123}
124
125static Value createI32Constant(ConversionPatternRewriter &rewriter,
126 Location loc, int32_t value) {
127 return LLVM::ConstantOp::create(rewriter, loc, rewriter.getI32Type(), value);
128}
129
130/// Convert an unsigned number `val` to i64.
131static Value convertUnsignedToI64(ConversionPatternRewriter &rewriter,
132 Location loc, Value val) {
133 IntegerType i64 = rewriter.getI64Type();
134 // Force check that `val` is of int type.
135 auto valTy = cast<IntegerType>(val.getType());
136 if (i64 == valTy)
137 return val;
138 return valTy.getWidth() > 64
139 ? Value(LLVM::TruncOp::create(rewriter, loc, i64, val))
140 : Value(LLVM::ZExtOp::create(rewriter, loc, i64, val));
141}
142
143static Value createI64Constant(ConversionPatternRewriter &rewriter,
144 Location loc, int64_t value) {
145 return LLVM::ConstantOp::create(rewriter, loc, rewriter.getI64Type(), value);
146}
147
148/// Returns the linear index used to access an element in the memref.
149static Value getLinearIndexI32(ConversionPatternRewriter &rewriter,
150 Location loc, MemRefDescriptor &memRefDescriptor,
152 IntegerType i32 = rewriter.getI32Type();
153 Value index;
154 for (auto [i, increment, stride] : llvm::enumerate(indices, strides)) {
155 if (stride != 1) { // Skip if stride is 1.
156 Value strideValue =
157 ShapedType::isDynamic(stride)
158 ? convertUnsignedToI32(rewriter, loc,
159 memRefDescriptor.stride(rewriter, loc, i))
160 : LLVM::ConstantOp::create(rewriter, loc, i32, stride);
161 increment = LLVM::MulOp::create(rewriter, loc, increment, strideValue);
162 }
163 index = index ? LLVM::AddOp::create(rewriter, loc, index, increment)
164 : increment;
165 }
166 return index ? index : createI32Constant(rewriter, loc, 0);
167}
168
169/// Compute the contents of the `num_records` field for a given memref
170/// descriptor - that is, the number of bytes that's one element past the
171/// greatest possible valid index into the memref.
172static Value getNumRecords(ConversionPatternRewriter &rewriter, Location loc,
173 MemRefType memrefType,
174 MemRefDescriptor &memrefDescriptor,
175 ArrayRef<int64_t> strides, int64_t elementByteWidth,
176 amdgpu::Chipset chipset, bool boundsCheck) {
177 if (has45BitNumRecordsBufferResource(chipset) && !boundsCheck) {
178 constexpr int64_t first45bits = (1ll << 45) - 1;
179 return createI64Constant(rewriter, loc, first45bits);
180 }
181 if (memrefType.hasStaticShape() &&
182 !llvm::any_of(strides, ShapedType::isDynamic)) {
183 int64_t size = memrefType.getRank() == 0 ? 1 : 0;
184 ArrayRef<int64_t> shape = memrefType.getShape();
185 for (uint32_t i = 0, e = memrefType.getRank(); i < e; ++i)
186 size = std::max(shape[i] * strides[i], size);
187 size = size * elementByteWidth;
188 return createI64Constant(rewriter, loc, size);
189 }
190 Value maxIndex;
191 for (uint32_t i = 0, e = memrefType.getRank(); i < e; ++i) {
192 Value size = memrefDescriptor.size(rewriter, loc, i);
193 Value stride = memrefDescriptor.stride(rewriter, loc, i);
194 Value maxThisDim = LLVM::MulOp::create(rewriter, loc, size, stride);
195 maxIndex = maxIndex
196 ? LLVM::UMaxOp::create(rewriter, loc, maxIndex, maxThisDim)
197 : maxThisDim;
198 }
199 Value maxIndexI64 = convertUnsignedToI64(rewriter, loc, maxIndex);
200 Value byteWidthConst = createI64Constant(rewriter, loc, elementByteWidth);
201 return LLVM::MulOp::create(rewriter, loc, maxIndexI64, byteWidthConst);
202}
203
204static Value makeBufferRsrc(ConversionPatternRewriter &rewriter, Location loc,
205 Value basePointer, Value numRecords,
206 bool boundsCheck, amdgpu::Chipset chipset,
207 Value cacheSwizzleStride = nullptr,
208 unsigned addressSpace = 8) {
209 // The stride value is generally 0. However, on MI-300 and onward, you can
210 // enable a cache swizzling mode by setting bit 14 of the stride field
211 // and setting that stride to a cache stride.
212 Type i16 = rewriter.getI16Type();
213 Value stride;
214 if (chipset.majorVersion == 9 && chipset >= kGfx942 && cacheSwizzleStride) {
215 Value cacheStrideZext =
216 LLVM::ZExtOp::create(rewriter, loc, i16, cacheSwizzleStride);
217 Value swizzleBit = LLVM::ConstantOp::create(
218 rewriter, loc, i16, rewriter.getI16IntegerAttr(1 << 14));
219 stride = LLVM::OrOp::create(rewriter, loc, cacheStrideZext, swizzleBit,
220 /*isDisjoint=*/true);
221 } else {
222 stride = LLVM::ConstantOp::create(rewriter, loc, i16,
223 rewriter.getI16IntegerAttr(0));
224 }
225
226 uint32_t flags = 0;
227 if (chipset >= kGfx1250) {
228 // Flag word:
229 // bit 0: swizzle
230 // bit 1: 0 means (total_offset + payload > numRecords)
231 // 1 means ((total_offset + payload >) numRecords) || ((offset +
232 // payload) > stride) only applied when swizzle_enable = 0. keep at
233 // zero.
234 // whether oob is done depends on numRecords.
235 // bits 2-3: Type (must be 0)
236 } else {
237 // Get the number of elements.
238 // Flag word:
239 // bits 0-11: dst sel, ignored by these intrinsics
240 // bits 12-14: data format (ignored, must be nonzero, 7=float)
241 // bits 15-18: data format (ignored, must be nonzero, 4=32bit)
242 // bit 19: In nested heap (0 here)
243 // bit 20: Behavior on unmap (0 means "return 0 / ignore")
244 // bits 21-22: Index stride for swizzles (N/A)
245 // bit 23: Add thread ID (0)
246 // bit 24: Reserved to 1 (RDNA) or 0 (CDNA)
247 // bits 25-26: Reserved (0)
248 // bit 27: Buffer is non-volatile (CDNA only)
249 // bits 28-29: Out of bounds select (0 = structured, 1 = check index, 2 =
250 // none, 3 = either swizzles or testing against offset field) RDNA only
251 // bits 30-31: Type (must be 0)
252 flags |= (7 << 12) | (4 << 15);
253 if (chipset.majorVersion >= 10) {
254 flags |= (1 << 24);
255 uint32_t oob = boundsCheck ? 3 : 2;
256 flags |= (oob << 28);
257 }
258 }
259 Value flagsConst = createI32Constant(rewriter, loc, flags);
260 numRecords = has45BitNumRecordsBufferResource(chipset)
261 ? convertUnsignedToI64(rewriter, loc, numRecords)
262 : convertUnsignedToI32(rewriter, loc, numRecords);
263 Type rsrcType =
264 LLVM::LLVMPointerType::get(rewriter.getContext(), addressSpace);
265 Value resource = rewriter.createOrFold<ROCDL::MakeBufferRsrcOp>(
266 loc, rsrcType, basePointer, stride, numRecords, flagsConst);
267 return resource;
268}
269
270namespace {
271struct FatRawBufferCastLowering
272 : public ConvertOpToLLVMPattern<FatRawBufferCastOp> {
273 FatRawBufferCastLowering(const LLVMTypeConverter &converter, Chipset chipset)
274 : ConvertOpToLLVMPattern<FatRawBufferCastOp>(converter),
275 chipset(chipset) {}
276
277 Chipset chipset;
278
279 LogicalResult
280 matchAndRewrite(FatRawBufferCastOp op, FatRawBufferCastOpAdaptor adaptor,
281 ConversionPatternRewriter &rewriter) const override {
282 Location loc = op.getLoc();
283 Value memRef = adaptor.getSource();
284 Value unconvertedMemref = op.getSource();
285 MemRefType memrefType = cast<MemRefType>(unconvertedMemref.getType());
286 MemRefDescriptor descriptor(memRef);
287
288 DataLayout dataLayout = DataLayout::closest(op);
289 int64_t elementByteWidth =
290 dataLayout.getTypeSizeInBits(memrefType.getElementType()) / 8;
291
292 int64_t unusedOffset = 0;
293 SmallVector<int64_t, 5> strideVals;
294 if (failed(memrefType.getStridesAndOffset(strideVals, unusedOffset)))
295 return op.emitOpError("Can't lower non-stride-offset memrefs");
296
297 Value numRecords = adaptor.getValidBytes();
298 if (!numRecords)
299 numRecords =
300 getNumRecords(rewriter, loc, memrefType, descriptor, strideVals,
301 elementByteWidth, chipset, adaptor.getBoundsCheck());
302
303 Value basePointer =
304 adaptor.getResetOffset()
305 ? descriptor.bufferPtr(rewriter, loc, *getTypeConverter(),
306 memrefType)
307 : descriptor.alignedPtr(rewriter, loc);
308
309 Value offset = adaptor.getResetOffset()
310 ? LLVM::ConstantOp::create(rewriter, loc, getIndexType(),
311 rewriter.getIndexAttr(0))
312 : descriptor.offset(rewriter, loc);
313
314 bool hasSizes = memrefType.getRank() > 0;
315 // No need to unpack() and pack() all the individual sizes and strides,
316 // so we'll just extract the arrays.
317 Value sizes = hasSizes
318 ? LLVM::ExtractValueOp::create(rewriter, loc, descriptor,
320 : Value{};
321 Value strides =
322 hasSizes ? LLVM::ExtractValueOp::create(rewriter, loc, descriptor,
324 : Value{};
325
326 Value fatPtr = makeBufferRsrc(
327 rewriter, loc, basePointer, numRecords, adaptor.getBoundsCheck(),
328 chipset, adaptor.getCacheSwizzleStride(), /*addressSpace=*/7);
329
330 Value result = MemRefDescriptor::poison(
331 rewriter, loc,
332 getTypeConverter()->convertType(op.getResult().getType()));
333 SmallVector<int64_t> pos{kAllocatedPtrPosInMemRefDescriptor};
334 result = LLVM::InsertValueOp::create(rewriter, loc, result, fatPtr, pos);
335 result = LLVM::InsertValueOp::create(rewriter, loc, result, fatPtr,
337 result = LLVM::InsertValueOp::create(rewriter, loc, result, offset,
339 if (hasSizes) {
340 result = LLVM::InsertValueOp::create(rewriter, loc, result, sizes,
342 result = LLVM::InsertValueOp::create(rewriter, loc, result, strides,
344 }
345 rewriter.replaceOp(op, result);
346 return success();
347 }
348};
349
350/// Define lowering patterns for raw buffer ops
351template <typename GpuOp, typename Intrinsic>
352struct RawBufferOpLowering : public ConvertOpToLLVMPattern<GpuOp> {
353 RawBufferOpLowering(const LLVMTypeConverter &converter, Chipset chipset)
354 : ConvertOpToLLVMPattern<GpuOp>(converter), chipset(chipset) {}
355
356 Chipset chipset;
357 static constexpr uint32_t maxVectorOpWidth = 128;
358
359 LogicalResult
360 matchAndRewrite(GpuOp gpuOp, typename GpuOp::Adaptor adaptor,
361 ConversionPatternRewriter &rewriter) const override {
362 Location loc = gpuOp.getLoc();
363 Value memref = adaptor.getMemref();
364 Value unconvertedMemref = gpuOp.getMemref();
365 MemRefType memrefType = cast<MemRefType>(unconvertedMemref.getType());
366
367 if (chipset.majorVersion < 9)
368 return gpuOp.emitOpError("raw buffer ops require GCN or higher");
369
370 Value storeData = adaptor.getODSOperands(0)[0];
371 if (storeData == memref) // no write component to this op
372 storeData = Value();
373 Type wantedDataType;
374 if (storeData)
375 wantedDataType = storeData.getType();
376 else
377 wantedDataType = gpuOp.getODSResults(0)[0].getType();
378
379 Value atomicCmpData = Value();
380 // Operand index 1 of a load is the indices, trying to read them can crash.
381 if (storeData) {
382 Value maybeCmpData = adaptor.getODSOperands(1)[0];
383 if (maybeCmpData != memref)
384 atomicCmpData = maybeCmpData;
385 }
386
387 Type llvmWantedDataType = this->typeConverter->convertType(wantedDataType);
388
389 Type i32 = rewriter.getI32Type();
390
391 // Get the type size in bytes.
392 DataLayout dataLayout = DataLayout::closest(gpuOp);
393 int64_t elementByteWidth =
394 dataLayout.getTypeSizeInBits(memrefType.getElementType()) / 8;
395 Value byteWidthConst = createI32Constant(rewriter, loc, elementByteWidth);
396
397 // If we want to load a vector<NxT> with total size <= 32
398 // bits, use a scalar load and bitcast it. Similarly, if bitsize(T) < 32
399 // and the total load size is >= 32, use a vector load of N / (bitsize(T) /
400 // 32) x i32 and bitcast. Also, the CAS intrinsic requires integer operands,
401 // so bitcast any floats to integers.
402 Type llvmBufferValType = llvmWantedDataType;
403 if (atomicCmpData) {
404 if (auto floatType = dyn_cast<FloatType>(wantedDataType))
405 llvmBufferValType = this->getTypeConverter()->convertType(
406 rewriter.getIntegerType(floatType.getWidth()));
407 }
408 if (auto dataVector = dyn_cast<VectorType>(wantedDataType)) {
409 uint32_t vecLen = dataVector.getNumElements();
410 uint32_t elemBits =
411 dataLayout.getTypeSizeInBits(dataVector.getElementType());
412 uint32_t totalBits = elemBits * vecLen;
413 bool usePackedFp16 =
414 isa_and_present<RawBufferAtomicFaddOp>(*gpuOp) && vecLen == 2;
415 if (totalBits > maxVectorOpWidth)
416 return gpuOp.emitOpError(
417 "Total width of loads or stores must be no more than " +
418 Twine(maxVectorOpWidth) + " bits, but we call for " +
419 Twine(totalBits) +
420 " bits. This should've been caught in validation");
421 if (!usePackedFp16 && elemBits < 32) {
422 if (totalBits > 32) {
423 if (totalBits % 32 != 0)
424 return gpuOp.emitOpError("Load or store of more than 32-bits that "
425 "doesn't fit into words. Can't happen\n");
426 llvmBufferValType = this->typeConverter->convertType(
427 VectorType::get(totalBits / 32, i32));
428 } else {
429 llvmBufferValType = this->typeConverter->convertType(
430 rewriter.getIntegerType(totalBits));
431 }
432 }
433 }
434 if (auto vecType = dyn_cast<VectorType>(llvmBufferValType)) {
435 // Buffer intrinsics doesn't support 1-element vectors, cast them to
436 // scalars.
437 if (vecType.getNumElements() == 1)
438 llvmBufferValType = vecType.getElementType();
439 }
440
441 SmallVector<Value, 6> args;
442 if (storeData) {
443 if (llvmBufferValType != llvmWantedDataType) {
444 Value castForStore = LLVM::BitcastOp::create(
445 rewriter, loc, llvmBufferValType, storeData);
446 args.push_back(castForStore);
447 } else {
448 args.push_back(storeData);
449 }
450 }
451
452 if (atomicCmpData) {
453 if (llvmBufferValType != llvmWantedDataType) {
454 Value castForCmp = LLVM::BitcastOp::create(
455 rewriter, loc, llvmBufferValType, atomicCmpData);
456 args.push_back(castForCmp);
457 } else {
458 args.push_back(atomicCmpData);
459 }
460 }
461
462 // Construct buffer descriptor from memref, attributes
463 int64_t offset = 0;
464 SmallVector<int64_t, 5> strides;
465 if (failed(memrefType.getStridesAndOffset(strides, offset)))
466 return gpuOp.emitOpError("Can't lower non-stride-offset memrefs");
467
468 MemRefDescriptor memrefDescriptor(memref);
469
470 Value ptr = memrefDescriptor.bufferPtr(
471 rewriter, loc, *this->getTypeConverter(), memrefType);
472 Value numRecords =
473 getNumRecords(rewriter, loc, memrefType, memrefDescriptor, strides,
474 elementByteWidth, chipset, adaptor.getBoundsCheck());
475 Value resource = makeBufferRsrc(rewriter, loc, ptr, numRecords,
476 adaptor.getBoundsCheck(), chipset);
477 args.push_back(resource);
478
479 // Indexing (voffset)
480 Value voffset = getLinearIndexI32(rewriter, loc, memrefDescriptor,
481 adaptor.getIndices(), strides);
482 if (std::optional<int32_t> indexOffset = adaptor.getIndexOffset();
483 indexOffset && *indexOffset > 0) {
484 Value extraOffsetConst = createI32Constant(rewriter, loc, *indexOffset);
485 voffset = voffset ? LLVM::AddOp::create(rewriter, loc, voffset,
486 extraOffsetConst)
487 : extraOffsetConst;
488 }
489 voffset = LLVM::MulOp::create(rewriter, loc, voffset, byteWidthConst);
490 args.push_back(voffset);
491
492 // SGPR offset.
493 Value sgprOffset = adaptor.getSgprOffset();
494 if (!sgprOffset)
495 sgprOffset = createI32Constant(rewriter, loc, 0);
496 sgprOffset = LLVM::MulOp::create(rewriter, loc, sgprOffset, byteWidthConst);
497 args.push_back(sgprOffset);
498
499 llvm::SmallVector<Type, 1> resultTypes(gpuOp->getNumResults(),
500 llvmBufferValType);
501 typename Intrinsic::Properties properties;
502 properties.aux = rewriter.getI32IntegerAttr(0);
503 Operation *lowered =
504 Intrinsic::create(rewriter, loc, resultTypes, args, properties);
505 if (lowered->getNumResults() == 1) {
506 Value replacement = lowered->getResult(0);
507 if (llvmBufferValType != llvmWantedDataType) {
508 replacement = LLVM::BitcastOp::create(rewriter, loc, llvmWantedDataType,
510 }
511 rewriter.replaceOp(gpuOp, replacement);
512 } else {
513 rewriter.eraseOp(gpuOp);
514 }
515 return success();
516 }
517};
518
519// TODO: AMDGPU backend already have all this bitpacking logic, we should move
520// it to some common place.
521/// Vmcnt, Expcnt and Lgkmcnt are decoded as follows:
522/// Vmcnt = Waitcnt[3:0] (pre-gfx9)
523/// Vmcnt = Waitcnt[15:14,3:0] (gfx9,10)
524/// Vmcnt = Waitcnt[15:10] (gfx11)
525/// Expcnt = Waitcnt[6:4] (pre-gfx11)
526/// Expcnt = Waitcnt[2:0] (gfx11)
527/// Lgkmcnt = Waitcnt[11:8] (pre-gfx10)
528/// Lgkmcnt = Waitcnt[13:8] (gfx10)
529/// Lgkmcnt = Waitcnt[9:4] (gfx11)
530static FailureOr<unsigned> encodeWaitcnt(Chipset chipset, unsigned vmcnt,
531 unsigned expcnt, unsigned lgkmcnt) {
532 if (chipset.majorVersion < 9) {
533 vmcnt = std::min(15u, vmcnt);
534 expcnt = std::min(7u, expcnt);
535 lgkmcnt = std::min(15u, lgkmcnt);
536 return vmcnt | (expcnt << 4) | (lgkmcnt << 8);
537 }
538 if (chipset.majorVersion == 9) {
539 vmcnt = std::min(63u, vmcnt);
540 expcnt = std::min(7u, expcnt);
541 lgkmcnt = std::min(15u, lgkmcnt);
542 unsigned lowBits = vmcnt & 0xF;
543 unsigned highBits = (vmcnt >> 4) << 14;
544 unsigned otherCnts = (expcnt << 4) | (lgkmcnt << 8);
545 return lowBits | highBits | otherCnts;
547 if (chipset.majorVersion == 10) {
548 vmcnt = std::min(63u, vmcnt);
549 expcnt = std::min(7u, expcnt);
550 lgkmcnt = std::min(63u, lgkmcnt);
551 unsigned lowBits = vmcnt & 0xF;
552 unsigned highBits = (vmcnt >> 4) << 14;
553 unsigned otherCnts = (expcnt << 4) | (lgkmcnt << 8);
554 return lowBits | highBits | otherCnts;
556 if (chipset.majorVersion == 11) {
557 vmcnt = std::min(63u, vmcnt);
558 expcnt = std::min(7u, expcnt);
559 lgkmcnt = std::min(63u, lgkmcnt);
560 return (vmcnt << 10) | expcnt | (lgkmcnt << 4);
562 return failure();
564
565struct MemoryCounterWaitOpLowering
566 : public ConvertOpToLLVMPattern<MemoryCounterWaitOp> {
567 MemoryCounterWaitOpLowering(const LLVMTypeConverter &converter,
569 : ConvertOpToLLVMPattern<MemoryCounterWaitOp>(converter),
570 chipset(chipset) {}
571
573
574 LogicalResult
575 matchAndRewrite(MemoryCounterWaitOp op, OpAdaptor adaptor,
576 ConversionPatternRewriter &rewriter) const override {
577 if (chipset.majorVersion >= 12) {
578 Location loc = op.getLoc();
579 if (std::optional<int> ds = adaptor.getDs())
580 ROCDL::WaitDscntOp::create(rewriter, loc, *ds);
581
582 if (std::optional<int> load = adaptor.getLoad())
583 ROCDL::WaitLoadcntOp::create(rewriter, loc, *load);
584
585 if (std::optional<int> store = adaptor.getStore())
586 ROCDL::WaitStorecntOp::create(rewriter, loc, *store);
587
588 if (std::optional<int> exp = adaptor.getExp())
589 ROCDL::WaitExpcntOp::create(rewriter, loc, *exp);
590
591 if (std::optional<int> tensor = adaptor.getTensor())
592 ROCDL::WaitTensorcntOp::create(rewriter, loc, *tensor);
593
594 rewriter.eraseOp(op);
595 return success();
597
598 if (adaptor.getTensor())
599 return op.emitOpError("unsupported chipset");
600
601 auto getVal = [](Attribute attr) -> unsigned {
602 if (attr)
603 return cast<IntegerAttr>(attr).getInt();
604
605 // This value will be clamped to the maximum value for the chipset.
606 return 1024;
607 };
608 unsigned ds = getVal(adaptor.getDsAttr());
609 unsigned exp = getVal(adaptor.getExpAttr());
610
611 unsigned vmcnt = 1024;
612 Attribute load = adaptor.getLoadAttr();
613 Attribute store = adaptor.getStoreAttr();
614 if (load && store) {
615 vmcnt = getVal(load) + getVal(store);
616 } else if (load) {
617 vmcnt = getVal(load);
618 } else if (store) {
619 vmcnt = getVal(store);
620 }
621
622 FailureOr<unsigned> waitcnt = encodeWaitcnt(chipset, vmcnt, exp, ds);
623 if (failed(waitcnt))
624 return op.emitOpError("unsupported chipset");
625
626 rewriter.replaceOpWithNewOp<ROCDL::SWaitcntOp>(op, *waitcnt);
627 return success();
628 }
629};
630
631struct LDSBarrierOpLowering : public ConvertOpToLLVMPattern<LDSBarrierOp> {
632 LDSBarrierOpLowering(const LLVMTypeConverter &converter, Chipset chipset)
633 : ConvertOpToLLVMPattern<LDSBarrierOp>(converter), chipset(chipset) {}
634
635 Chipset chipset;
636
637 LogicalResult
638 matchAndRewrite(LDSBarrierOp op, LDSBarrierOp::Adaptor adaptor,
639 ConversionPatternRewriter &rewriter) const override {
640 Location loc = op.getLoc();
641 // This ensures that waits on global memory aren't introduced on
642 // chips that don't have the BackOffBarrier feature enabled in LLVM.
643 bool requiresInlineAsm = chipset < kGfx90a;
644
645 Attribute mmra =
646 rewriter.getAttr<LLVM::MMRATagAttr>("amdgpu-synchronize-as", "local");
647 // Note: while there *is* a workgroup-one-as scope, this, when combined with
648 // the MMRA, will lead to the fence having no effect. This is because the
649 // codepaths for an atomic load or store will observe that a
650 // one-address-space atomic to LDS requires no synchronization because
651 // operations on LDS are totally ordered with respect to each other, and so
652 // will not emit the correct waitcnt operations that these fences are
653 // intended to produce. Therefore, we use a broader type of fence and rely
654 // on the MMRA to relax it to the semantics we want.
655 StringRef scope = "workgroup";
656
657 auto relFence = LLVM::FenceOp::create(rewriter, loc,
658 LLVM::AtomicOrdering::release, scope);
659 relFence->setDiscardableAttr(LLVM::LLVMDialect::getMmraAttrName(), mmra);
660 if (requiresInlineAsm) {
661 auto asmDialectAttr = LLVM::AsmDialectAttr::get(rewriter.getContext(),
662 LLVM::AsmDialect::AD_ATT);
663 const char *asmStr = ";;;WARNING: BREAKS DEBUG WATCHES\ns_barrier";
664 const char *constraints = "";
665 LLVM::InlineAsmOp::create(
666 rewriter, loc,
667 /*resultTypes=*/TypeRange(), /*operands=*/ValueRange(),
668 /*asm_string=*/asmStr, constraints, /*has_side_effects=*/true,
669 /*is_align_stack=*/false, LLVM::TailCallKind::None,
670 /*asm_dialect=*/asmDialectAttr,
671 /*operand_attrs=*/ArrayAttr());
672 } else if (chipset.majorVersion < 12) {
673 ROCDL::SBarrierOp::create(rewriter, loc);
674 } else {
675 ROCDL::BarrierSignalOp::create(rewriter, loc, -1);
676 ROCDL::BarrierWaitOp::create(rewriter, loc, -1);
677 }
678
679 auto acqFence = LLVM::FenceOp::create(rewriter, loc,
680 LLVM::AtomicOrdering::acquire, scope);
681 acqFence->setDiscardableAttr(LLVM::LLVMDialect::getMmraAttrName(), mmra);
682 rewriter.replaceOp(op, acqFence);
683 return success();
684 }
685};
686
687struct SchedBarrierOpLowering : public ConvertOpToLLVMPattern<SchedBarrierOp> {
688 SchedBarrierOpLowering(const LLVMTypeConverter &converter, Chipset chipset)
689 : ConvertOpToLLVMPattern<SchedBarrierOp>(converter), chipset(chipset) {}
690
691 Chipset chipset;
692
693 LogicalResult
694 matchAndRewrite(SchedBarrierOp op, SchedBarrierOp::Adaptor adaptor,
695 ConversionPatternRewriter &rewriter) const override {
696 rewriter.replaceOpWithNewOp<ROCDL::SchedBarrier>(op, op.getOptsAttr());
697 return success();
698 }
699};
700
701} // namespace
702
703/// Pack small float vector operands (fp4/fp6/fp8/bf16) into the format
704/// expected by scaled matrix multiply intrinsics (MFMA/WMMA).
705///
706/// Specifically:
707/// 1. If the element type is bfloat16, bitcast it to i16 unless rocdl intrinsic
708/// allows bf16. Newer MFMAs support bf16 types on operand, check
709/// IntrinsicsAMDGPU.td file for reference.
710/// 2. If instead we have a more than 64-bit quantity, use a <N / 4 x i32>
711/// instead, which is what the f8f6f4 intrinsics use.
712/// 3. If `input` is a vector of N <= 8 bytes, bitcast it to a (N * 8)-bit
713/// integer.
714///
715/// Note that the type of `input` has already been LLVM type converted:
716/// therefore 8-bit and smaller floats are represented as their corresponding
717/// `iN` integers.
718static Value packSmallFloatVectorOperand(ConversionPatternRewriter &rewriter,
719 Location loc, Value input,
720 bool allowBf16 = true) {
721 Type inputType = input.getType();
722 if (auto vectorType = dyn_cast<VectorType>(inputType)) {
723 if (vectorType.getElementType().isBF16() && !allowBf16)
724 return LLVM::BitcastOp::create(
725 rewriter, loc, vectorType.clone(rewriter.getI16Type()), input);
726 if (vectorType.getElementType().isInteger(8) &&
727 vectorType.getNumElements() <= 8)
728 return LLVM::BitcastOp::create(
729 rewriter, loc,
730 rewriter.getIntegerType(vectorType.getNumElements() * 8), input);
731 if (isa<IntegerType>(vectorType.getElementType()) &&
732 vectorType.getElementTypeBitWidth() <= 8) {
733 int64_t numWords = llvm::divideCeil(
734 vectorType.getNumElements() * vectorType.getElementTypeBitWidth(),
735 32);
736 return LLVM::BitcastOp::create(
737 rewriter, loc, VectorType::get(numWords, rewriter.getI32Type()),
738 input);
739 }
740 }
741 return input;
742}
743
744/// Converts packed vector operands to the expected ROCDL types.
745static Value convertPackedVectorOperand(ConversionPatternRewriter &rewriter,
746 Location loc, Value input,
747 bool allowBf16 = true) {
748 Type inputType = input.getType();
749 auto vectorType = cast<VectorType>(inputType);
750 // bf16 -> i16 when not allowed (pre-gfx950).
751 if (vectorType.getElementType().isBF16() && !allowBf16)
752 return LLVM::BitcastOp::create(
753 rewriter, loc, vectorType.clone(rewriter.getI16Type()), input);
754 // i8/fp8 vectors -> vector<Nxi32>.
755 if (isa<IntegerType>(vectorType.getElementType()) &&
756 vectorType.getElementTypeBitWidth() <= 8) {
757 int64_t numWords = llvm::divideCeil(
758 vectorType.getNumElements() * vectorType.getElementTypeBitWidth(), 32);
759 Type castType = (numWords > 1)
760 ? Type{VectorType::get(numWords, rewriter.getI32Type())}
761 : rewriter.getI32Type();
762 return LLVM::BitcastOp::create(rewriter, loc, castType, input);
763 }
764 return input;
765}
766
767/// Converts the scaled MFMA/WMMA operands, `scalesA` and `scalesB`, from MLIR
768/// AMDGPU dialect convention to ROCDL and LLVM AMDGPU intrinsics convention.
769///
770/// Specifically:
771/// 1. If `input` is a i8 value, zero extend it to i32
772/// 2. If `input` is a vector of length 4 or 8 and type i8, cast it to i32
773///
774/// Note that the type of `input` has already been LLVM type converted:
775/// therefore 8-bit and smaller floats are represented as their corresponding
776/// `iN` integers.
777static Value castScaleOperand(ConversionPatternRewriter &rewriter, Location loc,
778 Value input) {
779 return TypeSwitch<Type, Value>(input.getType())
780 .Case([&](IntegerType) {
781 // Handle scalar i8: zero extend to i32.
782 return LLVM::ZExtOp::create(rewriter, loc, rewriter.getI32Type(),
783 input);
784 })
785 .Case([&](VectorType vectorType) {
786 // Handle vector<4xi8> -> i32 or vector<8xi8> -> i64.
787 int64_t numElements = vectorType.getNumElements();
788 assert((numElements == 4 || numElements == 8) &&
789 "scale operand must be a vector of length 4 or 8");
790 IntegerType outputType =
791 (numElements == 4) ? rewriter.getI32Type() : rewriter.getI64Type();
792 return LLVM::BitcastOp::create(rewriter, loc, outputType, input);
793 })
794 .DefaultUnreachable("unexpected input type for scale operand");
795}
796
797/// Maps f8 scale element types to WMMA scale format codes.
798static std::optional<ROCDL::WMMAMatrixScaleFormat>
801 .Case([](Float8E8M0FNUType) { return ROCDL::WMMAMatrixScaleFormat::e8; })
802 .Case([](Float8E4M3FNType) { return ROCDL::WMMAMatrixScaleFormat::e4m3; })
803 .Default(std::nullopt);
804}
805
806/// Determines the ROCDL intrinsic name for scaled WMMA based on dimensions
807/// and scale block size (16 or 32).
808static std::optional<StringRef>
810 if (m == 16 && n == 16 && k == 128)
811 return isScale16
812 ? ROCDL::wmma_scale16_f32_16x16x128_f8f6f4::getOperationName()
813 : ROCDL::wmma_scale_f32_16x16x128_f8f6f4::getOperationName();
814
815 if (m == 32 && n == 16 && k == 128)
816 return isScale16 ? ROCDL::wmma_scale16_f32_32x16x128_f4::getOperationName()
817 : ROCDL::wmma_scale_f32_32x16x128_f4::getOperationName();
818
819 return std::nullopt;
820}
821
822/// Push an input operand. If it is a float type, nothing to do. If it is
823/// an integer type, then we need to also push its signdness (1 for signed, 0
824/// for unsigned) and we need to pack the input 16xi8 vector into a 4xi32
825/// vector (or the 8xi8 vector into a 2xi32 one for gfx12+).
826/// We also need to convert bfloat inputs to i16 to account for the bfloat
827/// intrinsics having been defined before the AMD backend supported bfloat. We
828/// similarly need to pack 8-bit float types into integers as if they were i8
829/// (which they are for the backend's purposes).
831 ConversionPatternRewriter &rewriter, Location loc,
832 const TypeConverter *typeConverter, bool isUnsigned, Value llvmInput,
833 Value mlirInput, SmallVectorImpl<Value> &operands,
834 SmallVectorImpl<NamedAttribute> &attrs, StringRef attrName) {
835 Type inputType = llvmInput.getType();
836 auto vectorType = dyn_cast<VectorType>(inputType);
837 if (!vectorType) {
838 operands.push_back(llvmInput);
839 return;
840 }
841 Type elemType = vectorType.getElementType();
842 if (elemType.getIntOrFloatBitWidth() > 8) {
843 operands.push_back(llvmInput);
844 return;
845 }
846
847 // We need to check the type of the input before conversion to properly test
848 // for int8. This is because, in LLVM, fp8 type is converted to int8, so the
849 // fp8/int8 information is lost during the conversion process.
850 auto mlirInputType = cast<VectorType>(mlirInput.getType());
851 bool isInputInteger = mlirInputType.getElementType().isInteger();
852 if (isInputInteger) {
853 // if element type is 8-bit signed or unsigned, ignore the isUnsigned flag
854 bool localIsUnsigned = isUnsigned;
855 if (elemType.isUnsignedInteger()) {
856 localIsUnsigned = true;
857 } else if (elemType.isSignedInteger()) {
858 localIsUnsigned = false;
859 }
860 attrs.push_back(
861 NamedAttribute(attrName, rewriter.getBoolAttr(!localIsUnsigned)));
862 }
863
864 int64_t numBits =
865 vectorType.getNumElements() * elemType.getIntOrFloatBitWidth();
866 Type i32 = rewriter.getI32Type();
867 Type intrinsicInType = numBits <= 32
868 ? (Type)rewriter.getIntegerType(numBits)
869 : (Type)VectorType::get(numBits / 32, i32);
870 auto llvmIntrinsicInType = typeConverter->convertType(intrinsicInType);
871 Value castInput = rewriter.createOrFold<LLVM::BitcastOp>(
872 loc, llvmIntrinsicInType, llvmInput);
873 // The wave64-mode 16x16x16 intrinsics that take 4-bit integers only need
874 // (256 / 64) * 4 = 16 bits of input (on gfx12+) but take i32 arguments.
875 // Add in the zeros here.
876 if (numBits < 32)
877 castInput = LLVM::ZExtOp::create(rewriter, loc, i32, castInput);
878 operands.push_back(castInput);
879}
880
881/// Push the output operand. For many cases this is only pushing the output in
882/// the operand list. But when we have f16 -> f16 or bf16 -> bf16 intrinsics,
883/// since the same numbers of VGPRs is used, we need to decide if to store the
884/// result in the upper 16 bits of the VGPRs or in the lower part. To store the
885/// result in the lower 16 bits, set subwordOffset to 1, otherwise result will
886/// be stored it in the upper part. The subwordOffset must not be set for gfx12,
887/// as the instructions have been changed to return fewer registers instead.
888static void wmmaPushOutputOperand(ConversionPatternRewriter &rewriter,
889 Location loc,
890 const TypeConverter *typeConverter,
891 Value output, int32_t subwordOffset,
892 bool clamp, SmallVectorImpl<Value> &operands,
894 Type inputType = output.getType();
895 auto vectorType = dyn_cast<VectorType>(inputType);
896 Type elemType = vectorType.getElementType();
897 operands.push_back(output);
898 if (elemType.isF16() || elemType.isBF16() || elemType.isInteger(16)) {
899 attrs.push_back(
900 NamedAttribute("opsel", rewriter.getBoolAttr(subwordOffset)));
901 } else if (elemType.isInteger(32)) {
902 attrs.push_back(NamedAttribute("clamp", rewriter.getBoolAttr(clamp)));
903 }
904}
905
906/// Return true if `type` is the E5M2 variant of an 8-bit float that is
907/// supported by the `_bf8` instructions on the given `chipset`.
908static bool typeIsExpectedBf8ForChipset(Chipset chipset, Type type) {
909 return (chipset == kGfx942 && isa<Float8E5M2FNUZType>(type)) ||
910 (hasOcpFp8(chipset) && isa<Float8E5M2Type>(type));
911}
912
913/// Return true if `type` is the E4M3FN variant of an 8-bit float that is
914/// supported by the `_fp8` instructions on the given `chipset`.
915static bool typeIsExpectedFp8ForChipset(Chipset chipset, Type type) {
916 return (chipset == kGfx942 && isa<Float8E4M3FNUZType>(type)) ||
917 (hasOcpFp8(chipset) && isa<Float8E4M3FNType>(type));
918}
919
920/// Return the `rocdl` intrinsic corresponding to a MFMA operation `mfma`
921/// if one exists. This includes checking to ensure the intrinsic is supported
922/// on the architecture you are compiling for.
923static std::optional<StringRef> mfmaOpToIntrinsic(MFMAOp mfma,
924 Chipset chipset) {
925 uint32_t m = mfma.getM(), n = mfma.getN(), k = mfma.getK(),
926 b = mfma.getBlocks();
927 Type sourceElem = getElementTypeOrSelf(mfma.getSourceA().getType());
928 Type destElem = getElementTypeOrSelf(mfma.getDestC().getType());
929
930 if (sourceElem.isF32() && destElem.isF32()) {
931 if (mfma.getReducePrecision() && chipset >= kGfx942) {
932 if (m == 32 && n == 32 && k == 4 && b == 1)
933 return ROCDL::mfma_f32_32x32x4_xf32::getOperationName();
934 if (m == 16 && n == 16 && k == 8 && b == 1)
935 return ROCDL::mfma_f32_16x16x8_xf32::getOperationName();
936 }
937 if (m == 32 && n == 32 && k == 1 && b == 2)
938 return ROCDL::mfma_f32_32x32x1f32::getOperationName();
939 if (m == 16 && n == 16 && k == 1 && b == 4)
940 return ROCDL::mfma_f32_16x16x1f32::getOperationName();
941 if (m == 4 && n == 4 && k == 1 && b == 16)
942 return ROCDL::mfma_f32_4x4x1f32::getOperationName();
943 if (m == 32 && n == 32 && k == 2 && b == 1)
944 return ROCDL::mfma_f32_32x32x2f32::getOperationName();
945 if (m == 16 && n == 16 && k == 4 && b == 1)
946 return ROCDL::mfma_f32_16x16x4f32::getOperationName();
947 }
948
949 if (sourceElem.isF16() && destElem.isF32()) {
950 if (chipset >= kGfx950) {
951 if (m == 32 && n == 32 && k == 16 && b == 1)
952 return ROCDL::mfma_f32_32x32x16_f16::getOperationName();
953 if (m == 16 && n == 16 && k == 32 && b == 1)
954 return ROCDL::mfma_f32_16x16x32_f16::getOperationName();
955 }
956 if (m == 32 && n == 32 && k == 4 && b == 2)
957 return ROCDL::mfma_f32_32x32x4f16::getOperationName();
958 if (m == 16 && n == 16 && k == 4 && b == 4)
959 return ROCDL::mfma_f32_16x16x4f16::getOperationName();
960 if (m == 4 && n == 4 && k == 4 && b == 16)
961 return ROCDL::mfma_f32_4x4x4f16::getOperationName();
962 if (m == 32 && n == 32 && k == 8 && b == 1)
963 return ROCDL::mfma_f32_32x32x8f16::getOperationName();
964 if (m == 16 && n == 16 && k == 16 && b == 1)
965 return ROCDL::mfma_f32_16x16x16f16::getOperationName();
966 }
967
968 if (sourceElem.isBF16() && destElem.isF32()) {
969 if (chipset >= kGfx950) {
970 if (m == 32 && n == 32 && k == 16 && b == 1)
971 return ROCDL::mfma_f32_32x32x16_bf16::getOperationName();
972 if (m == 16 && n == 16 && k == 32 && b == 1)
973 return ROCDL::mfma_f32_16x16x32_bf16::getOperationName();
974 }
975 if (chipset >= kGfx90a) {
976 if (m == 32 && n == 32 && k == 4 && b == 2)
977 return ROCDL::mfma_f32_32x32x4bf16_1k::getOperationName();
978 if (m == 16 && n == 16 && k == 4 && b == 4)
979 return ROCDL::mfma_f32_16x16x4bf16_1k::getOperationName();
980 if (m == 4 && n == 4 && k == 4 && b == 16)
981 return ROCDL::mfma_f32_4x4x4bf16_1k::getOperationName();
982 if (m == 32 && n == 32 && k == 8 && b == 1)
983 return ROCDL::mfma_f32_32x32x8bf16_1k::getOperationName();
984 if (m == 16 && n == 16 && k == 16 && b == 1)
985 return ROCDL::mfma_f32_16x16x16bf16_1k::getOperationName();
986 }
987 if (m == 32 && n == 32 && k == 2 && b == 2)
988 return ROCDL::mfma_f32_32x32x2bf16::getOperationName();
989 if (m == 16 && n == 16 && k == 2 && b == 4)
990 return ROCDL::mfma_f32_16x16x2bf16::getOperationName();
991 if (m == 4 && n == 4 && k == 2 && b == 16)
992 return ROCDL::mfma_f32_4x4x2bf16::getOperationName();
993 if (m == 32 && n == 32 && k == 4 && b == 1)
994 return ROCDL::mfma_f32_32x32x4bf16::getOperationName();
995 if (m == 16 && n == 16 && k == 8 && b == 1)
996 return ROCDL::mfma_f32_16x16x8bf16::getOperationName();
997 }
998
999 if (sourceElem.isInteger(8) && destElem.isInteger(32)) {
1000 if (chipset >= kGfx950) {
1001 if (m == 32 && n == 32 && k == 32 && b == 1)
1002 return ROCDL::mfma_i32_32x32x32_i8::getOperationName();
1003 if (m == 16 && n == 16 && k == 64 && b == 1)
1004 return ROCDL::mfma_i32_16x16x64_i8::getOperationName();
1005 }
1006 if (m == 32 && n == 32 && k == 4 && b == 2)
1007 return ROCDL::mfma_i32_32x32x4i8::getOperationName();
1008 if (m == 16 && n == 16 && k == 4 && b == 4)
1009 return ROCDL::mfma_i32_16x16x4i8::getOperationName();
1010 if (m == 4 && n == 4 && k == 4 && b == 16)
1011 return ROCDL::mfma_i32_4x4x4i8::getOperationName();
1012 if (m == 32 && n == 32 && k == 8 && b == 1)
1013 return ROCDL::mfma_i32_32x32x8i8::getOperationName();
1014 if (m == 16 && n == 16 && k == 16 && b == 1)
1015 return ROCDL::mfma_i32_16x16x16i8::getOperationName();
1016 if (m == 32 && n == 32 && k == 16 && b == 1 && chipset >= kGfx942)
1017 return ROCDL::mfma_i32_32x32x16_i8::getOperationName();
1018 if (m == 16 && n == 16 && k == 32 && b == 1 && chipset >= kGfx942)
1019 return ROCDL::mfma_i32_16x16x32_i8::getOperationName();
1020 }
1021
1022 if (sourceElem.isF64() && destElem.isF64() && chipset >= kGfx90a) {
1023 if (m == 16 && n == 16 && k == 4 && b == 1)
1024 return ROCDL::mfma_f64_16x16x4f64::getOperationName();
1025 if (m == 4 && n == 4 && k == 4 && b == 4)
1026 return ROCDL::mfma_f64_4x4x4f64::getOperationName();
1027 }
1028
1029 if (destElem.isF32() && typeIsExpectedBf8ForChipset(chipset, sourceElem)) {
1030 // Known to be correct because there are no scalar f8 instructions and
1031 // because a length mismatch will have been caught by the verifier.
1032 Type sourceBElem =
1033 cast<VectorType>(mfma.getSourceB().getType()).getElementType();
1034 if (m == 16 && n == 16 && k == 32 && b == 1) {
1035 if (typeIsExpectedBf8ForChipset(chipset, sourceBElem))
1036 return ROCDL::mfma_f32_16x16x32_bf8_bf8::getOperationName();
1037 if (typeIsExpectedFp8ForChipset(chipset, sourceBElem))
1038 return ROCDL::mfma_f32_16x16x32_bf8_fp8::getOperationName();
1039 }
1040 if (m == 32 && n == 32 && k == 16 && b == 1) {
1041 if (typeIsExpectedBf8ForChipset(chipset, sourceBElem))
1042 return ROCDL::mfma_f32_32x32x16_bf8_bf8::getOperationName();
1043 if (typeIsExpectedFp8ForChipset(chipset, sourceBElem))
1044 return ROCDL::mfma_f32_32x32x16_bf8_fp8::getOperationName();
1045 }
1046 }
1047
1048 if (destElem.isF32() && typeIsExpectedFp8ForChipset(chipset, sourceElem)) {
1049 Type sourceBElem =
1050 cast<VectorType>(mfma.getSourceB().getType()).getElementType();
1051 if (m == 16 && n == 16 && k == 32 && b == 1) {
1052 if (typeIsExpectedBf8ForChipset(chipset, sourceBElem))
1053 return ROCDL::mfma_f32_16x16x32_fp8_bf8::getOperationName();
1054 if (typeIsExpectedFp8ForChipset(chipset, sourceBElem))
1055 return ROCDL::mfma_f32_16x16x32_fp8_fp8::getOperationName();
1056 }
1057 if (m == 32 && n == 32 && k == 16 && b == 1) {
1058 if (typeIsExpectedBf8ForChipset(chipset, sourceBElem))
1059 return ROCDL::mfma_f32_32x32x16_fp8_bf8::getOperationName();
1060 if (typeIsExpectedFp8ForChipset(chipset, sourceBElem))
1061 return ROCDL::mfma_f32_32x32x16_fp8_fp8::getOperationName();
1062 }
1063 }
1064
1065 return std::nullopt;
1066}
1067
1068static std::optional<ROCDL::MatrixFormat>
1071 mlirElemType)
1072 .Case([](Float8E4M3FNType) { return ROCDL::MatrixFormat::fp8_e4m3; })
1073 .Case([](Float8E5M2Type) { return ROCDL::MatrixFormat::fp8_e5m2; })
1074 .Case([](Float6E2M3FNType) { return ROCDL::MatrixFormat::fp6_e2m3; })
1075 .Case([](Float6E3M2FNType) { return ROCDL::MatrixFormat::fp6_e3m2; })
1076 .Case([](Float4E2M1FNType) { return ROCDL::MatrixFormat::fp4_e2m1; })
1077 .Default(std::nullopt);
1078}
1079
1080/// If there is a scaled MFMA instruction for the input element types `aType`
1081/// and `bType`, output type `destType`, problem size M, N, K, and B (number of
1082/// blocks) on the given `chipset`, return a tuple consisting of the
1083/// OperationName of the intrinsic and the type codes that need to be passed to
1084/// that intrinsic. Note that this is also used to implement some un-scaled
1085/// MFMAs, since the compiler represents the ordinary instruction as a "scaled"
1086/// MFMA with a scale of 0.
1088 std::tuple<StringRef, ROCDL::MatrixFormat, ROCDL::MatrixFormat>;
1089
1090static std::optional<ScaledMFMAIntrinsic>
1091mfmaOpToScaledIntrinsic(Type aType, Type bType, Type destType, uint32_t m,
1092 uint32_t n, uint32_t k, uint32_t b, Chipset chipset) {
1093 aType = getElementTypeOrSelf(aType);
1094 bType = getElementTypeOrSelf(bType);
1095 destType = getElementTypeOrSelf(destType);
1096
1097 if (chipset < kGfx950)
1098 return std::nullopt;
1099 if (!isa<Float32Type>(destType))
1100 return std::nullopt;
1101
1102 std::optional<ROCDL::MatrixFormat> aTypeCode =
1104 std::optional<ROCDL::MatrixFormat> bTypeCode =
1106 if (!aTypeCode || !bTypeCode)
1107 return std::nullopt;
1108
1109 if (m == 32 && n == 32 && k == 64 && b == 1)
1110 return std::tuple{ROCDL::mfma_scale_f32_32x32x64_f8f6f4::getOperationName(),
1111 *aTypeCode, *bTypeCode};
1112 if (m == 16 && n == 16 && k == 128 && b == 1)
1113 return std::tuple{
1114 ROCDL::mfma_scale_f32_16x16x128_f8f6f4::getOperationName(), *aTypeCode,
1115 *bTypeCode};
1116
1117 return std::nullopt;
1118}
1119
1120static std::optional<ScaledMFMAIntrinsic>
1121mfmaOpToScaledIntrinsic(MFMAOp mfma, Chipset chipset) {
1123 mfma.getSourceA().getType(), mfma.getSourceB().getType(),
1124 mfma.getDestC().getType(), mfma.getM(), mfma.getN(), mfma.getK(),
1125 mfma.getBlocks(), chipset);
1126}
1127
1128static std::optional<ScaledMFMAIntrinsic>
1129mfmaOpToScaledIntrinsic(ScaledMFMAOp smfma, Chipset chipset) {
1130 return mfmaOpToScaledIntrinsic(smfma.getSourceA().getType(),
1131 smfma.getSourceB().getType(),
1132 smfma.getDestC().getType(), smfma.getM(),
1133 smfma.getN(), smfma.getK(), 1u, chipset);
1134}
1135
1136/// Returns the `rocdl` intrinsic corresponding to a WMMA operation `wmma`
1137/// for RDNA3/4 architectures.
1138static std::optional<StringRef>
1139wmmaOpToIntrinsicRDNA(Type elemSourceType, Type elemBSourceType,
1140 Type elemDestType, uint32_t k, bool isRDNA3) {
1141 using fp8 = Float8E4M3FNType;
1142 using bf8 = Float8E5M2Type;
1143
1144 // Handle k == 16 for RDNA3/4.
1145 if (k == 16) {
1146 // Common patterns for RDNA3 and RDNA4.
1147 if (elemSourceType.isF16() && elemDestType.isF32())
1148 return ROCDL::wmma_f32_16x16x16_f16::getOperationName();
1149 if (elemSourceType.isBF16() && elemDestType.isF32())
1150 return ROCDL::wmma_f32_16x16x16_bf16::getOperationName();
1151 if (elemSourceType.isF16() && elemDestType.isF16())
1152 return ROCDL::wmma_f16_16x16x16_f16::getOperationName();
1153 if (elemSourceType.isBF16() && elemDestType.isBF16())
1154 return ROCDL::wmma_bf16_16x16x16_bf16::getOperationName();
1155 if (elemSourceType.isInteger(8) && elemDestType.isInteger(32))
1156 return ROCDL::wmma_i32_16x16x16_iu8::getOperationName();
1157
1158 // RDNA3 specific patterns.
1159 if (isRDNA3) {
1160 if (elemSourceType.isInteger(4) && elemDestType.isInteger(32))
1161 return ROCDL::wmma_i32_16x16x16_iu4::getOperationName();
1162 return std::nullopt;
1163 }
1164
1165 // RDNA4 specific patterns (fp8/bf8).
1166 if (isa<fp8>(elemSourceType) && isa<fp8>(elemBSourceType) &&
1167 elemDestType.isF32())
1168 return ROCDL::wmma_f32_16x16x16_fp8_fp8::getOperationName();
1169 if (isa<fp8>(elemSourceType) && isa<bf8>(elemBSourceType) &&
1170 elemDestType.isF32())
1171 return ROCDL::wmma_f32_16x16x16_fp8_bf8::getOperationName();
1172 if (isa<bf8>(elemSourceType) && isa<bf8>(elemBSourceType) &&
1173 elemDestType.isF32())
1174 return ROCDL::wmma_f32_16x16x16_bf8_bf8::getOperationName();
1175 if (isa<bf8>(elemSourceType) && isa<fp8>(elemBSourceType) &&
1176 elemDestType.isF32())
1177 return ROCDL::wmma_f32_16x16x16_bf8_fp8::getOperationName();
1178 if (elemSourceType.isInteger(4) && elemDestType.isInteger(32))
1179 return ROCDL::wmma_i32_16x16x16_iu4::getOperationName();
1180
1181 return std::nullopt;
1182 }
1183
1184 // Handle k == 32 for RDNA4.
1185 if (k == 32 && !isRDNA3) {
1186 if (elemSourceType.isInteger(4) && elemDestType.isInteger(32))
1187 return ROCDL::wmma_i32_16x16x32_iu4::getOperationName();
1188 }
1189
1190 return std::nullopt;
1191}
1192
1193/// Return the `rocdl` intrinsic corresponding to a WMMA operation `wmma`
1194/// for the gfx1250 architecture.
1195static std::optional<StringRef> wmmaOpToIntrinsicGfx1250(Type elemSourceType,
1196 Type elemBSourceType,
1197 Type elemDestType,
1198 uint32_t k) {
1199 using fp8 = Float8E4M3FNType;
1200 using bf8 = Float8E5M2Type;
1201
1202 if (k == 4) {
1203 if (elemSourceType.isF32() && elemDestType.isF32())
1204 return ROCDL::wmma_f32_16x16x4_f32::getOperationName();
1205
1206 return std::nullopt;
1207 }
1208
1209 if (k == 32) {
1210 if (elemSourceType.isF16() && elemDestType.isF32())
1211 return ROCDL::wmma_f32_16x16x32_f16::getOperationName();
1212 if (elemSourceType.isBF16() && elemDestType.isF32())
1213 return ROCDL::wmma_f32_16x16x32_bf16::getOperationName();
1214 if (elemSourceType.isF16() && elemDestType.isF16())
1215 return ROCDL::wmma_f16_16x16x32_f16::getOperationName();
1216 if (elemSourceType.isBF16() && elemDestType.isBF16())
1217 return ROCDL::wmma_bf16_16x16x32_bf16::getOperationName();
1218
1219 return std::nullopt;
1220 }
1221
1222 if (k == 64) {
1223 if (isa<fp8>(elemSourceType) && isa<fp8>(elemBSourceType)) {
1224 if (elemDestType.isF32())
1225 return ROCDL::wmma_f32_16x16x64_fp8_fp8::getOperationName();
1226 if (elemDestType.isF16())
1227 return ROCDL::wmma_f16_16x16x64_fp8_fp8::getOperationName();
1228 }
1229 if (isa<fp8>(elemSourceType) && isa<bf8>(elemBSourceType)) {
1230 if (elemDestType.isF32())
1231 return ROCDL::wmma_f32_16x16x64_fp8_bf8::getOperationName();
1232 if (elemDestType.isF16())
1233 return ROCDL::wmma_f16_16x16x64_fp8_bf8::getOperationName();
1234 }
1235 if (isa<bf8>(elemSourceType) && isa<bf8>(elemBSourceType)) {
1236 if (elemDestType.isF32())
1237 return ROCDL::wmma_f32_16x16x64_bf8_bf8::getOperationName();
1238 if (elemDestType.isF16())
1239 return ROCDL::wmma_f16_16x16x64_bf8_bf8::getOperationName();
1240 }
1241 if (isa<bf8>(elemSourceType) && isa<fp8>(elemBSourceType)) {
1242 if (elemDestType.isF32())
1243 return ROCDL::wmma_f32_16x16x64_bf8_fp8::getOperationName();
1244 if (elemDestType.isF16())
1245 return ROCDL::wmma_f16_16x16x64_bf8_fp8::getOperationName();
1246 }
1247 if (elemSourceType.isInteger(8) && elemDestType.isInteger(32))
1248 return ROCDL::wmma_i32_16x16x64_iu8::getOperationName();
1249
1250 return std::nullopt;
1251 }
1252
1253 if (k == 128) {
1254 if (isa<fp8>(elemSourceType) && isa<fp8>(elemBSourceType)) {
1255 if (elemDestType.isF32())
1256 return ROCDL::wmma_f32_16x16x128_fp8_fp8::getOperationName();
1257 if (elemDestType.isF16())
1258 return ROCDL::wmma_f16_16x16x128_fp8_fp8::getOperationName();
1259 }
1260 if (isa<fp8>(elemSourceType) && isa<bf8>(elemBSourceType)) {
1261 if (elemDestType.isF32())
1262 return ROCDL::wmma_f32_16x16x128_fp8_bf8::getOperationName();
1263 if (elemDestType.isF16())
1264 return ROCDL::wmma_f16_16x16x128_fp8_bf8::getOperationName();
1265 }
1266 if (isa<bf8>(elemSourceType) && isa<bf8>(elemBSourceType)) {
1267 if (elemDestType.isF32())
1268 return ROCDL::wmma_f32_16x16x128_bf8_bf8::getOperationName();
1269 if (elemDestType.isF16())
1270 return ROCDL::wmma_f16_16x16x128_bf8_bf8::getOperationName();
1271 }
1272 if (isa<bf8>(elemSourceType) && isa<fp8>(elemBSourceType)) {
1273 if (elemDestType.isF32())
1274 return ROCDL::wmma_f32_16x16x128_bf8_fp8::getOperationName();
1275 if (elemDestType.isF16())
1276 return ROCDL::wmma_f16_16x16x128_bf8_fp8::getOperationName();
1277 }
1278
1279 return std::nullopt;
1280 }
1281
1282 return std::nullopt;
1283}
1284
1285/// Returns the `rocdl` intrinsic corresponding to a SparseMFMA (smfmac)
1286/// operation if one exists. This includes checking to ensure the intrinsic is
1287/// supported on the architecture you are compiling for.
1288static std::optional<StringRef> smfmacOpToIntrinsic(SparseMFMAOp op,
1289 Chipset chipset) {
1290 bool isGfx950 = chipset >= kGfx950;
1291 auto isFp8 = [&](Type t) { return typeIsExpectedFp8ForChipset(chipset, t); };
1292 auto isBf8 = [&](Type t) { return typeIsExpectedBf8ForChipset(chipset, t); };
1293
1294 uint32_t m = op.getM(), n = op.getN(), k = op.getK();
1295 Type sourceAElem = getElementTypeOrSelf(op.getSourceA().getType());
1296 Type sourceBElem = getElementTypeOrSelf(op.getSourceB().getType());
1297 Type destElem = getElementTypeOrSelf(op.getDestC().getType());
1298
1299 if (m == 16 && n == 16 && k == 32) {
1300 if (sourceAElem.isF16() && sourceBElem.isF16() && destElem.isF32())
1301 return ROCDL::smfmac_f32_16x16x32_f16::getOperationName();
1302 if (sourceAElem.isBF16() && sourceBElem.isBF16() && destElem.isF32())
1303 return ROCDL::smfmac_f32_16x16x32_bf16::getOperationName();
1304 }
1305
1306 if (m == 16 && n == 16 && k == 64) {
1307 if (isGfx950) {
1308 if (sourceAElem.isF16() && sourceBElem.isF16() && destElem.isF32())
1309 return ROCDL::smfmac_f32_16x16x64_f16::getOperationName();
1310 if (sourceAElem.isBF16() && sourceBElem.isBF16() && destElem.isF32())
1311 return ROCDL::smfmac_f32_16x16x64_bf16::getOperationName();
1312 }
1313 if (sourceAElem.isInteger(8) && sourceBElem.isInteger(8) &&
1314 destElem.isInteger(32))
1315 return ROCDL::smfmac_i32_16x16x64_i8::getOperationName();
1316 if (isFp8(sourceAElem) && isFp8(sourceBElem) && destElem.isF32())
1317 return ROCDL::smfmac_f32_16x16x64_fp8_fp8::getOperationName();
1318 if (isFp8(sourceAElem) && isBf8(sourceBElem) && destElem.isF32())
1319 return ROCDL::smfmac_f32_16x16x64_fp8_bf8::getOperationName();
1320 if (isBf8(sourceAElem) && isFp8(sourceBElem) && destElem.isF32())
1321 return ROCDL::smfmac_f32_16x16x64_bf8_fp8::getOperationName();
1322 if (isBf8(sourceAElem) && isBf8(sourceBElem) && destElem.isF32())
1323 return ROCDL::smfmac_f32_16x16x64_bf8_bf8::getOperationName();
1324 }
1325
1326 if (m == 16 && n == 16 && k == 128 && isGfx950) {
1327 if (sourceAElem.isInteger(8) && sourceBElem.isInteger(8) &&
1328 destElem.isInteger(32))
1329 return ROCDL::smfmac_i32_16x16x128_i8::getOperationName();
1330 if (isFp8(sourceAElem) && isFp8(sourceBElem) && destElem.isF32())
1331 return ROCDL::smfmac_f32_16x16x128_fp8_fp8::getOperationName();
1332 if (isFp8(sourceAElem) && isBf8(sourceBElem) && destElem.isF32())
1333 return ROCDL::smfmac_f32_16x16x128_fp8_bf8::getOperationName();
1334 if (isBf8(sourceAElem) && isFp8(sourceBElem) && destElem.isF32())
1335 return ROCDL::smfmac_f32_16x16x128_bf8_fp8::getOperationName();
1336 if (isBf8(sourceAElem) && isBf8(sourceBElem) && destElem.isF32())
1337 return ROCDL::smfmac_f32_16x16x128_bf8_bf8::getOperationName();
1338 }
1339
1340 if (m == 32 && n == 32 && k == 16) {
1341 if (sourceAElem.isF16() && sourceBElem.isF16() && destElem.isF32())
1342 return ROCDL::smfmac_f32_32x32x16_f16::getOperationName();
1343 if (sourceAElem.isBF16() && sourceBElem.isBF16() && destElem.isF32())
1344 return ROCDL::smfmac_f32_32x32x16_bf16::getOperationName();
1345 }
1346
1347 if (m == 32 && n == 32 && k == 32) {
1348 if (isGfx950) {
1349 if (sourceAElem.isF16() && sourceBElem.isF16() && destElem.isF32())
1350 return ROCDL::smfmac_f32_32x32x32_f16::getOperationName();
1351 if (sourceAElem.isBF16() && sourceBElem.isBF16() && destElem.isF32())
1352 return ROCDL::smfmac_f32_32x32x32_bf16::getOperationName();
1353 }
1354 if (sourceAElem.isInteger(8) && sourceBElem.isInteger(8) &&
1355 destElem.isInteger(32))
1356 return ROCDL::smfmac_i32_32x32x32_i8::getOperationName();
1357 if (isFp8(sourceAElem) && isFp8(sourceBElem) && destElem.isF32())
1358 return ROCDL::smfmac_f32_32x32x32_fp8_fp8::getOperationName();
1359 if (isFp8(sourceAElem) && isBf8(sourceBElem) && destElem.isF32())
1360 return ROCDL::smfmac_f32_32x32x32_fp8_bf8::getOperationName();
1361 if (isBf8(sourceAElem) && isFp8(sourceBElem) && destElem.isF32())
1362 return ROCDL::smfmac_f32_32x32x32_bf8_fp8::getOperationName();
1363 if (isBf8(sourceAElem) && isBf8(sourceBElem) && destElem.isF32())
1364 return ROCDL::smfmac_f32_32x32x32_bf8_bf8::getOperationName();
1365 }
1366
1367 if (m == 32 && n == 32 && k == 64 && isGfx950) {
1368 if (sourceAElem.isInteger(8) && sourceBElem.isInteger(8) &&
1369 destElem.isInteger(32))
1370 return ROCDL::smfmac_i32_32x32x64_i8::getOperationName();
1371 if (isFp8(sourceAElem) && isFp8(sourceBElem) && destElem.isF32())
1372 return ROCDL::smfmac_f32_32x32x64_fp8_fp8::getOperationName();
1373 if (isFp8(sourceAElem) && isBf8(sourceBElem) && destElem.isF32())
1374 return ROCDL::smfmac_f32_32x32x64_fp8_bf8::getOperationName();
1375 if (isBf8(sourceAElem) && isFp8(sourceBElem) && destElem.isF32())
1376 return ROCDL::smfmac_f32_32x32x64_bf8_fp8::getOperationName();
1377 if (isBf8(sourceAElem) && isBf8(sourceBElem) && destElem.isF32())
1378 return ROCDL::smfmac_f32_32x32x64_bf8_bf8::getOperationName();
1379 }
1380
1381 return std::nullopt;
1382}
1383
1384/// Returns the `rocdl` intrinsic corresponding to a WMMA operation `wmma`
1385/// if one exists. This includes checking to ensure the intrinsic is supported
1386/// on the architecture you are compiling for.
1387static std::optional<StringRef> wmmaOpToIntrinsic(WMMAOp wmma,
1388 Chipset chipset) {
1389 auto sourceVectorType = cast<VectorType>(wmma.getSourceA().getType());
1390 auto sourceBVectorType = cast<VectorType>(wmma.getSourceB().getType());
1391 auto destVectorType = cast<VectorType>(wmma.getDestC().getType());
1392 Type elemSourceType = sourceVectorType.getElementType();
1393 Type elemBSourceType = sourceBVectorType.getElementType();
1394 Type elemDestType = destVectorType.getElementType();
1395
1396 const uint32_t k = wmma.getK();
1397 const bool isRDNA3 = chipset.majorVersion == 11;
1398 const bool isRDNA4 = chipset.majorVersion == 12 && chipset.minorVersion == 0;
1399
1400 // Handle RDNA3 and RDNA4.
1401 if (isRDNA3 || isRDNA4)
1402 return wmmaOpToIntrinsicRDNA(elemSourceType, elemBSourceType, elemDestType,
1403 k, isRDNA3);
1404
1405 // Handle gfx1250.
1406 if (chipset == kGfx1250)
1407 return wmmaOpToIntrinsicGfx1250(elemSourceType, elemBSourceType,
1408 elemDestType, k);
1409
1410 return std::nullopt;
1411}
1412
1413/// Returns the `rocdl` intrinsic corresponding to a SparseWMMA operation
1414/// `swmmac` if one exists. This includes checking to ensure the intrinsic is
1415/// supported on the architecture you are compiling for.
1417 StringRef name;
1421};
1422
1423static std::optional<SparseWMMAOpInfo>
1424sparseWMMAOpToIntrinsic(SparseWMMAOp swmmac, Chipset chipset) {
1425 Type sourceAElem = getElementTypeOrSelf(swmmac.getSourceA().getType());
1426 Type sourceBElem = getElementTypeOrSelf(swmmac.getSourceB().getType());
1427 Type destElem = getElementTypeOrSelf(swmmac.getDestC().getType());
1428
1429 uint32_t m = swmmac.getM(), n = swmmac.getN(), k = swmmac.getK();
1430
1431 if ((m != 16) || (n != 16))
1432 return std::nullopt;
1433
1434 const bool isRDNA4 = chipset.majorVersion == 12 && chipset.minorVersion == 0;
1435 if (isRDNA4) {
1436 if (k == 32) {
1437 if (destElem.isF32() && sourceAElem.isF16() && sourceBElem.isF16())
1438 return SparseWMMAOpInfo{
1439 ROCDL::swmmac_f32_16x16x32_f16::getOperationName(), false, false,
1440 false};
1441 if (destElem.isF32() && sourceAElem.isBF16() && sourceBElem.isBF16())
1442 return SparseWMMAOpInfo{
1443 ROCDL::swmmac_f32_16x16x32_bf16::getOperationName(), false, false,
1444 false};
1445 if (destElem.isF16() && sourceAElem.isF16() && sourceBElem.isF16())
1446 return SparseWMMAOpInfo{
1447 ROCDL::swmmac_f16_16x16x32_f16::getOperationName(), false, false,
1448 false};
1449 if (destElem.isBF16() && sourceAElem.isBF16() && sourceBElem.isBF16())
1450 return SparseWMMAOpInfo{
1451 ROCDL::swmmac_bf16_16x16x32_bf16::getOperationName(), false, false,
1452 false};
1453 if (destElem.isInteger(32) && sourceAElem.isInteger(8) &&
1454 sourceBElem.isInteger(8))
1455 return SparseWMMAOpInfo{
1456 ROCDL::swmmac_i32_16x16x32_iu8::getOperationName(), true, false,
1457 true};
1458 if (destElem.isInteger(32) && sourceAElem.isInteger(4) &&
1459 sourceBElem.isInteger(4))
1460 return SparseWMMAOpInfo{
1461 ROCDL::swmmac_i32_16x16x32_iu4::getOperationName(), true, false,
1462 true};
1463 if (destElem.isF32() && sourceAElem.isF8E4M3FN() &&
1464 sourceBElem.isF8E4M3FN())
1465 return SparseWMMAOpInfo{
1466 ROCDL::swmmac_f32_16x16x32_fp8_fp8::getOperationName(), false,
1467 false, false};
1468 if (destElem.isF32() && sourceAElem.isF8E4M3FN() &&
1469 sourceBElem.isF8E5M2())
1470 return SparseWMMAOpInfo{
1471 ROCDL::swmmac_f32_16x16x32_fp8_bf8::getOperationName(), false,
1472 false, false};
1473 if (destElem.isF32() && sourceAElem.isF8E5M2() &&
1474 sourceBElem.isF8E4M3FN())
1475 return SparseWMMAOpInfo{
1476 ROCDL::swmmac_f32_16x16x32_bf8_fp8::getOperationName(), false,
1477 false, false};
1478 if (destElem.isF32() && sourceAElem.isF8E5M2() && sourceBElem.isF8E5M2())
1479 return SparseWMMAOpInfo{
1480 ROCDL::swmmac_f32_16x16x32_bf8_bf8::getOperationName(), false,
1481 false, false};
1482 }
1483 if (k == 64) {
1484 if (destElem.isInteger(32) && sourceAElem.isInteger(4) &&
1485 sourceBElem.isInteger(4))
1486 return SparseWMMAOpInfo{
1487 ROCDL::swmmac_i32_16x16x64_iu4::getOperationName(), true, false,
1488 true};
1489 }
1490 }
1491
1492 const bool isGFX1250 = chipset == kGfx1250;
1493 const bool isWavesize64 = swmmac.getWave64();
1494 if (isGFX1250 && !isWavesize64) {
1495 if (k == 64) {
1496 if (destElem.isF32() && sourceAElem.isF16() && sourceBElem.isF16())
1497 return SparseWMMAOpInfo{
1498 ROCDL::swmmac_f32_16x16x64_f16::getOperationName(), true, true,
1499 false};
1500 if (destElem.isF32() && sourceAElem.isBF16() && sourceBElem.isBF16())
1501 return SparseWMMAOpInfo{
1502 ROCDL::swmmac_f32_16x16x64_bf16::getOperationName(), true, true,
1503 false};
1504 if (destElem.isF16() && sourceAElem.isF16() && sourceBElem.isF16())
1505 return SparseWMMAOpInfo{
1506 ROCDL::swmmac_f16_16x16x64_f16::getOperationName(), true, true,
1507 false};
1508 if (destElem.isBF16() && sourceAElem.isBF16() && sourceBElem.isBF16())
1509 return SparseWMMAOpInfo{
1510 ROCDL::swmmac_bf16_16x16x64_bf16::getOperationName(), true, true,
1511 false};
1512 }
1513 if (k == 128) {
1514 if (destElem.isF32() && sourceAElem.isF8E4M3FN() &&
1515 sourceBElem.isF8E4M3FN())
1516 return SparseWMMAOpInfo{
1517 ROCDL::swmmac_f32_16x16x128_fp8_fp8::getOperationName(), false,
1518 true, false};
1519 if (destElem.isF32() && sourceAElem.isF8E4M3FN() &&
1520 sourceBElem.isF8E5M2())
1521 return SparseWMMAOpInfo{
1522 ROCDL::swmmac_f32_16x16x128_fp8_bf8::getOperationName(), false,
1523 true, false};
1524 if (destElem.isF32() && sourceAElem.isF8E5M2() &&
1525 sourceBElem.isF8E4M3FN())
1526 return SparseWMMAOpInfo{
1527 ROCDL::swmmac_f32_16x16x128_bf8_fp8::getOperationName(), false,
1528 true, false};
1529 if (destElem.isF32() && sourceAElem.isF8E5M2() && sourceBElem.isF8E5M2())
1530 return SparseWMMAOpInfo{
1531 ROCDL::swmmac_f32_16x16x128_bf8_bf8::getOperationName(), false,
1532 true, false};
1533 if (destElem.isF16() && sourceAElem.isF8E4M3FN() &&
1534 sourceBElem.isF8E4M3FN())
1535 return SparseWMMAOpInfo{
1536 ROCDL::swmmac_f16_16x16x128_fp8_fp8::getOperationName(), false,
1537 true, false};
1538 if (destElem.isF16() && sourceAElem.isF8E4M3FN() &&
1539 sourceBElem.isF8E5M2())
1540 return SparseWMMAOpInfo{
1541 ROCDL::swmmac_f16_16x16x128_fp8_bf8::getOperationName(), false,
1542 true, false};
1543 if (destElem.isF16() && sourceAElem.isF8E5M2() &&
1544 sourceBElem.isF8E4M3FN())
1545 return SparseWMMAOpInfo{
1546 ROCDL::swmmac_f16_16x16x128_bf8_fp8::getOperationName(), false,
1547 true, false};
1548 if (destElem.isF16() && sourceAElem.isF8E5M2() && sourceBElem.isF8E5M2())
1549 return SparseWMMAOpInfo{
1550 ROCDL::swmmac_f16_16x16x128_bf8_bf8::getOperationName(), false,
1551 true, false};
1552 if (destElem.isF16() && sourceAElem.isInteger(8) &&
1553 sourceBElem.isInteger(8))
1554 return SparseWMMAOpInfo{
1555 ROCDL::swmmac_f16_16x16x128_bf8_bf8::getOperationName(), false,
1556 true, false};
1557 if (destElem.isInteger(32) && sourceAElem.isInteger(8) &&
1558 sourceBElem.isInteger(8))
1559 return SparseWMMAOpInfo{
1560 ROCDL::swmmac_i32_16x16x128_iu8::getOperationName(), true, true,
1561 true};
1562 }
1563 }
1564
1565 return std::nullopt;
1566}
1567
1568namespace {
1569struct MFMAOpLowering : public ConvertOpToLLVMPattern<MFMAOp> {
1570 MFMAOpLowering(const LLVMTypeConverter &converter, Chipset chipset)
1571 : ConvertOpToLLVMPattern<MFMAOp>(converter), chipset(chipset) {}
1572
1573 Chipset chipset;
1574
1575 LogicalResult
1576 matchAndRewrite(MFMAOp op, MFMAOpAdaptor adaptor,
1577 ConversionPatternRewriter &rewriter) const override {
1578 Location loc = op.getLoc();
1579 Type destElem = getElementTypeOrSelf(op.getDestD().getType());
1580 Type outType = typeConverter->convertType(op.getDestD().getType());
1581 Type intrinsicOutType = outType;
1582 if (auto outVecType = dyn_cast<VectorType>(outType))
1583 if (outVecType.getElementType().isBF16())
1584 intrinsicOutType = outVecType.clone(rewriter.getI16Type());
1585
1586 if (chipset.majorVersion != 9 || chipset < kGfx908)
1587 return op->emitOpError("MFMA only supported on gfx908+");
1588 uint32_t getBlgpField = static_cast<uint32_t>(op.getBlgp());
1589 if (op.getNegateA() || op.getNegateB() || op.getNegateC()) {
1590 if (chipset < kGfx942)
1591 return op.emitOpError("negation unsupported on older than gfx942");
1592 getBlgpField |=
1593 op.getNegateA() | (op.getNegateB() << 1) | (op.getNegateC() << 2);
1594 }
1595 std::optional<StringRef> maybeIntrinsic = mfmaOpToIntrinsic(op, chipset);
1596 std::optional<ScaledMFMAIntrinsic> maybeScaledIntrinsic =
1597 mfmaOpToScaledIntrinsic(op, chipset);
1598 if (!maybeIntrinsic.has_value() && !maybeScaledIntrinsic.has_value())
1599 return op.emitOpError("no intrinsic matching MFMA size on given chipset");
1600
1601 bool isScaled =
1602 !maybeIntrinsic.has_value() && maybeScaledIntrinsic.has_value();
1603 if (isScaled &&
1604 (adaptor.getAbid() > 0 || getBlgpField > 0 || op.getCbsz() > 0)) {
1605 return op.emitOpError(
1606 "non-default abid, blgp, and cbsz aren't supported on MFMAs that can "
1607 "be scaled as those fields are used for type information");
1608 }
1609
1610 StringRef intrinsicName =
1611 isScaled ? std::get<0>(*maybeScaledIntrinsic) : *maybeIntrinsic;
1612 // Determine if we can use bf16 in the intrinsic. Newer MFMAs in gfx950+
1613 // allows bf16 as the input. For reference check IntrinsicsAMDGPU.td file.
1614 bool allowBf16 = [&]() {
1615 if (chipset < kGfx950)
1616 return false;
1617 if (isScaled)
1618 return true;
1619 return intrinsicName.contains("16x16x32.bf16") ||
1620 intrinsicName.contains("32x32x16.bf16");
1621 }();
1622 OperationState loweredOp(loc, intrinsicName);
1623 loweredOp.addTypes(intrinsicOutType);
1624 loweredOp.addOperands({packSmallFloatVectorOperand(
1625 rewriter, loc, adaptor.getSourceA(), allowBf16),
1627 rewriter, loc, adaptor.getSourceB(), allowBf16),
1628 adaptor.getDestC()});
1629 if (isScaled) {
1630 Value zero = createI32Constant(rewriter, loc, 0);
1631 auto [_scaledName, aTypeCode, bTypeCode] = *maybeScaledIntrinsic;
1632 loweredOp.addOperands({/*scale A=*/zero, /*scale B=*/zero});
1633 loweredOp.addAttributes(
1634 {{"cbsz",
1635 ROCDL::MatrixFormatAttr::get(rewriter.getContext(), aTypeCode)},
1636 {"blgp",
1637 ROCDL::MatrixFormatAttr::get(rewriter.getContext(), bTypeCode)},
1638 {"opselA", rewriter.getI32IntegerAttr(0)},
1639 {"opselB", rewriter.getI32IntegerAttr(0)}});
1640 } else {
1641 Attribute blgpAttr =
1642 destElem.isF64()
1643 ? Attribute(ROCDL::MFMANegModifierAttr::get(
1644 rewriter.getContext(),
1645 static_cast<ROCDL::MFMANegModifier>(getBlgpField)))
1646 : Attribute(ROCDL::MFMAPermBAttr::get(
1647 rewriter.getContext(),
1648 static_cast<ROCDL::MFMAPermB>(getBlgpField)));
1649 loweredOp.addAttributes(
1650 {{"cbsz", rewriter.getI32IntegerAttr(op.getCbsz())},
1651 {"abid", rewriter.getI32IntegerAttr(op.getAbid())},
1652 {"blgp", blgpAttr}});
1653 };
1654 Value lowered = rewriter.create(loweredOp)->getResult(0);
1655 if (outType != intrinsicOutType)
1656 lowered = LLVM::BitcastOp::create(rewriter, loc, outType, lowered);
1657 rewriter.replaceOp(op, lowered);
1658 return success();
1659 }
1660};
1661
1662struct ScaledMFMAOpLowering : public ConvertOpToLLVMPattern<ScaledMFMAOp> {
1663 ScaledMFMAOpLowering(const LLVMTypeConverter &converter, Chipset chipset)
1664 : ConvertOpToLLVMPattern(converter), chipset(chipset) {}
1665
1666 Chipset chipset;
1667
1668 LogicalResult
1669 matchAndRewrite(ScaledMFMAOp op, ScaledMFMAOpAdaptor adaptor,
1670 ConversionPatternRewriter &rewriter) const override {
1671 Location loc = op.getLoc();
1672 Type intrinsicOutType = typeConverter->convertType(op.getDestD().getType());
1673
1674 if (chipset.majorVersion != 9 || chipset < kGfx950)
1675 return op->emitOpError("scaled MFMA only supported on gfx908+");
1676 std::optional<ScaledMFMAIntrinsic> maybeScaledIntrinsic =
1677 mfmaOpToScaledIntrinsic(op, chipset);
1678 if (!maybeScaledIntrinsic.has_value())
1679 return op.emitOpError(
1680 "no intrinsic matching scaled MFMA size on given chipset");
1681
1682 auto [intrinsicName, aTypeCode, bTypeCode] = *maybeScaledIntrinsic;
1683 OperationState loweredOp(loc, intrinsicName);
1684 loweredOp.addTypes(intrinsicOutType);
1685 loweredOp.addOperands(
1686 {packSmallFloatVectorOperand(rewriter, loc, adaptor.getSourceA()),
1687 packSmallFloatVectorOperand(rewriter, loc, adaptor.getSourceB()),
1688 adaptor.getDestC()});
1689 loweredOp.addOperands(
1690 {/*scales A*/
1691 castScaleOperand(rewriter, loc, adaptor.getScalesA()),
1692 /*scales B*/
1693 castScaleOperand(rewriter, loc, adaptor.getScalesB())});
1694 loweredOp.addAttributes(
1695 {{"cbsz",
1696 ROCDL::MatrixFormatAttr::get(rewriter.getContext(), aTypeCode)},
1697 {"blgp",
1698 ROCDL::MatrixFormatAttr::get(rewriter.getContext(), bTypeCode)},
1699 {"opselA", rewriter.getI32IntegerAttr(adaptor.getScalesIdxA())},
1700 {"opselB", rewriter.getI32IntegerAttr(adaptor.getScalesIdxB())}});
1701
1702 Value lowered = rewriter.create(loweredOp)->getResult(0);
1703 rewriter.replaceOp(op, lowered);
1704 return success();
1705 }
1706};
1707
1708struct SparseMFMAOpLowering : public ConvertOpToLLVMPattern<SparseMFMAOp> {
1709 SparseMFMAOpLowering(const LLVMTypeConverter &converter, Chipset chipset)
1710 : ConvertOpToLLVMPattern<SparseMFMAOp>(converter), chipset(chipset) {}
1711
1712 Chipset chipset;
1713
1714 LogicalResult
1715 matchAndRewrite(SparseMFMAOp op, SparseMFMAOpAdaptor adaptor,
1716 ConversionPatternRewriter &rewriter) const override {
1717 Location loc = op.getLoc();
1718 auto outType =
1719 typeConverter->convertType<VectorType>(op.getDestC().getType());
1720 if (!outType)
1721 return rewriter.notifyMatchFailure(op, "type conversion failed");
1722
1723 // smfmac is supported on gfx942 and gfx950.
1724 if (chipset.majorVersion != 9 || chipset < kGfx942)
1725 return op->emitOpError("sparse MFMA (smfmac) only supported on gfx942+");
1726
1727 std::optional<StringRef> maybeIntrinsic = smfmacOpToIntrinsic(op, chipset);
1728 if (!maybeIntrinsic.has_value())
1729 return op.emitOpError(
1730 "no intrinsic matching sparse MFMA on the given chipset");
1731 bool isGfx942BF16 =
1732 (*maybeIntrinsic ==
1733 ROCDL::smfmac_f32_16x16x32_bf16::getOperationName() ||
1734 *maybeIntrinsic ==
1735 ROCDL::smfmac_f32_32x32x16_bf16::getOperationName());
1736 bool isGfx950 = (chipset >= kGfx950) && !isGfx942BF16;
1737
1738 Value a = convertPackedVectorOperand(rewriter, loc, adaptor.getSourceA(),
1739 isGfx950);
1740 Value b = convertPackedVectorOperand(rewriter, loc, adaptor.getSourceB(),
1741 isGfx950);
1742 Value c = adaptor.getDestC();
1743
1744 // Bitcast sparse indices from vector<4xi8> or vector<2xi16> to i32.
1745 // gfx950 8-bit variants already carry the index as i32; skip the bitcast.
1746 Value sparseIdx = adaptor.getSparseIdx();
1747 Type i32Type = rewriter.getI32Type();
1748 if (sparseIdx.getType() != i32Type)
1749 sparseIdx = LLVM::BitcastOp::create(rewriter, loc, i32Type, sparseIdx);
1750
1751 OperationState loweredOp(loc, maybeIntrinsic.value());
1752 loweredOp.addTypes(outType);
1753 loweredOp.addOperands({a, b, c, sparseIdx});
1754 loweredOp.addAttributes(
1755 {{"cbsz", rewriter.getI32IntegerAttr(op.getCbsz())},
1756 {"abid", rewriter.getI32IntegerAttr(op.getAbid())}});
1757 Value lowered = rewriter.create(loweredOp)->getResult(0);
1758 rewriter.replaceOp(op, lowered);
1759 return success();
1760 }
1761};
1762
1763struct WMMAOpLowering : public ConvertOpToLLVMPattern<WMMAOp> {
1764 WMMAOpLowering(const LLVMTypeConverter &converter, Chipset chipset)
1765 : ConvertOpToLLVMPattern<WMMAOp>(converter), chipset(chipset) {}
1766
1767 Chipset chipset;
1768
1769 LogicalResult
1770 matchAndRewrite(WMMAOp op, WMMAOpAdaptor adaptor,
1771 ConversionPatternRewriter &rewriter) const override {
1772 Location loc = op.getLoc();
1773 auto outType =
1774 typeConverter->convertType<VectorType>(op.getDestD().getType());
1775 if (!outType)
1776 return rewriter.notifyMatchFailure(op, "type conversion failed");
1777
1778 if (chipset.majorVersion != 11 && chipset.majorVersion != 12)
1779 return op->emitOpError("WMMA only supported on gfx11 and gfx12");
1780
1781 bool isGFX1250 = chipset >= kGfx1250;
1782
1783 // The WMMA operations represent vectors of bf16s as vectors of i16s
1784 // (except on gfx1250), so we need to bitcast bfloats to i16 and then
1785 // bitcast them back.
1786 auto aType = cast<VectorType>(adaptor.getSourceA().getType());
1787 auto bType = cast<VectorType>(adaptor.getSourceB().getType());
1788 auto destCType = cast<VectorType>(adaptor.getDestC().getType());
1789 bool castAToI16 = aType.getElementType().isBF16() && !isGFX1250;
1790 bool castBToI16 = bType.getElementType().isBF16() && !isGFX1250;
1791 bool castDestCToI16 = destCType.getElementType().isBF16() && !isGFX1250;
1792 bool castOutToI16 = outType.getElementType().isBF16() && !isGFX1250;
1793 VectorType rawOutType = outType;
1794 if (castOutToI16)
1795 rawOutType = outType.clone(rewriter.getI16Type());
1796 Value a = adaptor.getSourceA();
1797 if (castAToI16)
1798 a = LLVM::BitcastOp::create(rewriter, loc,
1799 aType.clone(rewriter.getI16Type()), a);
1800 Value b = adaptor.getSourceB();
1801 if (castBToI16)
1802 b = LLVM::BitcastOp::create(rewriter, loc,
1803 bType.clone(rewriter.getI16Type()), b);
1804 Value destC = adaptor.getDestC();
1805 if (castDestCToI16)
1806 destC = LLVM::BitcastOp::create(
1807 rewriter, loc, destCType.clone(rewriter.getI16Type()), destC);
1808
1809 std::optional<StringRef> maybeIntrinsic = wmmaOpToIntrinsic(op, chipset);
1810
1811 if (!maybeIntrinsic.has_value())
1812 return op.emitOpError("no intrinsic matching WMMA on the given chipset");
1813
1814 if (chipset.majorVersion >= 12 && op.getSubwordOffset() != 0)
1815 return op.emitOpError("subwordOffset not supported on gfx12+");
1816
1817 SmallVector<Value, 4> operands;
1818 SmallVector<NamedAttribute, 4> attrs;
1819 wmmaPushInputOperand(rewriter, loc, typeConverter, op.getUnsignedA(), a,
1820 op.getSourceA(), operands, attrs, "signA");
1821 wmmaPushInputOperand(rewriter, loc, typeConverter, op.getUnsignedB(), b,
1822 op.getSourceB(), operands, attrs, "signB");
1823 wmmaPushOutputOperand(rewriter, loc, typeConverter, destC,
1824 op.getSubwordOffset(), op.getClamp(), operands,
1825 attrs);
1826
1827 OperationState loweredOp(loc, *maybeIntrinsic);
1828 loweredOp.addTypes(rawOutType);
1829 loweredOp.addOperands(operands);
1830 loweredOp.addAttributes(attrs);
1831 Operation *lowered = rewriter.create(loweredOp);
1832
1833 Operation *maybeCastBack = lowered;
1834 if (rawOutType != outType)
1835 maybeCastBack = LLVM::BitcastOp::create(rewriter, loc, outType,
1836 lowered->getResult(0));
1837 rewriter.replaceOp(op, maybeCastBack->getResults());
1838
1839 return success();
1840 }
1841};
1842
1843enum class DotFamily {
1844 /// ROCDL_Dot_IntrOp: single `clamp` attribute.
1845 Clamp,
1846 /// ROCDL_Dot_NoClamp_IntrOp: no attributes.
1847 NoClamp,
1848 /// ROCDL_Sudot_IntrOp: `signA`, `signB`, and `clamp` attributes.
1849 Sudot,
1850};
1851
1852static std::optional<std::pair<StringRef, DotFamily>>
1853dotOpToIntrinsic(DotOp op, Chipset chipset) {
1854 Type aElem = cast<VectorType>(op.getSourceA().getType()).getElementType();
1855 Type bElem = cast<VectorType>(op.getSourceB().getType()).getElementType();
1856 Type dest = op.getDestC().getType();
1857 bool uA = op.getUnsignedA();
1858 bool uB = op.getUnsignedB();
1859
1860 // f16 x f16 -> f32 / f16.
1861 if (aElem.isF16() && bElem.isF16()) {
1862 if (dest.isF32() && hasDot10Insts(chipset))
1863 return {{ROCDL::fdot2::getOperationName(), DotFamily::Clamp}};
1864 if (dest.isF16() && hasDot9Insts(chipset))
1865 return {{ROCDL::fdot2_f16_f16::getOperationName(), DotFamily::NoClamp}};
1866 return std::nullopt;
1867 }
1868
1869 // bf16 x bf16 -> f32 / bf16.
1870 if (aElem.isBF16() && bElem.isBF16()) {
1871 if (dest.isF32() && hasDot12Insts(chipset))
1872 return {{ROCDL::fdot2_f32_bf16::getOperationName(), DotFamily::Clamp}};
1873 if (dest.isBF16() && hasDot9Insts(chipset))
1874 return {{ROCDL::fdot2_bf16_bf16::getOperationName(), DotFamily::NoClamp}};
1875 return std::nullopt;
1876 }
1877
1878 // Integer sources -> i32.
1879 if (isa<IntegerType>(aElem) && isa<IntegerType>(bElem) &&
1880 dest.isInteger(32)) {
1881 bool mixedSign = (uA != uB);
1882 unsigned elemWidth = aElem.getIntOrFloatBitWidth();
1883
1884 if (mixedSign) {
1885 if (!hasDot8Insts(chipset))
1886 return std::nullopt;
1887 StringRef name;
1888 switch (elemWidth) {
1889 case 8:
1890 name = ROCDL::sudot4::getOperationName();
1891 break;
1892 case 4:
1893 name = ROCDL::sudot8::getOperationName();
1894 break;
1895 default:
1896 return std::nullopt;
1897 }
1898 return {{name, DotFamily::Sudot}};
1899 }
1900
1901 StringRef name;
1902 bool supported = false;
1903 switch (elemWidth) {
1904 case 16:
1905 supported = hasDot2Insts(chipset);
1906 name = uA ? ROCDL::udot2::getOperationName()
1907 : ROCDL::sdot2::getOperationName();
1908 break;
1909 case 8:
1910 supported = uA ? hasDot7Insts(chipset)
1911 : hasDot1Insts(chipset) || hasDot8Insts(chipset);
1912 name = uA ? ROCDL::udot4::getOperationName()
1913 : ROCDL::sdot4::getOperationName();
1914 break;
1915 case 4:
1916 supported = uA ? hasDot7Insts(chipset)
1917 : hasDot1Insts(chipset) || hasDot8Insts(chipset);
1918 name = uA ? ROCDL::udot8::getOperationName()
1919 : ROCDL::sdot8::getOperationName();
1920 break;
1921 default:
1922 return std::nullopt;
1923 }
1924 if (!supported)
1925 return std::nullopt;
1926 return {{name, DotFamily::Clamp}};
1927 }
1928
1929 // fp8/bf8 x fp8/bf8 -> f32.
1930 bool aIsFp8 = isa<Float8E4M3FNType>(aElem);
1931 bool aIsBf8 = isa<Float8E5M2Type>(aElem);
1932 bool bIsFp8 = isa<Float8E4M3FNType>(bElem);
1933 bool bIsBf8 = isa<Float8E5M2Type>(bElem);
1934 if ((aIsFp8 || aIsBf8) && (bIsFp8 || bIsBf8) && dest.isF32()) {
1935 if (!hasDot11Insts(chipset))
1936 return std::nullopt;
1937 StringRef name;
1938 if (aIsFp8 && bIsFp8)
1939 name = ROCDL::dot4_f32_fp8_fp8::getOperationName();
1940 else if (aIsFp8 && bIsBf8)
1941 name = ROCDL::dot4_f32_fp8_bf8::getOperationName();
1942 else if (aIsBf8 && bIsFp8)
1943 name = ROCDL::dot4_f32_bf8_fp8::getOperationName();
1944 else
1945 name = ROCDL::dot4_f32_bf8_bf8::getOperationName();
1946 return {{name, DotFamily::NoClamp}};
1947 }
1948
1949 return std::nullopt;
1950}
1951
1952struct DotOpLowering : public ConvertOpToLLVMPattern<DotOp> {
1953 DotOpLowering(const LLVMTypeConverter &converter, Chipset chipset)
1954 : ConvertOpToLLVMPattern<DotOp>(converter), chipset(chipset) {}
1955
1956 Chipset chipset;
1957
1958 LogicalResult
1959 matchAndRewrite(DotOp op, DotOpAdaptor adaptor,
1960 ConversionPatternRewriter &rewriter) const override {
1961 Location loc = op.getLoc();
1962
1963 std::optional<std::pair<StringRef, DotFamily>> maybeIntrinsic =
1964 dotOpToIntrinsic(op, chipset);
1965 if (!maybeIntrinsic)
1966 return op.emitOpError("no intrinsic matching dot on the given chipset: ")
1967 << op.getSourceA().getType() << " * " << op.getSourceB().getType()
1968 << " + " << op.getDestC().getType();
1969
1970 auto [intrinsicName, family] = maybeIntrinsic.value();
1971
1972 Value a = convertPackedVectorOperand(rewriter, loc, adaptor.getSourceA());
1973 Value b = convertPackedVectorOperand(rewriter, loc, adaptor.getSourceB());
1974 Value c = adaptor.getDestC();
1975
1976 SmallVector<NamedAttribute, 3> attrs;
1977 if (family == DotFamily::Sudot) {
1978 attrs.push_back(rewriter.getNamedAttr(
1979 "signA", rewriter.getBoolAttr(!op.getUnsignedA())));
1980 attrs.push_back(rewriter.getNamedAttr(
1981 "signB", rewriter.getBoolAttr(!op.getUnsignedB())));
1982 }
1983
1984 if (family != DotFamily::NoClamp && op.getClamp())
1985 attrs.push_back(
1986 rewriter.getNamedAttr("clamp", rewriter.getBoolAttr(true)));
1987
1988 Type resultType = typeConverter->convertType(op.getDestD().getType());
1989
1990 OperationState loweredOp(loc, intrinsicName);
1991 loweredOp.addTypes(resultType);
1992 loweredOp.addOperands({a, b, c});
1993 loweredOp.addAttributes(attrs);
1994 Operation *lowered = rewriter.create(loweredOp);
1995 rewriter.replaceOp(op, lowered->getResults());
1996 return success();
1997 }
1998};
1999
2000struct SparseWMMAOpLowering : public ConvertOpToLLVMPattern<SparseWMMAOp> {
2001 SparseWMMAOpLowering(const LLVMTypeConverter &converter, Chipset chipset)
2002 : ConvertOpToLLVMPattern<SparseWMMAOp>(converter), chipset(chipset) {}
2003
2004 Chipset chipset;
2005
2006 LogicalResult
2007 matchAndRewrite(SparseWMMAOp op, SparseWMMAOpAdaptor adaptor,
2008 ConversionPatternRewriter &rewriter) const override {
2009 Location loc = op.getLoc();
2010 auto outType =
2011 typeConverter->convertType<VectorType>(op.getDestD().getType());
2012 if (!outType)
2013 return rewriter.notifyMatchFailure(op, "type conversion failed");
2014
2015 std::optional<SparseWMMAOpInfo> maybeIntrinsic =
2016 sparseWMMAOpToIntrinsic(op, chipset);
2017
2018 if (!maybeIntrinsic.has_value())
2019 return op.emitOpError(
2020 "no intrinsic matching Sparse WMMA on the given chipset");
2021 SparseWMMAOpInfo intrinsic = maybeIntrinsic.value();
2022
2023 SmallVector<NamedAttribute> attrs;
2024
2025 if ((op.getUnsignedA() || op.getUnsignedB()) && !intrinsic.useSign)
2026 return op->emitOpError("intrinsic doesn't support unsign");
2027 if (intrinsic.useSign) {
2028 if (auto attr = op.getUnsignedAAttr())
2029 attrs.push_back({"signA", attr});
2030 if (auto attr = op.getUnsignedBAttr())
2031 attrs.push_back({"signB", attr});
2032 }
2033
2034 if ((op.getReuseA() || op.getReuseB()) && !intrinsic.useReuse)
2035 return op->emitOpError("intrinsic doesn't support reuse");
2036 if (intrinsic.useReuse) {
2037 if (auto attr = op.getReuseAAttr())
2038 attrs.push_back({"reuseA", attr});
2039 if (auto attr = op.getReuseBAttr())
2040 attrs.push_back({"reuseB", attr});
2041 }
2042
2043 if (op.getClamp() && !intrinsic.useClamp)
2044 return op->emitOpError("intrinsic doesn't support clamp");
2045 if (intrinsic.useClamp && op.getClampAttr())
2046 attrs.push_back({"clamp", op.getClampAttr()});
2047
2048 const bool isGFX1250orHigher =
2049 chipset.majorVersion == 12 && chipset.minorVersion >= 5;
2050 Value a = convertPackedVectorOperand(rewriter, loc, adaptor.getSourceA(),
2051 isGFX1250orHigher);
2052 Value b = convertPackedVectorOperand(rewriter, loc, adaptor.getSourceB(),
2053 isGFX1250orHigher);
2054 Value c = adaptor.getDestC();
2055 VectorType rawOutType = outType;
2056 if (!isGFX1250orHigher) {
2057 c = convertPackedVectorOperand(rewriter, loc, adaptor.getDestC(), false);
2058 rawOutType = cast<VectorType>(c.getType());
2059 }
2060
2061 // Bitcast sparse indices from vector<4xi8> to i32.
2062 Value sparseIdx = LLVM::BitcastOp::create(
2063 rewriter, loc, rewriter.getI32Type(), adaptor.getSparseIdx());
2064
2065 OperationState loweredOp(loc, intrinsic.name);
2066 loweredOp.addTypes(rawOutType);
2067 loweredOp.addOperands({a, b, c, sparseIdx});
2068 loweredOp.addAttributes(attrs);
2069 Operation *lowered = rewriter.create(loweredOp);
2070
2071 Operation *maybeCastBack = lowered;
2072 if (rawOutType != outType)
2073 maybeCastBack = LLVM::BitcastOp::create(rewriter, loc, outType,
2074 lowered->getResult(0));
2075 rewriter.replaceOp(op, maybeCastBack->getResults());
2076
2077 return success();
2078 }
2079};
2080
2081struct ScaledWMMAOpLowering : public ConvertOpToLLVMPattern<ScaledWMMAOp> {
2082 ScaledWMMAOpLowering(const LLVMTypeConverter &converter, Chipset chipset)
2083 : ConvertOpToLLVMPattern<ScaledWMMAOp>(converter), chipset(chipset) {}
2084
2085 Chipset chipset;
2086
2087 LogicalResult
2088 matchAndRewrite(ScaledWMMAOp op, ScaledWMMAOpAdaptor adaptor,
2089 ConversionPatternRewriter &rewriter) const override {
2090 Location loc = op.getLoc();
2091 auto outType =
2092 typeConverter->convertType<VectorType>(op.getDestD().getType());
2093 if (!outType)
2094 return rewriter.notifyMatchFailure(op, "type conversion failed");
2095
2096 if (chipset < kGfx1250)
2097 return op->emitOpError("WMMA scale only supported on gfx1250+");
2098
2099 int64_t m = op.getM();
2100 int64_t n = op.getN();
2101 int64_t k = op.getK();
2102
2103 Type aElemType = getElementTypeOrSelf(op.getSourceA().getType());
2104 Type bElemType = getElementTypeOrSelf(op.getSourceB().getType());
2105
2106 std::optional<ROCDL::MatrixFormat> aFmtCode =
2108 std::optional<ROCDL::MatrixFormat> bFmtCode =
2110
2111 if (!aFmtCode || !bFmtCode)
2112 return op.emitOpError("unsupported element types for scaled_wmma");
2113
2114 // Get scale vector types and determine variant (scale vs scale16).
2115 auto scaleAVecType = cast<VectorType>(op.getScaleA().getType());
2116 auto scaleBVecType = cast<VectorType>(op.getScaleB().getType());
2117
2118 if (scaleAVecType.getNumElements() != scaleBVecType.getNumElements())
2119 return op.emitOpError("scaleA and scaleB must have equal vector length");
2120
2121 // Extract scale format from element types.
2122 Type scaleAElemType = scaleAVecType.getElementType();
2123 Type scaleBElemType = scaleBVecType.getElementType();
2124
2125 std::optional<ROCDL::WMMAMatrixScaleFormat> scaleAFmt =
2126 getWmmaScaleFormat(scaleAElemType);
2127 std::optional<ROCDL::WMMAMatrixScaleFormat> scaleBFmt =
2128 getWmmaScaleFormat(scaleBElemType);
2129
2130 if (!scaleAFmt || !scaleBFmt)
2131 return op.emitOpError("unsupported scale element types");
2132
2133 // Determine which intrinsic to use based on dimensions.
2134 bool isScale16 = (scaleAVecType.getNumElements() == 8);
2135 std::optional<StringRef> intrinsicName =
2136 getScaledWmmaIntrinsicName(m, n, k, isScale16);
2137 if (!intrinsicName)
2138 return op.emitOpError("unsupported scaled_wmma dimensions: ")
2139 << m << "x" << n << "x" << k;
2140
2141 SmallVector<NamedAttribute, 8> attrs;
2142
2143 // The f4 variant does not have fmtA and fmtB attributes.
2144 bool is32x16 = (m == 32 && n == 16 && k == 128);
2145 if (!is32x16) {
2146 attrs.emplace_back("fmtA", ROCDL::MatrixFormatAttr::get(
2147 rewriter.getContext(), *aFmtCode));
2148 attrs.emplace_back("fmtB", ROCDL::MatrixFormatAttr::get(
2149 rewriter.getContext(), *bFmtCode));
2150 }
2151
2152 // modC uses default value of 0.
2153 attrs.emplace_back(
2154 "modC", ROCDL::WMMACModifierAttr::get(rewriter.getContext(),
2155 ROCDL::WMMACModifier::none));
2156
2157 // Scale attributes. Convert user-facing firstScaleLane (0 or 16) to the
2158 // half of the wave that is being selected (0 or 1).
2159 attrs.emplace_back("scaleAType", ROCDL::WMMAMatrixScaleAttr::get(
2160 rewriter.getContext(),
2161 static_cast<ROCDL::WMMAMatrixScale>(
2162 op.getAFirstScaleLane() / 16)));
2163 attrs.emplace_back("fmtScaleA", ROCDL::WMMAMatrixScaleFormatAttr::get(
2164 rewriter.getContext(), *scaleAFmt));
2165 attrs.emplace_back("scaleBType", ROCDL::WMMAMatrixScaleAttr::get(
2166 rewriter.getContext(),
2167 static_cast<ROCDL::WMMAMatrixScale>(
2168 op.getBFirstScaleLane() / 16)));
2169 attrs.emplace_back("fmtScaleB", ROCDL::WMMAMatrixScaleFormatAttr::get(
2170 rewriter.getContext(), *scaleBFmt));
2171
2172 // Reuse flags use default value of false.
2173 attrs.emplace_back("reuseA", rewriter.getBoolAttr(false));
2174 attrs.emplace_back("reuseB", rewriter.getBoolAttr(false));
2175
2176 // Convert typed float vectors to packed format.
2177 Value sourceA =
2178 packSmallFloatVectorOperand(rewriter, loc, adaptor.getSourceA());
2179 Value sourceB =
2180 packSmallFloatVectorOperand(rewriter, loc, adaptor.getSourceB());
2181
2182 // Pack scale vectors into i32/i64.
2183 Value packedScaleA = castScaleOperand(rewriter, loc, adaptor.getScaleA());
2184 Value packedScaleB = castScaleOperand(rewriter, loc, adaptor.getScaleB());
2185
2186 // Create the intrinsic call.
2187 OperationState loweredOp(loc, *intrinsicName);
2188 loweredOp.addTypes(outType);
2189 loweredOp.addOperands(
2190 {sourceA, sourceB, adaptor.getDestC(), packedScaleA, packedScaleB});
2191 loweredOp.addAttributes(attrs);
2192
2193 Operation *lowered = rewriter.create(loweredOp);
2194 rewriter.replaceOp(op, lowered->getResults());
2195
2196 return success();
2197 }
2198};
2199
2200struct TransposeLoadOpLowering
2201 : public ConvertOpToLLVMPattern<TransposeLoadOp> {
2202 TransposeLoadOpLowering(const LLVMTypeConverter &converter, Chipset chipset)
2203 : ConvertOpToLLVMPattern<TransposeLoadOp>(converter), chipset(chipset) {}
2204
2205 Chipset chipset;
2206
2207 LogicalResult
2208 matchAndRewrite(TransposeLoadOp op, TransposeLoadOpAdaptor adaptor,
2209 ConversionPatternRewriter &rewriter) const override {
2210 if (chipset != kGfx950 && chipset < kGfx1250)
2211 return op.emitOpError(
2212 "transpose_load is only supported on gfx950 and gfx1250+");
2213
2214 Location loc = op.getLoc();
2215 auto srcMemRefType = cast<MemRefType>(op.getSrc().getType());
2216
2217 // Elements in subbyte memrefs are stored non-contiguously,
2218 // reject if source is sub-byte memref. Use emulated memrefs instead.
2219 size_t srcElementSize =
2220 srcMemRefType.getElementType().getIntOrFloatBitWidth();
2221 if (srcElementSize < 8)
2222 return op.emitOpError("Expect source memref to have at least 8 bits "
2223 "element size, got ")
2224 << srcElementSize;
2225
2226 auto resultType = cast<VectorType>(op.getResult().getType());
2227 Value srcPtr =
2228 getStridedElementPtr(rewriter, loc, srcMemRefType, adaptor.getSrc(),
2229 (adaptor.getSrcIndices()));
2230
2231 size_t numElements = resultType.getNumElements();
2232 size_t elementTypeSize =
2233 resultType.getElementType().getIntOrFloatBitWidth();
2234
2235 Type llvmResultType = typeConverter->convertType(resultType);
2236 // ROCDL transpose load intrinsics return vectors of 32-bit integers for
2237 // sub-16-bit element types, and otherwise return the converted result type.
2238 Type rocdlResultType =
2239 elementTypeSize < 16
2240 ? VectorType::get((numElements * elementTypeSize) / 32,
2241 rewriter.getIntegerType(32))
2242 : llvmResultType;
2243
2244 auto emitNumElementsError = [&](size_t expected, StringRef chipsetName) {
2245 return op.emitOpError()
2246 << elementTypeSize << "-bit transpose_load requires " << expected
2247 << " elements on " << chipsetName;
2248 };
2249
2250 Value intrinsic;
2251 if (chipset >= kGfx1250) {
2252 switch (elementTypeSize) {
2253 case 4: {
2254 if (numElements != 16)
2255 return emitNumElementsError(16, "gfx1250+");
2256 intrinsic =
2257 ROCDL::DsLoadTr4_B64::create(rewriter, loc, rocdlResultType, srcPtr)
2258 .getResult();
2259 break;
2260 }
2261 case 6: {
2262 if (numElements != 16)
2263 return emitNumElementsError(16, "gfx1250+");
2264 intrinsic =
2265 ROCDL::DsLoadTr6_B96::create(rewriter, loc, rocdlResultType, srcPtr)
2266 .getResult();
2267 break;
2268 }
2269 case 8: {
2270 if (numElements != 8)
2271 return emitNumElementsError(8, "gfx1250+");
2272 intrinsic =
2273 ROCDL::DsLoadTr8_B64::create(rewriter, loc, rocdlResultType, srcPtr)
2274 .getResult();
2275 break;
2276 }
2277 case 16: {
2278 if (numElements != 8)
2279 return emitNumElementsError(8, "gfx1250+");
2280 intrinsic = ROCDL::DsLoadTr16_B128::create(rewriter, loc,
2281 rocdlResultType, srcPtr)
2282 .getResult();
2283 break;
2284 }
2285 default:
2286 return op.emitOpError("Unsupported element size for transpose load");
2287 }
2288 } else {
2289 switch (elementTypeSize) {
2290 case 4: {
2291 if (numElements != 16)
2292 return emitNumElementsError(16, "gfx950");
2293 intrinsic = ROCDL::ds_read_tr4_b64::create(rewriter, loc,
2294 rocdlResultType, srcPtr)
2295 .getResult();
2296 break;
2297 }
2298 case 6: {
2299 if (numElements != 16)
2300 return emitNumElementsError(16, "gfx950");
2301 intrinsic = ROCDL::ds_read_tr6_b96::create(rewriter, loc,
2302 rocdlResultType, srcPtr)
2303 .getResult();
2304 break;
2305 }
2306 case 8: {
2307 if (numElements != 8)
2308 return emitNumElementsError(8, "gfx950");
2309 intrinsic = ROCDL::ds_read_tr8_b64::create(rewriter, loc,
2310 rocdlResultType, srcPtr)
2311 .getResult();
2312 break;
2313 }
2314 case 16: {
2315 if (numElements != 4)
2316 return emitNumElementsError(4, "gfx950");
2317 intrinsic = ROCDL::ds_read_tr16_b64::create(rewriter, loc,
2318 rocdlResultType, srcPtr)
2319 .getResult();
2320 break;
2321 }
2322 default:
2323 return op.emitOpError("Unsupported element size for transpose load");
2324 }
2325 }
2326
2327 assert(intrinsic && "expected ROCDL transpose load intrinsic");
2328 if (intrinsic.getType() == llvmResultType) {
2329 rewriter.replaceOp(op, intrinsic);
2330 return success();
2331 }
2332 rewriter.replaceOpWithNewOp<LLVM::BitcastOp>(op, llvmResultType, intrinsic);
2333 return success();
2334 }
2335};
2336
2337struct GlobalTransposeLoadOpLowering
2338 : public ConvertOpToLLVMPattern<GlobalTransposeLoadOp> {
2339 GlobalTransposeLoadOpLowering(const LLVMTypeConverter &converter,
2340 Chipset chipset)
2341 : ConvertOpToLLVMPattern<GlobalTransposeLoadOp>(converter),
2342 chipset(chipset) {}
2343
2344 Chipset chipset;
2345
2346 LogicalResult
2347 matchAndRewrite(GlobalTransposeLoadOp op,
2348 GlobalTransposeLoadOpAdaptor adaptor,
2349 ConversionPatternRewriter &rewriter) const override {
2350 if (chipset < kGfx1200)
2351 return op.emitOpError(
2352 "global_transpose_load is only supported on gfx1200+");
2353
2354 Location loc = op.getLoc();
2355 auto srcMemRefType = cast<MemRefType>(op.getSrc().getType());
2356 auto resultType = cast<VectorType>(op.getResult().getType());
2357
2358 Value srcPtr = getStridedElementPtr(
2359 rewriter, loc, srcMemRefType, adaptor.getSrc(), adaptor.getSrcIndices(),
2360 LLVM::GEPNoWrapFlags::inbounds | LLVM::GEPNoWrapFlags::nuw);
2361
2362 size_t numElements = resultType.getNumElements();
2363 size_t elementTypeSize =
2364 resultType.getElementType().getIntOrFloatBitWidth();
2365
2366 // ROCDL global transpose load intrinsics return vectors of i32 for
2367 // sub-16-bit elements, matching the LDS lowering convention.
2368 Type rocdlResultType =
2369 elementTypeSize < 16
2370 ? VectorType::get((numElements * elementTypeSize) / 32,
2371 rewriter.getIntegerType(32))
2372 : typeConverter->convertType(resultType);
2373 Type llvmResultType = typeConverter->convertType(resultType);
2374
2375 switch (elementTypeSize) {
2376 case 4: {
2377 assert(numElements == 16);
2378 if (chipset < kGfx1250)
2379 return op.emitOpError("4-bit global_transpose_load requires gfx1250+");
2380 auto rocdlOp = ROCDL::GlobalLoadTr4_B64::create(rewriter, loc,
2381 rocdlResultType, srcPtr);
2382 rewriter.replaceOpWithNewOp<LLVM::BitcastOp>(op, llvmResultType, rocdlOp);
2383 break;
2384 }
2385 case 6: {
2386 assert(numElements == 16);
2387 if (chipset < kGfx1250)
2388 return op.emitOpError("6-bit global_transpose_load requires gfx1250+");
2389 auto rocdlOp = ROCDL::GlobalLoadTr6_B96::create(rewriter, loc,
2390 rocdlResultType, srcPtr);
2391 rewriter.replaceOpWithNewOp<LLVM::BitcastOp>(op, llvmResultType, rocdlOp);
2392 break;
2393 }
2394 case 8: {
2395 assert(numElements == 8);
2396 auto rocdlOp = ROCDL::GlobalLoadTr8_B64::create(rewriter, loc,
2397 rocdlResultType, srcPtr);
2398 rewriter.replaceOpWithNewOp<LLVM::BitcastOp>(op, llvmResultType, rocdlOp);
2399 break;
2400 }
2401 case 16: {
2402 assert(numElements == 8);
2403 rewriter.replaceOpWithNewOp<ROCDL::GlobalLoadTr8_B128>(op, llvmResultType,
2404 srcPtr);
2405 break;
2406 }
2407 default:
2408 return op.emitOpError(
2409 "unsupported element size for global transpose load");
2410 }
2411 return success();
2412 }
2413};
2414
2415struct GatherToLDSOpLowering : public ConvertOpToLLVMPattern<GatherToLDSOp> {
2416 GatherToLDSOpLowering(const LLVMTypeConverter &converter, Chipset chipset)
2417 : ConvertOpToLLVMPattern<GatherToLDSOp>(converter), chipset(chipset) {}
2418
2419 Chipset chipset;
2420
2421 LogicalResult
2422 matchAndRewrite(GatherToLDSOp op, GatherToLDSOpAdaptor adaptor,
2423 ConversionPatternRewriter &rewriter) const override {
2424 if (chipset.majorVersion < 9 || chipset.majorVersion > 10)
2425 return op.emitOpError("pre-gfx9 and post-gfx10 not supported");
2426
2427 Location loc = op.getLoc();
2428
2429 auto srcMemRefType = cast<MemRefType>(op.getSrc().getType());
2430 auto dstMemRefType = cast<MemRefType>(op.getDst().getType());
2431
2432 // TODO: instead of only transfering one element per thread, we could
2433 // augment it to transfer multiple elements per thread by issuing multiple
2434 // `global_load_lds` instructions.
2435 Type transferType = op.getTransferType();
2436 int loadWidth = [&]() -> int {
2437 if (auto transferVectorType = dyn_cast<VectorType>(transferType)) {
2438 return (transferVectorType.getNumElements() *
2439 transferVectorType.getElementTypeBitWidth()) /
2440 8;
2441 }
2442 return transferType.getIntOrFloatBitWidth() / 8;
2443 }();
2444
2445 // Currently only 1, 2, 4, 12 and 16 byte loads are supported.
2446 if (!llvm::is_contained({1, 2, 4, 12, 16}, loadWidth))
2447 return op.emitOpError("chipset unsupported element size");
2448
2449 if (chipset != kGfx950 && llvm::is_contained({12, 16}, loadWidth))
2450 return op.emitOpError("Gather to LDS instructions with 12-byte and "
2451 "16-byte load widths are only supported on gfx950");
2452
2453 Value srcPtr =
2454 getStridedElementPtr(rewriter, loc, srcMemRefType, adaptor.getSrc(),
2455 (adaptor.getSrcIndices()));
2456 Value dstPtr =
2457 getStridedElementPtr(rewriter, loc, dstMemRefType, adaptor.getDst(),
2458 (adaptor.getDstIndices()));
2459
2460 if (op.getAsync()) {
2461 rewriter.replaceOpWithNewOp<ROCDL::LoadAsyncToLDSOp>(
2462 op, srcPtr, dstPtr, rewriter.getI32IntegerAttr(loadWidth),
2463 /*offset=*/rewriter.getI32IntegerAttr(0),
2464 /*aux=*/rewriter.getI32IntegerAttr(0), ArrayAttr{}, ArrayAttr{},
2465 ArrayAttr{});
2466 } else {
2467 rewriter.replaceOpWithNewOp<ROCDL::LoadToLDSOp>(
2468 op, srcPtr, dstPtr, rewriter.getI32IntegerAttr(loadWidth),
2469 /*offset=*/rewriter.getI32IntegerAttr(0),
2470 /*aux=*/rewriter.getI32IntegerAttr(0), ArrayAttr{}, ArrayAttr{},
2471 ArrayAttr{});
2472 }
2473
2474 return success();
2475 }
2476};
2477
2478struct GlobalLoadAsyncToLDSOpLowering
2479 : public ConvertOpToLLVMPattern<GlobalLoadAsyncToLDSOp> {
2480 GlobalLoadAsyncToLDSOpLowering(const LLVMTypeConverter &converter,
2481 Chipset chipset)
2482 : ConvertOpToLLVMPattern<GlobalLoadAsyncToLDSOp>(converter),
2483 chipset(chipset) {}
2484
2485 Chipset chipset;
2486
2487 LogicalResult
2488 matchAndRewrite(GlobalLoadAsyncToLDSOp op,
2489 GlobalLoadAsyncToLDSOpAdaptor adaptor,
2490 ConversionPatternRewriter &rewriter) const override {
2491 if (chipset < kGfx1250)
2492 return op.emitOpError(
2493 "global_load_async_to_lds is only supported on gfx1250+");
2494
2495 Location loc = op.getLoc();
2496 auto srcMemRefType = cast<MemRefType>(op.getSrc().getType());
2497 auto dstMemRefType = cast<MemRefType>(op.getDst().getType());
2498
2499 Type transferType = op.getTransferType();
2500 int transferBits =
2501 isa<VectorType>(transferType)
2502 ? cast<VectorType>(transferType).getNumElements() *
2503 cast<VectorType>(transferType).getElementTypeBitWidth()
2504 : transferType.getIntOrFloatBitWidth();
2505
2506 Value srcPtr =
2507 getStridedElementPtr(rewriter, loc, srcMemRefType, adaptor.getSrc(),
2508 adaptor.getSrcIndices());
2509 Value dstPtr =
2510 getStridedElementPtr(rewriter, loc, dstMemRefType, adaptor.getDst(),
2511 adaptor.getDstIndices());
2512
2513 if (op.getMask()) {
2514 Value mask = adaptor.getMask();
2515 int64_t nullptrVal =
2516 llvm::AMDGPU::getNullPointerValue(llvm::AMDGPUAS::LOCAL_ADDRESS);
2517 Value nullInt =
2518 createI32Constant(rewriter, loc, static_cast<int32_t>(nullptrVal));
2519 Value nullPtr =
2520 LLVM::IntToPtrOp::create(rewriter, loc, dstPtr.getType(), nullInt);
2521 dstPtr = LLVM::SelectOp::create(rewriter, loc, mask, dstPtr, nullPtr);
2522 }
2523
2524 auto offset = rewriter.getI32IntegerAttr(0);
2525 Attribute aux = rewriter.getI32IntegerAttr(0);
2526
2527 switch (transferBits) {
2528 case 8:
2529 rewriter.replaceOpWithNewOp<ROCDL::GlobalLoadAsyncToLDSB8Op>(
2530 op, srcPtr, dstPtr, offset, aux, ArrayAttr{}, ArrayAttr{},
2531 ArrayAttr{});
2532 break;
2533 case 32:
2534 rewriter.replaceOpWithNewOp<ROCDL::GlobalLoadAsyncToLDSB32Op>(
2535 op, srcPtr, dstPtr, offset, aux, ArrayAttr{}, ArrayAttr{},
2536 ArrayAttr{});
2537 break;
2538 case 64:
2539 rewriter.replaceOpWithNewOp<ROCDL::GlobalLoadAsyncToLDSB64Op>(
2540 op, srcPtr, dstPtr, offset, aux, ArrayAttr{}, ArrayAttr{},
2541 ArrayAttr{});
2542 break;
2543 case 128:
2544 rewriter.replaceOpWithNewOp<ROCDL::GlobalLoadAsyncToLDSB128Op>(
2545 op, srcPtr, dstPtr, offset, aux, ArrayAttr{}, ArrayAttr{},
2546 ArrayAttr{});
2547 break;
2548 default:
2549 return op.emitOpError("unsupported transfer width");
2550 }
2551 return success();
2552 }
2553};
2554
2555namespace {
2556struct ExtPackedFp8OpLowering final
2557 : public ConvertOpToLLVMPattern<ExtPackedFp8Op> {
2558 ExtPackedFp8OpLowering(const LLVMTypeConverter &converter, Chipset chipset)
2559 : ConvertOpToLLVMPattern<amdgpu::ExtPackedFp8Op>(converter),
2560 chipset(chipset) {}
2561 Chipset chipset;
2562
2563 LogicalResult
2564 matchAndRewrite(ExtPackedFp8Op op, ExtPackedFp8OpAdaptor adaptor,
2565 ConversionPatternRewriter &rewriter) const override;
2566};
2567
2568struct ScaledExtPackedMatrixOpLowering final
2569 : public ConvertOpToLLVMPattern<ScaledExtPackedMatrixOp> {
2570 ScaledExtPackedMatrixOpLowering(const LLVMTypeConverter &converter,
2571 Chipset chipset)
2572 : ConvertOpToLLVMPattern<amdgpu::ScaledExtPackedMatrixOp>(converter),
2573 chipset(chipset) {}
2574 Chipset chipset;
2575
2576 LogicalResult
2577 matchAndRewrite(ScaledExtPackedMatrixOp op,
2578 ScaledExtPackedMatrixOpAdaptor adaptor,
2579 ConversionPatternRewriter &rewriter) const override;
2580};
2581
2582struct PackedTrunc2xFp8OpLowering final
2583 : public ConvertOpToLLVMPattern<PackedTrunc2xFp8Op> {
2584 PackedTrunc2xFp8OpLowering(const LLVMTypeConverter &converter,
2585 Chipset chipset)
2586 : ConvertOpToLLVMPattern<amdgpu::PackedTrunc2xFp8Op>(converter),
2587 chipset(chipset) {}
2588 Chipset chipset;
2589
2590 LogicalResult
2591 matchAndRewrite(PackedTrunc2xFp8Op op, PackedTrunc2xFp8OpAdaptor adaptor,
2592 ConversionPatternRewriter &rewriter) const override;
2593};
2594
2595struct PackedStochRoundFp8OpLowering final
2596 : public ConvertOpToLLVMPattern<PackedStochRoundFp8Op> {
2597 PackedStochRoundFp8OpLowering(const LLVMTypeConverter &converter,
2598 Chipset chipset)
2599 : ConvertOpToLLVMPattern<amdgpu::PackedStochRoundFp8Op>(converter),
2600 chipset(chipset) {}
2601 Chipset chipset;
2602
2603 LogicalResult
2604 matchAndRewrite(PackedStochRoundFp8Op op,
2605 PackedStochRoundFp8OpAdaptor adaptor,
2606 ConversionPatternRewriter &rewriter) const override;
2607};
2608
2609struct ScaledExtPackedOpLowering final
2610 : public ConvertOpToLLVMPattern<ScaledExtPackedOp> {
2611 ScaledExtPackedOpLowering(const LLVMTypeConverter &converter, Chipset chipset)
2612 : ConvertOpToLLVMPattern<amdgpu::ScaledExtPackedOp>(converter),
2613 chipset(chipset) {}
2614 Chipset chipset;
2615
2616 LogicalResult
2617 matchAndRewrite(ScaledExtPackedOp op, ScaledExtPackedOpAdaptor adaptor,
2618 ConversionPatternRewriter &rewriter) const override;
2619};
2620
2621struct PackedScaledTruncOpLowering final
2622 : public ConvertOpToLLVMPattern<PackedScaledTruncOp> {
2623 PackedScaledTruncOpLowering(const LLVMTypeConverter &converter,
2624 Chipset chipset)
2625 : ConvertOpToLLVMPattern<amdgpu::PackedScaledTruncOp>(converter),
2626 chipset(chipset) {}
2627 Chipset chipset;
2628
2629 LogicalResult
2630 matchAndRewrite(PackedScaledTruncOp op, PackedScaledTruncOpAdaptor adaptor,
2631 ConversionPatternRewriter &rewriter) const override;
2632};
2633
2634} // end namespace
2635
2636LogicalResult ExtPackedFp8OpLowering::matchAndRewrite(
2637 ExtPackedFp8Op op, ExtPackedFp8OpAdaptor adaptor,
2638 ConversionPatternRewriter &rewriter) const {
2639 Location loc = op.getLoc();
2640 if (!(chipset == kGfx942 || hasOcpFp8(chipset)))
2641 return rewriter.notifyMatchFailure(
2642 loc, "Fp8 conversion instructions are not available on target "
2643 "architecture and their emulation is not implemented");
2644 Type v4i8 =
2645 getTypeConverter()->convertType(VectorType::get(4, rewriter.getI8Type()));
2646 Type i32 = getTypeConverter()->convertType(rewriter.getI32Type());
2647 Type f32 = getTypeConverter()->convertType(op.getResult().getType());
2648
2649 Value source = adaptor.getSource();
2650 auto sourceVecType = dyn_cast<VectorType>(op.getSource().getType());
2651 auto resultVecType = dyn_cast<VectorType>(op.getResult().getType());
2652 Type sourceElemType = getElementTypeOrSelf(op.getSource());
2653 // Extend to a v4i8
2654 if (!sourceVecType || sourceVecType.getNumElements() < 4) {
2655 Value longVec = LLVM::UndefOp::create(rewriter, loc, v4i8);
2656 if (!sourceVecType) {
2657 longVec = LLVM::InsertElementOp::create(
2658 rewriter, loc, longVec, source, createI32Constant(rewriter, loc, 0));
2659 } else {
2660 for (int32_t i = 0, e = sourceVecType.getNumElements(); i < e; ++i) {
2661 Value idx = createI32Constant(rewriter, loc, i);
2662 Value elem = LLVM::ExtractElementOp::create(rewriter, loc, source, idx);
2663 longVec =
2664 LLVM::InsertElementOp::create(rewriter, loc, longVec, elem, idx);
2665 }
2666 }
2667 source = longVec;
2668 }
2669 Value i32Source = LLVM::BitcastOp::create(rewriter, loc, i32, source);
2670 if (resultVecType) {
2671 if (typeIsExpectedBf8ForChipset(chipset, sourceElemType)) {
2672 rewriter.replaceOpWithNewOp<ROCDL::CvtPkF32Bf8Op>(op, f32, i32Source,
2673 op.getIndex());
2674 } else if (typeIsExpectedFp8ForChipset(chipset, sourceElemType)) {
2675 rewriter.replaceOpWithNewOp<ROCDL::CvtPkF32Fp8Op>(op, f32, i32Source,
2676 op.getIndex());
2677 }
2678 } else {
2679 if (typeIsExpectedBf8ForChipset(chipset, sourceElemType)) {
2680 rewriter.replaceOpWithNewOp<ROCDL::CvtF32Bf8Op>(op, f32, i32Source,
2681 op.getIndex());
2682 } else if (typeIsExpectedFp8ForChipset(chipset, sourceElemType)) {
2683 rewriter.replaceOpWithNewOp<ROCDL::CvtF32Fp8Op>(op, f32, i32Source,
2684 op.getIndex());
2685 }
2686 }
2687 return success();
2688}
2689
2690int32_t getScaleSel(int32_t blockSize, unsigned bitWidth, int32_t scaleWaveHalf,
2691 int32_t firstScaleByte) {
2692 // When lowering amdgpu.scaled_ext_packed_matrix to rocdl.cvt.scale.pk*.f*.f*
2693 // operations, the attributes blockSize, sourceType, scaleWaveHalf, and
2694 // firstScaleByte are merged into a single attribute scaleSel. This is how
2695 // those values are merged together. (Note: scaleWaveHalf isn't a high-level
2696 // attribute but is derifed from firstScaleLane).
2697 assert(llvm::is_contained({16, 32}, blockSize));
2698 assert(llvm::is_contained({4u, 6u, 8u}, bitWidth));
2699
2700 const bool isFp8 = bitWidth == 8;
2701 const bool isBlock16 = blockSize == 16;
2702
2703 if (!isFp8) {
2704 int32_t bit0 = isBlock16;
2705 assert(llvm::is_contained({0, 1, 2}, firstScaleByte));
2706 int32_t bit1 = (firstScaleByte == 2) << 1;
2707 assert(llvm::is_contained({0, 1}, scaleWaveHalf));
2708 int32_t bit2 = scaleWaveHalf << 2;
2709 return bit2 | bit1 | bit0;
2710 }
2711
2712 int32_t bit0 = isBlock16;
2713 // firstScaleByte is guaranteed to be defined by two bits.
2714 assert(llvm::is_contained({0, 1, 2, 3}, firstScaleByte));
2715 int32_t bits2and1 = firstScaleByte << 1;
2716 assert(llvm::is_contained({0, 1}, scaleWaveHalf));
2717 int32_t bit3 = scaleWaveHalf << 3;
2718 int32_t bits = bit3 | bits2and1 | bit0;
2719 // These are invalid cases.
2720 assert(!llvm::is_contained(
2721 {0b0011, 0b0101, 0b0111, 0b1000, 0b1001, 0b1011, 0b1111}, bits));
2722 return bits;
2723}
2724
2725static std::optional<StringRef>
2726scaledExtPacked816ToIntrinsic(Type srcElemType, Type destElemType) {
2727 using fp4 = Float4E2M1FNType;
2728 using fp8 = Float8E4M3FNType;
2729 using bf8 = Float8E5M2Type;
2730 using fp6 = Float6E2M3FNType;
2731 using bf6 = Float6E3M2FNType;
2732 if (isa<fp4>(srcElemType)) {
2733 if (destElemType.isF16())
2734 return ROCDL::CvtPkScalePk8F16Fp4Op::getOperationName();
2735 if (destElemType.isBF16())
2736 return ROCDL::CvtPkScalePk8Bf16Fp4Op::getOperationName();
2737 if (destElemType.isF32())
2738 return ROCDL::CvtPkScalePk8F32Fp4Op::getOperationName();
2739 return std::nullopt;
2740 }
2741 if (isa<fp8>(srcElemType)) {
2742 if (destElemType.isF16())
2743 return ROCDL::CvtPkScalePk8F16Fp8Op::getOperationName();
2744 if (destElemType.isBF16())
2745 return ROCDL::CvtPkScalePk8Bf16Fp8Op::getOperationName();
2746 if (destElemType.isF32())
2747 return ROCDL::CvtPkScalePk8F32Fp8Op::getOperationName();
2748 return std::nullopt;
2749 }
2750 if (isa<bf8>(srcElemType)) {
2751 if (destElemType.isF16())
2752 return ROCDL::CvtPkScalePk8F16Bf8Op::getOperationName();
2753 if (destElemType.isBF16())
2754 return ROCDL::CvtPkScalePk8Bf16Bf8Op::getOperationName();
2755 if (destElemType.isF32())
2756 return ROCDL::CvtPkScalePk8F32Bf8Op::getOperationName();
2757 return std::nullopt;
2758 }
2759 if (isa<fp6>(srcElemType)) {
2760 if (destElemType.isF16())
2761 return ROCDL::CvtPkScalePk16F16Fp6Op::getOperationName();
2762 if (destElemType.isBF16())
2763 return ROCDL::CvtPkScalePk16Bf16Fp6Op::getOperationName();
2764 if (destElemType.isF32())
2765 return ROCDL::CvtPkScalePk16F32Fp6Op::getOperationName();
2766 return std::nullopt;
2767 }
2768 if (isa<bf6>(srcElemType)) {
2769 if (destElemType.isF16())
2770 return ROCDL::CvtPkScalePk16F16Bf6Op::getOperationName();
2771 if (destElemType.isBF16())
2772 return ROCDL::CvtPkScalePk16Bf16Bf6Op::getOperationName();
2773 if (destElemType.isF32())
2774 return ROCDL::CvtPkScalePk16F32Bf6Op::getOperationName();
2775 return std::nullopt;
2776 }
2777 llvm_unreachable("invalid combination of element types for packed conversion "
2778 "instructions");
2779}
2780
2781LogicalResult ScaledExtPackedMatrixOpLowering::matchAndRewrite(
2782 ScaledExtPackedMatrixOp op, ScaledExtPackedMatrixOpAdaptor adaptor,
2783 ConversionPatternRewriter &rewriter) const {
2784 using fp4 = Float4E2M1FNType;
2785 using fp8 = Float8E4M3FNType;
2786 using bf8 = Float8E5M2Type;
2787 using fp6 = Float6E2M3FNType;
2788 using bf6 = Float6E3M2FNType;
2789 Location loc = op.getLoc();
2790 if (chipset != kGfx1250) {
2791 return rewriter.notifyMatchFailure(
2792 loc,
2793 "Scaled fp packed conversion instructions are not available on target "
2794 "architecture and their emulation is not implemented");
2795 }
2796 // Convert user-facing firstScaleLane (0 or 16) to the half of the wave that
2797 // is being selected.
2798 int32_t scaleWaveHalf = op.getFirstScaleLane() / 16;
2799 int32_t firstScaleByte = op.getFirstScaleByte();
2800 int32_t blockSize = op.getBlockSize();
2801 auto sourceType = cast<VectorType>(op.getSource().getType());
2802 auto srcElemType = cast<FloatType>(sourceType.getElementType());
2803 unsigned bitWidth = srcElemType.getWidth();
2804
2805 auto targetType = cast<VectorType>(op.getResult().getType());
2806 auto destElemType = cast<FloatType>(targetType.getElementType());
2807
2808 IntegerType i32 = rewriter.getI32Type();
2809 Value source = adaptor.getSource();
2810 Type llvmResultType = typeConverter->convertType(op.getResult().getType());
2811 Type packedType = nullptr;
2812 if (isa<fp4>(srcElemType)) {
2813 packedType = i32;
2814 packedType = getTypeConverter()->convertType(packedType);
2815 } else if (isa<fp8, bf8>(srcElemType)) {
2816 packedType = VectorType::get(2, i32);
2817 packedType = getTypeConverter()->convertType(packedType);
2818 } else if (isa<fp6, bf6>(srcElemType)) {
2819 packedType = VectorType::get(3, i32);
2820 packedType = getTypeConverter()->convertType(packedType);
2821 } else {
2822 llvm_unreachable("invalid element type for packed scaled ext");
2823 }
2824
2825 if (!packedType || !llvmResultType) {
2826 return rewriter.notifyMatchFailure(op, "type conversion failed");
2827 }
2828
2829 std::optional<StringRef> maybeIntrinsic =
2830 scaledExtPacked816ToIntrinsic(srcElemType, destElemType);
2831 if (!maybeIntrinsic.has_value())
2832 return op.emitOpError(
2833 "no intrinsic matching packed scaled conversion on the given chipset");
2834
2835 int32_t scaleSel =
2836 getScaleSel(blockSize, bitWidth, scaleWaveHalf, firstScaleByte);
2837 Value castedScale =
2838 LLVM::BitcastOp::create(rewriter, loc, i32, adaptor.getScale());
2839 Value castedSource =
2840 LLVM::BitcastOp::create(rewriter, loc, packedType, source);
2841
2842 OperationState loweredOp(loc, *maybeIntrinsic);
2843 loweredOp.addTypes({llvmResultType});
2844 loweredOp.addOperands({castedSource, castedScale});
2845
2846 SmallVector<NamedAttribute, 1> attrs;
2847 attrs.push_back(
2848 NamedAttribute("scaleSel", rewriter.getI32IntegerAttr(scaleSel)));
2849
2850 loweredOp.addAttributes(attrs);
2851 Operation *lowered = rewriter.create(loweredOp);
2852 rewriter.replaceOp(op, lowered);
2853
2854 return success();
2855}
2856
2857LogicalResult ScaledExtPackedOpLowering::matchAndRewrite(
2858 ScaledExtPackedOp op, ScaledExtPackedOpAdaptor adaptor,
2859 ConversionPatternRewriter &rewriter) const {
2860 Location loc = op.getLoc();
2861 if (chipset != kGfx950)
2862 return rewriter.notifyMatchFailure(
2863 loc, "Scaled fp conversion instructions are not available on target "
2864 "architecture and their emulation is not implemented");
2865 Type i32 = getTypeConverter()->convertType(rewriter.getI32Type());
2866
2867 Value source = adaptor.getSource();
2868 Value scale = adaptor.getScale();
2869
2870 VectorType sourceVecType = cast<VectorType>(op.getSource().getType());
2871 Type sourceElemType = sourceVecType.getElementType();
2872 VectorType destVecType = cast<VectorType>(op.getResult().getType());
2873 Type destElemType = destVecType.getElementType();
2874
2875 VectorType packedVecType;
2876 if (isa<Float8E5M2Type, Float8E4M3FNType>(sourceElemType)) {
2877 VectorType v4i8 = VectorType::get(4, rewriter.getI8Type());
2878 packedVecType = cast<VectorType>(getTypeConverter()->convertType(v4i8));
2879 } else if (isa<Float4E2M1FNType>(sourceElemType)) {
2880 VectorType v8i4 = VectorType::get(8, rewriter.getI4Type());
2881 packedVecType = cast<VectorType>(getTypeConverter()->convertType(v8i4));
2882 } else {
2883 llvm_unreachable("invalid element type for scaled ext");
2884 }
2885
2886 // Extend to a packedVectorType
2887 if (sourceVecType.getNumElements() < packedVecType.getNumElements()) {
2888 Value longVec = LLVM::ZeroOp::create(rewriter, loc, packedVecType);
2889 if (!sourceVecType) {
2890 longVec = LLVM::InsertElementOp::create(
2891 rewriter, loc, longVec, source, createI32Constant(rewriter, loc, 0));
2892 } else {
2893 for (int32_t i = 0, e = sourceVecType.getNumElements(); i < e; ++i) {
2894 Value idx = createI32Constant(rewriter, loc, i);
2895 Value elem = LLVM::ExtractElementOp::create(rewriter, loc, source, idx);
2896 longVec =
2897 LLVM::InsertElementOp::create(rewriter, loc, longVec, elem, idx);
2898 }
2899 }
2900 source = longVec;
2901 }
2902 Value i32Source = LLVM::BitcastOp::create(rewriter, loc, i32, source);
2903
2904 if (isa<Float8E5M2Type>(sourceElemType) && destElemType.isF32())
2905 rewriter.replaceOpWithNewOp<ROCDL::CvtScaleF32PkF32Bf8Op>(
2906 op, destVecType, i32Source, scale, op.getIndex());
2907 else if (isa<Float8E5M2Type>(sourceElemType) && destElemType.isF16())
2908 rewriter.replaceOpWithNewOp<ROCDL::CvtScaleF32PkF16Bf8Op>(
2909 op, destVecType, i32Source, scale, op.getIndex());
2910 else if (isa<Float8E5M2Type>(sourceElemType) && destElemType.isBF16())
2911 rewriter.replaceOpWithNewOp<ROCDL::CvtScaleF32PkBf16Bf8Op>(
2912 op, destVecType, i32Source, scale, op.getIndex());
2913 else if (isa<Float8E4M3FNType>(sourceElemType) && destElemType.isF32())
2914 rewriter.replaceOpWithNewOp<ROCDL::CvtScaleF32PkF32Fp8Op>(
2915 op, destVecType, i32Source, scale, op.getIndex());
2916 else if (isa<Float8E4M3FNType>(sourceElemType) && destElemType.isF16())
2917 rewriter.replaceOpWithNewOp<ROCDL::CvtScaleF32PkF16Fp8Op>(
2918 op, destVecType, i32Source, scale, op.getIndex());
2919 else if (isa<Float8E4M3FNType>(sourceElemType) && destElemType.isBF16())
2920 rewriter.replaceOpWithNewOp<ROCDL::CvtScaleF32PkBf16Fp8Op>(
2921 op, destVecType, i32Source, scale, op.getIndex());
2922 else if (isa<Float4E2M1FNType>(sourceElemType) && destElemType.isF32())
2923 rewriter.replaceOpWithNewOp<ROCDL::CvtScaleF32PkF32Fp4Op>(
2924 op, destVecType, i32Source, scale, op.getIndex());
2925 else if (isa<Float4E2M1FNType>(sourceElemType) && destElemType.isF16())
2926 rewriter.replaceOpWithNewOp<ROCDL::CvtScaleF32PkF16Fp4Op>(
2927 op, destVecType, i32Source, scale, op.getIndex());
2928 else if (isa<Float4E2M1FNType>(sourceElemType) && destElemType.isBF16())
2929 rewriter.replaceOpWithNewOp<ROCDL::CvtScaleF32PkBf16Fp4Op>(
2930 op, destVecType, i32Source, scale, op.getIndex());
2931 else
2932 return failure();
2933
2934 return success();
2935}
2936
2937LogicalResult PackedScaledTruncOpLowering::matchAndRewrite(
2938 PackedScaledTruncOp op, PackedScaledTruncOpAdaptor adaptor,
2939 ConversionPatternRewriter &rewriter) const {
2940 Location loc = op.getLoc();
2941 if (chipset != kGfx950)
2942 return rewriter.notifyMatchFailure(
2943 loc, "Scaled fp conversion instructions are not available on target "
2944 "architecture and their emulation is not implemented");
2945 Type v2i16 = getTypeConverter()->convertType(
2946 VectorType::get(2, rewriter.getI16Type()));
2947 Type i32 = getTypeConverter()->convertType(rewriter.getI32Type());
2948
2949 Type resultType = op.getResult().getType();
2950 Type resultElemType = getElementTypeOrSelf(resultType);
2951 VectorType sourceVecType = cast<VectorType>(op.getSource().getType());
2952 Type sourceElemType = sourceVecType.getElementType();
2953
2954 Type intResultType = isa<Float4E2M1FNType>(resultElemType) ? i32 : v2i16;
2955
2956 Value source = adaptor.getSource();
2957 Value scale = adaptor.getScale();
2958 Value existing = adaptor.getExisting();
2959 if (existing)
2960 existing = LLVM::BitcastOp::create(rewriter, loc, intResultType, existing);
2961 else
2962 existing = LLVM::ZeroOp::create(rewriter, loc, intResultType);
2963
2964 if (sourceVecType.getNumElements() < 2) {
2965 Value c0 = createI32Constant(rewriter, loc, 0);
2966 Value elem0 = LLVM::ExtractElementOp::create(rewriter, loc, source, c0);
2967 VectorType v2 = VectorType::get(2, sourceElemType);
2968 source = LLVM::ZeroOp::create(rewriter, loc, v2);
2969 source = LLVM::InsertElementOp::create(rewriter, loc, source, elem0, c0);
2970 }
2971
2972 Value sourceA, sourceB;
2973 if (sourceElemType.isF32()) {
2974 Value c0 = createI32Constant(rewriter, loc, 0);
2975 Value c1 = createI32Constant(rewriter, loc, 1);
2976 sourceA = LLVM::ExtractElementOp::create(rewriter, loc, source, c0);
2977 sourceB = LLVM::ExtractElementOp::create(rewriter, loc, source, c1);
2978 }
2979
2980 Value result;
2981 if (sourceElemType.isF32() && isa<Float8E5M2Type>(resultElemType))
2982 result = ROCDL::CvtScaleF32PkBf8F32Op::create(rewriter, loc, intResultType,
2983 existing, sourceA, sourceB,
2984 scale, op.getIndex());
2985 else if (sourceElemType.isF16() && isa<Float8E5M2Type>(resultElemType))
2986 result = ROCDL::CvtScaleF32PkBf8F16Op::create(
2987 rewriter, loc, intResultType, existing, source, scale, op.getIndex());
2988 else if (sourceElemType.isBF16() && isa<Float8E5M2Type>(resultElemType))
2989 result = ROCDL::CvtScaleF32PkBf8Bf16Op::create(
2990 rewriter, loc, intResultType, existing, source, scale, op.getIndex());
2991 else if (sourceElemType.isF32() && isa<Float8E4M3FNType>(resultElemType))
2992 result = ROCDL::CvtScaleF32PkFp8F32Op::create(rewriter, loc, intResultType,
2993 existing, sourceA, sourceB,
2994 scale, op.getIndex());
2995 else if (sourceElemType.isF16() && isa<Float8E4M3FNType>(resultElemType))
2996 result = ROCDL::CvtScaleF32PkFp8F16Op::create(
2997 rewriter, loc, intResultType, existing, source, scale, op.getIndex());
2998 else if (sourceElemType.isBF16() && isa<Float8E4M3FNType>(resultElemType))
2999 result = ROCDL::CvtScaleF32PkFp8Bf16Op::create(
3000 rewriter, loc, intResultType, existing, source, scale, op.getIndex());
3001 else if (sourceElemType.isF32() && isa<Float4E2M1FNType>(resultElemType))
3002 result = ROCDL::CvtScaleF32PkFp4F32Op::create(rewriter, loc, intResultType,
3003 existing, sourceA, sourceB,
3004 scale, op.getIndex());
3005 else if (sourceElemType.isF16() && isa<Float4E2M1FNType>(resultElemType))
3006 result = ROCDL::CvtScaleF32PkFp4F16Op::create(
3007 rewriter, loc, intResultType, existing, source, scale, op.getIndex());
3008 else if (sourceElemType.isBF16() && isa<Float4E2M1FNType>(resultElemType))
3009 result = ROCDL::CvtScaleF32PkFp4Bf16Op::create(
3010 rewriter, loc, intResultType, existing, source, scale, op.getIndex());
3011 else
3012 return failure();
3013
3014 result = rewriter.replaceOpWithNewOp<LLVM::BitcastOp>(
3015 op, getTypeConverter()->convertType(resultType), result);
3016 return success();
3017}
3018
3019LogicalResult PackedTrunc2xFp8OpLowering::matchAndRewrite(
3020 PackedTrunc2xFp8Op op, PackedTrunc2xFp8OpAdaptor adaptor,
3021 ConversionPatternRewriter &rewriter) const {
3022 Location loc = op.getLoc();
3023 if (!(chipset == kGfx942 || hasOcpFp8(chipset)))
3024 return rewriter.notifyMatchFailure(
3025 loc, "Fp8 conversion instructions are not available on target "
3026 "architecture and their emulation is not implemented");
3027 Type i32 = getTypeConverter()->convertType(rewriter.getI32Type());
3028
3029 Type resultType = op.getResult().getType();
3030 Type resultElemType = getElementTypeOrSelf(resultType);
3031
3032 Value sourceA = adaptor.getSourceA();
3033 Value sourceB = adaptor.getSourceB();
3034 if (!sourceB)
3035 sourceB = LLVM::UndefOp::create(rewriter, loc, sourceA.getType());
3036 Value existing = adaptor.getExisting();
3037 if (existing)
3038 existing = LLVM::BitcastOp::create(rewriter, loc, i32, existing);
3039 else
3040 existing = LLVM::UndefOp::create(rewriter, loc, i32);
3041
3042 Value result;
3043 if (typeIsExpectedBf8ForChipset(chipset, resultElemType))
3044 result = ROCDL::CvtPkBf8F32Op::create(rewriter, loc, i32, sourceA, sourceB,
3045 existing, op.getWordIndex());
3046 else if (typeIsExpectedFp8ForChipset(chipset, resultElemType))
3047 result = ROCDL::CvtPkFp8F32Op::create(rewriter, loc, i32, sourceA, sourceB,
3048 existing, op.getWordIndex());
3049
3050 result = rewriter.replaceOpWithNewOp<LLVM::BitcastOp>(
3051 op, getTypeConverter()->convertType(resultType), result);
3052 return success();
3053}
3054
3055LogicalResult PackedStochRoundFp8OpLowering::matchAndRewrite(
3056 PackedStochRoundFp8Op op, PackedStochRoundFp8OpAdaptor adaptor,
3057 ConversionPatternRewriter &rewriter) const {
3058 Location loc = op.getLoc();
3059 if (!(chipset == kGfx942 || hasOcpFp8(chipset)))
3060 return rewriter.notifyMatchFailure(
3061 loc, "Fp8 conversion instructions are not available on target "
3062 "architecture and their emulation is not implemented");
3063 Type i32 = getTypeConverter()->convertType(rewriter.getI32Type());
3064
3065 Type resultType = op.getResult().getType();
3066 Type resultElemType = getElementTypeOrSelf(resultType);
3067
3068 Value source = adaptor.getSource();
3069 Value stoch = adaptor.getStochiasticParam();
3070 Value existing = adaptor.getExisting();
3071 if (existing)
3072 existing = LLVM::BitcastOp::create(rewriter, loc, i32, existing);
3073 else
3074 existing = LLVM::UndefOp::create(rewriter, loc, i32);
3075
3076 Value result;
3077 if (typeIsExpectedBf8ForChipset(chipset, resultElemType))
3078 result = ROCDL::CvtSrBf8F32Op::create(rewriter, loc, i32, source, stoch,
3079 existing, op.getStoreIndex());
3080 else if (typeIsExpectedFp8ForChipset(chipset, resultElemType))
3081 result = ROCDL::CvtSrFp8F32Op::create(rewriter, loc, i32, source, stoch,
3082 existing, op.getStoreIndex());
3083
3084 result = rewriter.replaceOpWithNewOp<LLVM::BitcastOp>(
3085 op, getTypeConverter()->convertType(resultType), result);
3086 return success();
3087}
3088
3089// Implement the AMDGPU_DPPLowering class that will convert the amdgpu.dpp
3090// operation into the corresponding ROCDL instructions.
3091struct AMDGPUDPPLowering : public ConvertOpToLLVMPattern<DPPOp> {
3092 AMDGPUDPPLowering(const LLVMTypeConverter &converter, Chipset chipset)
3093 : ConvertOpToLLVMPattern<DPPOp>(converter), chipset(chipset) {}
3094 Chipset chipset;
3095
3096 LogicalResult
3097 matchAndRewrite(DPPOp DppOp, DPPOp::Adaptor adaptor,
3098 ConversionPatternRewriter &rewriter) const override {
3099
3100 // Convert the source operand to the corresponding LLVM type
3101 Location loc = DppOp.getLoc();
3102 Value src = adaptor.getSrc();
3103 Value old = adaptor.getOld();
3104 Type srcType = src.getType();
3105 Type oldType = old.getType();
3106 Type llvmType = nullptr;
3107 if (srcType.getIntOrFloatBitWidth() < 32) {
3108 llvmType = rewriter.getI32Type();
3109 } else if (isa<FloatType>(srcType)) {
3110 llvmType = (srcType.getIntOrFloatBitWidth() == 32)
3111 ? rewriter.getF32Type()
3112 : rewriter.getF64Type();
3113 } else if (isa<IntegerType>(srcType)) {
3114 llvmType = (srcType.getIntOrFloatBitWidth() == 32)
3115 ? rewriter.getI32Type()
3116 : rewriter.getI64Type();
3117 }
3118 auto llvmSrcIntType = typeConverter->convertType(
3119 rewriter.getIntegerType(srcType.getIntOrFloatBitWidth()));
3120
3121 // If the source type is less of 32, use bitcast to convert it to i32.
3122 auto convertOperand = [&](Value operand, Type operandType) {
3123 if (operandType.getIntOrFloatBitWidth() <= 16) {
3124 if (llvm::isa<FloatType>(operandType)) {
3125 operand =
3126 LLVM::BitcastOp::create(rewriter, loc, llvmSrcIntType, operand);
3127 }
3128 auto llvmVecType = typeConverter->convertType(mlir::VectorType::get(
3129 32 / operandType.getIntOrFloatBitWidth(), llvmSrcIntType));
3130 Value undefVec = LLVM::UndefOp::create(rewriter, loc, llvmVecType);
3131 operand =
3132 LLVM::InsertElementOp::create(rewriter, loc, undefVec, operand,
3133 createI32Constant(rewriter, loc, 0));
3134 operand = LLVM::BitcastOp::create(rewriter, loc, llvmType, operand);
3135 }
3136 return operand;
3137 };
3138
3139 src = convertOperand(src, srcType);
3140 old = convertOperand(old, oldType);
3141
3142 // This is taken from the following file llvm/lib/Target/AMDGPU/SIDefines.h
3143 enum DppCtrl : unsigned {
3144 ROW_SHL0 = 0x100,
3145 ROW_SHR0 = 0x110,
3146 ROW_ROR0 = 0x120,
3147 WAVE_SHL1 = 0x130,
3148 WAVE_ROL1 = 0x134,
3149 WAVE_SHR1 = 0x138,
3150 WAVE_ROR1 = 0x13C,
3151 ROW_MIRROR = 0x140,
3152 ROW_HALF_MIRROR = 0x141,
3153 BCAST15 = 0x142,
3154 BCAST31 = 0x143,
3155 };
3156
3157 auto kind = DppOp.getKind();
3158 auto permArgument = DppOp.getPermArgument();
3159 uint32_t DppCtrl = 0;
3160
3161 switch (kind) {
3162
3163 case DPPPerm::quad_perm: {
3164 auto quadPermAttr = cast<ArrayAttr>(*permArgument);
3165 int32_t i = 0;
3166 for (auto elem : quadPermAttr.getAsRange<IntegerAttr>()) {
3167 uint32_t num = elem.getInt();
3168 DppCtrl |= num << (i * 2);
3169 i++;
3170 }
3171 break;
3172 }
3173 case DPPPerm::row_shl: {
3174 auto intAttr = cast<IntegerAttr>(*permArgument);
3175 DppCtrl = intAttr.getInt() + DppCtrl::ROW_SHL0;
3176 break;
3177 }
3178 case DPPPerm::row_shr: {
3179 auto intAttr = cast<IntegerAttr>(*permArgument);
3180 DppCtrl = intAttr.getInt() + DppCtrl::ROW_SHR0;
3181 break;
3182 }
3183 case DPPPerm::row_ror: {
3184 auto intAttr = cast<IntegerAttr>(*permArgument);
3185 DppCtrl = intAttr.getInt() + DppCtrl::ROW_ROR0;
3186 break;
3187 }
3188 case DPPPerm::wave_shl:
3189 DppCtrl = DppCtrl::WAVE_SHL1;
3190 break;
3191 case DPPPerm::wave_shr:
3192 DppCtrl = DppCtrl::WAVE_SHR1;
3193 break;
3194 case DPPPerm::wave_rol:
3195 DppCtrl = DppCtrl::WAVE_ROL1;
3196 break;
3197 case DPPPerm::wave_ror:
3198 DppCtrl = DppCtrl::WAVE_ROR1;
3199 break;
3200 case DPPPerm::row_mirror:
3201 DppCtrl = DppCtrl::ROW_MIRROR;
3202 break;
3203 case DPPPerm::row_half_mirror:
3204 DppCtrl = DppCtrl::ROW_HALF_MIRROR;
3205 break;
3206 case DPPPerm::row_bcast_15:
3207 DppCtrl = DppCtrl::BCAST15;
3208 break;
3209 case DPPPerm::row_bcast_31:
3210 DppCtrl = DppCtrl::BCAST31;
3211 break;
3212 }
3213
3214 // Check for row_mask, bank_mask, bound_ctrl if they exist and create
3215 // constants
3216 auto rowMask = DppOp->getAttrOfType<IntegerAttr>("row_mask").getInt();
3217 auto bankMask = DppOp->getAttrOfType<IntegerAttr>("bank_mask").getInt();
3218 bool boundCtrl = DppOp->getAttrOfType<BoolAttr>("bound_ctrl").getValue();
3219
3220 // create a ROCDL_DPPMovOp instruction with the appropriate attributes
3221 auto dppMovOp =
3222 ROCDL::DPPUpdateOp::create(rewriter, loc, llvmType, old, src, DppCtrl,
3223 rowMask, bankMask, boundCtrl);
3224
3225 Value result = dppMovOp.getRes();
3226 if (srcType.getIntOrFloatBitWidth() < 32) {
3227 result = LLVM::TruncOp::create(rewriter, loc, llvmSrcIntType, result);
3228 if (!llvm::isa<IntegerType>(srcType)) {
3229 result = LLVM::BitcastOp::create(rewriter, loc, srcType, result);
3230 }
3231 }
3232
3233 // We are replacing the AMDGPU_DPPOp instruction with the new
3234 // ROCDL_DPPMovOp instruction
3235 rewriter.replaceOp(DppOp, ValueRange(result));
3236 return success();
3237 }
3238};
3239
3240struct AMDGPUSwizzleBitModeLowering
3241 : public ConvertOpToLLVMPattern<SwizzleBitModeOp> {
3243
3244 LogicalResult
3245 matchAndRewrite(SwizzleBitModeOp op, OpAdaptor adaptor,
3246 ConversionPatternRewriter &rewriter) const override {
3247 Location loc = op.getLoc();
3248 Type i32 = rewriter.getI32Type();
3249 Value src = adaptor.getSrc();
3250 SmallVector<Value> decomposed;
3251 if (failed(LLVM::decomposeValue(rewriter, loc, src, i32, decomposed)))
3252 return rewriter.notifyMatchFailure(op,
3253 "failed to decompose value to i32");
3254 unsigned andMask = op.getAndMask();
3255 unsigned orMask = op.getOrMask();
3256 unsigned xorMask = op.getXorMask();
3257
3258 // bit 15 is 0 for the BitMode swizzle.
3259 // https://gpuopen.com/learn/amd-gcn-assembly-cross-lane-operations/
3260 unsigned mask = andMask | (orMask << 5) | (xorMask << 10);
3261 Value maskValue = createI32Constant(rewriter, loc, mask);
3262 SmallVector<Value> swizzled;
3263 for (Value v : decomposed) {
3264 Value res =
3265 ROCDL::DsSwizzleOp::create(rewriter, loc, v.getType(), v, maskValue);
3266 swizzled.emplace_back(res);
3267 }
3268
3269 Value result = LLVM::composeValue(rewriter, loc, swizzled, src.getType());
3270 rewriter.replaceOp(op, result);
3271 return success();
3272 }
3273};
3274
3275struct AMDGPUPermlaneLowering : public ConvertOpToLLVMPattern<PermlaneSwapOp> {
3277
3278 AMDGPUPermlaneLowering(const LLVMTypeConverter &converter, Chipset chipset)
3279 : ConvertOpToLLVMPattern<PermlaneSwapOp>(converter), chipset(chipset) {}
3280 Chipset chipset;
3281
3282 LogicalResult
3283 matchAndRewrite(PermlaneSwapOp op, OpAdaptor adaptor,
3284 ConversionPatternRewriter &rewriter) const override {
3285 if (chipset < kGfx950)
3286 return op->emitOpError("permlane_swap is only supported on gfx950+");
3287
3288 Location loc = op.getLoc();
3289 Type i32 = rewriter.getI32Type();
3290 Value src = adaptor.getSrc();
3291 unsigned rowLength = op.getRowLength();
3292 bool fi = op.getFetchInactive();
3293 bool boundctrl = op.getBoundCtrl();
3294
3295 SmallVector<Value> decomposed;
3296 if (failed(LLVM::decomposeValue(rewriter, loc, src, i32, decomposed)))
3297 return rewriter.notifyMatchFailure(op,
3298 "failed to decompose value to i32");
3299
3300 SmallVector<Value> permuted;
3301 for (Value v : decomposed) {
3302 Value res;
3303 Type i32pair = LLVM::LLVMStructType::getLiteral(
3304 rewriter.getContext(), {v.getType(), v.getType()});
3305
3306 if (rowLength == 16)
3307 res = ROCDL::Permlane16SwapOp::create(rewriter, loc, i32pair, v, v, fi,
3308 boundctrl);
3309 else if (rowLength == 32)
3310 res = ROCDL::Permlane32SwapOp::create(rewriter, loc, i32pair, v, v, fi,
3311 boundctrl);
3312 else
3313 llvm_unreachable("unsupported row length");
3314
3315 Value vdst0 = LLVM::ExtractValueOp::create(rewriter, loc, res, {0});
3316 Value vdst1 = LLVM::ExtractValueOp::create(rewriter, loc, res, {1});
3317
3318 Value isEqual = LLVM::ICmpOp::create(rewriter, loc,
3319 LLVM::ICmpPredicate::eq, vdst0, v);
3320
3321 // Per `permlane(16|32)` semantics: if the first extracted element equals
3322 // 'v', the result is the second element; otherwise it is the first.
3323 Value vdstNew =
3324 LLVM::SelectOp::create(rewriter, loc, isEqual, vdst1, vdst0);
3325 permuted.emplace_back(vdstNew);
3326 }
3327
3328 Value result = LLVM::composeValue(rewriter, loc, permuted, src.getType());
3329 rewriter.replaceOp(op, result);
3330 return success();
3331 }
3332};
3333
3334struct AMDGPUPermlaneVarLowering
3335 : public ConvertOpToLLVMPattern<PermlaneVarOp> {
3337
3338 AMDGPUPermlaneVarLowering(const LLVMTypeConverter &converter, Chipset chipset)
3339 : ConvertOpToLLVMPattern<PermlaneVarOp>(converter), chipset(chipset) {}
3340 Chipset chipset;
3341
3342 LogicalResult
3343 matchAndRewrite(PermlaneVarOp op, OpAdaptor adaptor,
3344 ConversionPatternRewriter &rewriter) const override {
3345 if (chipset < kGfx1200)
3346 return op->emitOpError("permlane_var is only supported on GFX12+");
3347
3348 Location loc = op.getLoc();
3349 Type i32 = rewriter.getI32Type();
3350 Value src = adaptor.getSrc();
3351 Value selector = adaptor.getSelector();
3352 bool cross = op.getCross();
3353 bool fi = op.getFetchInactive();
3354 bool boundCtrl = op.getBoundCtrl();
3355
3356 SmallVector<Value> decomposed;
3357 if (failed(LLVM::decomposeValue(rewriter, loc, src, i32, decomposed)))
3358 return rewriter.notifyMatchFailure(op,
3359 "failed to decompose value to i32");
3360
3361 SmallVector<Value> permuted;
3362 for (Value v : decomposed) {
3363 Value res;
3364 if (cross)
3365 res = ROCDL::PermlaneX16VarOp::create(rewriter, loc, i32, v, v,
3366 selector, fi, boundCtrl);
3367 else
3368 res = ROCDL::Permlane16VarOp::create(rewriter, loc, i32, v, v, selector,
3369 fi, boundCtrl);
3370 permuted.emplace_back(res);
3371 }
3372
3373 Value result = LLVM::composeValue(rewriter, loc, permuted, src.getType());
3374 rewriter.replaceOp(op, result);
3375 return success();
3376 }
3377};
3378
3379//===----------------------------------------------------------------------===//
3380// In-LDS Barrier Operations
3381//===----------------------------------------------------------------------===//
3382
3383// Bit layout of ds_barrier_state (as i64):
3384// [63:32] init count (32 bits)
3385// [31:29] phase (3 bits)
3386// [28:0] pending count (29 bits)
3387constexpr int32_t kDsBarrierPendingCountBitWidth = 29;
3388constexpr int32_t kDsBarrierPhasePos = kDsBarrierPendingCountBitWidth;
3389constexpr int32_t kDsBarrierInitCountPos = 32;
3390constexpr int32_t kDsBarrierPendingCountMask =
3391 (1 << kDsBarrierPendingCountBitWidth) - 1;
3392
3393struct DsBarrierInitOpLowering
3394 : public ConvertOpToLLVMPattern<DsBarrierInitOp> {
3395 Chipset chipset;
3396
3397 DsBarrierInitOpLowering(const LLVMTypeConverter &converter, Chipset chipset)
3398 : ConvertOpToLLVMPattern<DsBarrierInitOp>(converter), chipset(chipset) {}
3399
3400 LogicalResult
3401 matchAndRewrite(DsBarrierInitOp op, OpAdaptor adaptor,
3402 ConversionPatternRewriter &rewriter) const override {
3403 if (chipset < kGfx1250)
3404 return op->emitOpError("only supported on gfx1250+");
3405
3406 Location loc = op.getLoc();
3407 Type i64 = rewriter.getI64Type();
3408
3409 MemRefType memrefType = cast<MemRefType>(op.getBase().getType());
3410 Value ptr = getStridedElementPtr(rewriter, loc, memrefType,
3411 adaptor.getBase(), adaptor.getIndices());
3412
3413 // Note: We give participants as the number of arrivals that have to occur
3414 // before the phase changes. Hardware changes the phase when updating the
3415 // pending count would underflow, so we subtract 1 to get the behavior we're
3416 // looking for.
3417 Value initCount =
3418 LLVM::SubOp::create(rewriter, loc, adaptor.getParticipants(),
3419 createI32Constant(rewriter, loc, 1));
3420
3421 // Just a bit of paranoia, but this also allows for configurable width if
3422 // that becomes a thing.
3423 Value countMask =
3424 createI32Constant(rewriter, loc, kDsBarrierPendingCountMask);
3425 Value maskedCount32 =
3426 LLVM::AndOp::create(rewriter, loc, initCount, countMask);
3427 Value maskedCount = LLVM::ZExtOp::create(rewriter, loc, i64, maskedCount32);
3428
3429 Value initCountShifted = LLVM::ShlOp::create(
3430 rewriter, loc, maskedCount,
3431 createI64Constant(rewriter, loc, kDsBarrierInitCountPos));
3432 Value barrierState =
3433 LLVM::OrOp::create(rewriter, loc, initCountShifted, maskedCount);
3434
3435 LLVM::StoreOp::create(
3436 rewriter, loc, barrierState, ptr, /*alignment=*/8, /*isVolatile=*/false,
3437 /*isNonTemporal=*/false,
3438 /*isInvariantGroup=*/false, LLVM::AtomicOrdering::release,
3439 /*syncscope=*/"workgroup");
3440
3441 rewriter.eraseOp(op);
3442 return success();
3443 }
3444};
3445
3446struct DsBarrierPollStateOpLowering
3447 : public ConvertOpToLLVMPattern<DsBarrierPollStateOp> {
3448 Chipset chipset;
3449
3450 DsBarrierPollStateOpLowering(const LLVMTypeConverter &converter,
3451 Chipset chipset)
3452 : ConvertOpToLLVMPattern<DsBarrierPollStateOp>(converter),
3453 chipset(chipset) {}
3454
3455 LogicalResult
3456 matchAndRewrite(DsBarrierPollStateOp op, OpAdaptor adaptor,
3457 ConversionPatternRewriter &rewriter) const override {
3458 if (chipset < kGfx1250)
3459 return op->emitOpError("only supported on gfx1250+");
3460
3461 Location loc = op.getLoc();
3462 Type i64 = rewriter.getI64Type();
3463
3464 MemRefType memrefType = cast<MemRefType>(op.getBase().getType());
3465 Value ptr = getStridedElementPtr(rewriter, loc, memrefType,
3466 adaptor.getBase(), adaptor.getIndices());
3467
3468 // Atomic load with workgroup scope and acquire ordering should be what
3469 // we're looking for.
3470 rewriter.replaceOpWithNewOp<LLVM::LoadOp>(
3471 op, i64, ptr, /*alignment=*/8, /*volatile_=*/false,
3472 /*nontemporal=*/false, /*invariant=*/false,
3473 /*invariantGroup=*/false, LLVM::AtomicOrdering::acquire,
3474 /*syncscope=*/"workgroup");
3475 return success();
3476 }
3477};
3478
3479struct DsAsyncBarrierArriveOpLowering
3480 : public ConvertOpToLLVMPattern<DsAsyncBarrierArriveOp> {
3481 Chipset chipset;
3482
3483 DsAsyncBarrierArriveOpLowering(const LLVMTypeConverter &converter,
3484 Chipset chipset)
3485 : ConvertOpToLLVMPattern<DsAsyncBarrierArriveOp>(converter),
3486 chipset(chipset) {}
3487
3488 LogicalResult
3489 matchAndRewrite(DsAsyncBarrierArriveOp op, OpAdaptor adaptor,
3490 ConversionPatternRewriter &rewriter) const override {
3491 if (chipset < kGfx1250)
3492 return op->emitOpError("only supported on gfx1250+");
3493
3494 Location loc = op.getLoc();
3495
3496 MemRefType memrefType = cast<MemRefType>(op.getBase().getType());
3497 Value ptr = getStridedElementPtr(rewriter, loc, memrefType,
3498 adaptor.getBase(), adaptor.getIndices());
3499
3500 rewriter.replaceOpWithNewOp<ROCDL::DsAtomicAsyncBarrierArriveOp>(
3501 op, ptr, /*alias_scopes=*/nullptr, /*noalias_scopes=*/nullptr,
3502 /*tbaa=*/nullptr);
3503 return success();
3504 }
3505};
3506
3507struct DsBarrierArriveOpLowering
3508 : public ConvertOpToLLVMPattern<DsBarrierArriveOp> {
3509 Chipset chipset;
3510
3511 DsBarrierArriveOpLowering(const LLVMTypeConverter &converter, Chipset chipset)
3512 : ConvertOpToLLVMPattern<DsBarrierArriveOp>(converter), chipset(chipset) {
3513 }
3514
3515 LogicalResult
3516 matchAndRewrite(DsBarrierArriveOp op, OpAdaptor adaptor,
3517 ConversionPatternRewriter &rewriter) const override {
3518 if (chipset < kGfx1250)
3519 return op->emitOpError("only supported on gfx1250+");
3520
3521 Location loc = op.getLoc();
3522 Type i64 = rewriter.getI64Type();
3523
3524 MemRefType memrefType = cast<MemRefType>(op.getBase().getType());
3525 Value ptr = getStridedElementPtr(rewriter, loc, memrefType,
3526 adaptor.getBase(), adaptor.getIndices());
3527
3528 rewriter.replaceOpWithNewOp<ROCDL::DsAtomicBarrierArriveRtnOp>(
3529 op, i64, ptr, adaptor.getCount(), /*alias_scopes=*/nullptr,
3530 /*noalias_scopes=*/nullptr, /*tbaa=*/nullptr);
3531 return success();
3532 }
3533};
3534
3535struct DsBarrierStatePhaseOpLowering
3536 : public ConvertOpToLLVMPattern<DsBarrierStatePhaseOp> {
3538
3539 LogicalResult
3540 matchAndRewrite(DsBarrierStatePhaseOp op, OpAdaptor adaptor,
3541 ConversionPatternRewriter &rewriter) const override {
3542 Location loc = op.getLoc();
3543 Type i32 = rewriter.getI32Type();
3544
3545 Value state = adaptor.getState();
3546
3547 Value noInitCount = LLVM::TruncOp::create(rewriter, loc, i32, state);
3548 Value phase = LLVM::LShrOp::create(
3549 rewriter, loc, noInitCount,
3550 createI32Constant(rewriter, loc, kDsBarrierPhasePos));
3551
3552 rewriter.replaceOp(op, phase);
3553 return success();
3554 }
3555};
3556
3557struct DsBarrierStatePendingCountOpLowering
3558 : public ConvertOpToLLVMPattern<DsBarrierStatePendingCountOp> {
3560
3561 LogicalResult
3562 matchAndRewrite(DsBarrierStatePendingCountOp op, OpAdaptor adaptor,
3563 ConversionPatternRewriter &rewriter) const override {
3564 Location loc = op.getLoc();
3565 Type i32 = rewriter.getI32Type();
3566
3567 Value state = adaptor.getState();
3568
3569 Value noInitCount = LLVM::TruncOp::create(rewriter, loc, i32, state);
3570 Value pendingCount = LLVM::AndOp::create(
3571 rewriter, loc, noInitCount,
3572 createI32Constant(rewriter, loc,
3573 static_cast<uint32_t>(kDsBarrierPendingCountMask)));
3574
3575 rewriter.replaceOp(op, pendingCount);
3576 return success();
3577 }
3578};
3579
3580struct DsBarrierStateInitCountOpLowering
3581 : public ConvertOpToLLVMPattern<DsBarrierStateInitCountOp> {
3583
3584 LogicalResult
3585 matchAndRewrite(DsBarrierStateInitCountOp op, OpAdaptor adaptor,
3586 ConversionPatternRewriter &rewriter) const override {
3587 Location loc = op.getLoc();
3588 Type i32 = rewriter.getI32Type();
3589
3590 Value state = adaptor.getState();
3591
3592 Value initCountI64 = LLVM::LShrOp::create(
3593 rewriter, loc, state,
3594 createI64Constant(rewriter, loc, kDsBarrierInitCountPos));
3595 Value initCount = LLVM::TruncOp::create(rewriter, loc, i32, initCountI64);
3596
3597 rewriter.replaceOp(op, initCount);
3598 return success();
3599 }
3600};
3601
3602struct DsBarrierStatePhaseParityLowering
3603 : public ConvertOpToLLVMPattern<DsBarrierStatePhaseParity> {
3605
3606 LogicalResult
3607 matchAndRewrite(DsBarrierStatePhaseParity op, OpAdaptor adaptor,
3608 ConversionPatternRewriter &rewriter) const override {
3609 Location loc = op.getLoc();
3610 Type i1 = rewriter.getI1Type();
3611
3612 Value state = adaptor.getState();
3613
3614 Value noInitCount =
3615 LLVM::TruncOp::create(rewriter, loc, rewriter.getI32Type(), state);
3616 Value phase = LLVM::LShrOp::create(
3617 rewriter, loc, noInitCount,
3618 createI32Constant(rewriter, loc, kDsBarrierPhasePos));
3619 Value parity = LLVM::TruncOp::create(rewriter, loc, i1, phase);
3620
3621 rewriter.replaceOp(op, parity);
3622 return success();
3623 }
3624};
3625
3626//===----------------------------------------------------------------------===//
3627// Tensor Data Mover (TDM)
3628//===----------------------------------------------------------------------===//
3629
3630static Value setValueAtOffset(ConversionPatternRewriter &rewriter, Location loc,
3631 Value accumulator, Value value, int64_t shift) {
3632 shift = shift % 32;
3633 Value shiftAmount;
3634 if (shift != 0) {
3635 shiftAmount = createI32Constant(rewriter, loc, shift % 32);
3636 value = LLVM::ShlOp::create(rewriter, loc, value, shiftAmount);
3637 }
3638
3639 if (matchPattern(accumulator, mlir::m_Zero()))
3640 return value;
3641
3642 constexpr bool isDisjoint = true;
3643 return LLVM::OrOp::create(rewriter, loc, accumulator, value, isDisjoint);
3644}
3645
3646template <typename BaseOp>
3647struct AMDGPUMakeDmaBaseLowering : public ConvertOpToLLVMPattern<BaseOp> {
3648 using ConvertOpToLLVMPattern<BaseOp>::ConvertOpToLLVMPattern;
3649 using Adaptor = typename ConvertOpToLLVMPattern<BaseOp>::OpAdaptor;
3650
3651 AMDGPUMakeDmaBaseLowering(const LLVMTypeConverter &converter, Chipset chipset)
3652 : ConvertOpToLLVMPattern<BaseOp>(converter), chipset(chipset) {}
3653 Chipset chipset;
3654
3655 LogicalResult
3656 matchAndRewrite(BaseOp op, Adaptor adaptor,
3657 ConversionPatternRewriter &rewriter) const override {
3658 if (chipset < kGfx1250)
3659 return op->emitOpError("make_dma_base is only supported on gfx1250");
3660
3661 Location loc = op.getLoc();
3662
3663 constexpr int32_t constlen = 4;
3664 Value consts[constlen];
3665 for (int64_t i = 0; i < constlen; ++i)
3666 consts[i] = createI32Constant(rewriter, loc, i);
3667
3668 constexpr int32_t sgprslen = constlen;
3669 Value sgprs[sgprslen];
3670 for (int64_t i = 0; i < sgprslen; ++i) {
3671 sgprs[i] = consts[0];
3672 }
3673
3674 sgprs[0] = consts[1];
3675
3676 if constexpr (BaseOp::isGather()) {
3677 sgprs[0] = setValueAtOffset(rewriter, loc, sgprs[0], consts[1], 30);
3678
3679 auto type = cast<TDMGatherBaseType>(op.getResult().getType());
3680 Type indexType = type.getIndexType();
3681 unsigned indexSize = indexType.getIntOrFloatBitWidth();
3682 assert(llvm::is_contained({16u, 32u}, indexSize) &&
3683 "expected index_size to be 16 or 32");
3684 unsigned idx = (indexSize / 16) - 1;
3685
3686 if (idx)
3687 sgprs[0] = setValueAtOffset(rewriter, loc, sgprs[0], consts[1], 31);
3688 }
3689
3690 ValueRange ldsIndices = adaptor.getLdsIndices();
3691 Value lds = adaptor.getLds();
3692 auto ldsMemRefType = cast<MemRefType>(op.getLds().getType());
3693
3695 rewriter, loc, ldsMemRefType, lds, ldsIndices);
3696
3697 ValueRange globalIndices = adaptor.getGlobalIndices();
3698 Value global = adaptor.getGlobal();
3699 auto globalMemRefType = cast<MemRefType>(op.getGlobal().getType());
3700
3702 rewriter, loc, globalMemRefType, global, globalIndices);
3703
3704 Type i32 = rewriter.getI32Type();
3705 Type i64 = rewriter.getI64Type();
3706
3707 sgprs[1] = LLVM::PtrToIntOp::create(rewriter, loc, i32, ldsPtr);
3708 Value castForGlobalAddr =
3709 LLVM::PtrToIntOp::create(rewriter, loc, i64, globalPtr);
3710
3711 sgprs[2] = LLVM::TruncOp::create(rewriter, loc, i32, castForGlobalAddr);
3712
3713 Value shift = LLVM::LShrOp::create(rewriter, loc, castForGlobalAddr,
3714 createI64Constant(rewriter, loc, 32));
3715
3716 Value highHalf = LLVM::TruncOp::create(rewriter, loc, i32, shift);
3717
3718 Value mask = createI32Constant(rewriter, loc, (1ull << 25) - 1);
3719 highHalf = LLVM::AndOp::create(rewriter, loc, highHalf, mask);
3720
3721 sgprs[3] = setValueAtOffset(rewriter, loc, highHalf, consts[2], 30);
3722
3723 Type v4i32 = this->typeConverter->convertType(VectorType::get(4, i32));
3724 assert(v4i32 && "expected type conversion to succeed");
3725 Value result = LLVM::PoisonOp::create(rewriter, loc, v4i32);
3726
3727 for (auto [sgpr, constant] : llvm::zip_equal(sgprs, consts))
3728 result =
3729 LLVM::InsertElementOp::create(rewriter, loc, result, sgpr, constant);
3730
3731 rewriter.replaceOp(op, result);
3732 return success();
3733 }
3734};
3735
3736template <typename DescriptorOp>
3737struct AMDGPULowerDescriptor : public ConvertOpToLLVMPattern<DescriptorOp> {
3738 using ConvertOpToLLVMPattern<DescriptorOp>::ConvertOpToLLVMPattern;
3739 using OpAdaptor = typename ConvertOpToLLVMPattern<DescriptorOp>::OpAdaptor;
3740
3741 AMDGPULowerDescriptor(const LLVMTypeConverter &converter, Chipset chipset)
3742 : ConvertOpToLLVMPattern<DescriptorOp>(converter), chipset(chipset) {}
3743 Chipset chipset;
3744
3745 Value getDGroup0(OpAdaptor adaptor) const { return adaptor.getBase(); }
3746
3747 Value setWorkgroupMask(DescriptorOp op, OpAdaptor adaptor,
3748 ConversionPatternRewriter &rewriter, Location loc,
3749 Value sgpr0) const {
3750 Value mask = op.getWorkgroupMask();
3751 if (!mask)
3752 return sgpr0;
3753
3754 Type i16 = rewriter.getI16Type();
3755 mask = LLVM::BitcastOp::create(rewriter, loc, i16, mask);
3756 Type i32 = rewriter.getI32Type();
3757 Value extendedMask = LLVM::ZExtOp::create(rewriter, loc, i32, mask);
3758 return setValueAtOffset(rewriter, loc, sgpr0, extendedMask, 0);
3759 }
3760
3761 Value setDataSize(DescriptorOp op, OpAdaptor adaptor,
3762 ConversionPatternRewriter &rewriter, Location loc,
3763 Value sgpr0, ArrayRef<Value> consts) const {
3764 unsigned elementTypeWidthInBits = op.getElementTypeWidth();
3765 assert(llvm::is_contained({8u, 16u, 32u, 64u}, elementTypeWidthInBits) &&
3766 "expected type width to be 8, 16, 32, or 64.");
3767 int64_t idx = llvm::Log2_32(elementTypeWidthInBits / 8);
3768 Value size = consts[idx];
3769 return setValueAtOffset(rewriter, loc, sgpr0, size, 16);
3770 }
3771
3772 Value setAtomicBarrier(DescriptorOp op, OpAdaptor adaptor,
3773 ConversionPatternRewriter &rewriter, Location loc,
3774 Value sgpr0, ArrayRef<Value> consts) const {
3775 if (!adaptor.getAtomicBarrierAddress())
3776 return sgpr0;
3777
3778 return setValueAtOffset(rewriter, loc, sgpr0, consts[1], 18);
3779 }
3780
3781 Value setIterateEnable(DescriptorOp op, OpAdaptor adaptor,
3782 ConversionPatternRewriter &rewriter, Location loc,
3783 Value sgpr0, ArrayRef<Value> consts) const {
3784 if (!adaptor.getGlobalIncrement())
3785 return sgpr0;
3786
3787 // Value is ignored when in gather mode.
3788 // TODO: emit error earlier?
3789 return setValueAtOffset(rewriter, loc, sgpr0, consts[1], 19);
3790 }
3791
3792 Value setPadEnable(DescriptorOp op, OpAdaptor adaptor,
3793 ConversionPatternRewriter &rewriter, Location loc,
3794 Value sgpr0, ArrayRef<Value> consts) const {
3795 if (!op.getPadAmount())
3796 return sgpr0;
3797
3798 return setValueAtOffset(rewriter, loc, sgpr0, consts[1], 20);
3799 }
3800
3801 Value setEarlyTimeout(DescriptorOp op, OpAdaptor adaptor,
3802 ConversionPatternRewriter &rewriter, Location loc,
3803 Value sgpr0, ArrayRef<Value> consts) const {
3804 if (!op.getWorkgroupMask())
3805 return sgpr0;
3806
3807 return setValueAtOffset(rewriter, loc, sgpr0, consts[1], 21);
3808 }
3809
3810 Value setPadInterval(DescriptorOp op, OpAdaptor adaptor,
3811 ConversionPatternRewriter &rewriter, Location loc,
3812 Value sgpr0, ArrayRef<Value> consts) const {
3813 if (!op.getPadAmount())
3814 return sgpr0;
3815
3816 // pre-condition: padInterval can be a power of two between 2 and 256.
3817 // TODO: Validation if the value breaks the pre-condition.
3818 // If the pre-condition fails, there is a possibility of
3819 // affecting the higher bits. In a following PR implement
3820 // RuntimeVerifiableOpInterface that instruments conditions that need to be
3821 // checked at runtime.
3822 IntegerType i32 = rewriter.getI32Type();
3823 Value padInterval = adaptor.getPadInterval();
3824 padInterval = LLVM::CountTrailingZerosOp::create(rewriter, loc, i32,
3825 padInterval, false);
3826 padInterval = LLVM::SubOp::create(rewriter, loc, padInterval, consts[1]);
3827 // post-condition: padInterval can be a value between 0 and 7.
3828 return setValueAtOffset(rewriter, loc, sgpr0, padInterval, 22);
3829 }
3830
3831 Value setPadAmount(DescriptorOp op, OpAdaptor adaptor,
3832 ConversionPatternRewriter &rewriter, Location loc,
3833 Value sgpr0, ArrayRef<Value> consts) const {
3834 if (!op.getPadAmount())
3835 return sgpr0;
3836
3837 // pre-condition: padAmount is a value between 1-128.
3838 // TODO: Validation if the value breaks the pre-condition.
3839 // If the pre-condition fails, there is a possibility of
3840 // affecting the higher bits. In a following PR implement
3841 // RuntimeVerifiableOpInterface that instruments conditions that need to be
3842 // checked at runtime.
3843 Value padAmount = adaptor.getPadAmount();
3844 padAmount = LLVM::SubOp::create(rewriter, loc, padAmount, consts[1]);
3845 // post-condition: padAmount is a value between 0-127.
3846 return setValueAtOffset(rewriter, loc, sgpr0, padAmount, 25);
3847 }
3848
3849 Value setAtomicBarrierAddress(DescriptorOp op, OpAdaptor adaptor,
3850 ConversionPatternRewriter &rewriter,
3851 Location loc, Value sgpr1,
3852 ArrayRef<Value> consts) const {
3853 if (!adaptor.getAtomicBarrierAddress())
3854 return sgpr1;
3855
3856 Value atomicBarrierAddress = adaptor.getAtomicBarrierAddress();
3857 auto barrierAddressTy =
3858 cast<MemRefType>(op.getAtomicBarrierAddress().getType());
3859 ValueRange atomicBarrierIndices = adaptor.getAtomicBarrierIndices();
3860 atomicBarrierAddress = ConvertToLLVMPattern::getStridedElementPtr(
3861 rewriter, loc, barrierAddressTy, atomicBarrierAddress,
3862 atomicBarrierIndices);
3863 IntegerType i32 = rewriter.getI32Type();
3864 // pre-condition: atomicBarrierAddress is aligned to 8 bytes which implies
3865 // that the 3 LSBs are zero.
3866 // TODO: Validation if the value breaks the pre-condition.
3867 // In a following PR implement RuntimeVerifiableOpInterface
3868 // that instruments conditions that need to be checked at runtime.
3869 atomicBarrierAddress =
3870 LLVM::PtrToIntOp::create(rewriter, loc, i32, atomicBarrierAddress);
3871 atomicBarrierAddress =
3872 LLVM::LShrOp::create(rewriter, loc, atomicBarrierAddress, consts[3]);
3873 Value mask = createI32Constant(rewriter, loc, 0xFFFF);
3874 atomicBarrierAddress =
3875 LLVM::AndOp::create(rewriter, loc, atomicBarrierAddress, mask);
3876 return setValueAtOffset(rewriter, loc, sgpr1, atomicBarrierAddress, 32);
3877 }
3878
3879 std::pair<Value, Value> setTensorDimX(DescriptorOp op, OpAdaptor adaptor,
3880 ConversionPatternRewriter &rewriter,
3881 Location loc, Value sgpr1, Value sgpr2,
3882 ArrayRef<Value> consts, uint64_t dimX,
3883 uint32_t offset) const {
3884 ArrayRef<int64_t> globalStaticSizes = adaptor.getGlobalStaticSizes();
3885 ValueRange globalDynamicSizes = adaptor.getGlobalDynamicSizes();
3886 SmallVector<OpFoldResult> mixedGlobalSizes =
3887 getMixedValues(globalStaticSizes, globalDynamicSizes, rewriter);
3888 if (mixedGlobalSizes.size() <= dimX)
3889 return {sgpr1, sgpr2};
3890
3891 OpFoldResult tensorDimXOpFoldResult = *(mixedGlobalSizes.rbegin() + dimX);
3892 // pre-condition: tensorDimX is less than 2^32-1
3893 // TODO: Validation if the value breaks the pre-condition.
3894 // In a following PR implement RuntimeVerifiableOpInterface that instruments
3895 // conditions that need to be checked at runtime. This could also be fixed
3896 // by saying that mixedGlobalSizes is a DynamicI32List.
3897 Value tensorDimX;
3898 if (auto attr = dyn_cast<Attribute>(tensorDimXOpFoldResult)) {
3899 tensorDimX =
3900 createI32Constant(rewriter, loc, cast<IntegerAttr>(attr).getInt());
3901 } else {
3902 IntegerType i32 = rewriter.getI32Type();
3903 tensorDimX = cast<Value>(tensorDimXOpFoldResult);
3904 tensorDimX = LLVM::TruncOp::create(rewriter, loc, i32, tensorDimX);
3905 }
3906
3907 sgpr1 = setValueAtOffset(rewriter, loc, sgpr1, tensorDimX, offset);
3908
3909 Value c16 = createI32Constant(rewriter, loc, 16);
3910 Value tensorDimXHigh = LLVM::LShrOp::create(rewriter, loc, tensorDimX, c16);
3911 sgpr2 = setValueAtOffset(rewriter, loc, sgpr2, tensorDimXHigh, offset + 16);
3912 return {sgpr1, sgpr2};
3913 }
3914
3915 std::pair<Value, Value> setTensorDim0(DescriptorOp op, OpAdaptor adaptor,
3916 ConversionPatternRewriter &rewriter,
3917 Location loc, Value sgpr1, Value sgpr2,
3918 ArrayRef<Value> consts) const {
3919 return setTensorDimX(op, adaptor, rewriter, loc, sgpr1, sgpr2, consts, 0,
3920 48);
3921 }
3922
3923 std::pair<Value, Value> setTensorDim1(DescriptorOp op, OpAdaptor adaptor,
3924 ConversionPatternRewriter &rewriter,
3925 Location loc, Value sgpr2, Value sgpr3,
3926 ArrayRef<Value> consts) const {
3927 return setTensorDimX(op, adaptor, rewriter, loc, sgpr2, sgpr3, consts, 1,
3928 80);
3929 }
3930
3931 Value setTileDimX(DescriptorOp op, OpAdaptor adaptor,
3932 ConversionPatternRewriter &rewriter, Location loc,
3933 Value sgpr, ArrayRef<Value> consts, size_t dimX,
3934 int64_t offset) const {
3935 ArrayRef<int64_t> sharedStaticSizes = adaptor.getSharedStaticSizes();
3936 ValueRange sharedDynamicSizes = adaptor.getSharedDynamicSizes();
3937 SmallVector<OpFoldResult> mixedSharedSizes =
3938 getMixedValues(sharedStaticSizes, sharedDynamicSizes, rewriter);
3939 if (mixedSharedSizes.size() <= dimX)
3940 return sgpr;
3941
3942 OpFoldResult tileDimXOpFoldResult = *(mixedSharedSizes.rbegin() + dimX);
3943 // pre-condition: tileDimX is less than 2^16-1
3944 // TODO: Validation if the value breaks the pre-condition.
3945 // If the pre-condition fails, there is a possibility of
3946 // affecting the higher bits. In a following PR implement
3947 // RuntimeVerifiableOpInterface that instruments conditions that need to be
3948 // checked at runtime. This could also be fixed by saying that
3949 // mixedSharedSizes is a DynamicI16List.
3950 Value tileDimX;
3951 if (auto attr = dyn_cast<Attribute>(tileDimXOpFoldResult)) {
3952 tileDimX =
3953 createI32Constant(rewriter, loc, cast<IntegerAttr>(attr).getInt());
3954 } else {
3955 IntegerType i32 = rewriter.getI32Type();
3956 tileDimX = cast<Value>(tileDimXOpFoldResult);
3957 tileDimX = LLVM::TruncOp::create(rewriter, loc, i32, tileDimX);
3958 }
3959
3960 return setValueAtOffset(rewriter, loc, sgpr, tileDimX, offset);
3961 }
3962
3963 Value setTileDim0(DescriptorOp op, OpAdaptor adaptor,
3964 ConversionPatternRewriter &rewriter, Location loc,
3965 Value sgpr3, ArrayRef<Value> consts) const {
3966 return setTileDimX(op, adaptor, rewriter, loc, sgpr3, consts, 0, 112);
3967 }
3968
3969 Value setTileDim1(DescriptorOp op, OpAdaptor adaptor,
3970 ConversionPatternRewriter &rewriter, Location loc,
3971 Value sgpr4, ArrayRef<Value> consts) const {
3972 return setTileDimX(op, adaptor, rewriter, loc, sgpr4, consts, 1, 128);
3973 }
3974
3975 Value setValidIndices(DescriptorOp op, OpAdaptor adaptor,
3976 ConversionPatternRewriter &rewriter, Location loc,
3977 Value sgpr4, ArrayRef<Value> consts) const {
3978 auto type = cast<VectorType>(op.getIndices().getType());
3979 ArrayRef<int64_t> shape = type.getShape();
3980 assert(shape.size() == 1 && "expected shape to be of rank 1.");
3981 unsigned length = shape.back();
3982 assert(0 < length && length <= 16 && "expected length to be at most 16.");
3983 Value value = createI32Constant(rewriter, loc, length);
3984 return setValueAtOffset(rewriter, loc, sgpr4, value, 128);
3985 }
3986
3987 Value setTileDim1OrValidIndices(DescriptorOp op, OpAdaptor adaptor,
3988 ConversionPatternRewriter &rewriter,
3989 Location loc, Value sgpr4,
3990 ArrayRef<Value> consts) const {
3991 if constexpr (DescriptorOp::isGather())
3992 return setValidIndices(op, adaptor, rewriter, loc, sgpr4, consts);
3993 return setTileDim1(op, adaptor, rewriter, loc, sgpr4, consts);
3994 }
3995
3996 Value setTileDim2(DescriptorOp op, OpAdaptor adaptor,
3997 ConversionPatternRewriter &rewriter, Location loc,
3998 Value sgpr4, ArrayRef<Value> consts) const {
3999 // Value is ignored when in gather mode.
4000 if constexpr (DescriptorOp::isGather())
4001 return sgpr4;
4002 return setTileDimX(op, adaptor, rewriter, loc, sgpr4, consts, 2, 144);
4003 }
4004
4005 std::pair<Value, Value>
4006 setTensorDimXStride(DescriptorOp op, OpAdaptor adaptor,
4007 ConversionPatternRewriter &rewriter, Location loc,
4008 Value sgprY, Value sgprZ, ArrayRef<Value> consts,
4009 size_t dimX, int64_t offset) const {
4010 ArrayRef<int64_t> globalStaticStrides = adaptor.getGlobalStaticStrides();
4011 ValueRange globalDynamicStrides = adaptor.getGlobalDynamicStrides();
4012 SmallVector<OpFoldResult> mixedGlobalStrides =
4013 getMixedValues(globalStaticStrides, globalDynamicStrides, rewriter);
4014
4015 if (mixedGlobalStrides.size() <= (dimX + 1))
4016 return {sgprY, sgprZ};
4017
4018 OpFoldResult tensorDimXStrideOpFoldResult =
4019 *(mixedGlobalStrides.rbegin() + dimX + 1);
4020 // pre-condition: tensorDimXStride is less than 2^48-1
4021 // TODO: Validation if the value breaks the pre-condition.
4022 // In a following PR implement RuntimeVerifiableOpInterface that instruments
4023 // conditions that need to be checked at runtime.
4024 Value tensorDimXStride;
4025 if (auto attr = dyn_cast<Attribute>(tensorDimXStrideOpFoldResult))
4026 tensorDimXStride =
4027 createI64Constant(rewriter, loc, cast<IntegerAttr>(attr).getInt());
4028 else
4029 tensorDimXStride = cast<Value>(tensorDimXStrideOpFoldResult);
4030
4031 constexpr int64_t first48bits = (1ll << 48) - 1;
4032 Value mask = createI64Constant(rewriter, loc, first48bits);
4033 tensorDimXStride =
4034 LLVM::AndOp::create(rewriter, loc, mask, tensorDimXStride);
4035 IntegerType i32 = rewriter.getI32Type();
4036 Value tensorDimXStrideLow =
4037 LLVM::TruncOp::create(rewriter, loc, i32, tensorDimXStride);
4038 sgprY = setValueAtOffset(rewriter, loc, sgprY, tensorDimXStrideLow, offset);
4039
4040 int64_t shift = (offset % 32) == 0 ? 32 : offset % 32;
4041 Value shiftVal = createI64Constant(rewriter, loc, shift);
4042 Value tensorDimXStrideHigh =
4043 LLVM::LShrOp::create(rewriter, loc, tensorDimXStride, shiftVal);
4044 tensorDimXStrideHigh =
4045 LLVM::TruncOp::create(rewriter, loc, i32, tensorDimXStrideHigh);
4046 sgprZ = setValueAtOffset(rewriter, loc, sgprZ, tensorDimXStrideHigh,
4047 offset + shift);
4048 return {sgprY, sgprZ};
4049 }
4050
4051 std::pair<Value, Value>
4052 setTensorDim0Stride(DescriptorOp op, OpAdaptor adaptor,
4053 ConversionPatternRewriter &rewriter, Location loc,
4054 Value sgpr5, Value sgpr6, ArrayRef<Value> consts) const {
4055 return setTensorDimXStride(op, adaptor, rewriter, loc, sgpr5, sgpr6, consts,
4056 0, 160);
4057 }
4058
4059 std::pair<Value, Value>
4060 setTensorDim1Stride(DescriptorOp op, OpAdaptor adaptor,
4061 ConversionPatternRewriter &rewriter, Location loc,
4062 Value sgpr5, Value sgpr6, ArrayRef<Value> consts) const {
4063 // Value is ignored when in gather mode.
4064 if constexpr (DescriptorOp::isGather())
4065 return {sgpr5, sgpr6};
4066 return setTensorDimXStride(op, adaptor, rewriter, loc, sgpr5, sgpr6, consts,
4067 1, 208);
4068 }
4069
4070 Value getDGroup1(DescriptorOp op, OpAdaptor adaptor,
4071 ConversionPatternRewriter &rewriter, Location loc,
4072 ArrayRef<Value> consts) const {
4073 Value sgprs[8];
4074 for (int64_t i = 0; i < 8; ++i) {
4075 sgprs[i] = consts[0];
4076 }
4077
4078 sgprs[0] = setWorkgroupMask(op, adaptor, rewriter, loc, sgprs[0]);
4079 sgprs[0] = setDataSize(op, adaptor, rewriter, loc, sgprs[0], consts);
4080 sgprs[0] = setAtomicBarrier(op, adaptor, rewriter, loc, sgprs[0], consts);
4081 sgprs[0] = setIterateEnable(op, adaptor, rewriter, loc, sgprs[0], consts);
4082 sgprs[0] = setPadEnable(op, adaptor, rewriter, loc, sgprs[0], consts);
4083 sgprs[0] = setEarlyTimeout(op, adaptor, rewriter, loc, sgprs[0], consts);
4084 sgprs[0] = setPadInterval(op, adaptor, rewriter, loc, sgprs[0], consts);
4085 sgprs[0] = setPadAmount(op, adaptor, rewriter, loc, sgprs[0], consts);
4086
4087 sgprs[1] =
4088 setAtomicBarrierAddress(op, adaptor, rewriter, loc, sgprs[1], consts);
4089 std::tie(sgprs[1], sgprs[2]) =
4090 setTensorDim0(op, adaptor, rewriter, loc, sgprs[1], sgprs[2], consts);
4091 std::tie(sgprs[2], sgprs[3]) =
4092 setTensorDim1(op, adaptor, rewriter, loc, sgprs[2], sgprs[3], consts);
4093
4094 sgprs[3] = setTileDim0(op, adaptor, rewriter, loc, sgprs[3], consts);
4095 sgprs[4] =
4096 setTileDim1OrValidIndices(op, adaptor, rewriter, loc, sgprs[4], consts);
4097 sgprs[4] = setTileDim2(op, adaptor, rewriter, loc, sgprs[4], consts);
4098 std::tie(sgprs[5], sgprs[6]) = setTensorDim0Stride(
4099 op, adaptor, rewriter, loc, sgprs[5], sgprs[6], consts);
4100 std::tie(sgprs[6], sgprs[7]) = setTensorDim1Stride(
4101 op, adaptor, rewriter, loc, sgprs[6], sgprs[7], consts);
4102
4103 IntegerType i32 = rewriter.getI32Type();
4104 Type v8i32 = this->typeConverter->convertType(VectorType::get(8, i32));
4105 assert(v8i32 && "expected type conversion to succeed");
4106 Value dgroup1 = LLVM::PoisonOp::create(rewriter, loc, v8i32);
4107
4108 for (auto [sgpr, constant] : llvm::zip_equal(sgprs, consts)) {
4109 dgroup1 =
4110 LLVM::InsertElementOp::create(rewriter, loc, dgroup1, sgpr, constant);
4111 }
4112
4113 return dgroup1;
4114 }
4115
4116 Value setTensorDimX(DescriptorOp op, OpAdaptor adaptor,
4117 ConversionPatternRewriter &rewriter, Location loc,
4118 Value sgpr0, ArrayRef<Value> consts, int64_t dimX,
4119 int64_t offset) const {
4120 ArrayRef<int64_t> globalStaticSizes = adaptor.getGlobalStaticSizes();
4121 ValueRange globalDynamicSizes = adaptor.getGlobalDynamicSizes();
4122 SmallVector<OpFoldResult> mixedGlobalSizes =
4123 getMixedValues(globalStaticSizes, globalDynamicSizes, rewriter);
4124 if (mixedGlobalSizes.size() <= static_cast<unsigned long>(dimX))
4125 return sgpr0;
4126
4127 OpFoldResult tensorDimXOpFoldResult = *(mixedGlobalSizes.rbegin() + dimX);
4128 Value tensorDimX;
4129 if (auto attr = dyn_cast<Attribute>(tensorDimXOpFoldResult)) {
4130 tensorDimX =
4131 createI32Constant(rewriter, loc, cast<IntegerAttr>(attr).getInt());
4132 } else {
4133 IntegerType i32 = rewriter.getI32Type();
4134 tensorDimX = cast<Value>(tensorDimXOpFoldResult);
4135 tensorDimX = LLVM::TruncOp::create(rewriter, loc, i32, tensorDimX);
4136 }
4137
4138 return setValueAtOffset(rewriter, loc, sgpr0, tensorDimX, offset);
4139 }
4140
4141 Value setTensorDim2(DescriptorOp op, OpAdaptor adaptor,
4142 ConversionPatternRewriter &rewriter, Location loc,
4143 Value sgpr0, ArrayRef<Value> consts) const {
4144 return setTensorDimX(op, adaptor, rewriter, loc, sgpr0, consts, 2, 0);
4145 }
4146
4147 Value truncateAndSetValueAtOffset(ConversionPatternRewriter &rewriter,
4148 Location loc, Value accumulator,
4149 Value value, int64_t shift) const {
4150
4151 IntegerType i32 = rewriter.getI32Type();
4152 value = LLVM::TruncOp::create(rewriter, loc, i32, value);
4153 return setValueAtOffset(rewriter, loc, accumulator, value, shift);
4154 }
4155
4156 Value setLDSAddrIncrement(DescriptorOp op, OpAdaptor adaptor,
4157 ConversionPatternRewriter &rewriter, Location loc,
4158 Value sgpr1, ArrayRef<Value> consts,
4159 int64_t offset) const {
4160 Value ldsAddrIncrement = adaptor.getLdsIncrement();
4161 return setValueAtOffset(rewriter, loc, sgpr1, ldsAddrIncrement, offset);
4162 }
4163
4164 std::pair<Value, Value>
4165 setGlobalAddrIncrement(DescriptorOp op, OpAdaptor adaptor,
4166 ConversionPatternRewriter &rewriter, Location loc,
4167 Value sgpr2, Value sgpr3, ArrayRef<Value> consts,
4168 int64_t offset) const {
4169 Value globalAddrIncrement = adaptor.getGlobalIncrement();
4170 sgpr2 = truncateAndSetValueAtOffset(rewriter, loc, sgpr2,
4171 globalAddrIncrement, offset);
4172 Value shift = createI64Constant(rewriter, loc, 32);
4173 globalAddrIncrement =
4174 LLVM::LShrOp::create(rewriter, loc, globalAddrIncrement, shift);
4175 constexpr int64_t first16BitsHigh = (1ll << 16) - 1;
4176 sgpr3 = truncateAndSetValueAtOffset(rewriter, loc, sgpr3,
4177 globalAddrIncrement, offset + 32);
4178 Value mask = createI32Constant(rewriter, loc, first16BitsHigh);
4179 sgpr3 = LLVM::AndOp::create(rewriter, loc, sgpr3, mask);
4180 return {sgpr2, sgpr3};
4181 }
4182
4183 Value setTensorDim3OrLDSAddrIncrement(DescriptorOp op, OpAdaptor adaptor,
4184 ConversionPatternRewriter &rewriter,
4185 Location loc, Value sgpr1,
4186 ArrayRef<Value> consts) const {
4187 Value ldsIncrement = op.getLdsIncrement();
4188 constexpr int64_t dim = 3;
4189 constexpr int64_t offset = 32;
4190 if (!ldsIncrement)
4191 return setTensorDimX(op, adaptor, rewriter, loc, sgpr1, consts, dim,
4192 offset);
4193 return setLDSAddrIncrement(op, adaptor, rewriter, loc, sgpr1, consts,
4194 offset);
4195 }
4196
4197 std::pair<Value, Value> setTensorDim2StrideOrGlobalAddrIncrement(
4198 DescriptorOp op, OpAdaptor adaptor, ConversionPatternRewriter &rewriter,
4199 Location loc, Value sgpr2, Value sgpr3, ArrayRef<Value> consts) const {
4200 Value globalIncrement = op.getGlobalIncrement();
4201 constexpr int32_t dim = 2;
4202 constexpr int32_t offset = 64;
4203 if (!globalIncrement)
4204 return setTensorDimXStride(op, adaptor, rewriter, loc, sgpr2, sgpr3,
4205 consts, dim, offset);
4206 return setGlobalAddrIncrement(op, adaptor, rewriter, loc, sgpr2, sgpr3,
4207 consts, offset);
4208 }
4209
4210 Value setIterateCount(DescriptorOp op, OpAdaptor adaptor,
4211 ConversionPatternRewriter &rewriter, Location loc,
4212 Value sgpr3, ArrayRef<Value> consts,
4213 int32_t offset) const {
4214 Value iterationCount = adaptor.getIterationCount();
4215 IntegerType i32 = rewriter.getI32Type();
4216 // pre-condition: iterationCount is in the inclusive interval [1, 256].
4217 // TODO: validation if the value breaks the pre-condition.
4218 // If the pre-condition fails, there is a possibility of
4219 // affecting the higher bits. In a following PR implement
4220 // RuntimeVerifiableOpInterface that instruments conditions that need to be
4221 // checked at runtime.
4222 iterationCount = LLVM::TruncOp::create(rewriter, loc, i32, iterationCount);
4223 iterationCount =
4224 LLVM::SubOp::create(rewriter, loc, iterationCount, consts[1]);
4225 return setValueAtOffset(rewriter, loc, sgpr3, iterationCount, offset);
4226 }
4227
4228 Value setTileDim3OrIterateCount(DescriptorOp op, OpAdaptor adaptor,
4229 ConversionPatternRewriter &rewriter,
4230 Location loc, Value sgpr3,
4231 ArrayRef<Value> consts) const {
4232 Value iterateCount = op.getIterationCount();
4233 constexpr int32_t dim = 2;
4234 constexpr int32_t offset = 112;
4235 if (!iterateCount)
4236 return setTileDimX(op, adaptor, rewriter, loc, sgpr3, consts, dim,
4237 offset);
4238
4239 return setIterateCount(op, adaptor, rewriter, loc, sgpr3, consts, offset);
4240 }
4241
4242 Value getDGroup2(DescriptorOp op, OpAdaptor adaptor,
4243 ConversionPatternRewriter &rewriter, Location loc,
4244 ArrayRef<Value> consts) const {
4245 if constexpr (DescriptorOp::isGather())
4246 return getDGroup2Gather(op, adaptor, rewriter, loc, consts);
4247 return getDGroup2NonGather(op, adaptor, rewriter, loc, consts);
4248 }
4249
4250 Value getDGroup2NonGather(DescriptorOp op, OpAdaptor adaptor,
4251 ConversionPatternRewriter &rewriter, Location loc,
4252 ArrayRef<Value> consts) const {
4253 IntegerType i32 = rewriter.getI32Type();
4254 Type v4i32 = this->typeConverter->convertType(VectorType::get(4, i32));
4255 assert(v4i32 && "expected type conversion to succeed.");
4256
4257 bool onlyNeedsTwoDescriptors = !op.getLdsIncrement() && op.getRank() <= 2;
4258 if (onlyNeedsTwoDescriptors)
4259 return LLVM::ZeroOp::create(rewriter, loc, v4i32);
4260
4261 constexpr int64_t sgprlen = 4;
4262 Value sgprs[sgprlen];
4263 for (int i = 0; i < sgprlen; ++i)
4264 sgprs[i] = consts[0];
4265
4266 sgprs[0] = setTensorDim2(op, adaptor, rewriter, loc, sgprs[0], consts);
4267 sgprs[1] = setTensorDim3OrLDSAddrIncrement(op, adaptor, rewriter, loc,
4268 sgprs[1], consts);
4269 std::tie(sgprs[2], sgprs[3]) = setTensorDim2StrideOrGlobalAddrIncrement(
4270 op, adaptor, rewriter, loc, sgprs[2], sgprs[3], consts);
4271 sgprs[3] =
4272 setTileDim3OrIterateCount(op, adaptor, rewriter, loc, sgprs[3], consts);
4273
4274 Value dgroup2 = LLVM::PoisonOp::create(rewriter, loc, v4i32);
4275 for (auto [sgpr, constant] : llvm::zip(sgprs, consts))
4276 dgroup2 =
4277 LLVM::InsertElementOp::create(rewriter, loc, dgroup2, sgpr, constant);
4278
4279 return dgroup2;
4280 }
4281
4282 Value getGatherIndices(DescriptorOp op, OpAdaptor adaptor,
4283 ConversionPatternRewriter &rewriter, Location loc,
4284 ArrayRef<Value> consts, bool firstHalf) const {
4285 IntegerType i32 = rewriter.getI32Type();
4286 Type v4i32 = this->typeConverter->convertType(VectorType::get(4, i32));
4287 assert(v4i32 && "expected type conversion to succeed.");
4288
4289 Value indices = adaptor.getIndices();
4290 auto vectorType = cast<VectorType>(indices.getType());
4291 unsigned length = vectorType.getShape().back();
4292 Type elementType = vectorType.getElementType();
4293 unsigned maxLength = elementType == i32 ? 4 : 8;
4294 int32_t offset = firstHalf ? 0 : maxLength;
4295 unsigned discountedLength =
4296 std::max(static_cast<int32_t>(length - offset), 0);
4297
4298 unsigned targetSize = std::min(maxLength, discountedLength);
4299
4300 SmallVector<Value> indicesVector;
4301 for (unsigned i = offset; i < targetSize + offset; ++i) {
4302 Value idx;
4303 if (i < consts.size())
4304 idx = consts[i];
4305 else
4306 idx = createI32Constant(rewriter, loc, i);
4307 Value elem = LLVM::ExtractElementOp::create(rewriter, loc, indices, idx);
4308 indicesVector.push_back(elem);
4309 }
4310
4311 SmallVector<Value> indicesI32Vector;
4312 if (elementType == i32) {
4313 indicesI32Vector = indicesVector;
4314 } else {
4315 for (unsigned i = 0; i < targetSize; ++i) {
4316 Value index = indicesVector[i];
4317 indicesI32Vector.push_back(
4318 LLVM::ZExtOp::create(rewriter, loc, i32, index));
4319 }
4320 if ((targetSize % 2) != 0)
4321 // Add padding when not divisible by two.
4322 indicesI32Vector.push_back(consts[0]);
4323 }
4324
4325 SmallVector<Value> indicesToInsert;
4326 if (elementType == i32) {
4327 indicesToInsert = indicesI32Vector;
4328 } else {
4329 unsigned size = indicesI32Vector.size() / 2;
4330 for (unsigned i = 0; i < size; ++i) {
4331 Value first = indicesI32Vector[2 * i];
4332 Value second = indicesI32Vector[2 * i + 1];
4333 Value joined = setValueAtOffset(rewriter, loc, first, second, 16);
4334 indicesToInsert.push_back(joined);
4335 }
4336 }
4337
4338 Value dgroup = LLVM::PoisonOp::create(rewriter, loc, v4i32);
4339 for (auto [sgpr, constant] : llvm::zip_first(indicesToInsert, consts))
4340 dgroup =
4341 LLVM::InsertElementOp::create(rewriter, loc, dgroup, sgpr, constant);
4342
4343 return dgroup;
4344 }
4345
4346 Value getDGroup2Gather(DescriptorOp op, OpAdaptor adaptor,
4347 ConversionPatternRewriter &rewriter, Location loc,
4348 ArrayRef<Value> consts) const {
4349 return getGatherIndices(op, adaptor, rewriter, loc, consts, true);
4350 }
4351
4352 std::pair<Value, Value>
4353 setTensorDim3Stride(DescriptorOp op, OpAdaptor adaptor,
4354 ConversionPatternRewriter &rewriter, Location loc,
4355 Value sgpr0, Value sgpr1, ArrayRef<Value> consts) const {
4356 constexpr int32_t dim = 3;
4357 constexpr int32_t offset = 0;
4358 return setTensorDimXStride(op, adaptor, rewriter, loc, sgpr0, sgpr1, consts,
4359 dim, offset);
4360 }
4361
4362 std::pair<Value, Value> setTensorDim4(DescriptorOp op, OpAdaptor adaptor,
4363 ConversionPatternRewriter &rewriter,
4364 Location loc, Value sgpr1, Value sgpr2,
4365 ArrayRef<Value> consts) const {
4366 constexpr int32_t dim = 4;
4367 constexpr int32_t offset = 48;
4368 return setTensorDimX(op, adaptor, rewriter, loc, sgpr1, sgpr2, consts, dim,
4369 offset);
4370 }
4371
4372 Value setTileDim4(DescriptorOp op, OpAdaptor adaptor,
4373 ConversionPatternRewriter &rewriter, Location loc,
4374 Value sgpr2, ArrayRef<Value> consts) const {
4375 constexpr int32_t dim = 4;
4376 constexpr int32_t offset = 80;
4377 return setTileDimX(op, adaptor, rewriter, loc, sgpr2, consts, dim, offset);
4378 }
4379
4380 Value getDGroup3(DescriptorOp op, OpAdaptor adaptor,
4381 ConversionPatternRewriter &rewriter, Location loc,
4382 ArrayRef<Value> consts) const {
4383 if constexpr (DescriptorOp::isGather())
4384 return getDGroup3Gather(op, adaptor, rewriter, loc, consts);
4385 return getDGroup3NonGather(op, adaptor, rewriter, loc, consts);
4386 }
4387
4388 Value getDGroup3NonGather(DescriptorOp op, OpAdaptor adaptor,
4389 ConversionPatternRewriter &rewriter, Location loc,
4390 ArrayRef<Value> consts) const {
4391 IntegerType i32 = rewriter.getI32Type();
4392 Type v4i32 = this->typeConverter->convertType(VectorType::get(4, i32));
4393 assert(v4i32 && "expected type conversion to succeed.");
4394 bool onlyNeedsTwoDescriptors = !op.getLdsIncrement() && op.getRank() <= 2;
4395 if (onlyNeedsTwoDescriptors)
4396 return LLVM::ZeroOp::create(rewriter, loc, v4i32);
4397
4398 constexpr int32_t sgprlen = 4;
4399 Value sgprs[sgprlen];
4400 for (int i = 0; i < sgprlen; ++i)
4401 sgprs[i] = consts[0];
4402
4403 std::tie(sgprs[0], sgprs[1]) = setTensorDim3Stride(
4404 op, adaptor, rewriter, loc, sgprs[0], sgprs[1], consts);
4405 std::tie(sgprs[1], sgprs[2]) =
4406 setTensorDim4(op, adaptor, rewriter, loc, sgprs[1], sgprs[2], consts);
4407 sgprs[2] = setTileDim4(op, adaptor, rewriter, loc, sgprs[2], consts);
4408
4409 Value dgroup3 = LLVM::PoisonOp::create(rewriter, loc, v4i32);
4410 for (auto [sgpr, constant] : llvm::zip(sgprs, consts))
4411 dgroup3 =
4412 LLVM::InsertElementOp::create(rewriter, loc, dgroup3, sgpr, constant);
4413
4414 return dgroup3;
4415 }
4416
4417 Value getDGroup3Gather(DescriptorOp op, OpAdaptor adaptor,
4418 ConversionPatternRewriter &rewriter, Location loc,
4419 ArrayRef<Value> consts) const {
4420 return getGatherIndices(op, adaptor, rewriter, loc, consts, false);
4421 }
4422
4423 LogicalResult
4424 matchAndRewrite(DescriptorOp op, OpAdaptor adaptor,
4425 ConversionPatternRewriter &rewriter) const override {
4426 if (chipset < kGfx1250)
4427 return op->emitOpError(
4428 "make_dma_descriptor is only supported on gfx1250");
4429
4430 Location loc = op.getLoc();
4431
4432 SmallVector<Value> consts;
4433 for (int64_t i = 0; i < 8; ++i)
4434 consts.push_back(createI32Constant(rewriter, loc, i));
4435
4436 Value dgroup0 = this->getDGroup0(adaptor);
4437 Value dgroup1 = this->getDGroup1(op, adaptor, rewriter, loc, consts);
4438 Value dgroup2 = this->getDGroup2(op, adaptor, rewriter, loc, consts);
4439 Value dgroup3 = this->getDGroup3(op, adaptor, rewriter, loc, consts);
4440 SmallVector<Value> results = {dgroup0, dgroup1, dgroup2, dgroup3};
4441 rewriter.replaceOpWithMultiple(op, {results});
4442 return success();
4443 }
4444};
4445
4446template <typename SourceOp, typename TargetOp>
4447struct AMDGPUTensorLoadStoreOpLowering
4448 : public ConvertOpToLLVMPattern<SourceOp> {
4449 using ConvertOpToLLVMPattern<SourceOp>::ConvertOpToLLVMPattern;
4451 AMDGPUTensorLoadStoreOpLowering(const LLVMTypeConverter &converter,
4452 Chipset chipset)
4453 : ConvertOpToLLVMPattern<SourceOp>(converter), chipset(chipset) {}
4454 Chipset chipset;
4455
4456 LogicalResult
4457 matchAndRewrite(SourceOp op, Adaptor adaptor,
4458 ConversionPatternRewriter &rewriter) const override {
4459 if (chipset < kGfx1250)
4460 return op->emitOpError("is only supported on gfx1250");
4461
4462 ValueRange desc = adaptor.getDesc();
4463 // Create a <v8 x i32> 0 as the fifth argument to match llvm intrinsic. It
4464 // will move into the TDM descriptor once it becomes relevant for future use
4465 auto v8i32 = VectorType::get(8, rewriter.getI32Type());
4466 Value dgroup4 = LLVM::ZeroOp::create(rewriter, op.getLoc(), v8i32);
4467 Attribute cachePolicy = rewriter.getI32IntegerAttr(0);
4468 rewriter.replaceOpWithNewOp<TargetOp>(op, desc[0], desc[1], desc[2],
4469 desc[3], dgroup4, cachePolicy,
4470 /*alias_scopes=*/nullptr,
4471 /*noalias_scopes=*/nullptr,
4472 /*tbaa=*/nullptr);
4473 return success();
4474 }
4475};
4476
4477struct GlobalPrefetchOpLowering
4478 : public ConvertOpToLLVMPattern<GlobalPrefetchOp> {
4479 GlobalPrefetchOpLowering(const LLVMTypeConverter &converter, Chipset chipset)
4480 : ConvertOpToLLVMPattern<GlobalPrefetchOp>(converter), chipset(chipset) {}
4481
4482 LogicalResult
4483 matchAndRewrite(GlobalPrefetchOp op, GlobalPrefetchOpAdaptor adaptor,
4484 ConversionPatternRewriter &rewriter) const override {
4485 if (chipset < kGfx1250)
4486 return op->emitOpError("is only supported on gfx1250+");
4487
4488 const bool isSpeculative = op.getSpeculative();
4489 const int32_t immArgValue = getGlobalPrefetchLLVMEncoding(
4490 op.getTemporalHint(), op.getCacheScope(), isSpeculative);
4491 // amdgpu.global_prefetch is gfx1250+, so its policy bits use gfx12
4492 // encoding.
4493 Attribute cachePolicy = ROCDL::Gfx12CachePolicyAttr::get(
4494 rewriter.getContext(),
4495 static_cast<ROCDL::Gfx12CachePolicy>(immArgValue));
4496
4497 ValueRange indices = adaptor.getIndices();
4498 Value memRef = adaptor.getSrc();
4499 MemRefDescriptor descriptor(memRef);
4500 MemRefType memRefType = op.getSrc().getType();
4501 Location loc = op->getLoc();
4502 auto inboundsFlags = isSpeculative ? LLVM::GEPNoWrapFlags::none
4503 : LLVM::GEPNoWrapFlags::inbounds |
4504 LLVM::GEPNoWrapFlags::nuw;
4505 Value prefetchPtr = getStridedElementPtr(
4506 rewriter, loc, memRefType, descriptor, indices, inboundsFlags);
4507
4508 rewriter.replaceOpWithNewOp<ROCDL::GlobalPrefetchOp>(
4509 op, prefetchPtr, cachePolicy, mlir::ArrayAttr{}, mlir::ArrayAttr{},
4510 mlir::ArrayAttr{});
4511 return success();
4512 }
4513
4514private:
4515 Chipset chipset;
4516};
4517
4518struct ConvertAMDGPUToROCDLPass
4519 : public impl::ConvertAMDGPUToROCDLPassBase<ConvertAMDGPUToROCDLPass> {
4520 using Base::Base;
4521
4522 void runOnOperation() override {
4523 MLIRContext *ctx = &getContext();
4524 FailureOr<Chipset> maybeChipset = Chipset::parse(chipset);
4525 if (failed(maybeChipset)) {
4526 emitError(UnknownLoc::get(ctx), "Invalid chipset name: " + chipset);
4527 return signalPassFailure();
4528 }
4529
4530 RewritePatternSet patterns(ctx);
4531 LLVMTypeConverter converter(ctx);
4532
4533 populateAMDGPUToROCDLConversionPatterns(converter, patterns, *maybeChipset);
4535 LLVMConversionTarget target(getContext());
4536 target.addIllegalDialect<::mlir::amdgpu::AMDGPUDialect>();
4537 target.addLegalDialect<::mlir::LLVM::LLVMDialect>();
4538 target.addLegalDialect<::mlir::ROCDL::ROCDLDialect>();
4539 if (failed(applyPartialConversion(getOperation(), target,
4540 std::move(patterns))))
4541 signalPassFailure();
4542 }
4543};
4544} // namespace
4545
4547 TypeConverter &typeConverter) {
4549 typeConverter, [](gpu::AddressSpace space) {
4550 switch (space) {
4551 case gpu::AddressSpace::Global:
4552 return ROCDL::ROCDLDialect::kGlobalMemoryAddressSpace;
4553 case gpu::AddressSpace::Workgroup:
4554 return ROCDL::ROCDLDialect::kSharedMemoryAddressSpace;
4555 case gpu::AddressSpace::Private:
4556 return ROCDL::ROCDLDialect::kPrivateMemoryAddressSpace;
4557 case gpu::AddressSpace::Constant:
4558 return ROCDL::ROCDLDialect::kConstantMemoryAddressSpace;
4559 }
4560 llvm_unreachable("unknown address space enum value");
4561 });
4562 typeConverter.addConversion([](gpu::NamedBarrierType type) {
4563 return LLVM::LLVMPointerType::get(
4564 type.getContext(), ROCDL::ROCDLDialect::kSharedMemoryAddressSpace);
4565 });
4566}
4567
4569 TypeConverter &typeConverter) {
4570 typeConverter.addTypeAttributeConversion(
4571 [](BaseMemRefType type, amdgpu::AddressSpaceAttr as)
4572 -> TypeConverter::AttributeConversionResult {
4573 MLIRContext *ctx = as.getContext();
4574 Type i64 = IntegerType::get(ctx, 64);
4575 switch (as.getValue()) {
4576 case amdgpu::AddressSpace::FatRawBuffer:
4577 return IntegerAttr::get(i64, 7);
4578 case amdgpu::AddressSpace::BufferRsrc:
4579 return IntegerAttr::get(i64, 8);
4580 case amdgpu::AddressSpace::FatStructuredBuffer:
4581 return IntegerAttr::get(i64, 9);
4582 }
4583 return TypeConverter::AttributeConversionResult::abort();
4584 });
4585 typeConverter.addConversion([&](DsBarrierStateType type) -> Type {
4586 return IntegerType::get(type.getContext(), 64);
4587 });
4588 typeConverter.addConversion([&](TDMBaseType type) -> Type {
4589 Type i32 = IntegerType::get(type.getContext(), 32);
4590 return typeConverter.convertType(VectorType::get(4, i32));
4591 });
4592 typeConverter.addConversion([&](TDMGatherBaseType type) -> Type {
4593 Type i32 = IntegerType::get(type.getContext(), 32);
4594 return typeConverter.convertType(VectorType::get(4, i32));
4595 });
4596 typeConverter.addConversion(
4597 [&](TDMDescriptorType type,
4598 SmallVectorImpl<Type> &result) -> std::optional<LogicalResult> {
4599 Type i32 = IntegerType::get(type.getContext(), 32);
4600 Type v4i32 = typeConverter.convertType(VectorType::get(4, i32));
4601 Type v8i32 = typeConverter.convertType(VectorType::get(8, i32));
4602 llvm::append_values(result, v4i32, v8i32, v4i32, v4i32);
4603 return success();
4604 });
4605
4606 auto addUnrealizedCast = [](OpBuilder &builder, TypeRange types,
4607 ValueRange inputs,
4609 // Only create unrealized_conversion_cast for TDMDescriptorType.
4610 // All other types which are not expected, should be
4611 // materialized by other target materialization functions.
4612 if (inputs.size() != 1)
4613 return {};
4614
4615 if (!isa<TDMDescriptorType>(inputs[0].getType()))
4616 return {};
4617
4618 auto cast = UnrealizedConversionCastOp::create(builder, loc, types, inputs);
4619 return cast.getResults();
4620 };
4621
4622 typeConverter.addTargetMaterialization(addUnrealizedCast);
4623}
4624
4626 RewritePatternSet &patterns,
4627 Chipset chipset) {
4629 patterns
4630 .add<FatRawBufferCastLowering,
4631 RawBufferOpLowering<RawBufferLoadOp, ROCDL::RawPtrBufferLoadOp>,
4632 RawBufferOpLowering<RawBufferStoreOp, ROCDL::RawPtrBufferStoreOp>,
4633 RawBufferOpLowering<RawBufferAtomicFaddOp,
4634 ROCDL::RawPtrBufferAtomicFaddOp>,
4635 RawBufferOpLowering<RawBufferAtomicFmaxOp,
4636 ROCDL::RawPtrBufferAtomicFmaxOp>,
4637 RawBufferOpLowering<RawBufferAtomicSmaxOp,
4638 ROCDL::RawPtrBufferAtomicSmaxOp>,
4639 RawBufferOpLowering<RawBufferAtomicUminOp,
4640 ROCDL::RawPtrBufferAtomicUminOp>,
4641 RawBufferOpLowering<RawBufferAtomicCmpswapOp,
4642 ROCDL::RawPtrBufferAtomicCmpSwap>,
4643 AMDGPUDPPLowering, MemoryCounterWaitOpLowering, LDSBarrierOpLowering,
4644 SchedBarrierOpLowering, MFMAOpLowering, ScaledMFMAOpLowering,
4645 SparseMFMAOpLowering, WMMAOpLowering, ScaledWMMAOpLowering,
4646 SparseWMMAOpLowering, DotOpLowering, ExtPackedFp8OpLowering,
4647 ScaledExtPackedMatrixOpLowering, ScaledExtPackedOpLowering,
4648 PackedScaledTruncOpLowering, PackedTrunc2xFp8OpLowering,
4649 PackedStochRoundFp8OpLowering, GatherToLDSOpLowering,
4650 GlobalLoadAsyncToLDSOpLowering, TransposeLoadOpLowering,
4651 GlobalTransposeLoadOpLowering, AMDGPUPermlaneLowering,
4652 AMDGPUPermlaneVarLowering, AMDGPUMakeDmaBaseLowering<MakeDmaBaseOp>,
4653 AMDGPUMakeDmaBaseLowering<MakeGatherDmaBaseOp>,
4654 AMDGPULowerDescriptor<MakeDmaDescriptorOp>,
4655 AMDGPULowerDescriptor<MakeGatherDmaDescriptorOp>,
4656 AMDGPUTensorLoadStoreOpLowering<TensorLoadToLDSOp,
4657 ROCDL::TensorLoadToLDSOp>,
4658 AMDGPUTensorLoadStoreOpLowering<TensorStoreFromLDSOp,
4659 ROCDL::TensorStoreFromLDSOp>,
4660 DsBarrierInitOpLowering, DsBarrierPollStateOpLowering,
4661 DsAsyncBarrierArriveOpLowering, DsBarrierArriveOpLowering,
4662 GlobalPrefetchOpLowering>(converter, chipset);
4663 patterns.add<AMDGPUSwizzleBitModeLowering, DsBarrierStatePhaseOpLowering,
4664 DsBarrierStatePendingCountOpLowering,
4665 DsBarrierStateInitCountOpLowering,
4666 DsBarrierStatePhaseParityLowering>(converter);
4667}
static bool typeIsExpectedFp8ForChipset(Chipset chipset, Type type)
Return true if type is the E4M3FN variant of an 8-bit float that is supported by the _fp8 instruction...
constexpr Chipset kGfx942
static std::optional< StringRef > wmmaOpToIntrinsicRDNA(Type elemSourceType, Type elemBSourceType, Type elemDestType, uint32_t k, bool isRDNA3)
Returns the rocdl intrinsic corresponding to a WMMA operation wmma for RDNA3/4 architectures.
static bool hasDot10Insts(const Chipset &chipset)
static bool hasDot7Insts(const Chipset &chipset)
static std::optional< SparseWMMAOpInfo > sparseWMMAOpToIntrinsic(SparseWMMAOp swmmac, Chipset chipset)
static std::optional< StringRef > mfmaOpToIntrinsic(MFMAOp mfma, Chipset chipset)
Return the rocdl intrinsic corresponding to a MFMA operation mfma if one exists.
constexpr Chipset kGfx908
static void wmmaPushInputOperand(ConversionPatternRewriter &rewriter, Location loc, const TypeConverter *typeConverter, bool isUnsigned, Value llvmInput, Value mlirInput, SmallVectorImpl< Value > &operands, SmallVectorImpl< NamedAttribute > &attrs, StringRef attrName)
Push an input operand.
static std::optional< ScaledMFMAIntrinsic > mfmaOpToScaledIntrinsic(Type aType, Type bType, Type destType, uint32_t m, uint32_t n, uint32_t k, uint32_t b, Chipset chipset)
constexpr Chipset kGfx1250
static Value castScaleOperand(ConversionPatternRewriter &rewriter, Location loc, Value input)
Converts the scaled MFMA/WMMA operands, scalesA and scalesB, from MLIR AMDGPU dialect convention to R...
constexpr Chipset kGfx90a
static std::optional< StringRef > getScaledWmmaIntrinsicName(int64_t m, int64_t n, int64_t k, bool isScale16)
Determines the ROCDL intrinsic name for scaled WMMA based on dimensions and scale block size (16 or 3...
static void wmmaPushOutputOperand(ConversionPatternRewriter &rewriter, Location loc, const TypeConverter *typeConverter, Value output, int32_t subwordOffset, bool clamp, SmallVectorImpl< Value > &operands, SmallVectorImpl< NamedAttribute > &attrs)
Push the output operand.
static bool typeIsExpectedBf8ForChipset(Chipset chipset, Type type)
Return true if type is the E5M2 variant of an 8-bit float that is supported by the _bf8 instructions ...
static std::optional< StringRef > wmmaOpToIntrinsic(WMMAOp wmma, Chipset chipset)
Returns the rocdl intrinsic corresponding to a WMMA operation wmma if one exists.
static bool hasDot11Insts(const Chipset &chipset)
static std::optional< StringRef > smfmacOpToIntrinsic(SparseMFMAOp op, Chipset chipset)
Returns the rocdl intrinsic corresponding to a SparseMFMA (smfmac) operation if one exists.
static Value makeBufferRsrc(ConversionPatternRewriter &rewriter, Location loc, Value basePointer, Value numRecords, bool boundsCheck, amdgpu::Chipset chipset, Value cacheSwizzleStride=nullptr, unsigned addressSpace=8)
static Value createI64Constant(ConversionPatternRewriter &rewriter, Location loc, int64_t value)
static bool hasDot9Insts(const Chipset &chipset)
static std::optional< StringRef > wmmaOpToIntrinsicGfx1250(Type elemSourceType, Type elemBSourceType, Type elemDestType, uint32_t k)
Return the rocdl intrinsic corresponding to a WMMA operation wmma for the gfx1250 architecture.
constexpr Chipset kGfx1200
static Value getNumRecords(ConversionPatternRewriter &rewriter, Location loc, MemRefType memrefType, MemRefDescriptor &memrefDescriptor, ArrayRef< int64_t > strides, int64_t elementByteWidth, amdgpu::Chipset chipset, bool boundsCheck)
Compute the contents of the num_records field for a given memref descriptor - that is,...
static Value packSmallFloatVectorOperand(ConversionPatternRewriter &rewriter, Location loc, Value input, bool allowBf16=true)
Pack small float vector operands (fp4/fp6/fp8/bf16) into the format expected by scaled matrix multipl...
static bool has45BitNumRecordsBufferResource(const Chipset &chipset)
static std::optional< ROCDL::WMMAMatrixScaleFormat > getWmmaScaleFormat(Type elemType)
Maps f8 scale element types to WMMA scale format codes.
static Value convertPackedVectorOperand(ConversionPatternRewriter &rewriter, Location loc, Value input, bool allowBf16=true)
Converts packed vector operands to the expected ROCDL types.
static Value getLinearIndexI32(ConversionPatternRewriter &rewriter, Location loc, MemRefDescriptor &memRefDescriptor, ValueRange indices, ArrayRef< int64_t > strides)
Returns the linear index used to access an element in the memref.
static Value convertUnsignedToI32(ConversionPatternRewriter &rewriter, Location loc, Value val)
Convert an unsigned number val to i32.
static bool hasDot8Insts(const Chipset &chipset)
static bool hasDot2Insts(const Chipset &chipset)
static Value createI32Constant(ConversionPatternRewriter &rewriter, Location loc, int32_t value)
static std::optional< ROCDL::MatrixFormat > smallFloatTypeToMatrixFormat(Type mlirElemType)
std::tuple< StringRef, ROCDL::MatrixFormat, ROCDL::MatrixFormat > ScaledMFMAIntrinsic
If there is a scaled MFMA instruction for the input element types aType and bType,...
static bool hasDot12Insts(const Chipset &chipset)
static Value convertUnsignedToI64(ConversionPatternRewriter &rewriter, Location loc, Value val)
Convert an unsigned number val to i64.
constexpr Chipset kGfx950
static bool hasDot1Insts(const Chipset &chipset)
return success()
b
Return true if permutation is a valid permutation of the outer_dims_perm (case OuterOrInnerPerm::Oute...
ArrayAttr()
b getContext())
auto load
*if copies could not be generated due to yet unimplemented cases *copyInPlacementStart and copyOutPlacementStart in copyPlacementBlock *specify the insertion points where the incoming copies and outgoing should be the output argument nBegin is set to its * replacement(set to `begin` if no invalidation happens). Since outgoing *copies could have been inserted at `end`
static constexpr unsigned kSizePosInMemRefDescriptor
static constexpr unsigned kStridePosInMemRefDescriptor
static constexpr unsigned kOffsetPosInMemRefDescriptor
static constexpr unsigned kAllocatedPtrPosInMemRefDescriptor
static constexpr unsigned kAlignedPtrPosInMemRefDescriptor
static Value clamp(ImplicitLocOpBuilder &builder, Value value, Value lowerBound, Value upperBound)
Attributes are known-constant values of operations.
Definition Attributes.h:25
This class provides a shared interface for ranked and unranked memref types.
Utility class for operation conversions targeting the LLVM dialect that match exactly one source oper...
Definition Pattern.h:227
ConvertOpToLLVMPattern(const LLVMTypeConverter &typeConverter, PatternBenefit benefit=1)
Definition Pattern.h:233
typename SourceOp::template GenericAdaptor< ArrayRef< ValueRange > > OneToNOpAdaptor
Definition Pattern.h:230
typename SourceOp::Adaptor OpAdaptor
Definition Pattern.h:229
Value getStridedElementPtr(ConversionPatternRewriter &rewriter, Location loc, MemRefType type, Value memRefDesc, ValueRange indices, LLVM::GEPNoWrapFlags noWrapFlags=LLVM::GEPNoWrapFlags::none) const
Convenience wrapper for the corresponding helper utility.
Definition Pattern.cpp:66
llvm::TypeSize getTypeSizeInBits(Type t) const
Returns the size in bits of the given type in the current scope.
Conversion from types to the LLVM IR dialect.
This class defines the main interface for locations in MLIR and acts as a non-nullable wrapper around...
Definition Location.h:76
MLIRContext is the top-level object for a collection of MLIR operations.
Definition MLIRContext.h:63
Helper class to produce LLVM dialect operations extracting or inserting elements of a MemRef descript...
Value stride(OpBuilder &builder, Location loc, unsigned pos)
Builds IR extracting the pos-th size from the descriptor.
Value size(OpBuilder &builder, Location loc, unsigned pos)
Builds IR extracting the pos-th size from the descriptor.
NamedAttribute represents a combination of a name and an Attribute value.
Definition Attributes.h:164
This class helps build Operations.
Definition Builders.h:210
OpResult getResult(unsigned idx)
Get the 'idx'th result of this operation.
Definition Operation.h:432
result_range getResults()
Definition Operation.h:440
unsigned getNumResults()
Return the number of results held by this operation.
Definition Operation.h:429
RewritePatternSet & add(ConstructorArg &&arg, ConstructorArgs &&...args)
Add an instance of each of the pattern types 'Ts' to the pattern list with the given arguments.
This class provides an abstraction over the various different ranges of value types.
Definition TypeRange.h:40
Instances of the Type class are uniqued, have an immutable identifier and an optional mutable compone...
Definition Types.h:74
bool isF64() const
Definition Types.cpp:41
MLIRContext * getContext() const
Return the MLIRContext in which this type was uniqued.
Definition Types.cpp:35
bool isF8E5M2() const
Definition Types.cpp:45
bool isSignedInteger() const
Return true if this is a signed integer type (with the specified width).
Definition Types.cpp:78
bool isF8E4M3FN() const
Definition Types.cpp:44
bool isF32() const
Definition Types.cpp:40
bool isUnsignedInteger() const
Return true if this is an unsigned integer type (with the specified width).
Definition Types.cpp:90
bool isInteger() const
Return true if this is an integer type (with the specified width).
Definition Types.cpp:58
bool isF16() const
Definition Types.cpp:38
unsigned getIntOrFloatBitWidth() const
Return the bit width of an integer or a float type, assert failure on other types.
Definition Types.cpp:124
bool isBF16() const
Definition Types.cpp:37
This class provides an abstraction over the different types of ranges over Values.
Definition ValueRange.h:389
This class represents an instance of an SSA value in the MLIR system, representing a computable value...
Definition Value.h:96
Type getType() const
Return the type of this value.
Definition Value.h:105
::mlir::Pass::Option< std::string > chipset
Value getStridedElementPtr(OpBuilder &builder, Location loc, const LLVMTypeConverter &converter, MemRefType type, Value memRefDesc, ValueRange indices, LLVM::GEPNoWrapFlags noWrapFlags=LLVM::GEPNoWrapFlags::none)
Performs the index computation to get to the element at indices of the memory pointed to by memRefDes...
Definition Pattern.cpp:603
LogicalResult decomposeValue(OpBuilder &builder, Location loc, Value src, Type dstType, SmallVectorImpl< Value > &result, bool permitVariablySizedScalars=false)
Decomposes a src value into a set of values of type dstType through series of bitcasts and vector ops...
Definition Pattern.cpp:495
Value composeValue(OpBuilder &builder, Location loc, ValueRange src, Type dstType)
Composes a set of src values into a single value of type dstType through series of bitcasts and vecto...
Definition Pattern.cpp:594
int32_t getGlobalPrefetchLLVMEncoding(amdgpu::LoadTemporalHint hint, amdgpu::Scope scope, bool isSpeculative)
Definition AMDGPUEnums.h:18
bool hasOcpFp8(const Chipset &chipset)
Definition Chipset.h:52
void populateCommonGPUTypeAndAttributeConversions(TypeConverter &typeConverter)
Remap common GPU memory spaces (Workgroup, Private, etc) to LLVM address spaces.
detail::InFlightRemark failed(Location loc, RemarkOpts opts)
Report an optimization remark that failed.
Definition Remarks.h:717
Include the generated interface declarations.
bool matchPattern(Value value, const Pattern &pattern)
Entry point for matching a pattern over a Value.
Definition Matchers.h:490
SmallVector< OpFoldResult > getMixedValues(ArrayRef< int64_t > staticValues, ValueRange dynamicValues, MLIRContext *context)
Return a vector of OpFoldResults with the same size a staticValues, but all elements for which Shaped...
Type getType(OpFoldResult ofr)
Returns the int type of the integer in ofr.
Definition Utils.cpp:307
InFlightDiagnostic emitError(Location loc)
Utility method to emit an error message using this location.
detail::constant_int_predicate_matcher m_Zero()
Matches a constant scalar / vector splat / tensor splat integer zero.
Definition Matchers.h:442
Type getElementTypeOrSelf(Type type)
Return the element type or return the type itself.
void populateGpuMemorySpaceAttributeConversions(TypeConverter &typeConverter, const MemorySpaceMapping &mapping)
Populates memory space attribute conversion rules for lowering gpu.address_space to integer values.
void populateAMDGPUToROCDLConversionPatterns(LLVMTypeConverter &converter, RewritePatternSet &patterns, amdgpu::Chipset chipset)
Note: This function will also add conversions for the AMDGPU-specific address spaces and types,...
llvm::TypeSwitch< T, ResultT > TypeSwitch
Definition LLVM.h:139
auto get(MLIRContext *context, Ts &&...params)
Helper method that injects context only if needed, this helps unify some of the attribute constructio...
void populateAMDGPUTypeAndAttributeConversions(TypeConverter &typeConverter)
Remap AMDGPU memory spaces to LLVM address spaces by mapping amdgpu::AddressSpace::fat_raw_buffer to ...
Returns the rocdl intrinsic corresponding to a SparseWMMA operation swmmac if one exists.
Represents the amdgpu gfx chipset version, e.g., gfx90a, gfx942, gfx1103.
Definition Chipset.h:22
unsigned majorVersion
Definition Chipset.h:23
unsigned minorVersion
Definition Chipset.h:24
static FailureOr< Chipset > parse(StringRef name)
Parses the chipset version string and returns the chipset on success, and failure otherwise.
Definition Chipset.cpp:14
unsigned steppingVersion
Definition Chipset.h:25