18#include "llvm/ADT/ArrayRef.h"
19#include "llvm/ADT/STLExtras.h"
20#include "llvm/Support/FormatVariadic.h"
21#include "llvm/Support/MathExtras.h"
29#include "llvm/ADT/TypeSwitch.h"
32#define GEN_PASS_DEF_CONVERTXEVMTOLLVMPASS
33#include "mlir/Conversion/Passes.h.inc"
41struct LLVMFuncAttributeOptions {
42 bool isConvergent =
false;
43 bool isNoUnwind =
false;
44 bool isWillReturn =
false;
45 LLVM::MemoryEffectsAttr memEffectsAttr{};
47static constexpr LLVMFuncAttributeOptions noUnwindAttrs = {
48 false,
true,
false, {}};
49static constexpr LLVMFuncAttributeOptions noUnwindWillReturnAttrs = {
50 false,
true,
true, {}};
51static constexpr LLVMFuncAttributeOptions convergentNoUnwindWillReturnAttrs = {
52 true,
true,
true, {}};
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);
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()) {
66 return isUnsigned ?
"h" :
"c";
68 return isUnsigned ?
"t" :
"s";
70 return isUnsigned ?
"j" :
"i";
72 return isUnsigned ?
"m" :
"l";
74 llvm_unreachable(
"unhandled integer type");
77 .DefaultUnreachable(
"unhandled type for mangling");
82 assert((isUnsigned.empty() || isUnsigned.size() == types.size()) &&
83 "Signedness info doesn't match");
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()) {
93 if (
unsigned firstIdx = it->getSecond(); firstIdx > 0)
97 if (!type.isIntOrFloat())
98 substitutions[type] = substitutions.size();
99 os << getTypeMangling(type, isUnsigned.empty() ?
false : isUnsigned[idx]);
109std::string getGenISATypeMangling(
Type ty) {
111 .Case([](VectorType ty) -> std::string {
112 return "v" + std::to_string(ty.getNumElements()) +
113 getGenISATypeMangling(ty.getElementType());
115 .Case([](IntegerType ty) -> std::string {
116 return "i" + std::to_string(ty.getWidth());
118 .DefaultUnreachable(
"unhandled type for GenISA mangling");
121std::string builtinElemType(ElemType elemType) {
134 return stringifyElemType(elemType).str();
138static int32_t getL1CacheControl(LoadCacheControl cc) {
141 case LoadCacheControl::USE_DEFAULT:
144 case LoadCacheControl::L1C_L2UC_L3UC:
145 case LoadCacheControl::L1C_L2UC_L3C:
146 case LoadCacheControl::L1C_L2C_L3UC:
147 case LoadCacheControl::L1C_L2C_L3C:
150 case LoadCacheControl::L1S_L2UC_L3UC:
151 case LoadCacheControl::L1S_L2UC_L3C:
152 case LoadCacheControl::L1S_L2C_L3UC:
153 case LoadCacheControl::L1S_L2C_L3C:
156 case LoadCacheControl::INVALIDATE_READ:
165static int32_t getL1CacheControl(StoreCacheControl cc) {
168 case StoreCacheControl::USE_DEFAULT:
171 case StoreCacheControl::L1WT_L2UC_L3UC:
172 case StoreCacheControl::L1WT_L2UC_L3WB:
173 case StoreCacheControl::L1WT_L2WB_L3UC:
174 case StoreCacheControl::L1WT_L2WB_L3WB:
177 case StoreCacheControl::L1WB_L2UC_L3UC:
178 case StoreCacheControl::L1WB_L2WB_L3UC:
179 case StoreCacheControl::L1WB_L2UC_L3WB:
182 case StoreCacheControl::L1S_L2UC_L3UC:
183 case StoreCacheControl::L1S_L2UC_L3WB:
184 case StoreCacheControl::L1S_L2WB_L3UC:
185 case StoreCacheControl::L1S_L2WB_L3WB:
194static int32_t getL3CacheControl(LoadCacheControl cc) {
197 case LoadCacheControl::USE_DEFAULT:
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:
208 case LoadCacheControl::INVALIDATE_READ:
217static int32_t getL3CacheControl(StoreCacheControl cc) {
220 case StoreCacheControl::USE_DEFAULT:
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:
238static std::optional<LoadCacheControl> getCacheControl(PrefetchOp op) {
239 return op.getCacheControl();
242static std::optional<LoadCacheControl> getCacheControl(BlockLoad2dOp op) {
243 return op.getCacheControl();
246static std::optional<LoadCacheControl> getCacheControl(BlockLoadOp op) {
247 return op.getCacheControl();
250static std::optional<LoadCacheControl> getCacheControl(BlockPrefetch2dOp op) {
251 return op.getCacheControl();
254static std::optional<StoreCacheControl> getCacheControl(BlockStore2dOp op) {
255 return op.getCacheControl();
258static std::optional<StoreCacheControl> getCacheControl(BlockStoreOp op) {
259 return op.getCacheControl();
262static std::optional<LoadCacheControl> getCacheControl(LLVM::LoadOp op) {
263 if (op->hasDiscardableAttr(
"cache_control")) {
264 auto attr = op->getDiscardableAttrOfType<xevm::LoadCacheControlAttr>(
268 return std::optional<LoadCacheControl>(attr.getValue());
273static std::optional<StoreCacheControl> getCacheControl(LLVM::StoreOp op) {
274 if (op->hasDiscardableAttr(
"cache_control")) {
275 auto attr = op->getDiscardableAttrOfType<xevm::StoreCacheControlAttr>(
279 return std::optional<StoreCacheControl>(attr.getValue());
284template <
typename OpType>
285int32_t getL1CacheControl(OpType op) {
286 return getL1CacheControl(*getCacheControl(op));
289template <
typename OpType>
290int32_t getL3CacheControl(OpType op) {
291 return getL3CacheControl(*getCacheControl(op));
294template <
typename OpType>
295static std::optional<ArrayAttr>
296getCacheControlMetadata(ConversionPatternRewriter &rewriter, OpType op) {
297 if (!getCacheControl(op))
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>;
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 "
318 if (getL1CacheControl<OpType>(op) == -1 &&
319 getL3CacheControl<OpType>(op) == -1)
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);
330 return rewriter.getArrayAttr(combinedAttrs);
352 llvm::StringMap<bool> seen;
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.");
360 auto tokenAttr = dyn_cast<IntegerAttr>(vals[0]);
361 auto secondAttr = dyn_cast<IntegerAttr>(vals[1]);
362 auto thirdAttr = dyn_cast<IntegerAttr>(vals[2]);
364 if (!tokenAttr || !secondAttr || !thirdAttr)
370 llvm::formatv(
"{{{0}:\"{1},{2}\"}", tokenAttr.getValue().getZExtValue(),
371 secondAttr.getValue().getZExtValue(),
372 thirdAttr.getValue().getZExtValue());
375 if (!seen.insert({entry, true}).second)
378 payloads.push_back(std::move(entry));
383static std::atomic<uint64_t> globalNameCounter{0};
388static Value createMetadataStringPtr(ConversionPatternRewriter &rewriter,
390 StringRef value, StringRef nameHint) {
392 std::string strWithNull = value.str();
393 strWithNull.push_back(
'\0');
394 StringRef strRef(strWithNull.data(), strWithNull.size());
396 auto as1PtrTy = LLVM::LLVMPointerType::get(rewriter.getContext(), 1);
400 if (
auto existingGlobal = dyn_cast<LLVM::GlobalOp>(&op)) {
401 if (!existingGlobal.getSection() ||
402 *existingGlobal.getSection() !=
"llvm.metadata")
405 dyn_cast_or_null<StringAttr>(existingGlobal.getValueOrNull())) {
406 if (strAttr.getValue() == strRef) {
407 return LLVM::AddressOfOp::create(rewriter, loc, as1PtrTy,
408 existingGlobal.getSymName());
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))
424 rewriter.setInsertionPointToStart(&moduleOp->
getRegion(0).
front());
427 LLVM::GlobalOp::create(rewriter, loc, arrayType,
428 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);
437 return LLVM::AddressOfOp::create(rewriter, loc, as1PtrTy, globalName);
460static Value annotatePtrWithCacheControl(ConversionPatternRewriter &rewriter,
465 buildCacheControlPayloads(cacheControls.getValue());
466 if (payloads.empty())
469 auto ptrType = cast<LLVM::LLVMPointerType>(
ptr.getType());
470 auto as1PtrTy = LLVM::LLVMPointerType::get(rewriter.getContext(), 1);
471 auto i32Ty = rewriter.getI32Type();
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);
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();
509template <
typename OpType>
511applyCacheControlAnnotation(ConversionPatternRewriter &rewriter,
Location loc,
513 Operation *moduleOp,
unsigned ptrIdx = 0) {
514 std::optional<ArrayAttr> optCacheControls =
515 getCacheControlMetadata(rewriter, op);
516 if (!optCacheControls)
519 Value annotatedPtr = annotatePtrWithCacheControl(rewriter, loc, args[ptrIdx],
520 *optCacheControls, moduleOp);
521 args[ptrIdx] = annotatedPtr;
528static LLVM::CallOp createDeviceFunctionCall(
529 ConversionPatternRewriter &rewriter, StringRef funcName,
Type retType,
532 LLVMFuncAttributeOptions funcAttributeOptions,
Operation *op) {
534 assert(moduleOp &&
"Expecting module");
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);
546 if (funcAttributeOptions.memEffectsAttr)
547 funcOp.setMemoryEffectsAttr(funcAttributeOptions.memEffectsAttr);
549 for (
auto [idx, attrName] : paramAttrs)
550 funcOp.setArgAttr(idx, attrName, rewriter.getUnitAttr());
552 auto callOp = LLVM::CallOp::create(rewriter, loc, funcOp, args);
554 auto copyAttr = [&](StringAttr name,
Attribute attr) {
555 if (callOp->getInherentAttr(name).has_value())
556 callOp->setInherentAttr(name, attr);
558 discardableAttrs.emplace_back(name, attr);
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);
566 callOp->setDiscardableAttrs(discardableAttrs);
571static unsigned getNumOperandsPerDword(xevm::ElemType pTy) {
573 case xevm::ElemType::F32:
574 case xevm::ElemType::TF32:
576 case xevm::ElemType::BF16:
577 case xevm::ElemType::F16:
579 case xevm::ElemType::U8:
580 case xevm::ElemType::S8:
581 case xevm::ElemType::BF8:
582 case xevm::ElemType::F8:
584 case xevm::ElemType::E2M1:
585 case xevm::ElemType::U4:
586 case xevm::ElemType::S4:
589 llvm_unreachable(
"unsupported xevm::ElemType");
593class MMAToOCLPattern :
public OpConversionPattern<xevm::MMAOp> {
594 using OpConversionPattern::OpConversionPattern;
596 matchAndRewrite(xevm::MMAOp op, xevm::MMAOp::Adaptor adaptor,
597 ConversionPatternRewriter &rewriter)
const override {
599 return rewriter.notifyMatchFailure(op,
"OCL requires C operand");
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");
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");
615 if (precisionA == xevm::ElemType::S32 ||
616 precisionA == xevm::ElemType::F32) {
617 return rewriter.notifyMatchFailure(op,
"type of A cannot be S32 or F32");
619 if (precisionB == xevm::ElemType::S32 ||
620 precisionB == xevm::ElemType::F32) {
621 return rewriter.notifyMatchFailure(op,
"type of B cannot be S32 or F32");
623 constexpr uint32_t bitWidthPackedA{16};
624 constexpr uint32_t bitWidthPackedB{32};
625 auto loc = op.getLoc();
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);
635 val = LLVM::BitcastOp::create(rewriter, loc, newTy, val);
640 Type packedAType = (op.getTypes().getA() == xevm::ElemType::TF32)
641 ? cast<Type>(rewriter.getF32Type())
642 : rewriter.getIntegerType(bitWidthPackedA);
643 a = castIfNeeded(a, packedAType);
646 Type packedBType = (op.getTypes().getB() == xevm::ElemType::TF32)
647 ? cast<Type>(rewriter.getF32Type())
648 : rewriter.getIntegerType(bitWidthPackedB);
649 b = castIfNeeded(
b, packedBType);
652 VectorType cOrigTy = cast<VectorType>(c.
getType());
653 VectorType resOrigTy = cast<VectorType>(op->getResultTypes()[0]);
654 assert(cOrigTy == resOrigTy &&
"Accumulator and result type mismatch");
657 cOrigTy.getElementType().isBF16()
658 ? VectorType::get(cOrigTy.getShape(), rewriter.getIntegerType(16))
660 VectorType resTy = cTy;
662 c = LLVM::BitcastOp::create(rewriter, loc, cTy, c);
664 constexpr int32_t systolicDepth{8};
666 llvm::formatv(
"intel_sub_group_{0}_{1}_matrix_mad_k{2}",
667 stringifyElemType(op.getTypes().getA()).str(),
668 stringifyElemType(op.getTypes().getB()).str(),
670 getNumOperandsPerDword(op.getTypes().getA()))
672 SmallVector<Type> argTypes{a.
getType(),
b.getType(), cTy};
673 fnName = mangle(fnName, argTypes);
674 SmallVector<Value> args{a,
b, c};
676 auto memAttr = rewriter.getAttr<LLVM::MemoryEffectsAttr>(
677 LLVM::ModRefInfo::NoModRef,
678 LLVM::ModRefInfo::NoModRef,
679 LLVM::ModRefInfo::NoModRef,
680 LLVM::ModRefInfo::NoModRef,
681 LLVM::ModRefInfo::NoModRef,
682 LLVM::ModRefInfo::NoModRef);
683 auto funcAttrs = convergentNoUnwindWillReturnAttrs;
684 funcAttrs.memEffectsAttr = memAttr;
686 createDeviceFunctionCall(rewriter, fnName, resTy, argTypes, args, {},
687 funcAttrs, op.getOperation())
690 if (resOrigTy != resTy)
691 result = LLVM::BitcastOp::create(rewriter, loc, resOrigTy,
result);
693 rewriter.replaceOp(op,
result);
698class PrefetchToOCLPattern :
public OpConversionPattern<PrefetchOp> {
699 using OpConversionPattern::OpConversionPattern;
701 matchAndRewrite(PrefetchOp op, PrefetchOp::Adaptor adaptor,
702 ConversionPatternRewriter &rewriter)
const override {
703 auto loc = op.getLoc();
706 const std::string fnName{
"_Z8prefetchPU3AS1Kcm"};
708 LLVM::ConstantOp::create(rewriter, loc, rewriter.getI64Type(), 1);
709 SmallVector<Value> args{op.getPtr(), one};
712 applyCacheControlAnnotation(rewriter, loc, op, args, moduleOp,
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 LLVM::ModRefInfo::NoModRef,
721 LLVM::ModRefInfo::Ref,
722 LLVM::ModRefInfo::NoModRef,
723 LLVM::ModRefInfo::NoModRef,
724 LLVM::ModRefInfo::NoModRef,
725 LLVM::ModRefInfo::NoModRef);
726 funcAttr.memEffectsAttr = memAttr;
728 createDeviceFunctionCall(rewriter, fnName,
729 LLVM::LLVMVoidType::get(rewriter.getContext()),
730 argTypes, args, {}, funcAttr, op.getOperation());
731 rewriter.eraseOp(op);
736class MemfenceToOCLPattern :
public OpConversionPattern<MemfenceOp> {
737 using OpConversionPattern::OpConversionPattern;
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:
748 case xevm::AddrSpace::GLOBAL:
753 return rewriter.notifyMatchFailure(
754 op,
"Fence only supports global and shared address spaces.");
756 switch (op.getScope()) {
757 case xevm::MemScope::WORKGROUP:
760 case xevm::MemScope::DEVICE:
765 return rewriter.notifyMatchFailure(
766 op,
"Fence only supports workgroup and device memory scopes.");
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,
780 rewriter.eraseOp(op);
784template <
typename OpType>
785class LoadStorePrefetchToOCLPattern :
public OpConversionPattern<OpType> {
786 using OpConversionPattern<OpType>::OpConversionPattern;
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>;
793 auto loc = op.getLoc();
794 auto *moduleOp = op->template getParentWithTrait<OpTrait::SymbolTable>();
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();
806 auto i32Type = rewriter.getI32Type();
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};
819 applyCacheControlAnnotation(rewriter, loc, op, args, moduleOp,
822 SmallVector<Type> retTypes;
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) {
829 funcName +=
"prefetch";
830 paramAttrs = {std::make_pair(0, LLVM::LLVMDialect::getNonNullAttrName())};
831 auto memAttr = rewriter.getAttr<LLVM::MemoryEffectsAttr>(
832 LLVM::ModRefInfo::NoModRef,
833 LLVM::ModRefInfo::Ref,
834 LLVM::ModRefInfo::NoModRef,
835 LLVM::ModRefInfo::NoModRef,
836 LLVM::ModRefInfo::NoModRef,
837 LLVM::ModRefInfo::NoModRef);
838 funcAttr = noUnwindAttrs;
839 funcAttr.memEffectsAttr = memAttr;
841 auto vecElemType = vecType.getElementType();
842 auto vecElemBitWidth = vecElemType.getIntOrFloatBitWidth();
843 auto vecNumElems = vecType.getNumElements();
849 if (op.getElemSizeInBits() == 8 && op.getTileWidth() == 32) {
850 vecElemBitWidth = 16;
851 vecElemType = rewriter.getI16Type();
852 vecNumElems = vecNumElems / 2;
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) {
862 bitWidthId = getTypeMangling(vecElemType,
true);
864 funcName +=
"_transform";
866 funcName +=
"_transpose";
867 spvLoadDstPtr = dstOrSrcPtr;
868 retTypes.push_back(vecType);
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()),
877 bitWidthId = (vecElemBitWidth == 32)
879 : ((vecElemBitWidth == 16) ?
"t" :
"h");
880 LLVM::StoreOp::create(rewriter, loc, op.getStoredVal(), dstOrSrcPtr);
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()),
891 llvm::formatv(
"{0}_{1}b_{2}r{3}x{4}c", funcName, op.getElemSizeInBits(),
892 op.getTileHeight(), op.getTileWidth(), op.getVBlocks())
894 std::string prefetchCode(
"");
897 funcName = llvm::formatv(
"_Z{0}{1}PU3AS1viiiDv2_i{2}{3}", funcName.size(),
898 funcName, prefetchCode, bitWidthId)
900 SmallVector<Type> argTypes;
901 for (
auto arg : args) {
902 argTypes.push_back(arg.getType());
904 createDeviceFunctionCall(
905 rewriter, funcName, LLVM::LLVMVoidType::get(rewriter.getContext()),
906 argTypes, args, paramAttrs, funcAttr, op.getOperation());
908 if constexpr (isLoad)
910 op, LLVM::LoadOp::create(rewriter, loc, vecType, spvLoadDstPtr));
912 rewriter.eraseOp(op);
917template <
typename OpType>
918class BlockLoadStore1DToOCLPattern :
public OpConversionPattern<OpType> {
919 using OpConversionPattern<OpType>::OpConversionPattern;
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>();
930 std::string funcName{
"intel_sub_group_block_"};
933 if constexpr (isStore) {
934 funcName +=
"write_u";
935 valOrResTy = op.getVal().getType();
937 funcName +=
"read_u";
938 valOrResTy = op.getType();
941 VectorType vecTy = dyn_cast<VectorType>(valOrResTy);
942 Type elemType = vecTy ? vecTy.getElementType() : valOrResTy;
943 funcName += getTypeMangling(elemType);
945 funcName += std::to_string(vecTy.getNumElements());
946 SmallVector<Type, 2> argTypes{};
950 SmallVector<bool, 2> isUnsigned{};
954 SmallVector<Value, 2> args{};
955 args.push_back(op.getPtr());
956 argTypes.push_back(op.getPtr().getType());
957 isUnsigned.push_back(
true);
960 applyCacheControlAnnotation(rewriter, loc, op, args, moduleOp,
964 argTypes[0] = args[0].getType();
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());
973 retType = valOrResTy;
975 funcName = std::string(
"_Z") + std::to_string(funcName.size()) + funcName +
977 std::to_string(op.getPtr().getType().getAddressSpace());
978 funcName += getTypeMangling(elemType,
true);
979 if constexpr (isStore)
980 funcName += getTypeMangling(valOrResTy,
true);
981 LLVMFuncAttributeOptions funcAttr{noUnwindWillReturnAttrs};
984 createDeviceFunctionCall(rewriter, funcName, retType, argTypes, args,
985 {}, funcAttr, op.getOperation());
987 if constexpr (isStore)
988 rewriter.eraseOp(op);
990 rewriter.replaceOp(op, call->getResult(0));
995template <
typename OpType>
996class LLVMLoadStoreToOCLPattern :
public OpConversionPattern<OpType> {
997 using OpConversionPattern<OpType>::OpConversionPattern;
999 matchAndRewrite(OpType op,
typename OpType::Adaptor adaptor,
1000 ConversionPatternRewriter &rewriter)
const override {
1001 if (!op->hasDiscardableAttr(
"cache_control"))
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"); });
1014 constexpr bool isStore = std::is_same_v<OpType, LLVM::StoreOp>;
1015 unsigned ptrIdx = isStore ? 1 : 0;
1016 Value ptr = op->getOperand(ptrIdx);
1019 Value annotatedPtr = annotatePtrWithCacheControl(
1020 rewriter, op->getLoc(), ptr, *optCacheControls, moduleOp);
1023 rewriter.modifyOpInPlace(op, [&]() {
1024 op->setOperand(ptrIdx, annotatedPtr);
1025 op->removeDiscardableAttr(
"cache_control");
1058static std::pair<StringRef, int64_t> getConfig(xevm::WorkitemIdXOp) {
1059 return {
"get_local_id", 0};
1061static std::pair<StringRef, int64_t> getConfig(xevm::WorkitemIdYOp) {
1062 return {
"get_local_id", 1};
1064static std::pair<StringRef, int64_t> getConfig(xevm::WorkitemIdZOp) {
1065 return {
"get_local_id", 2};
1067static std::pair<StringRef, int64_t> getConfig(xevm::WorkgroupDimXOp) {
1068 return {
"get_local_size", 0};
1070static std::pair<StringRef, int64_t> getConfig(xevm::WorkgroupDimYOp) {
1071 return {
"get_local_size", 1};
1073static std::pair<StringRef, int64_t> getConfig(xevm::WorkgroupDimZOp) {
1074 return {
"get_local_size", 2};
1076static std::pair<StringRef, int64_t> getConfig(xevm::WorkgroupIdXOp) {
1077 return {
"get_group_id", 0};
1079static std::pair<StringRef, int64_t> getConfig(xevm::WorkgroupIdYOp) {
1080 return {
"get_group_id", 1};
1082static std::pair<StringRef, int64_t> getConfig(xevm::WorkgroupIdZOp) {
1083 return {
"get_group_id", 2};
1085static std::pair<StringRef, int64_t> getConfig(xevm::GridDimXOp) {
1086 return {
"get_num_groups", 0};
1088static std::pair<StringRef, int64_t> getConfig(xevm::GridDimYOp) {
1089 return {
"get_num_groups", 1};
1091static std::pair<StringRef, int64_t> getConfig(xevm::GridDimZOp) {
1092 return {
"get_num_groups", 2};
1096template <
typename OpType>
1097class LaunchConfigOpToOCLPattern :
public OpConversionPattern<OpType> {
1098 using OpConversionPattern<OpType>::OpConversionPattern;
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();
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>(
1119 call.setMemoryEffectsAttr(memAttr);
1120 rewriter.replaceOp(op, call);
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";
1142template <
typename OpType>
1143class SubgroupOpWorkitemOpToOCLPattern :
public OpConversionPattern<OpType> {
1144 using OpConversionPattern<OpType>::OpConversionPattern;
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();
1151 createDeviceFunctionCall(rewriter, func, resTy, {}, {}, {},
1152 noUnwindWillReturnAttrs, op.getOperation());
1153 constexpr auto noModRef = LLVM::ModRefInfo::NoModRef;
1154 auto memAttr = rewriter.getAttr<LLVM::MemoryEffectsAttr>(
1160 call.setMemoryEffectsAttr(memAttr);
1161 rewriter.replaceOp(op, call);
1168static bool isSupportedSPIRVVectorLength(
int64_t numElements) {
1169 return llvm::is_contained({2, 3, 4, 8, 16}, numElements);
1173static Value castIfNeeded(ConversionPatternRewriter &rewriter,
Location loc,
1177 return LLVM::BitcastOp::create(rewriter, loc, ty, val);
1183static Value takeLeadingElements(ConversionPatternRewriter &rewriter,
1185 auto vecTy = cast<VectorType>(val.
getType());
1186 if (vecTy.getNumElements() == numElements)
1189 llvm::to_vector(llvm::seq<int32_t>(0,
static_cast<int32_t
>(numElements)));
1190 return LLVM::ShuffleVectorOp::create(rewriter, loc, val, val, mask);
1204class TruncfToOCLPattern :
public OpConversionPattern<TruncfOp> {
1205 using OpConversionPattern::OpConversionPattern;
1207 matchAndRewrite(TruncfOp op, TruncfOp::Adaptor adaptor,
1208 ConversionPatternRewriter &rewriter)
const override {
1210 auto srcEtype = op.getSrcEtype().getEtype();
1211 auto dstEtype = op.getDstEtype().getEtype();
1217 auto vecSrcTy = dyn_cast<VectorType>(op.getSrc().getType());
1219 return rewriter.notifyMatchFailure(op,
"Scalar src is not supported.");
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");
1227 Type dstTy = op.getDst().getType();
1228 Location loc = op.getLoc();
1229 Value src = op.getSrc();
1230 auto memAttr = rewriter.getAttr<LLVM::MemoryEffectsAttr>(
1231 LLVM::ModRefInfo::NoModRef,
1232 LLVM::ModRefInfo::NoModRef,
1233 LLVM::ModRefInfo::NoModRef,
1234 LLVM::ModRefInfo::NoModRef,
1235 LLVM::ModRefInfo::NoModRef,
1236 LLVM::ModRefInfo::NoModRef);
1237 auto funcAttrs = convergentNoUnwindWillReturnAttrs;
1238 funcAttrs.memEffectsAttr = memAttr;
1241 if (dstEtype == TruncfDstElemTypes::E2M1) {
1251 constexpr int kDnsclConvertToE2M1 = 1;
1252 constexpr int kDnsclModeBytes02 = 0;
1253 constexpr int kDnsclModeBytes13 = 2;
1256 int64_t numLanes = llvm::divideCeil(numElements, 2);
1258 Type i32Ty = rewriter.getI32Type();
1259 Type i8Ty = rewriter.getI8Type();
1263 if (numElements != numLanes * 2) {
1264 SmallVector<int32_t> mask = llvm::to_vector(
1265 llvm::seq<int32_t>(0,
static_cast<int32_t
>(numElements)));
1267 mask.append(
static_cast<size_t>(numLanes * 2 - numElements), 0);
1268 padded = LLVM::ShuffleVectorOp::create(rewriter, loc, src, src, mask);
1274 laneVec = LLVM::BitcastOp::create(
1275 rewriter, loc, VectorType::get(numLanes, i32Ty), padded);
1277 laneVec = LLVM::BitcastOp::create(rewriter, loc, i32Ty, padded);
1278 auto getLane = [&](int64_t idx) -> Value {
1282 LLVM::ConstantOp::create(rewriter, loc, rewriter.getI32Type(), idx);
1283 return LLVM::ExtractElementOp::create(rewriter, loc, laneVec, pos)
1287 std::string fnName =
"__builtin_IB_dnscl_";
1288 fnName += (srcEtype == TruncfSrcElemTypes::F16) ?
"hf16" :
"bf16";
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(),
1295 SmallVector<Value> args{lo, hi, convertTo, modeVal};
1296 return createDeviceFunctionCall(rewriter, fnName, i32Ty, argTypes, args,
1297 {}, funcAttrs, op.getOperation())
1302 if (numLanes <= 2) {
1305 Value lo = getLane(0);
1309 : LLVM::UndefOp::create(rewriter, loc, i32Ty)->getResult(0);
1310 Value dword = genDnscl(lo, hi, kDnsclModeBytes02);
1311 if (numLanes == 1) {
1313 result = LLVM::TruncOp::create(rewriter, loc, i8Ty, dword);
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});
1322 SmallVector<Value> dwords;
1323 for (int64_t base = 0; base < numLanes; base += 4) {
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));
1336 if (dwords.size() == 1) {
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);
1344 LLVM::InsertElementOp::create(rewriter, loc,
result, dword, pos)
1349 rewriter.replaceOp(op, castIfNeeded(rewriter, loc, dstTy,
result));
1356 std::string lenSuffix = std::to_string(numElements);
1359 if (srcEtype == TruncfSrcElemTypes::BF16) {
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())
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());
1380 createDeviceFunctionCall(rewriter, truncFnName, resTy, truncArgTypes,
1381 truncArgs, {}, funcAttrs, op.getOperation())
1384 if (dstEtype == TruncfDstElemTypes::BF8) {
1386 std::string fnName =
"__builtin_IB_hftobf8_" + lenSuffix;
1387 SmallVector<Type> argTypes{src.
getType()};
1388 SmallVector<Value> args{src};
1390 createDeviceFunctionCall(rewriter, fnName, dstTy, argTypes, args, {},
1391 funcAttrs, op.getOperation())
1394 rewriter.replaceOp(op,
result);
1395 }
else if (dstEtype == TruncfDstElemTypes::F8) {
1397 std::string fnName =
"__builtin_IB_hftohf8_" + lenSuffix;
1398 SmallVector<Type> argTypes{src.
getType()};
1399 SmallVector<Value> args{src};
1401 createDeviceFunctionCall(rewriter, fnName, dstTy, argTypes, args, {},
1402 funcAttrs, op.getOperation())
1405 rewriter.replaceOp(op,
result);
1407 return rewriter.notifyMatchFailure(
1408 op,
"Unsupported src, dst element type pair.");
1414class ExtfToOCLPattern :
public OpConversionPattern<ExtfOp> {
1415 using OpConversionPattern::OpConversionPattern;
1417 matchAndRewrite(ExtfOp op, ExtfOp::Adaptor adaptor,
1418 ConversionPatternRewriter &rewriter)
const override {
1421 auto srcEtype = op.getSrcEtype().getEtype();
1422 auto dstEtype = op.getDstEtype().getEtype();
1425 Type srcTy = op.getSrc().getType();
1427 auto vecDstTy = dyn_cast<VectorType>(op.getDst().getType());
1429 return rewriter.notifyMatchFailure(op,
"Scalar dst is not supported.");
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 LLVM::ModRefInfo::NoModRef,
1439 LLVM::ModRefInfo::NoModRef,
1440 LLVM::ModRefInfo::NoModRef,
1441 LLVM::ModRefInfo::NoModRef,
1442 LLVM::ModRefInfo::NoModRef,
1443 LLVM::ModRefInfo::NoModRef);
1444 auto funcAttrs = convergentNoUnwindWillReturnAttrs;
1445 funcAttrs.memEffectsAttr = memAttr;
1448 if (srcEtype == ExtfSrcElemTypes::E2M1) {
1461 int64_t numBytes = llvm::divideCeil(numElements, 2);
1462 constexpr int kLutE2M1ToF16 = 7;
1463 constexpr int kLutE2M1ToBF16 = 5;
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());
1470 createDeviceFunctionCall(rewriter,
"__builtin_IB_shfl_idx4_lut",
1471 lutTy, {lutIdx.
getType()}, {lutIdx}, {},
1472 funcAttrs, op.getOperation())
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) {
1482 packedResTy = i32Ty;
1484 fnName += std::to_string(numBytes) +
"_";
1485 argTy = VectorType::get(numBytes, i8Ty);
1486 packedResTy = VectorType::get(numBytes, i32Ty);
1489 SmallVector<Type> convArgTypes{lut.
getType(), argTy};
1490 SmallVector<Value> convArgs{lut, castIfNeeded(rewriter, loc, argTy, src)};
1492 createDeviceFunctionCall(rewriter, fnName, packedResTy, convArgTypes,
1493 convArgs, {}, funcAttrs, op.getOperation())
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);
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);
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};
1523 createDeviceFunctionCall(rewriter, fnName, f16Ty, argTypes, args, {},
1524 funcAttrs, op.getOperation())
1528 if (dstEtype == ExtfDstElemTypes::F16) {
1529 rewriter.replaceOp(op,
result);
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());
1543 createDeviceFunctionCall(rewriter, convFnName, f32Ty, convArgTypes,
1544 convArgs, {}, funcAttrs, op.getOperation())
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());
1553 createDeviceFunctionCall(rewriter, ftobfFnName, i16Ty, ftobfArgTypes,
1554 ftobfArgs, {}, funcAttrs, op.getOperation())
1557 result = LLVM::BitcastOp::create(rewriter, op.getLoc(), vecDstTy,
result);
1558 rewriter.replaceOp(op,
result);
1563class MMAMxToOCLPattern :
public OpConversionPattern<MMAMxOp> {
1564 using OpConversionPattern::OpConversionPattern;
1566 matchAndRewrite(MMAMxOp op, MMAMxOp::Adaptor adaptor,
1567 ConversionPatternRewriter &rewriter)
const override {
1569 return rewriter.notifyMatchFailure(op,
"OCL requires C operand");
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");
1577 constexpr uint32_t bitWidthPackedA{16};
1578 constexpr uint32_t bitWidthPackedB{32};
1579 auto loc = op.getLoc();
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);
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);
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);
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");
1611 cOrigTy.getElementType().isBF16()
1612 ? VectorType::get(cOrigTy.getShape(), rewriter.getIntegerType(16))
1614 VectorType resTy = cTy;
1616 c = LLVM::BitcastOp::create(rewriter, loc, cTy, c);
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()))
1625 auto scaleA = op.getScaleA();
1626 auto scaleB = op.getScaleB();
1627 SmallVector<Type> argTypes{cTy, a.
getType(),
b.getType(), scaleA.getType(),
1629 SmallVector<Value> args{c, a,
b, scaleA, scaleB};
1631 auto memAttr = rewriter.getAttr<LLVM::MemoryEffectsAttr>(
1632 LLVM::ModRefInfo::NoModRef,
1633 LLVM::ModRefInfo::NoModRef,
1634 LLVM::ModRefInfo::NoModRef,
1635 LLVM::ModRefInfo::NoModRef,
1636 LLVM::ModRefInfo::NoModRef,
1637 LLVM::ModRefInfo::NoModRef);
1638 auto funcAttrs = convergentNoUnwindWillReturnAttrs;
1639 funcAttrs.memEffectsAttr = memAttr;
1641 createDeviceFunctionCall(rewriter, fnName, resTy, argTypes, args, {},
1642 funcAttrs, op.getOperation())
1645 if (resOrigTy != resTy)
1646 result = LLVM::BitcastOp::create(rewriter, loc, resOrigTy,
result);
1648 rewriter.replaceOp(op,
result);
1661class BitcastShuffleToGenISAPattern
1662 :
public OpConversionPattern<BitcastShuffleOp> {
1663 using OpConversionPattern::OpConversionPattern;
1665 matchAndRewrite(BitcastShuffleOp op, BitcastShuffleOp::Adaptor adaptor,
1666 ConversionPatternRewriter &rewriter)
const override {
1667 Type srcTy = op.getSrc().getType();
1668 Type resTy = op.getRes().getType();
1670 std::string fnName =
"llvm.genx.GenISA.SubgroupBitcastShuffle." +
1671 getGenISATypeMangling(resTy) +
"." +
1672 getGenISATypeMangling(srcTy);
1674 Value
result = createDeviceFunctionCall(
1675 rewriter, fnName, resTy, {srcTy}, {adaptor.getSrc()}, {},
1676 convergentNoUnwindWillReturnAttrs, op.getOperation())
1679 rewriter.replaceOp(op,
result);
1684class AllocaToGlobalPattern :
public OpConversionPattern<LLVM::AllocaOp> {
1685 using OpConversionPattern::OpConversionPattern;
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();
1693 auto symTable = op->getParentWithTrait<OpTrait::SymbolTable>();
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();
1705 auto val = op.getArraySize();
1709 auto loc = op.getLoc();
1710 auto globalType = LLVM::LLVMArrayType::get(
1711 rewriter.getContext(), op.getElemType(), cst.getZExtValue());
1712 LLVM::GlobalOp globalVar;
1714 OpBuilder::InsertionGuard guard(rewriter);
1715 rewriter.setInsertionPointToStart(moduleBody);
1716 auto alignment = op.getAlignment();
1717 globalVar = LLVM::GlobalOp::create(
1718 rewriter, loc, globalType,
false,
1719 LLVM::Linkage::Internal,
1720 std::string(
"__global_alloca_") +
1721 std::to_string(getNextGlobalIdx()),
1723 alignment ? *alignment : 0, addrSpace);
1725 rewriter.replaceOpWithNewOp<LLVM::AddressOfOp>(op, globalVar);
1730 static unsigned getNextGlobalIdx() {
1731 static unsigned globalIdx = 0;
1744static bool isExtractingContiguousSlice(LLVM::ShuffleVectorOp op) {
1745 if (op.getV1() != op.getV2() &&
1746 !isa_and_present<LLVM::PoisonOp, LLVM::UndefOp>(
1747 op.getV2().getDefiningOp()))
1749 auto maskAttr = op.getMask();
1751 int64_t sourceSize = op.getV1().getType().getNumElements();
1752 if (maskSize > sourceSize)
1754 int64_t firstIndex = maskAttr[0];
1755 if (firstIndex < 0 || firstIndex >= sourceSize)
1757 for (
int64_t i = 1; i < maskSize; ++i) {
1759 if (
index != firstIndex + i)
1761 if (
index >= sourceSize)
1777 return rewriter.
create(state);
1791class HandleVectorExtractPattern
1793 using OpRewritePattern<LLVM::ShuffleVectorOp>::OpRewritePattern;
1795 void initialize() { setHasBoundedRewriteRecursion(); }
1797 LogicalResult matchAndRewrite(LLVM::ShuffleVectorOp op,
1798 PatternRewriter &rewriter)
const override {
1800 if (!isExtractingContiguousSlice(op))
1803 auto mask = op.getMask();
1804 auto loc = op.getLoc();
1805 auto ty = op.getType();
1807 auto src = op.getV1();
1810 if (isa<LLVM::FPExtOp>(srcOp) || isa<LLVM::FPTruncOp>(srcOp)) {
1811 Value srcInput = srcOp->getOperand(0);
1813 auto srcVecTy = dyn_cast<VectorType>(srcInput.
getType());
1816 auto newShuffleVecTy =
1817 VectorType::get(mask.size(), srcVecTy.getElementType());
1818 auto newShuffle = LLVM::ShuffleVectorOp::create(
1819 rewriter, loc, newShuffleVecTy, srcInput, srcInput, mask);
1822 if (isa<LLVM::FPExtOp>(srcOp)) {
1823 newUnaryOp = LLVM::FPExtOp::create(rewriter, loc, ty, newShuffle);
1825 newUnaryOp = LLVM::FPTruncOp::create(rewriter, loc, ty, newShuffle);
1828 }
else if (isa<LLVM::BitcastOp>(srcOp)) {
1829 Value srcInput = srcOp->getOperand(0);
1832 auto srcInputVecTy = dyn_cast<VectorType>(srcInput.
getType());
1833 auto srcResVecTy = dyn_cast<VectorType>(srcOp->getResult(0).getType());
1834 if (!srcInputVecTy || !srcResVecTy)
1836 auto srcInputSize = srcInputVecTy.getNumElements();
1837 auto srcResSize = srcResVecTy.getNumElements();
1838 auto maskSize =
static_cast<int32_t
>(mask.size());
1839 if (srcInputSize > srcResSize) {
1842 if (srcResSize % srcInputSize != 0) {
1845 auto maskScale = srcResSize / srcInputSize;
1850 SmallVector<int32_t> newMask;
1851 if (maskScale != 1) {
1854 if (mask[0] % maskScale != 0 || maskSize % maskScale != 0) {
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);
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);
1871 LLVM::BitcastOp::create(rewriter, loc, ty, newShuffle);
1873 }
else if (isa<LLVM::ShuffleVectorOp>(srcOp)) {
1878 auto srcShuffle = cast<LLVM::ShuffleVectorOp>(srcOp);
1879 if (!isExtractingContiguousSlice(srcShuffle))
1881 auto srcMask = srcShuffle.getMask();
1882 SmallVector<int32_t> combinedMask;
1883 for (
auto index : mask) {
1884 combinedMask.push_back(srcMask[index]);
1886 auto newShuffle = LLVM::ShuffleVectorOp::create(
1887 rewriter, loc, ty, srcShuffle.getV1(), srcShuffle.getV1(),
1890 }
else if (isa<LLVM::LoadOp>(srcOp)) {
1892 auto loadOp = cast<LLVM::LoadOp>(srcOp);
1893 auto loadPtr = loadOp.getAddr();
1894 auto loadAddrSpace = loadPtr.getType().getAddressSpace();
1895 if (loadAddrSpace != 0)
1897 auto loadTy = dyn_cast<VectorType>(loadOp.getType());
1900 auto elemTy = loadTy.getElementType();
1901 auto firstIndex = mask[0];
1902 auto newVecTy = VectorType::get(mask.size(), elemTy);
1905 auto newPtr = LLVM::GEPOp::create(
1907 LLVM::LLVMPointerType::get(rewriter.
getContext(), loadAddrSpace),
1908 elemTy, loadPtr, ArrayRef<LLVM::GEPArg>{firstIndex});
1909 auto newLoad = LLVM::LoadOp::create(rewriter, loc, newVecTy, newPtr);
1912 auto newLoad = LLVM::LoadOp::create(rewriter, loc, newVecTy, loadPtr);
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();
1928 SmallVector<Value> newOperands;
1929 newOperands.reserve(srcOp->getNumOperands());
1930 for (Value operand : srcOp->getOperands()) {
1931 auto operandTy = cast<VectorType>(operand.
getType());
1933 VectorType::get(mask.size(), operandTy.getElementType());
1934 newOperands.push_back(LLVM::ShuffleVectorOp::create(
1935 rewriter, loc, sliceTy, operand, operand, mask));
1955struct ConvertXeVMToLLVMPass
1959 void getDependentDialects(DialectRegistry ®istry)
const override {
1960 registry.
insert<LLVM::LLVMDialect, XeVMDialect>();
1963 void runOnOperation()
override {
1967 if (
failed(applyPartialConversion(getOperation(),
target,
1968 std::move(patterns))))
1969 signalPassFailure();
1973 RewritePatternSet vectorPatterns(&
getContext());
1974 vectorPatterns.add<HandleVectorExtractPattern>(&
getContext());
1975 GreedyRewriteConfig config{};
1980 config.enableFolding(
false);
1997 target.addDynamicallyLegalDialect<LLVM::LLVMDialect>([](
Operation *op) {
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;
2008 return !op->hasDiscardableAttr(
"cache_control");
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>(
LogicalResult initialize(unsigned origNumLoops, ArrayRef< ReassociationIndices > foldedIterationDims)
static Operation * cloneOpWithOperandsAndTypes(RewriterBase &rewriter, Location loc, Operation *op, ArrayRef< Value > operands, ArrayRef< Type > resultTypes)
Attributes are known-constant values of operations.
MLIRContext * getContext() const
This class defines the main interface for locations in MLIR and acts as a non-nullable wrapper around...
NamedAttribute represents a combination of a name and an Attribute value.
RAII guard to reset the insertion point of the builder when destroyed.
Operation * create(const OperationState &state)
Creates an operation given the fields represented as an OperationState.
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.
Region & getRegion(unsigned index)
Returns the region held by this operation at position 'index'.
OpResult getResult(unsigned idx)
Get the 'idx'th result of this operation.
Operation * getParentWithTrait()
Returns the closest surrounding parent operation with trait Trait.
Location getLoc()
The source location the operation was defined or derived from.
DictionaryAttr getRawDictionaryAttrs()
Return all attributes that are not stored as properties.
Attribute getPropertiesAsAttribute()
Return the properties converted to an attribute.
OperationName getName()
The name of an operation is the key identifier for it.
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...
This class represents an instance of an SSA value in the MLIR system, representing a computable value...
Type getType() const
Return the type of this value.
Operation * getDefiningOp() const
If this value is the result of an operation, return the operation that defines it.
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`.
Include the generated interface declarations.
bool matchPattern(Value value, const Pattern &pattern)
Entry point for matching a pattern over a Value.
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...
LogicalResult applyPatternsGreedily(Region ®ion, 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
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.