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 {
60static int32_t getNumericXeVMAddrSpace(xegpu::MemorySpace xeGpuMemspace) {
61 switch (xeGpuMemspace) {
62 case xegpu::MemorySpace::Global:
63 return static_cast<int>(xevm::AddrSpace::GLOBAL);
64 case xegpu::MemorySpace::SLM:
65 return static_cast<int>(xevm::AddrSpace::SHARED);
67 llvm_unreachable(
"Unknown XeGPU memory space");
77static FailureOr<unsigned> getNumericMemorySpace(
Attribute memSpace) {
80 if (
auto intAttr = llvm::dyn_cast<IntegerAttr>(memSpace))
81 return static_cast<unsigned>(intAttr.getInt());
82 if (
auto xevmSpace = llvm::dyn_cast<xevm::AddrSpaceAttr>(memSpace))
83 return static_cast<unsigned>(xevmSpace.getValue());
84 if (
auto gpuSpace = llvm::dyn_cast<gpu::AddressSpaceAttr>(memSpace)) {
85 switch (gpuSpace.getValue()) {
86 case gpu::AddressSpace::Global:
87 return static_cast<unsigned>(xevm::AddrSpace::GLOBAL);
88 case gpu::AddressSpace::Workgroup:
89 return static_cast<unsigned>(xevm::AddrSpace::SHARED);
90 case gpu::AddressSpace::Private:
91 return static_cast<unsigned>(xevm::AddrSpace::PRIVATE);
92 case gpu::AddressSpace::Constant:
93 return static_cast<unsigned>(xevm::AddrSpace::CONSTANT);
95 llvm_unreachable(
"Unknown GPU address space");
101static bool isSharedMemRef(
const MemRefType &memrefTy) {
102 FailureOr<unsigned> addrSpace =
103 getNumericMemorySpace(memrefTy.getMemorySpace());
104 return succeeded(addrSpace) &&
105 *addrSpace ==
static_cast<unsigned>(xevm::AddrSpace::SHARED);
109static VectorType encodeVectorTypeTo(VectorType currentVecType,
111 auto elemType = currentVecType.getElementType();
112 auto currentBitWidth = elemType.getIntOrFloatBitWidth();
115 currentVecType.getNumElements() * currentBitWidth / newBitWidth;
116 return VectorType::get(size, toElemType);
119static xevm::LoadCacheControl
120translateLoadXeGPUCacheHint(std::optional<xegpu::CachePolicy> L1hint,
121 std::optional<xegpu::CachePolicy> L3hint) {
123 if (!L1hint && !L3hint)
124 return xevm::LoadCacheControl::USE_DEFAULT;
126 auto L1hintVal = L1hint.value_or(xegpu::CachePolicy::CACHED);
127 auto L3hintVal = L3hint.value_or(xegpu::CachePolicy::CACHED);
129 case xegpu::CachePolicy::CACHED:
130 if (L3hintVal == xegpu::CachePolicy::CACHED)
131 return xevm::LoadCacheControl::L1C_L2UC_L3C;
132 else if (L3hintVal == xegpu::CachePolicy::UNCACHED)
133 return xevm::LoadCacheControl::L1C_L2UC_L3UC;
135 llvm_unreachable(
"Unsupported cache control.");
136 case xegpu::CachePolicy::UNCACHED:
137 if (L3hintVal == xegpu::CachePolicy::CACHED)
138 return xevm::LoadCacheControl::L1UC_L2UC_L3C;
139 else if (L3hintVal == xegpu::CachePolicy::UNCACHED)
140 return xevm::LoadCacheControl::L1UC_L2UC_L3UC;
142 llvm_unreachable(
"Unsupported cache control.");
143 case xegpu::CachePolicy::STREAMING:
144 if (L3hintVal == xegpu::CachePolicy::CACHED)
145 return xevm::LoadCacheControl::L1S_L2UC_L3C;
146 else if (L3hintVal == xegpu::CachePolicy::UNCACHED)
147 return xevm::LoadCacheControl::L1S_L2UC_L3UC;
149 llvm_unreachable(
"Unsupported cache control.");
150 case xegpu::CachePolicy::READ_INVALIDATE:
151 return xevm::LoadCacheControl::INVALIDATE_READ;
153 llvm_unreachable(
"Unsupported cache control.");
157static xevm::StoreCacheControl
158translateStoreXeGPUCacheHint(std::optional<xegpu::CachePolicy> L1hint,
159 std::optional<xegpu::CachePolicy> L3hint) {
161 if (!L1hint && !L3hint)
162 return xevm::StoreCacheControl::USE_DEFAULT;
164 auto L1hintVal = L1hint.value_or(xegpu::CachePolicy::UNCACHED);
165 auto L3hintVal = L3hint.value_or(xegpu::CachePolicy::WRITE_BACK);
167 case xegpu::CachePolicy::UNCACHED:
168 if (L3hintVal == xegpu::CachePolicy::UNCACHED)
169 return xevm::StoreCacheControl::L1UC_L2UC_L3UC;
170 else if (L3hintVal == xegpu::CachePolicy::WRITE_BACK)
171 return xevm::StoreCacheControl::L1UC_L2UC_L3WB;
173 llvm_unreachable(
"Unsupported cache control.");
174 case xegpu::CachePolicy::STREAMING:
175 if (L3hintVal == xegpu::CachePolicy::UNCACHED)
176 return xevm::StoreCacheControl::L1S_L2UC_L3UC;
177 else if (L3hintVal == xegpu::CachePolicy::WRITE_BACK)
178 return xevm::StoreCacheControl::L1S_L2UC_L3WB;
180 llvm_unreachable(
"Unsupported cache control.");
181 case xegpu::CachePolicy::WRITE_BACK:
182 if (L3hintVal == xegpu::CachePolicy::UNCACHED)
183 return xevm::StoreCacheControl::L1WB_L2UC_L3UC;
184 else if (L3hintVal == xegpu::CachePolicy::WRITE_BACK)
185 return xevm::StoreCacheControl::L1WB_L2UC_L3WB;
187 llvm_unreachable(
"Unsupported cache control.");
188 case xegpu::CachePolicy::WRITE_THROUGH:
189 if (L3hintVal == xegpu::CachePolicy::UNCACHED)
190 return xevm::StoreCacheControl::L1WT_L2UC_L3UC;
191 else if (L3hintVal == xegpu::CachePolicy::WRITE_BACK)
192 return xevm::StoreCacheControl::L1WT_L2UC_L3WB;
194 llvm_unreachable(
"Unsupported cache control.");
196 llvm_unreachable(
"Unsupported cache control.");
208class CreateNdDescToXeVMPattern
209 :
public OpConversionPattern<xegpu::CreateNdDescOp> {
210 using OpConversionPattern::OpConversionPattern;
212 matchAndRewrite(xegpu::CreateNdDescOp op,
213 xegpu::CreateNdDescOp::Adaptor adaptor,
214 ConversionPatternRewriter &rewriter)
const override {
215 auto loc = op.getLoc();
216 auto source = op.getSource();
220 Type payloadElemTy = rewriter.getI32Type();
221 VectorType payloadTy = VectorType::get(8, payloadElemTy);
222 Type i64Ty = rewriter.getI64Type();
224 VectorType payloadI64Ty = VectorType::get(4, i64Ty);
226 Value payload = arith::ConstantOp::create(
235 SmallVector<OpFoldResult> mixedSizes = op.getMixedSizes();
236 SmallVector<OpFoldResult> mixedStrides = op.getMixedStrides();
238 int64_t rank = mixedSizes.size();
239 auto sourceTy = source.getType();
240 auto sourceMemrefTy = dyn_cast<MemRefType>(sourceTy);
243 if (sourceMemrefTy) {
244 if (!sourceMemrefTy.hasRank()) {
245 return rewriter.notifyMatchFailure(op,
"Expected ranked Memref.");
249 baseAddr = adaptor.getSource();
251 baseAddr = adaptor.getSource();
252 if (baseAddr.
getType() != i64Ty) {
254 baseAddr = arith::ExtUIOp::create(rewriter, loc, i64Ty, baseAddr);
259 rewriter.replaceOp(op, baseAddr);
263 auto createOffset = [&](SmallVector<OpFoldResult> &ofrVec,
264 unsigned idx) -> Value {
271 baseShapeW = createOffset(mixedSizes, rank - 1);
272 baseShapeH = createOffset(mixedSizes, rank - 2);
274 Value basePitch = createOffset(mixedStrides, rank - 2);
277 vector::BitCastOp::create(rewriter, loc, payloadI64Ty, payload);
279 vector::InsertOp::create(rewriter, loc, baseAddr, payLoadAsI64,
280 static_cast<int>(NdTdescOffset::BasePtr));
281 payload = vector::BitCastOp::create(rewriter, loc, payloadTy, payLoadAsI64);
283 vector::InsertOp::create(rewriter, loc, baseShapeW, payload,
284 static_cast<int>(NdTdescOffset::BaseShapeW));
286 vector::InsertOp::create(rewriter, loc, baseShapeH, payload,
287 static_cast<int>(NdTdescOffset::BaseShapeH));
289 vector::InsertOp::create(rewriter, loc, basePitch, payload,
290 static_cast<int>(NdTdescOffset::BasePitch));
291 rewriter.replaceOp(op, payload);
298 typename = std::enable_if_t<llvm::is_one_of<
299 OpType, xegpu::LoadNdOp, xegpu::StoreNdOp, xegpu::PrefetchNdOp>::value>>
300class LoadStorePrefetchNdToXeVMPattern :
public OpConversionPattern<OpType> {
301 using OpConversionPattern<OpType>::OpConversionPattern;
303 matchAndRewrite(OpType op,
typename OpType::Adaptor adaptor,
304 ConversionPatternRewriter &rewriter)
const override {
305 auto mixedOffsets = op.getMixedOffsets();
306 int64_t opOffsetsSize = mixedOffsets.size();
307 auto loc = op.getLoc();
308 auto ctxt = rewriter.getContext();
310 auto tdesc = adaptor.getTensorDesc();
311 auto tdescTy = op.getTensorDescType();
312 auto tileRank = tdescTy.getRank();
313 if (opOffsetsSize != tileRank)
314 return rewriter.notifyMatchFailure(
315 op,
"Expected offset rank to match descriptor rank.");
316 auto elemType = tdescTy.getElementType();
317 auto elemBitSize = elemType.getIntOrFloatBitWidth();
318 bool isSubByte = elemBitSize < 8;
319 uint64_t wScaleFactor = 1;
321 if (!isSubByte && (elemBitSize % 8 != 0))
322 return rewriter.notifyMatchFailure(
323 op,
"Expected element type bit width to be multiple of 8.");
324 auto tileW = tdescTy.getDimSize(tileRank - 1);
327 if (elemBitSize != 4)
328 return rewriter.notifyMatchFailure(
329 op,
"Only sub byte types of 4bits are supported.");
331 return rewriter.notifyMatchFailure(
332 op,
"Sub byte types are only supported for 2D tensor descriptors.");
333 auto subByteFactor = 8 / elemBitSize;
334 auto tileH = tdescTy.getDimSize(0);
336 if constexpr (std::is_same_v<OpType, xegpu::LoadNdOp>) {
337 if (op.getPacked().value_or(
false)) {
339 if (tileH == systolicDepth * 4 &&
340 tileW == executionSize * subByteFactor) {
345 elemType = rewriter.getIntegerType(8);
346 tileW = executionSize;
347 wScaleFactor = subByteFactor;
352 if (wScaleFactor == 1) {
353 auto sub16BitFactor = subByteFactor * 2;
354 if (tileW == executionSize * sub16BitFactor) {
358 elemType = rewriter.getIntegerType(16);
359 tileW = executionSize;
360 wScaleFactor = sub16BitFactor;
362 return rewriter.notifyMatchFailure(
363 op,
"Unsupported tile shape for sub byte types.");
367 elemBitSize = elemType.getIntOrFloatBitWidth();
371 auto ptrTypeLLVM = LLVM::LLVMPointerType::get(
372 ctxt, getNumericXeVMAddrSpace(tdescTy.getMemorySpace()));
376 rewriter, loc, rewriter.getI32Type(), elemBitSize / 8);
377 VectorType payloadI64Ty = VectorType::get(4, rewriter.getI64Type());
379 vector::BitCastOp::create(rewriter, loc, payloadI64Ty, tdesc);
381 vector::ExtractOp::create(rewriter, loc, payLoadAsI64,
382 static_cast<int>(NdTdescOffset::BasePtr));
383 Value baseShapeW = vector::ExtractOp::create(
384 rewriter, loc, tdesc,
static_cast<int>(NdTdescOffset::BaseShapeW));
385 Value baseShapeH = vector::ExtractOp::create(
386 rewriter, loc, tdesc,
static_cast<int>(NdTdescOffset::BaseShapeH));
387 Value basePitch = vector::ExtractOp::create(
388 rewriter, loc, tdesc,
static_cast<int>(NdTdescOffset::BasePitch));
394 mixedOffsets[tileRank - 1]);
396 rewriter.getI32Type(), offsetW);
398 mixedOffsets[tileRank - 2]);
400 rewriter.getI32Type(), offsetH);
403 LLVM::IntToPtrOp::create(rewriter, loc, ptrTypeLLVM, basePtr);
407 Value baseShapeWInBytes =
408 arith::MulIOp::create(rewriter, loc, baseShapeW, elemByteSize);
410 Value basePitchBytes =
411 arith::MulIOp::create(rewriter, loc, basePitch, elemByteSize);
413 if (wScaleFactor > 1) {
417 rewriter, loc, rewriter.getI32Type(), llvm::Log2_64(wScaleFactor));
418 baseShapeWInBytes = arith::ShRSIOp::create(
419 rewriter, loc, baseShapeWInBytes, wScaleFactorValLog2);
420 basePitchBytes = arith::ShRSIOp::create(rewriter, loc, basePitchBytes,
421 wScaleFactorValLog2);
423 arith::ShRSIOp::create(rewriter, loc, offsetW, wScaleFactorValLog2);
426 auto tileH = tdescTy.getDimSize(tileRank - 2);
428 int32_t vblocks = tdescTy.getArrayLength();
429 if constexpr (std::is_same_v<OpType, xegpu::StoreNdOp>) {
430 Value src = adaptor.getValue();
436 VectorType srcVecTy = dyn_cast<VectorType>(src.
getType());
438 return rewriter.notifyMatchFailure(
439 op,
"Expected store value to be a vector type.");
441 VectorType newSrcVecTy =
442 encodeVectorTypeTo(srcVecTy, rewriter.getIntegerType(elemBitSize));
443 if (srcVecTy != newSrcVecTy)
444 src = vector::BitCastOp::create(rewriter, loc, newSrcVecTy, src);
445 auto storeCacheControl =
446 translateStoreXeGPUCacheHint(op.getL1Hint(), op.getL3Hint());
447 xevm::BlockStore2dOp::create(
448 rewriter, loc, basePtrLLVM, baseShapeWInBytes, baseShapeH,
449 basePitchBytes, offsetW, offsetH, elemBitSize, tileW, tileH, src,
450 xevm::StoreCacheControlAttr::get(ctxt, storeCacheControl));
451 rewriter.eraseOp(op);
453 auto loadCacheControl =
454 translateLoadXeGPUCacheHint(op.getL1Hint(), op.getL3Hint());
455 if constexpr (std::is_same_v<OpType, xegpu::PrefetchNdOp>) {
456 xevm::BlockPrefetch2dOp::create(
457 rewriter, loc, basePtrLLVM, baseShapeWInBytes, baseShapeH,
458 basePitchBytes, offsetW, offsetH, elemBitSize, tileW, tileH,
459 vblocks, xevm::LoadCacheControlAttr::get(ctxt, loadCacheControl));
460 rewriter.eraseOp(op);
462 VectorType dstVecTy = cast<VectorType>(op.getValue().getType());
463 bool vnni = op.getPacked().value_or(
false);
464 auto transposeValue = op.getTranspose();
466 transposeValue.has_value() && transposeValue.value()[0] == 1;
472 if (elemBitSize == 8 && tileW == 16 && tileH == 32 && !vnni &&
480 if (transpose && elemBitSize < 32) {
481 int32_t scale = 32 / elemBitSize;
483 rewriter, loc, rewriter.getI32Type(), llvm::Log2_64(scale));
484 offsetW = arith::ShRSIOp::create(rewriter, loc, offsetW, scaleLog2);
485 tileW = tileW * elemBitSize / 32;
488 VectorType loadedTy = encodeVectorTypeTo(
489 dstVecTy, vnni ? rewriter.getI32Type()
490 : rewriter.getIntegerType(elemBitSize));
492 Value resultFlatVec = xevm::BlockLoad2dOp::create(
493 rewriter, loc, loadedTy, basePtrLLVM, baseShapeWInBytes,
494 baseShapeH, basePitchBytes, offsetW, offsetH, elemBitSize, tileW,
495 tileH, vblocks, transpose, vnni,
496 xevm::LoadCacheControlAttr::get(ctxt, loadCacheControl));
497 resultFlatVec = vector::BitCastOp::create(
499 encodeVectorTypeTo(loadedTy, dstVecTy.getElementType()),
501 rewriter.replaceOp(op, resultFlatVec);
513 rewriter.getI64Type(), offset);
516 rewriter, loc, rewriter.getI64Type(), elemBitSize / 8);
518 rewriter.createOrFold<arith::MulIOp>(loc, offset, elemByteSize);
520 Value finalAddrI64 = rewriter.createOrFold<arith::AddIOp>(
526 LLVM::IntToPtrOp::create(rewriter, loc, ptrTypeLLVM, finalAddrI64);
527 if constexpr (std::is_same_v<OpType, xegpu::StoreNdOp>) {
528 Value src = adaptor.getValue();
534 VectorType srcVecTy = dyn_cast<VectorType>(src.
getType());
536 return rewriter.notifyMatchFailure(
537 op,
"Expected store value to be a vector type.");
539 VectorType newSrcVecTy =
540 encodeVectorTypeTo(srcVecTy, rewriter.getIntegerType(elemBitSize));
541 if (srcVecTy != newSrcVecTy)
542 src = vector::BitCastOp::create(rewriter, loc, newSrcVecTy, src);
543 auto storeCacheControl =
544 translateStoreXeGPUCacheHint(op.getL1Hint(), op.getL3Hint());
545 rewriter.replaceOpWithNewOp<xevm::BlockStoreOp>(
546 op, finalPtrLLVM, src,
547 xevm::StoreCacheControlAttr::get(ctxt, storeCacheControl));
548 }
else if constexpr (std::is_same_v<OpType, xegpu::LoadNdOp>) {
549 auto loadCacheControl =
550 translateLoadXeGPUCacheHint(op.getL1Hint(), op.getL3Hint());
551 VectorType resTy = cast<VectorType>(op.getValue().getType());
552 VectorType loadedTy =
553 encodeVectorTypeTo(resTy, rewriter.getIntegerType(elemBitSize));
554 Value
load = xevm::BlockLoadOp::create(
555 rewriter, loc, loadedTy, finalPtrLLVM,
556 xevm::LoadCacheControlAttr::get(ctxt, loadCacheControl));
557 if (loadedTy != resTy)
558 load = vector::BitCastOp::create(rewriter, loc, resTy,
load);
559 rewriter.replaceOp(op,
load);
561 return rewriter.notifyMatchFailure(
562 op,
"Unsupported operation: xegpu.prefetch_nd with tensor "
563 "descriptor rank == 1");
572static Value addOffsetToBaseAddr(ConversionPatternRewriter &rewriter,
576 rewriter, loc, baseAddr.
getType(), elemByteSize);
577 Value byteOffset = arith::MulIOp::create(rewriter, loc, offset, byteSize);
578 Value newAddr = arith::AddIOp::create(rewriter, loc, baseAddr, byteOffset);
582template <
typename OpType,
583 typename = std::enable_if_t<llvm::is_one_of<
584 OpType, xegpu::LoadGatherOp, xegpu::StoreScatterOp>::value>>
585class LoadStoreToXeVMPattern :
public OpConversionPattern<OpType> {
586 using OpConversionPattern<OpType>::OpConversionPattern;
588 matchAndRewrite(OpType op,
typename OpType::Adaptor adaptor,
589 ConversionPatternRewriter &rewriter)
const override {
590 Value offset = adaptor.getOffsets();
592 return rewriter.notifyMatchFailure(op,
"Expected offset to be provided.");
593 auto loc = op.getLoc();
594 auto ctxt = rewriter.getContext();
598 if constexpr (std::is_same_v<OpType, xegpu::LoadGatherOp>)
600 this->getTypeConverter()->convertType(op.getResult().getType());
602 valOrResTy = adaptor.getValue().getType();
603 VectorType valOrResVecTy = dyn_cast<VectorType>(valOrResTy);
604 bool hasScalarVal = !valOrResVecTy;
605 int64_t elemBitWidth =
607 : valOrResVecTy.getElementType().getIntOrFloatBitWidth();
609 if (elemBitWidth % 8 != 0)
610 return rewriter.notifyMatchFailure(
611 op,
"Expected element type bit width to be multiple of 8.");
612 int64_t elemByteSize = elemBitWidth / 8;
614 LLVM::LLVMPointerType ptrTypeLLVM = LLVM::LLVMPointerType::get(
615 ctxt, getNumericXeVMAddrSpace(xegpu::MemorySpace::Global));
618 if constexpr (std::is_same_v<OpType, xegpu::LoadGatherOp>) {
619 basePtrI64 = adaptor.getSource();
620 if (
auto memRefTy = dyn_cast<MemRefType>(op.getSource().getType())) {
621 FailureOr<unsigned> addrSpace =
622 getNumericMemorySpace(memRefTy.getMemorySpace());
624 return rewriter.notifyMatchFailure(
625 op,
"Unsupported memref memory space attribute.");
627 ptrTypeLLVM = LLVM::LLVMPointerType::get(ctxt, *addrSpace);
630 basePtrI64 = adaptor.getDest();
631 if (
auto memRefTy = dyn_cast<MemRefType>(op.getDest().getType())) {
632 FailureOr<unsigned> addrSpace =
633 getNumericMemorySpace(memRefTy.getMemorySpace());
635 return rewriter.notifyMatchFailure(
636 op,
"Unsupported memref memory space attribute.");
638 ptrTypeLLVM = LLVM::LLVMPointerType::get(ctxt, *addrSpace);
642 if (basePtrI64.
getType() != rewriter.getI64Type()) {
643 basePtrI64 = arith::ExtUIOp::create(rewriter, loc, rewriter.getI64Type(),
646 Value mask = adaptor.getMask();
647 if (dyn_cast<VectorType>(offset.
getType())) {
650 return rewriter.notifyMatchFailure(op,
"Expected offset to be a scalar.");
656 addOffsetToBaseAddr(rewriter, loc, basePtrI64, offset, elemByteSize);
660 LLVM::IntToPtrOp::create(rewriter, loc, ptrTypeLLVM, basePtrI64);
663 VectorType maskVecTy = dyn_cast<VectorType>(mask.
getType());
667 return rewriter.notifyMatchFailure(op,
"Expected mask to be a scalar.");
670 if constexpr (std::is_same_v<OpType, xegpu::LoadGatherOp>) {
671 scf::IfOp ifOp = scf::IfOp::create(rewriter, loc, {valOrResTy},
672 maskForLane,
true,
true);
674 rewriter.setInsertionPointToStart(&ifOp.getThenRegion().front());
676 valOrResTy = VectorType::get({valOrResVecTy.getNumElements()},
677 valOrResVecTy.getElementType());
679 LLVM::LoadOp::create(rewriter, loc, valOrResTy, basePtrLLVM);
682 "cache_control", xevm::LoadCacheControlAttr::get(
683 ctxt, translateLoadXeGPUCacheHint(
684 op.getL1Hint(), op.getL3Hint())));
685 scf::YieldOp::create(rewriter, loc,
ValueRange{loaded});
686 rewriter.setInsertionPointToStart(&ifOp.getElseRegion().front());
688 auto eTy = hasScalarVal ? valOrResTy : valOrResVecTy.getElementType();
691 eVal = FloatAttr::get(eTy, 0.0);
693 eVal = IntegerAttr::get(eTy, 0);
695 loaded = arith::ConstantOp::create(rewriter, loc, eVal);
697 loaded = arith::ConstantOp::create(
699 scf::YieldOp::create(rewriter, loc,
ValueRange{loaded});
700 rewriter.replaceOp(op, ifOp.getResult(0));
703 scf::IfOp ifOp = scf::IfOp::create(rewriter, loc, maskForLane,
false);
704 auto body = ifOp.getBody();
705 rewriter.setInsertionPointToStart(body);
707 LLVM::StoreOp::create(rewriter, loc, adaptor.getValue(), basePtrLLVM);
709 storeOp.getOperation()->setAttr(
710 "cache_control", xevm::StoreCacheControlAttr::get(
711 ctxt, translateStoreXeGPUCacheHint(
712 op.getL1Hint(), op.getL3Hint())));
713 rewriter.eraseOp(op);
719class CreateMemDescOpPattern final
720 :
public OpConversionPattern<xegpu::CreateMemDescOp> {
722 using OpConversionPattern<xegpu::CreateMemDescOp>::OpConversionPattern;
724 matchAndRewrite(xegpu::CreateMemDescOp op, OpAdaptor adaptor,
725 ConversionPatternRewriter &rewriter)
const override {
727 rewriter.replaceOp(op, adaptor.getSource());
732template <
typename OpType,
733 typename = std::enable_if_t<llvm::is_one_of<
734 OpType, xegpu::LoadMatrixOp, xegpu::StoreMatrixOp>::value>>
735class LoadStoreMatrixToXeVMPattern :
public OpConversionPattern<OpType> {
736 using OpConversionPattern<OpType>::OpConversionPattern;
738 matchAndRewrite(OpType op,
typename OpType::Adaptor adaptor,
739 ConversionPatternRewriter &rewriter)
const override {
741 SmallVector<OpFoldResult> offsets = op.getMixedOffsets();
743 return rewriter.notifyMatchFailure(op,
"Expected offset to be provided.");
745 auto loc = op.getLoc();
746 auto ctxt = rewriter.getContext();
747 Value baseAddr32 = adaptor.getMemDesc();
748 Value mdescVal = op.getMemDesc();
751 if constexpr (std::is_same_v<OpType, xegpu::LoadMatrixOp>) {
752 Type resType = op.getResult().getType();
755 if (
auto vecType = dyn_cast<VectorType>(resType)) {
756 assert(llvm::count_if(vecType.getShape(),
757 [](int64_t d) { return d != 1; }) <= 1 &&
758 "Expected either 1D vector or nD with unit dimensions");
759 resType = VectorType::get({vecType.getNumElements()},
760 vecType.getElementType());
764 dataTy = adaptor.getData().getType();
765 VectorType valOrResVecTy = dyn_cast<VectorType>(dataTy);
767 valOrResVecTy = VectorType::get(1, dataTy);
769 int64_t elemBitWidth =
770 valOrResVecTy.getElementType().getIntOrFloatBitWidth();
772 if (elemBitWidth % 8 != 0)
773 return rewriter.notifyMatchFailure(
774 op,
"Expected element type bit width to be multiple of 8.");
775 int64_t elemByteSize = elemBitWidth / 8;
778 LLVM::LLVMPointerType ptrTypeLLVM = LLVM::LLVMPointerType::get(
779 ctxt, getNumericXeVMAddrSpace(xegpu::MemorySpace::SLM));
781 auto mdescTy = cast<xegpu::MemDescType>(mdescVal.
getType());
783 Value linearOffset = mdescTy.getLinearOffsets(rewriter, loc, offsets);
784 linearOffset = arith::IndexCastUIOp::create(
785 rewriter, loc, rewriter.getI32Type(), linearOffset);
786 Value basePtrI32 = addOffsetToBaseAddr(rewriter, loc, baseAddr32,
787 linearOffset, elemByteSize);
791 LLVM::IntToPtrOp::create(rewriter, loc, ptrTypeLLVM, basePtrI32);
793 if (op.getSubgroupBlockIoAttr()) {
797 Type intElemTy = rewriter.getIntegerType(elemBitWidth);
798 VectorType intVecTy =
799 VectorType::get(valOrResVecTy.getShape(), intElemTy);
801 if constexpr (std::is_same_v<OpType, xegpu::LoadMatrixOp>) {
803 xevm::BlockLoadOp::create(rewriter, loc, intVecTy, basePtrLLVM);
804 if (intVecTy != valOrResVecTy) {
806 vector::BitCastOp::create(rewriter, loc, valOrResVecTy, loadOp);
808 rewriter.replaceOp(op, loadOp);
810 Value dataToStore = adaptor.getData();
811 if (valOrResVecTy != intVecTy) {
813 vector::BitCastOp::create(rewriter, loc, intVecTy, dataToStore);
815 xevm::BlockStoreOp::create(rewriter, loc, basePtrLLVM, dataToStore,
817 rewriter.eraseOp(op);
822 if (valOrResVecTy.getNumElements() >= 1) {
825 (*chipOpt !=
"pvc" && *chipOpt !=
"bmg" && *chipOpt !=
"cri")) {
827 return rewriter.notifyMatchFailure(
828 op,
"The lowering is specific to pvc, bmg or cri.");
832 if constexpr (std::is_same_v<OpType, xegpu::LoadMatrixOp>) {
839 this->getTypeConverter()->convertType(op.getResult().getType());
840 auto loadOp = LLVM::LoadOp::create(rewriter, loc, loadTy, basePtrLLVM);
841 rewriter.replaceOp(op, loadOp);
843 LLVM::StoreOp::create(rewriter, loc, adaptor.getData(), basePtrLLVM);
844 rewriter.eraseOp(op);
850class PrefetchToXeVMPattern :
public OpConversionPattern<xegpu::PrefetchOp> {
851 using OpConversionPattern::OpConversionPattern;
853 matchAndRewrite(xegpu::PrefetchOp op, xegpu::PrefetchOp::Adaptor adaptor,
854 ConversionPatternRewriter &rewriter)
const override {
855 auto loc = op.getLoc();
856 auto ctxt = rewriter.getContext();
857 Value basePtrI64 = adaptor.getSource();
859 if (basePtrI64.
getType() != rewriter.getI64Type())
860 basePtrI64 = arith::ExtUIOp::create(rewriter, loc, rewriter.getI64Type(),
862 Value offsets = adaptor.getOffsets();
864 VectorType offsetsVecTy = dyn_cast<VectorType>(offsets.
getType());
867 return rewriter.notifyMatchFailure(op,
868 "Expected offsets to be a scalar.");
870 int64_t elemBitWidth{0};
871 int64_t elemByteSize;
873 if (
auto memRefTy = dyn_cast<MemRefType>(op.getSourceType())) {
876 elemBitWidth = memRefTy.getElementType().getIntOrFloatBitWidth();
879 elemByteSize = *op.getOffsetAlignByte();
881 if (elemBitWidth != 0) {
882 if (elemBitWidth % 8 != 0)
883 return rewriter.notifyMatchFailure(
884 op,
"Expected element type bit width to be multiple of 8.");
885 elemByteSize = elemBitWidth / 8;
887 basePtrI64 = addOffsetToBaseAddr(rewriter, loc, basePtrI64, offsets,
892 LLVM::LLVMPointerType ptrTypeLLVM = LLVM::LLVMPointerType::get(
893 ctxt, getNumericXeVMAddrSpace(xegpu::MemorySpace::Global));
895 if (
auto memRefTy = dyn_cast<MemRefType>(op.getSource().getType())) {
896 FailureOr<unsigned> addrSpace =
897 getNumericMemorySpace(memRefTy.getMemorySpace());
899 return rewriter.notifyMatchFailure(
900 op,
"Unsupported memref memory space attribute.");
902 ptrTypeLLVM = LLVM::LLVMPointerType::get(ctxt, *addrSpace);
906 LLVM::IntToPtrOp::create(rewriter, loc, ptrTypeLLVM, basePtrI64);
908 xevm::PrefetchOp::create(
909 rewriter, loc, ptrLLVM,
910 xevm::LoadCacheControlAttr::get(
911 ctxt, translateLoadXeGPUCacheHint(op.getL1Hint(), op.getL3Hint())));
912 rewriter.eraseOp(op);
917class FenceToXeVMPattern :
public OpConversionPattern<xegpu::FenceOp> {
918 using OpConversionPattern::OpConversionPattern;
920 matchAndRewrite(xegpu::FenceOp op, xegpu::FenceOp::Adaptor adaptor,
921 ConversionPatternRewriter &rewriter)
const override {
922 auto loc = op.getLoc();
923 xevm::MemScope memScope{xevm::MemScope::WORKGROUP};
924 switch (op.getFenceScope()) {
925 case xegpu::FenceScope::Workgroup:
926 memScope = xevm::MemScope::WORKGROUP;
928 case xegpu::FenceScope::GPU:
929 memScope = xevm::MemScope::DEVICE;
932 xevm::AddrSpace addrSpace{xevm::AddrSpace::GLOBAL};
933 switch (op.getMemoryKind()) {
934 case xegpu::MemorySpace::Global:
935 addrSpace = xevm::AddrSpace::GLOBAL;
937 case xegpu::MemorySpace::SLM:
938 addrSpace = xevm::AddrSpace::SHARED;
941 xevm::MemfenceOp::create(rewriter, loc, memScope, addrSpace);
942 rewriter.eraseOp(op);
947static auto encodePrecision = [](
Type type) -> xevm::ElemType {
949 return xevm::ElemType::BF16;
950 else if (type.isF16())
951 return xevm::ElemType::F16;
952 else if (type.isTF32())
953 return xevm::ElemType::TF32;
954 else if (type.isInteger(8)) {
955 if (type.isUnsignedInteger())
956 return xevm::ElemType::U8;
957 return xevm::ElemType::S8;
958 }
else if (type.isF32())
959 return xevm::ElemType::F32;
960 else if (type.isInteger(32))
961 return xevm::ElemType::S32;
962 else if (type.isF8E5M2())
963 return xevm::ElemType::BF8;
964 else if (type.isF8E4M3FN())
965 return xevm::ElemType::F8;
966 else if (mlir::isa<Float4E2M1FNType>(type))
967 return xevm::ElemType::E2M1;
968 llvm_unreachable(
"add more support for ElemType");
971static unsigned getNumOperandsPerDword(xevm::ElemType pTy) {
973 case xevm::ElemType::TF32:
975 case xevm::ElemType::BF16:
976 case xevm::ElemType::F16:
978 case xevm::ElemType::U8:
979 case xevm::ElemType::S8:
980 case xevm::ElemType::F8:
981 case xevm::ElemType::BF8:
983 case xevm::ElemType::E2M1:
986 llvm_unreachable(
"unsupported xevm::ElemType");
990class DpasToXeVMPattern :
public OpConversionPattern<xegpu::DpasOp> {
991 using OpConversionPattern::OpConversionPattern;
993 matchAndRewrite(xegpu::DpasOp op, xegpu::DpasOp::Adaptor adaptor,
994 ConversionPatternRewriter &rewriter)
const override {
995 auto loc = op.getLoc();
996 auto ctxt = rewriter.getContext();
997 auto aTy = cast<VectorType>(op.getLhs().getType());
998 auto bTy = cast<VectorType>(op.getRhs().getType());
999 auto resultType = cast<VectorType>(op.getResultType());
1004 return rewriter.notifyMatchFailure(op,
"cannot determine target chip");
1008 return rewriter.notifyMatchFailure(op,
"unsupported target uArch");
1011 llvm::dyn_cast_or_null<xegpu::uArch::SubgroupMatrixMultiplyAcc>(
1012 uArch->getInstruction(
1013 xegpu::uArch::InstructionKind::SubgroupMatrixMultiplyAcc)));
1015 return rewriter.notifyMatchFailure(op,
1016 "DPAS not supported by target uArch");
1018 auto checkSupportedTypes = [&](VectorType vecTy,
1020 auto supported = dpasInst->getSupportedTypes(*ctxt, kind);
1021 return llvm::find(supported, vecTy.getElementType()) != supported.end();
1024 if (!checkSupportedTypes(aTy, xegpu::uArch::MMAOpndKind::MatrixA))
1025 return rewriter.notifyMatchFailure(
1026 op,
"A-matrix element type not supported by target uArch");
1027 if (!checkSupportedTypes(bTy, xegpu::uArch::MMAOpndKind::MatrixB))
1028 return rewriter.notifyMatchFailure(
1029 op,
"B-matrix element type not supported by target uArch");
1031 if (!checkSupportedTypes(resultType, xegpu::uArch::MMAOpndKind::MatrixD))
1032 return rewriter.notifyMatchFailure(
1033 op,
"result/accumulator element type not supported by target uArch");
1035 xevm::ElemType precATy = encodePrecision(aTy.getElementType());
1036 xevm::ElemType precBTy = encodePrecision(bTy.getElementType());
1037 Value c = op.getAcc();
1039 auto elementTy = resultType.getElementType();
1040 Attribute initValueAttr;
1041 if (isa<FloatType>(elementTy))
1042 initValueAttr = FloatAttr::get(elementTy, 0.0);
1044 initValueAttr = IntegerAttr::get(elementTy, 0);
1045 c = arith::ConstantOp::create(
1049 Value aVec = op.getLhs();
1050 Value bVec = op.getRhs();
1051 auto cvecty = cast<VectorType>(c.
getType());
1052 xevm::ElemType precCTy = encodePrecision(cvecty.getElementType());
1053 xevm::ElemType precDTy = encodePrecision(resultType.getElementType());
1055 VectorType::get(cvecty.getNumElements(), cvecty.getElementType());
1057 c = vector::ShapeCastOp::create(rewriter, loc, cNty, c);
1058 Value dpasRes = xevm::MMAOp::create(
1059 rewriter, loc, cNty, aVec, bVec, c,
1060 xevm::MMAShapeAttr::get(ctxt, cvecty.getNumElements(), executionSize,
1062 getNumOperandsPerDword(precATy)),
1063 xevm::MMATypesAttr::get(ctxt, precDTy, precATy, precBTy, precCTy));
1065 dpasRes = vector::ShapeCastOp::create(rewriter, loc, resultType, dpasRes);
1066 rewriter.replaceOp(op, dpasRes);
1071static std::optional<LLVM::AtomicBinOp>
1072matchSimpleAtomicOp(arith::AtomicRMWKind arithKind) {
1073 switch (arithKind) {
1074 case arith::AtomicRMWKind::addf:
1075 return LLVM::AtomicBinOp::fadd;
1076 case arith::AtomicRMWKind::addi:
1077 return LLVM::AtomicBinOp::add;
1078 case arith::AtomicRMWKind::assign:
1079 return LLVM::AtomicBinOp::xchg;
1080 case arith::AtomicRMWKind::maximumf:
1081 return LLVM::AtomicBinOp::fmax;
1082 case arith::AtomicRMWKind::maxs:
1083 return LLVM::AtomicBinOp::max;
1084 case arith::AtomicRMWKind::maxu:
1085 return LLVM::AtomicBinOp::umax;
1086 case arith::AtomicRMWKind::minimumf:
1087 return LLVM::AtomicBinOp::fmin;
1088 case arith::AtomicRMWKind::mins:
1089 return LLVM::AtomicBinOp::min;
1090 case arith::AtomicRMWKind::minu:
1091 return LLVM::AtomicBinOp::umin;
1092 case arith::AtomicRMWKind::ori:
1093 return LLVM::AtomicBinOp::_or;
1094 case arith::AtomicRMWKind::andi:
1095 return LLVM::AtomicBinOp::_and;
1097 return std::nullopt;
1101class AtomicRMWToXeVMPattern :
public OpConversionPattern<xegpu::AtomicRMWOp> {
1102 using OpConversionPattern::OpConversionPattern;
1104 matchAndRewrite(xegpu::AtomicRMWOp op, xegpu::AtomicRMWOp::Adaptor adaptor,
1105 ConversionPatternRewriter &rewriter)
const override {
1106 auto loc = op.getLoc();
1107 auto ctxt = rewriter.getContext();
1108 auto tdesc = op.getTensorDesc().getType();
1109 auto ptrTypeLLVM = LLVM::LLVMPointerType::get(
1110 ctxt, getNumericXeVMAddrSpace(tdesc.getMemorySpace()));
1111 Value basePtrI64 = arith::IndexCastOp::create(
1112 rewriter, loc, rewriter.getI64Type(), adaptor.getTensorDesc());
1114 LLVM::IntToPtrOp::create(rewriter, loc, ptrTypeLLVM, basePtrI64);
1115 VectorType srcOrDstVecTy = cast<VectorType>(op.getValue().getType());
1116 VectorType srcOrDstFlatVecTy = VectorType::get(
1117 srcOrDstVecTy.getNumElements(), srcOrDstVecTy.getElementType());
1118 Value srcFlatVec = vector::ShapeCastOp::create(
1119 rewriter, loc, srcOrDstFlatVecTy, op.getValue());
1120 auto atomicKind = matchSimpleAtomicOp(op.getKind());
1121 assert(atomicKind.has_value());
1122 Value resVec = srcFlatVec;
1123 for (
int i = 0; i < srcOrDstVecTy.getNumElements(); i++) {
1124 auto val = vector::ExtractOp::create(rewriter, loc, resVec, i);
1125 Value idx = LLVM::ConstantOp::create(rewriter, loc, rewriter.getI64Type(),
1126 rewriter.getIndexAttr(i));
1128 LLVM::GEPOp::create(rewriter, loc, ptrTypeLLVM,
1129 srcOrDstVecTy.getElementType(), basePtrLLVM, idx);
1131 LLVM::AtomicRMWOp::create(rewriter, loc, atomicKind.value(), currPtr,
1132 val, LLVM::AtomicOrdering::seq_cst);
1133 resVec = vector::InsertOp::create(rewriter, loc, newVal, resVec, i);
1135 rewriter.replaceOp(op, resVec);
1140class DpasMxToXeVMPattern :
public OpConversionPattern<xegpu::DpasMxOp> {
1141 using OpConversionPattern::OpConversionPattern;
1143 matchAndRewrite(xegpu::DpasMxOp op, xegpu::DpasMxOp::Adaptor adaptor,
1144 ConversionPatternRewriter &rewriter)
const override {
1145 auto loc = op.getLoc();
1146 auto ctxt = rewriter.getContext();
1147 auto aTy = op.getA().getType();
1148 auto bTy = op.getB().getType();
1150 cast<VectorType>(getTypeConverter()->convertType(op.getType()));
1154 return rewriter.notifyMatchFailure(op,
"cannot determine target chip");
1158 return rewriter.notifyMatchFailure(op,
"unsupported target uArch");
1162 xevm::ElemType precATy = encodePrecision(aTy.getElementType());
1163 xevm::ElemType precBTy = encodePrecision(bTy.getElementType());
1164 Value c = adaptor.getAcc();
1166 auto elementTy = resVecTy.getElementType();
1167 Attribute initValueAttr;
1168 if (isa<FloatType>(elementTy))
1169 initValueAttr = FloatAttr::get(elementTy, 0.0);
1171 initValueAttr = IntegerAttr::get(elementTy, 0);
1172 c = arith::ConstantOp::create(
1176 Value aVec = adaptor.getA();
1177 Value bVec = adaptor.getB();
1178 auto aVecTy = cast<VectorType>(aVec.
getType());
1179 auto bVecTy = cast<VectorType>(bVec.
getType());
1180 if (aVecTy.getElementTypeBitWidth() == 4)
1181 aVec = vector::BitCastOp::create(
1183 VectorType::get(aVecTy.getNumElements() / 2, rewriter.getI8Type()),
1185 if (bVecTy.getElementTypeBitWidth() == 4)
1186 bVec = vector::BitCastOp::create(
1188 VectorType::get(bVecTy.getNumElements() / 2, rewriter.getI8Type()),
1190 auto cVecTy = cast<VectorType>(c.
getType());
1191 xevm::ElemType precCTy = encodePrecision(cVecTy.getElementType());
1192 xevm::ElemType precDTy = encodePrecision(resVecTy.getElementType());
1193 Value scaleA = adaptor.getScaleA();
1194 Value scaleB = adaptor.getScaleB();
1195 Value dpasMxRes = xevm::MMAMxOp::create(
1196 rewriter, loc, resVecTy, aVec, bVec, scaleA, scaleB, c,
1197 xevm::MMAShapeAttr::get(ctxt, cVecTy.getNumElements(), executionSize,
1199 getNumOperandsPerDword(precATy)),
1200 xevm::MMATypesAttr::get(ctxt, precDTy, precATy, precBTy, precCTy));
1201 rewriter.replaceOp(op, dpasMxRes);
1220static constexpr int64_t kXeVMExtfTruncfNumElems = 16;
1223static std::optional<xevm::ExtfSrcElemTypes> getExtfNarrowType(
Type etype) {
1224 if (isa<Float8E5M2Type>(etype))
1225 return xevm::ExtfSrcElemTypes::BF8;
1226 if (isa<Float8E4M3FNType>(etype))
1227 return xevm::ExtfSrcElemTypes::F8;
1228 if (isa<Float4E2M1FNType>(etype))
1229 return xevm::ExtfSrcElemTypes::E2M1;
1230 return std::nullopt;
1234static std::optional<xevm::TruncfDstElemTypes> getTruncfNarrowType(
Type etype) {
1235 if (isa<Float8E5M2Type>(etype))
1236 return xevm::TruncfDstElemTypes::BF8;
1237 if (isa<Float8E4M3FNType>(etype))
1238 return xevm::TruncfDstElemTypes::F8;
1239 if (isa<Float4E2M1FNType>(etype))
1240 return xevm::TruncfDstElemTypes::E2M1;
1241 return std::nullopt;
1246static bool isXeVMExtf(arith::ExtFOp op) {
1247 auto srcTy = dyn_cast<VectorType>(op.getIn().getType());
1248 auto dstTy = dyn_cast<VectorType>(op.getType());
1249 if (!srcTy || !dstTy || srcTy.getRank() != 1 || dstTy.getRank() != 1)
1251 if (dstTy.getNumElements() != kXeVMExtfTruncfNumElems)
1253 Type dstETy = dstTy.getElementType();
1256 return getExtfNarrowType(srcTy.getElementType()).has_value();
1262static bool isXeVMTruncf(arith::TruncFOp op) {
1263 auto srcTy = dyn_cast<VectorType>(op.getIn().getType());
1264 auto dstTy = dyn_cast<VectorType>(op.getType());
1265 if (!srcTy || !dstTy || srcTy.getRank() != 1 || dstTy.getRank() != 1)
1267 if (srcTy.getNumElements() != kXeVMExtfTruncfNumElems)
1269 Type srcETy = srcTy.getElementType();
1272 return getTruncfNarrowType(dstTy.getElementType()).has_value();
1275class ExtfToXeVMPattern :
public OpConversionPattern<arith::ExtFOp> {
1276 using OpConversionPattern::OpConversionPattern;
1278 matchAndRewrite(arith::ExtFOp op, OpAdaptor adaptor,
1279 ConversionPatternRewriter &rewriter)
const override {
1280 if (!isXeVMExtf(op))
1281 return rewriter.notifyMatchFailure(op,
"not a xevm.extf compatible extf");
1282 Location loc = op.getLoc();
1283 MLIRContext *ctx = op.getContext();
1284 auto srcVecTy = cast<VectorType>(op.getIn().getType());
1285 auto dstVecTy = cast<VectorType>(op.getType());
1286 xevm::ExtfSrcElemTypes srcEnum =
1287 *getExtfNarrowType(srcVecTy.getElementType());
1288 xevm::ExtfDstElemTypes dstEnum = dstVecTy.getElementType().isF16()
1289 ? xevm::ExtfDstElemTypes::F16
1290 : xevm::ExtfDstElemTypes::BF16;
1294 Value src = adaptor.getIn();
1295 auto convSrcTy = cast<VectorType>(src.
getType());
1296 if (convSrcTy.getElementTypeBitWidth() == 4)
1297 src = vector::BitCastOp::create(
1299 VectorType::get(convSrcTy.getNumElements() / 2, rewriter.getI8Type()),
1301 Type resTy = getTypeConverter()->convertType(dstVecTy);
1302 Value res = xevm::ExtfOp::create(
1303 rewriter, loc, resTy, src, xevm::ExtfSrcElemTypeAttr::get(ctx, srcEnum),
1304 xevm::ExtfDstElemTypeAttr::get(ctx, dstEnum));
1305 rewriter.replaceOp(op, res);
1310class TruncfToXeVMPattern :
public OpConversionPattern<arith::TruncFOp> {
1311 using OpConversionPattern::OpConversionPattern;
1313 matchAndRewrite(arith::TruncFOp op, OpAdaptor adaptor,
1314 ConversionPatternRewriter &rewriter)
const override {
1315 if (!isXeVMTruncf(op))
1316 return rewriter.notifyMatchFailure(op,
1317 "not a xevm.truncf compatible truncf");
1318 Location loc = op.getLoc();
1319 MLIRContext *ctx = op.getContext();
1320 auto srcVecTy = cast<VectorType>(op.getIn().getType());
1321 auto dstVecTy = cast<VectorType>(op.getType());
1322 xevm::TruncfSrcElemTypes srcEnum = srcVecTy.getElementType().isF16()
1323 ? xevm::TruncfSrcElemTypes::F16
1324 : xevm::TruncfSrcElemTypes::BF16;
1325 xevm::TruncfDstElemTypes dstEnum =
1326 *getTruncfNarrowType(dstVecTy.getElementType());
1328 int64_t numNarrowBits =
1329 dstVecTy.getNumElements() * dstVecTy.getElementTypeBitWidth();
1330 Type packedTy = VectorType::get(numNarrowBits / 8, rewriter.getI8Type());
1332 xevm::TruncfOp::create(rewriter, loc, packedTy, adaptor.getIn(),
1333 xevm::TruncfSrcElemTypeAttr::get(ctx, srcEnum),
1334 xevm::TruncfDstElemTypeAttr::get(ctx, dstEnum));
1336 Type resTy = getTypeConverter()->convertType(dstVecTy);
1338 res = vector::BitCastOp::create(rewriter, loc, resTy, res);
1339 rewriter.replaceOp(op, res);
1348struct ConvertXeGPUToXeVMPass
1349 :
public impl::ConvertXeGPUToXeVMPassBase<ConvertXeGPUToXeVMPass> {
1352 void runOnOperation()
override {
1362 LowerToLLVMOptions
options(context);
1363 options.overrideIndexBitwidth(this->use64bitIndex ? 64 : 32);
1364 LLVMTypeConverter typeConverter(context,
options);
1366 Type xevmIndexType = typeConverter.convertType(IndexType::get(context));
1367 Type i32Type = IntegerType::get(context, 32);
1368 typeConverter.addConversion([&](VectorType type) -> Type {
1369 auto elemType = typeConverter.convertType(type.getElementType());
1371 unsigned rank = type.getRank();
1372 if (rank == 0 || type.getNumElements() == 1)
1375 int64_t sum = llvm::product_of(type.getShape());
1376 return VectorType::get(sum, elemType);
1378 typeConverter.addConversion([&](xegpu::TensorDescType type) -> Type {
1379 if (type.getRank() == 1)
1380 return xevmIndexType;
1381 return VectorType::get(8, i32Type);
1390 typeConverter.addConversion(
1391 [&](xegpu::MemDescType type) -> Type {
return i32Type; });
1393 typeConverter.addConversion([&](MemRefType type) -> Type {
1394 return isSharedMemRef(type) ? i32Type : xevmIndexType;
1404 auto memrefToIntMaterializationCast = [](OpBuilder &builder, Type type,
1406 Location loc) -> Value {
1407 if (inputs.size() != 1)
1409 auto input = inputs.front();
1410 if (
auto memrefTy = dyn_cast<MemRefType>(input.getType())) {
1411 unsigned rank = memrefTy.getRank();
1415 SmallVector<int64_t> intStrides;
1418 if (succeeded(memrefTy.getStridesAndOffset(intStrides, intOffsets)) &&
1419 ShapedType::isStatic(intOffsets)) {
1420 addr = memref::ExtractAlignedPointerAsIndexOp::create(builder, loc,
1422 offset = arith::ConstantOp::create(builder, loc,
1428 SmallVector<Type> resultTypes{
1429 MemRefType::get({}, memrefTy.getElementType(),
1430 MemRefLayoutAttrInterface(),
1431 memrefTy.getMemorySpace()),
1434 resultTypes.append(2 * rank, indexType);
1436 auto meta = memref::ExtractStridedMetadataOp::create(
1437 builder, loc, resultTypes, input);
1439 addr = memref::ExtractAlignedPointerAsIndexOp::create(
1440 builder, loc, meta.getBaseBuffer());
1441 offset = meta.getOffset();
1445 arith::IndexCastUIOp::create(builder, loc, type, addr);
1447 arith::IndexCastUIOp::create(builder, loc, type, offset);
1450 auto byteSize = arith::ConstantOp::create(
1453 memrefTy.getElementTypeBitWidth() / 8));
1455 arith::MulIOp::create(builder, loc, offsetCasted, byteSize);
1456 auto addrWithOffset =
1457 arith::AddIOp::create(builder, loc, addrCasted, byteOffset);
1459 return addrWithOffset.getResult();
1468 auto ui64ToI64MaterializationCast = [](OpBuilder &builder, Type type,
1470 Location loc) -> Value {
1471 if (inputs.size() != 1)
1473 auto input = inputs.front();
1476 index::CastUOp::create(builder, loc, builder.
getIndexType(), input)
1478 return arith::IndexCastUIOp::create(builder, loc, type, cast)
1488 auto ui32ToI32MaterializationCast = [](OpBuilder &builder, Type type,
1490 Location loc) -> Value {
1491 if (inputs.size() != 1)
1493 auto input = inputs.front();
1496 index::CastUOp::create(builder, loc, builder.
getIndexType(), input)
1498 return arith::IndexCastUIOp::create(builder, loc, type, cast)
1508 auto vectorToVectorMaterializationCast = [](OpBuilder &builder, Type type,
1510 Location loc) -> Value {
1511 if (inputs.size() != 1)
1513 auto input = inputs.front();
1514 if (
auto vecTy = dyn_cast<VectorType>(input.getType())) {
1515 if (
auto targetVecTy = dyn_cast<VectorType>(type)) {
1519 if (targetVecTy.getShape() != vecTy.getShape()) {
1520 cast = vector::ShapeCastOp::create(
1522 VectorType::get(targetVecTy.getShape(),
1523 vecTy.getElementType()),
1527 if (targetVecTy.getElementType() != vecTy.getElementType()) {
1528 cast = vector::BitCastOp::create(builder, loc, targetVecTy, cast)
1540 auto vectorToSingleElementMaterializationCast =
1541 [](OpBuilder &builder, Type type,
ValueRange inputs,
1542 Location loc) -> Value {
1543 if (inputs.size() != 1)
1545 auto input = inputs.front();
1546 if (
auto vecTy = dyn_cast<VectorType>(input.getType())) {
1548 auto rank = vecTy.getRank();
1549 if (rank != 0 && vecTy.getNumElements() != 1)
1551 auto inElemTy = vecTy.getElementType();
1555 cast = vector::ExtractOp::create(builder, loc, cast, {}).getResult();
1557 cast = vector::ExtractOp::create(builder, loc, cast,
1558 SmallVector<int64_t>(rank, 0))
1565 if (inElemTy.isIndex()) {
1566 cast = arith::IndexCastUIOp::create(builder, loc, type, cast)
1568 }
else if (inElemTy != type) {
1569 cast = arith::BitcastOp::create(builder, loc, type, cast).getResult();
1583 auto singleElementToVectorMaterializationCast =
1584 [](OpBuilder &builder, Type type,
ValueRange inputs,
1585 Location loc) -> Value {
1586 if (inputs.size() != 1)
1588 auto input = inputs.front();
1589 auto inTy = input.getType();
1590 if (!inTy.isIntOrFloat())
1594 if (
auto vecTy = dyn_cast<VectorType>(type)) {
1595 if (vecTy.getRank() != 0 && vecTy.getNumElements() != 1)
1597 auto outElemTy = vecTy.getElementType();
1599 if (outElemTy.isIndex()) {
1600 cast = arith::IndexCastUIOp::create(builder, loc,
1603 }
else if (inTy != outElemTy) {
1604 cast = arith::BitcastOp::create(builder, loc, outElemTy, cast)
1607 return vector::BroadcastOp::create(builder, loc, vecTy, cast)
1612 typeConverter.addSourceMaterialization(
1613 singleElementToVectorMaterializationCast);
1614 typeConverter.addSourceMaterialization(vectorToVectorMaterializationCast);
1615 typeConverter.addTargetMaterialization(memrefToIntMaterializationCast);
1616 typeConverter.addTargetMaterialization(ui32ToI32MaterializationCast);
1617 typeConverter.addTargetMaterialization(ui64ToI64MaterializationCast);
1618 typeConverter.addTargetMaterialization(
1619 vectorToSingleElementMaterializationCast);
1620 typeConverter.addTargetMaterialization(vectorToVectorMaterializationCast);
1621 ConversionTarget
target(*context);
1622 target.addLegalDialect<xevm::XeVMDialect, LLVM::LLVMDialect,
1623 vector::VectorDialect, arith::ArithDialect,
1624 memref::MemRefDialect, gpu::GPUDialect,
1625 index::IndexDialect>();
1626 target.addIllegalDialect<xegpu::XeGPUDialect>();
1629 target.addDynamicallyLegalOp<arith::ExtFOp>(
1630 [](arith::ExtFOp op) {
return !isXeVMExtf(op); });
1631 target.addDynamicallyLegalOp<arith::TruncFOp>(
1632 [](arith::TruncFOp op) {
return !isXeVMTruncf(op); });
1634 RewritePatternSet patterns(context);
1638 if (
failed(applyPartialConversion(getOperation(),
target,
1639 std::move(patterns))))
1640 signalPassFailure();
1650 patterns.
add<CreateNdDescToXeVMPattern,
1651 LoadStorePrefetchNdToXeVMPattern<xegpu::LoadNdOp>,
1652 LoadStorePrefetchNdToXeVMPattern<xegpu::StoreNdOp>,
1653 LoadStorePrefetchNdToXeVMPattern<xegpu::PrefetchNdOp>>(
1655 patterns.
add<AtomicRMWToXeVMPattern, PrefetchToXeVMPattern,
1656 LoadStoreToXeVMPattern<xegpu::LoadGatherOp>,
1657 LoadStoreToXeVMPattern<xegpu::StoreScatterOp>>(
1659 patterns.
add<LoadStoreMatrixToXeVMPattern<xegpu::LoadMatrixOp>,
1660 LoadStoreMatrixToXeVMPattern<xegpu::StoreMatrixOp>,
1661 CreateMemDescOpPattern>(typeConverter, patterns.
getContext());
1662 patterns.
add<FenceToXeVMPattern, DpasToXeVMPattern>(typeConverter,
1664 patterns.
add<DpasMxToXeVMPattern>(typeConverter, patterns.
getContext());
1665 patterns.
add<ExtfToXeVMPattern, TruncfToXeVMPattern>(typeConverter,
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 setAttr(StringAttr name, Attribute value)
If the an attribute exists with the specified name, change it to the new value.
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)
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.
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)