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