MLIR 24.0.0git
XeGPUToXeVM.cpp
Go to the documentation of this file.
1//===-- XeGPUToXeVM.cpp - XeGPU to XeVM dialect conversion ------*- C++ -*-===//
2//
3// This file is licensed under the Apache License v2.0 with LLVM Exceptions.
4// See https://llvm.org/LICENSE.txt for license information.
5// SPDX-License-Identifier: Apache-2.0 WITH LLVM-exception
6//
7//===----------------------------------------------------------------------===//
8
12
27#include "mlir/Pass/Pass.h"
28#include "mlir/Support/LLVM.h"
29#include "llvm/ADT/STLExtras.h"
30#include "llvm/Support/FormatVariadic.h"
31
33#include "mlir/IR/Types.h"
34
35#include "llvm/ADT/TypeSwitch.h"
36
37#include <numeric>
38
39namespace mlir {
40#define GEN_PASS_DEF_CONVERTXEGPUTOXEVMPASS
41#include "mlir/Conversion/Passes.h.inc"
42} // namespace mlir
43
44using namespace mlir;
45
46namespace {
47
48// TODO: Below are uArch dependent values, should move away from hardcoding
49static constexpr int32_t systolicDepth{8};
50static constexpr int32_t executionSize{16};
51
52// Offsets to individual fields of the 8xi32 layout nd tensor descriptor.
53enum class NdTdescOffset : uint32_t {
54 BasePtr = 0, // Base pointer (i64)
55 BaseShapeW = 2, // Base shape width (i32)
56 BaseShapeH = 3, // Base shape height (i32)
57 BasePitch = 4, // Base pitch/stride of dim rank-2 (i32)
58 LeadingStride0 = 5, // Row strides of the leading (batch) dims of a >2D
59 LeadingStride1 = 6, // descriptor (i32); added into offset_h by the load/store
60 LeadingStride2 = 7, // lowering. Left at 0 for 2D descriptors.
61};
62
63// Spare payload slots above, and the resulting max lowerable descriptor rank.
64static constexpr int64_t maxNdTdescLeadingDims{3};
65static constexpr int64_t maxNdTdescRank{2 + maxNdTdescLeadingDims};
66
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);
73 }
74 llvm_unreachable("Unknown XeGPU memory space");
75}
76
77/// Translates a memref memory space attribute into XeVM's numeric address
78/// space, which follows the OpenCL/SPIR-V convention (0 = private, 1 =
79/// global, 2 = constant, 3 = shared/local, 4 = generic). A null attribute,
80/// meaning the memory space was left unspecified, maps to the default space
81/// 0. Returns failure if `memSpace` is a representation this pass does not
82/// know how to translate (e.g. a SPIR-V storage class or an arbitrary string
83/// attribute), rather than assuming it is an `IntegerAttr` and asserting.
84static FailureOr<unsigned> getNumericMemorySpace(Attribute memSpace) {
85 if (!memSpace)
86 return 0u;
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);
101 }
102 llvm_unreachable("Unknown GPU address space");
103 }
104 return failure();
105}
106
107/// Checks if the given MemRefType refers to shared memory.
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);
113}
114
115// Get same bitwidth flat vector type of new element type.
116static VectorType encodeVectorTypeTo(VectorType currentVecType,
117 Type toElemType) {
118 auto elemType = currentVecType.getElementType();
119 auto currentBitWidth = elemType.getIntOrFloatBitWidth();
120 auto newBitWidth = toElemType.getIntOrFloatBitWidth();
121 const int size =
122 currentVecType.getNumElements() * currentBitWidth / newBitWidth;
123 return VectorType::get(size, toElemType);
124}
125
126static xevm::LoadCacheControl
127translateLoadXeGPUCacheHint(std::optional<xegpu::CachePolicy> L1hint,
128 std::optional<xegpu::CachePolicy> L3hint) {
129 // If no hints are provided, use the default cache control.
130 if (!L1hint && !L3hint)
131 return xevm::LoadCacheControl::USE_DEFAULT;
132 // If only one of the hints is provided, use the default for the other level.
133 auto L1hintVal = L1hint.value_or(xegpu::CachePolicy::CACHED);
134 auto L3hintVal = L3hint.value_or(xegpu::CachePolicy::CACHED);
135 switch (L1hintVal) {
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;
141 else
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;
148 else
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;
155 else
156 llvm_unreachable("Unsupported cache control.");
157 case xegpu::CachePolicy::READ_INVALIDATE:
158 return xevm::LoadCacheControl::INVALIDATE_READ;
159 default:
160 llvm_unreachable("Unsupported cache control.");
161 }
162}
163
164static xevm::StoreCacheControl
165translateStoreXeGPUCacheHint(std::optional<xegpu::CachePolicy> L1hint,
166 std::optional<xegpu::CachePolicy> L3hint) {
167 // If no hints are provided, use the default cache control.
168 if (!L1hint && !L3hint)
169 return xevm::StoreCacheControl::USE_DEFAULT;
170 // If only one of the hints is provided, use the default for the other level.
171 auto L1hintVal = L1hint.value_or(xegpu::CachePolicy::UNCACHED);
172 auto L3hintVal = L3hint.value_or(xegpu::CachePolicy::WRITE_BACK);
173 switch (L1hintVal) {
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;
179 else
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;
186 else
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;
193 else
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;
200 else
201 llvm_unreachable("Unsupported cache control.");
202 default:
203 llvm_unreachable("Unsupported cache control.");
204 }
205}
206
207//
208// Note:
209// Block operations for tile of sub byte element types are handled by
210// emulating with larger element types.
211// Tensor descriptor are keep intact and only ops consuming them are
212// emulated
213//
214
215//
216// High-D (>2D) nd descriptors are lowered by viewing the source as a single
217// flattened 2D plane, so the 2D-block surface covers every leading (batch)
218// plane at once and a batch position becomes a row offset into it. The surface
219// height is the row extent the leading dims reach:
220// base_height = size[R-2] + sum_d (size[d] - 1) * (stride[d] / stride[R-2])
221//
222// Leaving `base_ptr` at the true base means an out-of-range batch index lands
223// past `base_height`, where the HW boundary check handles it, instead of aiming
224// the surface at unmapped memory. Encoding the leading strides as row counts
225// (`stride[d] / stride[R-2]`) also makes them dimensionless, so they survive
226// element-type repacking (e.g. the f16 -> i32 transpose repack) without a unit
227// conversion.
228//
229// Limitations of the flattened-plane view:
230// 1. Each leading stride must be a whole number of rows. A source with gaps
231// between planes (`stride[d] % stride[R-2] != 0`) is not lowered. This can
232// only be checked when the strides are static; for dynamic strides the
233// divisibility is assumed.
234// 2. `base_height` grows with the leading extent, so a source with a large
235// batch x head x sequence extent can exceed the HW 2D-block surface height.
236// 3. Plane boundaries are invisible to the boundary check, which only knows
237// the surface: a tile whose rows run past `size[R-2]` spills into the
238// following rows of the surface -- the next plane, or the padding between
239// planes -- rather than being clipped there as it would be on a per-plane
240// surface. A load reads those rows instead of returning zeros, and a store
241// overwrites them, so it can corrupt the next plane. This only matters
242// when `size[R-2]` is not a multiple of the tile height.
243//
244
245class CreateNdDescToXeVMPattern
246 : public OpConversionPattern<xegpu::CreateNdDescOp> {
247 using OpConversionPattern::OpConversionPattern;
248 LogicalResult
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();
254
255 // Check all failure conditions before generating any IR, so nothing has to
256 // be rolled back.
257 int64_t rank = op.getType().getRank();
258 int64_t sourceRank;
259 auto memrefTy = dyn_cast<MemRefType>(source.getType());
260 if (memrefTy) {
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();
266 } else {
267 return rewriter.notifyMatchFailure(op,
268 "Expected ranked Memref or integer.");
269 }
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) +
279 ").");
280 // Limitation 1 above; dynamic strides are assumed to divide evenly.
281 if (rank > 2) {
282 SmallVector<std::optional<int64_t>> constStrides(rank, std::nullopt);
283 if (memrefTy) {
284 SmallVector<int64_t> staticStrides;
285 int64_t staticOffset;
286 if (succeeded(
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];
291 } else {
292 SmallVector<OpFoldResult> mixed = op.getMixedStrides();
293 for (int64_t d = 0; d < rank; ++d)
294 constStrides[d] = getConstantIntValue(mixed[d]);
295 }
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.");
303 }
304 }
305 }
306
307 Type payloadElemTy = rewriter.getI32Type();
308 Type i64Ty = rewriter.getI64Type();
309
310 // Access the adaptor only after the failure checks, so a bail-out leaves no
311 // materialization cast behind.
312 Value baseAddr = adaptor.getSource();
313 if (isa<IntegerType>(source.getType()) && baseAddr.getType() != i64Ty) {
314 // Pointer type may be i32. Cast to i64 if needed.
315 baseAddr = arith::ExtUIOp::create(rewriter, loc, i64Ty, baseAddr);
316 }
317 // 1D tensor descriptor is just the base address.
318 if (rank == 1) {
319 rewriter.replaceOp(op, baseAddr);
320 return success();
321 }
322
323 SmallVector<OpFoldResult> mixedSizes;
324 SmallVector<OpFoldResult> mixedStrides;
325 if (memrefTy && !xegpu::hasStaticShapeAndStrides(memrefTy)) {
326 auto meta =
327 memref::ExtractStridedMetadataOp::create(rewriter, loc, source);
328 mixedSizes = meta.getConstifiedMixedSizes();
329 mixedStrides = meta.getConstifiedMixedStrides();
330 } else {
331 mixedSizes = op.getMixedSizes();
332 mixedStrides = op.getMixedStrides();
333 }
334
335 // Op is lowered to a code sequence that populates payload.
336 // Payload is a 8xi32 vector. Offset to individual fields are defined in
337 // NdTdescOffset enum.
338 VectorType payloadTy = VectorType::get(8, payloadElemTy);
339 // 4xi64 view is used for inserting the base pointer.
340 VectorType payloadI64Ty = VectorType::get(4, i64Ty);
341 // Initialize payload to zero.
342 Value payload = arith::ConstantOp::create(
343 rewriter, loc,
344 DenseElementsAttr::get(payloadTy, IntegerAttr::get(payloadElemTy, 0)));
345
346 // Utility for creating offset values from op fold result.
347 auto createOffset = [&](SmallVector<OpFoldResult> &ofrVec,
348 unsigned idx) -> Value {
349 Value val = getValueOrCreateConstantIntOp(rewriter, loc, ofrVec[idx]);
350 val = getValueOrCreateCastToIndexLike(rewriter, loc, payloadElemTy, val);
351 return val;
352 };
353 // The descriptor's innermost 2 dims are the 2D tile (H, W).
354 Value baseShapeW = createOffset(mixedSizes, rank - 1);
355 // Pitch is the stride of dim rank-2 (the row stride of the 2D tile).
356 Value basePitch = createOffset(mixedStrides, rank - 2);
357 // Leading (batch) strides as a number of rows (`stride[d] / pitch`).
358 SmallVector<Value> leadingRowStrides;
359 for (int64_t d = 0; d < rank - 2; ++d) {
360 std::optional<int64_t> leading = getConstantIntValue(mixedStrides[d]);
361 std::optional<int64_t> pitch =
362 getConstantIntValue(mixedStrides[rank - 2]);
363 if (leading && pitch && *pitch != 0)
364 leadingRowStrides.push_back(arith::ConstantIntOp::create(
365 rewriter, loc, payloadElemTy, *leading / *pitch));
366 else
367 leadingRowStrides.push_back(arith::DivUIOp::create(
368 rewriter, loc, createOffset(mixedStrides, d), basePitch));
369 }
370 // Height of the flattened plane: every leading (batch) plane is stacked
371 // into the surface, so the boundary check covers an out-of-range batch.
372 // base_height = size[R-2] + sum_d (size[d] - 1) * leadingRowStride[d]
373 // For rank 2 this is just size[0].
374 Value baseShapeH = createOffset(mixedSizes, rank - 2);
375 if (rank > 2) {
376 Value one = arith::ConstantIntOp::create(rewriter, loc, payloadElemTy, 1);
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]);
382 baseShapeH =
383 rewriter.createOrFold<arith::AddIOp>(loc, baseShapeH, rows);
384 }
385 }
386 // Populate payload.
387 Value payLoadAsI64 =
388 vector::BitCastOp::create(rewriter, loc, payloadI64Ty, payload);
389 payLoadAsI64 =
390 vector::InsertOp::create(rewriter, loc, baseAddr, payLoadAsI64,
391 static_cast<int>(NdTdescOffset::BasePtr));
392 payload = vector::BitCastOp::create(rewriter, loc, payloadTy, payLoadAsI64);
393 payload =
394 vector::InsertOp::create(rewriter, loc, baseShapeW, payload,
395 static_cast<int>(NdTdescOffset::BaseShapeW));
396 payload =
397 vector::InsertOp::create(rewriter, loc, baseShapeH, payload,
398 static_cast<int>(NdTdescOffset::BaseShapeH));
399 payload =
400 vector::InsertOp::create(rewriter, loc, basePitch, payload,
401 static_cast<int>(NdTdescOffset::BasePitch));
402 // The leading (batch) row strides go into the spare payload slots; the
403 // load/store/prefetch lowering turns the batch offsets into a row offset
404 // with them.
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);
410 return success();
411 }
412};
413
414template <
415 typename OpType,
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;
420 LogicalResult
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();
427
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) +
441 ".");
442 auto elemType = tdescTy.getElementType();
443 auto elemBitSize = elemType.getIntOrFloatBitWidth();
444 bool isSubByte = elemBitSize < 8;
445 uint64_t wScaleFactor = 1;
446
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);
451 // For sub byte types, only 4bits are currently supported.
452 if (isSubByte) {
453 if (elemBitSize != 4)
454 return rewriter.notifyMatchFailure(
455 op, "Only sub byte types of 4bits are supported.");
456 if (tileRank != 2)
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);
461 // Handle special case for packed load.
462 if constexpr (std::is_same_v<OpType, xegpu::LoadNdOp>) {
463 if (op.getPacked().value_or(false)) {
464 // packed load is implemented as packed loads of 8bit elements.
465 if (tileH == systolicDepth * 4 &&
466 tileW == executionSize * subByteFactor) {
467 // Usage case for loading as Matrix B with pack request.
468 // source is assumed to pre-packed into 8bit elements
469 // Emulate with 8bit loads with pack request.
470 // scaled_tileW = executionSize
471 elemType = rewriter.getIntegerType(8);
472 tileW = executionSize;
473 wScaleFactor = subByteFactor;
474 }
475 }
476 }
477 // If not handled by packed load case above, handle other cases.
478 if (wScaleFactor == 1) {
479 auto sub16BitFactor = subByteFactor * 2;
480 if (tileW == executionSize * sub16BitFactor) {
481 // Usage case for loading as Matrix A operand
482 // Emulate with 16bit loads/stores.
483 // scaled_tileW = executionSize
484 elemType = rewriter.getIntegerType(16);
485 tileW = executionSize;
486 wScaleFactor = sub16BitFactor;
487 } else {
488 return rewriter.notifyMatchFailure(
489 op, "Unsupported tile shape for sub byte types.");
490 }
491 }
492 // recompute element bit size for emulation.
493 elemBitSize = elemType.getIntOrFloatBitWidth();
494 }
495
496 // Get address space from tensor descriptor memory space.
497 auto ptrTypeLLVM = LLVM::LLVMPointerType::get(
498 ctxt, getNumericXeVMAddrSpace(tdescTy.getMemorySpace()));
499 if (tileRank >= 2) {
500 // Compute element byte size.
501 Value elemByteSize = arith::ConstantIntOp::create(
502 rewriter, loc, rewriter.getI32Type(), elemBitSize / 8);
503 VectorType payloadI64Ty = VectorType::get(4, rewriter.getI64Type());
504 Value payLoadAsI64 =
505 vector::BitCastOp::create(rewriter, loc, payloadI64Ty, tdesc);
506 Value basePtr =
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));
515
516 Value offsetW = getValueOrCreateConstantIntOp(rewriter, loc,
517 mixedOffsets[tileRank - 1]);
518 offsetW = getValueOrCreateCastToIndexLike(rewriter, loc,
519 rewriter.getI32Type(), offsetW);
520 Value offsetH = getValueOrCreateConstantIntOp(rewriter, loc,
521 mixedOffsets[tileRank - 2]);
522 offsetH = getValueOrCreateCastToIndexLike(rewriter, loc,
523 rewriter.getI32Type(), offsetH);
524 // Turn the leading (batch) offsets into a row offset into the flattened
525 // plane, using the row-unit batch strides encoded at create time:
526 // offsetH += sum_d offset[d] * leadingRowStride[d]
527 // The base pointer stays at the true base, so an out-of-range batch index
528 // is caught by the HW boundary check instead of moving the surface to
529 // unmapped memory.
530 for (int64_t d = 0; d < tileRank - 2; ++d) {
531 Value off =
532 getValueOrCreateConstantIntOp(rewriter, loc, mixedOffsets[d]);
533 off = getValueOrCreateCastToIndexLike(rewriter, loc,
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);
540 }
541 // Convert base pointer (i64) to LLVM pointer type.
542 Value basePtrLLVM =
543 LLVM::IntToPtrOp::create(rewriter, loc, ptrTypeLLVM, basePtr);
544 // FIXME: width or pitch is not the same as baseShapeW it should be the
545 // stride of the second to last dimension in row major layout.
546 // Compute width in bytes.
547 Value baseShapeWInBytes =
548 arith::MulIOp::create(rewriter, loc, baseShapeW, elemByteSize);
549 // Compute pitch in bytes.
550 Value basePitchBytes =
551 arith::MulIOp::create(rewriter, loc, basePitch, elemByteSize);
552
553 if (wScaleFactor > 1) {
554 // Scale offsetW, baseShapeWInBytes for sub byte emulation.
555 // Note: tileW is already scaled above.
556 Value wScaleFactorValLog2 = arith::ConstantIntOp::create(
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);
562 offsetW =
563 arith::ShRSIOp::create(rewriter, loc, offsetW, wScaleFactorValLog2);
564 }
565 // Get tile height from the tensor descriptor type (second-to-last dim).
566 auto tileH = tdescTy.getDimSize(tileRank - 2);
567 // Get vblocks from the tensor descriptor type.
568 int32_t vblocks = tdescTy.getArrayLength();
569 if constexpr (std::is_same_v<OpType, xegpu::StoreNdOp>) {
570 Value src = adaptor.getValue();
571 // If store value is a scalar, get value from op instead of adaptor.
572 // Adaptor might have optimized away single element vector
573 if (src.getType().isIntOrFloat()) {
574 src = op.getValue();
575 }
576 VectorType srcVecTy = dyn_cast<VectorType>(src.getType());
577 if (!srcVecTy)
578 return rewriter.notifyMatchFailure(
579 op, "Expected store value to be a vector type.");
580 // Get flat vector type of integer type with matching element bit size.
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);
592 } else {
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);
601 } else {
602 VectorType dstVecTy = cast<VectorType>(op.getValue().getType());
603 bool vnni = op.getPacked().value_or(false);
604 auto transposeValue = op.getTranspose();
605 bool transpose =
606 transposeValue.has_value() && transposeValue.value()[0] == 1;
607 // Handle special case of 32x16 and 8bit element load
608 // with no vnni, no transpose, no vblocks.
609 // For this special case, vnni and non vnni yields the same output
610 // and only the vnni variant is supported by HW.
611 // Check and set vnni of the special case.
612 if (elemBitSize == 8 && tileW == 16 && tileH == 32 && !vnni &&
613 !transpose) {
614 vnni = true;
615 }
616 // Handle tranpose request on small element size
617 // Transpose needs to be requested on 32bit element type.
618 // offsetW and tileW needs to be adjusted to account for element type
619 // change.
620 if (transpose && elemBitSize < 32) {
621 int32_t scale = 32 / elemBitSize;
622 Value scaleLog2 = arith::ConstantIntOp::create(
623 rewriter, loc, rewriter.getI32Type(), llvm::Log2_64(scale));
624 offsetW = arith::ShRSIOp::create(rewriter, loc, offsetW, scaleLog2);
625 tileW = tileW * elemBitSize / 32;
626 elemBitSize = 32;
627 }
628 VectorType loadedTy = encodeVectorTypeTo(
629 dstVecTy, vnni ? rewriter.getI32Type()
630 : rewriter.getIntegerType(elemBitSize));
631
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(
638 rewriter, loc,
639 encodeVectorTypeTo(loadedTy, dstVecTy.getElementType()),
640 resultFlatVec);
641 rewriter.replaceOp(op, resultFlatVec);
642 }
643 }
644 } else {
645 // 1D tensor descriptor.
646 // `tdesc` represents base address as i64
647 // Offset in number of elements, need to multiply by element byte size.
648 // Compute byte offset.
649 // byteOffset = offset * elementByteSize
650 Value offset =
651 getValueOrCreateConstantIntOp(rewriter, loc, mixedOffsets[0]);
652 offset = getValueOrCreateCastToIndexLike(rewriter, loc,
653 rewriter.getI64Type(), offset);
654 // Compute element byte size.
655 Value elemByteSize = arith::ConstantIntOp::create(
656 rewriter, loc, rewriter.getI64Type(), elemBitSize / 8);
657 Value byteOffset =
658 rewriter.createOrFold<arith::MulIOp>(loc, offset, elemByteSize);
659 // Final address = basePtr + byteOffset
660 Value finalAddrI64 = rewriter.createOrFold<arith::AddIOp>(
661 loc, tdesc,
662 getValueOrCreateCastToIndexLike(rewriter, loc, rewriter.getI64Type(),
663 byteOffset));
664 // Convert base pointer (i64) to LLVM pointer type.
665 Value finalPtrLLVM =
666 LLVM::IntToPtrOp::create(rewriter, loc, ptrTypeLLVM, finalAddrI64);
667 if constexpr (std::is_same_v<OpType, xegpu::StoreNdOp>) {
668 Value src = adaptor.getValue();
669 // If store value is a scalar, get value from op instead of adaptor.
670 // Adaptor might have optimized away single element vector
671 if (src.getType().isIntOrFloat()) {
672 src = op.getValue();
673 }
674 VectorType srcVecTy = dyn_cast<VectorType>(src.getType());
675 if (!srcVecTy)
676 return rewriter.notifyMatchFailure(
677 op, "Expected store value to be a vector type.");
678 // Get flat vector type of integer type with matching element bit size.
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);
700 } else {
701 return rewriter.notifyMatchFailure(
702 op, "Unsupported operation: xegpu.prefetch_nd with tensor "
703 "descriptor rank == 1");
704 }
705 }
706 return success();
707 }
708};
709
710// Add a builder that creates
711// offset * elemByteSize + baseAddr
712static Value addOffsetToBaseAddr(ConversionPatternRewriter &rewriter,
713 Location loc, Value baseAddr, Value offset,
714 int64_t elemByteSize) {
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);
719 return newAddr;
720}
721
722// Returns true when every element of `mask` carries the same bit, so gating a
723// whole contiguous block on element 0 is equivalent. Splat constants, an
724// all-ones or all-zeros `vector.constant_mask`, broadcasts of a scalar, and a
725// `vector.from_elements` of one repeated value qualify.
726//
727// TODO: this recognizes a fixed list of producers. A general uniformity query
728// on the vector dialect would cover more forms and would not need updating
729// every time a new producer shows up here.
730static bool isUniformMask(Value mask) {
731 if (!isa<VectorType>(mask.getType()))
732 return true;
733 DenseElementsAttr splat;
734 if (matchPattern(mask, m_Constant(&splat)) && splat.isSplat())
735 return true;
736 if (auto constantMask = mask.getDefiningOp<vector::ConstantMaskOp>()) {
737 // All-ones and all-zeros are uniform. A zero dim size is only legal when
738 // every dim is zero, so those are the only two uniform cases.
739 return constantMask.isAllOnesMask() ||
740 llvm::all_of(constantMask.getMaskDimSizes(),
741 [](int64_t size) { return size == 0; });
742 }
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());
747 // Flattening a rank > 1 operand to rank 1 inserts a shape cast. It keeps the
748 // same elements, so it keeps uniformity.
749 if (auto shapeCast = mask.getDefiningOp<vector::ShapeCastOp>())
750 return isUniformMask(shapeCast.getSource());
751 return false;
752}
753
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;
759 LogicalResult
760 matchAndRewrite(OpType op, typename OpType::Adaptor adaptor,
761 ConversionPatternRewriter &rewriter) const override {
762 Value offset = adaptor.getOffsets();
763 if (!offset)
764 return rewriter.notifyMatchFailure(op, "Expected offset to be provided.");
765 auto loc = op.getLoc();
766 auto ctxt = rewriter.getContext();
767 Value basePtrI64;
768 // Load result or Store valye Type can be vector or scalar.
769 Type valOrResTy;
770 if constexpr (std::is_same_v<OpType, xegpu::LoadGatherOp>)
771 valOrResTy =
772 this->getTypeConverter()->convertType(op.getResult().getType());
773 else
774 valOrResTy = adaptor.getValue().getType();
775 VectorType valOrResVecTy = dyn_cast<VectorType>(valOrResTy);
776 bool hasScalarVal = !valOrResVecTy;
777 int64_t elemBitWidth =
778 hasScalarVal ? valOrResTy.getIntOrFloatBitWidth()
779 : valOrResVecTy.getElementType().getIntOrFloatBitWidth();
780 // Element type must be multiple of 8 bits.
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;
785 // Default memory space is global.
786 LLVM::LLVMPointerType ptrTypeLLVM = LLVM::LLVMPointerType::get(
787 ctxt, getNumericXeVMAddrSpace(xegpu::MemorySpace::Global));
788 // Base pointer can come from source (load) or dest (store).
789 // If they are memrefs, we use their memory space.
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());
795 if (failed(addrSpace))
796 return rewriter.notifyMatchFailure(
797 op, "Unsupported memref memory space attribute.");
798 if (*addrSpace != 0)
799 ptrTypeLLVM = LLVM::LLVMPointerType::get(ctxt, *addrSpace);
800 }
801 } else {
802 basePtrI64 = adaptor.getDest();
803 if (auto memRefTy = dyn_cast<MemRefType>(op.getDest().getType())) {
804 FailureOr<unsigned> addrSpace =
805 getNumericMemorySpace(memRefTy.getMemorySpace());
806 if (failed(addrSpace))
807 return rewriter.notifyMatchFailure(
808 op, "Unsupported memref memory space attribute.");
809 if (*addrSpace != 0)
810 ptrTypeLLVM = LLVM::LLVMPointerType::get(ctxt, *addrSpace);
811 }
812 }
813 // Base pointer is passed as i32 or i64 by adaptor, cast to i64 if needed.
814 if (basePtrI64.getType() != rewriter.getI64Type()) {
815 basePtrI64 = arith::ExtUIOp::create(rewriter, loc, rewriter.getI64Type(),
816 basePtrI64);
817 }
818 Value mask = adaptor.getMask();
819
820 // Coalesce a lane's multi-element access into one block access.
821 //
822 // Distribution gives every element its own offset and mask bit, so a lane
823 // that takes D neighbouring elements arrives with `vector<D>` offsets and
824 // mask. One block access needs one base offset and one mask bit, so take
825 // both from element 0. A non-uniform mask cannot be reduced to one bit, so
826 // it stays a vector and fails to match below.
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});
836 }
837
838 if (dyn_cast<VectorType>(offset.getType())) {
839 // Offset needs be scalar. Single element vector is converted to scalar
840 // by type converter.
841 return rewriter.notifyMatchFailure(op, "Expected offset to be a scalar.");
842 } else {
843 // If offset is provided, we add them to the base pointer.
844 // Offset is in number of elements, we need to multiply by
845 // element byte size.
846 basePtrI64 =
847 addOffsetToBaseAddr(rewriter, loc, basePtrI64, offset, elemByteSize);
848 }
849 // Convert base pointer (i64) to LLVM pointer type.
850 Value basePtrLLVM =
851 LLVM::IntToPtrOp::create(rewriter, loc, ptrTypeLLVM, basePtrI64);
852
853 Value maskForLane;
854 VectorType maskVecTy = dyn_cast<VectorType>(mask.getType());
855 if (maskVecTy) {
856 // Mask needs be scalar. Single element vector is converted to scalar by
857 // type converter.
858 return rewriter.notifyMatchFailure(op, "Expected mask to be a scalar.");
859 } else
860 maskForLane = mask;
861 if constexpr (std::is_same_v<OpType, xegpu::LoadGatherOp>) {
862 scf::IfOp ifOp = scf::IfOp::create(rewriter, loc, {valOrResTy},
863 maskForLane, true, true);
864 // If mask is true,- then clause - load from memory and yield.
865 rewriter.setInsertionPointToStart(&ifOp.getThenRegion().front());
866 if (!hasScalarVal)
867 valOrResTy = VectorType::get({valOrResVecTy.getNumElements()},
868 valOrResVecTy.getElementType());
869 Value loaded =
870 LLVM::LoadOp::create(rewriter, loc, valOrResTy, basePtrLLVM);
871 // Set cache control attribute on the load operation.
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());
878 // If mask is false - else clause -yield a vector of zeros.
879 auto eTy = hasScalarVal ? valOrResTy : valOrResVecTy.getElementType();
880 TypedAttr eVal;
881 if (eTy.isFloat())
882 eVal = FloatAttr::get(eTy, 0.0);
883 else
884 eVal = IntegerAttr::get(eTy, 0);
885 if (hasScalarVal)
886 loaded = arith::ConstantOp::create(rewriter, loc, eVal);
887 else
888 loaded = arith::ConstantOp::create(
889 rewriter, loc, DenseElementsAttr::get(valOrResVecTy, eVal));
890 scf::YieldOp::create(rewriter, loc, ValueRange{loaded});
891 rewriter.replaceOp(op, ifOp.getResult(0));
892 } else {
893 // If mask is true, perform the store.
894 scf::IfOp ifOp = scf::IfOp::create(rewriter, loc, maskForLane, false);
895 auto body = ifOp.getBody();
896 rewriter.setInsertionPointToStart(body);
897 auto storeOp =
898 LLVM::StoreOp::create(rewriter, loc, adaptor.getValue(), basePtrLLVM);
899 // Set cache control attribute on the store operation.
900 storeOp.getOperation()->setDiscardableAttr(
901 "cache_control", xevm::StoreCacheControlAttr::get(
902 ctxt, translateStoreXeGPUCacheHint(
903 op.getL1Hint(), op.getL3Hint())));
904 rewriter.eraseOp(op);
905 }
906 return success();
907 }
908};
909
910class CreateMemDescOpPattern final
911 : public OpConversionPattern<xegpu::CreateMemDescOp> {
912public:
913 using OpConversionPattern<xegpu::CreateMemDescOp>::OpConversionPattern;
914 LogicalResult
915 matchAndRewrite(xegpu::CreateMemDescOp op, OpAdaptor adaptor,
916 ConversionPatternRewriter &rewriter) const override {
917
918 rewriter.replaceOp(op, adaptor.getSource());
919 return success();
920 }
921};
922
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;
928 LogicalResult
929 matchAndRewrite(OpType op, typename OpType::Adaptor adaptor,
930 ConversionPatternRewriter &rewriter) const override {
931
932 SmallVector<OpFoldResult> offsets = op.getMixedOffsets();
933 if (offsets.empty())
934 return rewriter.notifyMatchFailure(op, "Expected offset to be provided.");
935
936 auto loc = op.getLoc();
937 auto ctxt = rewriter.getContext();
938 Value baseAddr32 = adaptor.getMemDesc();
939 Value mdescVal = op.getMemDesc();
940 // Load result or Store value Type can be vector or scalar.
941 Type dataTy;
942 if constexpr (std::is_same_v<OpType, xegpu::LoadMatrixOp>) {
943 Type resType = op.getResult().getType();
944 // Some transforms may leave unit dimension in the 2D vector, adaptors do
945 // not catch it for results.
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());
952 }
953 dataTy = resType;
954 } else
955 dataTy = adaptor.getData().getType();
956 VectorType valOrResVecTy = dyn_cast<VectorType>(dataTy);
957 if (!valOrResVecTy)
958 valOrResVecTy = VectorType::get(1, dataTy);
959
960 int64_t elemBitWidth =
961 valOrResVecTy.getElementType().getIntOrFloatBitWidth();
962 // Element type must be multiple of 8 bits.
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;
967
968 // Default memory space is SLM.
969 LLVM::LLVMPointerType ptrTypeLLVM = LLVM::LLVMPointerType::get(
970 ctxt, getNumericXeVMAddrSpace(xegpu::MemorySpace::SLM));
971
972 auto mdescTy = cast<xegpu::MemDescType>(mdescVal.getType());
973
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);
979
980 // convert base pointer (i32) to LLVM pointer type
981 Value basePtrLLVM =
982 LLVM::IntToPtrOp::create(rewriter, loc, ptrTypeLLVM, basePtrI32);
983
984 if (op.getSubgroupBlockIoAttr()) {
985 // if the attribute 'subgroup_block_io' is set to true, it lowers to
986 // xevm.blockload
987
988 Type intElemTy = rewriter.getIntegerType(elemBitWidth);
989 VectorType intVecTy =
990 VectorType::get(valOrResVecTy.getShape(), intElemTy);
991
992 if constexpr (std::is_same_v<OpType, xegpu::LoadMatrixOp>) {
993 Value loadOp = xevm::BlockLoadOp::create(
994 rewriter, loc, intVecTy, basePtrLLVM, /*cache_control=*/nullptr);
995 if (intVecTy != valOrResVecTy) {
996 loadOp =
997 vector::BitCastOp::create(rewriter, loc, valOrResVecTy, loadOp);
998 }
999 rewriter.replaceOp(op, loadOp);
1000 } else {
1001 Value dataToStore = adaptor.getData();
1002 if (valOrResVecTy != intVecTy) {
1003 dataToStore =
1004 vector::BitCastOp::create(rewriter, loc, intVecTy, dataToStore);
1005 }
1006 xevm::BlockStoreOp::create(rewriter, loc, basePtrLLVM, dataToStore,
1007 nullptr);
1008 rewriter.eraseOp(op);
1009 }
1010 return success();
1011 }
1012
1013 if (valOrResVecTy.getNumElements() >= 1) {
1014 auto chipOpt = xegpu::getChipStr(op);
1015 if (!chipOpt ||
1016 (*chipOpt != "pvc" && *chipOpt != "bmg" && *chipOpt != "cri")) {
1017 // the lowering for chunk load only works for pvc, bmg or cri
1018 return rewriter.notifyMatchFailure(
1019 op, "The lowering is specific to pvc, bmg or cri.");
1020 }
1021 }
1022
1023 if constexpr (std::is_same_v<OpType, xegpu::LoadMatrixOp>) {
1024 // The load result type is taken from the type converter. This maps
1025 // element types that are not directly representable in LLVM (e.g.
1026 // f8E8M0FNU) to an integer storage type of the same bit width, and
1027 // collapses single-element vectors to a scalar, since LLVM load/store
1028 // does not support vectors of size 1.
1029 Type loadTy =
1030 this->getTypeConverter()->convertType(op.getResult().getType());
1031 auto loadOp = LLVM::LoadOp::create(rewriter, loc, loadTy, basePtrLLVM);
1032 rewriter.replaceOp(op, loadOp);
1033 } else {
1034 LLVM::StoreOp::create(rewriter, loc, adaptor.getData(), basePtrLLVM);
1035 rewriter.eraseOp(op);
1036 }
1037 return success();
1038 }
1039};
1040
1041class PrefetchToXeVMPattern : public OpConversionPattern<xegpu::PrefetchOp> {
1042 using OpConversionPattern::OpConversionPattern;
1043 LogicalResult
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();
1049 // Base pointer is passed as i32 or i64 by adaptor, cast to i64 if needed.
1050 if (basePtrI64.getType() != rewriter.getI64Type())
1051 basePtrI64 = arith::ExtUIOp::create(rewriter, loc, rewriter.getI64Type(),
1052 basePtrI64);
1053 Value offsets = adaptor.getOffsets();
1054 if (offsets) {
1055 VectorType offsetsVecTy = dyn_cast<VectorType>(offsets.getType());
1056 if (offsetsVecTy) {
1057 // Offset needs be scalar.
1058 return rewriter.notifyMatchFailure(op,
1059 "Expected offsets to be a scalar.");
1060 } else {
1061 int64_t elemBitWidth{0};
1062 int64_t elemByteSize;
1063 // Element byte size can come from two sources:
1064 if (auto memRefTy = dyn_cast<MemRefType>(op.getSourceType())) {
1065 // If memref is available, we use its element type to
1066 // determine element byte size.
1067 elemBitWidth = memRefTy.getElementType().getIntOrFloatBitWidth();
1068 } else {
1069 // Otherwise, we use the provided offset byte alignment.
1070 elemByteSize = *op.getOffsetAlignByte();
1071 }
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;
1077 }
1078 basePtrI64 = addOffsetToBaseAddr(rewriter, loc, basePtrI64, offsets,
1079 elemByteSize);
1080 }
1081 }
1082 // Default memory space is global.
1083 LLVM::LLVMPointerType ptrTypeLLVM = LLVM::LLVMPointerType::get(
1084 ctxt, getNumericXeVMAddrSpace(xegpu::MemorySpace::Global));
1085 // If source is a memref, we use its memory space.
1086 if (auto memRefTy = dyn_cast<MemRefType>(op.getSource().getType())) {
1087 FailureOr<unsigned> addrSpace =
1088 getNumericMemorySpace(memRefTy.getMemorySpace());
1089 if (failed(addrSpace))
1090 return rewriter.notifyMatchFailure(
1091 op, "Unsupported memref memory space attribute.");
1092 if (*addrSpace != 0)
1093 ptrTypeLLVM = LLVM::LLVMPointerType::get(ctxt, *addrSpace);
1094 }
1095 // Convert base pointer (i64) to LLVM pointer type.
1096 Value ptrLLVM =
1097 LLVM::IntToPtrOp::create(rewriter, loc, ptrTypeLLVM, basePtrI64);
1098 // Create the prefetch op with cache control attribute.
1099 xevm::PrefetchOp::create(
1100 rewriter, loc, ptrLLVM,
1101 xevm::LoadCacheControlAttr::get(
1102 ctxt, translateLoadXeGPUCacheHint(op.getL1Hint(), op.getL3Hint())));
1103 rewriter.eraseOp(op);
1104 return success();
1105 }
1106};
1107
1108class FenceToXeVMPattern : public OpConversionPattern<xegpu::FenceOp> {
1109 using OpConversionPattern::OpConversionPattern;
1110 LogicalResult
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;
1118 break;
1119 case xegpu::FenceScope::GPU:
1120 memScope = xevm::MemScope::DEVICE;
1121 break;
1122 }
1123 xevm::AddrSpace addrSpace{xevm::AddrSpace::GLOBAL};
1124 switch (op.getMemoryKind()) {
1125 case xegpu::MemorySpace::Global:
1126 addrSpace = xevm::AddrSpace::GLOBAL;
1127 break;
1128 case xegpu::MemorySpace::SLM:
1129 addrSpace = xevm::AddrSpace::SHARED;
1130 break;
1131 }
1132 xevm::MemfenceOp::create(rewriter, loc, memScope, addrSpace);
1133 rewriter.eraseOp(op);
1134 return success();
1135 }
1136};
1137
1138static auto encodePrecision = [](Type type) -> xevm::ElemType {
1139 if (type.isBF16())
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");
1160};
1161
1162static unsigned getNumOperandsPerDword(xevm::ElemType pTy) {
1163 switch (pTy) {
1164 case xevm::ElemType::TF32:
1165 return 1;
1166 case xevm::ElemType::BF16:
1167 case xevm::ElemType::F16:
1168 return 2;
1169 case xevm::ElemType::U8:
1170 case xevm::ElemType::S8:
1171 case xevm::ElemType::F8:
1172 case xevm::ElemType::BF8:
1173 return 4;
1174 case xevm::ElemType::E2M1:
1175 return 8;
1176 default:
1177 llvm_unreachable("unsupported xevm::ElemType");
1178 }
1179}
1180
1181class DpasToXeVMPattern : public OpConversionPattern<xegpu::DpasOp> {
1182 using OpConversionPattern::OpConversionPattern;
1183 LogicalResult
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());
1191
1192 // get the correct dpasInst by getting info from chip
1193 auto chipStr = xegpu::getChipStr(op);
1194 if (!chipStr)
1195 return rewriter.notifyMatchFailure(op, "cannot determine target chip");
1196
1197 const auto *uArch = mlir::xegpu::uArch::getUArch(*chipStr);
1198 if (!uArch)
1199 return rewriter.notifyMatchFailure(op, "unsupported target uArch");
1200
1201 auto *dpasInst = const_cast<xegpu::uArch::SubgroupMatrixMultiplyAcc *>(
1202 llvm::dyn_cast_or_null<xegpu::uArch::SubgroupMatrixMultiplyAcc>(
1203 uArch->getInstruction(
1204 xegpu::uArch::InstructionKind::SubgroupMatrixMultiplyAcc)));
1205 if (!dpasInst)
1206 return rewriter.notifyMatchFailure(op,
1207 "DPAS not supported by target uArch");
1208
1209 auto checkSupportedTypes = [&](VectorType vecTy,
1210 xegpu::uArch::MMAOpndKind kind) -> bool {
1211 auto supported = dpasInst->getSupportedTypes(*ctxt, kind);
1212 return llvm::find(supported, vecTy.getElementType()) != supported.end();
1213 };
1214
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");
1221 // NOTE: Supported types for MatrixC and MatrixD are identical
1222 if (!checkSupportedTypes(resultType, xegpu::uArch::MMAOpndKind::MatrixD))
1223 return rewriter.notifyMatchFailure(
1224 op, "result/accumulator element type not supported by target uArch");
1225
1226 xevm::ElemType precATy = encodePrecision(aTy.getElementType());
1227 xevm::ElemType precBTy = encodePrecision(bTy.getElementType());
1228 Value c = op.getAcc();
1229 if (!c) {
1230 auto elementTy = resultType.getElementType();
1231 Attribute initValueAttr;
1232 if (isa<FloatType>(elementTy))
1233 initValueAttr = FloatAttr::get(elementTy, 0.0);
1234 else
1235 initValueAttr = IntegerAttr::get(elementTy, 0);
1236 c = arith::ConstantOp::create(
1237 rewriter, loc, DenseElementsAttr::get(resultType, initValueAttr));
1238 }
1239
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());
1245 VectorType cNty =
1246 VectorType::get(cvecty.getNumElements(), cvecty.getElementType());
1247 if (cvecty != cNty)
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,
1252 systolicDepth *
1253 getNumOperandsPerDword(precATy)),
1254 xevm::MMATypesAttr::get(ctxt, precDTy, precATy, precBTy, precCTy));
1255 if (cvecty != cNty)
1256 dpasRes = vector::ShapeCastOp::create(rewriter, loc, resultType, dpasRes);
1257 rewriter.replaceOp(op, dpasRes);
1258 return success();
1259 }
1260};
1261
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;
1287 default:
1288 return std::nullopt;
1289 }
1290}
1291
1292class AtomicRMWToXeVMPattern : public OpConversionPattern<xegpu::AtomicRMWOp> {
1293 using OpConversionPattern::OpConversionPattern;
1294 LogicalResult
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());
1304 Value basePtrLLVM =
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));
1318 Value currPtr =
1319 LLVM::GEPOp::create(rewriter, loc, ptrTypeLLVM,
1320 srcOrDstVecTy.getElementType(), basePtrLLVM, idx);
1321 Value newVal =
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);
1325 }
1326 rewriter.replaceOp(op, resVec);
1327 return success();
1328 }
1329};
1330
1331class DpasMxToXeVMPattern : public OpConversionPattern<xegpu::DpasMxOp> {
1332 using OpConversionPattern::OpConversionPattern;
1333 LogicalResult
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();
1340 auto resVecTy =
1341 cast<VectorType>(getTypeConverter()->convertType(op.getType()));
1342
1343 auto chipStr = xegpu::getChipStr(op);
1344 if (!chipStr)
1345 return rewriter.notifyMatchFailure(op, "cannot determine target chip");
1346
1347 const auto *uArch = xegpu::uArch::getUArch(*chipStr);
1348 if (!uArch)
1349 return rewriter.notifyMatchFailure(op, "unsupported target uArch");
1350
1351 // TODO: Add supported shape check
1352
1353 xevm::ElemType precATy = encodePrecision(aTy.getElementType());
1354 xevm::ElemType precBTy = encodePrecision(bTy.getElementType());
1355 Value c = adaptor.getAcc();
1356 if (!c) {
1357 auto elementTy = resVecTy.getElementType();
1358 Attribute initValueAttr;
1359 if (isa<FloatType>(elementTy))
1360 initValueAttr = FloatAttr::get(elementTy, 0.0);
1361 else
1362 initValueAttr = IntegerAttr::get(elementTy, 0);
1363 c = arith::ConstantOp::create(
1364 rewriter, loc, DenseElementsAttr::get(resVecTy, initValueAttr));
1365 }
1366
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(
1373 rewriter, loc,
1374 VectorType::get(aVecTy.getNumElements() / 2, rewriter.getI8Type()),
1375 aVec);
1376 if (bVecTy.getElementTypeBitWidth() == 4)
1377 bVec = vector::BitCastOp::create(
1378 rewriter, loc,
1379 VectorType::get(bVecTy.getNumElements() / 2, rewriter.getI8Type()),
1380 bVec);
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,
1389 systolicDepth *
1390 getNumOperandsPerDword(precATy)),
1391 xevm::MMATypesAttr::get(ctxt, precDTy, precATy, precBTy, precCTy));
1392 rewriter.replaceOp(op, dpasMxRes);
1393 return success();
1394 }
1395};
1396
1397//===----------------------------------------------------------------------===//
1398// arith.extf / arith.truncf to xevm.extf / xevm.truncf
1399//===----------------------------------------------------------------------===//
1400//
1401// Micro-scaling (MX) GEMM lowering breaks arith.scaling_extf/scaling_truncf
1402// into plain arith.extf/arith.truncf whose narrow side uses one of the MX float
1403// formats (f8E5M2, f8E4M3FN or f4E2M1FN). These narrow floats have no native
1404// LLVM support, so the conversions are mapped onto the dedicated xevm.extf /
1405// xevm.truncf ops which lower to hardware builtins. The f8E8M0FNU scale type is
1406// intentionally not handled here: it is expanded into integer arithmetic by
1407// arith-expand before this pass runs.
1408
1409// xevm.extf / xevm.truncf only convert between the MX narrow floats and
1410// f16/bf16, and the underlying builtins operate on exactly 16 f16/bf16 values.
1411static constexpr int64_t kXeVMExtfTruncfNumElems = 16;
1412
1413// Maps a narrow MX float element type to the matching xevm.extf source enum.
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;
1422}
1423
1424// Maps a narrow MX float element type to the matching xevm.truncf dest enum.
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;
1433}
1434
1435// Returns true if `op` is an arith.extf that can be lowered to xevm.extf, i.e.
1436// a rank-1 widening from an MX narrow float to a 16-element f16/bf16 vector.
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)
1441 return false;
1442 if (dstTy.getNumElements() != kXeVMExtfTruncfNumElems)
1443 return false;
1444 Type dstETy = dstTy.getElementType();
1445 if (!dstETy.isF16() && !dstETy.isBF16())
1446 return false;
1447 return getExtfNarrowType(srcTy.getElementType()).has_value();
1448}
1449
1450// Returns true if `op` is an arith.truncf that can be lowered to xevm.truncf,
1451// i.e. a rank-1 truncation from an f16/bf16 vector to an MX narrow float. The
1452// source has to hold a whole number of the fixed-size groups xevm.truncf
1453// converts at a time; wider vectors are converted in several steps.
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)
1458 return false;
1459 int64_t numElems = srcTy.getNumElements();
1460 if (numElems == 0 || numElems % kXeVMExtfTruncfNumElems != 0)
1461 return false;
1462 Type srcETy = srcTy.getElementType();
1463 if (!srcETy.isF16() && !srcETy.isBF16())
1464 return false;
1465 return getTruncfNarrowType(dstTy.getElementType()).has_value();
1466}
1467
1468class ExtfToXeVMPattern : public OpConversionPattern<arith::ExtFOp> {
1469 using OpConversionPattern::OpConversionPattern;
1470 LogicalResult
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;
1484 // The narrow float operand has already been type-converted to an integer
1485 // vector of the same bit width (i4 for fp4, i8 for fp8). xevm.extf takes
1486 // the values packed into an i8 vector, so re-pack fp4 (i4) operands.
1487 Value src = adaptor.getIn();
1488 auto convSrcTy = cast<VectorType>(src.getType());
1489 if (convSrcTy.getElementTypeBitWidth() == 4)
1490 src = vector::BitCastOp::create(
1491 rewriter, loc,
1492 VectorType::get(convSrcTy.getNumElements() / 2, rewriter.getI8Type()),
1493 src);
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);
1499 return success();
1500 }
1501};
1502
1503class TruncfToXeVMPattern : public OpConversionPattern<arith::TruncFOp> {
1504 using OpConversionPattern::OpConversionPattern;
1505 LogicalResult
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);
1522
1523 // xevm.truncf lowers to instructions that convert a fixed number of
1524 // elements at a time, so a wider source is converted one group at a time
1525 // and the packed results are concatenated. Each group produces the narrow
1526 // floats packed into an i8 vector.
1527 int64_t numGroups = srcVecTy.getNumElements() / kXeVMExtfTruncfNumElems;
1528 int64_t groupBytes =
1529 kXeVMExtfTruncfNumElems * dstVecTy.getElementTypeBitWidth() / 8;
1530 Type groupTy = VectorType::get(groupBytes, rewriter.getI8Type());
1531
1532 Value src = adaptor.getIn();
1533 Value packed;
1534 if (numGroups == 1) {
1535 packed = xevm::TruncfOp::create(rewriter, loc, groupTy, src, srcEnumAttr,
1536 dstEnumAttr);
1537 } else {
1538 auto packedTy =
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, /*strides=*/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,
1550 /*strides=*/1);
1551 }
1552 }
1553 // Re-shape to the type-converted result type (i4 vector for fp4).
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);
1558 return success();
1559 }
1560};
1561
1562// Lowers `xegpu.lane_shuffle` to `xevm.bitcast_shuffle`.
1563//
1564// `xevm.bitcast_shuffle` concatenates the components of its operand across the
1565// subgroup, the first component of every lane first, and then hands chunks the
1566// size of a result component back out to the lanes in order. Numbering the
1567// elements of a `vector<NxT>` fragment held by lane `i` of a subgroup of size
1568// `S` by their logical position, that concatenation is exactly the `pack` mode
1569// input numbering `j * S + i`. Taking the result as a single `N * width(T)` bit
1570// scalar then hands lane `i` the logical positions `i * N .. i * N + N - 1`,
1571// which is the `pack` mode output numbering.
1572//
1573// So `pack` is a vector-to-scalar `xevm.bitcast_shuffle` followed by a bitcast
1574// back to the fragment type, and `unpack`, being its inverse, is a bitcast to
1575// the scalar followed by a scalar-to-vector `xevm.bitcast_shuffle`.
1576//
1577// `xevm.bitcast_shuffle` only accepts the integer types `i8`, `i16`, `i32` and
1578// `i64`, since it is bit-preserving and so does not depend on how the bits are
1579// interpreted. A fragment of a floating point type is therefore bitcast to a
1580// same-width integer vector on the way in and back on the way out.
1581class LaneShuffleToXeVMPattern
1582 : public OpConversionPattern<xegpu::LaneShuffleOp> {
1583 using OpConversionPattern::OpConversionPattern;
1584 LogicalResult
1585 matchAndRewrite(xegpu::LaneShuffleOp op, OpAdaptor adaptor,
1586 ConversionPatternRewriter &rewriter) const override {
1587 auto vecTy = dyn_cast<VectorType>(adaptor.getSource().getType());
1588 if (!vecTy)
1589 return rewriter.notifyMatchFailure(op, "Expected a vector fragment.");
1590 // The shuffle redistributes whole bytes between the lanes, so sub-byte
1591 // element types, fp4 in particular, cannot be shuffled. Widths without a
1592 // matching integer type the op accepts are rejected for the same reason.
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.");
1601
1602 Location loc = op.getLoc();
1603 Type packedTy = rewriter.getIntegerType(fragmentBits);
1604 // The integer vector type the shuffle actually operates on. Equal to the
1605 // fragment type when that is already an integer vector.
1606 VectorType shuffleTy =
1607 VectorType::get(vecTy.getShape(), rewriter.getIntegerType(elemBits));
1608
1609 Value res;
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);
1616 } else {
1617 Value packed =
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);
1622 }
1623 rewriter.replaceOp(op, res);
1624 return success();
1625 }
1626};
1627
1628//===----------------------------------------------------------------------===//
1629// Pass Definition
1630//===----------------------------------------------------------------------===//
1631
1632struct ConvertXeGPUToXeVMPass
1633 : public impl::ConvertXeGPUToXeVMPassBase<ConvertXeGPUToXeVMPass> {
1634 using Base::Base;
1635
1636 void runOnOperation() override {
1637 MLIRContext *context = &getContext();
1638
1639 // XeVM type converter is based on LLVM type converter with the
1640 // following customizations.
1641 // First, type conversion rules are added for xegpu custom types,
1642 // TensorDescType and MemDescType.
1643 // Second, MemRefType is lowered to single integer type
1644 // Third, VectorType of single element or 0D is converted to vector
1645 // element type. Otherwise, vector type is flatten to 1D.
1646 LowerToLLVMOptions options(context);
1647 options.overrideIndexBitwidth(this->use64bitIndex ? 64 : 32);
1648 LLVMTypeConverter typeConverter(context, options);
1649
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());
1654 // If the vector rank is 0 or has a single element, return the element
1655 unsigned rank = type.getRank();
1656 if (rank == 0 || type.getNumElements() == 1)
1657 return elemType;
1658 // Otherwise, convert the vector to a flat vector type.
1659 int64_t sum = llvm::product_of(type.getShape());
1660 return VectorType::get(sum, elemType);
1661 });
1662 typeConverter.addConversion([&](xegpu::TensorDescType type) -> Type {
1663 if (type.getRank() == 1)
1664 return xevmIndexType;
1665 return VectorType::get(8, i32Type);
1666 });
1667 // SLM access related type conversions.
1668 // TODO: LLVM DLTI provides clean way of representing different pointer size
1669 // based on address space. Currently pointer size of SLM access is hard
1670 // coded to 32bit. Update to use DLTI when switching overall XeGPU lowering
1671 // to use DLTI instead of use64bitIndex option used above.
1672
1673 // Convert MemDescType into i32 for SLM
1674 typeConverter.addConversion(
1675 [&](xegpu::MemDescType type) -> Type { return i32Type; });
1676
1677 typeConverter.addConversion([&](MemRefType type) -> Type {
1678 return isSharedMemRef(type) ? i32Type : xevmIndexType;
1679 });
1680
1681 // LLVM type converter puts unrealized casts for the following cases:
1682 // add materialization casts to handle them.
1683
1684 // Materialization to convert memref to i64 or i32 depending on global/SLM
1685 // Applies only to target materialization.
1686 // Note: int type to memref materialization is not required as xegpu ops
1687 // currently do not produce memrefs as result.
1688 auto memrefToIntMaterializationCast = [](OpBuilder &builder, Type type,
1689 ValueRange inputs,
1690 Location loc) -> Value {
1691 if (inputs.size() != 1)
1692 return {};
1693 auto input = inputs.front();
1694 if (auto memrefTy = dyn_cast<MemRefType>(input.getType())) {
1695 unsigned rank = memrefTy.getRank();
1696 Type indexType = builder.getIndexType();
1697
1698 int64_t intOffsets;
1699 SmallVector<int64_t> intStrides;
1700 Value addr;
1701 Value offset;
1702 if (succeeded(memrefTy.getStridesAndOffset(intStrides, intOffsets)) &&
1703 ShapedType::isStatic(intOffsets)) {
1704 addr = memref::ExtractAlignedPointerAsIndexOp::create(builder, loc,
1705 input);
1706 offset = arith::ConstantOp::create(builder, loc,
1707 builder.getIndexAttr(intOffsets));
1708 } else {
1709
1710 // Result types: [base_memref, offset, stride0, stride1, ...,
1711 // strideN-1, size0, size1, ..., sizeN-1]
1712 SmallVector<Type> resultTypes{
1713 MemRefType::get({}, memrefTy.getElementType(),
1714 MemRefLayoutAttrInterface(),
1715 memrefTy.getMemorySpace()),
1716 indexType};
1717 // strides + sizes
1718 resultTypes.append(2 * rank, indexType);
1719
1720 auto meta = memref::ExtractStridedMetadataOp::create(
1721 builder, loc, resultTypes, input);
1722
1723 addr = memref::ExtractAlignedPointerAsIndexOp::create(
1724 builder, loc, meta.getBaseBuffer());
1725 offset = meta.getOffset();
1726 }
1727
1728 auto addrCasted =
1729 arith::IndexCastUIOp::create(builder, loc, type, addr);
1730 auto offsetCasted =
1731 arith::IndexCastUIOp::create(builder, loc, type, offset);
1732
1733 // Compute the final address: base address + byte offset
1734 auto byteSize = arith::ConstantOp::create(
1735 builder, loc, type,
1736 builder.getIntegerAttr(type,
1737 memrefTy.getElementTypeBitWidth() / 8));
1738 auto byteOffset =
1739 arith::MulIOp::create(builder, loc, offsetCasted, byteSize);
1740 auto addrWithOffset =
1741 arith::AddIOp::create(builder, loc, addrCasted, byteOffset);
1742
1743 return addrWithOffset.getResult();
1744 }
1745 return {};
1746 };
1747
1748 // Materialization to convert ui64 to i64
1749 // Applies only to target materialization.
1750 // Note: i64 to ui64 materialization is not required as xegpu ops
1751 // currently do not produce ui64 as result.
1752 auto ui64ToI64MaterializationCast = [](OpBuilder &builder, Type type,
1753 ValueRange inputs,
1754 Location loc) -> Value {
1755 if (inputs.size() != 1)
1756 return {};
1757 auto input = inputs.front();
1758 if (input.getType() == builder.getIntegerType(64, false)) {
1759 Value cast =
1760 index::CastUOp::create(builder, loc, builder.getIndexType(), input)
1761 .getResult();
1762 return arith::IndexCastUIOp::create(builder, loc, type, cast)
1763 .getResult();
1764 }
1765 return {};
1766 };
1767
1768 // Materialization to convert ui32 to i32
1769 // Applies only to target materialization.
1770 // Note: i32 to ui32 materialization is not required as xegpu ops
1771 // currently do not produce ui32 as result.
1772 auto ui32ToI32MaterializationCast = [](OpBuilder &builder, Type type,
1773 ValueRange inputs,
1774 Location loc) -> Value {
1775 if (inputs.size() != 1)
1776 return {};
1777 auto input = inputs.front();
1778 if (input.getType() == builder.getIntegerType(32, false)) {
1779 Value cast =
1780 index::CastUOp::create(builder, loc, builder.getIndexType(), input)
1781 .getResult();
1782 return arith::IndexCastUIOp::create(builder, loc, type, cast)
1783 .getResult();
1784 }
1785 return {};
1786 };
1787
1788 // Materialization to convert between vector types
1789 // - Add shape cast for different shapes
1790 // - Add bitcast for different element types
1791 // Applies to both source and target materialization.
1792 auto vectorToVectorMaterializationCast = [](OpBuilder &builder, Type type,
1793 ValueRange inputs,
1794 Location loc) -> Value {
1795 if (inputs.size() != 1)
1796 return {};
1797 auto input = inputs.front();
1798 if (auto vecTy = dyn_cast<VectorType>(input.getType())) {
1799 if (auto targetVecTy = dyn_cast<VectorType>(type)) {
1800 Value cast = input;
1801 // If the target type has a different shape, add a shape cast
1802 // If the target type has a different element type, add a bitcast
1803 if (targetVecTy.getShape() != vecTy.getShape()) {
1804 cast = vector::ShapeCastOp::create(
1805 builder, loc,
1806 VectorType::get(targetVecTy.getShape(),
1807 vecTy.getElementType()),
1808 cast)
1809 .getResult();
1810 }
1811 if (targetVecTy.getElementType() != vecTy.getElementType()) {
1812 cast = vector::BitCastOp::create(builder, loc, targetVecTy, cast)
1813 .getResult();
1814 }
1815 return cast;
1816 }
1817 }
1818 return {};
1819 };
1820
1821 // Materialization to convert
1822 // - single element vector to single element of vector element type
1823 // Applies only to target materialization.
1824 auto vectorToSingleElementMaterializationCast =
1825 [](OpBuilder &builder, Type type, ValueRange inputs,
1826 Location loc) -> Value {
1827 if (inputs.size() != 1)
1828 return {};
1829 auto input = inputs.front();
1830 if (auto vecTy = dyn_cast<VectorType>(input.getType())) {
1831 // Source needs to be single element vector
1832 auto rank = vecTy.getRank();
1833 if (rank != 0 && vecTy.getNumElements() != 1)
1834 return {};
1835 auto inElemTy = vecTy.getElementType();
1836 // extract scalar
1837 Value cast = input;
1838 if (rank == 0) {
1839 cast = vector::ExtractOp::create(builder, loc, cast, {}).getResult();
1840 } else {
1841 cast = vector::ExtractOp::create(builder, loc, cast,
1842 SmallVector<int64_t>(rank, 0))
1843 .getResult();
1844 }
1845 // Extracted element type may need conversion
1846 // Two cases
1847 // 1. Index type to integer type
1848 // 2. Other element type mismatch
1849 if (inElemTy.isIndex()) {
1850 cast = arith::IndexCastUIOp::create(builder, loc, type, cast)
1851 .getResult();
1852 } else if (inElemTy != type) {
1853 cast = arith::BitcastOp::create(builder, loc, type, cast).getResult();
1854 }
1855 return cast;
1856 }
1857 return {};
1858 };
1859
1860 // Materialization to convert
1861 // - single element of vector element type to single element vector
1862 // If result type of original op is single element vector and lowered type
1863 // is scalar. This materialization cast creates a single element vector by
1864 // First convert element type if needed and then broadcast to single
1865 // element vector.
1866 // Applies only to source materialization.
1867 auto singleElementToVectorMaterializationCast =
1868 [](OpBuilder &builder, Type type, ValueRange inputs,
1869 Location loc) -> Value {
1870 if (inputs.size() != 1)
1871 return {};
1872 auto input = inputs.front();
1873 auto inTy = input.getType();
1874 if (!inTy.isIntOrFloat())
1875 return {};
1876 // If the target type is a vector of rank 0 or single element vector
1877 // of element type matching input type, broadcast input to target type.
1878 if (auto vecTy = dyn_cast<VectorType>(type)) {
1879 if (vecTy.getRank() != 0 && vecTy.getNumElements() != 1)
1880 return {};
1881 auto outElemTy = vecTy.getElementType();
1882 Value cast = input;
1883 if (outElemTy.isIndex()) {
1884 cast = arith::IndexCastUIOp::create(builder, loc,
1885 builder.getIndexType(), cast)
1886 .getResult();
1887 } else if (inTy != outElemTy) {
1888 cast = arith::BitcastOp::create(builder, loc, outElemTy, cast)
1889 .getResult();
1890 }
1891 return vector::BroadcastOp::create(builder, loc, vecTy, cast)
1892 .getResult();
1893 }
1894 return {};
1895 };
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>();
1911 // arith.extf/arith.truncf between MX narrow floats and f16/bf16 are routed
1912 // to xevm.extf/xevm.truncf; all other arith float casts stay legal.
1913 target.addDynamicallyLegalOp<arith::ExtFOp>(
1914 [](arith::ExtFOp op) { return !isXeVMExtf(op); });
1915 target.addDynamicallyLegalOp<arith::TruncFOp>(
1916 [](arith::TruncFOp op) { return !isXeVMTruncf(op); });
1917
1918 RewritePatternSet patterns(context);
1919 populateXeGPUToXeVMConversionPatterns(typeConverter, patterns);
1921 patterns, target);
1922 if (failed(applyPartialConversion(getOperation(), target,
1923 std::move(patterns))))
1924 signalPassFailure();
1925 }
1926};
1927} // namespace
1928
1929//===----------------------------------------------------------------------===//
1930// Pattern Population
1931//===----------------------------------------------------------------------===//
1933 const LLVMTypeConverter &typeConverter, RewritePatternSet &patterns) {
1934 patterns.add<CreateNdDescToXeVMPattern,
1935 LoadStorePrefetchNdToXeVMPattern<xegpu::LoadNdOp>,
1936 LoadStorePrefetchNdToXeVMPattern<xegpu::StoreNdOp>,
1937 LoadStorePrefetchNdToXeVMPattern<xegpu::PrefetchNdOp>>(
1938 typeConverter, patterns.getContext());
1939 patterns.add<AtomicRMWToXeVMPattern, PrefetchToXeVMPattern,
1940 LoadStoreToXeVMPattern<xegpu::LoadGatherOp>,
1941 LoadStoreToXeVMPattern<xegpu::StoreScatterOp>>(
1942 typeConverter, patterns.getContext());
1943 patterns.add<LoadStoreMatrixToXeVMPattern<xegpu::LoadMatrixOp>,
1944 LoadStoreMatrixToXeVMPattern<xegpu::StoreMatrixOp>,
1945 CreateMemDescOpPattern>(typeConverter, patterns.getContext());
1946 patterns.add<FenceToXeVMPattern, DpasToXeVMPattern>(typeConverter,
1947 patterns.getContext());
1948 patterns.add<DpasMxToXeVMPattern>(typeConverter, patterns.getContext());
1949 patterns.add<ExtfToXeVMPattern, TruncfToXeVMPattern>(typeConverter,
1950 patterns.getContext());
1951 patterns.add<LaneShuffleToXeVMPattern>(typeConverter, patterns.getContext());
1952}
return success()
b getContext())
auto load
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.
Definition Attributes.h:25
IntegerAttr getIndexAttr(int64_t value)
Definition Builders.cpp:116
IntegerAttr getIntegerAttr(Type type, int64_t value)
Definition Builders.cpp:237
IntegerType getIntegerType(unsigned width)
Definition Builders.cpp:75
IndexType getIndexType()
Definition Builders.cpp:59
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...
Definition Location.h:76
void setDiscardableAttr(StringAttr name, Attribute value)
Set a discardable attribute by name.
Definition Operation.h:512
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...
Definition Types.h:74
bool isIntOrFloat() const
Return true if this is an integer (of any signedness) or a float type.
Definition Types.cpp:118
bool isF16() const
Definition Types.cpp:38
unsigned getIntOrFloatBitWidth() const
Return the bit width of an integer or a float type, assert failure on other types.
Definition Types.cpp:124
bool isBF16() const
Definition Types.cpp:37
This class represents an instance of an SSA value in the MLIR system, representing a computable value...
Definition Value.h:96
Type getType() const
Return the type of this value.
Definition Value.h:105
Operation * getDefiningOp() const
If this value is the result of an operation, return the operation that defines it.
Definition Value.cpp:18
static ConstantIntOp create(OpBuilder &builder, Location location, int64_t value, unsigned width)
Definition ArithOps.cpp:297
detail::InFlightRemark failed(Location loc, RemarkOpts opts)
Report an optimization remark that failed.
Definition Remarks.h:734
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...
const uArch * getUArch(llvm::StringRef archName)
Definition uArchCommon.h:24
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.
Definition Matchers.h:490
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.
Definition Utils.cpp:105
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.
Definition Utils.cpp:122
void populateXeGPUToXeVMConversionPatterns(const LLVMTypeConverter &typeConverter, RewritePatternSet &patterns)
detail::constant_op_matcher m_Constant()
Matches a constant foldable operation.
Definition Matchers.h:369