MLIR 24.0.0git
VectorToXeGPU.cpp
Go to the documentation of this file.
1//===- VectorToXeGPU.cpp - Convert vector to XeGPU dialect ------*- C++ -*-===//
2//
3// Part of the LLVM Project, 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//
9// This file implements lowering of vector operations to XeGPU dialect ops.
10//
11//===----------------------------------------------------------------------===//
12
15
26#include "mlir/IR/Matchers.h"
28#include "mlir/Pass/Pass.h"
30#include "llvm/ADT/TypeSwitch.h"
31
32#include <algorithm>
33#include <optional>
34
35namespace mlir {
36#define GEN_PASS_DEF_CONVERTVECTORTOXEGPU
37#include "mlir/Conversion/Passes.h.inc"
38} // namespace mlir
39
40using namespace mlir;
41
42namespace {
43
44// Return true if value represents a zero constant.
45static bool isZeroConstant(Value val) {
46 auto constant = val.getDefiningOp<arith::ConstantOp>();
47 if (!constant)
48 return false;
49
50 return TypeSwitch<Attribute, bool>(constant.getValue())
51 .Case([](FloatAttr floatAttr) { return floatAttr.getValue().isZero(); })
52 .Case([](IntegerAttr intAttr) { return intAttr.getValue().isZero(); })
53 .Default(false);
54}
55
56// Return true if the transfer padding value is compatible with the implicit
57// padding of an nd block load. LoadNdOp fills out-of-bounds elements with zero,
58// so a zero constant matches its semantics exactly. A poison padding means the
59// out-of-bounds elements are "don't care", so any implicit padding (including
60// zero) is also acceptable.
61static bool isZeroOrPoisonPadding(Value val) {
62 return isZeroConstant(val) || val.getDefiningOp<ub::PoisonOp>();
63}
64
65// Return true if the permutation map keeps every dimension in place except the
66// innermost two, which are swapped, e.g.:
67// (d0, d1) -> (d1, d0)
68// (d0, d1, d2) -> (d2, d1)
69// (d0, d1, d2, d3) -> (d0, d1, d3, d2)
70// This is the only non-identity permutation an nd block load can realize (by
71// loading the untransposed block and applying a trailing vector.transpose).
72static bool isInnermostTwoDimsTransposed(AffineMap map) {
73 unsigned numResults = map.getNumResults();
74 if (numResults < 2)
75 return false;
76 MLIRContext *ctx = map.getContext();
77 unsigned numInputs = map.getNumInputs();
78 // All but the innermost two results must match the minor-identity map.
79 for (unsigned i = 0; i + 2 < numResults; ++i)
80 if (map.getResult(i) != getAffineDimExpr(numInputs - numResults + i, ctx))
81 return false;
82 // The innermost two results must be the last two input dims, swapped.
83 return map.getResult(numResults - 2) ==
84 getAffineDimExpr(numInputs - 1, ctx) &&
85 map.getResult(numResults - 1) == getAffineDimExpr(numInputs - 2, ctx);
86}
87
88// Return true if `uArch` can transfer a `shape`-shaped tile of `elemTy`
89// elements with the subgroup 2D block instruction `instKind`.
90//
91// A 2D block instruction accesses the two innermost dimensions of the tile;
92// leading dimensions are unrolled into a sequence of 2D accesses. The
93// innermost 2D tile does not have to match a hardware block exactly, because
94// the later XeGPU passes (layout propagation and blocking) split it into
95// hardware-sized blocks - but that split only exists when each of the two
96// innermost extents is a multiple of a supported block extent. Transfers
97// without such a split cannot be expressed with block instructions at all and
98// must use a different lowering.
99static bool isSupportedBlockShape(const xegpu::uArch::uArch *uArch,
102 bool hasTranspose = false) {
103 if (shape.size() < 2)
104 return false;
105 if (!uArch || !uArch->isSupportedInstruction(instKind))
106 return false;
107
108 const auto *blockInst = dyn_cast<xegpu::uArch::BlockIOInstructionInterface>(
109 uArch->getInstruction(instKind));
110 if (!blockInst)
111 return false;
112
113 int width = static_cast<int>(shape.back());
114 int height = static_cast<int>(shape[shape.size() - 2]);
115
116 // A tile is accessible when both of its extents are multiples of a supported
117 // block extent, so that the later XeGPU passes can split it into hardware
118 // blocks. A missing entry means this variant does not support the element
119 // type.
120 auto fitsBlockShapes = [&](bool hasTransform) {
121 std::optional<xegpu::uArch::BlockIOInstructionInterface::BlockShapes>
122 blockShapes = blockInst->getBlockWidthHeightCount(elemTy, hasTransform,
123 hasTranspose);
124 if (!blockShapes)
125 return false;
126 auto [widths, heights, counts] = *blockShapes;
127 return xegpu::getLargestDivisor(width, widths) != -1 &&
128 xegpu::getLargestDivisor(height, heights) != -1;
129 };
130
131 // Whether a tile is eventually loaded with the transformed (VNNI) variant is
132 // only decided later, when layout propagation knows whether it feeds the B
133 // operand of a matrix operation, so accept a tile that either variant can
134 // access.
135 return fitsBlockShapes(/*hasTransform=*/false) ||
136 fitsBlockShapes(/*hasTransform=*/true);
137}
138
139static LogicalResult transferPreconditions(PatternRewriter &rewriter,
140 VectorTransferOpInterface xferOp) {
141 if (xferOp.getMask())
142 return rewriter.notifyMatchFailure(xferOp,
143 "Masked transfer is not supported");
144
145 auto srcTy = dyn_cast<MemRefType>(xferOp.getShapedType());
146 if (!srcTy)
147 return rewriter.notifyMatchFailure(xferOp, "Expects memref source");
148
149 // Validate further transfer op semantics.
150 SmallVector<int64_t> strides;
151 int64_t offset;
152 if (failed(srcTy.getStridesAndOffset(strides, offset)))
153 return rewriter.notifyMatchFailure(xferOp,
154 "The memref strides cannot be inferred");
155 if (strides.empty())
156 return rewriter.notifyMatchFailure(xferOp, "0D memref is not supported");
157 if (strides.back() != 1)
158 return rewriter.notifyMatchFailure(
159 xferOp, "Buffer must be contiguous in the innermost dimension");
160
161 VectorType vecTy = xferOp.getVectorType();
162 unsigned vecRank = vecTy.getRank();
163 if (vecRank == 0)
164 return rewriter.notifyMatchFailure(xferOp, "0D vectors are not supported");
165
166 AffineMap map = xferOp.getPermutationMap();
167 if (!map.isProjectedPermutation(/*allowZeroInResults=*/false))
168 return rewriter.notifyMatchFailure(xferOp, "Unsupported permutation map");
169 unsigned numInputDims = map.getNumInputs();
170 for (AffineExpr expr : map.getResults().take_back(vecRank)) {
171 auto dim = dyn_cast<AffineDimExpr>(expr);
172 if (dim.getPosition() < (numInputDims - vecRank))
173 return rewriter.notifyMatchFailure(
174 xferOp, "Only the innermost dimensions can be accessed");
175 }
176
177 return success();
178}
179
180// Adjusts the strides of a memref according to a given permutation map for
181// vector operations.
182//
183// This function updates the innermost strides in the `strides` array to
184// reflect the permutation specified by `permMap`. The permutation is computed
185// using the inverse and broadcasting-aware version of the permutation map,
186// and is applied to the relevant strides. This ensures that memory accesses
187// are consistent with the logical permutation of vector elements.
188//
189// Example:
190// Suppose we have a memref of rank 4 with strides `[s0, s1, s2, s3]`.
191// If the permutation map swaps the last two dimensions (e.g., [0, 1] -> [1,
192// 0]), then after calling this function, the last two strides will be
193// swapped:
194// Original strides: [s0, s1, s2, s3]
195// After permutation: [s0, s1, s3, s2]
196//
197static void adjustStridesForPermutation(AffineMap permMap,
198 SmallVectorImpl<Value> &strides) {
199
203 SmallVector<int64_t> perms64(perms.begin(), perms.end());
204 strides = applyPermutation(strides, perms64);
205}
206
207// Computes memory strides and a memref offset for vector transfer operations,
208// handling both static and dynamic memrefs while applying permutation
209// transformations for XeGPU lowering.
210template <
211 typename OpType,
212 typename = std::enable_if_t<llvm::is_one_of<
213 std::decay_t<OpType>, vector::TransferReadOp, vector::TransferWriteOp,
214 vector::GatherOp, vector::ScatterOp>::value>>
215static std::pair<SmallVector<Value>, Value>
216computeMemrefMeta(OpType xferOp, PatternRewriter &rewriter) {
217 SmallVector<Value> strides;
218 Value baseMemref = xferOp.getBase();
219 MemRefType memrefType = dyn_cast<MemRefType>(baseMemref.getType());
220
221 Location loc = xferOp.getLoc();
222 Value offsetVal = nullptr;
223 if (memrefType.hasStaticShape()) {
224 int64_t offset;
225 SmallVector<int64_t> intStrides;
226 if (failed(memrefType.getStridesAndOffset(intStrides, offset)))
227 return {{}, offsetVal};
228 bool hasDynamicStrides = llvm::any_of(intStrides, [](int64_t strideVal) {
229 return ShapedType::isDynamic(strideVal);
230 });
231
232 if (!hasDynamicStrides)
233 for (int64_t s : intStrides)
234 strides.push_back(arith::ConstantIndexOp::create(rewriter, loc, s));
235
236 if (!ShapedType::isDynamic(offset))
237 offsetVal = arith::ConstantIndexOp::create(rewriter, loc, offset);
238 }
239
240 if (strides.empty() || !offsetVal) {
241 // For dynamic shape memref, use memref.extract_strided_metadata to get
242 // stride values
243 unsigned rank = memrefType.getRank();
244 Type indexType = rewriter.getIndexType();
245
246 // Result types: [base_memref, offset, stride0, stride1, ..., strideN-1,
247 // size0, size1, ..., sizeN-1]
248 SmallVector<Type> resultTypes;
249 resultTypes.push_back(MemRefType::get(
250 {}, memrefType.getElementType())); // base memref (unranked)
251 resultTypes.push_back(indexType); // offset
252
253 for (unsigned i = 0; i < rank; ++i)
254 resultTypes.push_back(indexType); // strides
255
256 for (unsigned i = 0; i < rank; ++i)
257 resultTypes.push_back(indexType); // sizes
258
259 auto meta = memref::ExtractStridedMetadataOp::create(
260 rewriter, loc, resultTypes, baseMemref);
261
262 if (strides.empty())
263 strides.append(meta.getStrides().begin(), meta.getStrides().end());
264
265 if (!offsetVal)
266 offsetVal = meta.getOffset();
267 }
268
269 // Strides are returned in original memref order; permutation is applied in
270 // computeOffsets only where offsets are indexed in vector order.
271 return {strides, offsetVal};
272}
273
274// Adds the transfer's indices to `baseOffset`, the memref's own element offset:
275// the element at `indices` sits at `baseOffset + sum(indices[d] * strides[d])`.
276static Value computeBaseOffset(VectorTransferOpInterface xferOp,
277 PatternRewriter &rewriter,
278 ArrayRef<Value> strides, Value baseOffset) {
279 Location loc = xferOp.getLoc();
280 for (auto [index, stride] : llvm::zip_equal(xferOp.getIndices(), strides)) {
281 Value contrib = arith::MulIOp::create(rewriter, loc, index, stride);
282 baseOffset = arith::AddIOp::create(rewriter, loc, baseOffset, contrib);
283 }
284 return baseOffset;
285}
286
287// This function compute the vectors of localOffsets for scattered load/stores.
288// It is used in the lowering of vector.transfer_read/write to
289// load_gather/store_scatter Example:
290// %0 = vector.transfer_read %expand_shape[%block_id_y, %c0, %c0, %c0, %c0],
291// %cst {in_bounds = [true, true, true, true]}>} :
292// memref<8x4x2x6x32xbf16>, vector<4x2x6x32xbf16>
293//
294// %6 = vector.step: vector<4xindex>
295// %7 = vector.step: vector<2xindex>
296// %8 = vector.step: vector<6xindex>
297// %9 = vector.step: vector<32xindex>
298// %10 = arith.mul %6, 384
299// %11 = arith.mul %7, 192
300// %12 = arith.mul %8, 32
301// %13 = arith.mul %9, 1
302// %14 = vector.shape_cast %10: vector<4xindex> -> vector<4x1x1x1xbf16>
303// %15 = vector.shape_cast %11: vector<2xindex> -> vector<1x2x1x1xbf16>
304// %16 = vector.shape_cast %12: vector<6xindex> -> vector<1x1x6x1xbf16>
305// %17 = vector.shape_cast %13: vector<32xindex> -> vector<1x1x1x32xbf16>
306// %18 = vector.broadcast %14: vector<4x1x1x1xbf16> -> vector<4x2x6x32xindex>
307// %19 = vector.broadcast %15: vector<1x2x1x1xbf16> -> vector<4x2x6x32xindex>
308// %20 = vector.broadcast %16: vector<1x1x6x1xbf16> -> vector<4x2x6x32xindex>
309// %21 = vector.broadcast %17: vector<1x1x1x32xbf16> -> vector<4x2x6x32xindex>
310// %22 = arith.add %18, %19
311// %23 = arith.add %20, %21
312// %local_offsets = arith.add %22, %23
313// %orig_offset = %block_id_y * 4x2x6x32 // consider using affine map
314// %offsets = memref_offset + orig_offset + local_offsets
315static Value computeOffsets(VectorTransferOpInterface xferOp,
316 PatternRewriter &rewriter, ArrayRef<Value> strides,
317 Value baseOffset) {
318 Location loc = xferOp.getLoc();
319 VectorType vectorType = xferOp.getVectorType();
320 ArrayRef<int64_t> vectorShape = vectorType.getShape();
321
322 // Create vector.step operations for each dimension
323 SmallVector<Value> stepVectors;
324 llvm::map_to_vector(vectorShape, [&](int64_t dim) {
325 auto stepType = VectorType::get({dim}, rewriter.getIndexType());
326 auto stepOp = vector::StepOp::create(rewriter, loc, stepType);
327 stepVectors.push_back(stepOp);
328 return stepOp;
329 });
330
331 // Local offsets are indexed in vector order, so permute strides; the base
332 // offset below uses the original memref-order strides.
333 SmallVector<Value> permutedStrides(strides.begin(), strides.end());
334 adjustStridesForPermutation(xferOp.getPermutationMap(), permutedStrides);
335
336 // Multiply step vectors by corresponding strides
337 size_t memrefRank = permutedStrides.size();
338 size_t vectorRank = vectorShape.size();
339 SmallVector<Value> strideMultiplied;
340 for (size_t i = 0; i < vectorRank; ++i) {
341 size_t memrefDim = memrefRank - vectorRank + i;
342 Value strideValue = permutedStrides[memrefDim];
343 auto mulType = dyn_cast<VectorType>(stepVectors[i].getType());
344 auto bcastOp =
345 vector::BroadcastOp::create(rewriter, loc, mulType, strideValue);
346 auto mulOp = arith::MulIOp::create(rewriter, loc, stepVectors[i], bcastOp);
347 strideMultiplied.push_back(mulOp);
348 }
349
350 // Shape cast each multiplied vector to add singleton dimensions
351 SmallVector<Value> shapeCasted;
352 for (size_t i = 0; i < vectorRank; ++i) {
353 SmallVector<int64_t> newShape(vectorRank, 1);
354 newShape[i] = vectorShape[i];
355 auto newType = VectorType::get(newShape, rewriter.getIndexType());
356 auto castOp = vector::ShapeCastOp::create(rewriter, loc, newType,
357 strideMultiplied[i]);
358 shapeCasted.push_back(castOp);
359 }
360
361 // Broadcast each shape-casted vector to full vector shape
362 SmallVector<Value> broadcasted;
363 auto fullIndexVectorType =
364 VectorType::get(vectorShape, rewriter.getIndexType());
365 for (Value shapeCastVal : shapeCasted) {
366 auto broadcastOp = vector::BroadcastOp::create(
367 rewriter, loc, fullIndexVectorType, shapeCastVal);
368 broadcasted.push_back(broadcastOp);
369 }
370
371 // Add all broadcasted vectors together to compute local offsets
372 Value localOffsets = broadcasted[0];
373 for (size_t i = 1; i < broadcasted.size(); ++i)
374 localOffsets =
375 arith::AddIOp::create(rewriter, loc, localOffsets, broadcasted[i]);
376
377 // Broadcast base offset to match vector shape
378 baseOffset = computeBaseOffset(xferOp, rewriter, strides, baseOffset);
379 Value bcastBase = vector::BroadcastOp::create(
380 rewriter, loc, fullIndexVectorType, baseOffset);
381 localOffsets = arith::AddIOp::create(rewriter, loc, bcastBase, localOffsets);
382 return localOffsets;
383}
384
385// Returns the size of memref dimension `dim` as a value.
386static Value getMemrefDimSize(VectorTransferOpInterface xferOp, unsigned dim,
387 PatternRewriter &rewriter) {
388 Location loc = xferOp.getLoc();
389 auto memrefTy = cast<MemRefType>(xferOp.getShapedType());
390 if (memrefTy.isDynamicDim(dim))
391 return memref::DimOp::create(rewriter, loc, xferOp.getBase(), dim)
392 .getResult();
393 return arith::ConstantIndexOp::create(rewriter, loc, memrefTy.getDimSize(dim))
394 .getResult();
395}
396
397// Builds the mask a scattered access needs to stay within the source bounds.
398//
399// The element at vector position `i` along vector dimension `v` accesses memref
400// dimension `d` (the one `v` maps to) at `indices[d] + i`, so it is in bounds
401// iff `i < dim(d) - indices[d]`. Dimensions the transfer declares in-bounds are
402// skipped; non-vector dimensions are always in bounds by the op's contract. The
403// per-dimension masks are spread over the full vector shape and combined, the
404// same way `computeOffsets` spreads the per-dimension offsets.
405//
406// Example, for a `vector<8xf32>` read of a `memref<?xf32>` at `%off`:
407// %dim = memref.dim %src, %c0
408// %limit = arith.subi %dim, %off
409// %step = vector.step : vector<8xindex>
410// %bcast = vector.broadcast %limit : index to vector<8xindex>
411// %mask = arith.cmpi slt, %step, %bcast : vector<8xindex>
412static Value computeInBoundsMask(VectorTransferOpInterface xferOp,
413 PatternRewriter &rewriter) {
414 Location loc = xferOp.getLoc();
415 ArrayRef<int64_t> vectorShape = xferOp.getVectorType().getShape();
416 auto maskType = VectorType::get(vectorShape, rewriter.getI1Type());
417 AffineMap map = xferOp.getPermutationMap();
418 OperandRange indices = xferOp.getIndices();
419
420 Value mask;
421 for (unsigned v = 0, e = vectorShape.size(); v < e; ++v) {
422 if (xferOp.isDimInBounds(v))
423 continue;
424 unsigned d = cast<AffineDimExpr>(map.getResult(v)).getPosition();
425 Value bound = getMemrefDimSize(xferOp, d, rewriter);
426 // A negative limit masks the dimension off entirely, which is what a
427 // starting index past the end of the dimension should do.
428 Value limit = arith::SubIOp::create(rewriter, loc, bound, indices[d]);
429 auto stepType = VectorType::get({vectorShape[v]}, rewriter.getIndexType());
430 Value step = vector::StepOp::create(rewriter, loc, stepType);
431 Value limitVec =
432 vector::BroadcastOp::create(rewriter, loc, stepType, limit);
433 Value dimMask = arith::CmpIOp::create(
434 rewriter, loc, arith::CmpIPredicate::slt, step, limitVec);
435 if (vectorShape.size() > 1) {
436 SmallVector<int64_t> expandedShape(vectorShape.size(), 1);
437 expandedShape[v] = vectorShape[v];
438 dimMask = vector::ShapeCastOp::create(
439 rewriter, loc, VectorType::get(expandedShape, rewriter.getI1Type()),
440 dimMask);
441 dimMask = vector::BroadcastOp::create(rewriter, loc, maskType, dimMask);
442 }
443 mask = mask
444 ? arith::AndIOp::create(rewriter, loc, mask, dimMask).getResult()
445 : dimMask;
446 }
447 if (mask)
448 return mask;
449 return vector::ConstantMaskOp::create(rewriter, loc, maskType, vectorShape);
450}
451
452// Builds that mask as a scalar `i1`, for a transfer of a single element. The
453// element sits at `indices`, so `i < dim(d) - indices[d]` at `i == 0` reduces
454// to `indices[d] < dim(d)`, which needs neither the limit subtraction nor the
455// step vector. The dimensions checked are the same.
456//
457// Example, for a `vector<1xf32>` read of a `memref<?xf32>` at `%off`:
458// %dim = memref.dim %src, %c0
459// %mask = arith.cmpi slt, %off, %dim : index
460static Value computeUnitInBoundsMask(VectorTransferOpInterface xferOp,
461 PatternRewriter &rewriter) {
462 Location loc = xferOp.getLoc();
463 AffineMap map = xferOp.getPermutationMap();
464 OperandRange indices = xferOp.getIndices();
465
466 Value mask;
467 for (unsigned v = 0, e = xferOp.getVectorType().getRank(); v < e; ++v) {
468 if (xferOp.isDimInBounds(v))
469 continue;
470 unsigned d = cast<AffineDimExpr>(map.getResult(v)).getPosition();
471 Value bound = getMemrefDimSize(xferOp, d, rewriter);
472 Value dimMask = arith::CmpIOp::create(
473 rewriter, loc, arith::CmpIPredicate::slt, indices[d], bound);
474 mask = mask
475 ? arith::AndIOp::create(rewriter, loc, mask, dimMask).getResult()
476 : dimMask;
477 }
478 if (mask)
479 return mask;
480 return arith::ConstantOp::create(rewriter, loc, rewriter.getBoolAttr(true));
481}
482
483// Compute the element-wise offsets for vector.gather or vector.scatter ops.
484//
485// This function linearizes the base offsets of the gather/scatter operation
486// and combines them with the per-element indices to produce a final vector of
487// memory offsets.
488template <
489 typename OpType,
490 typename = std::enable_if_t<llvm::is_one_of<
491 std::decay_t<OpType>, vector::GatherOp, vector::ScatterOp>::value>>
492static Value computeOffsets(PatternRewriter &rewriter, OpType gatScatOp,
493 ArrayRef<Value> strides, Value baseOffset) {
494 Location loc = gatScatOp.getLoc();
495 SmallVector<Value> offsets = gatScatOp.getOffsets();
496 for (size_t i = 0; i < offsets.size(); ++i) {
497 Value offsetContrib =
498 arith::MulIOp::create(rewriter, loc, offsets[i], strides[i]);
499 baseOffset =
500 arith::AddIOp::create(rewriter, loc, baseOffset, offsetContrib);
501 }
502 Value indices = gatScatOp.getIndices();
503 VectorType vecType = cast<VectorType>(indices.getType());
504
505 Value strideVector =
506 vector::BroadcastOp::create(rewriter, loc, vecType, strides.back())
507 .getResult();
508 Value stridedIndices =
509 arith::MulIOp::create(rewriter, loc, strideVector, indices).getResult();
510
511 Value baseVector =
512 vector::BroadcastOp::create(
513 rewriter, loc,
514 VectorType::get(vecType.getShape(), rewriter.getIndexType()),
515 baseOffset)
516 .getResult();
517 return arith::AddIOp::create(rewriter, loc, baseVector, stridedIndices)
518 .getResult();
519}
520
521// Collapses shapes of a nD memref to the target rank while applying offsets for
522// the collapsed dimensions. Returns the new memref value and the remaining
523// offsets for the last targetRank dimensions. For example:
524// input: %memref = memref<2x4x8x32xf32>, offsets=[%i0, %i1, %i2, %i3],
525// output: %memref[%i0, %i1, 0, 0] -> memref<8x32xf32>, offsets: [%i2, %i3]
526static std::pair<Value, SmallVector<OpFoldResult>>
527convertMemrefAndOffsetsToTargetRank(PatternRewriter &rewriter, Location loc,
530 int64_t targetRank) {
531 auto memrefType = cast<MemRefType>(memref.getType());
532 unsigned rank = memrefType.getRank();
533
534 if (rank <= targetRank)
535 return {memref, offsets};
536
537 int64_t numCombinedDims = rank - targetRank;
538 SmallVector<OpFoldResult> subviewOffsets;
539 SmallVector<OpFoldResult> subviewSizes;
540 SmallVector<OpFoldResult> subviewStrides;
541
542 // For the combined dimensions: use the provided offsets, size=1, stride=1
543 for (unsigned i = 0; i < numCombinedDims; ++i) {
544 subviewOffsets.push_back(offsets[i]);
545 subviewSizes.push_back(rewriter.getI64IntegerAttr(1));
546 subviewStrides.push_back(rewriter.getI64IntegerAttr(1));
547 }
548
549 // For the last targetRank dimensions: offset=0, use full size, stride=1
550 SmallVector<int64_t> resultShape;
551 auto originalShape = memrefType.getShape();
552 auto meta = memref::ExtractStridedMetadataOp::create(rewriter, loc, memref);
553 for (unsigned i = numCombinedDims; i < rank; ++i) {
554 subviewOffsets.push_back(rewriter.getI64IntegerAttr(0));
555 if (ShapedType::isDynamic(originalShape[i])) {
556 subviewSizes.push_back(meta.getSizes()[i]);
557 resultShape.push_back(ShapedType::kDynamic);
558 } else {
559 subviewSizes.push_back(rewriter.getI64IntegerAttr(originalShape[i]));
560 resultShape.push_back(originalShape[i]);
561 }
562 subviewStrides.push_back(rewriter.getI64IntegerAttr(1));
563 }
564
565 auto resultType = memref::SubViewOp::inferRankReducedResultType(
566 resultShape, memrefType, subviewOffsets, subviewSizes, subviewStrides);
567 auto subviewOp =
568 memref::SubViewOp::create(rewriter, loc, resultType, memref,
569 subviewOffsets, subviewSizes, subviewStrides);
570
571 // Return the remaining offsets for the last targetRank dimensions
572 SmallVector<OpFoldResult> newOffsets(offsets.begin() + numCombinedDims,
573 offsets.end());
574 return {subviewOp.getResult(), newOffsets};
575}
576
577template <
578 typename OpType,
579 typename = std::enable_if_t<llvm::is_one_of<
580 std::decay_t<OpType>, vector::TransferReadOp, vector::TransferWriteOp,
581 vector::GatherOp, vector::ScatterOp>::value>>
582// Convert memref to i64 base pointer
583static Value memrefToIndexPtr(OpType xferOp, PatternRewriter &rewriter) {
584 Location loc = xferOp.getLoc();
585 auto indexPtr = memref::ExtractAlignedPointerAsIndexOp::create(
586 rewriter, loc, xferOp.getBase())
587 .getResult();
588 return arith::IndexCastOp::create(rewriter, loc, rewriter.getI64Type(),
589 indexPtr)
590 .getResult();
591}
592
593// Returns true if every use of `vec` extracts a scalar element from it. An
594// unused value does not qualify.
595static bool isUsedAsScalar(Value vec) {
596 if (vec.use_empty())
597 return false;
598 return llvm::all_of(vec.getUsers(), [](Operation *user) {
599 auto extractOp = dyn_cast<vector::ExtractOp>(user);
600 return extractOp && !isa<VectorType>(extractOp.getResult().getType());
601 });
602}
603
604// Lowers a transfer of a single element to a scalar `xegpu.load`.
605//
606// The transfer touches one location, which its own indices name, so the load
607// takes a scalar offset and a scalar mask - no descriptor, no lane vectors -
608// and the result is broadcast back to the transfer's unit-size vector type.
609//
610// %off = arith.addi %base, %contrib : index
611// %inb = arith.cmpi slt, %idx, %dim : index
612// %val = xegpu.load %src[%off], %inb : i64, index, i1 -> f32
613// %vec = vector.broadcast %val : f32 to vector<1xf32>
614static LogicalResult lowerToScalarLoadOp(vector::TransferReadOp readOp,
615 PatternRewriter &rewriter) {
616 Location loc = readOp.getLoc();
617 VectorType vectorType = readOp.getVectorType();
618 if (!isa<MemRefType>(readOp.getShapedType()))
619 return rewriter.notifyMatchFailure(readOp, "Expected memref source");
620
621 auto meta = computeMemrefMeta(readOp, rewriter);
622 if (meta.first.empty())
623 return rewriter.notifyMatchFailure(readOp, "Failed to compute strides");
624
625 Value offset = computeBaseOffset(readOp, rewriter, meta.first, meta.second);
626 Value flatMemref = memrefToIndexPtr(readOp, rewriter);
627
628 Value mask = computeUnitInBoundsMask(readOp, rewriter);
629 auto loadOp = xegpu::LoadGatherOp::create(
630 rewriter, loc, vectorType.getElementType(), flatMemref, offset, mask,
631 /*l1_hint=*/xegpu::CachePolicyAttr{},
632 /*l2_hint=*/xegpu::CachePolicyAttr{},
633 /*l3_hint=*/xegpu::CachePolicyAttr{},
634 /*layout=*/nullptr, /*contiguity=*/nullptr);
635
636 // A masked-off xegpu.load is unspecified, so the padding has to be applied
637 // explicitly, as in lowerToScatteredLoadOp. A poison padding is "don't care".
638 Value scalar = loadOp.getResult();
639 if (readOp.hasOutOfBoundsDim() &&
640 !readOp.getPadding().getDefiningOp<ub::PoisonOp>())
641 scalar = arith::SelectOp::create(rewriter, loc, mask, scalar,
642 readOp.getPadding());
643
644 rewriter.replaceOpWithNewOp<vector::BroadcastOp>(readOp, vectorType, scalar);
645 return success();
646}
647
648static LogicalResult lowerToScatteredLoadOp(vector::TransferReadOp readOp,
649 PatternRewriter &rewriter) {
650
651 Location loc = readOp.getLoc();
652 VectorType vectorType = readOp.getVectorType();
653 auto memrefType = dyn_cast<MemRefType>(readOp.getShapedType());
654 if (!memrefType)
655 return rewriter.notifyMatchFailure(readOp, "Expected memref source");
656
657 auto meta = computeMemrefMeta(readOp, rewriter);
658 if (meta.first.empty())
659 return rewriter.notifyMatchFailure(readOp, "Failed to compute strides");
660
661 Value localOffsets =
662 computeOffsets(readOp, rewriter, meta.first, meta.second);
663
664 Value flatMemref = memrefToIndexPtr(readOp, rewriter);
665
666 Value mask = computeInBoundsMask(readOp, rewriter);
667 auto gatherOp = xegpu::LoadGatherOp::create(
668 rewriter, loc, vectorType, flatMemref, localOffsets, mask,
669 /*l1_hint=*/xegpu::CachePolicyAttr{},
670 /*l2_hint=*/xegpu::CachePolicyAttr{},
671 /*l3_hint=*/xegpu::CachePolicyAttr{},
672 /*layout=*/nullptr, /*contiguity=*/nullptr);
673
674 // The masked-off lanes of an xegpu.load are unspecified, so the transfer's
675 // padding has to be applied explicitly. A poison padding leaves them "don't
676 // care", so the select can be skipped.
677 Value result = gatherOp.getResult();
678 if (readOp.hasOutOfBoundsDim() &&
679 !readOp.getPadding().getDefiningOp<ub::PoisonOp>()) {
680 Value padding = vector::BroadcastOp::create(rewriter, loc, vectorType,
681 readOp.getPadding());
682 result = arith::SelectOp::create(rewriter, loc, mask, result, padding);
683 }
684
685 rewriter.replaceOp(readOp, result);
686 return success();
687}
688
689static LogicalResult lowerToScatteredStoreOp(vector::TransferWriteOp writeOp,
690 PatternRewriter &rewriter) {
691
692 Location loc = writeOp.getLoc();
693 auto memrefType = dyn_cast<MemRefType>(writeOp.getShapedType());
694 if (!memrefType)
695 return rewriter.notifyMatchFailure(writeOp, "Expected memref source");
696
697 auto meta = computeMemrefMeta(writeOp, rewriter);
698 if (meta.first.empty())
699 return rewriter.notifyMatchFailure(writeOp, "Failed to compute strides");
700
701 Value localOffsets =
702 computeOffsets(writeOp, rewriter, meta.first, meta.second);
703
704 Value flatMemref = memrefToIndexPtr(writeOp, rewriter);
705
706 // Out-of-bounds elements are simply not stored, so the mask is all this
707 // needs - no counterpart to the read's padding.
708 Value mask = computeInBoundsMask(writeOp, rewriter);
709 xegpu::StoreScatterOp::create(rewriter, loc, writeOp.getVector(), flatMemref,
710 localOffsets, mask,
711 /*l1_hint=*/xegpu::CachePolicyAttr{},
712 /*l2_hint=*/xegpu::CachePolicyAttr{},
713 /*l3_hint=*/xegpu::CachePolicyAttr{},
714 /*layout=*/nullptr, /*contiguity=*/nullptr);
715 rewriter.eraseOp(writeOp);
716 return success();
717}
718
719struct TransferReadLowering : public OpRewritePattern<vector::TransferReadOp> {
720 using Base::Base;
721
722 LogicalResult matchAndRewrite(vector::TransferReadOp readOp,
723 PatternRewriter &rewriter) const override {
724 Location loc = readOp.getLoc();
725
726 if (failed(transferPreconditions(rewriter, readOp)))
727 return failure();
728 auto readMemTy = cast<MemRefType>(readOp.getShapedType());
729 VectorType loadedVecTy = readOp.getVectorType();
730 bool isOutOfBounds = readOp.hasOutOfBoundsDim();
731 // Check if the memref has address space 3 (shared local memory)
732 bool isSharedMemory = xegpu::XeGPUDialect::isSharedMemory(readMemTy);
733 // Handle the SLM case.
734 if (isSharedMemory) {
735 // load_matrix supports 1D and 2D loads from SLM.
736 if (loadedVecTy.getRank() != 1 && loadedVecTy.getRank() != 2)
737 return rewriter.notifyMatchFailure(
738 readOp, "Only 1D and 2D vector loads are supported for SLM");
739 AffineMap readMap = readOp.getPermutationMap();
740 if (!readMap.isMinorIdentity())
741 return rewriter.notifyMatchFailure(
742 readOp,
743 "Non identity transposition is not supported for SLM loads.");
744 // Out of bounds case is not supported for SLM loads.
745 if (isOutOfBounds)
746 return rewriter.notifyMatchFailure(
747 readOp, "Out-of-bounds access is not supported for SLM loads");
748
749 // Create mem_desc for SLM
750 auto memDescType =
751 xegpu::MemDescType::get(rewriter.getContext(), readMemTy.getShape(),
752 readMemTy.getElementType(),
753 /*mem_layout=*/nullptr);
754 auto createMemDescOp = xegpu::CreateMemDescOp::create(
755 rewriter, loc, memDescType, readOp.getBase());
756 // Convert indices to OpFoldResult for LoadMatrixOp
757 SmallVector<OpFoldResult> indices =
758 getAsOpFoldResult(readOp.getIndices());
759 auto loadMatrixOp = xegpu::LoadMatrixOp::create(
760 rewriter, loc, loadedVecTy, createMemDescOp.getResult(), indices,
761 /*layout=*/nullptr);
762
763 rewriter.replaceOp(readOp, loadMatrixOp.getResult());
764 return success();
765 }
766
767 // A transfer of a single element that is only ever extracted to a scalar
768 // becomes a scalar load, at any rank: the vector is a wrapper its consumers
769 // undo, so neither an nd descriptor nor a lane-vector gather buys anything.
770 // A unit-size vector genuinely used as a vector keeps the paths below.
771 if (loadedVecTy.getNumElements() == 1 && isUsedAsScalar(readOp.getResult()))
772 return lowerToScalarLoadOp(readOp, rewriter);
773
774 const xegpu::uArch::uArch *uArch =
775 xegpu::uArch::getUArch(xegpu::getChipStr(readOp).value_or(""));
776
777 // An nd block load can realize a minor-identity map directly, or an
778 // innermost-two-dims transpose via a trailing vector.transpose. Any other
779 // permutation (e.g. a mid-vector transpose of a high-dim load) is left to
780 // the scattered path, which permutes strides explicitly.
781 AffineMap readMap = readOp.getPermutationMap();
782 bool isTransposeLoad = isInnermostTwoDimsTransposed(readMap);
783
784 // A block load transfers the tile as it is laid out in memory, so a
785 // transposing read describes a tile whose innermost two dims are swapped.
786 Type elementType = loadedVecTy.getElementType();
787 SmallVector<int64_t> descShape(loadedVecTy.getShape());
788 if (isTransposeLoad) {
789 size_t rank = descShape.size();
790 assert(rank >= 2 && "Transpose requires at least 2 dimensions");
791 std::swap(descShape[rank - 1], descShape[rank - 2]);
792 }
793
794 // Prefer an nd block load. It requires a vector of rank >= 2 backed by a
795 // scalar-element memref, a map the block load can realize, and a tile shape
796 // the target's 2D block load can access. 1D vectors use the scattered
797 // xegpu.load path instead, which has a richer interface (e.g. layout
798 // capabilities). Out-of-bounds reads are allowed as long as the padding
799 // matches load_nd's implicit zero padding.
800 bool canLowerToLoadNd =
801 loadedVecTy.getRank() > 1 &&
802 (readMap.isMinorIdentity() || isTransposeLoad) &&
803 readMemTy.getElementType().isIntOrFloat() &&
804 (!isOutOfBounds || isZeroOrPoisonPadding(readOp.getPadding())) &&
805 isSupportedBlockShape(
806 uArch, xegpu::uArch::InstructionKind::Subgroup2DBlockLoad,
807 descShape, elementType, /*hasTranspose=*/isTransposeLoad);
808
809 if (canLowerToLoadNd) {
810 // The load produces the memory-ordered tile; the transpose below restores
811 // the shape the transfer_read asked for.
812 if (isTransposeLoad)
813 loadedVecTy = VectorType::get(descShape, elementType);
814 auto descType = xegpu::TensorDescType::get(
815 descShape, elementType, /*array_length=*/1,
816 /*boundary_check=*/isOutOfBounds, xegpu::MemorySpace::Global);
817 auto [src, indices] = convertMemrefAndOffsetsToTargetRank(
818 rewriter, loc, readOp.getBase(),
819 getAsOpFoldResult(readOp.getIndices()), loadedVecTy.getRank());
820 // By default, no specific caching policy is assigned.
821 xegpu::CachePolicyAttr hint = nullptr;
822 xegpu::CreateNdDescOp ndDesc = xegpu::CreateNdDescOp::create(
823 rewriter, loc, descType, dyn_cast<TypedValue<MemRefType>>(src));
824
825 Operation *loadedOp =
826 xegpu::LoadNdOp::create(rewriter, loc, loadedVecTy, ndDesc, indices,
827 /*packed=*/nullptr, /*transpose=*/nullptr,
828 /*l1_hint=*/hint,
829 /*l2_hint=*/hint, /*l3_hint=*/hint,
830 /*layout=*/nullptr);
831 if (isTransposeLoad) {
832 // Undo the innermost-two-dims swap with a trailing vector.transpose:
833 // keep the leading dimensions in place and interchange only the last
834 // two.
835 int64_t rank = loadedVecTy.getRank();
836 SmallVector<int64_t> perm(llvm::to_vector(llvm::seq<int64_t>(0, rank)));
837 std::swap(perm[rank - 1], perm[rank - 2]);
838 loadedOp = vector::TransposeOp::create(rewriter, loc,
839 loadedOp->getResult(0), perm);
840 }
841 rewriter.replaceOp(readOp, loadedOp);
842 return success();
843 }
844
845 // Fall back to a scattered load. It supports arbitrary permutations and any
846 // rank, and masks off the out-of-bounds elements.
847 return lowerToScatteredLoadOp(readOp, rewriter);
848 }
849};
850
851struct TransferWriteLowering
852 : public OpRewritePattern<vector::TransferWriteOp> {
853 using Base::Base;
854
855 LogicalResult matchAndRewrite(vector::TransferWriteOp writeOp,
856 PatternRewriter &rewriter) const override {
857 Location loc = writeOp.getLoc();
858
859 if (failed(transferPreconditions(rewriter, writeOp)))
860 return failure();
861 // Perform common data transfer checks.
862 VectorType vecTy = writeOp.getVectorType();
863 auto writeMemTy = cast<MemRefType>(writeOp.getShapedType());
864 // Check if the memref has address space 3 (shared local memory)
865 bool isSharedMemory = xegpu::XeGPUDialect::isSharedMemory(writeMemTy);
866
867 // For shared local memory (address space 3), use create_mem_desc +
868 // store_matrix
869 if (isSharedMemory) {
870 // store_matrix supports 1D and 2D stores to SLM.
871 if (vecTy.getRank() != 1 && vecTy.getRank() != 2)
872 return rewriter.notifyMatchFailure(
873 writeOp, "Only 1D and 2D vector stores are supported for SLM");
874 // Out of bounds case is not supported for SLM stores.
875 if (writeOp.hasOutOfBoundsDim())
876 return rewriter.notifyMatchFailure(
877 writeOp, "Out-of-bounds access is not supported for SLM stores");
878 // Create mem_desc for SLM
879 auto memDescType =
880 xegpu::MemDescType::get(rewriter.getContext(), writeMemTy.getShape(),
881 writeMemTy.getElementType(),
882 /*mem_layout=*/nullptr);
883
884 auto createMemDescOp = xegpu::CreateMemDescOp::create(
885 rewriter, loc, memDescType, writeOp.getBase());
886
887 // Convert indices to OpFoldResult for StoreMatrixOp
888 SmallVector<OpFoldResult> indices =
889 getAsOpFoldResult(writeOp.getIndices());
890
891 xegpu::StoreMatrixOp::create(rewriter, loc, writeOp.getVector(),
892 createMemDescOp.getResult(), indices,
893 /*layout=*/nullptr);
894
895 rewriter.eraseOp(writeOp);
896 return success();
897 }
898
899 const xegpu::uArch::uArch *uArch =
900 xegpu::uArch::getUArch(xegpu::getChipStr(writeOp).value_or(""));
901
902 // Prefer an nd block store. It requires a vector of rank >= 2 backed by a
903 // scalar-element memref, a minor-identity map (block stores have no
904 // transpose support), and a tile shape the target's 2D block store can
905 // access. 1D vectors use the scattered xegpu.store path instead, which has
906 // a richer interface. Out-of-bounds writes are handled by the descriptor's
907 // boundary check.
908 AffineMap map = writeOp.getPermutationMap();
909 bool canLowerToStoreNd =
910 vecTy.getRank() > 1 && map.isMinorIdentity() &&
911 writeMemTy.getElementType().isIntOrFloat() &&
912 isSupportedBlockShape(
913 uArch, xegpu::uArch::InstructionKind::Subgroup2DBlockStore,
914 vecTy.getShape(), vecTy.getElementType());
915
916 if (canLowerToStoreNd) {
917 auto [src, indices] = convertMemrefAndOffsetsToTargetRank(
918 rewriter, loc, writeOp.getBase(),
919 getAsOpFoldResult(writeOp.getIndices()), vecTy.getRank());
920
921 auto descType = xegpu::TensorDescType::get(
922 vecTy.getShape(), vecTy.getElementType(),
923 /*array_length=*/1, /*boundary_check=*/writeOp.hasOutOfBoundsDim(),
924 xegpu::MemorySpace::Global);
925 // By default, no specific caching policy is assigned.
926 xegpu::CachePolicyAttr hint = nullptr;
927 xegpu::CreateNdDescOp ndDesc = xegpu::CreateNdDescOp::create(
928 rewriter, loc, descType, dyn_cast<TypedValue<MemRefType>>(src));
929
930 auto storeOp = xegpu::StoreNdOp::create(
931 rewriter, loc, writeOp.getVector(), ndDesc, indices,
932 /*l1_hint=*/hint,
933 /*l2_hint=*/hint, /*l3_hint=*/hint,
934 /*layout=*/nullptr);
935 rewriter.replaceOp(writeOp, storeOp);
936 return success();
937 }
938
939 // Fall back to a scattered store. It supports arbitrary permutations and
940 // any rank, and masks off the out-of-bounds elements.
941 return lowerToScatteredStoreOp(writeOp, rewriter);
942 }
943};
944
945struct GatherLowering : public OpRewritePattern<vector::GatherOp> {
946 using Base::Base;
947
948 LogicalResult matchAndRewrite(vector::GatherOp gatherOp,
949 PatternRewriter &rewriter) const override {
950 auto srcTy = dyn_cast<MemRefType>(gatherOp.getBase().getType());
951 if (!srcTy)
952 return rewriter.notifyMatchFailure(gatherOp, "Expects memref source");
953
954 Location loc = gatherOp.getLoc();
955 VectorType vectorType = gatherOp.getVectorType();
956
957 auto meta = computeMemrefMeta(gatherOp, rewriter);
958 if (meta.first.empty())
959 return rewriter.notifyMatchFailure(gatherOp, "Failed to compute strides");
960
961 Value localOffsets =
962 computeOffsets(rewriter, gatherOp, meta.first, meta.second);
963 Value flatMemref = memrefToIndexPtr(gatherOp, rewriter);
964
965 auto xeGatherOp = xegpu::LoadGatherOp::create(
966 rewriter, loc, vectorType, flatMemref, localOffsets, gatherOp.getMask(),
967 /*l1_hint=*/xegpu::CachePolicyAttr{},
968 /*l2_hint=*/xegpu::CachePolicyAttr{},
969 /*l3_hint=*/xegpu::CachePolicyAttr{},
970 /*layout=*/nullptr, /*contiguity=*/nullptr);
971
972 auto selectOp =
973 arith::SelectOp::create(rewriter, loc, gatherOp.getMask(),
974 xeGatherOp.getResult(), gatherOp.getPassThru());
975 rewriter.replaceOp(gatherOp, selectOp.getResult());
976 return success();
977 }
978};
979
980struct ScatterLowering : public OpRewritePattern<vector::ScatterOp> {
981 using Base::Base;
982
983 LogicalResult matchAndRewrite(vector::ScatterOp scatterOp,
984 PatternRewriter &rewriter) const override {
985 auto srcTy = dyn_cast<MemRefType>(scatterOp.getBase().getType());
986 if (!srcTy)
987 return rewriter.notifyMatchFailure(scatterOp, "Expects memref source");
988
989 Location loc = scatterOp.getLoc();
990 auto meta = computeMemrefMeta(scatterOp, rewriter);
991 if (meta.first.empty())
992 return rewriter.notifyMatchFailure(scatterOp,
993 "Failed to compute strides");
994
995 Value localOffsets =
996 computeOffsets(rewriter, scatterOp, meta.first, meta.second);
997 Value flatMemref = memrefToIndexPtr(scatterOp, rewriter);
998
999 xegpu::StoreScatterOp::create(rewriter, loc, scatterOp.getValueToStore(),
1000 flatMemref, localOffsets, scatterOp.getMask(),
1001 /*l1_hint=*/xegpu::CachePolicyAttr{},
1002 /*l2_hint=*/xegpu::CachePolicyAttr{},
1003 /*l3_hint=*/xegpu::CachePolicyAttr{},
1004 /*layout=*/nullptr,
1005 /*contiguity=*/nullptr);
1006 rewriter.eraseOp(scatterOp);
1007 return success();
1008 }
1009};
1010
1011struct LoadLowering : public OpRewritePattern<vector::LoadOp> {
1012 using Base::Base;
1013
1014 LogicalResult matchAndRewrite(vector::LoadOp loadOp,
1015 PatternRewriter &rewriter) const override {
1016 Location loc = loadOp.getLoc();
1017
1018 VectorType vecTy = loadOp.getResult().getType();
1019 MemRefType memTy = loadOp.getBase().getType();
1020 // The plain vector.load lowering only supports 1D/2D block loads.
1021 if (vecTy.getRank() != 1 && vecTy.getRank() != 2)
1022 return rewriter.notifyMatchFailure(loadOp, "Expects 1D or 2D vector");
1023 if (!memTy.getElementType().isIntOrFloat())
1024 return rewriter.notifyMatchFailure(
1025 loadOp, "Unsupported memref element type: expected integer or float");
1026
1027 // Boundary check is available only for block instructions.
1028 bool boundaryCheck = vecTy.getRank() > 1;
1029 // By default, no specific caching policy is assigned.
1030 xegpu::CachePolicyAttr hint = nullptr;
1031
1032 auto [src, indices] = convertMemrefAndOffsetsToTargetRank(
1033 rewriter, loc, loadOp.getBase(), getAsOpFoldResult(loadOp.getIndices()),
1034 vecTy.getRank());
1035
1036 auto descType = xegpu::TensorDescType::get(
1037 vecTy.getShape(), vecTy.getElementType(), /*array_length=*/1,
1038 boundaryCheck, xegpu::MemorySpace::Global);
1039
1040 xegpu::CreateNdDescOp ndDesc = xegpu::CreateNdDescOp::create(
1041 rewriter, loc, descType, dyn_cast<TypedValue<MemRefType>>(src));
1042 auto loadNdOp =
1043 xegpu::LoadNdOp::create(rewriter, loc, vecTy, ndDesc, indices,
1044 /*packed=*/nullptr, /*transpose=*/nullptr,
1045 /*l1_hint=*/hint,
1046 /*l2_hint=*/hint, /*l3_hint=*/hint,
1047 /*layout=*/nullptr);
1048 rewriter.replaceOp(loadOp, loadNdOp);
1049
1050 return success();
1051 }
1052};
1053
1054struct StoreLowering : public OpRewritePattern<vector::StoreOp> {
1055 using Base::Base;
1056
1057 LogicalResult matchAndRewrite(vector::StoreOp storeOp,
1058 PatternRewriter &rewriter) const override {
1059 Location loc = storeOp.getLoc();
1060
1061 TypedValue<VectorType> vector = storeOp.getValueToStore();
1062 VectorType vecTy = vector.getType();
1063 MemRefType memTy = storeOp.getBase().getType();
1064 // The plain vector.store lowering only supports 1D/2D block stores.
1065 if (vecTy.getRank() != 1 && vecTy.getRank() != 2)
1066 return rewriter.notifyMatchFailure(storeOp, "Expects 1D or 2D vector");
1067 if (!memTy.getElementType().isIntOrFloat())
1068 return rewriter.notifyMatchFailure(
1069 storeOp,
1070 "Unsupported memref element type: expected integer or float");
1071
1072 // Boundary check is available only for block instructions.
1073 bool boundaryCheck = vecTy.getRank() > 1;
1074
1075 auto [src, indices] = convertMemrefAndOffsetsToTargetRank(
1076 rewriter, loc, storeOp.getBase(),
1077 getAsOpFoldResult(storeOp.getIndices()), vecTy.getRank());
1078
1079 auto descType = xegpu::TensorDescType::get(
1080 vecTy.getShape(), vecTy.getElementType(),
1081 /*array_length=*/1, boundaryCheck, xegpu::MemorySpace::Global);
1082
1083 // By default, no specific caching policy is assigned.
1084 xegpu::CachePolicyAttr hint = nullptr;
1085 xegpu::CreateNdDescOp ndDesc = xegpu::CreateNdDescOp::create(
1086 rewriter, loc, descType, dyn_cast<TypedValue<MemRefType>>(src));
1087
1088 auto storeNdOp =
1089 xegpu::StoreNdOp::create(rewriter, loc, vector, ndDesc, indices,
1090 /*l1_hint=*/hint,
1091 /*l2_hint=*/hint, /*l3_hint=*/hint,
1092 /*layout=*/nullptr);
1093
1094 rewriter.replaceOp(storeOp, storeNdOp);
1095
1096 return success();
1097 }
1098};
1099
1100// If `indexingMaps` describe a (batched) row-major matmul
1101// lhs[b..., m, k], rhs[b..., k, n], acc[b..., m, n]
1102// return the number of leading batch dims (0 for a plain 2D matmul);
1103// otherwise return std::nullopt.
1104static std::optional<int64_t>
1105getRowMajorMatmulBatchRank(ArrayAttr indexingMaps) {
1106 if (indexingMaps.size() != 3)
1107 return std::nullopt;
1108
1109 AffineMap mapA = cast<AffineMapAttr>(indexingMaps[0]).getValue();
1110 AffineMap mapB = cast<AffineMapAttr>(indexingMaps[1]).getValue();
1111 AffineMap mapC = cast<AffineMapAttr>(indexingMaps[2]).getValue();
1112
1113 // The result map exposes the batch dims followed by the core (m, n) dims.
1114 if (mapC.getNumResults() < 2)
1115 return std::nullopt;
1116 int64_t batchRank = mapC.getNumResults() - 2;
1117
1118 // A single `k` reduction gives batchRank + 3 iteration dims; each operand
1119 // map exposes batchRank + 2 dims (batch dims + 2 core dims).
1120 unsigned numDims = static_cast<unsigned>(batchRank) + 3;
1121 unsigned numOperandResults = static_cast<unsigned>(batchRank) + 2;
1122 if (mapA.getNumInputs() != numDims || mapB.getNumInputs() != numDims ||
1123 mapC.getNumInputs() != numDims)
1124 return std::nullopt;
1125 if (mapA.getNumResults() != numOperandResults ||
1126 mapB.getNumResults() != numOperandResults)
1127 return std::nullopt;
1128
1129 // Reconstruct the canonical maps from the batch/m/n dims of the result and
1130 // the k dim of lhs, then compare against the actual maps.
1131 MLIRContext *context = indexingMaps.getContext();
1132 ArrayRef<AffineExpr> batchDims = mapC.getResults().take_front(batchRank);
1133 AffineExpr m = mapC.getResult(batchRank);
1134 AffineExpr n = mapC.getResult(batchRank + 1);
1135 AffineExpr k = mapA.getResult(batchRank + 1);
1136
1137 SmallVector<AffineExpr> aDims = llvm::to_vector(batchDims);
1138 aDims.push_back(m);
1139 aDims.push_back(k);
1140 SmallVector<AffineExpr> bDims = llvm::to_vector(batchDims);
1141 bDims.push_back(k);
1142 bDims.push_back(n);
1143 SmallVector<AffineExpr> cDims = llvm::to_vector(batchDims);
1144 cDims.push_back(m);
1145 cDims.push_back(n);
1146
1147 auto expected = ArrayAttr::get(
1148 context,
1149 {AffineMapAttr::get(AffineMap::get(numDims, 0, aDims, context)),
1150 AffineMapAttr::get(AffineMap::get(numDims, 0, bDims, context)),
1151 AffineMapAttr::get(AffineMap::get(numDims, 0, cDims, context))});
1152 if (indexingMaps != expected)
1153 return std::nullopt;
1154 return batchRank;
1155}
1156
1157struct ContractionLowering : public OpRewritePattern<vector::ContractionOp> {
1158 using Base::Base;
1159
1160 LogicalResult matchAndRewrite(vector::ContractionOp contractOp,
1161 PatternRewriter &rewriter) const override {
1162 Location loc = contractOp.getLoc();
1163
1164 if (contractOp.getKind() != vector::CombiningKind::ADD)
1165 return rewriter.notifyMatchFailure(contractOp,
1166 "Expects add combining kind");
1167
1168 TypedValue<VectorType> lhs = contractOp.getLhs();
1169 TypedValue<VectorType> rhs = contractOp.getRhs();
1170 TypedValue<Type> acc = contractOp.getAcc();
1171 VectorType accType = dyn_cast<VectorType>(acc.getType());
1172 if (!accType)
1173 return rewriter.notifyMatchFailure(contractOp, "Expects vector acc");
1174
1175 std::optional<int64_t> batchRank =
1176 getRowMajorMatmulBatchRank(contractOp.getIndexingMapsAttr());
1177 if (!batchRank)
1178 return rewriter.notifyMatchFailure(
1179 contractOp,
1180 "Expects a (batched) row-major matmul: leading dims must "
1181 "be batch dims shared by lhs, rhs, and acc; innermost two "
1182 "dims must be (M, K), (K, N), and (M, N)");
1183
1184 // xegpu.dpas operands are limited to 2 batch + 2 core dims.
1185 if (*batchRank > 2)
1186 return rewriter.notifyMatchFailure(contractOp,
1187 "Expects operands of rank 4 or less");
1188
1189 auto dpasOp = xegpu::DpasOp::create(
1190 rewriter, loc, contractOp.getResultType(), lhs, rhs, acc,
1191 /*layout_a=*/nullptr, /*layout_b=*/nullptr, /*layout_cd=*/nullptr);
1192 rewriter.replaceOp(contractOp, dpasOp);
1193
1194 return success();
1195 }
1196};
1197
1198// Returns the `vector.shape_cast` that flattened a value of type `ndType` into
1199// `flat`, if that is how `flat` was produced.
1200static vector::ShapeCastOp getFlattenCast(Value flat, VectorType ndType) {
1201 auto shapeCast = flat.getDefiningOp<vector::ShapeCastOp>();
1202 if (shapeCast && shapeCast.getSourceVectorType() == ndType)
1203 return shapeCast;
1204 return nullptr;
1205}
1206
1207static DenseElementsAttr getDenseConstant(Value flat) {
1208 DenseElementsAttr elements;
1209 if (matchPattern(flat, m_Constant(&elements)))
1210 return elements;
1211 return nullptr;
1212}
1213
1214static vector::BroadcastOp getSplatBroadcast(Value flat) {
1215 auto broadcast = flat.getDefiningOp<vector::BroadcastOp>();
1216 if (broadcast && !isa<VectorType>(broadcast.getSourceType()))
1217 return broadcast;
1218 return nullptr;
1219}
1220
1221// Returns true if the cast `unflatten` creates will fold away, i.e. if `flat`
1222// is a flattening cast, a constant or a splat.
1223static bool canUnflatten(Value flat, VectorType ndType) {
1224 return getFlattenCast(flat, ndType) || getDenseConstant(flat) ||
1225 getSplatBroadcast(flat);
1226}
1227
1228// Reshape `flat` to `ndType`; `canUnflatten` must hold so the cast folds away.
1229static Value unflatten(PatternRewriter &rewriter, Value flat,
1230 VectorType ndType) {
1231 assert(canUnflatten(flat, ndType) && "expected the cast to fold away");
1232 return vector::ShapeCastOp::create(rewriter, flat.getLoc(), ndType, flat);
1233}
1234
1235// Restore the N-D form of a flattened `vector.gather` / `vector.scatter`.
1236//
1237// XeGPU layouts are expressed in terms of the N-D shape of the accessed data,
1238// so a flattened gather/scatter forces layout propagation to reason through the
1239// surrounding `vector.shape_cast` ops. That adds complexity and tends to yield
1240// layouts that lower to unoptimized code.
1241//
1242// Before:
1243// %cst = arith.constant dense<0.0> : vector<8192xbf16>
1244// %fi = vector.shape_cast %idx : vector<128x64xindex> to vector<8192xindex>
1245// %fm = vector.shape_cast %mask : vector<128x64xi1> to vector<8192xi1>
1246// %fr = vector.gather %src[%c0] [%fi], %fm, %cst : memref<?xbf16>,
1247// vector<8192xindex>, vector<8192xi1>, vector<8192xbf16>
1248// into vector<8192xbf16>
1249// %res = vector.shape_cast %fr : vector<8192xbf16> to vector<128x64xbf16>
1250//
1251// After:
1252// %cst = arith.constant dense<0.0> : vector<128x64xbf16>
1253// %res = vector.gather %src[%c0] [%idx], %mask, %cst : memref<?xbf16>,
1254// vector<128x64xindex>, vector<128x64xi1>, vector<128x64xbf16>
1255// into vector<128x64xbf16>
1256template <typename OpTy>
1257struct UnflattenGatherScatter : public OpRewritePattern<OpTy> {
1258 using OpRewritePattern<OpTy>::OpRewritePattern;
1259
1260 LogicalResult matchAndRewrite(OpTy op,
1261 PatternRewriter &rewriter) const override {
1262 constexpr bool isGather = std::is_same_v<OpTy, vector::GatherOp>;
1263
1264 if (!isa<MemRefType>(op.getBase().getType()))
1265 return rewriter.notifyMatchFailure(op, "expects a memref source");
1266
1267 if (op.getIndexVectorType().getRank() != 1)
1268 return rewriter.notifyMatchFailure(op, "index vector is not 1-D");
1269
1270 // The N-D shape comes from the index operand's producer: this only undoes
1271 // a flattening that already happened, it never invents a shape.
1272 auto indexCast =
1273 op.getIndices().template getDefiningOp<vector::ShapeCastOp>();
1274 if (!indexCast || indexCast.getSourceVectorType().getRank() < 2)
1275 return rewriter.notifyMatchFailure(
1276 op, "index vector is not a shape_cast of an N-D vector");
1277 VectorType ndIndexType = indexCast.getSourceVectorType();
1278 VectorType ndMaskType =
1279 ndIndexType.cloneWith(std::nullopt, rewriter.getI1Type());
1280 VectorType ndType = ndIndexType.cloneWith(
1281 std::nullopt, op.getVectorType().getElementType());
1282
1283 // Check everything before creating any IR: a partially applied rewrite
1284 // would leave dead ops behind.
1285 if (!canUnflatten(op.getMask(), ndMaskType))
1286 return rewriter.notifyMatchFailure(op, "cannot un-flatten the mask");
1287
1288 if constexpr (isGather) {
1289 if (!canUnflatten(op.getPassThru(), ndType))
1290 return rewriter.notifyMatchFailure(op,
1291 "cannot un-flatten the pass-thru");
1292
1293 Value mask = unflatten(rewriter, op.getMask(), ndMaskType);
1294 Value passThru = unflatten(rewriter, op.getPassThru(), ndType);
1295 auto ndGather = vector::GatherOp::create(
1296 rewriter, op.getLoc(), ndType, op.getBase(), op.getOffsets(),
1297 indexCast.getSource(), mask, passThru, op.getAlignmentAttr());
1298 ndGather->setDiscardableAttrs(op->getDiscardableAttrDictionary());
1299 rewriter.replaceOpWithNewOp<vector::ShapeCastOp>(op, op.getVectorType(),
1300 ndGather);
1301 } else {
1302 if (!canUnflatten(op.getValueToStore(), ndType))
1303 return rewriter.notifyMatchFailure(
1304 op, "cannot un-flatten the stored value");
1305
1306 Value mask = unflatten(rewriter, op.getMask(), ndMaskType);
1307 Value valueToStore = unflatten(rewriter, op.getValueToStore(), ndType);
1308 // Only operand types change, and this keeps the optional tensor result
1309 // untouched.
1310 rewriter.modifyOpInPlace(op, [&] {
1311 op.getIndicesMutable().assign(indexCast.getSource());
1312 op.getMaskMutable().assign(mask);
1313 op.getValueToStoreMutable().assign(valueToStore);
1314 });
1315 }
1316 return success();
1317 }
1318};
1319
1320// Un-flatten every gather/scatter that was flattened to 1-D operands, so that
1321// the conversion patterns and the XeGPU layouts downstream of them see the N-D
1322// shape of the accessed data.
1323static LogicalResult unflattenGatherScatter(Operation *root) {
1324 MLIRContext *ctx = root->getContext();
1325 RewritePatternSet patterns(ctx);
1326 patterns.add<UnflattenGatherScatter<vector::GatherOp>,
1327 UnflattenGatherScatter<vector::ScatterOp>>(ctx);
1328 vector::ShapeCastOp::getCanonicalizationPatterns(patterns, ctx);
1329 return applyPatternsGreedily(root, std::move(patterns));
1330}
1331
1332// Returns `memrefTy` with its memory space replaced by `newMemSpace`.
1333static MemRefType withMemorySpace(MemRefType memrefTy, Attribute newMemSpace) {
1334 return MemRefType::get(memrefTy.getShape(), memrefTy.getElementType(),
1335 memrefTy.getLayout(), newMemSpace);
1336}
1337
1338// Rewrite every `memref.alloca` not already in shared local memory (SLM) to
1339// be in SLM (address space 3), and propagate the new memory space through
1340// memref-producing aliasing users (e.g. memref.cast, memref.subview,
1341// memref.expand_shape, ...). Consumers that take a memref operand but
1342// produce a non-memref result (e.g. vector.transfer_read, vector.load) are
1343// left untouched: their operand type simply reflects the new memory space.
1344//
1345// This makes `xegpu.load_matrix`/`xegpu.store_matrix` lowering work end-to-end
1346// for IR coming from bufferization, which by default assigns memory space 0/1
1347// to allocations.
1348static void promoteAllocasToSLM(Operation *root) {
1349 MLIRContext *ctx = root->getContext();
1350 Attribute slmAttr = IntegerAttr::get(IntegerType::get(ctx, 64), 3);
1351
1352 // A user is treated as a memref-producing alias (e.g. memref.cast,
1353 // memref.subview, memref.expand_shape, ...) if it is side-effect free and
1354 // produces at least one memref result. This excludes ops like memref.copy
1355 // that have memory effects.
1356 auto isMemrefResultOp = [](Operation *op) {
1357 if (!isMemoryEffectFree(op))
1358 return false;
1359 return llvm::any_of(op->getResultTypes(),
1360 [](Type t) { return isa<MemRefType>(t); });
1361 };
1362
1363 // Update `v`'s type to have SLM memory space, then walk forward through
1364 // memref-producing users and update their result types accordingly.
1365 std::function<void(Value)> propagate = [&](Value v) {
1366 auto memrefTy = dyn_cast<MemRefType>(v.getType());
1367 if (!memrefTy || xegpu::XeGPUDialect::isSharedMemory(memrefTy))
1368 return;
1369 v.setType(withMemorySpace(memrefTy, slmAttr));
1370 for (Operation *user : v.getUsers()) {
1371 if (!isMemrefResultOp(user))
1372 continue;
1373 for (Value result : user->getResults())
1374 propagate(result);
1375 }
1376 };
1377
1379 root->walk([&](memref::AllocaOp op) {
1380 auto memrefTy = dyn_cast<MemRefType>(op.getResult().getType());
1381 if (!memrefTy || xegpu::XeGPUDialect::isSharedMemory(memrefTy))
1382 return;
1383 allocas.push_back(op);
1384 });
1385
1386 for (memref::AllocaOp alloca : allocas) {
1387 OpBuilder builder(alloca);
1388 auto memrefTy = cast<MemRefType>(alloca.getResult().getType());
1389 auto newTy = withMemorySpace(memrefTy, slmAttr);
1390 auto newOp = memref::AllocaOp::create(
1391 builder, alloca.getLoc(), newTy, alloca.getDynamicSizes(),
1392 alloca.getSymbolOperands(), alloca.getAlignmentAttr());
1393 alloca.getResult().replaceAllUsesWith(newOp.getResult());
1394 alloca.erase();
1395 // Propagate the new memory space through memref-producing consumers.
1396 for (Operation *user : newOp.getResult().getUsers()) {
1397 if (!isMemrefResultOp(user))
1398 continue;
1399 for (Value result : user->getResults())
1400 propagate(result);
1401 }
1402 }
1403}
1404
1405struct ConvertVectorToXeGPUPass
1406 : public impl::ConvertVectorToXeGPUBase<ConvertVectorToXeGPUPass> {
1407 void runOnOperation() override {
1408 // Promote local allocations to SLM (address space 3) so that
1409 // load_matrix/store_matrix lowerings have well-typed memref operands.
1410 promoteAllocasToSLM(getOperation());
1411
1412 // Undo any flattening of gather/scatter operands, so that the conversion
1413 // below sees the N-D shape the XeGPU layouts are expressed in.
1414 if (failed(unflattenGatherScatter(getOperation())))
1415 return signalPassFailure();
1416
1417 RewritePatternSet patterns(&getContext());
1420 if (failed(applyPatternsGreedily(getOperation(), std::move(patterns))))
1421 return signalPassFailure();
1422 }
1423};
1424
1425} // namespace
1426
1428 RewritePatternSet &patterns) {
1429 patterns
1430 .add<TransferReadLowering, TransferWriteLowering, LoadLowering,
1431 ScatterLowering, GatherLowering, StoreLowering, ContractionLowering>(
1432 patterns.getContext());
1433}
return success()
lhs
ArrayAttr()
b getContext())
static std::optional< VectorShape > vectorShape(Type type)
static Value broadcast(Location loc, Value toBroadcast, unsigned numElements, const TypeConverter &typeConverter, ConversionPatternRewriter &rewriter)
Broadcasts the value to vector with numElements number of elements.
static bool isSharedMemory(MemRefType type)
Return true if this is a shared memory memref type.
Base type for affine expression.
Definition AffineExpr.h:68
A multi-dimensional affine map Affine map's are immutable like Type's, and they are uniqued.
Definition AffineMap.h:46
MLIRContext * getContext() const
bool isMinorIdentity() const
Returns true if this affine map is a minor identity, i.e.
static AffineMap get(MLIRContext *context)
Returns a zero result affine map with no dimensions or symbols: () -> ().
bool isProjectedPermutation(bool allowZeroInResults=false) const
Returns true if the AffineMap represents a subset (i.e.
ArrayRef< AffineExpr > getResults() const
bool isPermutationOfMinorIdentityWithBroadcasting(SmallVectorImpl< unsigned > &permutedDims) const
Return true if this affine map can be converted to a minor identity with broadcast by doing a permute...
unsigned getNumResults() const
unsigned getNumInputs() const
AffineExpr getResult(unsigned idx) const
static AffineMap getPermutationMap(ArrayRef< unsigned > permutation, MLIRContext *context)
Returns an AffineMap representing a permutation.
Attributes are known-constant values of operations.
Definition Attributes.h:25
IntegerType getI64Type()
Definition Builders.cpp:73
IntegerAttr getI64IntegerAttr(int64_t value)
Definition Builders.cpp:120
BoolAttr getBoolAttr(bool value)
Definition Builders.cpp:108
IntegerType getI1Type()
Definition Builders.cpp:61
MLIRContext * getContext() const
Definition Builders.h:56
IndexType getIndexType()
Definition Builders.cpp:59
An attribute that represents a reference to a dense vector or tensor object.
This class defines the main interface for locations in MLIR and acts as a non-nullable wrapper around...
Definition Location.h:76
MLIRContext is the top-level object for a collection of MLIR operations.
Definition MLIRContext.h:63
This class helps build Operations.
Definition Builders.h:210
This class implements the operand iterators for the Operation class.
Definition ValueRange.h:44
Operation is the basic unit of execution within MLIR.
Definition Operation.h:87
OpResult getResult(unsigned idx)
Get the 'idx'th result of this operation.
Definition Operation.h:432
std::enable_if_t< llvm::function_traits< std::decay_t< FnT > >::num_args==1, RetT > walk(FnT &&callback)
Walk the operation by calling the callback for each nested operation (including this one),...
Definition Operation.h:849
user_range getUsers()
Returns a range of all users.
Definition Operation.h:925
MLIRContext * getContext()
Return the context this operation is associated with.
Definition Operation.h:233
A special type of RewriterBase that coordinates the application of a rewrite pattern on the current I...
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.
virtual void replaceOp(Operation *op, ValueRange newValues)
Replace the results of the given (original) operation with the specified list of values (replacements...
virtual void eraseOp(Operation *op)
This method erases an operation that is known to have no uses.
std::enable_if_t<!std::is_convertible< CallbackT, Twine >::value, LogicalResult > notifyMatchFailure(Location loc, CallbackT &&reasonCallback)
Used to notify the listener that the IR failed to be rewritten because of a match failure,...
void modifyOpInPlace(Operation *root, CallableT &&callable)
This method is a utility wrapper around an in-place modification of an operation.
OpTy replaceOpWithNewOp(Operation *op, Args &&...args)
Replace the results of the given (original) op with a new op that is created without verification (re...
Instances of the Type class are uniqued, have an immutable identifier and an optional mutable compone...
Definition Types.h:74
This class represents an instance of an SSA value in the MLIR system, representing a computable value...
Definition Value.h:96
bool use_empty() const
Returns true if this value has no uses.
Definition Value.h:208
Type getType() const
Return the type of this value.
Definition Value.h:105
user_range getUsers() const
Definition Value.h:218
Location getLoc() const
Return the location of this value.
Definition Value.cpp:24
Operation * getDefiningOp() const
If this value is the result of an operation, return the operation that defines it.
Definition Value.cpp:18
static ConstantIndexOp create(OpBuilder &builder, Location location, int64_t value)
Definition ArithOps.cpp:398
detail::InFlightRemark failed(Location loc, RemarkOpts opts)
Report an optimization remark that failed.
Definition Remarks.h:734
const uArch * getUArch(llvm::StringRef archName)
Definition uArchCommon.h:24
int getLargestDivisor(T dim, ArrayRef< T > candidates, ArrayRef< T > candidateMultiples={})
Helper Function to find a proper instruction multiple for the user-supplied sg-level data shape (dive...
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
void populatePrepareVectorToMMAPatterns(RewritePatternSet &patterns, bool useNvGpu=false)
Patterns to transform vector ops into a canonical form to convert to MMA matrix operations.
Type getType(OpFoldResult ofr)
Returns the int type of the integer in ofr.
Definition Utils.cpp:311
AffineMap inverseAndBroadcastProjectedPermutation(AffineMap map)
Return the reverse map of a projected permutation where the projected dimensions are transformed into...
SmallVector< T > applyPermutation(ArrayRef< T > input, ArrayRef< int64_t > permutation)
LogicalResult applyPatternsGreedily(Region &region, const FrozenRewritePatternSet &patterns, GreedyRewriteConfig config=GreedyRewriteConfig(), bool *changed=nullptr)
Rewrite ops in the given region, which must be isolated from above, by repeatedly applying the highes...
bool isMemoryEffectFree(Operation *op)
Returns true if the given operation is free of memory effects.
std::conditional_t< std::is_same_v< Ty, mlir::Type >, mlir::Value, detail::TypedValue< Ty > > TypedValue
If Ty is mlir::Type this will select Value instead of having a wrapper around it.
Definition Value.h:494
void populateVectorToXeGPUConversionPatterns(RewritePatternSet &patterns)
Collect a set of patterns to convert from the vector to XeGPU ops.
llvm::TypeSwitch< T, ResultT > TypeSwitch
Definition LLVM.h:139
OpFoldResult getAsOpFoldResult(Value val)
Given a value, try to extract a constant Attribute.
detail::constant_op_matcher m_Constant()
Matches a constant foldable operation.
Definition Matchers.h:369
AffineExpr getAffineDimExpr(unsigned position, MLIRContext *context)
These free functions allow clients of the API to not use classes in detail.
OpRewritePattern is a wrapper around RewritePattern that allows for matching and rewriting against an...
bool isSupportedInstruction(InstructionKind instr) const
Definition uArchBase.h:122
const Instruction * getInstruction(InstructionKind instKind) const
Definition uArchBase.h:115