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.");
245class CreateNdDescToXeVMPattern
246 :
public OpConversionPattern<xegpu::CreateNdDescOp> {
247 using OpConversionPattern::OpConversionPattern;
249 matchAndRewrite(xegpu::CreateNdDescOp op,
250 xegpu::CreateNdDescOp::Adaptor adaptor,
251 ConversionPatternRewriter &rewriter)
const override {
252 auto loc = op.getLoc();
253 auto source = op.getSource();
257 int64_t rank = op.getType().getRank();
259 auto memrefTy = dyn_cast<MemRefType>(source.getType());
261 if (!memrefTy.isStrided())
262 return rewriter.notifyMatchFailure(op,
"Expected strided Memref.");
263 sourceRank = memrefTy.getRank();
264 }
else if (isa<IntegerType>(source.getType())) {
265 sourceRank = op.getMixedSizes().size();
267 return rewriter.notifyMatchFailure(op,
268 "Expected ranked Memref or integer.");
270 if (sourceRank != rank)
271 return rewriter.notifyMatchFailure(
272 op,
"Expected descriptor rank to match source rank; subview the "
273 "source down to the descriptor rank.");
274 if (rank > maxNdTdescRank)
275 return rewriter.notifyMatchFailure(
276 op,
"Batched nd descriptor supports at most " +
277 std::to_string(maxNdTdescLeadingDims) +
278 " leading dims (rank <= " + std::to_string(maxNdTdescRank) +
282 SmallVector<std::optional<int64_t>> constStrides(rank, std::nullopt);
284 SmallVector<int64_t> staticStrides;
285 int64_t staticOffset;
287 memrefTy.getStridesAndOffset(staticStrides, staticOffset)))
288 for (int64_t d = 0; d < rank; ++d)
289 if (!ShapedType::isDynamic(staticStrides[d]))
290 constStrides[d] = staticStrides[d];
292 SmallVector<OpFoldResult> mixed = op.getMixedStrides();
293 for (int64_t d = 0; d < rank; ++d)
296 if (std::optional<int64_t> pitch = constStrides[rank - 2]) {
297 for (int64_t d = 0; d < rank - 2; ++d) {
298 std::optional<int64_t> leading = constStrides[d];
299 if (leading && (*pitch == 0 || *leading % *pitch != 0))
300 return rewriter.notifyMatchFailure(
301 op,
"Expected each leading (batch) stride to be a multiple of "
302 "the row stride; the source has gaps between planes.");
307 Type payloadElemTy = rewriter.getI32Type();
308 Type i64Ty = rewriter.getI64Type();
312 Value baseAddr = adaptor.getSource();
313 if (isa<IntegerType>(source.getType()) && baseAddr.
getType() != i64Ty) {
315 baseAddr = arith::ExtUIOp::create(rewriter, loc, i64Ty, baseAddr);
319 rewriter.replaceOp(op, baseAddr);
323 SmallVector<OpFoldResult> mixedSizes;
324 SmallVector<OpFoldResult> mixedStrides;
327 memref::ExtractStridedMetadataOp::create(rewriter, loc, source);
328 mixedSizes = meta.getConstifiedMixedSizes();
329 mixedStrides = meta.getConstifiedMixedStrides();
331 mixedSizes = op.getMixedSizes();
332 mixedStrides = op.getMixedStrides();
338 VectorType payloadTy = VectorType::get(8, payloadElemTy);
340 VectorType payloadI64Ty = VectorType::get(4, i64Ty);
342 Value payload = arith::ConstantOp::create(
347 auto createOffset = [&](SmallVector<OpFoldResult> &ofrVec,
348 unsigned idx) -> Value {
354 Value baseShapeW = createOffset(mixedSizes, rank - 1);
356 Value basePitch = createOffset(mixedStrides, rank - 2);
358 SmallVector<Value> leadingRowStrides;
359 for (int64_t d = 0; d < rank - 2; ++d) {
361 std::optional<int64_t> pitch =
363 if (leading && pitch && *pitch != 0)
365 rewriter, loc, payloadElemTy, *leading / *pitch));
367 leadingRowStrides.push_back(arith::DivUIOp::create(
368 rewriter, loc, createOffset(mixedStrides, d), basePitch));
374 Value baseShapeH = createOffset(mixedSizes, rank - 2);
377 for (int64_t d = 0; d < rank - 2; ++d) {
378 Value planesBelow = rewriter.createOrFold<arith::SubIOp>(
379 loc, createOffset(mixedSizes, d), one);
380 Value rows = rewriter.createOrFold<arith::MulIOp>(loc, planesBelow,
381 leadingRowStrides[d]);
383 rewriter.createOrFold<arith::AddIOp>(loc, baseShapeH, rows);
388 vector::BitCastOp::create(rewriter, loc, payloadI64Ty, payload);
390 vector::InsertOp::create(rewriter, loc, baseAddr, payLoadAsI64,
391 static_cast<int>(NdTdescOffset::BasePtr));
392 payload = vector::BitCastOp::create(rewriter, loc, payloadTy, payLoadAsI64);
394 vector::InsertOp::create(rewriter, loc, baseShapeW, payload,
395 static_cast<int>(NdTdescOffset::BaseShapeW));
397 vector::InsertOp::create(rewriter, loc, baseShapeH, payload,
398 static_cast<int>(NdTdescOffset::BaseShapeH));
400 vector::InsertOp::create(rewriter, loc, basePitch, payload,
401 static_cast<int>(NdTdescOffset::BasePitch));
405 for (int64_t d = 0; d < rank - 2; ++d)
406 payload = vector::InsertOp::create(
407 rewriter, loc, leadingRowStrides[d], payload,
408 static_cast<int>(NdTdescOffset::LeadingStride0) + d);
409 rewriter.replaceOp(op, payload);
416 typename = std::enable_if_t<llvm::is_one_of<
417 OpType, xegpu::LoadNdOp, xegpu::StoreNdOp, xegpu::PrefetchNdOp>::value>>
418class LoadStorePrefetchNdToXeVMPattern :
public OpConversionPattern<OpType> {
419 using OpConversionPattern<OpType>::OpConversionPattern;
421 matchAndRewrite(OpType op,
typename OpType::Adaptor adaptor,
422 ConversionPatternRewriter &rewriter)
const override {
423 auto mixedOffsets = op.getMixedOffsets();
424 int64_t opOffsetsSize = mixedOffsets.size();
425 auto loc = op.getLoc();
426 auto ctxt = rewriter.getContext();
428 auto tdesc = adaptor.getTensorDesc();
429 auto tdescTy = op.getTensorDescType();
430 auto tileRank = tdescTy.getRank();
431 if (opOffsetsSize != tileRank)
432 return rewriter.notifyMatchFailure(
433 op,
"Expected offset rank to match descriptor rank.");
434 if (tileRank > 2 && llvm::any_of(tdescTy.getShape().drop_back(2),
435 [](int64_t d) { return d != 1; }))
436 return rewriter.notifyMatchFailure(
437 op,
"Expected leading (batch) descriptor dims to be unit.");
438 if (tileRank > maxNdTdescRank)
439 return rewriter.notifyMatchFailure(
440 op,
"Expected descriptor rank <= " + std::to_string(maxNdTdescRank) +
442 auto elemType = tdescTy.getElementType();
443 auto elemBitSize = elemType.getIntOrFloatBitWidth();
444 bool isSubByte = elemBitSize < 8;
445 uint64_t wScaleFactor = 1;
447 if (!isSubByte && (elemBitSize % 8 != 0))
448 return rewriter.notifyMatchFailure(
449 op,
"Expected element type bit width to be multiple of 8.");
450 auto tileW = tdescTy.getDimSize(tileRank - 1);
453 if (elemBitSize != 4)
454 return rewriter.notifyMatchFailure(
455 op,
"Only sub byte types of 4bits are supported.");
457 return rewriter.notifyMatchFailure(
458 op,
"Sub byte types are only supported for 2D tensor descriptors.");
459 auto subByteFactor = 8 / elemBitSize;
460 auto tileH = tdescTy.getDimSize(0);
462 if constexpr (std::is_same_v<OpType, xegpu::LoadNdOp>) {
463 if (op.getPacked().value_or(
false)) {
465 if (tileH == systolicDepth * 4 &&
466 tileW == executionSize * subByteFactor) {
471 elemType = rewriter.getIntegerType(8);
472 tileW = executionSize;
473 wScaleFactor = subByteFactor;
478 if (wScaleFactor == 1) {
479 auto sub16BitFactor = subByteFactor * 2;
480 if (tileW == executionSize * sub16BitFactor) {
484 elemType = rewriter.getIntegerType(16);
485 tileW = executionSize;
486 wScaleFactor = sub16BitFactor;
488 return rewriter.notifyMatchFailure(
489 op,
"Unsupported tile shape for sub byte types.");
493 elemBitSize = elemType.getIntOrFloatBitWidth();
497 auto ptrTypeLLVM = LLVM::LLVMPointerType::get(
498 ctxt, getNumericXeVMAddrSpace(tdescTy.getMemorySpace()));
502 rewriter, loc, rewriter.getI32Type(), elemBitSize / 8);
503 VectorType payloadI64Ty = VectorType::get(4, rewriter.getI64Type());
505 vector::BitCastOp::create(rewriter, loc, payloadI64Ty, tdesc);
507 vector::ExtractOp::create(rewriter, loc, payLoadAsI64,
508 static_cast<int>(NdTdescOffset::BasePtr));
509 Value baseShapeW = vector::ExtractOp::create(
510 rewriter, loc, tdesc,
static_cast<int>(NdTdescOffset::BaseShapeW));
511 Value baseShapeH = vector::ExtractOp::create(
512 rewriter, loc, tdesc,
static_cast<int>(NdTdescOffset::BaseShapeH));
513 Value basePitch = vector::ExtractOp::create(
514 rewriter, loc, tdesc,
static_cast<int>(NdTdescOffset::BasePitch));
517 mixedOffsets[tileRank - 1]);
519 rewriter.getI32Type(), offsetW);
521 mixedOffsets[tileRank - 2]);
523 rewriter.getI32Type(), offsetH);
530 for (int64_t d = 0; d < tileRank - 2; ++d) {
534 rewriter.getI32Type(), off);
535 Value rowStride = vector::ExtractOp::create(
536 rewriter, loc, tdesc,
537 static_cast<int>(NdTdescOffset::LeadingStride0) + d);
538 Value term = arith::MulIOp::create(rewriter, loc, off, rowStride);
539 offsetH = arith::AddIOp::create(rewriter, loc, offsetH, term);
543 LLVM::IntToPtrOp::create(rewriter, loc, ptrTypeLLVM, basePtr);
547 Value baseShapeWInBytes =
548 arith::MulIOp::create(rewriter, loc, baseShapeW, elemByteSize);
550 Value basePitchBytes =
551 arith::MulIOp::create(rewriter, loc, basePitch, elemByteSize);
553 if (wScaleFactor > 1) {
557 rewriter, loc, rewriter.getI32Type(), llvm::Log2_64(wScaleFactor));
558 baseShapeWInBytes = arith::ShRSIOp::create(
559 rewriter, loc, baseShapeWInBytes, wScaleFactorValLog2);
560 basePitchBytes = arith::ShRSIOp::create(rewriter, loc, basePitchBytes,
561 wScaleFactorValLog2);
563 arith::ShRSIOp::create(rewriter, loc, offsetW, wScaleFactorValLog2);
566 auto tileH = tdescTy.getDimSize(tileRank - 2);
568 int32_t vblocks = tdescTy.getArrayLength();
569 if constexpr (std::is_same_v<OpType, xegpu::StoreNdOp>) {
570 Value src = adaptor.getValue();
576 VectorType srcVecTy = dyn_cast<VectorType>(src.
getType());
578 return rewriter.notifyMatchFailure(
579 op,
"Expected store value to be a vector type.");
581 VectorType newSrcVecTy =
582 encodeVectorTypeTo(srcVecTy, rewriter.getIntegerType(elemBitSize));
583 if (srcVecTy != newSrcVecTy)
584 src = vector::BitCastOp::create(rewriter, loc, newSrcVecTy, src);
585 auto storeCacheControl =
586 translateStoreXeGPUCacheHint(op.getL1Hint(), op.getL3Hint());
587 xevm::BlockStore2dOp::create(
588 rewriter, loc, basePtrLLVM, baseShapeWInBytes, baseShapeH,
589 basePitchBytes, offsetW, offsetH, elemBitSize, tileW, tileH, src,
590 xevm::StoreCacheControlAttr::get(ctxt, storeCacheControl));
591 rewriter.eraseOp(op);
593 auto loadCacheControl =
594 translateLoadXeGPUCacheHint(op.getL1Hint(), op.getL3Hint());
595 if constexpr (std::is_same_v<OpType, xegpu::PrefetchNdOp>) {
596 xevm::BlockPrefetch2dOp::create(
597 rewriter, loc, basePtrLLVM, baseShapeWInBytes, baseShapeH,
598 basePitchBytes, offsetW, offsetH, elemBitSize, tileW, tileH,
599 vblocks, xevm::LoadCacheControlAttr::get(ctxt, loadCacheControl));
600 rewriter.eraseOp(op);
602 VectorType dstVecTy = cast<VectorType>(op.getValue().getType());
603 bool vnni = op.getPacked().value_or(
false);
604 auto transposeValue = op.getTranspose();
606 transposeValue.has_value() && transposeValue.value()[0] == 1;
612 if (elemBitSize == 8 && tileW == 16 && tileH == 32 && !vnni &&
620 if (transpose && elemBitSize < 32) {
621 int32_t scale = 32 / elemBitSize;
623 rewriter, loc, rewriter.getI32Type(), llvm::Log2_64(scale));
624 offsetW = arith::ShRSIOp::create(rewriter, loc, offsetW, scaleLog2);
625 tileW = tileW * elemBitSize / 32;
628 VectorType loadedTy = encodeVectorTypeTo(
629 dstVecTy, vnni ? rewriter.getI32Type()
630 : rewriter.getIntegerType(elemBitSize));
632 Value resultFlatVec = xevm::BlockLoad2dOp::create(
633 rewriter, loc, loadedTy, basePtrLLVM, baseShapeWInBytes,
634 baseShapeH, basePitchBytes, offsetW, offsetH, elemBitSize, tileW,
635 tileH, vblocks, transpose, vnni,
636 xevm::LoadCacheControlAttr::get(ctxt, loadCacheControl));
637 resultFlatVec = vector::BitCastOp::create(
639 encodeVectorTypeTo(loadedTy, dstVecTy.getElementType()),
641 rewriter.replaceOp(op, resultFlatVec);
653 rewriter.getI64Type(), offset);
656 rewriter, loc, rewriter.getI64Type(), elemBitSize / 8);
658 rewriter.createOrFold<arith::MulIOp>(loc, offset, elemByteSize);
660 Value finalAddrI64 = rewriter.createOrFold<arith::AddIOp>(
666 LLVM::IntToPtrOp::create(rewriter, loc, ptrTypeLLVM, finalAddrI64);
667 if constexpr (std::is_same_v<OpType, xegpu::StoreNdOp>) {
668 Value src = adaptor.getValue();
674 VectorType srcVecTy = dyn_cast<VectorType>(src.
getType());
676 return rewriter.notifyMatchFailure(
677 op,
"Expected store value to be a vector type.");
679 VectorType newSrcVecTy =
680 encodeVectorTypeTo(srcVecTy, rewriter.getIntegerType(elemBitSize));
681 if (srcVecTy != newSrcVecTy)
682 src = vector::BitCastOp::create(rewriter, loc, newSrcVecTy, src);
683 auto storeCacheControl =
684 translateStoreXeGPUCacheHint(op.getL1Hint(), op.getL3Hint());
685 rewriter.replaceOpWithNewOp<xevm::BlockStoreOp>(
686 op, finalPtrLLVM, src,
687 xevm::StoreCacheControlAttr::get(ctxt, storeCacheControl));
688 }
else if constexpr (std::is_same_v<OpType, xegpu::LoadNdOp>) {
689 auto loadCacheControl =
690 translateLoadXeGPUCacheHint(op.getL1Hint(), op.getL3Hint());
691 VectorType resTy = cast<VectorType>(op.getValue().getType());
692 VectorType loadedTy =
693 encodeVectorTypeTo(resTy, rewriter.getIntegerType(elemBitSize));
694 Value
load = xevm::BlockLoadOp::create(
695 rewriter, loc, loadedTy, finalPtrLLVM,
696 xevm::LoadCacheControlAttr::get(ctxt, loadCacheControl));
697 if (loadedTy != resTy)
698 load = vector::BitCastOp::create(rewriter, loc, resTy,
load);
699 rewriter.replaceOp(op,
load);
701 return rewriter.notifyMatchFailure(
702 op,
"Unsupported operation: xegpu.prefetch_nd with tensor "
703 "descriptor rank == 1");
712static Value addOffsetToBaseAddr(ConversionPatternRewriter &rewriter,
716 rewriter, loc, baseAddr.
getType(), elemByteSize);
717 Value byteOffset = arith::MulIOp::create(rewriter, loc, offset, byteSize);
718 Value newAddr = arith::AddIOp::create(rewriter, loc, baseAddr, byteOffset);
730static bool isUniformMask(
Value mask) {
731 if (!isa<VectorType>(mask.getType()))
736 if (
auto constantMask = mask.getDefiningOp<vector::ConstantMaskOp>()) {
739 return constantMask.isAllOnesMask() ||
740 llvm::all_of(constantMask.getMaskDimSizes(),
741 [](
int64_t size) { return size == 0; });
743 if (
auto broadcast = mask.getDefiningOp<vector::BroadcastOp>())
744 return !isa<VectorType>(
broadcast.getSource().getType());
745 if (
auto fromElements = mask.getDefiningOp<vector::FromElementsOp>())
746 return llvm::all_equal(fromElements.getElements());
749 if (
auto shapeCast = mask.getDefiningOp<vector::ShapeCastOp>())
750 return isUniformMask(shapeCast.getSource());
754template <
typename OpType,
755 typename = std::enable_if_t<llvm::is_one_of<
756 OpType, xegpu::LoadGatherOp, xegpu::StoreScatterOp>::value>>
757class LoadStoreToXeVMPattern :
public OpConversionPattern<OpType> {
758 using OpConversionPattern<OpType>::OpConversionPattern;
760 matchAndRewrite(OpType op,
typename OpType::Adaptor adaptor,
761 ConversionPatternRewriter &rewriter)
const override {
762 Value offset = adaptor.getOffsets();
764 return rewriter.notifyMatchFailure(op,
"Expected offset to be provided.");
765 auto loc = op.getLoc();
766 auto ctxt = rewriter.getContext();
770 if constexpr (std::is_same_v<OpType, xegpu::LoadGatherOp>)
772 this->getTypeConverter()->convertType(op.getResult().getType());
774 valOrResTy = adaptor.getValue().getType();
775 VectorType valOrResVecTy = dyn_cast<VectorType>(valOrResTy);
776 bool hasScalarVal = !valOrResVecTy;
777 int64_t elemBitWidth =
779 : valOrResVecTy.getElementType().getIntOrFloatBitWidth();
781 if (elemBitWidth % 8 != 0)
782 return rewriter.notifyMatchFailure(
783 op,
"Expected element type bit width to be multiple of 8.");
784 int64_t elemByteSize = elemBitWidth / 8;
786 LLVM::LLVMPointerType ptrTypeLLVM = LLVM::LLVMPointerType::get(
787 ctxt, getNumericXeVMAddrSpace(xegpu::MemorySpace::Global));
790 if constexpr (std::is_same_v<OpType, xegpu::LoadGatherOp>) {
791 basePtrI64 = adaptor.getSource();
792 if (
auto memRefTy = dyn_cast<MemRefType>(op.getSource().getType())) {
793 FailureOr<unsigned> addrSpace =
794 getNumericMemorySpace(memRefTy.getMemorySpace());
796 return rewriter.notifyMatchFailure(
797 op,
"Unsupported memref memory space attribute.");
799 ptrTypeLLVM = LLVM::LLVMPointerType::get(ctxt, *addrSpace);
802 basePtrI64 = adaptor.getDest();
803 if (
auto memRefTy = dyn_cast<MemRefType>(op.getDest().getType())) {
804 FailureOr<unsigned> addrSpace =
805 getNumericMemorySpace(memRefTy.getMemorySpace());
807 return rewriter.notifyMatchFailure(
808 op,
"Unsupported memref memory space attribute.");
810 ptrTypeLLVM = LLVM::LLVMPointerType::get(ctxt, *addrSpace);
814 if (basePtrI64.
getType() != rewriter.getI64Type()) {
815 basePtrI64 = arith::ExtUIOp::create(rewriter, loc, rewriter.getI64Type(),
818 Value mask = adaptor.getMask();
827 auto origOffsetsTy = dyn_cast<VectorType>(op.getOffsets().getType());
828 if (isa<VectorType>(offset.
getType()) && origOffsetsTy && valOrResVecTy &&
829 origOffsetsTy.getNumElements() == valOrResVecTy.getNumElements() &&
830 isUniformMask(op.getMask())) {
831 offset = vector::ExtractOp::create(rewriter, loc, offset,
832 ArrayRef<int64_t>{0});
833 if (isa<VectorType>(mask.
getType()))
834 mask = vector::ExtractOp::create(rewriter, loc, mask,
835 ArrayRef<int64_t>{0});
838 if (dyn_cast<VectorType>(offset.
getType())) {
841 return rewriter.notifyMatchFailure(op,
"Expected offset to be a scalar.");
847 addOffsetToBaseAddr(rewriter, loc, basePtrI64, offset, elemByteSize);
851 LLVM::IntToPtrOp::create(rewriter, loc, ptrTypeLLVM, basePtrI64);
854 VectorType maskVecTy = dyn_cast<VectorType>(mask.
getType());
858 return rewriter.notifyMatchFailure(op,
"Expected mask to be a scalar.");
861 if constexpr (std::is_same_v<OpType, xegpu::LoadGatherOp>) {
862 scf::IfOp ifOp = scf::IfOp::create(rewriter, loc, {valOrResTy},
863 maskForLane,
true,
true);
865 rewriter.setInsertionPointToStart(&ifOp.getThenRegion().front());
867 valOrResTy = VectorType::get({valOrResVecTy.getNumElements()},
868 valOrResVecTy.getElementType());
870 LLVM::LoadOp::create(rewriter, loc, valOrResTy, basePtrLLVM);
873 "cache_control", xevm::LoadCacheControlAttr::get(
874 ctxt, translateLoadXeGPUCacheHint(
875 op.getL1Hint(), op.getL3Hint())));
876 scf::YieldOp::create(rewriter, loc,
ValueRange{loaded});
877 rewriter.setInsertionPointToStart(&ifOp.getElseRegion().front());
879 auto eTy = hasScalarVal ? valOrResTy : valOrResVecTy.getElementType();
882 eVal = FloatAttr::get(eTy, 0.0);
884 eVal = IntegerAttr::get(eTy, 0);
886 loaded = arith::ConstantOp::create(rewriter, loc, eVal);
888 loaded = arith::ConstantOp::create(
890 scf::YieldOp::create(rewriter, loc,
ValueRange{loaded});
891 rewriter.replaceOp(op, ifOp.getResult(0));
894 scf::IfOp ifOp = scf::IfOp::create(rewriter, loc, maskForLane,
false);
895 auto body = ifOp.getBody();
896 rewriter.setInsertionPointToStart(body);
898 LLVM::StoreOp::create(rewriter, loc, adaptor.getValue(), basePtrLLVM);
900 storeOp.getOperation()->setDiscardableAttr(
901 "cache_control", xevm::StoreCacheControlAttr::get(
902 ctxt, translateStoreXeGPUCacheHint(
903 op.getL1Hint(), op.getL3Hint())));
904 rewriter.eraseOp(op);
910class CreateMemDescOpPattern final
911 :
public OpConversionPattern<xegpu::CreateMemDescOp> {
913 using OpConversionPattern<xegpu::CreateMemDescOp>::OpConversionPattern;
915 matchAndRewrite(xegpu::CreateMemDescOp op, OpAdaptor adaptor,
916 ConversionPatternRewriter &rewriter)
const override {
918 rewriter.replaceOp(op, adaptor.getSource());
923template <
typename OpType,
924 typename = std::enable_if_t<llvm::is_one_of<
925 OpType, xegpu::LoadMatrixOp, xegpu::StoreMatrixOp>::value>>
926class LoadStoreMatrixToXeVMPattern :
public OpConversionPattern<OpType> {
927 using OpConversionPattern<OpType>::OpConversionPattern;
929 matchAndRewrite(OpType op,
typename OpType::Adaptor adaptor,
930 ConversionPatternRewriter &rewriter)
const override {
932 SmallVector<OpFoldResult> offsets = op.getMixedOffsets();
934 return rewriter.notifyMatchFailure(op,
"Expected offset to be provided.");
936 auto loc = op.getLoc();
937 auto ctxt = rewriter.getContext();
938 Value baseAddr32 = adaptor.getMemDesc();
939 Value mdescVal = op.getMemDesc();
942 if constexpr (std::is_same_v<OpType, xegpu::LoadMatrixOp>) {
943 Type resType = op.getResult().getType();
946 if (
auto vecType = dyn_cast<VectorType>(resType)) {
947 assert(llvm::count_if(vecType.getShape(),
948 [](int64_t d) { return d != 1; }) <= 1 &&
949 "Expected either 1D vector or nD with unit dimensions");
950 resType = VectorType::get({vecType.getNumElements()},
951 vecType.getElementType());
955 dataTy = adaptor.getData().getType();
956 VectorType valOrResVecTy = dyn_cast<VectorType>(dataTy);
958 valOrResVecTy = VectorType::get(1, dataTy);
960 int64_t elemBitWidth =
961 valOrResVecTy.getElementType().getIntOrFloatBitWidth();
963 if (elemBitWidth % 8 != 0)
964 return rewriter.notifyMatchFailure(
965 op,
"Expected element type bit width to be multiple of 8.");
966 int64_t elemByteSize = elemBitWidth / 8;
969 LLVM::LLVMPointerType ptrTypeLLVM = LLVM::LLVMPointerType::get(
970 ctxt, getNumericXeVMAddrSpace(xegpu::MemorySpace::SLM));
972 auto mdescTy = cast<xegpu::MemDescType>(mdescVal.
getType());
974 Value linearOffset = mdescTy.getLinearOffsets(rewriter, loc, offsets);
975 linearOffset = arith::IndexCastUIOp::create(
976 rewriter, loc, rewriter.getI32Type(), linearOffset);
977 Value basePtrI32 = addOffsetToBaseAddr(rewriter, loc, baseAddr32,
978 linearOffset, elemByteSize);
982 LLVM::IntToPtrOp::create(rewriter, loc, ptrTypeLLVM, basePtrI32);
984 if (op.getSubgroupBlockIoAttr()) {
988 Type intElemTy = rewriter.getIntegerType(elemBitWidth);
989 VectorType intVecTy =
990 VectorType::get(valOrResVecTy.getShape(), intElemTy);
992 if constexpr (std::is_same_v<OpType, xegpu::LoadMatrixOp>) {
993 Value loadOp = xevm::BlockLoadOp::create(
994 rewriter, loc, intVecTy, basePtrLLVM,
nullptr);
995 if (intVecTy != valOrResVecTy) {
997 vector::BitCastOp::create(rewriter, loc, valOrResVecTy, loadOp);
999 rewriter.replaceOp(op, loadOp);
1001 Value dataToStore = adaptor.getData();
1002 if (valOrResVecTy != intVecTy) {
1004 vector::BitCastOp::create(rewriter, loc, intVecTy, dataToStore);
1006 xevm::BlockStoreOp::create(rewriter, loc, basePtrLLVM, dataToStore,
1008 rewriter.eraseOp(op);
1013 if (valOrResVecTy.getNumElements() >= 1) {
1016 (*chipOpt !=
"pvc" && *chipOpt !=
"bmg" && *chipOpt !=
"cri")) {
1018 return rewriter.notifyMatchFailure(
1019 op,
"The lowering is specific to pvc, bmg or cri.");
1023 if constexpr (std::is_same_v<OpType, xegpu::LoadMatrixOp>) {
1030 this->getTypeConverter()->convertType(op.getResult().getType());
1031 auto loadOp = LLVM::LoadOp::create(rewriter, loc, loadTy, basePtrLLVM);
1032 rewriter.replaceOp(op, loadOp);
1034 LLVM::StoreOp::create(rewriter, loc, adaptor.getData(), basePtrLLVM);
1035 rewriter.eraseOp(op);
1041class PrefetchToXeVMPattern :
public OpConversionPattern<xegpu::PrefetchOp> {
1042 using OpConversionPattern::OpConversionPattern;
1044 matchAndRewrite(xegpu::PrefetchOp op, xegpu::PrefetchOp::Adaptor adaptor,
1045 ConversionPatternRewriter &rewriter)
const override {
1046 auto loc = op.getLoc();
1047 auto ctxt = rewriter.getContext();
1048 Value basePtrI64 = adaptor.getSource();
1050 if (basePtrI64.
getType() != rewriter.getI64Type())
1051 basePtrI64 = arith::ExtUIOp::create(rewriter, loc, rewriter.getI64Type(),
1053 Value offsets = adaptor.getOffsets();
1055 VectorType offsetsVecTy = dyn_cast<VectorType>(offsets.
getType());
1058 return rewriter.notifyMatchFailure(op,
1059 "Expected offsets to be a scalar.");
1061 int64_t elemBitWidth{0};
1062 int64_t elemByteSize;
1064 if (
auto memRefTy = dyn_cast<MemRefType>(op.getSourceType())) {
1067 elemBitWidth = memRefTy.getElementType().getIntOrFloatBitWidth();
1070 elemByteSize = *op.getOffsetAlignByte();
1072 if (elemBitWidth != 0) {
1073 if (elemBitWidth % 8 != 0)
1074 return rewriter.notifyMatchFailure(
1075 op,
"Expected element type bit width to be multiple of 8.");
1076 elemByteSize = elemBitWidth / 8;
1078 basePtrI64 = addOffsetToBaseAddr(rewriter, loc, basePtrI64, offsets,
1083 LLVM::LLVMPointerType ptrTypeLLVM = LLVM::LLVMPointerType::get(
1084 ctxt, getNumericXeVMAddrSpace(xegpu::MemorySpace::Global));
1086 if (
auto memRefTy = dyn_cast<MemRefType>(op.getSource().getType())) {
1087 FailureOr<unsigned> addrSpace =
1088 getNumericMemorySpace(memRefTy.getMemorySpace());
1090 return rewriter.notifyMatchFailure(
1091 op,
"Unsupported memref memory space attribute.");
1092 if (*addrSpace != 0)
1093 ptrTypeLLVM = LLVM::LLVMPointerType::get(ctxt, *addrSpace);
1097 LLVM::IntToPtrOp::create(rewriter, loc, ptrTypeLLVM, basePtrI64);
1099 xevm::PrefetchOp::create(
1100 rewriter, loc, ptrLLVM,
1101 xevm::LoadCacheControlAttr::get(
1102 ctxt, translateLoadXeGPUCacheHint(op.getL1Hint(), op.getL3Hint())));
1103 rewriter.eraseOp(op);
1108class FenceToXeVMPattern :
public OpConversionPattern<xegpu::FenceOp> {
1109 using OpConversionPattern::OpConversionPattern;
1111 matchAndRewrite(xegpu::FenceOp op, xegpu::FenceOp::Adaptor adaptor,
1112 ConversionPatternRewriter &rewriter)
const override {
1113 auto loc = op.getLoc();
1114 xevm::MemScope memScope{xevm::MemScope::WORKGROUP};
1115 switch (op.getFenceScope()) {
1116 case xegpu::FenceScope::Workgroup:
1117 memScope = xevm::MemScope::WORKGROUP;
1119 case xegpu::FenceScope::GPU:
1120 memScope = xevm::MemScope::DEVICE;
1123 xevm::AddrSpace addrSpace{xevm::AddrSpace::GLOBAL};
1124 switch (op.getMemoryKind()) {
1125 case xegpu::MemorySpace::Global:
1126 addrSpace = xevm::AddrSpace::GLOBAL;
1128 case xegpu::MemorySpace::SLM:
1129 addrSpace = xevm::AddrSpace::SHARED;
1132 xevm::MemfenceOp::create(rewriter, loc, memScope, addrSpace);
1133 rewriter.eraseOp(op);
1138static auto encodePrecision = [](
Type type) -> xevm::ElemType {
1140 return xevm::ElemType::BF16;
1141 else if (type.isF16())
1142 return xevm::ElemType::F16;
1143 else if (type.isTF32())
1144 return xevm::ElemType::TF32;
1145 else if (type.isInteger(8)) {
1146 if (type.isUnsignedInteger())
1147 return xevm::ElemType::U8;
1148 return xevm::ElemType::S8;
1149 }
else if (type.isF32())
1150 return xevm::ElemType::F32;
1151 else if (type.isInteger(32))
1152 return xevm::ElemType::S32;
1153 else if (type.isF8E5M2())
1154 return xevm::ElemType::BF8;
1155 else if (type.isF8E4M3FN())
1156 return xevm::ElemType::F8;
1157 else if (mlir::isa<Float4E2M1FNType>(type))
1158 return xevm::ElemType::E2M1;
1159 llvm_unreachable(
"add more support for ElemType");
1162static unsigned getNumOperandsPerDword(xevm::ElemType pTy) {
1164 case xevm::ElemType::TF32:
1166 case xevm::ElemType::BF16:
1167 case xevm::ElemType::F16:
1169 case xevm::ElemType::U8:
1170 case xevm::ElemType::S8:
1171 case xevm::ElemType::F8:
1172 case xevm::ElemType::BF8:
1174 case xevm::ElemType::E2M1:
1177 llvm_unreachable(
"unsupported xevm::ElemType");
1181class DpasToXeVMPattern :
public OpConversionPattern<xegpu::DpasOp> {
1182 using OpConversionPattern::OpConversionPattern;
1184 matchAndRewrite(xegpu::DpasOp op, xegpu::DpasOp::Adaptor adaptor,
1185 ConversionPatternRewriter &rewriter)
const override {
1186 auto loc = op.getLoc();
1187 auto ctxt = rewriter.getContext();
1188 auto aTy = cast<VectorType>(op.getLhs().getType());
1189 auto bTy = cast<VectorType>(op.getRhs().getType());
1190 auto resultType = cast<VectorType>(op.getResultType());
1195 return rewriter.notifyMatchFailure(op,
"cannot determine target chip");
1199 return rewriter.notifyMatchFailure(op,
"unsupported target uArch");
1202 llvm::dyn_cast_or_null<xegpu::uArch::SubgroupMatrixMultiplyAcc>(
1203 uArch->getInstruction(
1204 xegpu::uArch::InstructionKind::SubgroupMatrixMultiplyAcc)));
1206 return rewriter.notifyMatchFailure(op,
1207 "DPAS not supported by target uArch");
1209 auto checkSupportedTypes = [&](VectorType vecTy,
1211 auto supported = dpasInst->getSupportedTypes(*ctxt, kind);
1212 return llvm::find(supported, vecTy.getElementType()) != supported.end();
1215 if (!checkSupportedTypes(aTy, xegpu::uArch::MMAOpndKind::MatrixA))
1216 return rewriter.notifyMatchFailure(
1217 op,
"A-matrix element type not supported by target uArch");
1218 if (!checkSupportedTypes(bTy, xegpu::uArch::MMAOpndKind::MatrixB))
1219 return rewriter.notifyMatchFailure(
1220 op,
"B-matrix element type not supported by target uArch");
1222 if (!checkSupportedTypes(resultType, xegpu::uArch::MMAOpndKind::MatrixD))
1223 return rewriter.notifyMatchFailure(
1224 op,
"result/accumulator element type not supported by target uArch");
1226 xevm::ElemType precATy = encodePrecision(aTy.getElementType());
1227 xevm::ElemType precBTy = encodePrecision(bTy.getElementType());
1228 Value c = op.getAcc();
1230 auto elementTy = resultType.getElementType();
1231 Attribute initValueAttr;
1232 if (isa<FloatType>(elementTy))
1233 initValueAttr = FloatAttr::get(elementTy, 0.0);
1235 initValueAttr = IntegerAttr::get(elementTy, 0);
1236 c = arith::ConstantOp::create(
1240 Value aVec = op.getLhs();
1241 Value bVec = op.getRhs();
1242 auto cvecty = cast<VectorType>(c.
getType());
1243 xevm::ElemType precCTy = encodePrecision(cvecty.getElementType());
1244 xevm::ElemType precDTy = encodePrecision(resultType.getElementType());
1246 VectorType::get(cvecty.getNumElements(), cvecty.getElementType());
1248 c = vector::ShapeCastOp::create(rewriter, loc, cNty, c);
1249 Value dpasRes = xevm::MMAOp::create(
1250 rewriter, loc, cNty, aVec, bVec, c,
1251 xevm::MMAShapeAttr::get(ctxt, cvecty.getNumElements(), executionSize,
1253 getNumOperandsPerDword(precATy)),
1254 xevm::MMATypesAttr::get(ctxt, precDTy, precATy, precBTy, precCTy));
1256 dpasRes = vector::ShapeCastOp::create(rewriter, loc, resultType, dpasRes);
1257 rewriter.replaceOp(op, dpasRes);
1262static std::optional<LLVM::AtomicBinOp>
1263matchSimpleAtomicOp(arith::AtomicRMWKind arithKind) {
1264 switch (arithKind) {
1265 case arith::AtomicRMWKind::addf:
1266 return LLVM::AtomicBinOp::fadd;
1267 case arith::AtomicRMWKind::addi:
1268 return LLVM::AtomicBinOp::add;
1269 case arith::AtomicRMWKind::assign:
1270 return LLVM::AtomicBinOp::xchg;
1271 case arith::AtomicRMWKind::maximumf:
1272 return LLVM::AtomicBinOp::fmax;
1273 case arith::AtomicRMWKind::maxs:
1274 return LLVM::AtomicBinOp::max;
1275 case arith::AtomicRMWKind::maxu:
1276 return LLVM::AtomicBinOp::umax;
1277 case arith::AtomicRMWKind::minimumf:
1278 return LLVM::AtomicBinOp::fmin;
1279 case arith::AtomicRMWKind::mins:
1280 return LLVM::AtomicBinOp::min;
1281 case arith::AtomicRMWKind::minu:
1282 return LLVM::AtomicBinOp::umin;
1283 case arith::AtomicRMWKind::ori:
1284 return LLVM::AtomicBinOp::_or;
1285 case arith::AtomicRMWKind::andi:
1286 return LLVM::AtomicBinOp::_and;
1288 return std::nullopt;
1292class AtomicRMWToXeVMPattern :
public OpConversionPattern<xegpu::AtomicRMWOp> {
1293 using OpConversionPattern::OpConversionPattern;
1295 matchAndRewrite(xegpu::AtomicRMWOp op, xegpu::AtomicRMWOp::Adaptor adaptor,
1296 ConversionPatternRewriter &rewriter)
const override {
1297 auto loc = op.getLoc();
1298 auto ctxt = rewriter.getContext();
1299 auto tdesc = op.getTensorDesc().getType();
1300 auto ptrTypeLLVM = LLVM::LLVMPointerType::get(
1301 ctxt, getNumericXeVMAddrSpace(tdesc.getMemorySpace()));
1302 Value basePtrI64 = arith::IndexCastOp::create(
1303 rewriter, loc, rewriter.getI64Type(), adaptor.getTensorDesc());
1305 LLVM::IntToPtrOp::create(rewriter, loc, ptrTypeLLVM, basePtrI64);
1306 VectorType srcOrDstVecTy = cast<VectorType>(op.getValue().getType());
1307 VectorType srcOrDstFlatVecTy = VectorType::get(
1308 srcOrDstVecTy.getNumElements(), srcOrDstVecTy.getElementType());
1309 Value srcFlatVec = vector::ShapeCastOp::create(
1310 rewriter, loc, srcOrDstFlatVecTy, op.getValue());
1311 auto atomicKind = matchSimpleAtomicOp(op.getKind());
1312 assert(atomicKind.has_value());
1313 Value resVec = srcFlatVec;
1314 for (
int i = 0; i < srcOrDstVecTy.getNumElements(); i++) {
1315 auto val = vector::ExtractOp::create(rewriter, loc, resVec, i);
1316 Value idx = LLVM::ConstantOp::create(rewriter, loc, rewriter.getI64Type(),
1317 rewriter.getI64IntegerAttr(i));
1319 LLVM::GEPOp::create(rewriter, loc, ptrTypeLLVM,
1320 srcOrDstVecTy.getElementType(), basePtrLLVM, idx);
1322 LLVM::AtomicRMWOp::create(rewriter, loc, atomicKind.value(), currPtr,
1323 val, LLVM::AtomicOrdering::seq_cst);
1324 resVec = vector::InsertOp::create(rewriter, loc, newVal, resVec, i);
1326 rewriter.replaceOp(op, resVec);
1331class DpasMxToXeVMPattern :
public OpConversionPattern<xegpu::DpasMxOp> {
1332 using OpConversionPattern::OpConversionPattern;
1334 matchAndRewrite(xegpu::DpasMxOp op, xegpu::DpasMxOp::Adaptor adaptor,
1335 ConversionPatternRewriter &rewriter)
const override {
1336 auto loc = op.getLoc();
1337 auto ctxt = rewriter.getContext();
1338 auto aTy = op.getA().getType();
1339 auto bTy = op.getB().getType();
1341 cast<VectorType>(getTypeConverter()->convertType(op.getType()));
1345 return rewriter.notifyMatchFailure(op,
"cannot determine target chip");
1349 return rewriter.notifyMatchFailure(op,
"unsupported target uArch");
1353 xevm::ElemType precATy = encodePrecision(aTy.getElementType());
1354 xevm::ElemType precBTy = encodePrecision(bTy.getElementType());
1355 Value c = adaptor.getAcc();
1357 auto elementTy = resVecTy.getElementType();
1358 Attribute initValueAttr;
1359 if (isa<FloatType>(elementTy))
1360 initValueAttr = FloatAttr::get(elementTy, 0.0);
1362 initValueAttr = IntegerAttr::get(elementTy, 0);
1363 c = arith::ConstantOp::create(
1367 Value aVec = adaptor.getA();
1368 Value bVec = adaptor.getB();
1369 auto aVecTy = cast<VectorType>(aVec.
getType());
1370 auto bVecTy = cast<VectorType>(bVec.
getType());
1371 if (aVecTy.getElementTypeBitWidth() == 4)
1372 aVec = vector::BitCastOp::create(
1374 VectorType::get(aVecTy.getNumElements() / 2, rewriter.getI8Type()),
1376 if (bVecTy.getElementTypeBitWidth() == 4)
1377 bVec = vector::BitCastOp::create(
1379 VectorType::get(bVecTy.getNumElements() / 2, rewriter.getI8Type()),
1381 auto cVecTy = cast<VectorType>(c.
getType());
1382 xevm::ElemType precCTy = encodePrecision(cVecTy.getElementType());
1383 xevm::ElemType precDTy = encodePrecision(resVecTy.getElementType());
1384 Value scaleA = adaptor.getScaleA();
1385 Value scaleB = adaptor.getScaleB();
1386 Value dpasMxRes = xevm::MMAMxOp::create(
1387 rewriter, loc, resVecTy, aVec, bVec, scaleA, scaleB, c,
1388 xevm::MMAShapeAttr::get(ctxt, cVecTy.getNumElements(), executionSize,
1390 getNumOperandsPerDword(precATy)),
1391 xevm::MMATypesAttr::get(ctxt, precDTy, precATy, precBTy, precCTy));
1392 rewriter.replaceOp(op, dpasMxRes);
1411static constexpr int64_t kXeVMExtfTruncfNumElems = 16;
1414static std::optional<xevm::ExtfSrcElemTypes> getExtfNarrowType(
Type etype) {
1415 if (isa<Float8E5M2Type>(etype))
1416 return xevm::ExtfSrcElemTypes::BF8;
1417 if (isa<Float8E4M3FNType>(etype))
1418 return xevm::ExtfSrcElemTypes::F8;
1419 if (isa<Float4E2M1FNType>(etype))
1420 return xevm::ExtfSrcElemTypes::E2M1;
1421 return std::nullopt;
1425static std::optional<xevm::TruncfDstElemTypes> getTruncfNarrowType(
Type etype) {
1426 if (isa<Float8E5M2Type>(etype))
1427 return xevm::TruncfDstElemTypes::BF8;
1428 if (isa<Float8E4M3FNType>(etype))
1429 return xevm::TruncfDstElemTypes::F8;
1430 if (isa<Float4E2M1FNType>(etype))
1431 return xevm::TruncfDstElemTypes::E2M1;
1432 return std::nullopt;
1437static bool isXeVMExtf(arith::ExtFOp op) {
1438 auto srcTy = dyn_cast<VectorType>(op.getIn().getType());
1439 auto dstTy = dyn_cast<VectorType>(op.getType());
1440 if (!srcTy || !dstTy || srcTy.getRank() != 1 || dstTy.getRank() != 1)
1442 if (dstTy.getNumElements() != kXeVMExtfTruncfNumElems)
1444 Type dstETy = dstTy.getElementType();
1447 return getExtfNarrowType(srcTy.getElementType()).has_value();
1454static bool isXeVMTruncf(arith::TruncFOp op) {
1455 auto srcTy = dyn_cast<VectorType>(op.getIn().getType());
1456 auto dstTy = dyn_cast<VectorType>(op.getType());
1457 if (!srcTy || !dstTy || srcTy.getRank() != 1 || dstTy.getRank() != 1)
1459 int64_t numElems = srcTy.getNumElements();
1460 if (numElems == 0 || numElems % kXeVMExtfTruncfNumElems != 0)
1462 Type srcETy = srcTy.getElementType();
1465 return getTruncfNarrowType(dstTy.getElementType()).has_value();
1468class ExtfToXeVMPattern :
public OpConversionPattern<arith::ExtFOp> {
1469 using OpConversionPattern::OpConversionPattern;
1471 matchAndRewrite(arith::ExtFOp op, OpAdaptor adaptor,
1472 ConversionPatternRewriter &rewriter)
const override {
1473 if (!isXeVMExtf(op))
1474 return rewriter.notifyMatchFailure(op,
"not a xevm.extf compatible extf");
1475 Location loc = op.getLoc();
1476 MLIRContext *ctx = op.getContext();
1477 auto srcVecTy = cast<VectorType>(op.getIn().getType());
1478 auto dstVecTy = cast<VectorType>(op.getType());
1479 xevm::ExtfSrcElemTypes srcEnum =
1480 *getExtfNarrowType(srcVecTy.getElementType());
1481 xevm::ExtfDstElemTypes dstEnum = dstVecTy.getElementType().isF16()
1482 ? xevm::ExtfDstElemTypes::F16
1483 : xevm::ExtfDstElemTypes::BF16;
1487 Value src = adaptor.getIn();
1488 auto convSrcTy = cast<VectorType>(src.
getType());
1489 if (convSrcTy.getElementTypeBitWidth() == 4)
1490 src = vector::BitCastOp::create(
1492 VectorType::get(convSrcTy.getNumElements() / 2, rewriter.getI8Type()),
1494 Type resTy = getTypeConverter()->convertType(dstVecTy);
1495 Value res = xevm::ExtfOp::create(
1496 rewriter, loc, resTy, src, xevm::ExtfSrcElemTypeAttr::get(ctx, srcEnum),
1497 xevm::ExtfDstElemTypeAttr::get(ctx, dstEnum));
1498 rewriter.replaceOp(op, res);
1503class TruncfToXeVMPattern :
public OpConversionPattern<arith::TruncFOp> {
1504 using OpConversionPattern::OpConversionPattern;
1506 matchAndRewrite(arith::TruncFOp op, OpAdaptor adaptor,
1507 ConversionPatternRewriter &rewriter)
const override {
1508 if (!isXeVMTruncf(op))
1509 return rewriter.notifyMatchFailure(op,
1510 "not a xevm.truncf compatible truncf");
1511 Location loc = op.getLoc();
1512 MLIRContext *ctx = op.getContext();
1513 auto srcVecTy = cast<VectorType>(op.getIn().getType());
1514 auto dstVecTy = cast<VectorType>(op.getType());
1515 xevm::TruncfSrcElemTypes srcEnum = srcVecTy.getElementType().isF16()
1516 ? xevm::TruncfSrcElemTypes::F16
1517 : xevm::TruncfSrcElemTypes::BF16;
1518 xevm::TruncfDstElemTypes dstEnum =
1519 *getTruncfNarrowType(dstVecTy.getElementType());
1520 auto srcEnumAttr = xevm::TruncfSrcElemTypeAttr::get(ctx, srcEnum);
1521 auto dstEnumAttr = xevm::TruncfDstElemTypeAttr::get(ctx, dstEnum);
1527 int64_t numGroups = srcVecTy.getNumElements() / kXeVMExtfTruncfNumElems;
1528 int64_t groupBytes =
1529 kXeVMExtfTruncfNumElems * dstVecTy.getElementTypeBitWidth() / 8;
1530 Type groupTy = VectorType::get(groupBytes, rewriter.getI8Type());
1532 Value src = adaptor.getIn();
1534 if (numGroups == 1) {
1535 packed = xevm::TruncfOp::create(rewriter, loc, groupTy, src, srcEnumAttr,
1539 VectorType::get(groupBytes * numGroups, rewriter.getI8Type());
1540 packed = arith::ConstantOp::create(rewriter, loc, packedTy,
1541 rewriter.getZeroAttr(packedTy));
1542 for (int64_t group = 0; group < numGroups; group++) {
1543 Value slice = vector::ExtractStridedSliceOp::create(
1544 rewriter, loc, src, group * kXeVMExtfTruncfNumElems,
1545 kXeVMExtfTruncfNumElems, 1);
1546 Value converted = xevm::TruncfOp::create(rewriter, loc, groupTy, slice,
1547 srcEnumAttr, dstEnumAttr);
1548 packed = vector::InsertStridedSliceOp::create(
1549 rewriter, loc, converted, packed, group * groupBytes,
1554 Type resTy = getTypeConverter()->convertType(dstVecTy);
1555 if (packed.
getType() != resTy)
1556 packed = vector::BitCastOp::create(rewriter, loc, resTy, packed);
1557 rewriter.replaceOp(op, packed);
1581class LaneShuffleToXeVMPattern
1582 :
public OpConversionPattern<xegpu::LaneShuffleOp> {
1583 using OpConversionPattern::OpConversionPattern;
1585 matchAndRewrite(xegpu::LaneShuffleOp op, OpAdaptor adaptor,
1586 ConversionPatternRewriter &rewriter)
const override {
1587 auto vecTy = dyn_cast<VectorType>(adaptor.getSource().getType());
1589 return rewriter.notifyMatchFailure(op,
"Expected a vector fragment.");
1593 unsigned elemBits = vecTy.getElementTypeBitWidth();
1594 if (elemBits != 8 && elemBits != 16 && elemBits != 32 && elemBits != 64)
1595 return rewriter.notifyMatchFailure(
1596 op,
"Expected an element type of 8, 16, 32 or 64 bits.");
1597 int64_t fragmentBits = vecTy.getNumElements() * elemBits;
1598 if (fragmentBits > 64 || !llvm::isPowerOf2_64(fragmentBits))
1599 return rewriter.notifyMatchFailure(
1600 op,
"Expected a fragment of 8, 16, 32 or 64 bits.");
1602 Location loc = op.getLoc();
1603 Type packedTy = rewriter.getIntegerType(fragmentBits);
1606 VectorType shuffleTy =
1607 VectorType::get(vecTy.getShape(), rewriter.getIntegerType(elemBits));
1610 if (op.getMode() == xegpu::LaneShuffleMode::Pack) {
1611 Value src = adaptor.getSource();
1612 if (shuffleTy != vecTy)
1613 src = LLVM::BitcastOp::create(rewriter, loc, shuffleTy, src);
1614 res = xevm::BitcastShuffleOp::create(rewriter, loc, packedTy, src);
1615 res = LLVM::BitcastOp::create(rewriter, loc, vecTy, res);
1618 LLVM::BitcastOp::create(rewriter, loc, packedTy, adaptor.getSource());
1619 res = xevm::BitcastShuffleOp::create(rewriter, loc, shuffleTy, packed);
1620 if (shuffleTy != vecTy)
1621 res = LLVM::BitcastOp::create(rewriter, loc, vecTy, res);
1623 rewriter.replaceOp(op, res);
1632struct ConvertXeGPUToXeVMPass
1636 void runOnOperation()
override {
1646 LowerToLLVMOptions
options(context);
1647 options.overrideIndexBitwidth(this->use64bitIndex ? 64 : 32);
1648 LLVMTypeConverter typeConverter(context,
options);
1650 Type xevmIndexType = typeConverter.convertType(IndexType::get(context));
1651 Type i32Type = IntegerType::get(context, 32);
1652 typeConverter.addConversion([&](VectorType type) -> Type {
1653 auto elemType = typeConverter.convertType(type.getElementType());
1655 unsigned rank = type.getRank();
1656 if (rank == 0 || type.getNumElements() == 1)
1659 int64_t sum = llvm::product_of(type.getShape());
1660 return VectorType::get(sum, elemType);
1662 typeConverter.addConversion([&](xegpu::TensorDescType type) -> Type {
1663 if (type.getRank() == 1)
1664 return xevmIndexType;
1665 return VectorType::get(8, i32Type);
1674 typeConverter.addConversion(
1675 [&](xegpu::MemDescType type) -> Type {
return i32Type; });
1677 typeConverter.addConversion([&](MemRefType type) -> Type {
1678 return isSharedMemRef(type) ? i32Type : xevmIndexType;
1688 auto memrefToIntMaterializationCast = [](OpBuilder &builder, Type type,
1690 Location loc) -> Value {
1691 if (inputs.size() != 1)
1693 auto input = inputs.front();
1694 if (
auto memrefTy = dyn_cast<MemRefType>(input.getType())) {
1695 unsigned rank = memrefTy.getRank();
1699 SmallVector<int64_t> intStrides;
1702 if (succeeded(memrefTy.getStridesAndOffset(intStrides, intOffsets)) &&
1703 ShapedType::isStatic(intOffsets)) {
1704 addr = memref::ExtractAlignedPointerAsIndexOp::create(builder, loc,
1706 offset = arith::ConstantOp::create(builder, loc,
1712 SmallVector<Type> resultTypes{
1713 MemRefType::get({}, memrefTy.getElementType(),
1714 MemRefLayoutAttrInterface(),
1715 memrefTy.getMemorySpace()),
1718 resultTypes.append(2 * rank, indexType);
1720 auto meta = memref::ExtractStridedMetadataOp::create(
1721 builder, loc, resultTypes, input);
1723 addr = memref::ExtractAlignedPointerAsIndexOp::create(
1724 builder, loc, meta.getBaseBuffer());
1725 offset = meta.getOffset();
1729 arith::IndexCastUIOp::create(builder, loc, type, addr);
1731 arith::IndexCastUIOp::create(builder, loc, type, offset);
1734 auto byteSize = arith::ConstantOp::create(
1737 memrefTy.getElementTypeBitWidth() / 8));
1739 arith::MulIOp::create(builder, loc, offsetCasted, byteSize);
1740 auto addrWithOffset =
1741 arith::AddIOp::create(builder, loc, addrCasted, byteOffset);
1743 return addrWithOffset.getResult();
1752 auto ui64ToI64MaterializationCast = [](OpBuilder &builder, Type type,
1754 Location loc) -> Value {
1755 if (inputs.size() != 1)
1757 auto input = inputs.front();
1760 index::CastUOp::create(builder, loc, builder.
getIndexType(), input)
1762 return arith::IndexCastUIOp::create(builder, loc, type, cast)
1772 auto ui32ToI32MaterializationCast = [](OpBuilder &builder, Type type,
1774 Location loc) -> Value {
1775 if (inputs.size() != 1)
1777 auto input = inputs.front();
1780 index::CastUOp::create(builder, loc, builder.
getIndexType(), input)
1782 return arith::IndexCastUIOp::create(builder, loc, type, cast)
1792 auto vectorToVectorMaterializationCast = [](OpBuilder &builder, Type type,
1794 Location loc) -> Value {
1795 if (inputs.size() != 1)
1797 auto input = inputs.front();
1798 if (
auto vecTy = dyn_cast<VectorType>(input.getType())) {
1799 if (
auto targetVecTy = dyn_cast<VectorType>(type)) {
1803 if (targetVecTy.getShape() != vecTy.getShape()) {
1804 cast = vector::ShapeCastOp::create(
1806 VectorType::get(targetVecTy.getShape(),
1807 vecTy.getElementType()),
1811 if (targetVecTy.getElementType() != vecTy.getElementType()) {
1812 cast = vector::BitCastOp::create(builder, loc, targetVecTy, cast)
1824 auto vectorToSingleElementMaterializationCast =
1825 [](OpBuilder &builder, Type type,
ValueRange inputs,
1826 Location loc) -> Value {
1827 if (inputs.size() != 1)
1829 auto input = inputs.front();
1830 if (
auto vecTy = dyn_cast<VectorType>(input.getType())) {
1832 auto rank = vecTy.getRank();
1833 if (rank != 0 && vecTy.getNumElements() != 1)
1835 auto inElemTy = vecTy.getElementType();
1839 cast = vector::ExtractOp::create(builder, loc, cast, {}).getResult();
1841 cast = vector::ExtractOp::create(builder, loc, cast,
1842 SmallVector<int64_t>(rank, 0))
1849 if (inElemTy.isIndex()) {
1850 cast = arith::IndexCastUIOp::create(builder, loc, type, cast)
1852 }
else if (inElemTy != type) {
1853 cast = arith::BitcastOp::create(builder, loc, type, cast).getResult();
1867 auto singleElementToVectorMaterializationCast =
1868 [](OpBuilder &builder, Type type,
ValueRange inputs,
1869 Location loc) -> Value {
1870 if (inputs.size() != 1)
1872 auto input = inputs.front();
1873 auto inTy = input.getType();
1874 if (!inTy.isIntOrFloat())
1878 if (
auto vecTy = dyn_cast<VectorType>(type)) {
1879 if (vecTy.getRank() != 0 && vecTy.getNumElements() != 1)
1881 auto outElemTy = vecTy.getElementType();
1883 if (outElemTy.isIndex()) {
1884 cast = arith::IndexCastUIOp::create(builder, loc,
1887 }
else if (inTy != outElemTy) {
1888 cast = arith::BitcastOp::create(builder, loc, outElemTy, cast)
1891 return vector::BroadcastOp::create(builder, loc, vecTy, cast)
1896 typeConverter.addSourceMaterialization(
1897 singleElementToVectorMaterializationCast);
1898 typeConverter.addSourceMaterialization(vectorToVectorMaterializationCast);
1899 typeConverter.addTargetMaterialization(memrefToIntMaterializationCast);
1900 typeConverter.addTargetMaterialization(ui32ToI32MaterializationCast);
1901 typeConverter.addTargetMaterialization(ui64ToI64MaterializationCast);
1902 typeConverter.addTargetMaterialization(
1903 vectorToSingleElementMaterializationCast);
1904 typeConverter.addTargetMaterialization(vectorToVectorMaterializationCast);
1905 ConversionTarget
target(*context);
1906 target.addLegalDialect<xevm::XeVMDialect, LLVM::LLVMDialect,
1907 vector::VectorDialect, arith::ArithDialect,
1908 memref::MemRefDialect, gpu::GPUDialect,
1909 index::IndexDialect>();
1910 target.addIllegalDialect<xegpu::XeGPUDialect>();
1913 target.addDynamicallyLegalOp<arith::ExtFOp>(
1914 [](arith::ExtFOp op) {
return !isXeVMExtf(op); });
1915 target.addDynamicallyLegalOp<arith::TruncFOp>(
1916 [](arith::TruncFOp op) {
return !isXeVMTruncf(op); });
1918 RewritePatternSet patterns(context);
1922 if (
failed(applyPartialConversion(getOperation(),
target,
1923 std::move(patterns))))
1924 signalPassFailure();
1934 patterns.
add<CreateNdDescToXeVMPattern,
1935 LoadStorePrefetchNdToXeVMPattern<xegpu::LoadNdOp>,
1936 LoadStorePrefetchNdToXeVMPattern<xegpu::StoreNdOp>,
1937 LoadStorePrefetchNdToXeVMPattern<xegpu::PrefetchNdOp>>(
1939 patterns.
add<AtomicRMWToXeVMPattern, PrefetchToXeVMPattern,
1940 LoadStoreToXeVMPattern<xegpu::LoadGatherOp>,
1941 LoadStoreToXeVMPattern<xegpu::StoreScatterOp>>(
1943 patterns.
add<LoadStoreMatrixToXeVMPattern<xegpu::LoadMatrixOp>,
1944 LoadStoreMatrixToXeVMPattern<xegpu::StoreMatrixOp>,
1945 CreateMemDescOpPattern>(typeConverter, patterns.
getContext());
1946 patterns.
add<FenceToXeVMPattern, DpasToXeVMPattern>(typeConverter,
1948 patterns.
add<DpasMxToXeVMPattern>(typeConverter, patterns.
getContext());
1949 patterns.
add<ExtfToXeVMPattern, TruncfToXeVMPattern>(typeConverter,
1951 patterns.
add<LaneShuffleToXeVMPattern>(typeConverter, patterns.
getContext());
static llvm::ManagedStatic< PassManagerOptions > options
static Value broadcast(Location loc, Value toBroadcast, unsigned numElements, const TypeConverter &typeConverter, ConversionPatternRewriter &rewriter)
Broadcasts the value to vector with numElements number of elements.
Attributes are known-constant values of operations.
IntegerAttr getIndexAttr(int64_t value)
IntegerAttr getIntegerAttr(Type type, int64_t value)
IntegerType getIntegerType(unsigned width)
An attribute that represents a reference to a dense vector or tensor object.
bool isSplat() const
Returns true if this attribute corresponds to a splat, i.e.
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.
bool matchPattern(Value value, const Pattern &pattern)
Entry point for matching a pattern over a Value.
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)
detail::constant_op_matcher m_Constant()
Matches a constant foldable operation.