MLIR 24.0.0git
XeGPUUtils.cpp
Go to the documentation of this file.
1//===---- XeGPUUtils.cpp - MLIR Utilities for XeGPUOps ------------------===//
2//
3// Part of the MLIR 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 utility methods for working with the XeGPU dialect.
10//
11//===----------------------------------------------------------------------===//
12
22#include "mlir/IR/Builders.h"
23#include "mlir/IR/BuiltinOps.h"
24#include "mlir/IR/Operation.h"
25#include "mlir/IR/ValueRange.h"
29#include "llvm/Support/Casting.h"
30#include "llvm/Support/FormatVariadic.h"
31#include <cstdint>
32#include <numeric>
33
34using namespace mlir;
35
36/// convert ArrayRef<ValueRange> into SmallVector<Value>
39 for (const auto &vals : values)
40 llvm::append_range(result, vals);
41 return result;
42}
43
44FailureOr<VectorType>
45mlir::xegpu::getDistributedVectorType(xegpu::TensorDescType tdescTy) {
46 auto layout = llvm::dyn_cast_if_present<LayoutAttr>(tdescTy.getLayout());
47 // It only works for subgroup level layout, which only has lane_layout
48 // and lane_data, and is to distribute a SIMD code into SIMT code.
49 if (!layout || !layout.isForSubgroup())
50 return failure();
51
52 SmallVector<int64_t> laneData(layout.getLaneData().asArrayRef());
53 SmallVector<int64_t> laneLayout(layout.getLaneLayout().asArrayRef());
54 auto tdescShape = tdescTy.getShape();
55 auto elementType = tdescTy.getElementType();
56
57 // compute sgSize by multiply elements of laneLayout
58 // e.g. for 2D layout, sgSize = laneLayout[0] * laneLayout[1]
59 // e.g. for 1D layout, sgSize = laneLayout[0]
60 int64_t sgSize = llvm::product_of(laneLayout);
61
62 // Check if the tensor descriptor shape is distributable.
63 int64_t tensorSize = 1;
64 for (auto [tdescDim, laneDim, laneDataDim] :
65 llvm::zip_equal(tdescShape, laneLayout, laneData)) {
66 assert((tdescDim % (laneDim * laneDataDim) == 0) &&
67 "tensor descriptor shape is not distributable");
68 tensorSize *= tdescDim;
69 }
70 // tensorSize must be adjusted for array_length.
71 tensorSize *= tdescTy.getArrayLength();
72
73 return VectorType::get({tensorSize / sgSize}, elementType);
74}
75
76FailureOr<VectorType>
77mlir::xegpu::getDistributedVectorType(VectorType originalType,
78 xegpu::LayoutAttr layout) {
79 int64_t rank = originalType.getRank();
80 if (rank < 1)
81 return failure();
82 ArrayRef<int64_t> shape = originalType.getShape();
83 // For rank > 2, leading dimensions are treated as batch/array dimensions.
84 // Drop them and use the product as arrayLength.
85 int arrayLength = 1;
86 while (shape.size() > 2) {
87 arrayLength *= shape[0];
88 shape = shape.drop_front();
89 }
90 // Drop matching leading dims from layout if the layout rank exceeds the
91 // remaining shape rank.
92 auto laneLayout = layout.getEffectiveLaneLayoutAsInt();
93 auto laneData = layout.getEffectiveLaneDataAsInt();
94 while (!laneLayout.empty() && laneLayout.size() > shape.size()) {
95 laneLayout.erase(laneLayout.begin());
96 laneData.erase(laneData.begin());
97 }
98 auto trimmedLayout = xegpu::LayoutAttr::get(
99 layout.getContext(),
100 SmallVector<int32_t>(laneLayout.begin(), laneLayout.end()),
101 SmallVector<int32_t>(laneData.begin(), laneData.end()));
102 auto helperTdescTy = xegpu::TensorDescType::get(
103 shape, originalType.getElementType(), arrayLength,
104 /*boundary_check=*/true,
105 /*memory_space=*/xegpu::MemorySpace::Global, trimmedLayout);
106 return xegpu::getDistributedVectorType(helperTdescTy);
107}
108
109FailureOr<VectorType>
110xegpu::getDistVecTypeBasedOnLaneLayout(xegpu::DistributeLayoutAttr layout,
111 VectorType originalType) {
112 if (!layout)
113 return failure();
114 assert((isa<xegpu::LayoutAttr>(layout) || isa<xegpu::SliceAttr>(layout)) &&
115 "Expecting a valid layout.");
116
117 int64_t vectorRank = originalType.getRank();
118 int64_t layoutRank = layout.getRank();
119 assert(vectorRank >= layoutRank && "Vector rank must be >= layout rank.");
120
121 // When the vector has more dimensions than the layout, only the trailing
122 // dimensions are distributed. Leading dimensions are preserved as-is.
123 int64_t offset = vectorRank - layoutRank;
124 ArrayRef<int64_t> fullShape = originalType.getShape();
125 SmallVector<int64_t> trailingShape(fullShape.begin() + offset,
126 fullShape.end());
127 auto distributedShapeOrFailure =
128 layout.computeDistributedShape(trailingShape);
129 if (failed(distributedShapeOrFailure))
130 return failure();
131
132 SmallVector<int64_t> resultShape(fullShape.begin(),
133 fullShape.begin() + offset);
134 resultShape.append(distributedShapeOrFailure->begin(),
135 distributedShapeOrFailure->end());
136 return VectorType::get(resultShape, originalType.getElementType());
137}
138
139std::string xegpu::getTemporaryLayoutName(const OpOperand &operand) {
140 const StringRef prefix("layout_operand_");
141 unsigned idx = const_cast<OpOperand &>(operand).getOperandNumber();
142 return llvm::formatv("{0}{1}", prefix, idx).str();
143}
144
146 const StringRef prefix = "layout_result_";
147 return llvm::formatv("{0}{1}", prefix, result.getResultNumber()).str();
148}
149
150xegpu::DistributeLayoutAttr xegpu::getDistributeLayoutAttr(const Value value) {
151 if (!value)
152 return nullptr;
153
154 if (auto result = dyn_cast<OpResult>(value)) {
155 Operation *defOp = result.getDefiningOp();
156 assert(defOp && "result must have a defining op");
157
158 if (auto anchorOp = dyn_cast<xegpu::AnchorLayoutInterface>(defOp)) {
159 auto layout = anchorOp.getAnchorLayout();
160 return layout;
161 }
162
163 std::string layoutName = getTemporaryLayoutName(result);
164 if (defOp->hasDiscardableAttr(layoutName)) {
165 auto layout =
166 defOp->getDiscardableAttrOfType<xegpu::DistributeLayoutAttr>(
167 layoutName);
168 return layout;
169 }
170 }
171
172 if (auto arg = dyn_cast<BlockArgument>(value)) {
173 auto *parentOp = arg.getOwner()->getParentOp();
174 auto loop = dyn_cast_if_present<LoopLikeOpInterface>(parentOp);
175 if (loop)
176 if (OpOperand *tiedInit = loop.getTiedLoopInit(arg))
177 return getTemporaryLayout(*tiedInit);
178 // An scf.while "after" argument is tied to no init operand; scf.condition
179 // feeds it. Only a pass-through is supported: the forwarded value must be
180 // the matching "before" argument, whose tied init operand carries the
181 // layout.
182 if (auto whileOp = dyn_cast_if_present<scf::WhileOp>(parentOp);
183 whileOp && arg.getOwner()->getParent() == &whileOp.getAfter()) {
184 Value forwarded = whileOp.getConditionOp().getArgs()[arg.getArgNumber()];
185 if (auto beforeArg = dyn_cast<BlockArgument>(forwarded))
186 if (OpOperand *tiedInit = whileOp.getTiedLoopInit(beforeArg))
187 return getTemporaryLayout(*tiedInit);
188 }
189 }
190
191 if (auto tdescTy =
192 dyn_cast_if_present<xegpu::TensorDescType>(value.getType()))
193 return tdescTy.getLayoutAttr();
194
195 return nullptr;
196}
197xegpu::DistributeLayoutAttr
199 Operation *op = opr.getOwner();
200 unsigned idx = const_cast<OpOperand &>(opr).getOperandNumber();
201
202 if (auto anchorOp = dyn_cast<xegpu::AnchorLayoutInterface>(op)) {
203 if (auto dpasOp = dyn_cast<xegpu::DpasOp>(op)) {
204 if (idx == 0) {
205 return dpasOp.getLayoutAAttr();
206 } else if (idx == 1) {
207 return dpasOp.getLayoutBAttr();
208 } else if (idx == 2) {
209 return dpasOp.getLayoutCdAttr();
210 }
211 }
212 if (auto dpasMxOp = dyn_cast<xegpu::DpasMxOp>(op)) {
213 // DpasMxOp has operands: a, b, optional acc, optional scale_a, optional
214 // scale_b
215 unsigned currentIdx = 0;
216
217 if (idx == currentIdx++)
218 return dpasMxOp.getLayoutAAttr();
219
220 if (idx == currentIdx++)
221 return dpasMxOp.getLayoutBAttr();
222
223 if (dpasMxOp.getAcc())
224 if (idx == currentIdx++)
225 return dpasMxOp.getLayoutCdAttr();
226
227 if (dpasMxOp.getScaleA())
228 if (idx == currentIdx++)
229 return dpasMxOp.getLayoutAScaleAttr();
230
231 if (dpasMxOp.getScaleB())
232 if (idx == currentIdx++)
233 return dpasMxOp.getLayoutBScaleAttr();
234
235 return nullptr;
236 }
237 if (auto convertOp = dyn_cast<xegpu::ConvertLayoutOp>(op)) {
238 return convertOp.getEffectiveInputLayout();
239 }
240 auto layout = anchorOp.getAnchorLayout();
241
242 if (idx == 0)
243 return layout;
244
245 // For StoreNdOp and StoreMatrixOp,
246 // the layout is valid for the first two operands: value and memref/tdesc.
247 if (isa<xegpu::StoreNdOp, xegpu::StoreMatrixOp>(op) && (idx < 2))
248 return layout;
249
250 // For gather/scatter ops the mask and offsets share the value's layout.
251 if (isa<xegpu::StoreScatterOp, xegpu::LoadGatherOp>(op))
252 return layout;
253 }
254
255 std::string layoutName = xegpu::getTemporaryLayoutName(opr);
256 if (op->hasDiscardableAttr(layoutName)) {
257 auto layout =
258 op->getDiscardableAttrOfType<xegpu::DistributeLayoutAttr>(layoutName);
259 return layout;
260 }
261
262 return nullptr;
263}
264
265// Returns the permanent layout attribute for the given result if it's
266// available on the defining op. Otherwise returns the provided layout.
267xegpu::DistributeLayoutAttr
268maybePickPermanentLayout(xegpu::DistributeLayoutAttr layout,
269 const OpResult &result, mlir::Operation *owner,
270 const std::string &name) {
271 xegpu::DistributeLayoutAttr candidate = layout;
272
273 if (auto loadOp = dyn_cast<xegpu::LoadGatherOp>(owner)) {
274 if (auto perm = loadOp.getLayoutAttr())
275 candidate = perm;
276 }
277
278 return candidate;
279}
280
281// Returns the permanent layout attribute for the given operand if it's
282// available on the defining op. Otherwise returns the provided layout.
283xegpu::DistributeLayoutAttr
284maybePickPermanentLayout(xegpu::DistributeLayoutAttr layout,
285 const OpOperand &operand, mlir::Operation *owner,
286 const std::string &name) {
287 xegpu::DistributeLayoutAttr candidate = layout;
288 unsigned idx = const_cast<OpOperand &>(operand).getOperandNumber();
289
290 if (auto storeOp = dyn_cast<xegpu::StoreScatterOp>(owner)) {
291 if (idx == 0) {
292 if (auto perm = storeOp.getLayoutAttr())
293 candidate = perm;
294 }
295 }
296
297 return candidate;
298}
299
300// TODO-LayoutRefactor: Remove this function after replacing use
301// with setTemporaryLayout or setAnchorLayout
303 const mlir::OpResult &result,
304 const mlir::xegpu::DistributeLayoutAttr layout) {
305 Operation *owner = result.getOwner();
306
307 if (auto anchorOp = dyn_cast<xegpu::AnchorLayoutInterface>(owner)) {
308 if (anchorOp.getAnchorLayout() == layout)
309 return;
310 anchorOp.setAnchorLayout(layout);
311 return;
312 }
313
314 std::string name = xegpu::getTemporaryLayoutName(result);
315 if (owner->hasDiscardableAttrOfType<DistributeLayoutAttr>(name)) {
316 return;
317 }
318 if (layout) {
319 owner->setDiscardableAttr(name, layout);
320 }
321}
322
323// TODO-LayoutRefactor: Remove this function after replacing use
324// with setTemporaryLayout or setAnchorLayout
326 const DistributeLayoutAttr layout) {
327 Operation *owner = operand.getOwner();
328 unsigned idx = const_cast<OpOperand &>(operand).getOperandNumber();
329
330 if (!layout) {
331 return;
332 }
333 if (auto anchorOp = dyn_cast<xegpu::AnchorLayoutInterface>(owner)) {
334 if (auto dpasOp = dyn_cast<xegpu::DpasOp>(owner)) {
335 if (idx == 0) {
336 return dpasOp.setLayoutAAttr(layout);
337 } else if (idx == 1) {
338 return dpasOp.setLayoutBAttr(layout);
339 } else if (idx == 2) {
340 return dpasOp.setLayoutCdAttr(layout);
341 }
342 }
343 if (auto convertOp = dyn_cast<xegpu::ConvertLayoutOp>(owner)) {
344 return convertOp.setInputLayoutAttr(layout);
345 }
346
347 // For store operations (StoreScatterOp, StoreNdOp, StoreMatrixOp),
348 // the layout is valid for the first two operands: value and memref/tdesc.
349 // For other operations, the layout applies to the first operand only.
350 if (isa<xegpu::StoreScatterOp, xegpu::StoreNdOp, xegpu::StoreMatrixOp>(
351 owner)) {
352 if (idx < 2) {
353 anchorOp.setAnchorLayout(layout);
354 }
355 } else {
356 if (idx == 0) {
357 anchorOp.setAnchorLayout(layout);
358 }
359 }
360 }
361
362 std::string name = xegpu::getTemporaryLayoutName(operand);
363 if (owner->hasDiscardableAttrOfType<DistributeLayoutAttr>(name)) {
364 return;
365 }
366 if (layout) {
367 owner->setDiscardableAttr(name, layout);
368 }
369}
370
371template <typename T, typename>
372xegpu::DistributeLayoutAttr
373xegpu::getTemporaryLayout(const T &operandOrResult) {
374 Operation *op = operandOrResult.getOwner();
375
376 std::string layoutName = xegpu::getTemporaryLayoutName(operandOrResult);
377 if (op->hasDiscardableAttr(layoutName)) {
378 auto layout =
379 op->getDiscardableAttrOfType<xegpu::DistributeLayoutAttr>(layoutName);
380 return layout;
381 }
382
383 return nullptr;
384}
385
386template xegpu::DistributeLayoutAttr
388template xegpu::DistributeLayoutAttr
390
391template <typename T, typename>
392void xegpu::setTemporaryLayout(const T &operandOrResult,
393 const xegpu::DistributeLayoutAttr layout) {
394 Operation *owner = operandOrResult.getOwner();
395 std::string name = xegpu::getTemporaryLayoutName(operandOrResult);
396 if (owner->hasDiscardableAttrOfType<xegpu::DistributeLayoutAttr>(name)) {
397 return;
398 }
399 if (layout) {
400 owner->setDiscardableAttr(name, layout);
401 }
402}
403
405 const mlir::OpResult &result,
406 const mlir::xegpu::DistributeLayoutAttr layout);
407
409 const mlir::OpOperand &operand,
410 const mlir::xegpu::DistributeLayoutAttr layout);
411
415 auto vecTy = dyn_cast<VectorType>(value.getType());
416 if (!vecTy)
417 return {value};
418
419 ArrayRef<int64_t> srcShape = vecTy.getShape();
420 if (!computeShapeRatio(srcShape, shape))
421 return {value};
422
423 int64_t srcShapeRank = srcShape.size();
424 int64_t targetShapeRank = shape.size();
425
426 SmallVector<int64_t> adjustedTargetShape(srcShape.size());
427 int64_t rankDiff = srcShapeRank - targetShapeRank;
428 std::fill(adjustedTargetShape.begin(), adjustedTargetShape.begin() + rankDiff,
429 1);
430 llvm::copy(shape, adjustedTargetShape.begin() + rankDiff);
431
433 for (SmallVector<int64_t> offsets :
434 StaticTileOffsetRange(srcShape, adjustedTargetShape)) {
435 SmallVector<int64_t> staticStrides(offsets.size(), 1);
436 Value slice = vector::ExtractStridedSliceOp::create(
437 builder, loc, value, offsets, adjustedTargetShape, staticStrides);
438
439 // Reshape to remove leading unit dims if needed
440 if (srcShapeRank > targetShapeRank) {
441 auto targetTy = VectorType::get(shape, vecTy.getElementType());
442 slice = vector::ShapeCastOp::create(builder, loc, targetTy, slice);
443 }
444 result.push_back(slice);
445 }
446
447 return result;
448}
449
451 ValueRange values,
453 VectorType inputTy = dyn_cast<VectorType>(values[0].getType());
454 assert(llvm::all_of(values.getTypes(),
455 [&](Type type) { return type == inputTy; }) &&
456 "values must be of the same VectorType");
457
458 Type elemTy = inputTy.getElementType();
459 ArrayRef<int64_t> tileShape = inputTy.getShape();
460
461 VectorType resultTy = VectorType::get(shape, elemTy);
462 auto zeroAttr = builder.getZeroAttr(elemTy);
463 Value result = arith::ConstantOp::create(
464 builder, loc, resultTy, DenseElementsAttr::get(resultTy, zeroAttr));
465
466 for (auto [src, offsets] :
467 llvm::zip_equal(values, StaticTileOffsetRange(shape, tileShape))) {
468 SmallVector<int64_t> staticStrides(tileShape.size(), 1);
469 result = vector::InsertStridedSliceOp::create(builder, loc, src, result,
470 offsets, staticStrides);
471 }
472 return result;
473}
474
475std::optional<std::string> xegpu::getChipStr(Operation *op) {
476 auto gpuModuleOp = op->getParentOfType<gpu::GPUModuleOp>();
477
478 if (!gpuModuleOp)
479 return std::nullopt;
480
481 auto targetAttrs = gpuModuleOp.getTargets();
482 if (targetAttrs) {
483 for (auto &attr : *targetAttrs) {
484 auto xevmAttr = llvm::dyn_cast<xevm::XeVMTargetAttr>(attr);
485 if (xevmAttr)
486 return xevmAttr.getChip().str();
487 }
488 }
489
490 return std::nullopt;
491}
492
494 int64_t subgroupSize) {
495 auto gpuFunc = op->getParentOfType<gpu::GPUFuncOp>();
496 if (!gpuFunc)
497 return failure();
498 std::optional<ArrayRef<int32_t>> blockSize = gpuFunc.getKnownBlockSize();
499 if (!blockSize)
500 return failure();
501 if (!llvm::all_of(*blockSize, [](int32_t dim) {
502 return dim > 0 && llvm::isPowerOf2_32(dim);
503 }))
504 return failure();
505 int64_t numSubgroups = llvm::product_of(*blockSize) / subgroupSize;
506 if (numSubgroups < 1)
507 return failure();
508 return numSubgroups;
509}
510
511/// Generates element-wise addition ops of two arrays with same length.
513 Location loc,
516 assert(lhs.size() == rhs.size() && "lhs and rhs must have the same size");
518 for (auto [l, r] : llvm::zip_equal(lhs, rhs)) {
519 auto lval = getValueOrCreateConstantIndexOp(builder, loc, l);
520 auto rval = getValueOrCreateConstantIndexOp(builder, loc, r);
521 results.push_back(builder.createOrFold<arith::AddIOp>(loc, lval, rval));
522 }
523 return results;
524}
525
526/// Generates element-wise addition ops of two arrays with automatic alignment.
527/// When the input arrays have different sizes, the shorter array is
528/// right-aligned with the longer array, and the unmatched leading elements from
529/// the longer array are preserved unchanged. This is commonly used for offset
530/// computation where higher-dimensional offsets need to be added to
531/// lower-dimensional adjustments.
532///
533/// Example:
534/// lhs = [l1, l2, l3], rhs = [r1, r2]
535/// Result: [11, l2+r1, l3+r2]
540 // ensure a is longer than b
541 ArrayRef<OpFoldResult> a = lhs.size() >= rhs.size() ? lhs : rhs;
542 ArrayRef<OpFoldResult> b = lhs.size() >= rhs.size() ? rhs : lhs;
543 SmallVector<OpFoldResult> results(a.take_front(a.size() - b.size()));
544 a = a.slice(a.size() - b.size());
545 results.append(addElementwise(builder, loc, a, b));
546 return results;
547}
548
549template <typename T>
551 ArrayRef<T> candidateMultiples) {
552 static_assert(std::is_integral<T>::value, "T must be an integer type");
553 int largest = -1;
554 SmallVector<T> multiples = {1};
555 if (!candidateMultiples.empty())
556 multiples =
557 SmallVector<T>(candidateMultiples.begin(), candidateMultiples.end());
558 for (T candidate : candidates) {
559 for (T multiple : multiples) {
560 int value = static_cast<int>(candidate * multiple);
561 if (value != 0 && dim % value == 0 && value > largest)
562 largest = value;
563 }
564 }
565 return largest;
566}
567
569 vector::CombiningKind kind, uint32_t size) {
570 // First reduce on a single thread to get per lane reduction value.
571 Value laneVal = vector::ReductionOp::create(builder, loc, kind, input);
572 // Parallel reduction using butterfly shuffles.
573 for (uint64_t i = 1; i < size; i <<= 1) {
574 Value shuffled =
575 gpu::ShuffleOp::create(builder, loc, laneVal, i, /** width = **/ size,
576 /** mode = **/ gpu::ShuffleMode::XOR)
577 .getShuffleResult();
578 laneVal = makeArithReduction(builder, loc, kind, laneVal, shuffled);
579 }
580 return laneVal;
581}
582
585 vector::CombiningKind kind,
586 int64_t reductionDim, Location loc,
587 PatternRewriter &rewriter) {
588 VectorType sourceType = src.getType();
589 int64_t sourceRank = sourceType.getRank();
590 // Expecting at least a 2D source vector. Leading dimensions (all except the
591 // last two) must be unit.
592 assert(sourceRank >= 2 && "expected at least a 2D source vector");
593 for (int64_t i = 0; i < sourceRank - 2; ++i)
594 assert(sourceType.getShape()[i] == 1 &&
595 "expected leading dimensions to be unit");
596 int64_t rowIdx = sourceRank - 2;
597 int64_t columnIdx = sourceRank - 1;
598 int64_t sourceH = sourceType.getShape()[rowIdx];
599 int64_t sourceW = sourceType.getShape()[columnIdx];
600 int nSlices = (reductionDim == rowIdx) ? sourceW : sourceH;
601 // Create a constant vector to hold the result of the reduction.
602 TypedAttr zeroAttr = rewriter.getZeroAttr(sourceType.getElementType());
603 Value reductionResult = arith::ConstantOp::create(
604 rewriter, loc, acc.getType(),
605 DenseElementsAttr::get(acc.getType(), zeroAttr));
606 auto srcLayout = xegpu::getTemporaryLayout(dyn_cast<OpResult>(src));
607 auto accLayout = xegpu::getTemporaryLayout(dyn_cast<OpResult>(acc));
608 // Reduction result should have the same layout as the accumulator.
609 xegpu::setTemporaryLayout(cast<OpResult>(reductionResult), accLayout);
610 // For each slice of the source, extract the slice vector, do a reduction
611 // and, insert the reduced value back to the result vector.
612 int64_t accRank = acc.getType().getRank();
613 for (int i = 0; i < nSlices; ++i) {
614 // Build nD offsets, sizes, and strides. Leading unit dims get
615 // offset=0, size=1. The last two dims are set based on reductionDim.
616 SmallVector<int64_t> sliceOffsets(sourceRank, 0);
617 SmallVector<int64_t> sliceSizes(sourceRank, 1);
618 SmallVector<int64_t> strides(sourceRank, 1);
619 if (reductionDim == columnIdx) {
620 sliceOffsets[rowIdx] = i;
621 sliceSizes[columnIdx] = sourceW;
622 } else {
623 sliceOffsets[columnIdx] = i;
624 sliceSizes[rowIdx] = sourceH;
625 }
626
627 vector::ExtractStridedSliceOp extractOp =
628 vector::ExtractStridedSliceOp::create(rewriter, loc, src, sliceOffsets,
629 sliceSizes, strides);
630 // Extract strided slice has the same layout as src.
631 xegpu::setTemporaryLayout(extractOp->getOpResult(0), srcLayout);
632
633 int64_t nSliceElements = extractOp.getResult().getType().getNumElements();
634
635 vector::ShapeCastOp slice = vector::ShapeCastOp::create(
636 rewriter, loc,
637 VectorType::get({nSliceElements}, sourceType.getElementType()),
638 extractOp.getResult());
639
640 // Shape cast output has the same layout as the accumulator. Shape cast
641 // source has the same layout as the original reduction source.
642 xegpu::setTemporaryLayout(slice->getOpOperand(0), srcLayout);
643 xegpu::setTemporaryLayout(slice->getOpResult(0), accLayout);
644 // Extract and reduction results in scalars, so no result layout is needed.
645 // Build multi-dim index into acc (sourceRank-1 dims, i.e. source shape with
646 // the reduction dim removed). Leading unit dims get index 0.
647 SmallVector<int64_t> accIdx(accRank, 0);
648 accIdx[accRank - 1] = i;
649 Value accExtract = vector::ExtractOp::create(rewriter, loc, acc, accIdx);
650 Value reduction = vector::ReductionOp::create(
651 rewriter, loc, kind, slice.getResult(), accExtract);
652 reductionResult = vector::InsertOp::create(rewriter, loc, reduction,
653 reductionResult, accIdx);
654 // Insert op should have the same layout as the accumulator.
655 xegpu::setTemporaryLayout(cast<OpResult>(reductionResult), accLayout);
656 }
657 return reductionResult;
658}
659
662 vector::CombiningKind kind, int64_t reductionDim, int64_t reductionSize,
663 Location loc, PatternRewriter &rewriter) {
664 VectorType sourceType = src.getType();
665 int64_t sourceRank = sourceType.getRank();
666 // Expecting at least a 2D source vector. Leading dimensions (all except the
667 // last two) must be unit.
668 assert(sourceRank >= 2 && "expected at least a 2D source vector");
669 for (int64_t i = 0; i < sourceRank - 2; ++i)
670 assert(sourceType.getShape()[i] == 1 &&
671 "expected leading dimensions to be unit");
672 int64_t rowIdx = sourceRank - 2;
673 int64_t columnIdx = sourceRank - 1;
674 int64_t sourceH = sourceType.getShape()[rowIdx];
675 int64_t sourceW = sourceType.getShape()[columnIdx];
676
677 // Create a constant vector to hold the result of the reduction.
678 TypedAttr zeroAttr = rewriter.getZeroAttr(sourceType.getElementType());
679 Value reductionResult = arith::ConstantOp::create(
680 rewriter, loc, acc.getType(),
681 DenseElementsAttr::get(acc.getType(), zeroAttr));
682
683 // nSlices is the number of reduction operations needed to reduce the entire
684 // source vector. For example, if reductionDim is the row dim, we are
685 // reducing across rows, and each slice is a column. So the number of slices
686 // is the number of columns, which is sourceW.
687 int nSlices = (reductionDim == rowIdx) ? sourceW : sourceH;
688
689 // For each slice of the source, extract the slice vector, do a reduction
690 // and, insert the reduced value back to the result vector.
691 int64_t accRank = acc.getType().getRank();
692 for (int i = 0; i < nSlices; ++i) {
693 // Build nD offsets, sizes, and strides. Leading unit dims get
694 // offset=0, size=1. The last two dims are set based on reductionDim.
695 SmallVector<int64_t> sliceOffsets(sourceRank, 0);
696 SmallVector<int64_t> sliceSizes(sourceRank, 1);
697 SmallVector<int64_t> strides(sourceRank, 1);
698 if (reductionDim == columnIdx) {
699 sliceOffsets[rowIdx] = i;
700 sliceSizes[columnIdx] = sourceW;
701 } else {
702 sliceOffsets[columnIdx] = i;
703 sliceSizes[rowIdx] = sourceH;
704 }
705
706 vector::ExtractStridedSliceOp extractOp =
707 vector::ExtractStridedSliceOp::create(rewriter, loc, src, sliceOffsets,
708 sliceSizes, strides);
709 int64_t nSliceElements = extractOp.getResult().getType().getNumElements();
710 vector::ShapeCastOp slice = vector::ShapeCastOp::create(
711 rewriter, loc,
712 VectorType::get({nSliceElements}, sourceType.getElementType()),
713 extractOp.getResult());
714
715 SmallVector<int64_t> accIdx(accRank, 0);
716 accIdx[accRank - 1] = i;
717 Value accExtract = vector::ExtractOp::create(rewriter, loc, acc, accIdx);
718 Value fullReduce =
719 xegpu::subgroupReduction(loc, rewriter, slice, kind, reductionSize);
720 fullReduce =
721 vector::makeArithReduction(rewriter, loc, kind, fullReduce, accExtract);
722 reductionResult = vector::InsertOp::create(rewriter, loc, fullReduce,
723 reductionResult, accIdx);
724 }
725 return reductionResult;
726}
727
729 Type type,
730 vector::CombiningKind kind) {
731 auto vecTy = dyn_cast<VectorType>(type);
732 Type elemTy = vecTy ? vecTy.getElementType() : type;
733
734 // Helper to create either a splat vector or scalar constant from an attr.
735 auto makeConst = [&](Attribute scalarAttr) -> Value {
736 if (vecTy)
737 return arith::ConstantOp::create(
738 builder, loc, vecTy, DenseElementsAttr::get(vecTy, scalarAttr));
739 return arith::ConstantOp::create(builder, loc, cast<TypedAttr>(scalarAttr));
740 };
741
742 switch (kind) {
743 case vector::CombiningKind::ADD:
744 case vector::CombiningKind::XOR:
745 case vector::CombiningKind::OR:
746 case vector::CombiningKind::MAXUI:
747 return makeConst(builder.getZeroAttr(elemTy));
748
749 case vector::CombiningKind::MUL:
750 case vector::CombiningKind::AND:
751 return makeConst(builder.getOneAttr(elemTy));
752
753 case vector::CombiningKind::MINSI:
754 if (auto intTy = dyn_cast<IntegerType>(elemTy))
755 return makeConst(builder.getIntegerAttr(
756 elemTy, APInt::getSignedMaxValue(intTy.getWidth())));
757 return nullptr;
758
759 case vector::CombiningKind::MINUI:
760 if (auto intTy = dyn_cast<IntegerType>(elemTy))
761 return makeConst(
762 builder.getIntegerAttr(elemTy, APInt::getMaxValue(intTy.getWidth())));
763 return nullptr;
764
765 case vector::CombiningKind::MAXSI:
766 if (auto intTy = dyn_cast<IntegerType>(elemTy))
767 return makeConst(builder.getIntegerAttr(
768 elemTy, APInt::getSignedMinValue(intTy.getWidth())));
769 return nullptr;
770
771 case vector::CombiningKind::MINIMUMF:
772 if (auto floatTy = dyn_cast<FloatType>(elemTy))
773 return makeConst(builder.getFloatAttr(
774 elemTy, APFloat::getInf(floatTy.getFloatSemantics())));
775 return nullptr;
776
777 case vector::CombiningKind::MAXIMUMF:
778 if (auto floatTy = dyn_cast<FloatType>(elemTy))
779 return makeConst(builder.getFloatAttr(
780 elemTy,
781 APFloat::getInf(floatTy.getFloatSemantics(), /*Negative=*/true)));
782 return nullptr;
783
784 case vector::CombiningKind::MINNUMF:
785 case vector::CombiningKind::MINIMUMNUMF:
786 case vector::CombiningKind::MAXNUMF:
787 case vector::CombiningKind::MAXIMUMNUMF:
788 if (auto floatTy = dyn_cast<FloatType>(elemTy))
789 return makeConst(builder.getFloatAttr(
790 elemTy, APFloat::getQNaN(floatTy.getFloatSemantics())));
791 return nullptr;
792 }
793 return nullptr;
794}
795
796/// Explicit instantiations
797template int xegpu::getLargestDivisor<int>(int dim, ArrayRef<int> candidates,
798 ArrayRef<int> candidateMultiples);
799template int
801 ArrayRef<unsigned> candidateMultiples);
802
803std::optional<SmallVector<int64_t>>
805 if (vals.size() < 2)
806 return std::nullopt;
807 if (llvm::any_of(vals.drop_back(2), [](int64_t v) { return v != 1; }))
808 return std::nullopt;
809 return SmallVector<int64_t>(vals.take_back(2));
810}
811
812bool xegpu::requirePacked(const xegpu::DistributeLayoutAttr layout) {
813 if (!layout)
814 return false;
815 auto laneData =
816 getInner2DIfUnitLeadingDims(layout.getEffectiveLaneDataAsInt());
817 return laneData && (*laneData)[0] != 1;
818}
819
820bool xegpu::requireTranspose(const xegpu::DistributeLayoutAttr layout,
821 const xegpu::uArch::uArch *uArch) {
822 // Return false for unsupported targets.
823 // TODO: Add more support or move to target info.
824 if (!isa<xegpu::uArch::Xe2>(uArch) && !isa<xegpu::uArch::Xe3>(uArch))
825 return false;
826 if (!layout)
827 return false;
828 auto laneLayout =
829 getInner2DIfUnitLeadingDims(layout.getEffectiveLaneLayoutAsInt());
830 return laneLayout && (*laneLayout)[0] == uArch->getSubgroupSize() &&
831 (*laneLayout)[1] == 1;
832}
833
834bool xegpu::hasStaticShapeAndStrides(MemRefType type) {
835 if (!type.hasStaticShape())
836 return false;
837 SmallVector<int64_t> strides;
838 int64_t offset;
839 return succeeded(type.getStridesAndOffset(strides, offset)) &&
840 llvm::none_of(strides, ShapedType::isDynamic);
841}
842
843// Check if dst shape is an expansion of src shape by inserting unit dimensions.
844// Returns true if all dimensions in src match corresponding dimensions in dst
845// (after skipping unit dimensions), and populates expandedUnitDims with the
846// indices of the unit dimensions in dst that were added (not present in src).
847// Example: src=[2,3], dst=[1,2,3,1] -> true, expandedUnitDims=[0,3]
849 SmallVector<int64_t> &expandedUnitDims) {
850 // All unit dimensions in dst that don't appear in src are the expanded
851 // unit dimensions
852 size_t srcIdx = 0;
853 for (size_t dstIdx = 0; dstIdx < dst.size(); ++dstIdx)
854 if (srcIdx < src.size() && src[srcIdx] == dst[dstIdx])
855 srcIdx++;
856 else if (dst[dstIdx] == 1)
857 expandedUnitDims.push_back(dstIdx);
858 else
859 return false;
860 return srcIdx == src.size();
861}
862
863// Checks if dst shape is an expansion of src shape where each dimension in src
864// is split into one or more consecutive dimensions in dst whose product equals
865// the original dimension. Populates splitDimGroups with groups of dst indices
866// that correspond to each src dimension. Example: src=[6,4], dst=[2,3,2,2] ->
867// true
870 SmallVector<SmallVector<int64_t>> &splitDimGroups) {
871 // each dim in src can be mapped to one or more dims in dst whose product
872 // equals to the src dim
873 size_t srcIdx = 0;
874 int64_t accumulatedSize = 1;
875 SmallVector<int64_t> currentDstDims;
876
877 splitDimGroups.clear();
878 for (size_t dstIdx = 0; dstIdx < dst.size(); ++dstIdx) {
879 if (srcIdx >= src.size())
880 return false;
881 accumulatedSize *= dst[dstIdx];
882 currentDstDims.push_back(dstIdx);
883
884 if (accumulatedSize == src[srcIdx]) {
885 // Also collect trailing unit dims in destination, if any.
886 // Leading unit dims were implicitly collected.
887 if (srcIdx == src.size() - 1) {
888 while (++dstIdx < dst.size() && dst[dstIdx] == 1)
889 currentDstDims.push_back(dstIdx);
890 }
891 // Record the mapping: srcIdx -> currentDstDims
892 splitDimGroups.push_back(currentDstDims);
893 // move to next src dim
894 srcIdx++;
895 accumulatedSize = 1;
896 currentDstDims.clear();
897 } else if (accumulatedSize > src[srcIdx]) {
898 return false;
899 }
900 }
901 return srcIdx == src.size();
902}
903
904//===----------------------------------------------------------------------===//
905// Context-aware type conversion utilities
906//===----------------------------------------------------------------------===//
907
908// Pre-computes distributed VectorType mappings for every value carried through
909// an SCF region-branch op (scf.while, scf.for, scf.if): block args (iter_args /
910// before-/after-args), op results, and the terminator operands feeding them.
911// These positions share one logical value and must convert identically, so each
912// is derived from a single source -- the layout of the feeding value (loop
913// init, `scf.condition` operand, or `scf.if` result) -- via
914// `getDistributeLayoutAttr(Value)`, and keyed by `Value`. Keying by Value is
915// required because the SCF converters
916// detach/replace the loop body mid-conversion (scf.while detaches before/after
917// blocks -> a detached-arg layout query trips an ilist assertion; scf.for
918// rebuilds the op, which loses the temporary `layout_operand_N` attrs -> the
919// query returns null). Recording results and terminator operands lets a 1:N
920// pass resolve them from the map after stripping the loop op's transient attrs
921// (see XeGPUBlocking).
924 SubShapeAndCountFn getSubShapeAndCount) {
926 // Derive the distributed types from the feeding value's layout (the single
927 // authoritative source) and record them for every value that shares this
928 // loop-carried position.
929 auto recordTypes = [&](Value layoutSrc, ArrayRef<Value> dests) {
930 auto vecTy = dyn_cast<VectorType>(layoutSrc.getType());
931 if (!vecTy)
932 return;
933 auto layout = xegpu::getDistributeLayoutAttr(layoutSrc);
934 if (!layout)
935 return;
936 auto [subShape, count] = getSubShapeAndCount(vecTy, layout);
937 if (count <= 0)
938 return;
939 auto newTy = VectorType::get(subShape, vecTy.getElementType());
940 for (Value dest : dests)
941 loopArgTypes[dest] = SmallVector<Type>(count, newTy);
942 };
943 topLevelOp->walk([&](Operation *op) {
944 if (auto whileOp = dyn_cast<scf::WhileOp>(op)) {
945 // "before" args (and the after-region yield operands that feed them)
946 // correspond to the while `inits` operands.
947 auto yieldOp =
948 cast<scf::YieldOp>(whileOp.getAfterBody()->getTerminator());
949 for (auto [init, beforeArg, yieldVal] :
950 llvm::zip(whileOp.getInits(), whileOp.getBeforeArguments(),
951 yieldOp.getOperands()))
952 recordTypes(init, {beforeArg, yieldVal});
953 // "after" args and the while results correspond to the operands of the
954 // embedded `scf.condition` op (not the `inits`).
955 scf::ConditionOp condOp = whileOp.getConditionOp();
956 for (auto [condArg, afterArg, res] :
957 llvm::zip(condOp.getArgs(), whileOp.getAfterArguments(),
958 whileOp.getResults()))
959 recordTypes(condArg, {afterArg, res});
960 return;
961 }
962 if (auto forOp = dyn_cast<scf::ForOp>(op)) {
963 // Each loop-carried position pairs an init operand with its iter_arg,
964 // its loop result, and the yield operand that feeds the next iteration.
965 auto yieldOp = cast<scf::YieldOp>(forOp.getBody()->getTerminator());
966 for (auto [init, arg, res, yieldVal] :
967 llvm::zip(forOp.getInitArgs(), forOp.getRegionIterArgs(),
968 forOp.getResults(), yieldOp.getOperands()))
969 recordTypes(init, {arg, res, yieldVal});
970 return;
971 }
972 if (auto ifOp = dyn_cast<scf::IfOp>(op)) {
973 // Each result and its then/else yield operands share one position and
974 // must convert identically; derive all from the result's layout.
975 scf::YieldOp thenYield = ifOp.thenYield();
976 scf::YieldOp elseYield = ifOp.elseBlock() ? ifOp.elseYield() : nullptr;
977 for (auto [idx, res] : llvm::enumerate(ifOp.getResults())) {
978 SmallVector<Value> dests{res, thenYield.getOperand(idx)};
979 if (elseYield)
980 dests.push_back(elseYield.getOperand(idx));
981 recordTypes(res, dests);
982 }
983 return;
984 }
985 });
986 return loopArgTypes;
987}
988
990 TypeConverter &converter, SubShapeAndCountFn getSubShapeAndCount,
991 DenseMap<Value, SmallVector<Type>> loopArgTypes) {
992 // Context-aware VectorType conversion (1:1 shape-changing or 1:N). For
993 // SCF loop block arguments (scf.while, scf.for), uses the pre-computed
994 // map. For all other Values, retrieves the layout directly via
995 // getDistributeLayoutAttr.
996 auto loopArgTypeMap = std::make_shared<DenseMap<Value, SmallVector<Type>>>(
997 std::move(loopArgTypes));
998 converter.addConversion(
999 [loopArgTypeMap, getSubShapeAndCount](
1000 Value v,
1001 SmallVectorImpl<Type> &result) -> std::optional<LogicalResult> {
1002 if (!isa<VectorType>(v.getType()))
1003 return std::nullopt;
1004
1005 // Check the pre-computed map first. It covers every value carried
1006 // through an SCF loop (operands, block args, results, yield
1007 // operands), all keyed by Value identity.
1008 auto it = loopArgTypeMap->find(v);
1009 if (it != loopArgTypeMap->end()) {
1010 result.append(it->second.begin(), it->second.end());
1011 return success();
1012 }
1013
1014 // For all other Values, retrieve the layout directly.
1015 auto layout = xegpu::getDistributeLayoutAttr(v);
1016 if (!layout)
1017 return std::nullopt;
1018
1019 auto vecType = cast<VectorType>(v.getType());
1020 auto [subShape, count] = getSubShapeAndCount(vecType, layout);
1021 if (count <= 0)
1022 return std::nullopt;
1023
1024 auto newTy = VectorType::get(subShape, vecType.getElementType());
1025 result.append(count, newTy);
1026 return success();
1027 });
1028}
1029
1031 Operation *root,
1032 const llvm::SmallSetVector<UnrealizedConversionCastOp, 8> &existingCasts) {
1033 // Structural type conversion can generate some redundant
1034 // UnrealizedConversionCastOps to materialize the original type from the
1035 // type converted (sub-tile) type. These are redundant at this point and
1036 // can be eliminated by either folding the cancelling cast chain or, when
1037 // the original and final shapes differ but their element counts match,
1038 // inserting a vector.shape_cast instead.
1039 //
1040 // Example (shape differs but element count matches -> shape_cast):
1041 // %1 = UnrealizedConversionCastOp %0 : vector<16x1xf32>
1042 // to vector<16x16xf32>
1043 // %2 = UnrealizedConversionCastOp %1 : vector<16x16xf32>
1044 // to vector<16xf32>
1045 // becomes:
1046 // %2 = vector.shape_cast %0 : vector<16x1xf32> to vector<16xf32>
1047 //
1048 // For unpaired casts that emulate a pack (1:N) or unpack (N:1) between a
1049 // single large VectorType and N identically-typed smaller VectorTypes,
1050 // lower to vector.extract_strided_slice / vector.insert_strided_slice.
1051 auto hasIdenticalVectorTypes = [](ValueRange values) {
1052 auto types = values.getTypes();
1053 return !types.empty() && llvm::all_of(types, [&](Type type) {
1054 return isa<VectorType>(type) && type == types.front();
1055 });
1056 };
1057 OpBuilder builder(root);
1058 root->walk([&](UnrealizedConversionCastOp op) {
1059 if (existingCasts.contains(op))
1060 return;
1061 // Handle N:1 cast (N >= 1) where all inputs come from a single 1:N cast.
1062 if (op.getNumResults() == 1 && op.getNumOperands() >= 1) {
1063 auto defOp =
1064 op.getInputs()[0].getDefiningOp<UnrealizedConversionCastOp>();
1065 if (defOp && !existingCasts.contains(defOp) &&
1066 defOp.getNumOperands() == 1 &&
1067 defOp.getNumResults() == op.getNumOperands() &&
1068 llvm::all_of(op.getInputs(),
1069 [&](Value v) { return v.getDefiningOp() == defOp; })) {
1070 Value orig = defOp.getInputs()[0];
1071 auto origTy = dyn_cast<VectorType>(orig.getType());
1072 auto resTy = dyn_cast<VectorType>(op.getResult(0).getType());
1073 if (origTy && resTy &&
1074 origTy.getNumElements() == resTy.getNumElements() &&
1075 origTy != resTy) {
1076 builder.setInsertionPoint(op);
1077 auto shapeCast =
1078 vector::ShapeCastOp::create(builder, op.getLoc(), resTy, orig);
1079 op.replaceAllUsesWith(ValueRange{shapeCast.getResult()});
1080 } else {
1081 op.replaceAllUsesWith(ValueRange{orig});
1082 }
1083 return;
1084 }
1085 // Unpaired N:1 cast emulating unpack: stitch inputs into the output
1086 // shape via vector.insert_strided_slice.
1087 auto outputTy = dyn_cast<VectorType>(op.getResult(0).getType());
1088 if (op.getNumOperands() > 1 && outputTy &&
1089 hasIdenticalVectorTypes(op.getInputs())) {
1090 builder.setInsertionPoint(op);
1092 builder, op.getLoc(), op.getInputs(), outputTy.getShape());
1093 op->replaceAllUsesWith(ValueRange(result));
1094 }
1095 return;
1096 }
1097 // Handle 1:N cast where the single input comes from an N:1 cast.
1098 if (op.getNumOperands() == 1 && op.getNumResults() > 1) {
1099 auto defOp =
1100 op.getInputs()[0].getDefiningOp<UnrealizedConversionCastOp>();
1101 if (defOp && !existingCasts.contains(defOp) &&
1102 defOp.getNumResults() == 1 &&
1103 defOp.getNumOperands() == op.getNumResults() &&
1104 llvm::equal(ValueRange(defOp.getInputs()).getTypes(),
1105 op->getResultTypes())) {
1106 op.replaceAllUsesWith(defOp.getInputs());
1107 return;
1108 }
1109 // Unpaired 1:N cast emulating pack: split the input into the output
1110 // tile shape via vector.extract_strided_slice.
1111 auto tileTy = dyn_cast<VectorType>(op.getResult(0).getType());
1112 if (tileTy && hasIdenticalVectorTypes(op.getResults())) {
1113 builder.setInsertionPoint(op);
1115 builder, op.getLoc(), op.getInputs()[0], tileTy.getShape());
1116 op->replaceAllUsesWith(results);
1117 }
1118 return;
1119 }
1120 });
1121
1122 // Erase dead casts iteratively.
1123 bool changed = true;
1124 while (changed) {
1125 changed = false;
1126 root->walk([&](UnrealizedConversionCastOp op) {
1127 if (existingCasts.contains(op))
1128 return;
1129 if (op.use_empty()) {
1130 op.erase();
1131 changed = true;
1132 }
1133 });
1134 }
1135}
1136
1137// Checks if dst shape is a collapse of src shape where each dim in dst is
1138// produced by one or more consecutive dims in src whose product equals the dst
1139// dim. Populates collapseDims with one group per dst dim listing the src
1140// indices collapsed into it. Unit dims in dst that have no backing src dim
1141// (leading, in-between, or trailing) get empty groups; src unit dims that
1142// fall past the last consumed dst dim are absorbed into the most-recent
1143// non-empty group.
1144// Examples:
1145// src=[8,16,32], dst=[1,4096] -> true, collapseDims=[[],[0,1,2]]
1146// src=[8,16,32], dst=[4096,1] -> true, collapseDims=[[0,1,2],[]]
1147// src=[2,3,4], dst=[6,4] -> true, collapseDims=[[0,1],[2]]
1148// src=[64], dst=[64] -> true, collapseDims=[[0]]
1150 SmallVector<SmallVector<int64_t>> &collapseDims) {
1151 collapseDims.clear();
1152 collapseDims.resize(dst.size());
1153
1154 // Cheap precondition: src and dst must describe the same number of
1155 // elements. Bails out early on mismatched shapes without walking the dims.
1156 int64_t srcProd = std::accumulate(src.begin(), src.end(), int64_t{1},
1157 std::multiplies<int64_t>());
1158 int64_t dstProd = std::accumulate(dst.begin(), dst.end(), int64_t{1},
1159 std::multiplies<int64_t>());
1160 if (srcProd != dstProd)
1161 return false;
1162
1163 // Step 1: validate the partition on the unit-dim-stripped (compact) shapes.
1164 // Unit dims play no role in the matching decision — they only need to be
1165 // placed somewhere in the final groups (handled in step 2).
1166 SmallVector<int64_t> srcCompact, dstCompact;
1167 for (int64_t s : src)
1168 if (s != 1)
1169 srcCompact.push_back(s);
1170 for (int64_t d : dst)
1171 if (d != 1)
1172 dstCompact.push_back(d);
1173
1174 size_t s = 0;
1175 for (int64_t need : dstCompact) {
1176 int64_t acc = 1;
1177 while (s < srcCompact.size() && acc < need)
1178 acc *= srcCompact[s++];
1179 if (acc != need)
1180 return false;
1181 }
1182 if (s != srcCompact.size())
1183 return false;
1184
1185 // Step 2: assign each original src index to the correct original dst group.
1186 // Walk dst in original order, advancing past unit dst dims (they keep their
1187 // pre-initialized empty group). Walk src in original order; non-unit src
1188 // dims accumulate into the current dst group, unit src dims attach to the
1189 // current group when one is open or to the most-recent non-empty group
1190 // after dst is exhausted (leading unit src dims with no group yet are
1191 // dropped).
1192 size_t dstIdx = 0;
1193 while (dstIdx < dst.size() && dst[dstIdx] == 1)
1194 dstIdx++;
1195
1196 int64_t lastNonEmpty = -1;
1197 int64_t acc = 1;
1198 for (size_t srcIdx = 0; srcIdx < src.size(); ++srcIdx) {
1199 if (dstIdx >= dst.size()) {
1200 // dst exhausted; remaining src dims are unit (validated above) and
1201 // attach to the last non-empty group, if any.
1202 if (lastNonEmpty >= 0)
1203 collapseDims[lastNonEmpty].push_back(srcIdx);
1204 continue;
1205 }
1206 acc *= src[srcIdx];
1207 collapseDims[dstIdx].push_back(srcIdx);
1208 lastNonEmpty = dstIdx;
1209 if (acc == dst[dstIdx]) {
1210 acc = 1;
1211 ++dstIdx;
1212 while (dstIdx < dst.size() && dst[dstIdx] == 1)
1213 ++dstIdx;
1214 }
1215 }
1216 return true;
1217}
return success()
b
Return true if permutation is a valid permutation of the outer_dims_perm (case OuterOrInnerPerm::Oute...
xegpu::DistributeLayoutAttr maybePickPermanentLayout(xegpu::DistributeLayoutAttr layout, const OpResult &result, mlir::Operation *owner, const std::string &name)
Attributes are known-constant values of operations.
Definition Attributes.h:25
IntegerAttr getIntegerAttr(Type type, int64_t value)
Definition Builders.cpp:237
FloatAttr getFloatAttr(Type type, double value)
Definition Builders.cpp:263
TypedAttr getZeroAttr(Type type)
Definition Builders.cpp:333
TypedAttr getOneAttr(Type type)
Definition Builders.cpp:351
static DenseElementsAttr get(ShapedType type, ArrayRef< Attribute > values)
Constructs a dense elements attribute from an array of element values.
This class defines the main interface for locations in MLIR and acts as a non-nullable wrapper around...
Definition Location.h:76
This class helps build Operations.
Definition Builders.h:210
void setInsertionPoint(Block *block, Block::iterator insertPoint)
Set the insertion point to the specified location.
Definition Builders.h:401
void createOrFold(SmallVectorImpl< Value > &results, Location location, Args &&...args)
Create an operation of specific op type at the current insertion point, and immediately try to fold i...
Definition Builders.h:528
This class represents an operand of an operation.
Definition Value.h:254
This is a value defined by a result of an operation.
Definition Value.h:454
Operation is the basic unit of execution within MLIR.
Definition Operation.h:87
bool hasDiscardableAttrOfType(NameT &&name)
Definition Operation.h:506
void setDiscardableAttr(StringAttr name, Attribute value)
Set a discardable attribute by name.
Definition Operation.h:512
OpTy getParentOfType()
Return the closest surrounding parent operation that is of type 'OpTy'.
Definition Operation.h:255
bool hasDiscardableAttr(StringRef name)
Return true if this operation has a discardable attribute with the provided name.
Definition Operation.h:503
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
AttrClass getDiscardableAttrOfType(StringRef name)
Access a discardable attribute by name and cast it to AttrClass.
Definition Operation.h:493
A special type of RewriterBase that coordinates the application of a rewrite pattern on the current I...
A range-style iterator that allows for iterating over the offsets of all potential tiles of size tile...
Instances of the Type class are uniqued, have an immutable identifier and an optional mutable compone...
Definition Types.h:74
This class provides an abstraction over the different types of ranges over Values.
Definition ValueRange.h:389
type_range getTypes() const
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 * getOwner() const
Return the owner of this operand.
Definition UseDefLists.h:38
detail::InFlightRemark failed(Location loc, RemarkOpts opts)
Report an optimization remark that failed.
Definition Remarks.h:734
Value makeArithReduction(OpBuilder &b, Location loc, CombiningKind kind, Value v1, Value acc, arith::FastMathFlagsAttr fastmath=nullptr, Value mask=nullptr)
Returns the result value of reducing two scalar/vector values with the corresponding arith operation.
bool matchDimCollapse(ArrayRef< int64_t > src, ArrayRef< int64_t > dst, SmallVector< SmallVector< int64_t > > &collapseDims)
Value createVectorWithShapeFromValues(OpBuilder &builder, Location loc, ValueRange values, ArrayRef< int64_t > shape)
Create a vector of shape from a set of values using vector.insert_stride_slice.
bool requirePacked(const DistributeLayoutAttr layout)
Helper function to check if the layout is packed.
void setTemporaryLayout(const T &operandOrResult, const DistributeLayoutAttr layout)
Value createReductionNeutralValue(OpBuilder &builder, Location loc, Type type, vector::CombiningKind kind)
Creates a constant filled with the neutral (identity) value for the given reduction kind.
void setDistributeLayoutAttr(const OpResult &Result, const DistributeLayoutAttr layout)
[to-be-deprecated] Sets the DistributeLayoutAttr for a given OpResult user should use setAnchorLayout...
Value subgroupReduction(Location loc, OpBuilder &builder, Value input, vector::CombiningKind kind, uint32_t size)
Given an input value representing per-lane data, this function returns the result after performing a ...
bool matchUnitDimExpansion(ArrayRef< int64_t > src, ArrayRef< int64_t > dst, SmallVector< int64_t > &expandedUnitDims)
std::optional< SmallVector< int64_t > > getInner2DIfUnitLeadingDims(ArrayRef< int64_t > vals)
Returns the innermost 2 entries of vals if it is at least 2D and all of its leading entries are unit;...
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...
FailureOr< int64_t > getNumSubgroupsFromBlockSize(Operation *op, int64_t subgroupSize)
Returns the number of subgroups the kernel enclosing op runs, derived from the known_block_size of it...
bool hasStaticShapeAndStrides(MemRefType type)
Returns true if type has a static shape and static strides.
FailureOr< VectorType > getDistVecTypeBasedOnLaneLayout(DistributeLayoutAttr layout, VectorType originalType)
Helper function to get distributed vector type for a source vector type according to the lane_layout.
Value lowerToVectorReductions(TypedValue< VectorType > src, TypedValue< VectorType > acc, vector::CombiningKind kind, int64_t reductionDim, Location loc, PatternRewriter &rewriter)
Given a src and an acc argumments from a vector::MultiDimReductionOp, lower to a set of vector::Reduc...
bool requireTranspose(const DistributeLayoutAttr layout, const uArch::uArch *uArch)
Helper function to check if the layout requires a transpose effect.
bool matchSplitDimExpansion(ArrayRef< int64_t > src, ArrayRef< int64_t > dst, SmallVector< SmallVector< int64_t > > &splitDimGroups)
DistributeLayoutAttr getDistributeLayoutAttr(const Value value)
Retrieves the DistributeLayoutAttr associated with a given Value, or nullptr if none is found.
DenseMap< Value, SmallVector< Type > > precomputeLoopBlockArgTypes(Operation *topLevelOp, SubShapeAndCountFn getSubShapeAndCount)
Pre-computes distributed VectorType mappings for every value carried through an SCF loop under topLev...
std::string getTemporaryLayoutName(const OpOperand &operand)
Return the attribute name for the OpOperand to attach DistributeLayoutAttr.
std::optional< std::string > getChipStr(Operation *op)
Retrieves the chip string from the XeVM target attribute of the parent GPU module operation.
void addVectorTypeConversion(TypeConverter &converter, SubShapeAndCountFn getSubShapeAndCount, DenseMap< Value, SmallVector< Type > > loopArgTypes)
Adds a context-aware VectorType conversion to converter (1:1 shape-changing or 1:N,...
SmallVector< Value > extractVectorsWithShapeFromValue(OpBuilder &builder, Location loc, Value value, ArrayRef< int64_t > shape)
Extract a set of small vectors from a value with a given shape using vector.extract_stride_slice.
DistributeLayoutAttr getTemporaryLayout(const T &operandOrResult)
get and set distribute layout attribute for non-anchor operations (and offsets/masks of load/store op...
Value lowerCrossLaneReductionToShuffles(TypedValue< VectorType > src, TypedValue< VectorType > acc, vector::CombiningKind kind, int64_t reductionDim, int64_t reductionSize, Location loc, PatternRewriter &rewriter)
Lowers cross-lane reductions to shuffle operations on a 2D vector.
std::function< std::pair< SmallVector< int64_t >, int >( VectorType, DistributeLayoutAttr)> SubShapeAndCountFn
Callback type for computing sub-shape and count for 1:N (or 1:1 shape-changing) VectorType conversion...
Definition XeGPUUtils.h:260
void cleanupUnrealizedConversionCasts(Operation *root, const llvm::SmallSetVector< UnrealizedConversionCastOp, 8 > &existingCasts)
Cleans up UnrealizedConversionCastOps inserted during SCF structural type conversion and/or XeGPU unr...
SmallVector< Value > flattenValues(ArrayRef< ValueRange > values)
Flatten a set of ValueRange into a single SmallVector<Value>
SmallVector< OpFoldResult > addWithRightAligned(OpBuilder &builder, Location loc, ArrayRef< OpFoldResult > lhs, ArrayRef< OpFoldResult > rhs)
Generates element-wise addition ops of two arrays with automatic alignment.
SmallVector< OpFoldResult > addElementwise(OpBuilder &builder, Location loc, ArrayRef< OpFoldResult > lhs, ArrayRef< OpFoldResult > rhs)
Generates element-wise addition ops of two arrays with same length.
FailureOr< VectorType > getDistributedVectorType(xegpu::TensorDescType tdescTy)
If tensor descriptor has a layout attribute it is used in SIMT mode.
Include the generated interface declarations.
Type getType(OpFoldResult ofr)
Returns the int type of the integer in ofr.
Definition Utils.cpp:311
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
Value getValueOrCreateConstantIndexOp(OpBuilder &b, Location loc, OpFoldResult ofr)
Converts an OpFoldResult to a Value.
Definition Utils.cpp:114
llvm::DenseMap< KeyT, ValueT, KeyInfoT, BucketT > DenseMap
Definition LLVM.h:120
std::optional< SmallVector< int64_t > > computeShapeRatio(ArrayRef< int64_t > shape, ArrayRef< int64_t > subShape)
Return the multi-dimensional integral ratio of subShape to the trailing dimensions of shape.
virtual int getSubgroupSize() const =0