29#include "llvm/ADT/STLExtras.h"
30#include "llvm/Support/FormatVariadic.h"
35#include "llvm/ADT/TypeSwitch.h"
40#define GEN_PASS_DEF_CONVERTXEGPUTOXEVMPASS
41#include "mlir/Conversion/Passes.h.inc"
49static constexpr int32_t systolicDepth{8};
50static constexpr int32_t executionSize{16};
53enum class NdTdescOffset : uint32_t {
64static constexpr int64_t maxNdTdescLeadingDims{3};
65static constexpr int64_t maxNdTdescRank{2 + maxNdTdescLeadingDims};
67static int32_t getNumericXeVMAddrSpace(xegpu::MemorySpace xeGpuMemspace) {
68 switch (xeGpuMemspace) {
69 case xegpu::MemorySpace::Global:
70 return static_cast<int>(xevm::AddrSpace::GLOBAL);
71 case xegpu::MemorySpace::SLM:
72 return static_cast<int>(xevm::AddrSpace::SHARED);
74 llvm_unreachable(
"Unknown XeGPU memory space");
84static FailureOr<unsigned> getNumericMemorySpace(
Attribute memSpace) {
87 if (
auto intAttr = llvm::dyn_cast<IntegerAttr>(memSpace))
88 return static_cast<unsigned>(intAttr.getInt());
89 if (
auto xevmSpace = llvm::dyn_cast<xevm::AddrSpaceAttr>(memSpace))
90 return static_cast<unsigned>(xevmSpace.getValue());
91 if (
auto gpuSpace = llvm::dyn_cast<gpu::AddressSpaceAttr>(memSpace)) {
92 switch (gpuSpace.getValue()) {
93 case gpu::AddressSpace::Global:
94 return static_cast<unsigned>(xevm::AddrSpace::GLOBAL);
95 case gpu::AddressSpace::Workgroup:
96 return static_cast<unsigned>(xevm::AddrSpace::SHARED);
97 case gpu::AddressSpace::Private:
98 return static_cast<unsigned>(xevm::AddrSpace::PRIVATE);
99 case gpu::AddressSpace::Constant:
100 return static_cast<unsigned>(xevm::AddrSpace::CONSTANT);
102 llvm_unreachable(
"Unknown GPU address space");
108static bool isSharedMemRef(
const MemRefType &memrefTy) {
109 FailureOr<unsigned> addrSpace =
110 getNumericMemorySpace(memrefTy.getMemorySpace());
111 return succeeded(addrSpace) &&
112 *addrSpace ==
static_cast<unsigned>(xevm::AddrSpace::SHARED);
116static VectorType encodeVectorTypeTo(VectorType currentVecType,
118 auto elemType = currentVecType.getElementType();
119 auto currentBitWidth = elemType.getIntOrFloatBitWidth();
122 currentVecType.getNumElements() * currentBitWidth / newBitWidth;
123 return VectorType::get(size, toElemType);
126static xevm::LoadCacheControl
127translateLoadXeGPUCacheHint(std::optional<xegpu::CachePolicy> L1hint,
128 std::optional<xegpu::CachePolicy> L3hint) {
130 if (!L1hint && !L3hint)
131 return xevm::LoadCacheControl::USE_DEFAULT;
133 auto L1hintVal = L1hint.value_or(xegpu::CachePolicy::CACHED);
134 auto L3hintVal = L3hint.value_or(xegpu::CachePolicy::CACHED);
136 case xegpu::CachePolicy::CACHED:
137 if (L3hintVal == xegpu::CachePolicy::CACHED)
138 return xevm::LoadCacheControl::L1C_L2UC_L3C;
139 else if (L3hintVal == xegpu::CachePolicy::UNCACHED)
140 return xevm::LoadCacheControl::L1C_L2UC_L3UC;
142 llvm_unreachable(
"Unsupported cache control.");
143 case xegpu::CachePolicy::UNCACHED:
144 if (L3hintVal == xegpu::CachePolicy::CACHED)
145 return xevm::LoadCacheControl::L1UC_L2UC_L3C;
146 else if (L3hintVal == xegpu::CachePolicy::UNCACHED)
147 return xevm::LoadCacheControl::L1UC_L2UC_L3UC;
149 llvm_unreachable(
"Unsupported cache control.");
150 case xegpu::CachePolicy::STREAMING:
151 if (L3hintVal == xegpu::CachePolicy::CACHED)
152 return xevm::LoadCacheControl::L1S_L2UC_L3C;
153 else if (L3hintVal == xegpu::CachePolicy::UNCACHED)
154 return xevm::LoadCacheControl::L1S_L2UC_L3UC;
156 llvm_unreachable(
"Unsupported cache control.");
157 case xegpu::CachePolicy::READ_INVALIDATE:
158 return xevm::LoadCacheControl::INVALIDATE_READ;
160 llvm_unreachable(
"Unsupported cache control.");
164static xevm::StoreCacheControl
165translateStoreXeGPUCacheHint(std::optional<xegpu::CachePolicy> L1hint,
166 std::optional<xegpu::CachePolicy> L3hint) {
168 if (!L1hint && !L3hint)
169 return xevm::StoreCacheControl::USE_DEFAULT;
171 auto L1hintVal = L1hint.value_or(xegpu::CachePolicy::UNCACHED);
172 auto L3hintVal = L3hint.value_or(xegpu::CachePolicy::WRITE_BACK);
174 case xegpu::CachePolicy::UNCACHED:
175 if (L3hintVal == xegpu::CachePolicy::UNCACHED)
176 return xevm::StoreCacheControl::L1UC_L2UC_L3UC;
177 else if (L3hintVal == xegpu::CachePolicy::WRITE_BACK)
178 return xevm::StoreCacheControl::L1UC_L2UC_L3WB;
180 llvm_unreachable(
"Unsupported cache control.");
181 case xegpu::CachePolicy::STREAMING:
182 if (L3hintVal == xegpu::CachePolicy::UNCACHED)
183 return xevm::StoreCacheControl::L1S_L2UC_L3UC;
184 else if (L3hintVal == xegpu::CachePolicy::WRITE_BACK)
185 return xevm::StoreCacheControl::L1S_L2UC_L3WB;
187 llvm_unreachable(
"Unsupported cache control.");
188 case xegpu::CachePolicy::WRITE_BACK:
189 if (L3hintVal == xegpu::CachePolicy::UNCACHED)
190 return xevm::StoreCacheControl::L1WB_L2UC_L3UC;
191 else if (L3hintVal == xegpu::CachePolicy::WRITE_BACK)
192 return xevm::StoreCacheControl::L1WB_L2UC_L3WB;
194 llvm_unreachable(
"Unsupported cache control.");
195 case xegpu::CachePolicy::WRITE_THROUGH:
196 if (L3hintVal == xegpu::CachePolicy::UNCACHED)
197 return xevm::StoreCacheControl::L1WT_L2UC_L3UC;
198 else if (L3hintVal == xegpu::CachePolicy::WRITE_BACK)
199 return xevm::StoreCacheControl::L1WT_L2UC_L3WB;
201 llvm_unreachable(
"Unsupported cache control.");
203 llvm_unreachable(
"Unsupported cache control.");
240class CreateNdDescToXeVMPattern
241 :
public OpConversionPattern<xegpu::CreateNdDescOp> {
242 using OpConversionPattern::OpConversionPattern;
244 matchAndRewrite(xegpu::CreateNdDescOp op,
245 xegpu::CreateNdDescOp::Adaptor adaptor,
246 ConversionPatternRewriter &rewriter)
const override {
247 auto loc = op.getLoc();
248 auto source = op.getSource();
252 int64_t rank = op.getType().getRank();
254 auto memrefTy = dyn_cast<MemRefType>(source.getType());
256 if (!memrefTy.isStrided())
257 return rewriter.notifyMatchFailure(op,
"Expected strided Memref.");
258 sourceRank = memrefTy.getRank();
259 }
else if (isa<IntegerType>(source.getType())) {
260 sourceRank = op.getMixedSizes().size();
262 return rewriter.notifyMatchFailure(op,
263 "Expected ranked Memref or integer.");
265 if (sourceRank != rank)
266 return rewriter.notifyMatchFailure(
267 op,
"Expected descriptor rank to match source rank; subview the "
268 "source down to the descriptor rank.");
269 if (rank > maxNdTdescRank)
270 return rewriter.notifyMatchFailure(
271 op,
"Batched nd descriptor supports at most " +
272 std::to_string(maxNdTdescLeadingDims) +
273 " leading dims (rank <= " + std::to_string(maxNdTdescRank) +
277 SmallVector<std::optional<int64_t>> constStrides(rank, std::nullopt);
279 SmallVector<int64_t> staticStrides;
280 int64_t staticOffset;
282 memrefTy.getStridesAndOffset(staticStrides, staticOffset)))
283 for (int64_t d = 0; d < rank; ++d)
284 if (!ShapedType::isDynamic(staticStrides[d]))
285 constStrides[d] = staticStrides[d];
287 SmallVector<OpFoldResult> mixed = op.getMixedStrides();
288 for (int64_t d = 0; d < rank; ++d)
291 if (std::optional<int64_t> pitch = constStrides[rank - 2]) {
292 for (int64_t d = 0; d < rank - 2; ++d) {
293 std::optional<int64_t> leading = constStrides[d];
294 if (leading && (*pitch == 0 || *leading % *pitch != 0))
295 return rewriter.notifyMatchFailure(
296 op,
"Expected each leading (batch) stride to be a multiple of "
297 "the row stride; the source has gaps between planes.");
302 Type payloadElemTy = rewriter.getI32Type();
303 Type i64Ty = rewriter.getI64Type();
307 Value baseAddr = adaptor.getSource();
308 if (isa<IntegerType>(source.getType()) && baseAddr.
getType() != i64Ty) {
310 baseAddr = arith::ExtUIOp::create(rewriter, loc, i64Ty, baseAddr);
314 rewriter.replaceOp(op, baseAddr);
318 SmallVector<OpFoldResult> mixedSizes;
319 SmallVector<OpFoldResult> mixedStrides;
322 memref::ExtractStridedMetadataOp::create(rewriter, loc, source);
323 mixedSizes = meta.getConstifiedMixedSizes();
324 mixedStrides = meta.getConstifiedMixedStrides();
326 mixedSizes = op.getMixedSizes();
327 mixedStrides = op.getMixedStrides();
333 VectorType payloadTy = VectorType::get(8, payloadElemTy);
335 VectorType payloadI64Ty = VectorType::get(4, i64Ty);
337 Value payload = arith::ConstantOp::create(
342 auto createOffset = [&](SmallVector<OpFoldResult> &ofrVec,
343 unsigned idx) -> Value {
349 Value baseShapeW = createOffset(mixedSizes, rank - 1);
353 Value baseShapeH = createOffset(mixedSizes, rank - 2);
354 for (int64_t d = 0; d < rank - 2; ++d)
355 baseShapeH = arith::MulIOp::create(rewriter, loc, baseShapeH,
356 createOffset(mixedSizes, d));
358 Value basePitch = createOffset(mixedStrides, rank - 2);
361 vector::BitCastOp::create(rewriter, loc, payloadI64Ty, payload);
363 vector::InsertOp::create(rewriter, loc, baseAddr, payLoadAsI64,
364 static_cast<int>(NdTdescOffset::BasePtr));
365 payload = vector::BitCastOp::create(rewriter, loc, payloadTy, payLoadAsI64);
367 vector::InsertOp::create(rewriter, loc, baseShapeW, payload,
368 static_cast<int>(NdTdescOffset::BaseShapeW));
370 vector::InsertOp::create(rewriter, loc, baseShapeH, payload,
371 static_cast<int>(NdTdescOffset::BaseShapeH));
373 vector::InsertOp::create(rewriter, loc, basePitch, payload,
374 static_cast<int>(NdTdescOffset::BasePitch));
379 for (int64_t d = 0; d < rank - 2; ++d) {
381 std::optional<int64_t> pitch =
383 Value leadingRowStride;
384 if (leading && pitch && *pitch != 0) {
386 rewriter, loc, payloadElemTy, *leading / *pitch);
388 leadingRowStride = arith::DivUIOp::create(
389 rewriter, loc, createOffset(mixedStrides, d), basePitch);
391 payload = vector::InsertOp::create(
392 rewriter, loc, leadingRowStride, payload,
393 static_cast<int>(NdTdescOffset::LeadingStride0) + d);
395 rewriter.replaceOp(op, payload);
402 typename = std::enable_if_t<llvm::is_one_of<
403 OpType, xegpu::LoadNdOp, xegpu::StoreNdOp, xegpu::PrefetchNdOp>::value>>
404class LoadStorePrefetchNdToXeVMPattern :
public OpConversionPattern<OpType> {
405 using OpConversionPattern<OpType>::OpConversionPattern;
407 matchAndRewrite(OpType op,
typename OpType::Adaptor adaptor,
408 ConversionPatternRewriter &rewriter)
const override {
409 auto mixedOffsets = op.getMixedOffsets();
410 int64_t opOffsetsSize = mixedOffsets.size();
411 auto loc = op.getLoc();
412 auto ctxt = rewriter.getContext();
414 auto tdesc = adaptor.getTensorDesc();
415 auto tdescTy = op.getTensorDescType();
416 auto tileRank = tdescTy.getRank();
417 if (opOffsetsSize != tileRank)
418 return rewriter.notifyMatchFailure(
419 op,
"Expected offset rank to match descriptor rank.");
420 if (tileRank > 2 && llvm::any_of(tdescTy.getShape().drop_back(2),
421 [](int64_t d) { return d != 1; }))
422 return rewriter.notifyMatchFailure(
423 op,
"Expected leading (batch) descriptor dims to be unit.");
424 if (tileRank > maxNdTdescRank)
425 return rewriter.notifyMatchFailure(
426 op,
"Expected descriptor rank <= " + std::to_string(maxNdTdescRank) +
428 auto elemType = tdescTy.getElementType();
429 auto elemBitSize = elemType.getIntOrFloatBitWidth();
430 bool isSubByte = elemBitSize < 8;
431 uint64_t wScaleFactor = 1;
433 if (!isSubByte && (elemBitSize % 8 != 0))
434 return rewriter.notifyMatchFailure(
435 op,
"Expected element type bit width to be multiple of 8.");
436 auto tileW = tdescTy.getDimSize(tileRank - 1);
439 if (elemBitSize != 4)
440 return rewriter.notifyMatchFailure(
441 op,
"Only sub byte types of 4bits are supported.");
443 return rewriter.notifyMatchFailure(
444 op,
"Sub byte types are only supported for 2D tensor descriptors.");
445 auto subByteFactor = 8 / elemBitSize;
446 auto tileH = tdescTy.getDimSize(0);
448 if constexpr (std::is_same_v<OpType, xegpu::LoadNdOp>) {
449 if (op.getPacked().value_or(
false)) {
451 if (tileH == systolicDepth * 4 &&
452 tileW == executionSize * subByteFactor) {
457 elemType = rewriter.getIntegerType(8);
458 tileW = executionSize;
459 wScaleFactor = subByteFactor;
464 if (wScaleFactor == 1) {
465 auto sub16BitFactor = subByteFactor * 2;
466 if (tileW == executionSize * sub16BitFactor) {
470 elemType = rewriter.getIntegerType(16);
471 tileW = executionSize;
472 wScaleFactor = sub16BitFactor;
474 return rewriter.notifyMatchFailure(
475 op,
"Unsupported tile shape for sub byte types.");
479 elemBitSize = elemType.getIntOrFloatBitWidth();
483 auto ptrTypeLLVM = LLVM::LLVMPointerType::get(
484 ctxt, getNumericXeVMAddrSpace(tdescTy.getMemorySpace()));
488 rewriter, loc, rewriter.getI32Type(), elemBitSize / 8);
489 VectorType payloadI64Ty = VectorType::get(4, rewriter.getI64Type());
491 vector::BitCastOp::create(rewriter, loc, payloadI64Ty, tdesc);
493 vector::ExtractOp::create(rewriter, loc, payLoadAsI64,
494 static_cast<int>(NdTdescOffset::BasePtr));
495 Value baseShapeW = vector::ExtractOp::create(
496 rewriter, loc, tdesc,
static_cast<int>(NdTdescOffset::BaseShapeW));
497 Value baseShapeH = vector::ExtractOp::create(
498 rewriter, loc, tdesc,
static_cast<int>(NdTdescOffset::BaseShapeH));
499 Value basePitch = vector::ExtractOp::create(
500 rewriter, loc, tdesc,
static_cast<int>(NdTdescOffset::BasePitch));
503 mixedOffsets[tileRank - 1]);
505 rewriter.getI32Type(), offsetW);
507 mixedOffsets[tileRank - 2]);
509 rewriter.getI32Type(), offsetH);
516 for (int64_t d = 0; d < tileRank - 2; ++d) {
520 rewriter.getI32Type(), off);
521 Value rowStride = vector::ExtractOp::create(
522 rewriter, loc, tdesc,
523 static_cast<int>(NdTdescOffset::LeadingStride0) + d);
524 Value term = arith::MulIOp::create(rewriter, loc, off, rowStride);
525 offsetH = arith::AddIOp::create(rewriter, loc, offsetH, term);
529 LLVM::IntToPtrOp::create(rewriter, loc, ptrTypeLLVM, basePtr);
533 Value baseShapeWInBytes =
534 arith::MulIOp::create(rewriter, loc, baseShapeW, elemByteSize);
536 Value basePitchBytes =
537 arith::MulIOp::create(rewriter, loc, basePitch, elemByteSize);
539 if (wScaleFactor > 1) {
543 rewriter, loc, rewriter.getI32Type(), llvm::Log2_64(wScaleFactor));
544 baseShapeWInBytes = arith::ShRSIOp::create(
545 rewriter, loc, baseShapeWInBytes, wScaleFactorValLog2);
546 basePitchBytes = arith::ShRSIOp::create(rewriter, loc, basePitchBytes,
547 wScaleFactorValLog2);
549 arith::ShRSIOp::create(rewriter, loc, offsetW, wScaleFactorValLog2);
552 auto tileH = tdescTy.getDimSize(tileRank - 2);
554 int32_t vblocks = tdescTy.getArrayLength();
555 if constexpr (std::is_same_v<OpType, xegpu::StoreNdOp>) {
556 Value src = adaptor.getValue();
562 VectorType srcVecTy = dyn_cast<VectorType>(src.
getType());
564 return rewriter.notifyMatchFailure(
565 op,
"Expected store value to be a vector type.");
567 VectorType newSrcVecTy =
568 encodeVectorTypeTo(srcVecTy, rewriter.getIntegerType(elemBitSize));
569 if (srcVecTy != newSrcVecTy)
570 src = vector::BitCastOp::create(rewriter, loc, newSrcVecTy, src);
571 auto storeCacheControl =
572 translateStoreXeGPUCacheHint(op.getL1Hint(), op.getL3Hint());
573 xevm::BlockStore2dOp::create(
574 rewriter, loc, basePtrLLVM, baseShapeWInBytes, baseShapeH,
575 basePitchBytes, offsetW, offsetH, elemBitSize, tileW, tileH, src,
576 xevm::StoreCacheControlAttr::get(ctxt, storeCacheControl));
577 rewriter.eraseOp(op);
579 auto loadCacheControl =
580 translateLoadXeGPUCacheHint(op.getL1Hint(), op.getL3Hint());
581 if constexpr (std::is_same_v<OpType, xegpu::PrefetchNdOp>) {
582 xevm::BlockPrefetch2dOp::create(
583 rewriter, loc, basePtrLLVM, baseShapeWInBytes, baseShapeH,
584 basePitchBytes, offsetW, offsetH, elemBitSize, tileW, tileH,
585 vblocks, xevm::LoadCacheControlAttr::get(ctxt, loadCacheControl));
586 rewriter.eraseOp(op);
588 VectorType dstVecTy = cast<VectorType>(op.getValue().getType());
589 bool vnni = op.getPacked().value_or(
false);
590 auto transposeValue = op.getTranspose();
592 transposeValue.has_value() && transposeValue.value()[0] == 1;
598 if (elemBitSize == 8 && tileW == 16 && tileH == 32 && !vnni &&
606 if (transpose && elemBitSize < 32) {
607 int32_t scale = 32 / elemBitSize;
609 rewriter, loc, rewriter.getI32Type(), llvm::Log2_64(scale));
610 offsetW = arith::ShRSIOp::create(rewriter, loc, offsetW, scaleLog2);
611 tileW = tileW * elemBitSize / 32;
614 VectorType loadedTy = encodeVectorTypeTo(
615 dstVecTy, vnni ? rewriter.getI32Type()
616 : rewriter.getIntegerType(elemBitSize));
618 Value resultFlatVec = xevm::BlockLoad2dOp::create(
619 rewriter, loc, loadedTy, basePtrLLVM, baseShapeWInBytes,
620 baseShapeH, basePitchBytes, offsetW, offsetH, elemBitSize, tileW,
621 tileH, vblocks, transpose, vnni,
622 xevm::LoadCacheControlAttr::get(ctxt, loadCacheControl));
623 resultFlatVec = vector::BitCastOp::create(
625 encodeVectorTypeTo(loadedTy, dstVecTy.getElementType()),
627 rewriter.replaceOp(op, resultFlatVec);
639 rewriter.getI64Type(), offset);
642 rewriter, loc, rewriter.getI64Type(), elemBitSize / 8);
644 rewriter.createOrFold<arith::MulIOp>(loc, offset, elemByteSize);
646 Value finalAddrI64 = rewriter.createOrFold<arith::AddIOp>(
652 LLVM::IntToPtrOp::create(rewriter, loc, ptrTypeLLVM, finalAddrI64);
653 if constexpr (std::is_same_v<OpType, xegpu::StoreNdOp>) {
654 Value src = adaptor.getValue();
660 VectorType srcVecTy = dyn_cast<VectorType>(src.
getType());
662 return rewriter.notifyMatchFailure(
663 op,
"Expected store value to be a vector type.");
665 VectorType newSrcVecTy =
666 encodeVectorTypeTo(srcVecTy, rewriter.getIntegerType(elemBitSize));
667 if (srcVecTy != newSrcVecTy)
668 src = vector::BitCastOp::create(rewriter, loc, newSrcVecTy, src);
669 auto storeCacheControl =
670 translateStoreXeGPUCacheHint(op.getL1Hint(), op.getL3Hint());
671 rewriter.replaceOpWithNewOp<xevm::BlockStoreOp>(
672 op, finalPtrLLVM, src,
673 xevm::StoreCacheControlAttr::get(ctxt, storeCacheControl));
674 }
else if constexpr (std::is_same_v<OpType, xegpu::LoadNdOp>) {
675 auto loadCacheControl =
676 translateLoadXeGPUCacheHint(op.getL1Hint(), op.getL3Hint());
677 VectorType resTy = cast<VectorType>(op.getValue().getType());
678 VectorType loadedTy =
679 encodeVectorTypeTo(resTy, rewriter.getIntegerType(elemBitSize));
680 Value
load = xevm::BlockLoadOp::create(
681 rewriter, loc, loadedTy, finalPtrLLVM,
682 xevm::LoadCacheControlAttr::get(ctxt, loadCacheControl));
683 if (loadedTy != resTy)
684 load = vector::BitCastOp::create(rewriter, loc, resTy,
load);
685 rewriter.replaceOp(op,
load);
687 return rewriter.notifyMatchFailure(
688 op,
"Unsupported operation: xegpu.prefetch_nd with tensor "
689 "descriptor rank == 1");
698static Value addOffsetToBaseAddr(ConversionPatternRewriter &rewriter,
702 rewriter, loc, baseAddr.
getType(), elemByteSize);
703 Value byteOffset = arith::MulIOp::create(rewriter, loc, offset, byteSize);
704 Value newAddr = arith::AddIOp::create(rewriter, loc, baseAddr, byteOffset);
708template <
typename OpType,
709 typename = std::enable_if_t<llvm::is_one_of<
710 OpType, xegpu::LoadGatherOp, xegpu::StoreScatterOp>::value>>
711class LoadStoreToXeVMPattern :
public OpConversionPattern<OpType> {
712 using OpConversionPattern<OpType>::OpConversionPattern;
714 matchAndRewrite(OpType op,
typename OpType::Adaptor adaptor,
715 ConversionPatternRewriter &rewriter)
const override {
716 Value offset = adaptor.getOffsets();
718 return rewriter.notifyMatchFailure(op,
"Expected offset to be provided.");
719 auto loc = op.getLoc();
720 auto ctxt = rewriter.getContext();
724 if constexpr (std::is_same_v<OpType, xegpu::LoadGatherOp>)
726 this->getTypeConverter()->convertType(op.getResult().getType());
728 valOrResTy = adaptor.getValue().getType();
729 VectorType valOrResVecTy = dyn_cast<VectorType>(valOrResTy);
730 bool hasScalarVal = !valOrResVecTy;
731 int64_t elemBitWidth =
733 : valOrResVecTy.getElementType().getIntOrFloatBitWidth();
735 if (elemBitWidth % 8 != 0)
736 return rewriter.notifyMatchFailure(
737 op,
"Expected element type bit width to be multiple of 8.");
738 int64_t elemByteSize = elemBitWidth / 8;
740 LLVM::LLVMPointerType ptrTypeLLVM = LLVM::LLVMPointerType::get(
741 ctxt, getNumericXeVMAddrSpace(xegpu::MemorySpace::Global));
744 if constexpr (std::is_same_v<OpType, xegpu::LoadGatherOp>) {
745 basePtrI64 = adaptor.getSource();
746 if (
auto memRefTy = dyn_cast<MemRefType>(op.getSource().getType())) {
747 FailureOr<unsigned> addrSpace =
748 getNumericMemorySpace(memRefTy.getMemorySpace());
750 return rewriter.notifyMatchFailure(
751 op,
"Unsupported memref memory space attribute.");
753 ptrTypeLLVM = LLVM::LLVMPointerType::get(ctxt, *addrSpace);
756 basePtrI64 = adaptor.getDest();
757 if (
auto memRefTy = dyn_cast<MemRefType>(op.getDest().getType())) {
758 FailureOr<unsigned> addrSpace =
759 getNumericMemorySpace(memRefTy.getMemorySpace());
761 return rewriter.notifyMatchFailure(
762 op,
"Unsupported memref memory space attribute.");
764 ptrTypeLLVM = LLVM::LLVMPointerType::get(ctxt, *addrSpace);
768 if (basePtrI64.
getType() != rewriter.getI64Type()) {
769 basePtrI64 = arith::ExtUIOp::create(rewriter, loc, rewriter.getI64Type(),
772 Value mask = adaptor.getMask();
773 if (dyn_cast<VectorType>(offset.
getType())) {
776 return rewriter.notifyMatchFailure(op,
"Expected offset to be a scalar.");
782 addOffsetToBaseAddr(rewriter, loc, basePtrI64, offset, elemByteSize);
786 LLVM::IntToPtrOp::create(rewriter, loc, ptrTypeLLVM, basePtrI64);
789 VectorType maskVecTy = dyn_cast<VectorType>(mask.
getType());
793 return rewriter.notifyMatchFailure(op,
"Expected mask to be a scalar.");
796 if constexpr (std::is_same_v<OpType, xegpu::LoadGatherOp>) {
797 scf::IfOp ifOp = scf::IfOp::create(rewriter, loc, {valOrResTy},
798 maskForLane,
true,
true);
800 rewriter.setInsertionPointToStart(&ifOp.getThenRegion().front());
802 valOrResTy = VectorType::get({valOrResVecTy.getNumElements()},
803 valOrResVecTy.getElementType());
805 LLVM::LoadOp::create(rewriter, loc, valOrResTy, basePtrLLVM);
808 "cache_control", xevm::LoadCacheControlAttr::get(
809 ctxt, translateLoadXeGPUCacheHint(
810 op.getL1Hint(), op.getL3Hint())));
811 scf::YieldOp::create(rewriter, loc,
ValueRange{loaded});
812 rewriter.setInsertionPointToStart(&ifOp.getElseRegion().front());
814 auto eTy = hasScalarVal ? valOrResTy : valOrResVecTy.getElementType();
817 eVal = FloatAttr::get(eTy, 0.0);
819 eVal = IntegerAttr::get(eTy, 0);
821 loaded = arith::ConstantOp::create(rewriter, loc, eVal);
823 loaded = arith::ConstantOp::create(
825 scf::YieldOp::create(rewriter, loc,
ValueRange{loaded});
826 rewriter.replaceOp(op, ifOp.getResult(0));
829 scf::IfOp ifOp = scf::IfOp::create(rewriter, loc, maskForLane,
false);
830 auto body = ifOp.getBody();
831 rewriter.setInsertionPointToStart(body);
833 LLVM::StoreOp::create(rewriter, loc, adaptor.getValue(), basePtrLLVM);
835 storeOp.getOperation()->setDiscardableAttr(
836 "cache_control", xevm::StoreCacheControlAttr::get(
837 ctxt, translateStoreXeGPUCacheHint(
838 op.getL1Hint(), op.getL3Hint())));
839 rewriter.eraseOp(op);
845class CreateMemDescOpPattern final
846 :
public OpConversionPattern<xegpu::CreateMemDescOp> {
848 using OpConversionPattern<xegpu::CreateMemDescOp>::OpConversionPattern;
850 matchAndRewrite(xegpu::CreateMemDescOp op, OpAdaptor adaptor,
851 ConversionPatternRewriter &rewriter)
const override {
853 rewriter.replaceOp(op, adaptor.getSource());
858template <
typename OpType,
859 typename = std::enable_if_t<llvm::is_one_of<
860 OpType, xegpu::LoadMatrixOp, xegpu::StoreMatrixOp>::value>>
861class LoadStoreMatrixToXeVMPattern :
public OpConversionPattern<OpType> {
862 using OpConversionPattern<OpType>::OpConversionPattern;
864 matchAndRewrite(OpType op,
typename OpType::Adaptor adaptor,
865 ConversionPatternRewriter &rewriter)
const override {
867 SmallVector<OpFoldResult> offsets = op.getMixedOffsets();
869 return rewriter.notifyMatchFailure(op,
"Expected offset to be provided.");
871 auto loc = op.getLoc();
872 auto ctxt = rewriter.getContext();
873 Value baseAddr32 = adaptor.getMemDesc();
874 Value mdescVal = op.getMemDesc();
877 if constexpr (std::is_same_v<OpType, xegpu::LoadMatrixOp>) {
878 Type resType = op.getResult().getType();
881 if (
auto vecType = dyn_cast<VectorType>(resType)) {
882 assert(llvm::count_if(vecType.getShape(),
883 [](int64_t d) { return d != 1; }) <= 1 &&
884 "Expected either 1D vector or nD with unit dimensions");
885 resType = VectorType::get({vecType.getNumElements()},
886 vecType.getElementType());
890 dataTy = adaptor.getData().getType();
891 VectorType valOrResVecTy = dyn_cast<VectorType>(dataTy);
893 valOrResVecTy = VectorType::get(1, dataTy);
895 int64_t elemBitWidth =
896 valOrResVecTy.getElementType().getIntOrFloatBitWidth();
898 if (elemBitWidth % 8 != 0)
899 return rewriter.notifyMatchFailure(
900 op,
"Expected element type bit width to be multiple of 8.");
901 int64_t elemByteSize = elemBitWidth / 8;
904 LLVM::LLVMPointerType ptrTypeLLVM = LLVM::LLVMPointerType::get(
905 ctxt, getNumericXeVMAddrSpace(xegpu::MemorySpace::SLM));
907 auto mdescTy = cast<xegpu::MemDescType>(mdescVal.
getType());
909 Value linearOffset = mdescTy.getLinearOffsets(rewriter, loc, offsets);
910 linearOffset = arith::IndexCastUIOp::create(
911 rewriter, loc, rewriter.getI32Type(), linearOffset);
912 Value basePtrI32 = addOffsetToBaseAddr(rewriter, loc, baseAddr32,
913 linearOffset, elemByteSize);
917 LLVM::IntToPtrOp::create(rewriter, loc, ptrTypeLLVM, basePtrI32);
919 if (op.getSubgroupBlockIoAttr()) {
923 Type intElemTy = rewriter.getIntegerType(elemBitWidth);
924 VectorType intVecTy =
925 VectorType::get(valOrResVecTy.getShape(), intElemTy);
927 if constexpr (std::is_same_v<OpType, xegpu::LoadMatrixOp>) {
929 xevm::BlockLoadOp::create(rewriter, loc, intVecTy, basePtrLLVM);
930 if (intVecTy != valOrResVecTy) {
932 vector::BitCastOp::create(rewriter, loc, valOrResVecTy, loadOp);
934 rewriter.replaceOp(op, loadOp);
936 Value dataToStore = adaptor.getData();
937 if (valOrResVecTy != intVecTy) {
939 vector::BitCastOp::create(rewriter, loc, intVecTy, dataToStore);
941 xevm::BlockStoreOp::create(rewriter, loc, basePtrLLVM, dataToStore,
943 rewriter.eraseOp(op);
948 if (valOrResVecTy.getNumElements() >= 1) {
951 (*chipOpt !=
"pvc" && *chipOpt !=
"bmg" && *chipOpt !=
"cri")) {
953 return rewriter.notifyMatchFailure(
954 op,
"The lowering is specific to pvc, bmg or cri.");
958 if constexpr (std::is_same_v<OpType, xegpu::LoadMatrixOp>) {
965 this->getTypeConverter()->convertType(op.getResult().getType());
966 auto loadOp = LLVM::LoadOp::create(rewriter, loc, loadTy, basePtrLLVM);
967 rewriter.replaceOp(op, loadOp);
969 LLVM::StoreOp::create(rewriter, loc, adaptor.getData(), basePtrLLVM);
970 rewriter.eraseOp(op);
976class PrefetchToXeVMPattern :
public OpConversionPattern<xegpu::PrefetchOp> {
977 using OpConversionPattern::OpConversionPattern;
979 matchAndRewrite(xegpu::PrefetchOp op, xegpu::PrefetchOp::Adaptor adaptor,
980 ConversionPatternRewriter &rewriter)
const override {
981 auto loc = op.getLoc();
982 auto ctxt = rewriter.getContext();
983 Value basePtrI64 = adaptor.getSource();
985 if (basePtrI64.
getType() != rewriter.getI64Type())
986 basePtrI64 = arith::ExtUIOp::create(rewriter, loc, rewriter.getI64Type(),
988 Value offsets = adaptor.getOffsets();
990 VectorType offsetsVecTy = dyn_cast<VectorType>(offsets.
getType());
993 return rewriter.notifyMatchFailure(op,
994 "Expected offsets to be a scalar.");
996 int64_t elemBitWidth{0};
997 int64_t elemByteSize;
999 if (
auto memRefTy = dyn_cast<MemRefType>(op.getSourceType())) {
1002 elemBitWidth = memRefTy.getElementType().getIntOrFloatBitWidth();
1005 elemByteSize = *op.getOffsetAlignByte();
1007 if (elemBitWidth != 0) {
1008 if (elemBitWidth % 8 != 0)
1009 return rewriter.notifyMatchFailure(
1010 op,
"Expected element type bit width to be multiple of 8.");
1011 elemByteSize = elemBitWidth / 8;
1013 basePtrI64 = addOffsetToBaseAddr(rewriter, loc, basePtrI64, offsets,
1018 LLVM::LLVMPointerType ptrTypeLLVM = LLVM::LLVMPointerType::get(
1019 ctxt, getNumericXeVMAddrSpace(xegpu::MemorySpace::Global));
1021 if (
auto memRefTy = dyn_cast<MemRefType>(op.getSource().getType())) {
1022 FailureOr<unsigned> addrSpace =
1023 getNumericMemorySpace(memRefTy.getMemorySpace());
1025 return rewriter.notifyMatchFailure(
1026 op,
"Unsupported memref memory space attribute.");
1027 if (*addrSpace != 0)
1028 ptrTypeLLVM = LLVM::LLVMPointerType::get(ctxt, *addrSpace);
1032 LLVM::IntToPtrOp::create(rewriter, loc, ptrTypeLLVM, basePtrI64);
1034 xevm::PrefetchOp::create(
1035 rewriter, loc, ptrLLVM,
1036 xevm::LoadCacheControlAttr::get(
1037 ctxt, translateLoadXeGPUCacheHint(op.getL1Hint(), op.getL3Hint())));
1038 rewriter.eraseOp(op);
1043class FenceToXeVMPattern :
public OpConversionPattern<xegpu::FenceOp> {
1044 using OpConversionPattern::OpConversionPattern;
1046 matchAndRewrite(xegpu::FenceOp op, xegpu::FenceOp::Adaptor adaptor,
1047 ConversionPatternRewriter &rewriter)
const override {
1048 auto loc = op.getLoc();
1049 xevm::MemScope memScope{xevm::MemScope::WORKGROUP};
1050 switch (op.getFenceScope()) {
1051 case xegpu::FenceScope::Workgroup:
1052 memScope = xevm::MemScope::WORKGROUP;
1054 case xegpu::FenceScope::GPU:
1055 memScope = xevm::MemScope::DEVICE;
1058 xevm::AddrSpace addrSpace{xevm::AddrSpace::GLOBAL};
1059 switch (op.getMemoryKind()) {
1060 case xegpu::MemorySpace::Global:
1061 addrSpace = xevm::AddrSpace::GLOBAL;
1063 case xegpu::MemorySpace::SLM:
1064 addrSpace = xevm::AddrSpace::SHARED;
1067 xevm::MemfenceOp::create(rewriter, loc, memScope, addrSpace);
1068 rewriter.eraseOp(op);
1073static auto encodePrecision = [](
Type type) -> xevm::ElemType {
1075 return xevm::ElemType::BF16;
1076 else if (type.isF16())
1077 return xevm::ElemType::F16;
1078 else if (type.isTF32())
1079 return xevm::ElemType::TF32;
1080 else if (type.isInteger(8)) {
1081 if (type.isUnsignedInteger())
1082 return xevm::ElemType::U8;
1083 return xevm::ElemType::S8;
1084 }
else if (type.isF32())
1085 return xevm::ElemType::F32;
1086 else if (type.isInteger(32))
1087 return xevm::ElemType::S32;
1088 else if (type.isF8E5M2())
1089 return xevm::ElemType::BF8;
1090 else if (type.isF8E4M3FN())
1091 return xevm::ElemType::F8;
1092 else if (mlir::isa<Float4E2M1FNType>(type))
1093 return xevm::ElemType::E2M1;
1094 llvm_unreachable(
"add more support for ElemType");
1097static unsigned getNumOperandsPerDword(xevm::ElemType pTy) {
1099 case xevm::ElemType::TF32:
1101 case xevm::ElemType::BF16:
1102 case xevm::ElemType::F16:
1104 case xevm::ElemType::U8:
1105 case xevm::ElemType::S8:
1106 case xevm::ElemType::F8:
1107 case xevm::ElemType::BF8:
1109 case xevm::ElemType::E2M1:
1112 llvm_unreachable(
"unsupported xevm::ElemType");
1116class DpasToXeVMPattern :
public OpConversionPattern<xegpu::DpasOp> {
1117 using OpConversionPattern::OpConversionPattern;
1119 matchAndRewrite(xegpu::DpasOp op, xegpu::DpasOp::Adaptor adaptor,
1120 ConversionPatternRewriter &rewriter)
const override {
1121 auto loc = op.getLoc();
1122 auto ctxt = rewriter.getContext();
1123 auto aTy = cast<VectorType>(op.getLhs().getType());
1124 auto bTy = cast<VectorType>(op.getRhs().getType());
1125 auto resultType = cast<VectorType>(op.getResultType());
1130 return rewriter.notifyMatchFailure(op,
"cannot determine target chip");
1134 return rewriter.notifyMatchFailure(op,
"unsupported target uArch");
1137 llvm::dyn_cast_or_null<xegpu::uArch::SubgroupMatrixMultiplyAcc>(
1138 uArch->getInstruction(
1139 xegpu::uArch::InstructionKind::SubgroupMatrixMultiplyAcc)));
1141 return rewriter.notifyMatchFailure(op,
1142 "DPAS not supported by target uArch");
1144 auto checkSupportedTypes = [&](VectorType vecTy,
1146 auto supported = dpasInst->getSupportedTypes(*ctxt, kind);
1147 return llvm::find(supported, vecTy.getElementType()) != supported.end();
1150 if (!checkSupportedTypes(aTy, xegpu::uArch::MMAOpndKind::MatrixA))
1151 return rewriter.notifyMatchFailure(
1152 op,
"A-matrix element type not supported by target uArch");
1153 if (!checkSupportedTypes(bTy, xegpu::uArch::MMAOpndKind::MatrixB))
1154 return rewriter.notifyMatchFailure(
1155 op,
"B-matrix element type not supported by target uArch");
1157 if (!checkSupportedTypes(resultType, xegpu::uArch::MMAOpndKind::MatrixD))
1158 return rewriter.notifyMatchFailure(
1159 op,
"result/accumulator element type not supported by target uArch");
1161 xevm::ElemType precATy = encodePrecision(aTy.getElementType());
1162 xevm::ElemType precBTy = encodePrecision(bTy.getElementType());
1163 Value c = op.getAcc();
1165 auto elementTy = resultType.getElementType();
1166 Attribute initValueAttr;
1167 if (isa<FloatType>(elementTy))
1168 initValueAttr = FloatAttr::get(elementTy, 0.0);
1170 initValueAttr = IntegerAttr::get(elementTy, 0);
1171 c = arith::ConstantOp::create(
1175 Value aVec = op.getLhs();
1176 Value bVec = op.getRhs();
1177 auto cvecty = cast<VectorType>(c.
getType());
1178 xevm::ElemType precCTy = encodePrecision(cvecty.getElementType());
1179 xevm::ElemType precDTy = encodePrecision(resultType.getElementType());
1181 VectorType::get(cvecty.getNumElements(), cvecty.getElementType());
1183 c = vector::ShapeCastOp::create(rewriter, loc, cNty, c);
1184 Value dpasRes = xevm::MMAOp::create(
1185 rewriter, loc, cNty, aVec, bVec, c,
1186 xevm::MMAShapeAttr::get(ctxt, cvecty.getNumElements(), executionSize,
1188 getNumOperandsPerDword(precATy)),
1189 xevm::MMATypesAttr::get(ctxt, precDTy, precATy, precBTy, precCTy));
1191 dpasRes = vector::ShapeCastOp::create(rewriter, loc, resultType, dpasRes);
1192 rewriter.replaceOp(op, dpasRes);
1197static std::optional<LLVM::AtomicBinOp>
1198matchSimpleAtomicOp(arith::AtomicRMWKind arithKind) {
1199 switch (arithKind) {
1200 case arith::AtomicRMWKind::addf:
1201 return LLVM::AtomicBinOp::fadd;
1202 case arith::AtomicRMWKind::addi:
1203 return LLVM::AtomicBinOp::add;
1204 case arith::AtomicRMWKind::assign:
1205 return LLVM::AtomicBinOp::xchg;
1206 case arith::AtomicRMWKind::maximumf:
1207 return LLVM::AtomicBinOp::fmax;
1208 case arith::AtomicRMWKind::maxs:
1209 return LLVM::AtomicBinOp::max;
1210 case arith::AtomicRMWKind::maxu:
1211 return LLVM::AtomicBinOp::umax;
1212 case arith::AtomicRMWKind::minimumf:
1213 return LLVM::AtomicBinOp::fmin;
1214 case arith::AtomicRMWKind::mins:
1215 return LLVM::AtomicBinOp::min;
1216 case arith::AtomicRMWKind::minu:
1217 return LLVM::AtomicBinOp::umin;
1218 case arith::AtomicRMWKind::ori:
1219 return LLVM::AtomicBinOp::_or;
1220 case arith::AtomicRMWKind::andi:
1221 return LLVM::AtomicBinOp::_and;
1223 return std::nullopt;
1227class AtomicRMWToXeVMPattern :
public OpConversionPattern<xegpu::AtomicRMWOp> {
1228 using OpConversionPattern::OpConversionPattern;
1230 matchAndRewrite(xegpu::AtomicRMWOp op, xegpu::AtomicRMWOp::Adaptor adaptor,
1231 ConversionPatternRewriter &rewriter)
const override {
1232 auto loc = op.getLoc();
1233 auto ctxt = rewriter.getContext();
1234 auto tdesc = op.getTensorDesc().getType();
1235 auto ptrTypeLLVM = LLVM::LLVMPointerType::get(
1236 ctxt, getNumericXeVMAddrSpace(tdesc.getMemorySpace()));
1237 Value basePtrI64 = arith::IndexCastOp::create(
1238 rewriter, loc, rewriter.getI64Type(), adaptor.getTensorDesc());
1240 LLVM::IntToPtrOp::create(rewriter, loc, ptrTypeLLVM, basePtrI64);
1241 VectorType srcOrDstVecTy = cast<VectorType>(op.getValue().getType());
1242 VectorType srcOrDstFlatVecTy = VectorType::get(
1243 srcOrDstVecTy.getNumElements(), srcOrDstVecTy.getElementType());
1244 Value srcFlatVec = vector::ShapeCastOp::create(
1245 rewriter, loc, srcOrDstFlatVecTy, op.getValue());
1246 auto atomicKind = matchSimpleAtomicOp(op.getKind());
1247 assert(atomicKind.has_value());
1248 Value resVec = srcFlatVec;
1249 for (
int i = 0; i < srcOrDstVecTy.getNumElements(); i++) {
1250 auto val = vector::ExtractOp::create(rewriter, loc, resVec, i);
1251 Value idx = LLVM::ConstantOp::create(rewriter, loc, rewriter.getI64Type(),
1252 rewriter.getI64IntegerAttr(i));
1254 LLVM::GEPOp::create(rewriter, loc, ptrTypeLLVM,
1255 srcOrDstVecTy.getElementType(), basePtrLLVM, idx);
1257 LLVM::AtomicRMWOp::create(rewriter, loc, atomicKind.value(), currPtr,
1258 val, LLVM::AtomicOrdering::seq_cst);
1259 resVec = vector::InsertOp::create(rewriter, loc, newVal, resVec, i);
1261 rewriter.replaceOp(op, resVec);
1266class DpasMxToXeVMPattern :
public OpConversionPattern<xegpu::DpasMxOp> {
1267 using OpConversionPattern::OpConversionPattern;
1269 matchAndRewrite(xegpu::DpasMxOp op, xegpu::DpasMxOp::Adaptor adaptor,
1270 ConversionPatternRewriter &rewriter)
const override {
1271 auto loc = op.getLoc();
1272 auto ctxt = rewriter.getContext();
1273 auto aTy = op.getA().getType();
1274 auto bTy = op.getB().getType();
1276 cast<VectorType>(getTypeConverter()->convertType(op.getType()));
1280 return rewriter.notifyMatchFailure(op,
"cannot determine target chip");
1284 return rewriter.notifyMatchFailure(op,
"unsupported target uArch");
1288 xevm::ElemType precATy = encodePrecision(aTy.getElementType());
1289 xevm::ElemType precBTy = encodePrecision(bTy.getElementType());
1290 Value c = adaptor.getAcc();
1292 auto elementTy = resVecTy.getElementType();
1293 Attribute initValueAttr;
1294 if (isa<FloatType>(elementTy))
1295 initValueAttr = FloatAttr::get(elementTy, 0.0);
1297 initValueAttr = IntegerAttr::get(elementTy, 0);
1298 c = arith::ConstantOp::create(
1302 Value aVec = adaptor.getA();
1303 Value bVec = adaptor.getB();
1304 auto aVecTy = cast<VectorType>(aVec.
getType());
1305 auto bVecTy = cast<VectorType>(bVec.
getType());
1306 if (aVecTy.getElementTypeBitWidth() == 4)
1307 aVec = vector::BitCastOp::create(
1309 VectorType::get(aVecTy.getNumElements() / 2, rewriter.getI8Type()),
1311 if (bVecTy.getElementTypeBitWidth() == 4)
1312 bVec = vector::BitCastOp::create(
1314 VectorType::get(bVecTy.getNumElements() / 2, rewriter.getI8Type()),
1316 auto cVecTy = cast<VectorType>(c.
getType());
1317 xevm::ElemType precCTy = encodePrecision(cVecTy.getElementType());
1318 xevm::ElemType precDTy = encodePrecision(resVecTy.getElementType());
1319 Value scaleA = adaptor.getScaleA();
1320 Value scaleB = adaptor.getScaleB();
1321 Value dpasMxRes = xevm::MMAMxOp::create(
1322 rewriter, loc, resVecTy, aVec, bVec, scaleA, scaleB, c,
1323 xevm::MMAShapeAttr::get(ctxt, cVecTy.getNumElements(), executionSize,
1325 getNumOperandsPerDword(precATy)),
1326 xevm::MMATypesAttr::get(ctxt, precDTy, precATy, precBTy, precCTy));
1327 rewriter.replaceOp(op, dpasMxRes);
1346static constexpr int64_t kXeVMExtfTruncfNumElems = 16;
1349static std::optional<xevm::ExtfSrcElemTypes> getExtfNarrowType(
Type etype) {
1350 if (isa<Float8E5M2Type>(etype))
1351 return xevm::ExtfSrcElemTypes::BF8;
1352 if (isa<Float8E4M3FNType>(etype))
1353 return xevm::ExtfSrcElemTypes::F8;
1354 if (isa<Float4E2M1FNType>(etype))
1355 return xevm::ExtfSrcElemTypes::E2M1;
1356 return std::nullopt;
1360static std::optional<xevm::TruncfDstElemTypes> getTruncfNarrowType(
Type etype) {
1361 if (isa<Float8E5M2Type>(etype))
1362 return xevm::TruncfDstElemTypes::BF8;
1363 if (isa<Float8E4M3FNType>(etype))
1364 return xevm::TruncfDstElemTypes::F8;
1365 if (isa<Float4E2M1FNType>(etype))
1366 return xevm::TruncfDstElemTypes::E2M1;
1367 return std::nullopt;
1372static bool isXeVMExtf(arith::ExtFOp op) {
1373 auto srcTy = dyn_cast<VectorType>(op.getIn().getType());
1374 auto dstTy = dyn_cast<VectorType>(op.getType());
1375 if (!srcTy || !dstTy || srcTy.getRank() != 1 || dstTy.getRank() != 1)
1377 if (dstTy.getNumElements() != kXeVMExtfTruncfNumElems)
1379 Type dstETy = dstTy.getElementType();
1382 return getExtfNarrowType(srcTy.getElementType()).has_value();
1389static bool isXeVMTruncf(arith::TruncFOp op) {
1390 auto srcTy = dyn_cast<VectorType>(op.getIn().getType());
1391 auto dstTy = dyn_cast<VectorType>(op.getType());
1392 if (!srcTy || !dstTy || srcTy.getRank() != 1 || dstTy.getRank() != 1)
1394 int64_t numElems = srcTy.getNumElements();
1395 if (numElems == 0 || numElems % kXeVMExtfTruncfNumElems != 0)
1397 Type srcETy = srcTy.getElementType();
1400 return getTruncfNarrowType(dstTy.getElementType()).has_value();
1403class ExtfToXeVMPattern :
public OpConversionPattern<arith::ExtFOp> {
1404 using OpConversionPattern::OpConversionPattern;
1406 matchAndRewrite(arith::ExtFOp op, OpAdaptor adaptor,
1407 ConversionPatternRewriter &rewriter)
const override {
1408 if (!isXeVMExtf(op))
1409 return rewriter.notifyMatchFailure(op,
"not a xevm.extf compatible extf");
1410 Location loc = op.getLoc();
1411 MLIRContext *ctx = op.getContext();
1412 auto srcVecTy = cast<VectorType>(op.getIn().getType());
1413 auto dstVecTy = cast<VectorType>(op.getType());
1414 xevm::ExtfSrcElemTypes srcEnum =
1415 *getExtfNarrowType(srcVecTy.getElementType());
1416 xevm::ExtfDstElemTypes dstEnum = dstVecTy.getElementType().isF16()
1417 ? xevm::ExtfDstElemTypes::F16
1418 : xevm::ExtfDstElemTypes::BF16;
1422 Value src = adaptor.getIn();
1423 auto convSrcTy = cast<VectorType>(src.
getType());
1424 if (convSrcTy.getElementTypeBitWidth() == 4)
1425 src = vector::BitCastOp::create(
1427 VectorType::get(convSrcTy.getNumElements() / 2, rewriter.getI8Type()),
1429 Type resTy = getTypeConverter()->convertType(dstVecTy);
1430 Value res = xevm::ExtfOp::create(
1431 rewriter, loc, resTy, src, xevm::ExtfSrcElemTypeAttr::get(ctx, srcEnum),
1432 xevm::ExtfDstElemTypeAttr::get(ctx, dstEnum));
1433 rewriter.replaceOp(op, res);
1438class TruncfToXeVMPattern :
public OpConversionPattern<arith::TruncFOp> {
1439 using OpConversionPattern::OpConversionPattern;
1441 matchAndRewrite(arith::TruncFOp op, OpAdaptor adaptor,
1442 ConversionPatternRewriter &rewriter)
const override {
1443 if (!isXeVMTruncf(op))
1444 return rewriter.notifyMatchFailure(op,
1445 "not a xevm.truncf compatible truncf");
1446 Location loc = op.getLoc();
1447 MLIRContext *ctx = op.getContext();
1448 auto srcVecTy = cast<VectorType>(op.getIn().getType());
1449 auto dstVecTy = cast<VectorType>(op.getType());
1450 xevm::TruncfSrcElemTypes srcEnum = srcVecTy.getElementType().isF16()
1451 ? xevm::TruncfSrcElemTypes::F16
1452 : xevm::TruncfSrcElemTypes::BF16;
1453 xevm::TruncfDstElemTypes dstEnum =
1454 *getTruncfNarrowType(dstVecTy.getElementType());
1455 auto srcEnumAttr = xevm::TruncfSrcElemTypeAttr::get(ctx, srcEnum);
1456 auto dstEnumAttr = xevm::TruncfDstElemTypeAttr::get(ctx, dstEnum);
1462 int64_t numGroups = srcVecTy.getNumElements() / kXeVMExtfTruncfNumElems;
1463 int64_t groupBytes =
1464 kXeVMExtfTruncfNumElems * dstVecTy.getElementTypeBitWidth() / 8;
1465 Type groupTy = VectorType::get(groupBytes, rewriter.getI8Type());
1467 Value src = adaptor.getIn();
1469 if (numGroups == 1) {
1470 packed = xevm::TruncfOp::create(rewriter, loc, groupTy, src, srcEnumAttr,
1474 VectorType::get(groupBytes * numGroups, rewriter.getI8Type());
1475 packed = arith::ConstantOp::create(rewriter, loc, packedTy,
1476 rewriter.getZeroAttr(packedTy));
1477 for (int64_t group = 0; group < numGroups; group++) {
1478 Value slice = vector::ExtractStridedSliceOp::create(
1479 rewriter, loc, src, group * kXeVMExtfTruncfNumElems,
1480 kXeVMExtfTruncfNumElems, 1);
1481 Value converted = xevm::TruncfOp::create(rewriter, loc, groupTy, slice,
1482 srcEnumAttr, dstEnumAttr);
1483 packed = vector::InsertStridedSliceOp::create(
1484 rewriter, loc, converted, packed, group * groupBytes,
1489 Type resTy = getTypeConverter()->convertType(dstVecTy);
1490 if (packed.
getType() != resTy)
1491 packed = vector::BitCastOp::create(rewriter, loc, resTy, packed);
1492 rewriter.replaceOp(op, packed);
1516class LaneShuffleToXeVMPattern
1517 :
public OpConversionPattern<xegpu::LaneShuffleOp> {
1518 using OpConversionPattern::OpConversionPattern;
1520 matchAndRewrite(xegpu::LaneShuffleOp op, OpAdaptor adaptor,
1521 ConversionPatternRewriter &rewriter)
const override {
1522 auto vecTy = dyn_cast<VectorType>(adaptor.getSource().getType());
1524 return rewriter.notifyMatchFailure(op,
"Expected a vector fragment.");
1528 unsigned elemBits = vecTy.getElementTypeBitWidth();
1529 if (elemBits != 8 && elemBits != 16 && elemBits != 32 && elemBits != 64)
1530 return rewriter.notifyMatchFailure(
1531 op,
"Expected an element type of 8, 16, 32 or 64 bits.");
1532 int64_t fragmentBits = vecTy.getNumElements() * elemBits;
1533 if (fragmentBits > 64 || !llvm::isPowerOf2_64(fragmentBits))
1534 return rewriter.notifyMatchFailure(
1535 op,
"Expected a fragment of 8, 16, 32 or 64 bits.");
1537 Location loc = op.getLoc();
1538 Type packedTy = rewriter.getIntegerType(fragmentBits);
1541 VectorType shuffleTy =
1542 VectorType::get(vecTy.getShape(), rewriter.getIntegerType(elemBits));
1545 if (op.getMode() == xegpu::LaneShuffleMode::Pack) {
1546 Value src = adaptor.getSource();
1547 if (shuffleTy != vecTy)
1548 src = LLVM::BitcastOp::create(rewriter, loc, shuffleTy, src);
1549 res = xevm::BitcastShuffleOp::create(rewriter, loc, packedTy, src);
1550 res = LLVM::BitcastOp::create(rewriter, loc, vecTy, res);
1553 LLVM::BitcastOp::create(rewriter, loc, packedTy, adaptor.getSource());
1554 res = xevm::BitcastShuffleOp::create(rewriter, loc, shuffleTy, packed);
1555 if (shuffleTy != vecTy)
1556 res = LLVM::BitcastOp::create(rewriter, loc, vecTy, res);
1558 rewriter.replaceOp(op, res);
1567struct ConvertXeGPUToXeVMPass
1568 :
public impl::ConvertXeGPUToXeVMPassBase<ConvertXeGPUToXeVMPass> {
1571 void runOnOperation()
override {
1581 LowerToLLVMOptions
options(context);
1582 options.overrideIndexBitwidth(this->use64bitIndex ? 64 : 32);
1583 LLVMTypeConverter typeConverter(context,
options);
1585 Type xevmIndexType = typeConverter.convertType(IndexType::get(context));
1586 Type i32Type = IntegerType::get(context, 32);
1587 typeConverter.addConversion([&](VectorType type) -> Type {
1588 auto elemType = typeConverter.convertType(type.getElementType());
1590 unsigned rank = type.getRank();
1591 if (rank == 0 || type.getNumElements() == 1)
1594 int64_t sum = llvm::product_of(type.getShape());
1595 return VectorType::get(sum, elemType);
1597 typeConverter.addConversion([&](xegpu::TensorDescType type) -> Type {
1598 if (type.getRank() == 1)
1599 return xevmIndexType;
1600 return VectorType::get(8, i32Type);
1609 typeConverter.addConversion(
1610 [&](xegpu::MemDescType type) -> Type {
return i32Type; });
1612 typeConverter.addConversion([&](MemRefType type) -> Type {
1613 return isSharedMemRef(type) ? i32Type : xevmIndexType;
1623 auto memrefToIntMaterializationCast = [](OpBuilder &builder, Type type,
1625 Location loc) -> Value {
1626 if (inputs.size() != 1)
1628 auto input = inputs.front();
1629 if (
auto memrefTy = dyn_cast<MemRefType>(input.getType())) {
1630 unsigned rank = memrefTy.getRank();
1634 SmallVector<int64_t> intStrides;
1637 if (succeeded(memrefTy.getStridesAndOffset(intStrides, intOffsets)) &&
1638 ShapedType::isStatic(intOffsets)) {
1639 addr = memref::ExtractAlignedPointerAsIndexOp::create(builder, loc,
1641 offset = arith::ConstantOp::create(builder, loc,
1647 SmallVector<Type> resultTypes{
1648 MemRefType::get({}, memrefTy.getElementType(),
1649 MemRefLayoutAttrInterface(),
1650 memrefTy.getMemorySpace()),
1653 resultTypes.append(2 * rank, indexType);
1655 auto meta = memref::ExtractStridedMetadataOp::create(
1656 builder, loc, resultTypes, input);
1658 addr = memref::ExtractAlignedPointerAsIndexOp::create(
1659 builder, loc, meta.getBaseBuffer());
1660 offset = meta.getOffset();
1664 arith::IndexCastUIOp::create(builder, loc, type, addr);
1666 arith::IndexCastUIOp::create(builder, loc, type, offset);
1669 auto byteSize = arith::ConstantOp::create(
1672 memrefTy.getElementTypeBitWidth() / 8));
1674 arith::MulIOp::create(builder, loc, offsetCasted, byteSize);
1675 auto addrWithOffset =
1676 arith::AddIOp::create(builder, loc, addrCasted, byteOffset);
1678 return addrWithOffset.getResult();
1687 auto ui64ToI64MaterializationCast = [](OpBuilder &builder, Type type,
1689 Location loc) -> Value {
1690 if (inputs.size() != 1)
1692 auto input = inputs.front();
1695 index::CastUOp::create(builder, loc, builder.
getIndexType(), input)
1697 return arith::IndexCastUIOp::create(builder, loc, type, cast)
1707 auto ui32ToI32MaterializationCast = [](OpBuilder &builder, Type type,
1709 Location loc) -> Value {
1710 if (inputs.size() != 1)
1712 auto input = inputs.front();
1715 index::CastUOp::create(builder, loc, builder.
getIndexType(), input)
1717 return arith::IndexCastUIOp::create(builder, loc, type, cast)
1727 auto vectorToVectorMaterializationCast = [](OpBuilder &builder, Type type,
1729 Location loc) -> Value {
1730 if (inputs.size() != 1)
1732 auto input = inputs.front();
1733 if (
auto vecTy = dyn_cast<VectorType>(input.getType())) {
1734 if (
auto targetVecTy = dyn_cast<VectorType>(type)) {
1738 if (targetVecTy.getShape() != vecTy.getShape()) {
1739 cast = vector::ShapeCastOp::create(
1741 VectorType::get(targetVecTy.getShape(),
1742 vecTy.getElementType()),
1746 if (targetVecTy.getElementType() != vecTy.getElementType()) {
1747 cast = vector::BitCastOp::create(builder, loc, targetVecTy, cast)
1759 auto vectorToSingleElementMaterializationCast =
1760 [](OpBuilder &builder, Type type,
ValueRange inputs,
1761 Location loc) -> Value {
1762 if (inputs.size() != 1)
1764 auto input = inputs.front();
1765 if (
auto vecTy = dyn_cast<VectorType>(input.getType())) {
1767 auto rank = vecTy.getRank();
1768 if (rank != 0 && vecTy.getNumElements() != 1)
1770 auto inElemTy = vecTy.getElementType();
1774 cast = vector::ExtractOp::create(builder, loc, cast, {}).getResult();
1776 cast = vector::ExtractOp::create(builder, loc, cast,
1777 SmallVector<int64_t>(rank, 0))
1784 if (inElemTy.isIndex()) {
1785 cast = arith::IndexCastUIOp::create(builder, loc, type, cast)
1787 }
else if (inElemTy != type) {
1788 cast = arith::BitcastOp::create(builder, loc, type, cast).getResult();
1802 auto singleElementToVectorMaterializationCast =
1803 [](OpBuilder &builder, Type type,
ValueRange inputs,
1804 Location loc) -> Value {
1805 if (inputs.size() != 1)
1807 auto input = inputs.front();
1808 auto inTy = input.getType();
1809 if (!inTy.isIntOrFloat())
1813 if (
auto vecTy = dyn_cast<VectorType>(type)) {
1814 if (vecTy.getRank() != 0 && vecTy.getNumElements() != 1)
1816 auto outElemTy = vecTy.getElementType();
1818 if (outElemTy.isIndex()) {
1819 cast = arith::IndexCastUIOp::create(builder, loc,
1822 }
else if (inTy != outElemTy) {
1823 cast = arith::BitcastOp::create(builder, loc, outElemTy, cast)
1826 return vector::BroadcastOp::create(builder, loc, vecTy, cast)
1831 typeConverter.addSourceMaterialization(
1832 singleElementToVectorMaterializationCast);
1833 typeConverter.addSourceMaterialization(vectorToVectorMaterializationCast);
1834 typeConverter.addTargetMaterialization(memrefToIntMaterializationCast);
1835 typeConverter.addTargetMaterialization(ui32ToI32MaterializationCast);
1836 typeConverter.addTargetMaterialization(ui64ToI64MaterializationCast);
1837 typeConverter.addTargetMaterialization(
1838 vectorToSingleElementMaterializationCast);
1839 typeConverter.addTargetMaterialization(vectorToVectorMaterializationCast);
1840 ConversionTarget
target(*context);
1841 target.addLegalDialect<xevm::XeVMDialect, LLVM::LLVMDialect,
1842 vector::VectorDialect, arith::ArithDialect,
1843 memref::MemRefDialect, gpu::GPUDialect,
1844 index::IndexDialect>();
1845 target.addIllegalDialect<xegpu::XeGPUDialect>();
1848 target.addDynamicallyLegalOp<arith::ExtFOp>(
1849 [](arith::ExtFOp op) {
return !isXeVMExtf(op); });
1850 target.addDynamicallyLegalOp<arith::TruncFOp>(
1851 [](arith::TruncFOp op) {
return !isXeVMTruncf(op); });
1853 RewritePatternSet patterns(context);
1857 if (
failed(applyPartialConversion(getOperation(),
target,
1858 std::move(patterns))))
1859 signalPassFailure();
1869 patterns.
add<CreateNdDescToXeVMPattern,
1870 LoadStorePrefetchNdToXeVMPattern<xegpu::LoadNdOp>,
1871 LoadStorePrefetchNdToXeVMPattern<xegpu::StoreNdOp>,
1872 LoadStorePrefetchNdToXeVMPattern<xegpu::PrefetchNdOp>>(
1874 patterns.
add<AtomicRMWToXeVMPattern, PrefetchToXeVMPattern,
1875 LoadStoreToXeVMPattern<xegpu::LoadGatherOp>,
1876 LoadStoreToXeVMPattern<xegpu::StoreScatterOp>>(
1878 patterns.
add<LoadStoreMatrixToXeVMPattern<xegpu::LoadMatrixOp>,
1879 LoadStoreMatrixToXeVMPattern<xegpu::StoreMatrixOp>,
1880 CreateMemDescOpPattern>(typeConverter, patterns.
getContext());
1881 patterns.
add<FenceToXeVMPattern, DpasToXeVMPattern>(typeConverter,
1883 patterns.
add<DpasMxToXeVMPattern>(typeConverter, patterns.
getContext());
1884 patterns.
add<ExtfToXeVMPattern, TruncfToXeVMPattern>(typeConverter,
1886 patterns.
add<LaneShuffleToXeVMPattern>(typeConverter, patterns.
getContext());
static llvm::ManagedStatic< PassManagerOptions > options
Attributes are known-constant values of operations.
IntegerAttr getIndexAttr(int64_t value)
IntegerAttr getIntegerAttr(Type type, int64_t value)
IntegerType getIntegerType(unsigned width)
static DenseElementsAttr get(ShapedType type, ArrayRef< Attribute > values)
Constructs a dense elements attribute from an array of element values.
Conversion from types to the LLVM IR dialect.
This class defines the main interface for locations in MLIR and acts as a non-nullable wrapper around...
void setDiscardableAttr(StringAttr name, Attribute value)
Set a discardable attribute by name.
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.
Instances of the Type class are uniqued, have an immutable identifier and an optional mutable compone...
bool isIntOrFloat() const
Return true if this is an integer (of any signedness) or a float type.
unsigned getIntOrFloatBitWidth() const
Return the bit width of an integer or a float type, assert failure on other types.
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 ConstantIntOp create(OpBuilder &builder, Location location, int64_t value, unsigned width)
void populateSCFStructuralTypeConversionsAndLegality(const TypeConverter &typeConverter, RewritePatternSet &patterns, ConversionTarget &target, PatternBenefit benefit=1)
Populates patterns for SCF structural type conversions and sets up the provided ConversionTarget with...
@ SubgroupMatrixMultiplyAcc
const uArch * getUArch(llvm::StringRef archName)
bool hasStaticShapeAndStrides(MemRefType type)
Returns true if type has a static shape and static strides.
std::optional< std::string > getChipStr(Operation *op)
Retrieves the chip string from the XeVM target attribute of the parent GPU module operation.
Include the generated interface declarations.
std::optional< int64_t > getConstantIntValue(OpFoldResult ofr)
If ofr is a constant integer or an IntegerAttr, return the integer.
Value getValueOrCreateConstantIntOp(OpBuilder &b, Location loc, OpFoldResult ofr)
Converts an OpFoldResult to a Value.
Value getValueOrCreateCastToIndexLike(OpBuilder &b, Location loc, Type targetType, Value value)
Create a cast from an index-like value (index or integer) to another index-like value.
void populateXeGPUToXeVMConversionPatterns(const LLVMTypeConverter &typeConverter, RewritePatternSet &patterns)