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