MLIR 24.0.0git
XeVMToLLVM.cpp
Go to the documentation of this file.
1//===-- XeVMToLLVM.cpp - XeVM to LLVM dialect conversion --------*- C++ -*-===//
2//
3// This file is licensed 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
16#include "mlir/Pass/Pass.h"
17#include "mlir/Support/LLVM.h"
18#include "llvm/ADT/ArrayRef.h"
19#include "llvm/ADT/STLExtras.h"
20#include "llvm/Support/FormatVariadic.h"
21#include "llvm/Support/MathExtras.h"
22
24#include "mlir/IR/Matchers.h"
25#include "mlir/IR/Types.h"
28
29#include "llvm/ADT/TypeSwitch.h"
30
31namespace mlir {
32#define GEN_PASS_DEF_CONVERTXEVMTOLLVMPASS
33#include "mlir/Conversion/Passes.h.inc"
34} // namespace mlir
35
36using namespace mlir;
37using namespace xevm;
38
39namespace {
40
41struct LLVMFuncAttributeOptions {
42 bool isConvergent = false;
43 bool isNoUnwind = false;
44 bool isWillReturn = false;
45 LLVM::MemoryEffectsAttr memEffectsAttr{};
46};
47static constexpr LLVMFuncAttributeOptions noUnwindAttrs = {
48 false, true, false, {}};
49static constexpr LLVMFuncAttributeOptions noUnwindWillReturnAttrs = {
50 false, true, true, {}};
51static constexpr LLVMFuncAttributeOptions convergentNoUnwindWillReturnAttrs = {
52 true, true, true, {}};
53
54std::string getTypeMangling(Type ty, bool isUnsigned = false) {
56 .Case([isUnsigned](VectorType ty) -> std::string {
57 return "Dv" + std::to_string(ty.getNumElements()) + "_" +
58 getTypeMangling(ty.getElementType(), isUnsigned);
59 })
60 .Case([](Float16Type) -> std::string { return "Dh"; })
61 .Case([](Float32Type) -> std::string { return "f"; })
62 .Case([](Float64Type) -> std::string { return "d"; })
63 .Case([isUnsigned](IntegerType ty) -> std::string {
64 switch (ty.getWidth()) {
65 case 8:
66 return isUnsigned ? "h" : "c";
67 case 16:
68 return isUnsigned ? "t" : "s";
69 case 32:
70 return isUnsigned ? "j" : "i";
71 case 64:
72 return isUnsigned ? "m" : "l";
73 default:
74 llvm_unreachable("unhandled integer type");
75 }
76 })
77 .DefaultUnreachable("unhandled type for mangling");
78}
79
80std::string mangle(StringRef baseName, ArrayRef<Type> types,
81 ArrayRef<bool> isUnsigned = {}) {
82 assert((isUnsigned.empty() || isUnsigned.size() == types.size()) &&
83 "Signedness info doesn't match");
84 std::string s;
85 llvm::raw_string_ostream os(s);
86 llvm::SmallDenseMap<Type, unsigned> substitutions;
87 os << "_Z" << baseName.size() << baseName;
88 for (auto [idx, type] : llvm::enumerate(types)) {
89 auto it = substitutions.find(type);
90 if (it != substitutions.end()) {
91 os << "S";
92 // First substitution is `S_`, second is `S0_`, and so on.
93 if (unsigned firstIdx = it->getSecond(); firstIdx > 0)
94 os << firstIdx - 1;
95 os << "_";
96 } else {
97 if (!type.isIntOrFloat())
98 substitutions[type] = substitutions.size();
99 os << getTypeMangling(type, isUnsigned.empty() ? false : isUnsigned[idx]);
100 }
101 }
102 return os.str();
103}
104
105// Returns the mangling of `ty` used to name an overloaded `llvm.genx.GenISA.*`
106// intrinsic: `i32`, `v8i16`, ... Note that this is IGC's own scheme for its
107// intrinsics, not the Itanium mangling used for the SPIR-V friendly and OCL
108// builtins that `mangle` above produces.
109std::string getGenISATypeMangling(Type ty) {
111 .Case([](VectorType ty) -> std::string {
112 return "v" + std::to_string(ty.getNumElements()) +
113 getGenISATypeMangling(ty.getElementType());
114 })
115 .Case([](IntegerType ty) -> std::string {
116 return "i" + std::to_string(ty.getWidth());
117 })
118 .DefaultUnreachable("unhandled type for GenISA mangling");
119}
120
121std::string builtinElemType(ElemType elemType) {
122 switch (elemType) {
123 case ElemType::BF8:
124 return "bf8";
125 case ElemType::F8:
126 return "hf8";
127 case ElemType::BF16:
128 return "bf";
129 case ElemType::F16:
130 return "hf";
131 case ElemType::F32:
132 return "f";
133 default:
134 return stringifyElemType(elemType).str();
135 }
136}
137
138static int32_t getL1CacheControl(LoadCacheControl cc) {
139 int32_t control = 0;
140 switch (cc) {
141 case LoadCacheControl::USE_DEFAULT:
142 control = -1;
143 break;
144 case LoadCacheControl::L1C_L2UC_L3UC:
145 case LoadCacheControl::L1C_L2UC_L3C:
146 case LoadCacheControl::L1C_L2C_L3UC:
147 case LoadCacheControl::L1C_L2C_L3C:
148 control = 1;
149 break;
150 case LoadCacheControl::L1S_L2UC_L3UC:
151 case LoadCacheControl::L1S_L2UC_L3C:
152 case LoadCacheControl::L1S_L2C_L3UC:
153 case LoadCacheControl::L1S_L2C_L3C:
154 control = 2;
155 break;
156 case LoadCacheControl::INVALIDATE_READ:
157 control = 3;
158 break;
159 default:
160 break;
161 }
162 return control;
163}
164
165static int32_t getL1CacheControl(StoreCacheControl cc) {
166 int32_t control = 0;
167 switch (cc) {
168 case StoreCacheControl::USE_DEFAULT:
169 control = -1;
170 break;
171 case StoreCacheControl::L1WT_L2UC_L3UC:
172 case StoreCacheControl::L1WT_L2UC_L3WB:
173 case StoreCacheControl::L1WT_L2WB_L3UC:
174 case StoreCacheControl::L1WT_L2WB_L3WB:
175 control = 1;
176 break;
177 case StoreCacheControl::L1WB_L2UC_L3UC:
178 case StoreCacheControl::L1WB_L2WB_L3UC:
179 case StoreCacheControl::L1WB_L2UC_L3WB:
180 control = 2;
181 break;
182 case StoreCacheControl::L1S_L2UC_L3UC:
183 case StoreCacheControl::L1S_L2UC_L3WB:
184 case StoreCacheControl::L1S_L2WB_L3UC:
185 case StoreCacheControl::L1S_L2WB_L3WB:
186 control = 3;
187 break;
188 default:
189 break;
190 }
191 return control;
192}
193
194static int32_t getL3CacheControl(LoadCacheControl cc) {
195 int32_t control = 0;
196 switch (cc) {
197 case LoadCacheControl::USE_DEFAULT:
198 control = -1;
199 break;
200 case LoadCacheControl::L1UC_L2UC_L3C:
201 case LoadCacheControl::L1UC_L2C_L3C:
202 case LoadCacheControl::L1C_L2UC_L3C:
203 case LoadCacheControl::L1C_L2C_L3C:
204 case LoadCacheControl::L1S_L2UC_L3C:
205 case LoadCacheControl::L1S_L2C_L3C:
206 control = 1;
207 break;
208 case LoadCacheControl::INVALIDATE_READ:
209 control = 3;
210 break;
211 default:
212 break;
213 }
214 return control;
215}
216
217static int32_t getL3CacheControl(StoreCacheControl cc) {
218 int32_t control = 0;
219 switch (cc) {
220 case StoreCacheControl::USE_DEFAULT:
221 control = -1;
222 break;
223 case StoreCacheControl::L1UC_L2UC_L3WB:
224 case StoreCacheControl::L1UC_L2WB_L3WB:
225 case StoreCacheControl::L1WT_L2UC_L3WB:
226 case StoreCacheControl::L1WT_L2WB_L3WB:
227 case StoreCacheControl::L1S_L2UC_L3WB:
228 case StoreCacheControl::L1S_L2WB_L3WB:
229 case StoreCacheControl::L1WB_L2UC_L3WB:
230 control = 2;
231 break;
232 default:
233 break;
234 }
235 return control;
236}
237
238static std::optional<LoadCacheControl> getCacheControl(PrefetchOp op) {
239 return op.getCacheControl();
240}
241
242static std::optional<LoadCacheControl> getCacheControl(BlockLoad2dOp op) {
243 return op.getCacheControl();
244}
245
246static std::optional<LoadCacheControl> getCacheControl(BlockLoadOp op) {
247 return op.getCacheControl();
248}
249
250static std::optional<LoadCacheControl> getCacheControl(BlockPrefetch2dOp op) {
251 return op.getCacheControl();
252}
253
254static std::optional<StoreCacheControl> getCacheControl(BlockStore2dOp op) {
255 return op.getCacheControl();
256}
257
258static std::optional<StoreCacheControl> getCacheControl(BlockStoreOp op) {
259 return op.getCacheControl();
260}
261
262static std::optional<LoadCacheControl> getCacheControl(LLVM::LoadOp op) {
263 if (op->hasDiscardableAttr("cache_control")) {
264 auto attr = op->getDiscardableAttrOfType<xevm::LoadCacheControlAttr>(
265 "cache_control");
266 if (!attr)
267 return std::nullopt;
268 return std::optional<LoadCacheControl>(attr.getValue());
269 }
270 return std::nullopt;
271}
272
273static std::optional<StoreCacheControl> getCacheControl(LLVM::StoreOp op) {
274 if (op->hasDiscardableAttr("cache_control")) {
275 auto attr = op->getDiscardableAttrOfType<xevm::StoreCacheControlAttr>(
276 "cache_control");
277 if (!attr)
278 return std::nullopt;
279 return std::optional<StoreCacheControl>(attr.getValue());
280 }
281 return std::nullopt;
282}
283
284template <typename OpType>
285int32_t getL1CacheControl(OpType op) {
286 return getL1CacheControl(*getCacheControl(op));
287}
288
289template <typename OpType>
290int32_t getL3CacheControl(OpType op) {
291 return getL3CacheControl(*getCacheControl(op));
292}
293
294template <typename OpType>
295static std::optional<ArrayAttr>
296getCacheControlMetadata(ConversionPatternRewriter &rewriter, OpType op) {
297 if (!getCacheControl(op))
298 return {};
299
300 constexpr int32_t decorationCacheControlArity{3};
301 constexpr int32_t loadCacheControlKey{6442};
302 constexpr int32_t storeCacheControlKey{6443};
303 constexpr bool isLoad = std::is_same_v<OpType, BlockLoad2dOp> ||
304 std::is_same_v<OpType, BlockPrefetch2dOp> ||
305 std::is_same_v<OpType, LLVM::LoadOp> ||
306 std::is_same_v<OpType, BlockLoadOp> ||
307 std::is_same_v<OpType, PrefetchOp>;
308
309 // If the cache control is USE_DEFAULT, then we don’t emit any metadata.
310 // Assert that if one of the L1 or L3 cache control values is USE_DEFAULT
311 // (represented as -1), then both must be USE_DEFAULT; otherwise there is a
312 // bug.
313 assert(((getL1CacheControl<OpType>(op) == -1) ==
314 (getL3CacheControl<OpType>(op) == -1)) &&
315 "If one of L1 or L3 cache control is USE_DEFAULT, both must be "
316 "USE_DEFAULT");
317
318 if (getL1CacheControl<OpType>(op) == -1 &&
319 getL3CacheControl<OpType>(op) == -1)
320 return {};
321 const int32_t controlKey{isLoad ? loadCacheControlKey : storeCacheControlKey};
323 controlKey, 0, getL1CacheControl<OpType>(op)};
325 controlKey, 1, getL3CacheControl<OpType>(op)};
326 auto arrayAttrL1 = rewriter.getI32ArrayAttr(decorationsL1);
327 auto arrayAttrL3 = rewriter.getI32ArrayAttr(decorationsL3);
328
329 SmallVector<Attribute, 2> combinedAttrs = {arrayAttrL1, arrayAttrL3};
330 return rewriter.getArrayAttr(combinedAttrs);
331}
332
333//===----------------------------------------------------------------------===//
334// Cache control annotation utilities
335//
336// Instead of attaching cache control as MLIR attributes and handling them
337// during LLVM translation, we directly emit llvm.intr.ptr.annotation op in
338// MLIR.
339//===----------------------------------------------------------------------===//
340
341/// Build one cache-control payload string per attribute.
342///
343/// Each Attribute is expected to be an ArrayAttr of 3 IntegerAttr values:
344/// [SPIR-V decoration token, cache level, cache control value]
345///
346/// A single entry produces a string like: {6442:"0,1"}
347/// where the quote characters (0x22) will appear as \22 in LLVM IR textual
348/// form.
350buildCacheControlPayloads(ArrayRef<Attribute> attrs) {
352 llvm::StringMap<bool> seen;
353
354 for (auto arr : llvm::make_isa_range<ArrayAttr>(attrs)) {
355 auto vals = arr.getValue();
356 assert(vals.size() == 3 &&
357 "Expected exactly 3 integer values (Token, CacheLevel, "
358 "ControlValue) in cache control attribute.");
359
360 auto tokenAttr = dyn_cast<IntegerAttr>(vals[0]);
361 auto secondAttr = dyn_cast<IntegerAttr>(vals[1]);
362 auto thirdAttr = dyn_cast<IntegerAttr>(vals[2]);
363
364 if (!tokenAttr || !secondAttr || !thirdAttr)
365 continue;
366
367 // Produce: {SPIR-V decoration token:"L1 cache control,L3 cache control"}
368 // The quote char (0x22) is embedded literally; LLVM IR prints it as \22.
369 std::string entry =
370 llvm::formatv("{{{0}:\"{1},{2}\"}", tokenAttr.getValue().getZExtValue(),
371 secondAttr.getValue().getZExtValue(),
372 thirdAttr.getValue().getZExtValue());
373
374 // Deduplicate identical annotations.
375 if (!seen.insert({entry, true}).second)
376 continue;
377
378 payloads.push_back(std::move(entry));
379 }
380 return payloads;
381}
382/// Counter for generating unique global variable names.
383static std::atomic<uint64_t> globalNameCounter{0};
384
385/// Get or create a global metadata string and return a !llvm.ptr<1> value
386/// pointing to it. The AddressOfOp is created at the current rewriter
387/// insertion point; the GlobalOp is created at the module start.
388static Value createMetadataStringPtr(ConversionPatternRewriter &rewriter,
389 Operation *moduleOp, Location loc,
390 StringRef value, StringRef nameHint) {
391 // Build null-terminated string.
392 std::string strWithNull = value.str();
393 strWithNull.push_back('\0');
394 StringRef strRef(strWithNull.data(), strWithNull.size());
395
396 auto as1PtrTy = LLVM::LLVMPointerType::get(rewriter.getContext(), 1);
397
398 // Search for an existing global with the same content.
399 for (auto &op : moduleOp->getRegion(0).front()) {
400 if (auto existingGlobal = dyn_cast<LLVM::GlobalOp>(&op)) {
401 if (!existingGlobal.getSection() ||
402 *existingGlobal.getSection() != "llvm.metadata")
403 continue;
404 if (auto strAttr =
405 dyn_cast_or_null<StringAttr>(existingGlobal.getValueOrNull())) {
406 if (strAttr.getValue() == strRef) {
407 return LLVM::AddressOfOp::create(rewriter, loc, as1PtrTy,
408 existingGlobal.getSymName());
409 }
410 }
411 }
412 }
413
414 // Create new global at module start.
415 auto i8Type = rewriter.getI8Type();
416 auto arrayType = LLVM::LLVMArrayType::get(i8Type, strWithNull.size());
417 std::string globalName =
418 llvm::formatv("{0}.{1}", nameHint,
419 globalNameCounter.fetch_add(1, std::memory_order_relaxed))
420 .str();
421
422 {
423 OpBuilder::InsertionGuard guard(rewriter);
424 rewriter.setInsertionPointToStart(&moduleOp->getRegion(0).front());
425
426 auto globalOp =
427 LLVM::GlobalOp::create(rewriter, loc, arrayType,
428 /*isConstant=*/true, LLVM::Linkage::Private,
429 globalName, rewriter.getStringAttr(strRef));
430 globalOp.setSection(StringRef("llvm.metadata"));
431 globalOp.setUnnamedAddr(LLVM::UnnamedAddr::Global);
432 globalOp.setAlignment(1);
433 globalOp.setAddrSpace(1);
434 }
435 // InsertionGuard restores the original insertion point here.
436
437 return LLVM::AddressOfOp::create(rewriter, loc, as1PtrTy, globalName);
438}
439
440/// Annotate a pointer value with cache control metadata by emitting chained
441/// `llvm.intr.ptr.annotation` ops (LLVM::PtrAnnotation).
442///
443/// This is the MLIR-level equivalent of handleDecorationCacheControl() from
444/// the LLVM translation layer. For each cache control attribute, it emits:
445///
446/// %ann = llvm.intr.ptr.annotation %ptr, @".str.cachecontrol.N",
447/// @".str.file.N", 0, null : !llvm.ptr<AS>
448///
449/// Multiple annotations are chained: the result of each annotation op is
450/// fed as the pointer input to the next one.
451///
452/// \param rewriter The pattern rewriter.
453/// \param loc Source location for created ops.
454/// \param ptr The pointer value to annotate.
455/// \param cacheControls The cache control ArrayAttr (from
456/// getCacheControlMetadata).
457/// \param moduleOp The enclosing module (for creating globals).
458/// \returns The annotated pointer value (or the original ptr if no
459/// annotations).
460static Value annotatePtrWithCacheControl(ConversionPatternRewriter &rewriter,
461 Location loc, Value ptr,
462 ArrayAttr cacheControls,
463 Operation *moduleOp) {
464 SmallVector<std::string> payloads =
465 buildCacheControlPayloads(cacheControls.getValue());
466 if (payloads.empty())
467 return ptr;
468
469 auto ptrType = cast<LLVM::LLVMPointerType>(ptr.getType());
470 auto as1PtrTy = LLVM::LLVMPointerType::get(rewriter.getContext(), 1);
471 auto i32Ty = rewriter.getI32Type();
472
473 // Create shared constants for all annotations on this pointer.
474 Value fileStr =
475 createMetadataStringPtr(rewriter, moduleOp, loc, "", ".str.file");
476 Value lineVal = LLVM::ConstantOp::create(rewriter, loc, i32Ty, 0);
477 Value nullAS1 = LLVM::ZeroOp::create(rewriter, loc, as1PtrTy);
478
479 // Chain: each annotation takes the result of the previous one as its
480 // pointer operand.
481 Value curPtr = ptr;
482 for (const std::string &payload : payloads) {
483 Value annStr = createMetadataStringPtr(rewriter, moduleOp, loc, payload,
484 ".str.cachecontrol");
485 auto annOp = LLVM::PtrAnnotation::create(rewriter, loc, ptrType, curPtr,
486 annStr, fileStr, lineVal, nullAS1);
487 curPtr = annOp.getResult();
488 }
489
490 return curPtr;
491}
492
493/// Helper to apply cache control annotation on a pointer operand of a call.
494/// Replaces the pointer argument of the call with an annotated version.
495///
496/// For operations that produce a call (like block load/store/prefetch), the
497/// pointer is typically the first argument. This function:
498/// 1. Builds the annotation chain on the pointer.
499/// 2. Replaces the pointer operand in the provided args list.
500///
501/// \param rewriter The pattern rewriter.
502/// \param loc Source location.
503/// \param ptr The original pointer value (first arg to the call).
504/// \param cacheControls The cache control metadata.
505/// \param moduleOp The enclosing module.
506/// \param args The argument list (modified in place: args[ptrIdx] is
507/// replaced).
508/// \param ptrIdx Index of the pointer in the args list (default 0).
509template <typename OpType>
510static void
511applyCacheControlAnnotation(ConversionPatternRewriter &rewriter, Location loc,
512 OpType op, SmallVectorImpl<Value> &args,
513 Operation *moduleOp, unsigned ptrIdx = 0) {
514 std::optional<ArrayAttr> optCacheControls =
515 getCacheControlMetadata(rewriter, op);
516 if (!optCacheControls)
517 return;
518
519 Value annotatedPtr = annotatePtrWithCacheControl(rewriter, loc, args[ptrIdx],
520 *optCacheControls, moduleOp);
521 args[ptrIdx] = annotatedPtr;
522}
523
524//===----------------------------------------------------------------------===//
525// End cache control annotation utilities
526//===----------------------------------------------------------------------===//
527
528static LLVM::CallOp createDeviceFunctionCall(
529 ConversionPatternRewriter &rewriter, StringRef funcName, Type retType,
530 ArrayRef<Type> argTypes, ArrayRef<Value> args,
531 mlir::ArrayRef<std::pair<unsigned, mlir::StringRef>> paramAttrs,
532 LLVMFuncAttributeOptions funcAttributeOptions, Operation *op) {
533 auto *moduleOp = op->getParentWithTrait<OpTrait::SymbolTable>();
534 assert(moduleOp && "Expecting module");
535 Location loc = op->getLoc();
536
537 auto funcOpRes =
538 LLVM::lookupOrCreateFn(rewriter, moduleOp, funcName, argTypes, retType);
539 assert(!failed(funcOpRes));
540 LLVM::LLVMFuncOp funcOp = funcOpRes.value();
541 funcOp.setCConv(LLVM::cconv::CConv::SPIR_FUNC);
542 funcOp.setConvergent(funcAttributeOptions.isConvergent);
543 funcOp.setNoUnwind(funcAttributeOptions.isNoUnwind);
544 funcOp.setWillReturn(funcAttributeOptions.isWillReturn);
545
546 if (funcAttributeOptions.memEffectsAttr)
547 funcOp.setMemoryEffectsAttr(funcAttributeOptions.memEffectsAttr);
548
549 for (auto [idx, attrName] : paramAttrs)
550 funcOp.setArgAttr(idx, attrName, rewriter.getUnitAttr());
551
552 auto callOp = LLVM::CallOp::create(rewriter, loc, funcOp, args);
553 SmallVector<NamedAttribute> discardableAttrs;
554 auto copyAttr = [&](StringAttr name, Attribute attr) {
555 if (callOp->getInherentAttr(name).has_value())
556 callOp->setInherentAttr(name, attr);
557 else
558 discardableAttrs.emplace_back(name, attr);
559 };
560 for (NamedAttribute attr : funcOp->getDiscardableAttrDictionary())
561 copyAttr(attr.getName(), attr.getValue());
562 funcOp->getName().walkInherentAttrs(
563 funcOp, [&](StringRef name, Attribute &attr) {
564 copyAttr(rewriter.getStringAttr(name), attr);
565 });
566 callOp->setDiscardableAttrs(discardableAttrs);
567
568 return callOp;
569}
570
571static unsigned getNumOperandsPerDword(xevm::ElemType pTy) {
572 switch (pTy) {
573 case xevm::ElemType::F32:
574 case xevm::ElemType::TF32:
575 return 1;
576 case xevm::ElemType::BF16:
577 case xevm::ElemType::F16:
578 return 2;
579 case xevm::ElemType::U8:
580 case xevm::ElemType::S8:
581 case xevm::ElemType::BF8:
582 case xevm::ElemType::F8:
583 return 4;
584 case xevm::ElemType::E2M1:
585 case xevm::ElemType::U4:
586 case xevm::ElemType::S4:
587 return 8;
588 default:
589 llvm_unreachable("unsupported xevm::ElemType");
590 }
591}
592
593class MMAToOCLPattern : public OpConversionPattern<xevm::MMAOp> {
594 using OpConversionPattern::OpConversionPattern;
595 LogicalResult
596 matchAndRewrite(xevm::MMAOp op, xevm::MMAOp::Adaptor adaptor,
597 ConversionPatternRewriter &rewriter) const override {
598 if (!op.getC()) {
599 return rewriter.notifyMatchFailure(op, "OCL requires C operand");
600 }
601 auto precisionA = op.getTypes().getA();
602 auto precisionB = op.getTypes().getB();
603 auto precisionC = op.getTypes().getC();
604 auto precisionD = op.getTypes().getD();
605 if (precisionC != precisionD) {
606 return rewriter.notifyMatchFailure(op, "type of C and D need to match");
607 }
608 if (precisionC != xevm::ElemType::S32 &&
609 precisionC != xevm::ElemType::F32 &&
610 precisionC != xevm::ElemType::F16 &&
611 precisionC != xevm::ElemType::BF16) {
612 return rewriter.notifyMatchFailure(
613 op, "type of C and D must be S32, F32, F16 or BF16");
614 }
615 if (precisionA == xevm::ElemType::S32 ||
616 precisionA == xevm::ElemType::F32) {
617 return rewriter.notifyMatchFailure(op, "type of A cannot be S32 or F32");
618 }
619 if (precisionB == xevm::ElemType::S32 ||
620 precisionB == xevm::ElemType::F32) {
621 return rewriter.notifyMatchFailure(op, "type of B cannot be S32 or F32");
622 }
623 constexpr uint32_t bitWidthPackedA{16};
624 constexpr uint32_t bitWidthPackedB{32};
625 auto loc = op.getLoc();
626
627 auto castIfNeeded = [&](Value val, Type packedType) -> Value {
628 VectorType origTy = cast<VectorType>(val.getType());
629 const uint32_t vecBitSize =
630 origTy.getNumElements() *
631 origTy.getElementType().getIntOrFloatBitWidth();
632 VectorType newTy = VectorType::get(
633 vecBitSize / packedType.getIntOrFloatBitWidth(), packedType);
634 if (origTy != newTy)
635 val = LLVM::BitcastOp::create(rewriter, loc, newTy, val);
636 return val;
637 };
638
639 Value a = op.getA();
640 Type packedAType = (op.getTypes().getA() == xevm::ElemType::TF32)
641 ? cast<Type>(rewriter.getF32Type())
642 : rewriter.getIntegerType(bitWidthPackedA);
643 a = castIfNeeded(a, packedAType);
644
645 Value b = op.getB();
646 Type packedBType = (op.getTypes().getB() == xevm::ElemType::TF32)
647 ? cast<Type>(rewriter.getF32Type())
648 : rewriter.getIntegerType(bitWidthPackedB);
649 b = castIfNeeded(b, packedBType);
650
651 Value c = op.getC();
652 VectorType cOrigTy = cast<VectorType>(c.getType());
653 VectorType resOrigTy = cast<VectorType>(op->getResultTypes()[0]);
654 assert(cOrigTy == resOrigTy && "Accumulator and result type mismatch");
655 // OCL builtins encode bfloat16 as int16
656 VectorType cTy =
657 cOrigTy.getElementType().isBF16()
658 ? VectorType::get(cOrigTy.getShape(), rewriter.getIntegerType(16))
659 : cOrigTy;
660 VectorType resTy = cTy;
661 if (cOrigTy != cTy)
662 c = LLVM::BitcastOp::create(rewriter, loc, cTy, c);
663
664 constexpr int32_t systolicDepth{8};
665 std::string fnName =
666 llvm::formatv("intel_sub_group_{0}_{1}_matrix_mad_k{2}",
667 stringifyElemType(op.getTypes().getA()).str(),
668 stringifyElemType(op.getTypes().getB()).str(),
669 systolicDepth *
670 getNumOperandsPerDword(op.getTypes().getA()))
671 .str();
672 SmallVector<Type> argTypes{a.getType(), b.getType(), cTy};
673 fnName = mangle(fnName, argTypes);
674 SmallVector<Value> args{a, b, c};
675
676 auto memAttr = rewriter.getAttr<LLVM::MemoryEffectsAttr>(
677 /*other=*/LLVM::ModRefInfo::NoModRef,
678 /*argMem=*/LLVM::ModRefInfo::NoModRef,
679 /*inaccessibleMem=*/LLVM::ModRefInfo::NoModRef,
680 /*errnoMem=*/LLVM::ModRefInfo::NoModRef,
681 /*targetMem0=*/LLVM::ModRefInfo::NoModRef,
682 /*targetMem1=*/LLVM::ModRefInfo::NoModRef);
683 auto funcAttrs = convergentNoUnwindWillReturnAttrs;
684 funcAttrs.memEffectsAttr = memAttr;
685 Value result =
686 createDeviceFunctionCall(rewriter, fnName, resTy, argTypes, args, {},
687 funcAttrs, op.getOperation())
688 ->getResult(0);
689
690 if (resOrigTy != resTy)
691 result = LLVM::BitcastOp::create(rewriter, loc, resOrigTy, result);
692
693 rewriter.replaceOp(op, result);
694 return success();
695 }
696};
697
698class PrefetchToOCLPattern : public OpConversionPattern<PrefetchOp> {
699 using OpConversionPattern::OpConversionPattern;
700 LogicalResult
701 matchAndRewrite(PrefetchOp op, PrefetchOp::Adaptor adaptor,
702 ConversionPatternRewriter &rewriter) const override {
703 auto loc = op.getLoc();
704 auto *moduleOp = op->getParentWithTrait<OpTrait::SymbolTable>();
705
706 const std::string fnName{"_Z8prefetchPU3AS1Kcm"};
707 Value one =
708 LLVM::ConstantOp::create(rewriter, loc, rewriter.getI64Type(), 1);
709 SmallVector<Value> args{op.getPtr(), one};
710
711 // Annotate pointer with cache control before passing to the call.
712 applyCacheControlAnnotation(rewriter, loc, op, args, moduleOp,
713 /*ptrIdx=*/0);
714
715 SmallVector<Type> argTypes;
716 for (auto arg : args)
717 argTypes.push_back(arg.getType());
718 auto funcAttr = noUnwindAttrs;
719 auto memAttr = rewriter.getAttr<LLVM::MemoryEffectsAttr>(
720 /*other=*/LLVM::ModRefInfo::NoModRef,
721 /*argMem=*/LLVM::ModRefInfo::Ref,
722 /*inaccessibleMem=*/LLVM::ModRefInfo::NoModRef,
723 /*errnoMem=*/LLVM::ModRefInfo::NoModRef,
724 /*targetMem0=*/LLVM::ModRefInfo::NoModRef,
725 /*targetMem1=*/LLVM::ModRefInfo::NoModRef);
726 funcAttr.memEffectsAttr = memAttr;
727
728 createDeviceFunctionCall(rewriter, fnName,
729 LLVM::LLVMVoidType::get(rewriter.getContext()),
730 argTypes, args, {}, funcAttr, op.getOperation());
731 rewriter.eraseOp(op);
732 return success();
733 }
734};
735
736class MemfenceToOCLPattern : public OpConversionPattern<MemfenceOp> {
737 using OpConversionPattern::OpConversionPattern;
738 LogicalResult
739 matchAndRewrite(MemfenceOp op, MemfenceOp::Adaptor adaptor,
740 ConversionPatternRewriter &rewriter) const override {
741 auto loc = op.getLoc();
742 const std::string fnName{"atomic_work_item_fence"};
743 int memScope, addrSpace;
744 switch (op.getAddrspace()) {
745 case xevm::AddrSpace::SHARED:
746 addrSpace = 1; // CLK_LOCAL_MEM_FENCE
747 break;
748 case xevm::AddrSpace::GLOBAL:
749 addrSpace = 2; // CLK_GLOBAL_MEM_FENCE
750 break;
751 default:
752 // GENERIC is not supported in OpenCL
753 return rewriter.notifyMatchFailure(
754 op, "Fence only supports global and shared address spaces.");
755 }
756 switch (op.getScope()) {
757 case xevm::MemScope::WORKGROUP:
758 memScope = 1;
759 break;
760 case xevm::MemScope::DEVICE:
761 memScope = 2;
762 break;
763 default:
764 // CLUSTER and SYSTEM are not supported in OpenCL
765 return rewriter.notifyMatchFailure(
766 op, "Fence only supports workgroup and device memory scopes.");
767 }
768 Type i32Type = rewriter.getI32Type();
769 Value acqRel = LLVM::ConstantOp::create(rewriter, loc, i32Type, 4);
770 Value memScopeConst =
771 LLVM::ConstantOp::create(rewriter, loc, i32Type, memScope);
772 Value addrSpaceConst =
773 LLVM::ConstantOp::create(rewriter, loc, i32Type, addrSpace);
774 SmallVector<Value> args{addrSpaceConst, acqRel, memScopeConst};
775 SmallVector<Type> argTypes{3, i32Type};
776 createDeviceFunctionCall(rewriter, mangle(fnName, argTypes),
777 LLVM::LLVMVoidType::get(rewriter.getContext()),
778 argTypes, args, {}, noUnwindAttrs,
779 op.getOperation());
780 rewriter.eraseOp(op);
781 return success();
782 }
783};
784template <typename OpType>
785class LoadStorePrefetchToOCLPattern : public OpConversionPattern<OpType> {
786 using OpConversionPattern<OpType>::OpConversionPattern;
787 LogicalResult
788 matchAndRewrite(OpType op, typename OpType::Adaptor adaptor,
789 ConversionPatternRewriter &rewriter) const override {
790 constexpr bool isLoad = std::is_same_v<OpType, BlockLoad2dOp>;
791 constexpr bool isPrefetch = std::is_same_v<OpType, BlockPrefetch2dOp>;
792
793 auto loc = op.getLoc();
794 auto *moduleOp = op->template getParentWithTrait<OpTrait::SymbolTable>();
795 VectorType vecType;
796 bool packReg = false;
797 bool transpose = false;
798 if constexpr (isLoad) {
799 vecType = op.getRes().getType();
800 packReg = op.getPackRegister();
801 transpose = op.getTranspose();
802 } else if constexpr (!isPrefetch) {
803 vecType = op.getStoredVal().getType();
804 }
805
806 auto i32Type = rewriter.getI32Type();
807 Value byteCoord =
808 LLVM::UndefOp::create(rewriter, loc, VectorType::get(2, i32Type));
809 Value zero = LLVM::ConstantOp::create(rewriter, loc, i32Type, 0);
810 Value one = LLVM::ConstantOp::create(rewriter, loc, i32Type, 1);
811 byteCoord = LLVM::InsertElementOp::create(
812 rewriter, loc, VectorType::get(2, i32Type), byteCoord, op.getX(), zero);
813 byteCoord = LLVM::InsertElementOp::create(
814 rewriter, loc, VectorType::get(2, i32Type), byteCoord, op.getY(), one);
815 SmallVector<Value> args{op.getPtr(), op.getBaseWidth(), op.getBaseHeight(),
816 op.getBasePitch(), byteCoord};
817
818 // Annotate pointer (args[0]) with cache control before the call.
819 applyCacheControlAnnotation(rewriter, loc, op, args, moduleOp,
820 /*ptrIdx=*/0);
821
822 SmallVector<Type> retTypes;
823 Value spvLoadDstPtr;
824 std::string funcName{"intel_sub_group_2d_block_"};
825 std::string bitWidthId;
826 LLVMFuncAttributeOptions funcAttr{noUnwindWillReturnAttrs};
827 SmallVector<std::pair<unsigned, StringRef>, 4> paramAttrs;
828 if constexpr (isPrefetch) { // Prefetch
829 funcName += "prefetch";
830 paramAttrs = {std::make_pair(0, LLVM::LLVMDialect::getNonNullAttrName())};
831 auto memAttr = rewriter.getAttr<LLVM::MemoryEffectsAttr>(
832 /*other=*/LLVM::ModRefInfo::NoModRef,
833 /*argMem=*/LLVM::ModRefInfo::Ref,
834 /*inaccessibleMem=*/LLVM::ModRefInfo::NoModRef,
835 /*errnoMem=*/LLVM::ModRefInfo::NoModRef,
836 /*targetMem0=*/LLVM::ModRefInfo::NoModRef,
837 /*targetMem1=*/LLVM::ModRefInfo::NoModRef);
838 funcAttr = noUnwindAttrs;
839 funcAttr.memEffectsAttr = memAttr;
840 } else {
841 auto vecElemType = vecType.getElementType();
842 auto vecElemBitWidth = vecElemType.getIntOrFloatBitWidth();
843 auto vecNumElems = vecType.getNumElements();
844 // OpenCL Intel 2D block load has a special case
845 // when element bit size is 8 and tile width is 32, which is twice
846 // the subgroup size, loaded element is packed as i16.
847 // To reflect this, element bit size is updated to 16 and
848 // vector length is reduced by half.
849 if (op.getElemSizeInBits() == 8 && op.getTileWidth() == 32) {
850 vecElemBitWidth = 16;
851 vecElemType = rewriter.getI16Type();
852 vecNumElems = vecNumElems / 2;
853 }
854 Value numElems =
855 LLVM::ConstantOp::create(rewriter, loc, i32Type, vecNumElems);
856 auto dstOrSrcPtr = LLVM::AllocaOp::create(
857 rewriter, loc, LLVM::LLVMPointerType::get(rewriter.getContext()),
858 vecElemType, numElems);
859 args.push_back(dstOrSrcPtr);
860 if constexpr (isLoad) { // Load
861 funcName += "read";
862 bitWidthId = getTypeMangling(vecElemType, /*isUnsigned=*/true);
863 if (packReg)
864 funcName += "_transform";
865 else if (transpose)
866 funcName += "_transpose";
867 spvLoadDstPtr = dstOrSrcPtr;
868 retTypes.push_back(vecType);
869 paramAttrs = {
870 std::make_pair(0, LLVM::LLVMDialect::getNonNullAttrName()),
871 std::make_pair(0, LLVM::LLVMDialect::getReadonlyAttrName()),
872 std::make_pair(5, LLVM::LLVMDialect::getNonNullAttrName()),
873 std::make_pair(5, LLVM::LLVMDialect::getWriteOnlyAttrName()),
874 };
875 } else { // Store
876 funcName += "write";
877 bitWidthId = (vecElemBitWidth == 32)
878 ? "j"
879 : ((vecElemBitWidth == 16) ? "t" : "h");
880 LLVM::StoreOp::create(rewriter, loc, op.getStoredVal(), dstOrSrcPtr);
881 paramAttrs = {
882 std::make_pair(0, LLVM::LLVMDialect::getNonNullAttrName()),
883 std::make_pair(0, LLVM::LLVMDialect::getWriteOnlyAttrName()),
884 std::make_pair(5, LLVM::LLVMDialect::getNonNullAttrName()),
885 std::make_pair(5, LLVM::LLVMDialect::getReadonlyAttrName()),
886 };
887 }
888 }
889
890 funcName =
891 llvm::formatv("{0}_{1}b_{2}r{3}x{4}c", funcName, op.getElemSizeInBits(),
892 op.getTileHeight(), op.getTileWidth(), op.getVBlocks())
893 .str();
894 std::string prefetchCode("");
895 if (!isPrefetch)
896 prefetchCode += "P";
897 funcName = llvm::formatv("_Z{0}{1}PU3AS1viiiDv2_i{2}{3}", funcName.size(),
898 funcName, prefetchCode, bitWidthId)
899 .str();
900 SmallVector<Type> argTypes;
901 for (auto arg : args) {
902 argTypes.push_back(arg.getType());
903 }
904 createDeviceFunctionCall(
905 rewriter, funcName, LLVM::LLVMVoidType::get(rewriter.getContext()),
906 argTypes, args, paramAttrs, funcAttr, op.getOperation());
907
908 if constexpr (isLoad)
909 rewriter.replaceOp(
910 op, LLVM::LoadOp::create(rewriter, loc, vecType, spvLoadDstPtr));
911 else
912 rewriter.eraseOp(op);
913 return success();
914 }
915};
916
917template <typename OpType>
918class BlockLoadStore1DToOCLPattern : public OpConversionPattern<OpType> {
919 using OpConversionPattern<OpType>::OpConversionPattern;
920 LogicalResult
921 matchAndRewrite(OpType op, typename OpType::Adaptor adaptor,
922 ConversionPatternRewriter &rewriter) const override {
923 constexpr bool isStore = std::is_same_v<OpType, xevm::BlockStoreOp>;
924 auto loc = op.getLoc();
925 auto *moduleOp = op->template getParentWithTrait<OpTrait::SymbolTable>();
926
927 // Get OpenCL function name
928 // https://registry.khronos.org/OpenCL/extensions/
929 // intel/cl_intel_subgroup_local_block_io.html
930 std::string funcName{"intel_sub_group_block_"};
931 // Value or Result type can be vector or scalar
932 Type valOrResTy;
933 if constexpr (isStore) {
934 funcName += "write_u";
935 valOrResTy = op.getVal().getType();
936 } else {
937 funcName += "read_u";
938 valOrResTy = op.getType();
939 }
940 // Get element type of the vector/scalar
941 VectorType vecTy = dyn_cast<VectorType>(valOrResTy);
942 Type elemType = vecTy ? vecTy.getElementType() : valOrResTy;
943 funcName += getTypeMangling(elemType);
944 if (vecTy)
945 funcName += std::to_string(vecTy.getNumElements());
946 SmallVector<Type, 2> argTypes{};
947 // XeVM BlockLoad/StoreOp always use signless integer types
948 // but OpenCL builtins expect unsigned types
949 // use unsigned types for mangling
950 SmallVector<bool, 2> isUnsigned{};
951 // arg0: pointer to the src/dst address
952 // arg1 - only if store : vector to store
953 // Prepare arguments
954 SmallVector<Value, 2> args{};
955 args.push_back(op.getPtr());
956 argTypes.push_back(op.getPtr().getType());
957 isUnsigned.push_back(true);
958
959 // Annotate pointer (args[0]) with cache control.
960 applyCacheControlAnnotation(rewriter, loc, op, args, moduleOp,
961 /*ptrIdx=*/0);
962 // Update argTypes[0] in case the pointer type changed (it shouldn't
963 // change type, but the value is now the annotated pointer).
964 argTypes[0] = args[0].getType();
965
966 Type retType;
967 if constexpr (isStore) {
968 args.push_back(op.getVal());
969 argTypes.push_back(op.getVal().getType());
970 isUnsigned.push_back(true);
971 retType = LLVM::LLVMVoidType::get(rewriter.getContext());
972 } else {
973 retType = valOrResTy;
974 }
975 funcName = std::string("_Z") + std::to_string(funcName.size()) + funcName +
976 "PU3AS" +
977 std::to_string(op.getPtr().getType().getAddressSpace());
978 funcName += getTypeMangling(elemType, /*isUnsigned=*/true);
979 if constexpr (isStore)
980 funcName += getTypeMangling(valOrResTy, /*isUnsigned=*/true);
981 LLVMFuncAttributeOptions funcAttr{noUnwindWillReturnAttrs};
982
983 LLVM::CallOp call =
984 createDeviceFunctionCall(rewriter, funcName, retType, argTypes, args,
985 {}, funcAttr, op.getOperation());
986
987 if constexpr (isStore)
988 rewriter.eraseOp(op);
989 else
990 rewriter.replaceOp(op, call->getResult(0));
991 return success();
992 }
993};
994
995template <typename OpType>
996class LLVMLoadStoreToOCLPattern : public OpConversionPattern<OpType> {
997 using OpConversionPattern<OpType>::OpConversionPattern;
998 LogicalResult
999 matchAndRewrite(OpType op, typename OpType::Adaptor adaptor,
1000 ConversionPatternRewriter &rewriter) const override {
1001 if (!op->hasDiscardableAttr("cache_control"))
1002 return failure();
1003
1004 auto *moduleOp = op->template getParentWithTrait<OpTrait::SymbolTable>();
1005 std::optional<ArrayAttr> optCacheControls =
1006 getCacheControlMetadata(rewriter, op);
1007 if (!optCacheControls) {
1008 rewriter.modifyOpInPlace(
1009 op, [&]() { op->removeDiscardableAttr("cache_control"); });
1010 return success();
1011 }
1012
1013 // Determine which operand is the pointer.
1014 constexpr bool isStore = std::is_same_v<OpType, LLVM::StoreOp>;
1015 unsigned ptrIdx = isStore ? 1 : 0;
1016 Value ptr = op->getOperand(ptrIdx);
1017
1018 // Emit annotation intrinsic calls on the pointer.
1019 Value annotatedPtr = annotatePtrWithCacheControl(
1020 rewriter, op->getLoc(), ptr, *optCacheControls, moduleOp);
1021
1022 // Replace the pointer operand with the annotated one.
1023 rewriter.modifyOpInPlace(op, [&]() {
1024 op->setOperand(ptrIdx, annotatedPtr);
1025 op->removeDiscardableAttr("cache_control");
1026 });
1027 return success();
1028 }
1029};
1030
1031//===----------------------------------------------------------------------===//
1032// GPU index id operations
1033//===----------------------------------------------------------------------===//
1034/*
1035// Launch Config ops
1036// dimidx - x, y, z - is fixed to i32
1037// return type is set by XeVM type converter
1038// get_local_id
1039xevm::WorkitemIdXOp;
1040xevm::WorkitemIdYOp;
1041xevm::WorkitemIdZOp;
1042// get_local_size
1043xevm::WorkgroupDimXOp;
1044xevm::WorkgroupDimYOp;
1045xevm::WorkgroupDimZOp;
1046// get_group_id
1047xevm::WorkgroupIdXOp;
1048xevm::WorkgroupIdYOp;
1049xevm::WorkgroupIdZOp;
1050// get_num_groups
1051xevm::GridDimXOp;
1052xevm::GridDimYOp;
1053xevm::GridDimZOp;
1054// get_global_id : to be added if needed
1055*/
1056
1057// Helpers to get the OpenCL function name and dimension argument for each op.
1058static std::pair<StringRef, int64_t> getConfig(xevm::WorkitemIdXOp) {
1059 return {"get_local_id", 0};
1060}
1061static std::pair<StringRef, int64_t> getConfig(xevm::WorkitemIdYOp) {
1062 return {"get_local_id", 1};
1063}
1064static std::pair<StringRef, int64_t> getConfig(xevm::WorkitemIdZOp) {
1065 return {"get_local_id", 2};
1066}
1067static std::pair<StringRef, int64_t> getConfig(xevm::WorkgroupDimXOp) {
1068 return {"get_local_size", 0};
1069}
1070static std::pair<StringRef, int64_t> getConfig(xevm::WorkgroupDimYOp) {
1071 return {"get_local_size", 1};
1072}
1073static std::pair<StringRef, int64_t> getConfig(xevm::WorkgroupDimZOp) {
1074 return {"get_local_size", 2};
1075}
1076static std::pair<StringRef, int64_t> getConfig(xevm::WorkgroupIdXOp) {
1077 return {"get_group_id", 0};
1078}
1079static std::pair<StringRef, int64_t> getConfig(xevm::WorkgroupIdYOp) {
1080 return {"get_group_id", 1};
1081}
1082static std::pair<StringRef, int64_t> getConfig(xevm::WorkgroupIdZOp) {
1083 return {"get_group_id", 2};
1084}
1085static std::pair<StringRef, int64_t> getConfig(xevm::GridDimXOp) {
1086 return {"get_num_groups", 0};
1087}
1088static std::pair<StringRef, int64_t> getConfig(xevm::GridDimYOp) {
1089 return {"get_num_groups", 1};
1090}
1091static std::pair<StringRef, int64_t> getConfig(xevm::GridDimZOp) {
1092 return {"get_num_groups", 2};
1093}
1094/// Replace `xevm.*` with an `llvm.call` to the corresponding OpenCL func with
1095/// a constant argument for the dimension - x, y or z.
1096template <typename OpType>
1097class LaunchConfigOpToOCLPattern : public OpConversionPattern<OpType> {
1098 using OpConversionPattern<OpType>::OpConversionPattern;
1099 LogicalResult
1100 matchAndRewrite(OpType op, typename OpType::Adaptor adaptor,
1101 ConversionPatternRewriter &rewriter) const override {
1102 Location loc = op->getLoc();
1103 auto [baseName, dim] = getConfig(op);
1104 Type dimTy = rewriter.getI32Type();
1105 Value dimVal = LLVM::ConstantOp::create(rewriter, loc, dimTy,
1106 static_cast<int64_t>(dim));
1107 std::string func = mangle(baseName, {dimTy}, {true});
1108 Type resTy = op.getType();
1109 auto call =
1110 createDeviceFunctionCall(rewriter, func, resTy, {dimTy}, {dimVal}, {},
1111 noUnwindWillReturnAttrs, op.getOperation());
1112 constexpr auto noModRef = LLVM::ModRefInfo::NoModRef;
1113 auto memAttr = rewriter.getAttr<LLVM::MemoryEffectsAttr>(
1114 /*other=*/noModRef,
1115 /*argMem=*/noModRef, /*inaccessibleMem=*/noModRef,
1116 /*errnoMem=*/noModRef,
1117 /*targetMem0=*/noModRef,
1118 /*targetMem1=*/noModRef);
1119 call.setMemoryEffectsAttr(memAttr);
1120 rewriter.replaceOp(op, call);
1121 return success();
1122 }
1123};
1124
1125/*
1126// Subgroup ops
1127// get_sub_group_local_id
1128xevm::LaneIdOp;
1129// get_sub_group_id
1130xevm::SubgroupIdOp;
1131// get_sub_group_size
1132xevm::SubgroupSizeOp;
1133// get_num_sub_groups : to be added if needed
1134*/
1135
1136// Helpers to get the OpenCL function name for each op.
1137static StringRef getConfig(xevm::LaneIdOp) { return "get_sub_group_local_id"; }
1138static StringRef getConfig(xevm::SubgroupIdOp) { return "get_sub_group_id"; }
1139static StringRef getConfig(xevm::SubgroupSizeOp) {
1140 return "get_sub_group_size";
1141}
1142template <typename OpType>
1143class SubgroupOpWorkitemOpToOCLPattern : public OpConversionPattern<OpType> {
1144 using OpConversionPattern<OpType>::OpConversionPattern;
1145 LogicalResult
1146 matchAndRewrite(OpType op, typename OpType::Adaptor adaptor,
1147 ConversionPatternRewriter &rewriter) const override {
1148 std::string func = mangle(getConfig(op).str(), {});
1149 Type resTy = op.getType();
1150 auto call =
1151 createDeviceFunctionCall(rewriter, func, resTy, {}, {}, {},
1152 noUnwindWillReturnAttrs, op.getOperation());
1153 constexpr auto noModRef = LLVM::ModRefInfo::NoModRef;
1154 auto memAttr = rewriter.getAttr<LLVM::MemoryEffectsAttr>(
1155 /*other=*/noModRef,
1156 /*argMem=*/noModRef, /*inaccessibleMem=*/noModRef,
1157 /*errnoMem=*/noModRef,
1158 /*targetMem0=*/noModRef,
1159 /*targetMem1=*/noModRef);
1160 call.setMemoryEffectsAttr(memAttr);
1161 rewriter.replaceOp(op, call);
1162 return success();
1163 }
1164};
1165
1166/// SPIR-V, and so the OpenCL builtins the float conversions call into, only
1167/// provides vector types of 2, 3, 4, 8 and 16 elements.
1168static bool isSupportedSPIRVVectorLength(int64_t numElements) {
1169 return llvm::is_contained({2, 3, 4, 8, 16}, numElements);
1170}
1171
1172/// Bitcasts `val` to `ty` unless it already has that type.
1173static Value castIfNeeded(ConversionPatternRewriter &rewriter, Location loc,
1174 Type ty, Value val) {
1175 if (val.getType() == ty)
1176 return val;
1177 return LLVM::BitcastOp::create(rewriter, loc, ty, val);
1178}
1179
1180/// Selects `numElements` leading elements of the vector `val`. Used to drop the
1181/// padding the 3 element case needs, as the hardware conversions work on whole
1182/// pairs of elements.
1183static Value takeLeadingElements(ConversionPatternRewriter &rewriter,
1184 Location loc, Value val, int64_t numElements) {
1185 auto vecTy = cast<VectorType>(val.getType());
1186 if (vecTy.getNumElements() == numElements)
1187 return val;
1189 llvm::to_vector(llvm::seq<int32_t>(0, static_cast<int32_t>(numElements)));
1190 return LLVM::ShuffleVectorOp::create(rewriter, loc, val, val, mask);
1191}
1192
1193//
1194// Note: TruncfToOCLPattern and ExtfToOCLPattern does not lower to OpenCL API
1195// calls as there are not official ones yet. They are lowered directly to Intel
1196// graphics compiler built in functions.
1197// See
1198// https://github.com/intel/intel-graphics-compiler/tree/master/IGC/BiFModule/Implementation/SPV_INTEL_fp_conversions
1199// for builtin function usage. The folder contains implementation of
1200// experimental SPIR-V extension for truncf and extf using builtin functions.
1201// TODO: Move to OpenCL API call once they are available.
1202//
1203
1204class TruncfToOCLPattern : public OpConversionPattern<TruncfOp> {
1205 using OpConversionPattern::OpConversionPattern;
1206 LogicalResult
1207 matchAndRewrite(TruncfOp op, TruncfOp::Adaptor adaptor,
1208 ConversionPatternRewriter &rewriter) const override {
1209 // Supported source and result types are resticted for now.
1210 auto srcEtype = op.getSrcEtype().getEtype();
1211 auto dstEtype = op.getDstEtype().getEtype();
1212 // The conversions are provided as OpenCL builtins, one per vector length,
1213 // so only the SPIR-V vector lengths can be lowered. A wider conversion has
1214 // to be split into several ops before reaching this pattern.
1215 //
1216 // Scalar case is not supported until usage case become clear.
1217 auto vecSrcTy = dyn_cast<VectorType>(op.getSrc().getType());
1218 if (!vecSrcTy) {
1219 return rewriter.notifyMatchFailure(op, "Scalar src is not supported.");
1220 }
1221 int64_t numElements = vecSrcTy.getNumElements();
1222 if (!isSupportedSPIRVVectorLength(numElements))
1223 return rewriter.notifyMatchFailure(
1224 op, "src vector length must be 2, 3, 4, 8 or 16");
1225 // The destination is scalar only where the packed values fit in one byte,
1226 // which SPIR-V spells as a scalar rather than a one element vector.
1227 Type dstTy = op.getDst().getType();
1228 Location loc = op.getLoc();
1229 Value src = op.getSrc();
1230 auto memAttr = rewriter.getAttr<LLVM::MemoryEffectsAttr>(
1231 /*other=*/LLVM::ModRefInfo::NoModRef,
1232 /*argMem=*/LLVM::ModRefInfo::NoModRef,
1233 /*inaccessibleMem=*/LLVM::ModRefInfo::NoModRef,
1234 /*errnoMem=*/LLVM::ModRefInfo::NoModRef,
1235 /*targetMem0=*/LLVM::ModRefInfo::NoModRef,
1236 /*targetMem1=*/LLVM::ModRefInfo::NoModRef);
1237 auto funcAttrs = convergentNoUnwindWillReturnAttrs;
1238 funcAttrs.memEffectsAttr = memAttr;
1239
1240 // Handle the case where dst type is fp4 first.
1241 if (dstEtype == TruncfDstElemTypes::E2M1) {
1242 // `__builtin_IB_dnscl_{hf16,bf16}(uint a, uint b, convert_to, mode)`
1243 // takes two dwords, each holding two source elements, and packs each pair
1244 // into one byte of the result dword. `mode` picks which bytes of that
1245 // dword are written: mode 0 writes bytes 0 and 2, mode 2 writes bytes 1
1246 // and 3. Two calls with complementary modes therefore OR together into
1247 // one fully packed dword covering eight source elements.
1248 //
1249 // A pair of elements is the conversion granularity, so an odd length is
1250 // padded up and the spare nibble left undefined.
1251 constexpr int kDnsclConvertToE2M1 = 1;
1252 constexpr int kDnsclModeBytes02 = 0;
1253 constexpr int kDnsclModeBytes13 = 2;
1254 // One dword lane per element pair, and one result byte per lane. The op
1255 // verifier has already checked that the destination is exactly that wide.
1256 int64_t numLanes = llvm::divideCeil(numElements, 2);
1257
1258 Type i32Ty = rewriter.getI32Type();
1259 Type i8Ty = rewriter.getI8Type();
1260 // Pad an odd length up to a whole number of pairs, then view the source
1261 // as dword lanes.
1262 Value padded = src;
1263 if (numElements != numLanes * 2) {
1264 SmallVector<int32_t> mask = llvm::to_vector(
1265 llvm::seq<int32_t>(0, static_cast<int32_t>(numElements)));
1266 // The padding element is never read back, so any valid index will do.
1267 mask.append(static_cast<size_t>(numLanes * 2 - numElements), 0);
1268 padded = LLVM::ShuffleVectorOp::create(rewriter, loc, src, src, mask);
1269 }
1270 // A single lane is passed as a bare i32 rather than a one element vector,
1271 // which SPIR-V has no type for.
1272 Value laneVec;
1273 if (numLanes > 1)
1274 laneVec = LLVM::BitcastOp::create(
1275 rewriter, loc, VectorType::get(numLanes, i32Ty), padded);
1276 else
1277 laneVec = LLVM::BitcastOp::create(rewriter, loc, i32Ty, padded);
1278 auto getLane = [&](int64_t idx) -> Value {
1279 if (numLanes == 1)
1280 return laneVec;
1281 Value pos =
1282 LLVM::ConstantOp::create(rewriter, loc, rewriter.getI32Type(), idx);
1283 return LLVM::ExtractElementOp::create(rewriter, loc, laneVec, pos)
1284 ->getResult(0);
1285 };
1286
1287 std::string fnName = "__builtin_IB_dnscl_";
1288 fnName += (srcEtype == TruncfSrcElemTypes::F16) ? "hf16" : "bf16";
1289 Value convertTo =
1290 LLVM::ConstantOp::create(rewriter, loc, i32Ty, kDnsclConvertToE2M1);
1291 auto genDnscl = [&](Value lo, Value hi, int mode) -> Value {
1292 Value modeVal = LLVM::ConstantOp::create(rewriter, loc, i32Ty, mode);
1293 SmallVector<Type> argTypes{lo.getType(), hi.getType(),
1294 convertTo.getType(), modeVal.getType()};
1295 SmallVector<Value> args{lo, hi, convertTo, modeVal};
1296 return createDeviceFunctionCall(rewriter, fnName, i32Ty, argTypes, args,
1297 {}, funcAttrs, op.getOperation())
1298 ->getResult(0);
1299 };
1300
1301 Value result;
1302 if (numLanes <= 2) {
1303 // Fewer than four lanes cannot fill a dword, so a single call is made
1304 // and the written bytes, 0 and 2, are compacted afterwards.
1305 Value lo = getLane(0);
1306 Value hi =
1307 numLanes == 2
1308 ? getLane(1)
1309 : LLVM::UndefOp::create(rewriter, loc, i32Ty)->getResult(0);
1310 Value dword = genDnscl(lo, hi, kDnsclModeBytes02);
1311 if (numLanes == 1) {
1312 // A single byte, so the low one, is all that is kept.
1313 result = LLVM::TruncOp::create(rewriter, loc, i8Ty, dword);
1314 } else {
1315 Value bytes = LLVM::BitcastOp::create(
1316 rewriter, loc, VectorType::get(4, i8Ty), dword);
1317 result = LLVM::ShuffleVectorOp::create(rewriter, loc, bytes, bytes,
1318 ArrayRef<int32_t>{0, 2});
1319 }
1320 } else {
1321 // Four lanes, eight source elements, per fully packed dword.
1322 SmallVector<Value> dwords;
1323 for (int64_t base = 0; base < numLanes; base += 4) {
1324 // Each lane is bound to a name first: `getLane` builds ops, and the
1325 // order of evaluation within an argument list is unspecified, which
1326 // would otherwise leave the order of the emitted ops up to the host
1327 // compiler.
1328 Value lane0 = getLane(base);
1329 Value lane2 = getLane(base + 2);
1330 Value even = genDnscl(lane0, lane2, kDnsclModeBytes02);
1331 Value lane1 = getLane(base + 1);
1332 Value lane3 = getLane(base + 3);
1333 Value odd = genDnscl(lane1, lane3, kDnsclModeBytes13);
1334 dwords.push_back(LLVM::OrOp::create(rewriter, loc, even, odd));
1335 }
1336 if (dwords.size() == 1) {
1337 result = dwords.front();
1338 } else {
1339 Type packedTy = VectorType::get(dwords.size(), i32Ty);
1340 result = LLVM::UndefOp::create(rewriter, loc, packedTy);
1341 for (auto [idx, dword] : llvm::enumerate(dwords)) {
1342 Value pos = LLVM::ConstantOp::create(rewriter, loc, i32Ty, idx);
1343 result =
1344 LLVM::InsertElementOp::create(rewriter, loc, result, dword, pos)
1345 ->getResult(0);
1346 }
1347 }
1348 }
1349 rewriter.replaceOp(op, castIfNeeded(rewriter, loc, dstTy, result));
1350 return success();
1351 }
1352
1353 // Handle the case where dst type is fp8.
1354 // The fp8 conversions come as one builtin per vector length, so the length
1355 // is simply appended to the builtin name.
1356 std::string lenSuffix = std::to_string(numElements);
1357 // BF16 type needs some preprocessing before conversion,
1358 // First extended to F32 and then truncated to F16.
1359 if (srcEtype == TruncfSrcElemTypes::BF16) {
1360 // Step 1: Extend to F32
1361 // Use floatN __builtin_IB_bftof_N(shortN)
1362 src = LLVM::BitcastOp::create(
1363 rewriter, op.getLoc(),
1364 VectorType::get(vecSrcTy.getShape(), rewriter.getI16Type()), src);
1365 std::string fnName = "__builtin_IB_bftof_" + lenSuffix;
1366 SmallVector<Type> argTypes{src.getType()};
1367 SmallVector<Value> args{src};
1368 Type resTy = VectorType::get(vecSrcTy.getShape(), rewriter.getF32Type());
1369 src = createDeviceFunctionCall(rewriter, fnName, resTy, argTypes, args,
1370 {}, funcAttrs, op.getOperation())
1371 ->getResult(0);
1372 // Step 2: Truncf to F16
1373 // Use halfN convert_halfN(floatN)
1374 std::string truncFnName = "convert_half" + lenSuffix;
1375 SmallVector<Type> truncArgTypes{src.getType()};
1376 SmallVector<Value> truncArgs{src};
1377 truncFnName = mangle(truncFnName, truncArgTypes);
1378 resTy = VectorType::get(vecSrcTy.getShape(), rewriter.getF16Type());
1379 src =
1380 createDeviceFunctionCall(rewriter, truncFnName, resTy, truncArgTypes,
1381 truncArgs, {}, funcAttrs, op.getOperation())
1382 ->getResult(0);
1383 }
1384 if (dstEtype == TruncfDstElemTypes::BF8) { // Float8E5M2Type
1385 // Use charN __builtin_IB_hftobf8_N(halfN)
1386 std::string fnName = "__builtin_IB_hftobf8_" + lenSuffix;
1387 SmallVector<Type> argTypes{src.getType()};
1388 SmallVector<Value> args{src};
1389 Value result =
1390 createDeviceFunctionCall(rewriter, fnName, dstTy, argTypes, args, {},
1391 funcAttrs, op.getOperation())
1392 ->getResult(0);
1393
1394 rewriter.replaceOp(op, result);
1395 } else if (dstEtype == TruncfDstElemTypes::F8) { // Float8E4M3FNType
1396 // Use charN __builtin_IB_hftohf8_N(halfN)
1397 std::string fnName = "__builtin_IB_hftohf8_" + lenSuffix;
1398 SmallVector<Type> argTypes{src.getType()};
1399 SmallVector<Value> args{src};
1400 Value result =
1401 createDeviceFunctionCall(rewriter, fnName, dstTy, argTypes, args, {},
1402 funcAttrs, op.getOperation())
1403 ->getResult(0);
1404
1405 rewriter.replaceOp(op, result);
1406 } else {
1407 return rewriter.notifyMatchFailure(
1408 op, "Unsupported src, dst element type pair.");
1409 }
1410 return success();
1411 }
1412};
1413
1414class ExtfToOCLPattern : public OpConversionPattern<ExtfOp> {
1415 using OpConversionPattern::OpConversionPattern;
1416 LogicalResult
1417 matchAndRewrite(ExtfOp op, ExtfOp::Adaptor adaptor,
1418 ConversionPatternRewriter &rewriter) const override {
1419 // `xevm.extf` is the inverse of `xevm.truncf`. Supported source and result
1420 // types are restricted for now, mirroring the truncf lowering.
1421 auto srcEtype = op.getSrcEtype().getEtype();
1422 auto dstEtype = op.getDstEtype().getEtype();
1423 // The source is scalar only where the packed values fit in one byte, which
1424 // SPIR-V spells as a scalar rather than a one element vector.
1425 Type srcTy = op.getSrc().getType();
1426 // Scalar dst is not supported until usage case become clear.
1427 auto vecDstTy = dyn_cast<VectorType>(op.getDst().getType());
1428 if (!vecDstTy)
1429 return rewriter.notifyMatchFailure(op, "Scalar dst is not supported.");
1430 // As for truncf, one builtin exists per SPIR-V vector length.
1431 int64_t numElements = vecDstTy.getNumElements();
1432 if (!isSupportedSPIRVVectorLength(numElements))
1433 return rewriter.notifyMatchFailure(
1434 op, "dst vector length must be 2, 3, 4, 8 or 16");
1435 Location loc = op.getLoc();
1436 Value src = op.getSrc();
1437 auto memAttr = rewriter.getAttr<LLVM::MemoryEffectsAttr>(
1438 /*other=*/LLVM::ModRefInfo::NoModRef,
1439 /*argMem=*/LLVM::ModRefInfo::NoModRef,
1440 /*inaccessibleMem=*/LLVM::ModRefInfo::NoModRef,
1441 /*errnoMem=*/LLVM::ModRefInfo::NoModRef,
1442 /*targetMem0=*/LLVM::ModRefInfo::NoModRef,
1443 /*targetMem1=*/LLVM::ModRefInfo::NoModRef);
1444 auto funcAttrs = convergentNoUnwindWillReturnAttrs;
1445 funcAttrs.memEffectsAttr = memAttr;
1446
1447 // Handle the case where src type is fp4 (e2m1) first.
1448 if (srcEtype == ExtfSrcElemTypes::E2M1) {
1449 // Two fp4 values are packed per source byte, and one builtin exists per
1450 // source byte count:
1451 // uint16 __builtin_IB_shfl_idx4_lut(int lut_index)
1452 // uint __builtin_IB_shfl_idx4_to_fp16_packed(uint16 lut, char src)
1453 // uintN __builtin_IB_shfl_idx4_to_fp16_N_packed(uint16 lut, charN src)
1454 // Each returns one dword, holding two f16/bf16 values, per source byte.
1455 // The lookup table selects the target format:
1456 // 7 = e2m1 -> f16, 5 = e2m1 -> bf16.
1457 //
1458 // A byte is the conversion granularity, so an odd length reads one spare
1459 // value that is dropped afterwards. The op verifier has already checked
1460 // that the source is exactly as wide as those bytes.
1461 int64_t numBytes = llvm::divideCeil(numElements, 2);
1462 constexpr int kLutE2M1ToF16 = 7;
1463 constexpr int kLutE2M1ToBF16 = 5;
1464 int lutIndex =
1465 (dstEtype == ExtfDstElemTypes::F16) ? kLutE2M1ToF16 : kLutE2M1ToBF16;
1466 Value lutIdx = LLVM::ConstantOp::create(rewriter, loc,
1467 rewriter.getI32Type(), lutIndex);
1468 Type lutTy = VectorType::get(16, rewriter.getI32Type());
1469 Value lut =
1470 createDeviceFunctionCall(rewriter, "__builtin_IB_shfl_idx4_lut",
1471 lutTy, {lutIdx.getType()}, {lutIdx}, {},
1472 funcAttrs, op.getOperation())
1473 ->getResult(0);
1474 // A single byte is passed as a bare i8, and one dword returned as a bare
1475 // i32, rather than as one element vectors SPIR-V has no type for.
1476 Type i8Ty = rewriter.getI8Type();
1477 Type i32Ty = rewriter.getI32Type();
1478 std::string fnName = "__builtin_IB_shfl_idx4_to_fp16_";
1479 Type argTy, packedResTy;
1480 if (numBytes == 1) {
1481 argTy = i8Ty;
1482 packedResTy = i32Ty;
1483 } else {
1484 fnName += std::to_string(numBytes) + "_";
1485 argTy = VectorType::get(numBytes, i8Ty);
1486 packedResTy = VectorType::get(numBytes, i32Ty);
1487 }
1488 fnName += "packed";
1489 SmallVector<Type> convArgTypes{lut.getType(), argTy};
1490 SmallVector<Value> convArgs{lut, castIfNeeded(rewriter, loc, argTy, src)};
1491 Value result =
1492 createDeviceFunctionCall(rewriter, fnName, packedResTy, convArgTypes,
1493 convArgs, {}, funcAttrs, op.getOperation())
1494 ->getResult(0);
1495 // The builtin returns the f16/bf16 bits packed as i32, bitcast to the
1496 // f16/bf16 dst type and drop the padding an odd length produced.
1497 Type wideTy = VectorType::get(numBytes * 2, vecDstTy.getElementType());
1498 result = LLVM::BitcastOp::create(rewriter, loc, wideTy, result);
1499 result = takeLeadingElements(rewriter, loc, result, numElements);
1500 rewriter.replaceOp(op, result);
1501 return success();
1502 }
1503
1504 // Handle the case where src type is fp8 (bf8/hf8). One fp8 value per source
1505 // byte, so source and destination lengths match.
1506 auto vecSrcTy = dyn_cast<VectorType>(srcTy);
1507 if (!vecSrcTy || vecSrcTy.getNumElements() != numElements)
1508 return rewriter.notifyMatchFailure(
1509 op, "fp8 src and dst must have the same number of elements");
1510 std::string lenSuffix = std::to_string(numElements);
1511
1512 // Step 1: Extend fp8 (bf8/hf8) to F16.
1513 // bf8 -> half: halfN __builtin_IB_bf8tohf_N(charN)
1514 // hf8 -> half: halfN __builtin_IB_hf8tohf_N(charN)
1515 std::string fnName = (srcEtype == ExtfSrcElemTypes::BF8)
1516 ? "__builtin_IB_bf8tohf_"
1517 : "__builtin_IB_hf8tohf_";
1518 fnName += lenSuffix;
1519 Type f16Ty = VectorType::get(vecSrcTy.getShape(), rewriter.getF16Type());
1520 SmallVector<Type> argTypes{src.getType()};
1521 SmallVector<Value> args{src};
1522 Value result =
1523 createDeviceFunctionCall(rewriter, fnName, f16Ty, argTypes, args, {},
1524 funcAttrs, op.getOperation())
1525 ->getResult(0);
1526
1527 // When the destination is F16, we are done.
1528 if (dstEtype == ExtfDstElemTypes::F16) {
1529 rewriter.replaceOp(op, result);
1530 return success();
1531 }
1532
1533 // BF16 destination needs some postprocessing.
1534 // First extend F16 to F32 and then truncate to BF16.
1535 // Step 2: Extend to F32.
1536 // Use floatN convert_floatN(halfN)
1537 std::string convFnName = "convert_float" + lenSuffix;
1538 SmallVector<Type> convArgTypes{result.getType()};
1539 SmallVector<Value> convArgs{result};
1540 convFnName = mangle(convFnName, convArgTypes);
1541 Type f32Ty = VectorType::get(vecSrcTy.getShape(), rewriter.getF32Type());
1542 result =
1543 createDeviceFunctionCall(rewriter, convFnName, f32Ty, convArgTypes,
1544 convArgs, {}, funcAttrs, op.getOperation())
1545 ->getResult(0);
1546 // Step 3: Truncate F32 to BF16.
1547 // Use shortN __builtin_IB_ftobf_N(floatN)
1548 std::string ftobfFnName = "__builtin_IB_ftobf_" + lenSuffix;
1549 SmallVector<Type> ftobfArgTypes{result.getType()};
1550 SmallVector<Value> ftobfArgs{result};
1551 Type i16Ty = VectorType::get(vecSrcTy.getShape(), rewriter.getI16Type());
1552 result =
1553 createDeviceFunctionCall(rewriter, ftobfFnName, i16Ty, ftobfArgTypes,
1554 ftobfArgs, {}, funcAttrs, op.getOperation())
1555 ->getResult(0);
1556 // The builtin returns the bf16 bits as i16, bitcast to the bf16 dst type.
1557 result = LLVM::BitcastOp::create(rewriter, op.getLoc(), vecDstTy, result);
1558 rewriter.replaceOp(op, result);
1559 return success();
1560 }
1561};
1562
1563class MMAMxToOCLPattern : public OpConversionPattern<MMAMxOp> {
1564 using OpConversionPattern::OpConversionPattern;
1565 LogicalResult
1566 matchAndRewrite(MMAMxOp op, MMAMxOp::Adaptor adaptor,
1567 ConversionPatternRewriter &rewriter) const override {
1568 if (!op.getC()) {
1569 return rewriter.notifyMatchFailure(op, "OCL requires C operand");
1570 }
1571 auto precisionC = op.getTypes().getC();
1572 auto precisionD = op.getTypes().getD();
1573 if (precisionC != precisionD) {
1574 return rewriter.notifyMatchFailure(op, "type of C and D need to match");
1575 }
1576
1577 constexpr uint32_t bitWidthPackedA{16};
1578 constexpr uint32_t bitWidthPackedB{32};
1579 auto loc = op.getLoc();
1580
1581 auto castIfNeeded = [&](Value val, Type packedType) -> Value {
1582 VectorType origTy = cast<VectorType>(val.getType());
1583 const uint32_t vecBitSize =
1584 origTy.getNumElements() *
1585 origTy.getElementType().getIntOrFloatBitWidth();
1586 VectorType newTy = VectorType::get(
1587 vecBitSize / packedType.getIntOrFloatBitWidth(), packedType);
1588 if (origTy != newTy)
1589 val = LLVM::BitcastOp::create(rewriter, loc, newTy, val);
1590 return val;
1591 };
1592
1593 Value a = op.getA();
1594 Type packedAType = (op.getTypes().getA() == xevm::ElemType::TF32)
1595 ? cast<Type>(rewriter.getF32Type())
1596 : rewriter.getIntegerType(bitWidthPackedA);
1597 a = castIfNeeded(a, packedAType);
1598
1599 Value b = op.getB();
1600 Type packedBType = (op.getTypes().getB() == xevm::ElemType::TF32)
1601 ? cast<Type>(rewriter.getF32Type())
1602 : rewriter.getIntegerType(bitWidthPackedB);
1603 b = castIfNeeded(b, packedBType);
1604
1605 Value c = op.getC();
1606 VectorType cOrigTy = cast<VectorType>(c.getType());
1607 VectorType resOrigTy = cast<VectorType>(op->getResultTypes()[0]);
1608 assert(cOrigTy == resOrigTy && "Accumulator and result type mismatch");
1609 // OCL builtins encode bfloat16 as int16
1610 VectorType cTy =
1611 cOrigTy.getElementType().isBF16()
1612 ? VectorType::get(cOrigTy.getShape(), rewriter.getIntegerType(16))
1613 : cOrigTy;
1614 VectorType resTy = cTy;
1615 if (cOrigTy != cTy)
1616 c = LLVM::BitcastOp::create(rewriter, loc, cTy, c);
1617
1618 std::string fnName =
1619 llvm::formatv("__builtin_IB_sub_group16_bdpas_{0}_{1}_{2}_{3}_8_8",
1620 builtinElemType(op.getTypes().getD()),
1621 builtinElemType(op.getTypes().getC()),
1622 builtinElemType(op.getTypes().getA()),
1623 builtinElemType(op.getTypes().getB()))
1624 .str();
1625 auto scaleA = op.getScaleA();
1626 auto scaleB = op.getScaleB();
1627 SmallVector<Type> argTypes{cTy, a.getType(), b.getType(), scaleA.getType(),
1628 scaleB.getType()};
1629 SmallVector<Value> args{c, a, b, scaleA, scaleB};
1630
1631 auto memAttr = rewriter.getAttr<LLVM::MemoryEffectsAttr>(
1632 /*other=*/LLVM::ModRefInfo::NoModRef,
1633 /*argMem=*/LLVM::ModRefInfo::NoModRef,
1634 /*inaccessibleMem=*/LLVM::ModRefInfo::NoModRef,
1635 /*errnoMem=*/LLVM::ModRefInfo::NoModRef,
1636 /*targetMem0=*/LLVM::ModRefInfo::NoModRef,
1637 /*targetMem1=*/LLVM::ModRefInfo::NoModRef);
1638 auto funcAttrs = convergentNoUnwindWillReturnAttrs;
1639 funcAttrs.memEffectsAttr = memAttr;
1640 Value result =
1641 createDeviceFunctionCall(rewriter, fnName, resTy, argTypes, args, {},
1642 funcAttrs, op.getOperation())
1643 ->getResult(0);
1644
1645 if (resOrigTy != resTy)
1646 result = LLVM::BitcastOp::create(rewriter, loc, resOrigTy, result);
1647
1648 rewriter.replaceOp(op, result);
1649 return success();
1650 }
1651};
1652
1653// Lowers `xevm.bitcast_shuffle` to a call to the IGC intrinsic
1654// `llvm.genx.GenISA.SubgroupBitcastShuffle`, which is overloaded on both the
1655// result and the operand type. E.g. a `vector<4xi8>` -> `vector<2xi16>` shuffle
1656// becomes a call to
1657// `llvm.genx.GenISA.SubgroupBitcastShuffle.v2i16.v4i8`.
1658//
1659// Only integer types reach here: the op accepts nothing else, so a producer
1660// holding floating point data bitcasts it to a same-width integer beforehand.
1661class BitcastShuffleToGenISAPattern
1662 : public OpConversionPattern<BitcastShuffleOp> {
1663 using OpConversionPattern::OpConversionPattern;
1664 LogicalResult
1665 matchAndRewrite(BitcastShuffleOp op, BitcastShuffleOp::Adaptor adaptor,
1666 ConversionPatternRewriter &rewriter) const override {
1667 Type srcTy = op.getSrc().getType();
1668 Type resTy = op.getRes().getType();
1669
1670 std::string fnName = "llvm.genx.GenISA.SubgroupBitcastShuffle." +
1671 getGenISATypeMangling(resTy) + "." +
1672 getGenISATypeMangling(srcTy);
1673
1674 Value result = createDeviceFunctionCall(
1675 rewriter, fnName, resTy, {srcTy}, {adaptor.getSrc()}, {},
1676 convergentNoUnwindWillReturnAttrs, op.getOperation())
1677 ->getResult(0);
1678
1679 rewriter.replaceOp(op, result);
1680 return success();
1681 }
1682};
1683
1684class AllocaToGlobalPattern : public OpConversionPattern<LLVM::AllocaOp> {
1685 using OpConversionPattern::OpConversionPattern;
1686 LogicalResult
1687 matchAndRewrite(LLVM::AllocaOp op, LLVM::AllocaOp::Adaptor adaptor,
1688 ConversionPatternRewriter &rewriter) const override {
1689 auto ptrType = cast<LLVM::LLVMPointerType>(op.getType());
1690 auto addrSpace = ptrType.getAddressSpace();
1691 if (addrSpace != 3)
1692 return failure();
1693 auto symTable = op->getParentWithTrait<OpTrait::SymbolTable>();
1694 if (!symTable)
1695 return failure();
1696 Block *moduleBody;
1697 if (ModuleOp mod = dyn_cast<ModuleOp>(*symTable)) {
1698 moduleBody = mod.getBody();
1699 } else if (gpu::GPUModuleOp gpuMod =
1700 dyn_cast<gpu::GPUModuleOp>(*symTable)) {
1701 moduleBody = gpuMod.getBody();
1702 } else {
1703 return failure();
1704 }
1705 auto val = op.getArraySize();
1706 APInt cst;
1707 if (!matchPattern(val, m_ConstantInt(&cst)))
1708 return failure();
1709 auto loc = op.getLoc();
1710 auto globalType = LLVM::LLVMArrayType::get(
1711 rewriter.getContext(), op.getElemType(), cst.getZExtValue());
1712 LLVM::GlobalOp globalVar;
1713 {
1714 OpBuilder::InsertionGuard guard(rewriter);
1715 rewriter.setInsertionPointToStart(moduleBody);
1716 auto alignment = op.getAlignment();
1717 globalVar = LLVM::GlobalOp::create(
1718 rewriter, loc, globalType, /*isConstant=*/false,
1719 /*linkage=*/LLVM::Linkage::Internal,
1720 /*name=*/std::string("__global_alloca_") +
1721 std::to_string(getNextGlobalIdx()),
1722 /*value=*/Attribute(),
1723 /*alignment=*/alignment ? *alignment : 0, /*addrSpace=*/addrSpace);
1724 }
1725 rewriter.replaceOpWithNewOp<LLVM::AddressOfOp>(op, globalVar);
1726 return success();
1727 }
1728
1729private:
1730 static unsigned getNextGlobalIdx() {
1731 static unsigned globalIdx = 0;
1732 return globalIdx++;
1733 }
1734};
1735
1736// Checks if shufflevector is used as a way to extract a contiguous slice
1737// from a vector.
1738// - source vector V2 is either the same as V1, or a poison/undef value. In
1739// both cases the mask can only meaningfully address elements of V1, which
1740// the mask checks below enforce.
1741// - mask size is not greater than the source vector size
1742// - mask values represent a sequence of consecutive increasing numbers
1743// that stay in bounds of the source vector when used for indexing.
1744static bool isExtractingContiguousSlice(LLVM::ShuffleVectorOp op) {
1745 if (op.getV1() != op.getV2() &&
1746 !isa_and_present<LLVM::PoisonOp, LLVM::UndefOp>(
1747 op.getV2().getDefiningOp()))
1748 return false;
1749 auto maskAttr = op.getMask();
1750 int64_t maskSize = static_cast<int64_t>(maskAttr.size());
1751 int64_t sourceSize = op.getV1().getType().getNumElements();
1752 if (maskSize > sourceSize)
1753 return false;
1754 int64_t firstIndex = maskAttr[0];
1755 if (firstIndex < 0 || firstIndex >= sourceSize)
1756 return false;
1757 for (int64_t i = 1; i < maskSize; ++i) {
1758 int64_t index = maskAttr[i];
1759 if (index != firstIndex + i)
1760 return false;
1761 if (index >= sourceSize)
1762 return false;
1763 }
1764 return true;
1765}
1766
1767// Clones `op` with new operands and result types, preserving its attributes
1768// and properties (e.g. an `icmp` predicate or fastmath flags). Same idiom used
1769// by the Vector dialect's unroll patterns.
1771 Location loc, Operation *op,
1772 ArrayRef<Value> operands,
1773 ArrayRef<Type> resultTypes) {
1774 OperationState state(loc, op->getName().getStringRef(), operands, resultTypes,
1775 op->getRawDictionaryAttrs().getValue());
1776 state.propertiesAttr = op->getPropertiesAsAttribute();
1777 return rewriter.create(state);
1778}
1779
1780// Input vector of a shuffle vector op extracting a contiguous slice is an
1781// illegal vector in SPIRV kernel if the vector size is > 16 elements.
1782// To legalize this case, keep applying the following transformations until no
1783// more match:
1784// 1. keep hoisting the shuffle vector op past unary element-wise operations
1785// start with fpext, fptrunc and bitcast for now.
1786// 2. merge with another shuffle vector op
1787// 3. merge with load as a smaller load
1788// 4. hoist past n-ary element-wise operations (e.g. fdiv, select, icmp) by
1789// slicing every same-length vector operand with the same mask and cloning
1790// the op at the narrow width.
1791class HandleVectorExtractPattern
1792 : public OpRewritePattern<LLVM::ShuffleVectorOp> {
1793 using OpRewritePattern<LLVM::ShuffleVectorOp>::OpRewritePattern;
1794
1795 void initialize() { setHasBoundedRewriteRecursion(); }
1796
1797 LogicalResult matchAndRewrite(LLVM::ShuffleVectorOp op,
1798 PatternRewriter &rewriter) const override {
1799
1800 if (!isExtractingContiguousSlice(op))
1801 return failure();
1802
1803 auto mask = op.getMask();
1804 auto loc = op.getLoc();
1805 auto ty = op.getType();
1806 // Check source operand to determine rewrite pattern.
1807 auto src = op.getV1();
1808 // 1. Hoist past unary element-wise operations
1809 if (auto srcOp = src.getDefiningOp()) {
1810 if (isa<LLVM::FPExtOp>(srcOp) || isa<LLVM::FPTruncOp>(srcOp)) {
1811 Value srcInput = srcOp->getOperand(0);
1812 // Create new shuffle vector op with unary input as source.
1813 auto srcVecTy = dyn_cast<VectorType>(srcInput.getType());
1814 if (!srcVecTy)
1815 return failure();
1816 auto newShuffleVecTy =
1817 VectorType::get(mask.size(), srcVecTy.getElementType());
1818 auto newShuffle = LLVM::ShuffleVectorOp::create(
1819 rewriter, loc, newShuffleVecTy, srcInput, srcInput, mask);
1820 // Create new unary op with new shuffle as input.
1821 Value newUnaryOp;
1822 if (isa<LLVM::FPExtOp>(srcOp)) {
1823 newUnaryOp = LLVM::FPExtOp::create(rewriter, loc, ty, newShuffle);
1824 } else {
1825 newUnaryOp = LLVM::FPTruncOp::create(rewriter, loc, ty, newShuffle);
1826 }
1827 rewriter.replaceOp(op, newUnaryOp);
1828 } else if (isa<LLVM::BitcastOp>(srcOp)) {
1829 Value srcInput = srcOp->getOperand(0);
1830 // Create new shuffle vector op with unary input as source. A bitcast
1831 // from a scalar has no slice to rewrite in terms of.
1832 auto srcInputVecTy = dyn_cast<VectorType>(srcInput.getType());
1833 auto srcResVecTy = dyn_cast<VectorType>(srcOp->getResult(0).getType());
1834 if (!srcInputVecTy || !srcResVecTy)
1835 return failure();
1836 auto srcInputSize = srcInputVecTy.getNumElements();
1837 auto srcResSize = srcResVecTy.getNumElements();
1838 auto maskSize = static_cast<int32_t>(mask.size());
1839 if (srcInputSize > srcResSize) {
1840 return failure();
1841 }
1842 if (srcResSize % srcInputSize != 0) {
1843 return failure();
1844 }
1845 auto maskScale = srcResSize / srcInputSize;
1846 // Storage for the rescaled mask. `mask` is an ArrayRef that gets
1847 // rebound to this buffer below, so it has to outlive the `if` that
1848 // fills it - otherwise the uses after the `if` read a destroyed
1849 // SmallVector.
1850 SmallVector<int32_t> newMask;
1851 if (maskScale != 1) {
1852 // The slice has to start at, and cover, whole source elements to be
1853 // expressible in terms of the bitcast source.
1854 if (mask[0] % maskScale != 0 || maskSize % maskScale != 0) {
1855 return failure();
1856 }
1857 // Create a new mask that maps to the source vector
1858 int32_t newMaskSize = maskSize / maskScale;
1859 int32_t maskStart = mask[0] / maskScale;
1860 for (int32_t i = 0; i < newMaskSize; ++i) {
1861 newMask.push_back(maskStart + i);
1862 }
1863 mask = newMask;
1864 }
1865 auto newShuffleVecTy = VectorType::get(
1866 static_cast<int64_t>(mask.size()), srcInputVecTy.getElementType());
1867 auto newShuffle = LLVM::ShuffleVectorOp::create(
1868 rewriter, loc, newShuffleVecTy, srcInput, srcInput, mask);
1869 // Create new unary op with new shuffle as input.
1870 auto newBitcast =
1871 LLVM::BitcastOp::create(rewriter, loc, ty, newShuffle);
1872 rewriter.replaceOp(op, newBitcast);
1873 } else if (isa<LLVM::ShuffleVectorOp>(srcOp)) {
1874 // 2. Merge with source shuffle vector op if, the source op is
1875 // also extracting a contigous slice and create a new
1876 // shuffle vector op directly from the source of
1877 // the first shuffle.
1878 auto srcShuffle = cast<LLVM::ShuffleVectorOp>(srcOp);
1879 if (!isExtractingContiguousSlice(srcShuffle))
1880 return failure();
1881 auto srcMask = srcShuffle.getMask();
1882 SmallVector<int32_t> combinedMask;
1883 for (auto index : mask) {
1884 combinedMask.push_back(srcMask[index]);
1885 }
1886 auto newShuffle = LLVM::ShuffleVectorOp::create(
1887 rewriter, loc, ty, srcShuffle.getV1(), srcShuffle.getV1(),
1888 DenseI32ArrayAttr::get(rewriter.getContext(), combinedMask));
1889 rewriter.replaceOp(op, newShuffle);
1890 } else if (isa<LLVM::LoadOp>(srcOp)) {
1891 // 3. Merge with load as a smaller load
1892 auto loadOp = cast<LLVM::LoadOp>(srcOp);
1893 auto loadPtr = loadOp.getAddr();
1894 auto loadAddrSpace = loadPtr.getType().getAddressSpace();
1895 if (loadAddrSpace != 0)
1896 return failure();
1897 auto loadTy = dyn_cast<VectorType>(loadOp.getType());
1898 if (!loadTy)
1899 return failure();
1900 auto elemTy = loadTy.getElementType();
1901 auto firstIndex = mask[0];
1902 auto newVecTy = VectorType::get(mask.size(), elemTy);
1903 // GEPOp is needed if first index is not zero
1904 if (firstIndex) {
1905 auto newPtr = LLVM::GEPOp::create(
1906 rewriter, loc,
1907 LLVM::LLVMPointerType::get(rewriter.getContext(), loadAddrSpace),
1908 elemTy, loadPtr, ArrayRef<LLVM::GEPArg>{firstIndex});
1909 auto newLoad = LLVM::LoadOp::create(rewriter, loc, newVecTy, newPtr);
1910 rewriter.replaceOp(op, newLoad);
1911 } else {
1912 auto newLoad = LLVM::LoadOp::create(rewriter, loc, newVecTy, loadPtr);
1913 rewriter.replaceOp(op, newLoad);
1914 }
1915 } else if (isMemoryEffectFree(srcOp) && srcOp->getNumResults() == 1 &&
1916 srcOp->getNumOperands() >= 1 &&
1917 llvm::all_of(srcOp->getOperands(), [&](Value operand) {
1918 auto operandTy = dyn_cast<VectorType>(operand.getType());
1919 auto srcTy = cast<VectorType>(src.getType());
1920 return operandTy && operandTy.getRank() == 1 &&
1921 operandTy.getNumElements() == srcTy.getNumElements();
1922 })) {
1923 // 4. Hoist past an n-ary element-wise op: slice each same-length vector
1924 // operand with the same contiguous mask and clone the op narrow.
1925 // The freshly created operand shuffles keep the greedy driver going,
1926 // hoisting them further up (or merging via case 2 / into loads via
1927 // case 3) until they reach operands that are already legal width.
1928 SmallVector<Value> newOperands;
1929 newOperands.reserve(srcOp->getNumOperands());
1930 for (Value operand : srcOp->getOperands()) {
1931 auto operandTy = cast<VectorType>(operand.getType());
1932 auto sliceTy =
1933 VectorType::get(mask.size(), operandTy.getElementType());
1934 newOperands.push_back(LLVM::ShuffleVectorOp::create(
1935 rewriter, loc, sliceTy, operand, operand, mask));
1936 }
1937 Operation *newOp =
1938 cloneOpWithOperandsAndTypes(rewriter, loc, srcOp, newOperands, ty);
1939 rewriter.replaceOp(op, newOp->getResult(0));
1940 } else {
1941 return failure();
1942 }
1943 } else {
1944 // No defining op (e.g. function argument): nothing to hoist/merge.
1945 return failure();
1946 }
1947 return success();
1948 }
1949};
1950
1951//===----------------------------------------------------------------------===//
1952// Pass Definition
1953//===----------------------------------------------------------------------===//
1954
1955struct ConvertXeVMToLLVMPass
1956 : public impl::ConvertXeVMToLLVMPassBase<ConvertXeVMToLLVMPass> {
1957 using Base::Base;
1958
1959 void getDependentDialects(DialectRegistry &registry) const override {
1960 registry.insert<LLVM::LLVMDialect, XeVMDialect>();
1961 }
1962
1963 void runOnOperation() override {
1964 ConversionTarget target(getContext());
1965 RewritePatternSet patterns(&getContext());
1967 if (failed(applyPartialConversion(getOperation(), target,
1968 std::move(patterns))))
1969 signalPassFailure();
1970
1971 // Apply in-dialect lowerings to handle illegal vectors
1972 {
1973 RewritePatternSet vectorPatterns(&getContext());
1974 vectorPatterns.add<HandleVectorExtractPattern>(&getContext());
1975 GreedyRewriteConfig config{};
1976 // folding can remove ops with temporary attributes used to
1977 // represent LLVM metadata, so disable it here.
1978 // Effectively just this single pattern is applied without any
1979 // op folding patterns from dialects.
1980 config.enableFolding(false);
1981 // config.setMaxIterations(GreedyRewriteConfig::kNoLimit);
1982 // config.setMaxNumRewrites(GreedyRewriteConfig::kNoLimit);
1983 (void)applyPatternsGreedily(getOperation(), std::move(vectorPatterns),
1984 config);
1985 }
1986 }
1987};
1988} // namespace
1989
1990//===----------------------------------------------------------------------===//
1991// Pattern Population
1992//===----------------------------------------------------------------------===//
1993
1994void ::mlir::populateXeVMToLLVMConversionPatterns(ConversionTarget &target,
1995 RewritePatternSet &patterns) {
1996 // some LLVM operations need to be converted.
1997 target.addDynamicallyLegalDialect<LLVM::LLVMDialect>([](Operation *op) {
1998 // llvm alloca op with addrspace 3 for OpenCL (Workgroup) is not handled
1999 // properly by SPIRV backend. It needs to be rewritten as a sequence with
2000 // llvm global.
2001 if (isa<LLVM::AllocaOp>(op)) {
2002 LLVM::AllocaOp aOp = cast<LLVM::AllocaOp>(op);
2003 LLVM::LLVMPointerType pTy = cast<LLVM::LLVMPointerType>(aOp.getType());
2004 auto addrSpace = pTy.getAddressSpace();
2005 return addrSpace != 3;
2006 }
2007 // cache_control attribute should be converted.
2008 return !op->hasDiscardableAttr("cache_control");
2009 });
2010 target.addIllegalDialect<XeVMDialect>();
2011 patterns.add<LoadStorePrefetchToOCLPattern<BlockLoad2dOp>,
2012 LoadStorePrefetchToOCLPattern<BlockStore2dOp>,
2013 LoadStorePrefetchToOCLPattern<BlockPrefetch2dOp>,
2014 MMAToOCLPattern, MemfenceToOCLPattern, PrefetchToOCLPattern,
2015 LLVMLoadStoreToOCLPattern<LLVM::LoadOp>,
2016 LLVMLoadStoreToOCLPattern<LLVM::StoreOp>,
2017 BlockLoadStore1DToOCLPattern<BlockLoadOp>,
2018 BlockLoadStore1DToOCLPattern<BlockStoreOp>,
2019 LaunchConfigOpToOCLPattern<WorkitemIdXOp>,
2020 LaunchConfigOpToOCLPattern<WorkitemIdYOp>,
2021 LaunchConfigOpToOCLPattern<WorkitemIdZOp>,
2022 LaunchConfigOpToOCLPattern<WorkgroupDimXOp>,
2023 LaunchConfigOpToOCLPattern<WorkgroupDimYOp>,
2024 LaunchConfigOpToOCLPattern<WorkgroupDimZOp>,
2025 LaunchConfigOpToOCLPattern<WorkgroupIdXOp>,
2026 LaunchConfigOpToOCLPattern<WorkgroupIdYOp>,
2027 LaunchConfigOpToOCLPattern<WorkgroupIdZOp>,
2028 LaunchConfigOpToOCLPattern<GridDimXOp>,
2029 LaunchConfigOpToOCLPattern<GridDimYOp>,
2030 LaunchConfigOpToOCLPattern<GridDimZOp>,
2031 SubgroupOpWorkitemOpToOCLPattern<LaneIdOp>,
2032 SubgroupOpWorkitemOpToOCLPattern<SubgroupIdOp>,
2033 SubgroupOpWorkitemOpToOCLPattern<SubgroupSizeOp>,
2034 TruncfToOCLPattern, ExtfToOCLPattern, MMAMxToOCLPattern,
2035 BitcastShuffleToGenISAPattern, AllocaToGlobalPattern>(
2036 patterns.getContext());
2037}
return success()
LogicalResult initialize(unsigned origNumLoops, ArrayRef< ReassociationIndices > foldedIterationDims)
b
Return true if permutation is a valid permutation of the outer_dims_perm (case OuterOrInnerPerm::Oute...
ArrayAttr()
b getContext())
static Operation * cloneOpWithOperandsAndTypes(RewriterBase &rewriter, Location loc, Operation *op, ArrayRef< Value > operands, ArrayRef< Type > resultTypes)
Attributes are known-constant values of operations.
Definition Attributes.h:25
MLIRContext * getContext() const
Definition Builders.h:56
This class defines the main interface for locations in MLIR and acts as a non-nullable wrapper around...
Definition Location.h:76
NamedAttribute represents a combination of a name and an Attribute value.
Definition Attributes.h:164
RAII guard to reset the insertion point of the builder when destroyed.
Definition Builders.h:351
Operation * create(const OperationState &state)
Creates an operation given the fields represented as an OperationState.
Definition Builders.cpp:466
A trait used to provide symbol table functionalities to a region operation.
StringRef getStringRef() const
Return the name of this operation. This always succeeds.
Operation is the basic unit of execution within MLIR.
Definition Operation.h:87
Region & getRegion(unsigned index)
Returns the region held by this operation at position 'index'.
Definition Operation.h:738
OpResult getResult(unsigned idx)
Get the 'idx'th result of this operation.
Definition Operation.h:432
Operation * getParentWithTrait()
Returns the closest surrounding parent operation with trait Trait.
Definition Operation.h:273
Location getLoc()
The source location the operation was defined or derived from.
Definition Operation.h:240
DictionaryAttr getRawDictionaryAttrs()
Return all attributes that are not stored as properties.
Definition Operation.h:561
Attribute getPropertiesAsAttribute()
Return the properties converted to an attribute.
OperationName getName()
The name of an operation is the key identifier for it.
Definition Operation.h:115
Block & front()
Definition Region.h:65
MLIRContext * getContext() const
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 coordinates the application of a rewrite on a set of IR, providing a way for clients to tr...
virtual void replaceOp(Operation *op, ValueRange newValues)
Replace the results of the given (original) operation with the specified list of values (replacements...
Instances of the Type class are uniqued, have an immutable identifier and an optional mutable compone...
Definition Types.h:74
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
Operation * getDefiningOp() const
If this value is the result of an operation, return the operation that defines it.
Definition Value.cpp:18
static DenseArrayAttrImpl get(MLIRContext *context, ArrayRef< int32_t > content)
FailureOr< LLVM::LLVMFuncOp > lookupOrCreateFn(OpBuilder &b, Operation *moduleOp, StringRef name, ArrayRef< Type > paramTypes={}, Type resultType={}, bool isVarArg=false, bool isReserved=false, SymbolTableCollection *symbolTables=nullptr)
Create a FuncOp with signature resultType(paramTypes) and name name`.
detail::InFlightRemark failed(Location loc, RemarkOpts opts)
Report an optimization remark that failed.
Definition Remarks.h:734
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
detail::constant_int_value_binder m_ConstantInt(IntegerAttr::ValueType *bind_value)
Matches a constant holding a scalar/vector/tensor integer (splat) and writes the integer value to bin...
Definition Matchers.h:527
LogicalResult applyPatternsGreedily(Region &region, const FrozenRewritePatternSet &patterns, GreedyRewriteConfig config=GreedyRewriteConfig(), bool *changed=nullptr)
Rewrite ops in the given region, which must be isolated from above, by repeatedly applying the highes...
bool isMemoryEffectFree(Operation *op)
Returns true if the given operation is free of memory effects.
void populateXeVMToLLVMConversionPatterns(ConversionTarget &target, RewritePatternSet &patterns)
llvm::TypeSwitch< T, ResultT > TypeSwitch
Definition LLVM.h:139
OpRewritePattern is a wrapper around RewritePattern that allows for matching and rewriting against an...
This represents an operation in an abstracted form, suitable for use with the builder APIs.