MLIR 24.0.0git
XeGPULayoutImpl.cpp
Go to the documentation of this file.
1//===---- XeGPULayoutImpl.cpp - MLIR Utilities for XeGPUOps
2//------------------===//
3//
4// Part of the MLIR Project, under the Apache License v2.0 with LLVM Exceptions.
5// See https://llvm.org/LICENSE.txt for license information.
6// SPDX-License-Identifier: Apache-2.0 WITH LLVM-exception
7//
8//===----------------------------------------------------------------------===//
9//
10// This file implements layout utility functions for XeGPU dialect
11// transformation.
12//
13//===----------------------------------------------------------------------===//
14
23#include "mlir/IR/Builders.h"
24#include "mlir/IR/Operation.h"
25#include "mlir/IR/ValueRange.h"
30#include "llvm/ADT/PostOrderIterator.h"
31#include "llvm/Support/FormatVariadic.h"
32#include <cstdint>
33#include <numeric>
34
35using namespace mlir;
36
40 out.reserve(attrs.size());
41
42 for (auto attr : attrs) {
43 if (auto dist = dyn_cast<xegpu::DistributeLayoutAttr>(attr.getValue())) {
44 auto newLayout = dist.dropSgLayoutAndData();
45 if (newLayout)
46 out.emplace_back(attr.getName(), newLayout);
47 } else {
48 out.push_back(attr);
49 }
50 }
51
52 return out;
53}
54
58 out.reserve(attrs.size());
59
60 for (auto attr : attrs) {
61 if (auto dist = dyn_cast<xegpu::DistributeLayoutAttr>(attr.getValue())) {
62 auto newLayout = dist.dropInstData();
63 if (newLayout)
64 out.emplace_back(attr.getName(), newLayout);
65 } else {
66 out.push_back(attr);
67 }
68 }
69
70 return out;
71}
72
74 op->getName().walkInherentAttrs(op, [](StringRef, Attribute &attr) {
75 if (auto dist = dyn_cast<xegpu::DistributeLayoutAttr>(attr))
76 attr = dist.dropInstData();
77 });
78}
79
80// Sets the layout on a TensorDesc value by updating its type to include
81// the given layout, if the type does not already have a layout attached.
82static void setTensorDescLayout(Value val, xegpu::DistributeLayoutAttr layout) {
83 auto tensorDescTy = dyn_cast<xegpu::TensorDescType>(val.getType());
84 if (!tensorDescTy || tensorDescTy.getLayoutAttr())
85 return;
86 auto typeWithLayout = xegpu::TensorDescType::get(
87 tensorDescTy.getContext(), tensorDescTy.getShape(),
88 tensorDescTy.getElementType(), tensorDescTy.getEncoding(), layout);
89 val.setType(typeWithLayout);
90}
91
92// the walkRegionBackward() is a recursive function
93// the input rootOp is the function operation, which is also a region op.
94// it recursively processes the region op in reverse topological order.
95static void walkRegionBackward(Region &region,
97
98 // Use post-order traversal to process blocks in reverse topological order.
99 // This ensures that use blocks are visited before def blocks, which is
100 // required for backward layout propagation.
101 if (region.empty())
102 return;
103 llvm::ReversePostOrderTraversal<Region *> rpot(&region);
104 SmallVector<Block *> blocks(rpot.begin(), rpot.end());
105 for (Block *block : llvm::reverse(blocks)) {
106 // ops: back -> front
107 for (Operation &op : llvm::reverse(*block)) {
108 // make sure we first visit inside the region op (so yield op first)
109 // and then move to region op itself
110 // Regions are iterated in forward order so that for multi-region ops
111 // like scf.while, earlier regions (e.g., "before/cond") are processed
112 // first. This ensures that when a later region's terminator (e.g., "do"
113 // yield) needs the layout of an earlier region's block args, those
114 // layouts are already available from use points.
115 for (Region &nested : op.getRegions())
116 walkRegionBackward(nested, visit);
117
118 visit(&op);
119 }
120 }
121}
122
123static xegpu::DistributeLayoutAttr getLayoutFromUsePoints(Value result) {
124 xegpu::DistributeLayoutAttr layout = nullptr;
125 for (OpOperand &use : result.getUses()) {
126 if (auto tmpLayout = xegpu::getDistributeLayoutAttr(use)) {
127 if (!layout)
128 layout = tmpLayout;
129 break;
130 }
131 }
132 return layout;
133}
134
135// Returns true if `op` is safe and cheap to clone (no side effects, no
136// regions, and all operands are themselves trivially rematerializable, e.g.
137// block-arg-free pure value generators such as `vector.step`, splat
138// `arith.constant`, or `vector.create_mask` whose operands are constants).
140 if (!op || op->getNumRegions() != 0)
141 return false;
142 if (!isMemoryEffectFree(op))
143 return false;
144 for (Value v : op->getOperands()) {
145 Operation *defOp = v.getDefiningOp();
146 if (!defOp)
147 return false;
148 if (!isTriviallyRematerializable(defOp))
149 return false;
150 }
151 return true;
152}
153
154// For regular operations: First the result layouts are propagated from uses.
155// Then the result layouts are propagated to uses (operands).
157 if (op->getNumResults() == 0)
158 return;
159 if (op->getNumResults() > 1 && !isa<vector::DeinterleaveOp>(op))
160 return;
161 OpResult result = op->getResult(0);
162 xegpu::DistributeLayoutAttr resLayout = getLayoutFromUsePoints(result);
163 Type resultType = result.getType();
164
165 if (!resLayout)
166 return;
167
168 // Recover layout for TensorDesc type results by updating the type to include
169 // the layout. For vector type
170 if (isa<xegpu::TensorDescType>(resultType))
171 setTensorDescLayout(result, resLayout);
172
173 // Recover layout for vector type results, or for multi-reduction ops which
174 // may reduce to a scalar that still needs a layout.
175 if (isa<VectorType>(resultType) || isa<vector::MultiDimReductionOp>(op))
177
178 if (isa<vector::DeinterleaveOp>(op))
179 xegpu::setTemporaryLayout(op->getResult(1), resLayout);
180
181 for (OpOperand &opr : op->getOpOperands()) {
182 xegpu::DistributeLayoutAttr operandLayout =
184 if (isa<VectorType>(opr.get().getType()) && operandLayout)
185 xegpu::setTemporaryLayout(opr, operandLayout);
186 }
187}
188
189// Propagate layout from region op results and sibling region block args
190// to yield/condition operands. For each successor of this terminator:
191// - Parent successor: propagate from parent op's result layouts (use points).
192// - Region successor: propagate from target region's block arg layouts (use
193// points), e.g., scf.yield in "after/do" region propagates to "before/cond"
194// block args.
196 mlir::RegionBranchTerminatorOpInterface yieldOp) {
197 auto regionBranchOp =
198 dyn_cast<RegionBranchOpInterface>(yieldOp->getParentOp());
199 if (!regionBranchOp)
200 return;
201
203 SmallVector<Attribute> operandAttrs(yieldOp->getNumOperands(), nullptr);
204 yieldOp.getSuccessorRegions(operandAttrs, successors);
205
206 for (const RegionSuccessor &successor : successors) {
207 OperandRange succOps = yieldOp.getSuccessorOperands(successor);
208 if (succOps.empty())
209 continue;
210 unsigned beginIdx = succOps.getBeginOperandIndex();
211 ValueRange successorInputs = regionBranchOp.getSuccessorInputs(successor);
212 unsigned count = std::min<unsigned>(succOps.size(), successorInputs.size());
213
214 for (unsigned i = 0; i < count; ++i) {
215 xegpu::DistributeLayoutAttr layout;
216 if (successor.isOperation()) {
217 // For parent successor, get layout from external use points of the
218 // parent op's results.
219 auto regionResult = regionBranchOp->getResult(i);
220 layout = getLayoutFromUsePoints(regionResult);
221 if (layout) {
222 // set layout for the region op, like scf.loop
223 xegpu::setTemporaryLayout(regionResult, layout);
224 if (isa<xegpu::TensorDescType>(regionResult.getType()))
225 setTensorDescLayout(regionResult, layout);
226 }
227 } else {
228 // For region successor, get layout from the target region's block
229 // arg use points (e.g., "before/cond" region args for scf.while
230 // "after/do" yield).
231 layout = getLayoutFromUsePoints(successorInputs[i]);
232 }
233 if (!layout)
234 continue;
235 auto operandType = succOps[i].getType();
236 if (isa<VectorType>(operandType) ||
237 dyn_cast<xegpu::TensorDescType>(operandType))
238 // recover layout for yield op operands
239 xegpu::setTemporaryLayout(yieldOp->getOpOperand(beginIdx + i), layout);
240 }
241 }
242}
243
244/// Assign a layout to a region op's results (e.g. scf.for) using the layout of
245/// the terminator operands that the region forwards to them. For each operand a
246/// terminator (e.g. scf.yield) forwards to a successor input, if that input is
247/// a region op result, the operand's layout is written onto the result.
248/// clang-format off
249/// Example: scf.for ... iter_args(...) -> (out types) {
250/// ...
251/// scf.yield ... : (yield types)
252/// }
253/// clang-format on
254/// Having a layout on the region op result lets a later step attach a
255/// convert_layout as a use to resolve the region op's no-use case.
256/// Block-argument successors are left untouched.
258 mlir::RegionBranchTerminatorOpInterface terminator,
259 xegpu::GetLayoutFnTy getLayoutOfValue) {
260 // Only process if the terminator is inside a region branch op.
261 auto branchOp = dyn_cast<RegionBranchOpInterface>(terminator->getParentOp());
262 if (!branchOp)
263 return success();
264
266 branchOp.getSuccessorOperandInputMapping(mapping,
267 RegionBranchPoint(terminator));
268 for (const auto &[successorOperand, successorInputs] : mapping) {
269 for (Value successorInput : successorInputs) {
270 Type inputType = successorInput.getType();
271 // We only need to operate on vector types.
272 if (!isa<VectorType>(inputType))
273 continue;
274 xegpu::DistributeLayoutAttr successorOperandLayout =
275 getLayoutOfValue(successorOperand->get());
276
277 // The forwarded operand must carry a layout to propagate.
278 if (!successorOperandLayout)
279 return failure();
280 // Assign the yield operand's layout to the region op result it feeds.
281 if (auto result = dyn_cast<OpResult>(successorInput))
282 xegpu::setDistributeLayoutAttr(result, successorOperandLayout);
283 // Restrict the input IR: a successor argument that is not tied to an init
284 // operand (scf.while's "after" arguments) must be fed by a pass-through,
285 // because nothing else identifies which value the region carries. Its
286 // layout is then that of the forwarded argument's init operand.
287 if (auto arg = dyn_cast<BlockArgument>(successorInput)) {
288 auto loop =
289 dyn_cast<LoopLikeOpInterface>(arg.getOwner()->getParentOp());
290 bool tiedToInit = loop && loop.getTiedLoopInit(arg);
291 if (!tiedToInit && !isa<BlockArgument>(successorOperand->get()))
292 return terminator->emitError(
293 "unsupported region structure: the successor argument it feeds "
294 "is not tied to an init operand, so its value must be passed "
295 "through from predecessor region argument.");
296 }
297 }
298 }
299 return success();
300}
301
302// Propagate layout from region arguments to region op's init operands. This
303// sets the temporary layout for region arguments and init operands.
304LogicalResult
305xegpu::propagateRegionArgsToInits(mlir::RegionBranchOpInterface regionOp,
306 xegpu::GetLayoutFnTy getLayoutOfValue) {
307 // Iterate all regions of the region op. For each block argument that has a
308 // layout (obtained via `getLayoutOfValue`), trace back to find the
309 // corresponding init operand of the regionOp and set the layout on it.
310 // This works generically for scf.for, scf.while, and other
311 // RegionBranchOpInterface ops.
312 for (Region &region : regionOp->getRegions()) {
313 RegionSuccessor regionSuccessor(&region);
314 // Use getSuccessorInputs to get the block arguments that correspond to
315 // predecessor operands. This correctly handles ops like scf.for where
316 // the induction variable is a block arg but not a successor input.
317 ValueRange successorInputs = regionOp.getSuccessorInputs(regionSuccessor);
318 for (auto [inputIdx, regionArg] : llvm::enumerate(successorInputs)) {
319 auto layout = getLayoutOfValue(regionArg);
320 if (!layout)
321 continue;
322
323 // Recover layout for tensor_desc block args by updating the type.
324 if (isa<xegpu::TensorDescType>(regionArg.getType()))
325 setTensorDescLayout(regionArg, layout);
326
327 // Recover layout for region op operands, like scf.for's init operands.
328 // Find all predecessor values that flow into this block argument.
329 SmallVector<Value> predValues;
330 regionOp.getPredecessorValues(regionSuccessor, inputIdx, predValues);
331 for (Value predVal : predValues) {
332 // Match predecessor value to an operand of the regionOp.
333 for (OpOperand &operand : regionOp->getOpOperands()) {
334 if (operand.get() == predVal)
335 xegpu::setTemporaryLayout(operand, layout);
336 }
337 }
338 }
339 }
340 return success();
341}
342
343// Prerequisite for Layout Recovery
344// It relies on the following invariant:
345// 1. there is no layout conflict between different uses of the same definition.
346// 2. each definition has a well-defined layout requirement at its use point.
347// - Every definition must have at least one use that appears after it in
348// topological order.
349// - TODO: If a definition has no such use (e.g., a loop result or region
350// output), an explicit convert_layout operation is inserted to create a
351// use.
352// - Only the result of convert_layout is permitted to have no subsequent
353// use.
354//
355// The recovery proceeds by scanning the operation in reverse topological order
356// as follows:
357// For regular operations: First the result layouts are propagated from uses.
358// Then the result layouts are propagated to operands.
359//
360// For region operations (e.g., loops):
361// - When backward propagation reaches a region op, it sets the layout of
362// the region op’s results according to use points like regular ops.
363// - Then, the result layouts (such as a loop output) are propagated to
364// their corresponding operands in the yield.
365// - When backward propagation reaches the first operation inside the
366// region, the pass examines the region op’s initialization list,
367// propagating from region arguments to the corresponding initialization
368// operands.
369// - This ensures that layouts are consistently propagated
370// across region boundaries while preserving a single well-defined use for
371// each definition at the region-op level.
373 auto processFunc = [&](Region &body, StringRef funcName) {
374 walkRegionBackward(body, [&](Operation *op) {
375 if (auto regionOp = dyn_cast<mlir::RegionBranchOpInterface>(op)) {
378 } else if (auto yieldOp =
379 dyn_cast<mlir::RegionBranchTerminatorOpInterface>(op)) {
381 } else if (!dyn_cast<xegpu::AnchorLayoutInterface>(op)) {
383 }
384 });
385 };
387 rootOp->walk([&](func::FuncOp func) {
388 processFunc(func.getBody(), func.getSymName());
389 });
390 rootOp->walk([&](gpu::GPUFuncOp func) {
391 processFunc(func.getBody(), func.getName());
392 });
393
394 return true;
395}
396
397template <typename T, typename>
398void xegpu::removeLayoutAttr(const T &operandOrResult) {
399 Operation *owner = operandOrResult.getOwner();
400 std::string name = xegpu::getTemporaryLayoutName(operandOrResult);
401 if (owner->hasDiscardableAttrOfType<DistributeLayoutAttr>(name))
402 owner->removeDiscardableAttr(name);
403}
404
405// Explicit instantiation for OpResult
406template void
408
409// Explicit instantiation for OpOperand
410template void
412
414 op->walk([&](Operation *nestOp) {
415 // Remove all attributes of DistributeLayoutAttr type
416 SmallVector<StringAttr> attrsToRemove;
417 for (auto namedAttr : nestOp->getDiscardableAttrDictionary().getValue()) {
418 if (isa<DistributeLayoutAttr>(namedAttr.getValue()))
419 attrsToRemove.push_back(namedAttr.getName());
420 }
421 for (auto attrName : attrsToRemove)
422 nestOp->removeDiscardableAttr(attrName);
423 });
424}
425
427 op->walk([&](Operation *nestOp) {
428 SmallVector<StringAttr> attrsToRemove;
429 for (auto namedAttr : nestOp->getDiscardableAttrDictionary().getValue()) {
430 if (isa<xegpu::DistributeLayoutAttr>(namedAttr.getValue()))
431 attrsToRemove.push_back(namedAttr.getName());
432 }
433 for (auto attrName : attrsToRemove)
434 nestOp->removeDiscardableAttr(attrName);
435 });
436}
437
438/// Returns true if every dimension of `shape` except the innermost
439/// `numInnerDims` is a unit (size-1) dimension.
440[[maybe_unused]] static bool leadingDimsAreUnit(ArrayRef<int64_t> shape,
441 int numInnerDims) {
442 int numLeading = static_cast<int>(shape.size()) - numInnerDims;
443 if (numLeading <= 0)
444 return true;
445 return llvm::all_of(shape.take_front(numLeading),
446 [](int64_t dim) { return dim == 1; });
447}
448
449static xegpu::LayoutAttr buildInstDataLayoutWithLane(
450 mlir::MLIRContext *context, ArrayRef<int64_t> instData,
451 ArrayRef<int64_t> laneLayout, ArrayRef<int64_t> laneData,
452 DenseI32ArrayAttr orderAttr = nullptr) {
453 auto toI32Attr = [&](auto range) {
454 SmallVector<int32_t> v(range.begin(), range.end());
455 return DenseI32ArrayAttr::get(context, v);
456 };
457 return xegpu::LayoutAttr::get(context, /*sg_layout=*/nullptr,
458 /*sg_data=*/nullptr, toI32Attr(instData),
459 toI32Attr(laneLayout), toI32Attr(laneData),
460 orderAttr);
461}
462
464 ArrayRef<int64_t> laneLayout,
465 ArrayRef<int64_t> laneData) {
466 return !llvm::any_of(llvm::seq<int>(0, dataShape.size()), [&](int dim) {
467 return dataShape[dim] % (laneLayout[dim] * laneData[dim]) != 0;
468 });
469}
470
471static xegpu::LayoutAttr
473 ArrayRef<int64_t> laneData,
474 DenseI32ArrayAttr orderAttr = nullptr) {
475 auto toI32Attr = [&](auto range) {
476 SmallVector<int32_t> v(range.begin(), range.end());
477 return DenseI32ArrayAttr::get(context, v);
478 };
479 return xegpu::LayoutAttr::get(context, /*sg_layout=*/nullptr,
480 /*sg_data=*/nullptr,
481 /*inst_data=*/nullptr, toI32Attr(laneLayout),
482 toI32Attr(laneData), orderAttr);
483}
484
485static xegpu::LayoutAttr
487 ArrayRef<int64_t> sgData, ArrayRef<int64_t> instData,
488 ArrayRef<int64_t> laneLayout, ArrayRef<int64_t> laneData,
489 DenseI32ArrayAttr orderAttr = nullptr) {
490 auto toI32Attr = [&](auto range) {
491 SmallVector<int32_t> v(range.begin(), range.end());
492 return DenseI32ArrayAttr::get(context, v);
493 };
494 return xegpu::LayoutAttr::get(
495 context, sgLayout.empty() ? nullptr : toI32Attr(sgLayout),
496 sgData.empty() ? nullptr : toI32Attr(sgData),
497 instData.empty() ? nullptr : toI32Attr(instData),
498 laneLayout.empty() ? nullptr : toI32Attr(laneLayout),
499 laneData.empty() ? nullptr : toI32Attr(laneData), orderAttr);
500}
501
502static xegpu::LayoutAttr buildSgLayout(mlir::MLIRContext *context,
503 ArrayRef<int64_t> wgTileShape,
504 ArrayRef<int64_t> sgLayout,
505 int dimK = -1,
506 DenseI32ArrayAttr orderAttr = nullptr) {
507 SmallVector<int64_t> sgData(sgLayout.size());
508 for (int dim = 0; dim < (int)sgLayout.size(); ++dim) {
509 if (dim == dimK)
510 sgData[dim] = wgTileShape[dim];
511 else
512 sgData[dim] = wgTileShape[dim] / sgLayout[dim];
513 }
514 return buildLayout(context, sgLayout, sgData,
515 /*inst_data=*/{}, /*lane_layout=*/{},
516 /*lane_data=*/{}, /*order=*/nullptr);
517}
518
519/// Infers the source layout attribute for a broadcast operation given the
520/// result layout attribute, result shape, source shape.
521xegpu::DistributeLayoutAttr
522xegpu::inferBroadcastSourceLayout(xegpu::DistributeLayoutAttr resLayout,
523 ArrayRef<int64_t> resShape,
524 ArrayRef<int64_t> srcShape) {
525
526 SmallVector<int64_t> bcastDims;
527 size_t dimDiff = resShape.size() - srcShape.size();
528 auto bcastSourceLayout = resLayout;
529
530 // Right-aligned source in result, look for stretched unit dims.
531 for (size_t i = dimDiff; i < resShape.size(); i++) {
532 if ((srcShape[i - dimDiff] == 1) && (resShape[i] != 1))
533 bcastDims.push_back(i);
534 }
535
536 // Case UnitDimStretch (e.g., 1x4 -> 4x4): the source layout data field must
537 // be 1.
538 if (!bcastDims.empty())
539 bcastSourceLayout = bcastSourceLayout.setUnitDimData(bcastDims);
540
541 // Case RankDiff:
542 if (dimDiff) {
543 SmallVector<int64_t> sliceDims;
544 bool isOuterDimDiffUnitDims = llvm::all_of(
545 resShape.take_front(dimDiff), [&](int64_t dim) { return dim == 1; });
546 if (dimDiff && bcastDims.size() == dimDiff && isOuterDimDiffUnitDims) {
547 // Case RankDiffInnerDims (e.g., 1x4 -> 1x16x4):
548 // slice the expanded inner dims
549 sliceDims.assign(bcastDims.begin(), bcastDims.end());
550 } else {
551 // Case RankDiffOuterDims (e.g., 1x4 -> 1x1x4):
552 // slice the outer dims
553 llvm::append_range(sliceDims, llvm::seq<int64_t>(0, dimDiff));
554 }
555 bcastSourceLayout = xegpu::SliceAttr::get(
556 resLayout.getContext(), bcastSourceLayout,
557 DenseI64ArrayAttr::get(resLayout.getContext(), sliceDims));
558 }
559 return bcastSourceLayout;
560}
561
562/// Infers the source layout attribute for a reduction operation given the
563/// result layout attribute and reduced dims.
564xegpu::DistributeLayoutAttr
565xegpu::inferMultiReductionSourceLayout(xegpu::DistributeLayoutAttr resLayout,
566 SmallVector<int64_t> reduceDims) {
567
568 assert(isa<xegpu::SliceAttr>(resLayout) &&
569 "reduction result layout must be slice layout");
570
571 xegpu::SliceAttr sliceLayout = dyn_cast<xegpu::SliceAttr>(resLayout);
572
573 assert((reduceDims == sliceLayout.getDims().asArrayRef()) &&
574 "reduction dims must match with slice dims");
575
576 return sliceLayout.getParent();
577}
578
579xegpu::DistributeLayoutAttr
580xegpu::inferReductionSourceLayout(xegpu::DistributeLayoutAttr resLayout) {
581 return xegpu::inferMultiReductionSourceLayout(resLayout, {0});
582}
583
584/// Infers the source layout attribute for a transpose operation given the
585/// result layout attribute and permutation.
586///
587/// vector.transpose semantics is `result[i] = source[permutation[i]]`, so
588/// `result_layout[i] = source_layout[permutation[i]]`. To recover the source
589/// layout from the result layout we must apply the inverse permutation.
590xegpu::DistributeLayoutAttr
591xegpu::inferTransposeSourceLayout(xegpu::DistributeLayoutAttr resLayout,
592 ArrayRef<int64_t> permutation) {
594 invertPermutationVector(permutation);
595 return resLayout.transposeDims(inversePermutation);
596}
597
598/// Infers the source layout attribute for a bitcast operation given the
599/// result layout attribute, result element type bitwidth, and source element
600/// type bitwidth.
601xegpu::DistributeLayoutAttr
602xegpu::inferBitCastSourceLayout(xegpu::DistributeLayoutAttr resLayout,
603 int resElemTyBitWidth, int srcElemTyBitWidth) {
604
605 SmallVector<int64_t> sgData = resLayout.getEffectiveSgDataAsInt();
606 SmallVector<int64_t> instData = resLayout.getEffectiveInstDataAsInt();
607 SmallVector<int64_t> laneData = resLayout.getEffectiveLaneDataAsInt();
608 size_t sgDataSize = sgData.size();
609 size_t instDataSize = instData.size();
610 size_t laneDataSize = laneData.size();
611 int64_t sgDataValue = -1;
612 int64_t instDataValue = -1;
613 int64_t laneDataValue = -1;
614 int64_t dim = resLayout.getRank() - 1;
615
616 if (srcElemTyBitWidth <= resElemTyBitWidth) {
617 int bitWidthRatio = resElemTyBitWidth / srcElemTyBitWidth;
618 if (sgDataSize)
619 sgDataValue = sgData.back() * bitWidthRatio;
620 if (instDataSize)
621 instDataValue = instData.back() * bitWidthRatio;
622 if (laneDataSize)
623 laneDataValue = laneData.back() * bitWidthRatio;
624 } else {
625 int bitWidthRatio = srcElemTyBitWidth / resElemTyBitWidth;
626 if (sgDataSize) {
627 assert((sgData.back() % bitWidthRatio) == 0 &&
628 "sgData not divisible by bitWidthRatio");
629 sgDataValue = sgData.back() / bitWidthRatio;
630 }
631 if (instDataSize) {
632 assert((instData.back() % bitWidthRatio) == 0 &&
633 "instData not divisible by bitWidthRatio");
634 instDataValue = instData.back() / bitWidthRatio;
635 }
636 if (laneDataSize) {
637 assert((laneData.back() % bitWidthRatio) == 0 &&
638 "laneData not divisible by bitWidthRatio");
639 laneDataValue = laneData.back() / bitWidthRatio;
640 }
641 }
642
643 xegpu::DistributeLayoutAttr finalSrcLayout;
644 finalSrcLayout =
645 resLayout.setDimData(dim, sgDataValue, instDataValue, laneDataValue);
646
647 return finalSrcLayout;
648}
649
650/// Infers the source layout attribute for an interleave operation given the
651/// result layout attribute. Interleave doubles the size of the innermost
652/// dimension, so the layout inference is similar to bitcast where the source
653/// element type is larger than the result element type (ratio = 2).
654xegpu::DistributeLayoutAttr
655xegpu::inferInterleaveSourceLayout(xegpu::DistributeLayoutAttr resLayout) {
656
657 SmallVector<int64_t> sgData = resLayout.getEffectiveSgDataAsInt();
658 SmallVector<int64_t> instData = resLayout.getEffectiveInstDataAsInt();
659 SmallVector<int64_t> laneData = resLayout.getEffectiveLaneDataAsInt();
660 size_t sgDataSize = sgData.size();
661 size_t instDataSize = instData.size();
662 size_t laneDataSize = laneData.size();
663 int64_t sgDataValue = -1;
664 int64_t instDataValue = -1;
665 int64_t laneDataValue = -1;
666 int64_t dim = resLayout.getRank() - 1;
667
668 // Interleave doubles the innermost dimension, so we need to halve the
669 // layout values (similar to bitcast with ratio = 2)
670 constexpr int ratio = 2;
671 if (sgDataSize) {
672 assert((sgData.back() % ratio) == 0 &&
673 "sgData not divisible by interleave ratio");
674 sgDataValue = sgData.back() / ratio;
675 }
676 if (instDataSize) {
677 assert((instData.back() % ratio) == 0 &&
678 "instData not divisible by interleave ratio");
679 instDataValue = instData.back() / ratio;
680 }
681 if (laneDataSize) {
682 assert((laneData.back() % ratio) == 0 &&
683 "laneData not divisible by interleave ratio");
684 laneDataValue = laneData.back() / ratio;
685 }
686
687 return resLayout.setDimData(dim, sgDataValue, instDataValue, laneDataValue);
688}
689
690/// Infers the source layout attribute for a deinterleave operation given the
691/// result layout attribute. Deinterleave halves the size of the innermost
692/// dimension, so the layout inference is similar to bitcast where the source
693/// element type is smaller than the result element type (ratio = 2).
694xegpu::DistributeLayoutAttr
695xegpu::inferDeinterleaveSourceLayout(xegpu::DistributeLayoutAttr resLayout) {
696
697 SmallVector<int64_t> sgData = resLayout.getEffectiveSgDataAsInt();
698 SmallVector<int64_t> instData = resLayout.getEffectiveInstDataAsInt();
699 SmallVector<int64_t> laneData = resLayout.getEffectiveLaneDataAsInt();
700 size_t sgDataSize = sgData.size();
701 size_t instDataSize = instData.size();
702 size_t laneDataSize = laneData.size();
703 int64_t sgDataValue = -1;
704 int64_t instDataValue = -1;
705 int64_t laneDataValue = -1;
706 int64_t dim = resLayout.getRank() - 1;
707
708 // Deinterleave halves the innermost dimension, so we need to double the
709 // layout values (similar to bitcast with ratio = 2)
710 constexpr int ratio = 2;
711 if (sgDataSize)
712 sgDataValue = sgData.back() * ratio;
713 if (instDataSize)
714 instDataValue = instData.back() * ratio;
715 if (laneDataSize)
716 laneDataValue = laneData.back() * ratio;
717
718 return resLayout.setDimData(dim, sgDataValue, instDataValue, laneDataValue);
719}
720
721/// Infers the source layout attribute for an insert strided slice operation
722/// given the result layout attribute, result shape, and source shape. Removes
723/// leading dimensions from the result layout to match the source shape size.
724xegpu::DistributeLayoutAttr xegpu::inferInsertStridedSliceSourceLayout(
725 xegpu::DistributeLayoutAttr resLayout, ArrayRef<int64_t> resShape,
726 ArrayRef<int64_t> srcShape) {
727
728 int srcShapeSize = srcShape.size();
729 int resShapeSize = resShape.size();
730 int dimDiff = resShapeSize - srcShapeSize;
731
732 if (dimDiff > 0) {
733 // assert that the leading dimensions being sliced off are not distributed
734 // (i.e. sg_layout and lane_layout for those dimensions are all 1)
735 auto resSgLayout = resLayout.getEffectiveSgLayoutAsInt();
736 auto resLaneLayout = resLayout.getEffectiveLaneLayoutAsInt();
737 for (int i = 0; i < dimDiff; i++) {
738 assert((resSgLayout.size() == 0 || resSgLayout[i] == 1) &&
739 (resLaneLayout.size() == 0 || resLaneLayout[i] == 1) &&
740 "Leading dimensions being sliced off must not be distributed");
741 }
742 return resLayout.dropDims(llvm::to_vector(llvm::seq<int64_t>(0, dimDiff)));
743 }
744 return resLayout;
745}
746
747/// Infers the source layout attribute for an insert operation
748/// given the result layout attribute, result shape, and source shape. Removes
749/// leading dimensions from the result layout to match the source shape size.
750// TODO: add propagation support for insert op
751xegpu::DistributeLayoutAttr
752xegpu::inferInsertSourceLayout(xegpu::DistributeLayoutAttr resLayout,
753 ArrayRef<int64_t> resShape,
754 ArrayRef<int64_t> srcShape) {
755
756 int srcShapeSize = srcShape.size();
757 int resShapeSize = resShape.size();
758 int dimDiff = resShapeSize - srcShapeSize;
759
760 if (dimDiff > 0) {
761 // assert that the leading dimensions being sliced off are not distributed
762 // (i.e. sg_layout and lane_layout for those dimensions are all 1)
763 auto resSgLayout = resLayout.getEffectiveSgLayoutAsInt();
764 auto resLaneLayout = resLayout.getEffectiveLaneLayoutAsInt();
765 for (int i = 0; i < dimDiff; i++) {
766 assert((resSgLayout.size() == 0 || resSgLayout[i] == 1) &&
767 (resLaneLayout.size() == 0 || resLaneLayout[i] == 1) &&
768 "Leading dimensions being sliced off must not be distributed");
769 }
770 return resLayout.dropDims(llvm::to_vector(llvm::seq<int64_t>(0, dimDiff)));
771 }
772 return resLayout;
773}
774
775/// Infers the source layout attribute for extract operation
776/// given the result layout attribute, result shape, and source shape. Adds
777/// leading dimensions to the source layout to match the source shape size.
778// TODO: add layout attribute interface: expandDim() and use it here.
779// TODO: add propagation support for extract op
780xegpu::DistributeLayoutAttr
781xegpu::inferExtractSourceLayout(xegpu::DistributeLayoutAttr resLayout,
782 ArrayRef<int64_t> resShape,
783 ArrayRef<int64_t> srcShape) {
784
785 int srcShapeSize = srcShape.size();
786 int resShapeSize = resShape.size();
787 int dimDiff = srcShapeSize - resShapeSize;
788 auto context = resLayout.getContext();
789 // construct the source layout by adding unit dimensions to the front of
790 // result layout
791 if (dimDiff > 0) {
792 auto sgLayout = resLayout.getEffectiveSgLayoutAsInt();
793 auto sgData = resLayout.getEffectiveSgDataAsInt();
794 auto instData = resLayout.getEffectiveInstDataAsInt();
795 auto laneLayout = resLayout.getEffectiveLaneLayoutAsInt();
796 auto laneData = resLayout.getEffectiveLaneDataAsInt();
797 auto order = resLayout.getEffectiveOrderAsInt();
798
799 // Example: result shape is 3D with order [1, 2, 0], source shape is 5D
800 // (adding 2 leading dimensions). Expected source order: [3, 4, 2, 1, 0]
801 // Step 1: shift existing order by dimDiff: [1, 2, 0] -> [3, 4, 2]
802 // Step 2: append new leading dims in reverse (slowest first): [3, 4, 2, 1,
803 // 0]
804
805 // Shift existing dimension indices in order by dimDiff to account for the
806 // new leading dimensions being added to the source shape
807 for (auto &o : order)
808 o += dimDiff;
809
810 // Add unit dimensions to the front of non-empty layout vectors and append
811 // the new dimension indices to the order array in reverse (slowest
812 // dimension has the lowest index and appears last in the order array)
813 for (int i = 0; i < dimDiff; i++) {
814 if (!sgLayout.empty())
815 sgLayout.insert(sgLayout.begin(), 1);
816 if (!sgData.empty())
817 sgData.insert(sgData.begin(), 1);
818 if (!instData.empty())
819 instData.insert(instData.begin(), 1);
820 if (!laneLayout.empty())
821 laneLayout.insert(laneLayout.begin(), 1);
822 if (!laneData.empty())
823 laneData.insert(laneData.begin(), 1);
824 order.push_back(dimDiff - 1 - i);
825 }
826
828 context, SmallVector<int32_t>(order.begin(), order.end()));
829 if (!resLayout.getOrder())
830 orderAttr = nullptr;
831
832 return buildLayout(context, sgLayout, sgData, instData, laneLayout,
833 laneData, orderAttr);
834 }
835 return resLayout;
836}
837
838/// Infers the source layout attribute for a shape cast operation given the
839/// result layout attribute, result shape, and source shape.
840xegpu::DistributeLayoutAttr
841xegpu::inferShapeCastSourceLayout(xegpu::DistributeLayoutAttr resLayout,
842 ArrayRef<int64_t> resShape,
843 ArrayRef<int64_t> srcShape) {
844
845 // There are three use cases:
846 // 1. expand dims of low-rank dimensions (e.g., 1D to 2D): to set up the
847 // tensor before broadcast
848 // 2. split dim of a high-rank dimension (e.g., 1D to 2D): to setup tensor
849 // for multi-stage reduction
850 // 3. combines all dims to a single dim and put in the innermost dim in 2d as
851 // [1, combinedData] or [combinedData]. Say, [2, 4, 8] -> [1, 64] or [64]
852 // Use cases are only supported after workgroup distribution,
853 // like cross-sg reduction saves multidimension data to
854 // 1D slm buffer, shapecast inserted by cse/canonicalization passes.
855
856 // Use case 1: Shapes only differ by expanding unit dimensions, for broadcast
857 SmallVector<int64_t> expandedUnitDims;
858
859 if (xegpu::matchUnitDimExpansion(srcShape, resShape, expandedUnitDims)) {
860 // create a slice layout for the source by removing the expanded unit dims
861 auto sliceDimsAttr = DenseI64ArrayAttr::get(
862 resLayout.getContext(), ArrayRef<int64_t>(expandedUnitDims));
863 auto srcLayout =
864 xegpu::SliceAttr::get(resLayout.getContext(), resLayout, sliceDimsAttr);
865 return srcLayout;
866 }
867
868 // Use case 2: Dim split from source to result, for multi-stage reduction
869 SmallVector<SmallVector<int64_t>> splitDimGroups;
870 if (xegpu::matchSplitDimExpansion(srcShape, resShape, splitDimGroups)) {
871 auto srcLayout = resLayout;
872 for (const auto &dimGroup : splitDimGroups)
873 srcLayout = srcLayout.collapseDims(dimGroup);
874
875 return srcLayout;
876 }
877
878 // Use case 3: General dim collapse, for cross-sg reduction to SLM and other
879 // shape casts where consecutive src dims fold into a single dst dim.
881 if (xegpu::matchDimCollapse(srcShape, resShape, collapseDims)) {
882 auto srcLayout = resLayout;
883 for (int64_t dstIdx = static_cast<int64_t>(collapseDims.size()) - 1;
884 dstIdx >= 0; --dstIdx) {
885 ArrayRef<int64_t> srcDims = collapseDims[dstIdx];
886 if (srcDims.empty()) {
887 srcLayout = srcLayout.dropDims({dstIdx});
888 continue;
889 }
890 if (srcDims.size() == 1)
891 continue;
892 SmallVector<int64_t> targetShape;
893 targetShape.reserve(srcDims.size());
894 for (int64_t d : srcDims)
895 targetShape.push_back(srcShape[d]);
896 srcLayout = srcLayout.expandDim(dstIdx, targetShape);
897 }
898 return srcLayout;
899 }
900 return nullptr;
901}
902
903//===----------------------------------------------------------------------===//
904// Forward layout inference (source layout -> result layout)
905//===----------------------------------------------------------------------===//
906
907/// Infers the result layout attribute for a transpose operation given the
908/// source layout attribute and permutation.
909///
910/// vector.transpose semantics is `result[i] = source[permutation[i]]`, so
911/// `result_layout[i] = source_layout[permutation[i]]`, which is exactly
912/// `srcLayout.transposeDims(permutation)`. This is the inverse of
913/// inferTransposeSourceLayout (which applies the inverse permutation).
914xegpu::DistributeLayoutAttr
915xegpu::inferTransposeResultLayout(xegpu::DistributeLayoutAttr srcLayout,
916 ArrayRef<int64_t> permutation) {
917 return srcLayout.transposeDims(permutation);
918}
919
920/// Infers the result layout attribute for a shape cast operation given the
921/// source layout attribute, source shape, and result shape. This is the
922/// inverse of inferShapeCastSourceLayout: a dim-split (src -> res) is undone by
923/// collapsing the split groups, and a dim-collapse (src -> res) is undone by
924/// expanding the collapsed groups. The unit-dim-expansion case is not inverted
925/// here because recovering which result dims are the expanded unit dims would
926/// require the SliceAttr the backward direction produces; such patterns return
927/// nullptr (leaving the result un-laid-out).
928xegpu::DistributeLayoutAttr
929xegpu::inferShapeCastResultLayout(xegpu::DistributeLayoutAttr srcLayout,
930 ArrayRef<int64_t> srcShape,
931 ArrayRef<int64_t> resShape) {
932 // Case: source dims were split into result dims (forward of use case 2 in
933 // inferShapeCastSourceLayout). Undo by expanding each source dim into its
934 // group of result dims.
935 SmallVector<SmallVector<int64_t>> splitDimGroups;
936 if (xegpu::matchSplitDimExpansion(srcShape, resShape, splitDimGroups)) {
937 auto resLayout = srcLayout;
938 // Process source dims from innermost to outermost so that expanding a dim
939 // does not shift the indices of dims not yet processed.
940 for (int64_t srcIdx = static_cast<int64_t>(splitDimGroups.size()) - 1;
941 srcIdx >= 0; --srcIdx) {
942 ArrayRef<int64_t> resDims = splitDimGroups[srcIdx];
943 if (resDims.size() <= 1)
944 continue;
945 SmallVector<int64_t> targetShape;
946 targetShape.reserve(resDims.size());
947 for (int64_t d : resDims)
948 targetShape.push_back(resShape[d]);
949 resLayout = resLayout.expandDim(srcIdx, targetShape);
950 }
951 return resLayout;
952 }
953
954 // Case: source dims were collapsed into result dims (forward of use case 3).
955 // Undo by collapsing each group of source dims into its single result dim.
957 if (xegpu::matchDimCollapse(srcShape, resShape, collapseDims)) {
958 auto resLayout = srcLayout;
959 // Process result dims from innermost to outermost so that collapsing a
960 // group does not shift the indices of groups not yet processed.
961 for (int64_t dstIdx = static_cast<int64_t>(collapseDims.size()) - 1;
962 dstIdx >= 0; --dstIdx) {
963 ArrayRef<int64_t> srcDims = collapseDims[dstIdx];
964 // A result dim with no backing source dims is a trailing/leading unit
965 // dim; its forward inference is ambiguous, so bail out.
966 if (srcDims.empty())
967 return nullptr;
968 if (srcDims.size() == 1)
969 continue;
970 resLayout = resLayout.collapseDims(llvm::to_vector(srcDims));
971 }
972 return resLayout;
973 }
974
975 return nullptr;
976}
977
978/// Infers the result layout attribute for a non-anchor operation from the
979/// layouts of its source operands. Forward counterpart of
980/// inferSourceLayoutFromResultForNonAnchorOp.
981xegpu::DistributeLayoutAttr xegpu::inferResultLayoutFromSourceForNonAnchorOp(
983 if (op->getNumResults() != 1)
984 return nullptr;
985
986 // For vector::TransposeOp, infer the result layout from the source layout.
987 if (auto transpose = dyn_cast<vector::TransposeOp>(op)) {
988 if (!operandLayouts[0])
989 return nullptr;
990 return xegpu::inferTransposeResultLayout(operandLayouts[0],
991 transpose.getPermutation());
992 }
993
994 // For vector::ShapeCastOp, infer the result layout from the source layout.
995 if (auto shapeCast = dyn_cast<vector::ShapeCastOp>(op)) {
996 if (!operandLayouts[0])
997 return nullptr;
999 operandLayouts[0], shapeCast.getSourceVectorType().getShape(),
1000 shapeCast.getResultVectorType().getShape());
1001 }
1002
1003 // For elementwise operations, all operands and the result share the same
1004 // layout. Use the first operand that carries a layout.
1006 for (xegpu::DistributeLayoutAttr layout : operandLayouts)
1007 if (layout)
1008 return layout;
1009 return nullptr;
1010 }
1011
1012 // TODO: add forward inference rules for the remaining ops; their result is
1013 // left un-laid-out until then.
1014 // - vector::BroadcastOp: the forward direction is under-determined. The
1015 // backward rule (inferBroadcastSourceLayout) either sets broadcast dims to
1016 // unit data (losing the original data on those dims) or wraps the result
1017 // in a SliceAttr; neither is generally invertible from the source layout
1018 // alone, so a forward rule must decide how to distribute the new/stretched
1019 // dims.
1020 // - vector::BitCastOp, vector::MultiDimReductionOp / vector::ReductionOp,
1021 // vector::InterleaveOp / vector::DeinterleaveOp, and the insert / extract
1022 // / strided-slice family.
1023 return nullptr;
1024}
1025
1026//===----------------------------------------------------------------------===//
1027// Layout derivation helpers: factorize sgCount into
1028// sg_layout candidates, then
1029// compute per-subgroup (sgData) and per-lane
1030// (lane_layout/lane_data/inst_data).
1031//===----------------------------------------------------------------------===//
1032
1034
1035/// Enumerates all ways to split `total` into `rank` factors whose product
1036/// equals `total`. Returns the list of all such factorizations.
1038 int64_t rank) {
1040 SmallVector<int64_t> current(rank, 0);
1041
1042 // Returns all divisors of `n` in ascending order.
1043 auto getDivisors = [](int64_t n) {
1045 for (int64_t i = 1; i * i <= n; ++i) {
1046 if (n % i == 0) {
1047 divs.push_back(i);
1048 if (i != n / i)
1049 divs.push_back(n / i);
1050 }
1051 }
1052 llvm::sort(divs);
1053 return divs;
1054 };
1055
1056 std::function<void(int64_t, int64_t)> generate = [&](int64_t dim,
1057 int64_t remaining) {
1058 if (dim == rank - 1) {
1059 current[dim] = remaining;
1060 results.push_back(LayoutRepresentation(current));
1061 return;
1062 }
1063 for (int64_t factor : getDivisors(remaining)) {
1064 current[dim] = factor;
1065 generate(dim + 1, remaining / factor);
1066 }
1067 };
1068
1069 generate(0, total);
1070 return results;
1071}
1072
1073// Computes all valid N-dimensional sg_layout candidates for the given
1074// sgCount, whose sgData (= wgShape / sgLayout):
1075// 1. Evenly divides wgShape (i.e., wgShape[d] % sgLayout[d] == 0).
1076// 2. Is a multiple of instData (i.e., sgData[d] % instData[d] == 0).
1077// Results are sorted by balance (smallest max-min spread first), with
1078// lexicographic order as a tiebreaker.
1079//
1080// `broadcastDim` (default -1 = none) marks a dimension broadcast across
1081// subgroups rather than distributed (e.g. the K/contraction dim of a DPAS
1082// operand). Its full extent stays in every subgroup, so rule 1 is skipped for
1083// it, but rule 2 (multiple of instData) still applies.
1084//
1085// Example (2D):
1086// wgShape = [128, 64], instData = [8, 16], sgCount = 32
1087// Returns: [[8,4], [16,2]], corresponding to sgData [16,16] and [8,32].
1090 int64_t sgCount, int64_t broadcastDim = -1) {
1091 int64_t rank = wgShape.size();
1092 assert(rank > 0 && "wgShape must be non-empty");
1093 assert(static_cast<int64_t>(instData.size()) == rank &&
1094 "instData rank must match wgShape rank");
1095
1096 // Step 1: Get all N-D factorizations of sgCount.
1097 auto allFactorizations = enumerateFactorizations(sgCount, rank);
1098
1099 // Step 2: Filter to keep only valid candidates.
1101 for (const auto &sgLayout : allFactorizations) {
1102 bool valid = true;
1103 for (int64_t dim = 0; dim < rank; ++dim) {
1104 // A broadcast dim keeps its full extent in every subgroup; others are
1105 // split evenly by sgLayout[dim].
1106 int64_t sgData;
1107 if (dim == broadcastDim) {
1108 sgData = wgShape[dim];
1109 } else {
1110 if (wgShape[dim] % sgLayout[dim] != 0) {
1111 valid = false;
1112 break;
1113 }
1114 sgData = wgShape[dim] / sgLayout[dim];
1115 }
1116 if (sgData % instData[dim] != 0) {
1117 valid = false;
1118 break;
1119 }
1120 }
1121 if (valid)
1122 candidates.push_back(sgLayout);
1123 }
1124
1125 // Step 3: Sort by balance (smallest max-min spread), then lexicographic.
1126 llvm::sort(candidates, [](const LayoutRepresentation &lhs,
1127 const LayoutRepresentation &rhs) {
1128 int64_t spreadLhs = *llvm::max_element(lhs) - *llvm::min_element(lhs);
1129 int64_t spreadRhs = *llvm::max_element(rhs) - *llvm::min_element(rhs);
1130 if (spreadLhs != spreadRhs)
1131 return spreadLhs < spreadRhs;
1132 return lhs < rhs;
1133 });
1134 return candidates;
1135}
1136
1137/// Helper function to compute inst_data vectors for DPAS operands A, B, and
1138/// C/D.
1139static std::optional<SmallVector<int64_t>> get2DBlockIOInstDataLayout(
1140 ArrayRef<int64_t> dataShape, Type elemTy,
1141 const xegpu::uArch::BlockIOInstructionInterface *uArchInstruction,
1142 bool transform = false, bool transpose = false) {
1143 int rank = dataShape.size();
1144 auto blockWHC =
1145 uArchInstruction->getBlockWidthHeightCount(elemTy, transform, transpose);
1146 if (!blockWHC)
1147 return std::nullopt;
1148 auto [bWidths, bHeights, bCounts] = blockWHC.value();
1149 // Compute inst_data from hardware block params. For Nd ops, the lane
1150 // factorization above (laneLayout / laneData) is rigid; inst_data must be
1151 // a multiple of lane_layout * lane_data on each dim (Category A
1152 // invariant).
1153 SmallVector<int64_t> instData(rank, 1);
1154 assert(rank >= 2 && "dataShape must be at least 2D for 2D-block IO");
1155 int instWidth =
1156 xegpu::getLargestDivisor(static_cast<int>(dataShape.back()), bWidths);
1157 int instHeight =
1158 xegpu::getLargestDivisor(static_cast<int>(dataShape[rank - 2]), bHeights);
1159 // No supported hardware block size divides the data dim (e.g. innermost dim
1160 // of 1 vs. minimum block width 16): not realizable as a 2D-block instruction.
1161 if (instWidth < 0 || instHeight < 0)
1162 return std::nullopt;
1163 instData.back() = instWidth;
1164 instData[rank - 2] = instHeight;
1165
1166 return instData;
1167}
1168
1169/// Helper function to compute inst_data vectors for DPAS operands A, B, and
1170/// C/D. Look up the uArch table and search for the largest supported block size
1171/// that divides the data shape
1172static std::optional<std::tuple<SmallVector<int64_t>, SmallVector<int64_t>,
1175 VectorType aTy, VectorType bTy, VectorType cdTy,
1176 const xegpu::uArch::MMAInstructionInterface *uArchInstruction) {
1177
1178 // M dimension is the second-to-last dim of A (handles batch dims).
1179 const unsigned dataALen = aTy.getShape()[aTy.getRank() - 2];
1180 auto supportedALen = uArchInstruction->getSupportedM(aTy.getElementType());
1181 const int maxALen =
1182 xegpu::getLargestDivisor(dataALen, ArrayRef<unsigned>(supportedALen));
1183
1184 // N dimension is the last dim of B.
1185 const unsigned dataBLen = bTy.getShape().back();
1186 auto supportedBLen = uArchInstruction->getSupportedN(bTy.getElementType());
1187 const int maxBLen =
1188 xegpu::getLargestDivisor(dataBLen, ArrayRef<unsigned>(supportedBLen));
1189
1190 auto supportedCLen = uArchInstruction->getSupportedN(cdTy.getElementType());
1191 const int maxCLen =
1192 xegpu::getLargestDivisor(dataBLen, ArrayRef<unsigned>(supportedCLen));
1193 if (maxALen == -1 || maxBLen == -1 || maxCLen == -1)
1194 return std::nullopt;
1195
1196 auto supportedKLen = uArchInstruction->getSupportedK(aTy.getElementType());
1197 if (supportedKLen.empty())
1198 return std::nullopt;
1199 auto kDimSize = supportedKLen[0];
1200
1201 SmallVector<int64_t> instDataA(aTy.getRank(), 1);
1202 instDataA[aTy.getRank() - 2] = maxALen;
1203 instDataA[aTy.getRank() - 1] = kDimSize;
1204 SmallVector<int64_t> instDataB(bTy.getRank(), 1);
1205 instDataB[bTy.getRank() - 2] = kDimSize;
1206 instDataB[bTy.getRank() - 1] = maxBLen;
1207 SmallVector<int64_t> instDataCD(cdTy.getRank(), 1);
1208 instDataCD[cdTy.getRank() - 2] = maxALen;
1209 instDataCD[cdTy.getRank() - 1] = maxCLen;
1210 return std::make_tuple(instDataA, instDataB, instDataCD);
1211}
1212
1213/// Computes lane_layout and lane_data for scatter-style store anchor layouts
1214/// (store scatter, store matrix). Lanes and the per-lane vector both live on
1215/// the innermost dim:
1216/// - laneLayout[innermost] = min(subgroupSize, srcShape[innermost])
1217/// - laneData[innermost] = min(srcShape[innermost] / laneLayout[innermost],
1218/// maxChunkSize)
1219/// All other entries are 1.
1220static std::pair<SmallVector<int64_t>, SmallVector<int64_t>>
1222 int64_t subgroupSize, int64_t maxChunkSize) {
1223 int64_t rank = instShape.size();
1224 SmallVector<int64_t> laneLayout(rank, 1), laneData(rank, 1);
1225 int64_t innermost = rank - 1;
1226 laneLayout[innermost] = std::min(subgroupSize, instShape[innermost]);
1227 laneData[innermost] =
1228 std::min(instShape[innermost] / laneLayout[innermost], maxChunkSize);
1229 return {laneLayout, laneData};
1230}
1231
1232// Computes the (lane_layout, lane_data, order) a 2D block instruction can
1233// deliver over `instShape`, from the element `bitwidth` and the variant's
1234// `packingSize`. lane_data packs elements narrower than the packing size along
1235// the packing dim: the innermost one, or the one above it when transformed
1236// (VNNI). Lanes are spread along the innermost dim, or the one above it when
1237// transposed, which makes that dim fastest-changing and so returns a reversed
1238// `order`. `order` is empty for the variants that keep the default. Leading
1239// dims are unit.
1240static std::tuple<SmallVector<int64_t>, SmallVector<int64_t>,
1243 int64_t bitwidth, int64_t packingSize,
1244 bool transform = false, bool transpose = false) {
1245 int64_t rank = instShape.size();
1246 assert(rank >= 2 && "Expected at least a 2D shape for a 2D block op");
1247 SmallVector<int64_t> laneLayout(rank, 1), laneData(rank, 1), order;
1248
1249 int64_t packDim = transform ? rank - 2 : rank - 1;
1250 laneData[packDim] = llvm::divideCeil(packingSize, bitwidth);
1251
1252 int64_t laneDim = transpose ? rank - 2 : rank - 1;
1253 laneLayout[laneDim] =
1254 std::min(subgroupSize, instShape[laneDim] / laneData[laneDim]);
1255
1256 if (transpose) {
1257 for (int64_t i = 0; i < rank; ++i)
1258 order.push_back(rank - 1 - i);
1259 std::swap(order[0], order[1]);
1260 }
1261
1262 // assert that the lane layout and data fit in the inst shape
1263 for (int64_t i = 0; i < rank; ++i) {
1264 int64_t laneProduct = laneLayout[i] * laneData[i];
1265 assert(instShape[i] % laneProduct == 0 &&
1266 "lane_layout * lane_data must evenly divide the inst shape");
1267 (void)laneProduct;
1268 }
1269 return {laneLayout, laneData, order};
1270}
1271
1272/// Computes the (lane_layout, lane_data) for a multi-reduction's source layout.
1273/// Only the innermost two dims are distributed; leading dims are assumed unit.
1274/// `subgroupSize` lanes go on one dim; up to `maxReduceVectorSize` elements are
1275/// packed into lane_data on the other. To minimize cross-lane reduction, lanes
1276/// are spread across a non-reduction dim when possible so the reduction happens
1277/// within a lane. inst_data is the element-wise product lane_layout *
1278/// lane_data.
1279///
1280/// e.g. with srcShape=[32, 128], subgroupSize=16, maxReduceVectorSize=2:
1281/// - Switch: reductionDims=[1] and consumerReductionDims=[] -> lanes move
1282/// to the non-reduction dim 0: lane_layout=[16, 1], lane_data=[1, 2].
1283/// - Default: reductionDims=[0, 1] (both reduced) -> lanes stay on the
1284/// innermost dim: lane_layout=[1, 16], lane_data=[2, 1].
1285static std::pair<SmallVector<int64_t>, SmallVector<int64_t>>
1287 ArrayRef<int64_t> reductionDims,
1288 int subgroupSize, int64_t maxReduceVectorSize,
1289 bool verticalLaneLayout = false) {
1290 int srcRank = srcShape.size();
1291 SmallVector<int64_t> laneLayout(srcRank, 1), laneData(srcRank, 1);
1292
1293 int innermost = srcRank - 1;
1294 int secondInnermost = srcRank - 2;
1295
1296 if (verticalLaneLayout && secondInnermost >= 0) {
1297 std::swap(innermost, secondInnermost);
1298 }
1299 int laneDim = innermost;
1300 int vectorDim = secondInnermost; // negative for rank 1
1301
1302 laneLayout[laneDim] =
1303 std::min(static_cast<int64_t>(subgroupSize), srcShape[laneDim]);
1304 if (vectorDim >= 0)
1305 laneData[vectorDim] = std::min(maxReduceVectorSize, srcShape[vectorDim]);
1306
1307 return {laneLayout, laneData};
1308}
1309
1310//===----------------------------------------------------------------------===//
1311// Result/anchor-layout setup. Each op category derives lane_layout/lane_data
1312// (and inst_data / sgData) differently. Two things vary across ops:
1313//
1314// * Consumer dependence: consumer-driven ops prefer the layout requested by
1315// their downstream uses and fall back to uArch defaults only when it is
1316// absent/invalid; sinks (StoreNd, PrefetchNd) have no consumer and always
1317// pick their own layout from uArch.
1318//
1319// * Derivation direction between inst_data and lane_layout/lane_data. Both
1320// obey the invariant inst_data = k * lane_layout * lane_data, where `k` is
1321// a per-dim integer >= 1 giving how many times each lane repeats its
1322// access to cover one instruction's data tile (k == 1 means one lane
1323// position per element; k > 1 means the instruction loads/stores several
1324// elements per lane along that dim). Ops solve this invariant from
1325// opposite ends:
1326// - Rigid-lane ops (Nd block IO, DPAS): hardware fixes lane_layout /
1327// lane_data first, then inst_data is built as a multiple of their
1328// product (using get2DBlockIOInstDataLayout / getDpasInstDataLayouts).
1329// - inst_data-first ops (scatter load): take inst_data from the consumer
1330// and derive lane_layout/lane_data underneath it.
1331//
1332// - DPAS (+DPAS_MX) : rigid lanes — inst_data from HW block dims; A/B/C/D
1333// lanes/data follow each operand's matmul role; DPAS_MX
1334// additionally lays out the scale operand.
1335// - LoadNd : consumer-driven, rigid lanes — honors the consumer's
1336// inst_data / lane / sg_layout (incl. transpose & VNNI
1337// packing) when it satisfies uArch block constraints,
1338// else falls back to the default 2D-block scheme (lanes
1339// on the last dim, rank-2 if transposed). The fallback
1340// picks the LARGEST uArch block that divides the data
1341// shape, so the resulting inst_data block can be bigger
1342// than what the consumer asked for (fewer, wider
1343// loads).
1344// - StoreNd/PrefetchNd: data sinks, no consumer, rigid lanes — pick the
1345// 2D-block layout directly from uArch (no VNNI
1346// packing).
1347// - Load (scatter) : load_gather / load_matrix, consumer-driven,
1348// inst_data-first — reuse the consumer's inst_data and
1349// derive lane_layout/lane_data, else default to lanes +
1350// per-lane chunk on the innermost dim (chunk capped by
1351// maxChunkSize).
1352// - Store (scatter) : store_scatter / store_matrix — same scatter scheme,
1353// but always self-derived from the scatter default.
1354// - Reduction : (multi_)reduction, consumer-driven — distribute the
1355// inner two dims, with lanes on the innermost dim by
1356// default (reducing across lanes) and switched to a
1357// non-reduction dim only when that keeps the reduction
1358// within a lane. Reuses the consumer's slice layout
1359// when it slices exactly the reduction dims, otherwise
1360// re-derives. See setupMultiReductionResultLayout for
1361// the exact switch condition and worked examples.
1362// - BitCast/Interleave: scale the innermost data field by the bitwidth /
1363// interleave ratio so the source layout divides back
1364// out.
1365// - InsertStridedSlice: clamp lane_data per dim to fit the inserted slice
1366// (Lane kind only; sg/inst layouts unsupported).
1367//===----------------------------------------------------------------------===//
1368
1369/// Helper function to set up subgroup layouts for DPAS operands A, B, and
1370/// C/D. Compute subgroup layout candidates based on wgtile and instData, and
1371/// then pick the best one that satisfies all operands and the consumer (if
1372/// specified).
1373static std::optional<
1374 std::tuple<xegpu::DistributeLayoutAttr, xegpu::DistributeLayoutAttr,
1375 xegpu::DistributeLayoutAttr>>
1377 mlir::MLIRContext *context, VectorType aTy, VectorType bTy, VectorType cdTy,
1378 xegpu::DistributeLayoutAttr consumerLayout, int numSg,
1380 instDataVecs) {
1381 auto [instDataA, instDataB, instDataCD] = instDataVecs;
1382
1383 std::optional<LayoutRepresentation> consumerSgLayout = std::nullopt;
1384 if (consumerLayout && consumerLayout.isForWorkgroup()) {
1385 consumerSgLayout = consumerLayout.getEffectiveSgLayoutAsInt();
1386 }
1387
1388 // Get all valid layouts for A, B and C/D operands
1389 auto layoutsA = getSgLayoutCandidates(aTy.getShape(), instDataA, numSg,
1390 /*broadcastDim=*/aTy.getRank() - 1);
1391 auto layoutsB = getSgLayoutCandidates(bTy.getShape(), instDataB, numSg,
1392 /*broadcastDim=*/bTy.getRank() - 2);
1393 auto layoutsCD = getSgLayoutCandidates(cdTy.getShape(), instDataCD, numSg);
1394 if (layoutsA.empty() || layoutsB.empty() || layoutsCD.empty())
1395 return std::nullopt;
1396
1397 // Pick the best subgroup layout
1398 std::optional<LayoutRepresentation> bestPick;
1399 for (auto &sgLayout : layoutsB) {
1400 if (llvm::is_contained(layoutsA, sgLayout) &&
1401 llvm::is_contained(layoutsCD, sgLayout)) {
1402 // Is in (A and B and CD) and matches consumer -> best pick
1403 if (consumerSgLayout.has_value() && sgLayout == *consumerSgLayout) {
1404 bestPick = sgLayout;
1405 break;
1406 }
1407 // Is in (A and B and CD) layoutsB is ordered from most
1408 // balanced to least. So the first one we see is the most balanced one,
1409 // remember it and later only update if there is one that matches the
1410 // consumer.
1411 if (!bestPick)
1412 bestPick = sgLayout;
1413 }
1414 }
1415 if (!bestPick)
1416 return std::nullopt;
1417
1418 const auto &picked = *bestPick;
1419
1420 auto dpasALayout = buildSgLayout(context, aTy.getShape(), picked,
1421 /*dimK=*/aTy.getRank() - 1);
1422 auto dpasBLayout = buildSgLayout(context, bTy.getShape(), picked,
1423 /*dimK=*/bTy.getRank() - 2);
1424 auto dpasCDLayout = buildSgLayout(context, cdTy.getShape(), picked);
1425 return std::make_tuple(dpasALayout, dpasBLayout, dpasCDLayout);
1426}
1427
1428/// Sets up the anchor layouts for dpas operands (A, B, and C/D).
1429/// The numSg and consumerLayout (optional) are only used by sg layout
1430/// creation.
1431std::optional<
1432 std::tuple<xegpu::DistributeLayoutAttr, xegpu::DistributeLayoutAttr,
1433 xegpu::DistributeLayoutAttr>>
1434xegpu::setupDpasLayout(xegpu::LayoutKind layoutKind, VectorType aTy,
1435 VectorType bTy, VectorType cdTy,
1436 xegpu::DistributeLayoutAttr consumerLayout, int numSg,
1437 const xegpu::uArch::uArch *uArch) {
1438 auto context = aTy.getContext();
1439 const auto *uArchInstruction =
1440 dyn_cast<xegpu::uArch::SubgroupMatrixMultiplyAcc>(uArch->getInstruction(
1442 if (!uArchInstruction)
1443 return std::nullopt;
1444 auto subgroupSize = uArch->getSubgroupSize();
1445
1446 auto [laneLayoutA, laneDataA, orderA] =
1447 compute2DBlockIOLaneLayout(aTy.getShape(), subgroupSize,
1448 aTy.getElementType().getIntOrFloatBitWidth(),
1449 uArchInstruction->getPackedFormatBitSizeA());
1450 auto [laneLayoutB, laneDataB, orderB] = compute2DBlockIOLaneLayout(
1451 bTy.getShape(), subgroupSize,
1452 bTy.getElementType().getIntOrFloatBitWidth(),
1453 uArchInstruction->getPackedFormatBitSizeB(), /*vnni=*/true);
1454 auto [laneLayoutCD, laneDataCD, orderCD] =
1455 compute2DBlockIOLaneLayout(cdTy.getShape(), subgroupSize,
1456 cdTy.getElementType().getIntOrFloatBitWidth(),
1457 cdTy.getElementType().getIntOrFloatBitWidth());
1458
1459 auto instDataVecs = getDpasInstDataLayouts(aTy, bTy, cdTy, uArchInstruction);
1460 if (!instDataVecs)
1461 return std::nullopt;
1462
1463 if (layoutKind == xegpu::LayoutKind::Subgroup) {
1464 assert(numSg > 0 &&
1465 "Number of subgroups must be provided for sg layout creation.");
1466 return getDpasSubgroupLayouts(context, aTy, bTy, cdTy, consumerLayout,
1467 numSg, *instDataVecs);
1468 } else if (layoutKind == xegpu::LayoutKind::InstData) {
1469 auto [instDataA, instDataB, instDataCD] = *instDataVecs;
1470 return std::make_tuple(
1471 buildInstDataLayoutWithLane(context, instDataA, laneLayoutA, laneDataA),
1472 buildInstDataLayoutWithLane(context, instDataB, laneLayoutB, laneDataB),
1473 buildInstDataLayoutWithLane(context, instDataCD, laneLayoutCD,
1474 laneDataCD));
1475 } else if (layoutKind == xegpu::LayoutKind::Lane) {
1476 auto aLayout = buildLaneLayout(context, laneLayoutA, laneDataA);
1477 auto bLayout = buildLaneLayout(context, laneLayoutB, laneDataB);
1478 auto cdLayout = buildLaneLayout(context, laneLayoutCD, laneDataCD);
1479 return std::make_tuple(aLayout, bLayout, cdLayout);
1480 }
1481 return std::nullopt;
1482}
1483
1484/// Helper to create a scale layout derived from a matrix operand layout.
1485/// The scale layout is computed by mapping each dimension of the matrix
1486/// layout to the corresponding scale tensor dimension using the ratio
1487/// between the matrix and scale shapes.
1488static xegpu::DistributeLayoutAttr
1489createScaleLayout(mlir::MLIRContext *context, VectorType matrixTy,
1490 VectorType scaleTy, xegpu::DistributeLayoutAttr matrixLayout,
1491 bool isBScale, const xegpu::uArch::uArch *uArch) {
1492 if (!scaleTy || !matrixLayout)
1493 return nullptr;
1494
1495 // Calculate scaling factor by dividing matrix shape by scale shape
1496 ArrayRef<int64_t> matrixShape = matrixTy.getShape();
1497 ArrayRef<int64_t> scaleShape = scaleTy.getShape();
1498
1499 // Scale shapes can be 1D or 2D, handle both cases
1500 if (scaleShape.empty())
1501 return nullptr;
1502
1503 auto uArchInstruction =
1504 dyn_cast<xegpu::uArch::SubgroupScaledMatrixMultiplyAcc>(
1505 uArch->getInstruction(
1507
1508 int64_t rank = matrixLayout.getRank();
1509 assert(rank >= 2 && "dpas layouts must be at least two dimensions");
1510
1511 SmallVector<int64_t> sgLayout = matrixLayout.getEffectiveSgLayoutAsInt();
1512 SmallVector<int64_t> sgData = matrixLayout.getEffectiveSgDataAsInt();
1513 SmallVector<int64_t> instData = matrixLayout.getEffectiveInstDataAsInt();
1514 SmallVector<int64_t> laneLayout = matrixLayout.getEffectiveLaneLayoutAsInt();
1515 SmallVector<int64_t> laneData = matrixLayout.getEffectiveLaneDataAsInt();
1516 auto order = matrixLayout.getOrder();
1517
1518 SmallVector<int64_t> scaleSgLayout;
1519 SmallVector<int64_t> scaleSgData;
1520 if (!sgLayout.empty() && !sgData.empty()) {
1521 scaleSgLayout.assign(sgLayout.begin(), sgLayout.end());
1522 scaleSgData.assign(sgData.begin(), sgData.end());
1523 scaleSgData[rank - 2] = std::max<int64_t>(
1524 scaleShape[rank - 2] / (matrixShape[rank - 2] / sgData[rank - 2]), 1);
1525 scaleSgData[rank - 1] = std::max<int64_t>(
1526 scaleShape[rank - 1] / (matrixShape[rank - 1] / sgData[rank - 1]), 1);
1527 }
1528
1529 // For DPAS_MX scales: if matrix has inst_data, scale needs adjusted
1530 // inst_data. Scale inst_data is derived from matrix inst_data divided by
1531 // scale factor.
1532 SmallVector<int64_t> scaleInstData;
1533 if (!instData.empty()) {
1534 scaleInstData.assign(instData.begin(), instData.end());
1535 if (isBScale)
1536 scaleInstData[rank - 2] = std::max<int64_t>(
1537 scaleShape[rank - 2] / (matrixShape[rank - 2] / instData[rank - 2]),
1538 1);
1539 else
1540 scaleInstData[rank - 1] = std::max<int64_t>(
1541 scaleShape[rank - 1] / (matrixShape[rank - 1] / instData[rank - 1]),
1542 1);
1543 }
1544
1545 SmallVector<int64_t> scaleLaneLayout;
1546 SmallVector<int64_t> scaleLaneData;
1547 if (!laneLayout.empty() && !laneData.empty()) {
1548 scaleLaneLayout.assign(laneLayout.begin(), laneLayout.end());
1549 scaleLaneData.assign(laneData.size(), 1);
1550
1551 bool isRowMajor = uArchInstruction->isLaneLayoutRowMajorOrder();
1552 if (isBScale ^ isRowMajor)
1553 std::swap(scaleLaneLayout[rank - 2], scaleLaneLayout[rank - 1]);
1554 // Cap lane_layout by the per-instruction tile (inst_data) on each dim.
1555 // Then derive lane_data = inst_data / lane_layout so the Category A
1556 // invariant inst_data = lane_layout * lane_data * k (with k = 1) holds
1557 // for the scale operand's load_nd consumer.
1558 auto layoutCap = scaleInstData.empty() ? scaleShape : scaleInstData;
1559 for (int64_t d = rank - 2; d < rank; ++d)
1560 scaleLaneLayout[d] = std::min<int64_t>(layoutCap[d], scaleLaneLayout[d]);
1561 }
1562 return buildLayout(context, scaleSgLayout, scaleSgData, scaleInstData,
1563 scaleLaneLayout, scaleLaneData, order);
1564}
1565
1566/// Sets up the anchor layouts for dpas_mx operands (A, B, C/D, A_scale, and
1567/// B_scale). The numSg and consumerLayout (optional) are only used by sg
1568/// layout creation.
1569std::optional<
1570 std::tuple<xegpu::DistributeLayoutAttr, xegpu::DistributeLayoutAttr,
1571 xegpu::DistributeLayoutAttr, xegpu::DistributeLayoutAttr,
1572 xegpu::DistributeLayoutAttr>>
1573xegpu::setupDpasMxLayout(xegpu::LayoutKind layoutKind, VectorType aTy,
1574 VectorType bTy, VectorType cdTy, VectorType aScaleTy,
1575 VectorType bScaleTy,
1576 xegpu::DistributeLayoutAttr consumerLayout, int numSg,
1577 const xegpu::uArch::uArch *uArch) {
1578 auto context = aTy.getContext();
1579 const auto *uArchInstruction =
1580 dyn_cast<xegpu::uArch::SubgroupMatrixMultiplyAcc>(uArch->getInstruction(
1582 if (!uArchInstruction)
1583 return std::nullopt;
1584 auto subgroupSize = uArch->getSubgroupSize();
1585
1586 auto [laneLayoutA, laneDataA, orderA] =
1587 compute2DBlockIOLaneLayout(aTy.getShape(), subgroupSize,
1588 aTy.getElementType().getIntOrFloatBitWidth(),
1589 uArchInstruction->getPackedFormatBitSizeA());
1590 auto [laneLayoutB, laneDataB, orderB] = compute2DBlockIOLaneLayout(
1591 bTy.getShape(), subgroupSize,
1592 bTy.getElementType().getIntOrFloatBitWidth(),
1593 uArchInstruction->getPackedFormatBitSizeB(), /*vnni=*/true);
1594 auto [laneLayoutCD, laneDataCD, orderCD] =
1595 compute2DBlockIOLaneLayout(cdTy.getShape(), subgroupSize,
1596 cdTy.getElementType().getIntOrFloatBitWidth(),
1597 cdTy.getElementType().getIntOrFloatBitWidth());
1598 auto instDataVecs = getDpasInstDataLayouts(aTy, bTy, cdTy, uArchInstruction);
1599 if (!instDataVecs)
1600 return std::nullopt;
1601
1602 if (layoutKind == xegpu::LayoutKind::Subgroup) {
1603 assert(numSg > 0 &&
1604 "Number of subgroups must be provided for sg layout creation.");
1605 auto dpasLayouts = getDpasSubgroupLayouts(
1606 context, aTy, bTy, cdTy, consumerLayout, numSg, *instDataVecs);
1607 if (!dpasLayouts)
1608 return std::nullopt;
1609
1610 auto [dpasALayout, dpasBLayout, dpasCDLayout] = *dpasLayouts;
1611
1612 // Create scale layouts
1613 auto aScaleLayout =
1614 createScaleLayout(context, aTy, aScaleTy, dpasALayout, false, uArch);
1615
1616 auto bScaleLayout =
1617 createScaleLayout(context, bTy, bScaleTy, dpasBLayout, true, uArch);
1618
1619 return std::make_tuple(dpasALayout, dpasBLayout, dpasCDLayout, aScaleLayout,
1620 bScaleLayout);
1621 } else if (layoutKind == xegpu::LayoutKind::InstData) {
1622
1623 auto [instDataA, instDataB, instDataCD] = *instDataVecs;
1624
1625 auto dpasALayout =
1626 buildInstDataLayoutWithLane(context, instDataA, laneLayoutA, laneDataA);
1627 auto dpasBLayout =
1628 buildInstDataLayoutWithLane(context, instDataB, laneLayoutB, laneDataB);
1629 auto dpasCDLayout = buildInstDataLayoutWithLane(context, instDataCD,
1630 laneLayoutCD, laneDataCD);
1631
1632 auto aScaleLayout =
1633 createScaleLayout(context, aTy, aScaleTy, dpasALayout, false, uArch);
1634 auto bScaleLayout =
1635 createScaleLayout(context, bTy, bScaleTy, dpasBLayout, true, uArch);
1636
1637 return std::make_tuple(dpasALayout, dpasBLayout, dpasCDLayout, aScaleLayout,
1638 bScaleLayout);
1639 } else if (layoutKind == xegpu::LayoutKind::Lane) {
1640 auto dpasALayout = buildLaneLayout(context, laneLayoutA, laneDataA);
1641 auto dpasBLayout = buildLaneLayout(context, laneLayoutB, laneDataB);
1642 auto dpasCDLayout = buildLaneLayout(context, laneLayoutCD, laneDataCD);
1643
1644 auto aScaleLayout =
1645 createScaleLayout(context, aTy, aScaleTy, dpasALayout, false, uArch);
1646 auto bScaleLayout =
1647 createScaleLayout(context, bTy, bScaleTy, dpasBLayout, true, uArch);
1648
1649 return std::make_tuple(dpasALayout, dpasBLayout, dpasCDLayout, aScaleLayout,
1650 bScaleLayout);
1651 }
1652 return std::nullopt;
1653}
1654
1655/// Sets up the anchor layout for a store_nd operation. StoreNd picks its
1656/// own layout based on uArch block parameters (it does not take a consumer
1657/// layout, since it is a data sink).
1658xegpu::DistributeLayoutAttr
1660 VectorType srcVecTy, int numSg,
1661 const xegpu::uArch::uArch *uArch) {
1662 const auto *uArchInstruction =
1663 dyn_cast<xegpu::uArch::Subgroup2DBlockStoreInstruction>(
1664 uArch->getInstruction(
1666 if (!uArchInstruction)
1667 return nullptr;
1668
1669 auto context = srcVecTy.getContext();
1670 Type elemTy = srcVecTy.getElementType();
1671 auto subgroupSize = uArch->getSubgroupSize();
1672 auto dataShape = srcVecTy.getShape();
1673 [[maybe_unused]] int rank = srcVecTy.getRank();
1674 assert(rank >= 2 && "Expected at least 2D shape for ND op");
1675
1676 // Compute the default 2D block IO lane layout / lane data.
1677 unsigned bitwidth = elemTy.getIntOrFloatBitWidth();
1678 auto [laneLayout, laneData, order] =
1679 compute2DBlockIOLaneLayout(dataShape, subgroupSize, bitwidth,
1680 uArchInstruction->getPackedFormatBitSize());
1681
1682 if (layoutKind == xegpu::LayoutKind::Lane)
1683 return buildLaneLayout(context, laneLayout, laneData);
1684
1685 auto instData =
1686 get2DBlockIOInstDataLayout(dataShape, elemTy, uArchInstruction);
1687 // Shape not realizable as a 2D-block instruction; let the caller report it.
1688 if (!instData)
1689 return nullptr;
1690
1691 if (layoutKind == xegpu::LayoutKind::InstData) {
1692 assert(isValidLaneLayout(*instData, laneLayout, laneData) &&
1693 "Expected the store layout to satisfy uArch block constraints");
1694 return buildInstDataLayoutWithLane(context, *instData, laneLayout,
1695 laneData);
1696 }
1697
1698 if (layoutKind == xegpu::LayoutKind::Subgroup) {
1699 assert(numSg > 0 &&
1700 "Number of subgroups must be provided for sg layout creation.");
1701 auto sgLayouts = getSgLayoutCandidates(dataShape, *instData, numSg);
1702 if (sgLayouts.empty())
1703 return nullptr;
1704 return buildSgLayout(context, dataShape, sgLayouts.front(), /*dimK=*/-1);
1705 }
1706
1707 return nullptr;
1708}
1709
1710/// Sets up the anchor layout for a prefetch_nd operation. PrefetchNd has no
1711/// consumer (it produces no value), so it picks its own layout from uArch
1712/// block parameters.
1713xegpu::DistributeLayoutAttr
1715 xegpu::TensorDescType tdescTy, int numSg,
1716 const xegpu::uArch::uArch *uArch) {
1717
1718 const auto *uArchInstruction =
1719 dyn_cast<xegpu::uArch::Subgroup2DBlockPrefetchInstruction>(
1720 uArch->getInstruction(
1722 if (!uArchInstruction)
1723 return nullptr;
1724
1725 auto context = tdescTy.getContext();
1726 Type elemTy = tdescTy.getElementType();
1727 auto subgroupSize = uArch->getSubgroupSize();
1728 auto dataShape = tdescTy.getShape();
1729 [[maybe_unused]] int rank = tdescTy.getRank();
1730 assert(rank >= 2 && "Expected at least 2D shape for ND op");
1731
1732 // Compute the default 2D block IO lane layout / lane data.
1733 unsigned bitwidth = elemTy.getIntOrFloatBitWidth();
1734 auto [laneLayout, laneData, order] =
1735 compute2DBlockIOLaneLayout(dataShape, subgroupSize, bitwidth,
1736 uArchInstruction->getPackedFormatBitSize());
1737
1738 if (layoutKind == xegpu::LayoutKind::Lane)
1739 return buildLaneLayout(context, laneLayout, laneData);
1740
1741 auto instData =
1742 get2DBlockIOInstDataLayout(dataShape, elemTy, uArchInstruction);
1743 // Shape not realizable as a 2D-block instruction; let the caller report it.
1744 if (!instData)
1745 return nullptr;
1746
1747 if (layoutKind == xegpu::LayoutKind::InstData) {
1748 assert(isValidLaneLayout(*instData, laneLayout, laneData) &&
1749 "Expected the prefetch layout to satisfy uArch block constraints");
1750 return buildInstDataLayoutWithLane(context, *instData, laneLayout,
1751 laneData);
1752 }
1753
1754 if (layoutKind == xegpu::LayoutKind::Subgroup) {
1755 assert(numSg > 0 &&
1756 "Number of subgroups must be provided for sg layout creation.");
1757 auto sgLayouts = getSgLayoutCandidates(dataShape, *instData, numSg);
1758 if (sgLayouts.empty())
1759 return nullptr;
1760 return buildSgLayout(context, dataShape, sgLayouts.front(), /*dimK=*/-1);
1761 }
1762
1763 return nullptr;
1764}
1765
1766/// Sets up the anchor layout for a load_nd operation. LoadNd takes a
1767/// consumer layout (from its result's downstream uses) and validates it
1768/// against uArch constraints; if valid, the consumer's `inst_data` /
1769/// `sg_layout` are honored. Otherwise the helper falls back to defaults
1770/// derived from uArch block parameters.
1771xegpu::DistributeLayoutAttr
1773 VectorType resVecTy,
1774 xegpu::DistributeLayoutAttr consumerLayout,
1775 int numSg, const xegpu::uArch::uArch *uArch) {
1776
1777 assert(consumerLayout && "Expected a valid consumer layout");
1778 if (layoutKind == xegpu::LayoutKind::Subgroup) {
1779 assert(consumerLayout.isForWorkgroup() &&
1780 "Expected consumer layout to be a complete workgroup-level layout");
1781 return consumerLayout;
1782 }
1783
1784 auto context = resVecTy.getContext();
1785 Type elemTy = resVecTy.getElementType();
1786 auto subgroupSize = uArch->getSubgroupSize();
1787 auto dataShape = resVecTy.getShape();
1788 const auto *uArchInstruction =
1789 dyn_cast<xegpu::uArch::Subgroup2DBlockLoadInstruction>(
1790 uArch->getInstruction(
1792 if (!uArchInstruction)
1793 return nullptr;
1794
1795 int rank = resVecTy.getRank();
1796 SmallVector<int64_t> consumerInstData =
1797 consumerLayout.getEffectiveInstDataAsInt();
1798 SmallVector<int64_t> consumerLaneLayout =
1799 consumerLayout.getEffectiveLaneLayoutAsInt();
1800 SmallVector<int64_t> consumerLaneData =
1801 consumerLayout.getEffectiveLaneDataAsInt();
1802
1803 assert(!consumerLaneLayout.empty() && !consumerLaneData.empty() &&
1804 "Expected consumer layout to have lane_layout and lane_data");
1805
1806 // vertical lane layout means that the blockload must be transposed
1807 // note scaleA on PVC has vertical lane layout even without transposed order
1808 // attr
1809 bool hasTranspose =
1810 consumerLaneLayout[rank - 2] > 1 && consumerLaneLayout[rank - 1] == 1;
1811 bool hasTransform = !hasTranspose && consumerLaneData[rank - 2] > 1 &&
1812 consumerLaneData[rank - 1] == 1;
1813 assert((consumerLaneData[rank - 2] == 1 || consumerLaneData[rank - 1] == 1) &&
1814 "Expected consumer lane data to have at most one non-unit dim");
1815
1816 if (layoutKind == xegpu::LayoutKind::InstData) {
1817 auto blockWHC = uArchInstruction->getBlockWidthHeightCount(
1818 elemTy, hasTransform, hasTranspose,
1819 /*upConv=*/false);
1820 if (!blockWHC)
1821 return nullptr;
1822 auto [bWidths, bHeights, bCounts] = blockWHC.value();
1823
1824 // lane_layout and lane_data are the block load's own, not the consumer's: a
1825 // consumer may ask for more elements per lane than the hardware can pack,
1826 // e.g. an f32 multi_reduction wanting 16.
1827 unsigned packingSize = hasTransform || hasTranspose
1829 : uArchInstruction->getPackedFormatBitSize();
1830 auto [laneLayout, laneData, order] = compute2DBlockIOLaneLayout(
1831 dataShape, subgroupSize, elemTy.getIntOrFloatBitWidth(), packingSize,
1832 hasTransform, hasTranspose);
1833 DenseI32ArrayAttr orderAttr =
1834 order.empty()
1835 ? nullptr
1837 context, SmallVector<int32_t>(order.begin(), order.end()));
1838
1839 // See whether the consumer's inst_data satisfies the block constraints.
1840 int64_t height = consumerInstData[rank - 2];
1841 int64_t width = consumerInstData[rank - 1];
1842 auto maxBlockCount = *llvm::max_element(bCounts);
1843 auto maxWidth = *llvm::max_element(bWidths);
1844 if (llvm::is_contained(bWidths, static_cast<int>(width)) ||
1845 (width % maxWidth == 0 && width / maxWidth < maxBlockCount)) {
1846 if (llvm::is_contained(bHeights, static_cast<int>(height))) {
1847 assert(isValidLaneLayout(consumerInstData, laneLayout, laneData) &&
1848 "Expected the load layout to satisfy uArch block constraints");
1849 return buildInstDataLayoutWithLane(context, consumerInstData,
1850 laneLayout, laneData, orderAttr);
1851 }
1852 }
1853
1854 // The consumer's inst_data is not a usable block, so take the uArch's.
1855 // DPAS_MX scales land here. A scale whose innermost dim is too narrow to
1856 // spread packed lanes over (width 16 for an 8-bit scale needing 16 lanes of
1857 // 2 elements) cannot be loaded regularly; transformed, it packs down the
1858 // column instead. A vertical lane_layout is already transposed.
1859 if (!hasTranspose && !hasTransform &&
1860 consumerInstData[rank - 1] %
1861 (laneLayout[rank - 1] * laneData[rank - 1]) !=
1862 0) {
1863 hasTransform = true;
1864 std::tie(laneLayout, laneData, order) = compute2DBlockIOLaneLayout(
1865 dataShape, subgroupSize, elemTy.getIntOrFloatBitWidth(),
1866 uArch->getGeneralPackedFormatBitSize(), hasTransform, hasTranspose);
1867 }
1868
1869 auto instData = get2DBlockIOInstDataLayout(
1870 dataShape, elemTy, uArchInstruction, hasTransform, hasTranspose);
1871 // Shape not realizable as a 2D-block instruction; let the caller report it.
1872 if (!instData)
1873 return nullptr;
1874 assert(isValidLaneLayout(*instData, laneLayout, laneData) &&
1875 "Expected the load layout to satisfy uArch block constraints");
1876 return buildInstDataLayoutWithLane(context, *instData, laneLayout, laneData,
1877 orderAttr);
1878 }
1879 if (layoutKind == xegpu::LayoutKind::Lane) {
1880 assert(isValidLaneLayout(dataShape, consumerLaneLayout, consumerLaneData) &&
1881 "Expected the lane layout to satisfy uArch block constraints");
1882 return consumerLayout;
1883 }
1884 return nullptr;
1885}
1886
1887/// Sets up the anchor layout for load gather and load matrix operation.
1888/// load matrix lowers to load gather and 1d block load. All of them share the
1889/// same layout setup logic.
1890///
1891/// For Subgroup layout, uses the consumer layout directly.
1892///
1893/// For InstData layout, takes consumer's inst_data as-is; lane_layout and
1894/// lane_data are taken from the consumer.
1895///
1896/// For Lane layout, lane_layout/lane_data are taken from the consumer.
1897///
1898/// A consumer layout that carries lane_layout and lane_data is required.
1899/// `maxChunkSize` is not read yet: the consumer's lane_data already fixes the
1900/// per-lane chunk.
1901///
1902/// TODO: derive lane_layout/lane_data here when the consumer has none, capped
1903/// by `maxChunkSize`, the way `setupGenericStoreAnchorLayout` does via
1904/// `computeScatterIOLaneLayoutAndData`. That path is missing today, so the
1905/// assert below stands in for it.
1906static xegpu::DistributeLayoutAttr setupGenericLoadAnchorLayout(
1907 xegpu::LayoutKind layoutKind, mlir::MLIRContext *context,
1908 xegpu::DistributeLayoutAttr consumerLayout, int maxChunkSize,
1909 ArrayRef<int64_t> resShape, int subgroupSize) {
1910
1911 if (layoutKind == xegpu::LayoutKind::Subgroup)
1912 return consumerLayout;
1913
1914 SmallVector<int64_t> consumerInstData =
1915 consumerLayout.getEffectiveInstDataAsInt();
1916 SmallVector<int64_t> consumerLaneLayout =
1917 consumerLayout.getEffectiveLaneLayoutAsInt();
1918 SmallVector<int64_t> consumerLaneData =
1919 consumerLayout.getEffectiveLaneDataAsInt();
1920
1921 SmallVector<int64_t> laneLayout;
1922 SmallVector<int64_t> laneData;
1923 assert(!consumerLaneLayout.empty() && !consumerLaneData.empty() &&
1924 "Expected consumer layout to have lane_layout and lane_data");
1925 laneLayout.assign(consumerLaneLayout.begin(), consumerLaneLayout.end());
1926 laneData.assign(consumerLaneData.begin(), consumerLaneData.end());
1927
1928 if (layoutKind == xegpu::LayoutKind::InstData) {
1929 SmallVector<int64_t> instData;
1930 instData.resize(resShape.size());
1931 for (size_t i = 0; i < resShape.size(); ++i)
1932 instData[i] = laneLayout[i] * laneData[i];
1933 return buildInstDataLayoutWithLane(context, instData, laneLayout, laneData);
1934 }
1935 if (layoutKind == xegpu::LayoutKind::Lane)
1936 return buildLaneLayout(context, laneLayout, laneData);
1937 return nullptr;
1938}
1939
1940/// Sets up the anchor layout for a load gather operation.
1941xegpu::DistributeLayoutAttr xegpu::setupLoadGatherAnchorLayout(
1942 xegpu::LayoutKind layoutKind, VectorType resVecTy, int contigChunkSize,
1943 xegpu::DistributeLayoutAttr consumerLayout, const uArch::uArch *uArch) {
1944
1945 const int subgroupSize = uArch->getSubgroupSize();
1946 ArrayRef<int64_t> resShape = resVecTy.getShape();
1947 auto context = resVecTy.getContext();
1948
1949 // The per-lane chunk is bounded by what the offsets prove contiguous
1950 // (`contigChunkSize`) and by what one lane can access in one instruction.
1951 const auto *uArchInstruction = dyn_cast<xegpu::uArch::LoadGatherInstruction>(
1952 uArch->getInstruction(xegpu::uArch::InstructionKind::LoadGather));
1953 int maxChunkSize =
1954 std::min(uArchInstruction->getMaxLaneAccessSizeBytes(), contigChunkSize);
1955
1956 return setupGenericLoadAnchorLayout(layoutKind, context, consumerLayout,
1957 maxChunkSize, resShape, subgroupSize);
1958}
1959
1960/// Sets up the anchor layout for load matrix operation.
1961/// TODO: enhance load matrix to indicate lowering to chunked load or not.
1962xegpu::DistributeLayoutAttr
1964 VectorType resVecTy, int contigChunkSize,
1965 xegpu::DistributeLayoutAttr consumerLayout,
1966 const xegpu::uArch::uArch *uArch) {
1967
1968 const int subgroupSize = uArch->getSubgroupSize();
1969 ArrayRef<int64_t> resShape = resVecTy.getShape();
1970 auto context = resVecTy.getContext();
1971
1972 const auto *uArchInstruction = dyn_cast<xegpu::uArch::LoadGatherInstruction>(
1974 int maxChunkSize =
1975 std::min(uArchInstruction->getMaxLaneAccessSizeBytes(), contigChunkSize);
1976 return setupGenericLoadAnchorLayout(layoutKind, context, consumerLayout,
1977 maxChunkSize, resShape, subgroupSize);
1978}
1979
1980/// Picks the subgroup layout for a scatter-style store (store_scatter /
1981/// store_matrix): the most balanced `numSg` factorization that divides
1982/// `wgShape` with sg_data a multiple of `instData`. A store has no consumer.
1983static xegpu::DistributeLayoutAttr
1985 ArrayRef<int64_t> instData, int numSg) {
1986 auto candidates = getSgLayoutCandidates(wgShape, instData, numSg);
1987 if (candidates.empty())
1988 return nullptr;
1989 // Candidates are ordered most-balanced first.
1990 return buildSgLayout(context, wgShape, candidates.front(), /*dimK=*/-1);
1991}
1992
1993/// Sets up the anchor layout for store scatter and store matrix operation,
1994/// which share the same logic. Lane layout comes from
1995/// `computeScatterIOLaneLayoutAndData`; inst_data is lane_layout * lane_data.
1996static xegpu::DistributeLayoutAttr setupGenericStoreAnchorLayout(
1997 xegpu::LayoutKind layoutKind, mlir::MLIRContext *context, int maxChunkSize,
1998 ArrayRef<int64_t> srcShape, int subgroupSize, int numSg) {
1999
2000 auto [laneLayout, laneData] =
2001 computeScatterIOLaneLayoutAndData(srcShape, subgroupSize, maxChunkSize);
2002
2003 SmallVector<int64_t> instData(srcShape.size());
2004 for (size_t i = 0; i < srcShape.size(); ++i)
2005 instData[i] = laneLayout[i] * laneData[i];
2006
2007 if (layoutKind == xegpu::LayoutKind::Subgroup) {
2008 assert(numSg > 0 &&
2009 "Number of subgroups must be provided for sg layout creation.");
2010 return getStoreSubgroupLayouts(context, srcShape, instData, numSg);
2011 }
2012 if (layoutKind == xegpu::LayoutKind::InstData) {
2013 return buildInstDataLayoutWithLane(context, instData, laneLayout, laneData);
2014 }
2015 if (layoutKind == xegpu::LayoutKind::Lane) {
2016 return buildLaneLayout(context, laneLayout, laneData);
2017 }
2018 return nullptr;
2019}
2020
2021/// Sets up the anchor layout for a store scatter operation.
2022xegpu::DistributeLayoutAttr
2024 VectorType srcVecTy, int contigChunkSize,
2025 int numSg, const uArch::uArch *uArch) {
2026
2027 const int subgroupSize = uArch->getSubgroupSize();
2028 ArrayRef<int64_t> srcShape = srcVecTy.getShape();
2029 auto context = srcVecTy.getContext();
2030
2031 // The per-lane chunk is bounded by what the offsets prove contiguous
2032 // (`contigChunkSize`) and by what one lane can access in one instruction.
2033 const auto *uArchInstruction =
2034 dyn_cast<xegpu::uArch::StoreScatterInstruction>(
2036 int maxChunkSize =
2037 std::min(uArchInstruction->getMaxLaneAccessSizeBytes(), contigChunkSize);
2038 return setupGenericStoreAnchorLayout(layoutKind, context, maxChunkSize,
2039 srcShape, subgroupSize, numSg);
2040}
2041
2042/// Sets up the anchor layout for a store matrix operation.
2043xegpu::DistributeLayoutAttr xegpu::setupStoreMatrixAnchorLayout(
2044 xegpu::LayoutKind layoutKind, VectorType srcVecTy, int contigChunkSize,
2045 int numSg, const xegpu::uArch::uArch *uArch) {
2046
2047 const int subgroupSize = uArch->getSubgroupSize();
2048 ArrayRef<int64_t> srcShape = srcVecTy.getShape();
2049 auto context = srcVecTy.getContext();
2050
2051 const auto *uArchInstruction =
2052 dyn_cast<xegpu::uArch::StoreScatterInstruction>(
2054 int maxChunkSize =
2055 std::min(uArchInstruction->getMaxLaneAccessSizeBytes(), contigChunkSize);
2056
2057 return setupGenericStoreAnchorLayout(layoutKind, context, maxChunkSize,
2058 srcShape, subgroupSize, numSg);
2059}
2060
2061/// Completes a scatter IO layout by deriving lane_layout and lane_data from
2062/// `specifiedLayout`'s inst_data when they are missing. The layout is returned
2063/// unchanged if `specifiedLayout` is null, carries no inst_data, or already has
2064/// both lane_layout and lane_data.
2065///
2066/// When lane info is absent, inst_data is treated as the effective shape and
2067/// the lane factorization is filled in as follows:
2068/// - If `consumerLayout` is present and its lane_layout / lane_data are a
2069/// valid factorization of inst_data, that consumer lane info is reused so
2070/// the completed layout matches the consumer (avoiding a relayout).
2071/// - Otherwise a standard scatter-style factorization is computed via
2072/// `computeScatterIOLaneLayoutAndData`, bounded by `maxChunkSize` — the
2073/// per-lane load width reported by the uArch's LoadGather instruction
2074/// (`getMaxLaneAccessSizeBytes`).
2075///
2076std::optional<xegpu::DistributeLayoutAttr>
2078 xegpu::DistributeLayoutAttr specifiedLayout,
2079 xegpu::DistributeLayoutAttr consumerLayout, Type elemTy,
2080 const xegpu::uArch::LoadGatherInstruction *uArchInstruction,
2081 const int subgroupSize) {
2082 if (!specifiedLayout)
2083 return specifiedLayout;
2084 SmallVector<int64_t> specifiedInstData =
2085 specifiedLayout.getEffectiveInstDataAsInt();
2086 if (specifiedInstData.empty())
2087 return specifiedLayout;
2088 if (!specifiedLayout.getEffectiveLaneLayoutAsInt().empty() &&
2089 !specifiedLayout.getEffectiveLaneDataAsInt().empty())
2090 return specifiedLayout;
2091
2092 // Reuse the load-side setup with inst_data as the destination shape.
2093 auto *context = specifiedLayout.getContext();
2094 int maxChunkSize = uArchInstruction->getMaxLaneAccessSizeBytes();
2095 if (consumerLayout) {
2096 auto consumerLaneLayout = consumerLayout.getEffectiveLaneLayoutAsInt();
2097 auto consumerLaneData = consumerLayout.getEffectiveLaneDataAsInt();
2098 if (!consumerLaneLayout.empty() && !consumerLaneData.empty() &&
2099 isValidLaneLayout(specifiedInstData, consumerLaneLayout,
2100 consumerLaneData))
2101 return buildInstDataLayoutWithLane(context, specifiedInstData,
2102 consumerLaneLayout, consumerLaneData);
2103 }
2104 auto [defLaneLayout, defLaneData] = computeScatterIOLaneLayoutAndData(
2105 specifiedInstData, subgroupSize, maxChunkSize);
2106 if (!isValidLaneLayout(specifiedInstData, defLaneLayout, defLaneData))
2107 return std::nullopt;
2108 return buildInstDataLayoutWithLane(context, specifiedInstData, defLaneLayout,
2109 defLaneData);
2110}
2111
2112/// Like completeScatterLoadLaneLayoutFromInstData, but for scatter stores. A
2113/// store is a data sink, so lane info is derived purely from inst_data (bounded
2114/// by the uArch's per-lane store width); there is no consumer layout to reuse.
2115std::optional<xegpu::DistributeLayoutAttr>
2117 xegpu::DistributeLayoutAttr specifiedLayout, Type elemTy,
2118 const xegpu::uArch::StoreScatterInstruction *uArchInstruction,
2119 const int subgroupSize) {
2120 if (!specifiedLayout)
2121 return specifiedLayout;
2122 SmallVector<int64_t> specifiedInstData =
2123 specifiedLayout.getEffectiveInstDataAsInt();
2124 if (specifiedInstData.empty())
2125 return specifiedLayout;
2126 if (!specifiedLayout.getEffectiveLaneLayoutAsInt().empty() &&
2127 !specifiedLayout.getEffectiveLaneDataAsInt().empty())
2128 return specifiedLayout;
2129
2130 // Reuse the store-side setup with inst_data as the source shape.
2131 auto *context = specifiedLayout.getContext();
2132 int maxChunkSize = uArchInstruction->getMaxLaneAccessSizeBytes();
2133 auto [defLaneLayout, defLaneData] = computeScatterIOLaneLayoutAndData(
2134 specifiedInstData, subgroupSize, maxChunkSize);
2135 if (!isValidLaneLayout(specifiedInstData, defLaneLayout, defLaneData))
2136 return std::nullopt;
2137 return buildInstDataLayoutWithLane(context, specifiedInstData, defLaneLayout,
2138 defLaneData);
2139}
2140
2141/// Completes a 2D-block store/prefetch layout from its inst_data. store_nd and
2142/// prefetch_nd are data sinks, so lane info is derived purely from inst_data
2143/// (no consumer to reuse). One helper serves both via
2144/// BlockIOInstructionInterface.
2145std::optional<xegpu::DistributeLayoutAttr>
2147 xegpu::DistributeLayoutAttr specifiedLayout, Type elemTy,
2148 const xegpu::uArch::BlockIOInstructionInterface *uArchInstruction,
2149 const int subgroupSize) {
2150 if (!specifiedLayout)
2151 return specifiedLayout;
2152 SmallVector<int64_t> specifiedInstData =
2153 specifiedLayout.getEffectiveInstDataAsInt();
2154 if (specifiedInstData.empty())
2155 return specifiedLayout;
2156 if (!specifiedLayout.getEffectiveLaneLayoutAsInt().empty() &&
2157 !specifiedLayout.getEffectiveLaneDataAsInt().empty())
2158 return specifiedLayout;
2159
2160 auto *context = specifiedLayout.getContext();
2161 auto [laneLayout, laneData, order] = compute2DBlockIOLaneLayout(
2162 specifiedInstData, subgroupSize, elemTy.getIntOrFloatBitWidth(),
2163 uArchInstruction->getPackedFormatBitSize());
2164 if (!isValidLaneLayout(specifiedInstData, laneLayout, laneData))
2165 return std::nullopt;
2166 return buildInstDataLayoutWithLane(context, specifiedInstData, laneLayout,
2167 laneData);
2168}
2169
2170/// Like completeBlockStoreLaneLayoutFromInstData, but for load_nd. The
2171/// consumer's lane_data and order are reused as-is; lane_layout is rebuilt from
2172/// the consumer's lane_layout, bumping every non-unit dim up to the subgroup
2173/// size. The user-provided inst_data is preserved.
2174std::optional<xegpu::DistributeLayoutAttr>
2176 xegpu::DistributeLayoutAttr specifiedLayout,
2177 xegpu::DistributeLayoutAttr consumerLayout, Type elemTy,
2178 const xegpu::uArch::BlockIOInstructionInterface *uArchInstruction,
2179 const int subgroupSize) {
2180 if (!specifiedLayout)
2181 return specifiedLayout;
2182 SmallVector<int64_t> specifiedInstData =
2183 specifiedLayout.getEffectiveInstDataAsInt();
2184 if (specifiedInstData.empty())
2185 return specifiedLayout;
2186 if (!specifiedLayout.getEffectiveLaneLayoutAsInt().empty() &&
2187 !specifiedLayout.getEffectiveLaneDataAsInt().empty())
2188 return specifiedLayout;
2189 if (!consumerLayout)
2190 return specifiedLayout;
2191 SmallVector<int64_t> consumerLaneLayout =
2192 consumerLayout.getEffectiveLaneLayoutAsInt();
2193 SmallVector<int64_t> consumerLaneData =
2194 consumerLayout.getEffectiveLaneDataAsInt();
2195 if (consumerLaneLayout.empty() || consumerLaneData.empty())
2196 return specifiedLayout;
2197
2198 auto *context = specifiedLayout.getContext();
2199 int rank = specifiedInstData.size();
2200
2201 SmallVector<int64_t> laneLayout;
2202 // set the laneLayout to use consumer's LaneLayout as base, but adjust its
2203 // size to match the subgroupsize in case its original value is larger than 1
2204 for (int i = 0; i < rank; i++) {
2205 if (consumerLaneLayout[i] > 1) {
2206 laneLayout.push_back(
2207 std::max(static_cast<int64_t>(subgroupSize), consumerLaneLayout[i]));
2208 } else {
2209 laneLayout.push_back(1);
2210 }
2211 }
2212
2213 if (!isValidLaneLayout(specifiedInstData, laneLayout, consumerLaneData))
2214 return std::nullopt;
2215 return buildInstDataLayoutWithLane(context, specifiedInstData, laneLayout,
2216 consumerLaneData,
2217 consumerLayout.getOrder());
2218}
2219
2220/// Completes user-provided DPAS A/B/C-D anchors that carry only inst_data by
2221/// filling in lane_layout / lane_data. The lane factorization mirrors the
2222/// InstData branch of `setupDpasLayout` (derived from each operand's shape and
2223/// matmul role, B using VNNI packing); the user's inst_data is preserved.
2224std::optional<
2225 std::tuple<xegpu::DistributeLayoutAttr, xegpu::DistributeLayoutAttr,
2226 xegpu::DistributeLayoutAttr>>
2227xegpu::completeDpasLaneLayoutFromInstData(xegpu::DistributeLayoutAttr aLayout,
2228 xegpu::DistributeLayoutAttr bLayout,
2229 xegpu::DistributeLayoutAttr cdLayout,
2230 VectorType aTy, VectorType bTy,
2231 VectorType cdTy,
2232 const xegpu::uArch::uArch *uArch) {
2233 auto context = aTy.getContext();
2234 const auto *uArchInstruction =
2235 dyn_cast<xegpu::uArch::SubgroupMatrixMultiplyAcc>(uArch->getInstruction(
2237 if (!uArchInstruction)
2238 return std::nullopt;
2239 auto subgroupSize = uArch->getSubgroupSize();
2240 llvm::SmallVector<int64_t> laneLayoutA, laneDataA, laneLayoutB, laneDataB,
2241 laneLayoutCD, laneDataCD;
2242 // DPAS operands are never transposed, so the order comes back empty.
2243 llvm::SmallVector<int64_t> orderA, orderB, orderCD;
2244 SmallVector<int64_t> instDataA = aLayout.getEffectiveInstDataAsInt();
2245 SmallVector<int64_t> instDataB = bLayout.getEffectiveInstDataAsInt();
2246 SmallVector<int64_t> instDataCD = cdLayout.getEffectiveInstDataAsInt();
2247
2248 if (isa<xegpu::uArch::Xe2, xegpu::uArch::Xe3>(uArch)) {
2249 std::tie(laneLayoutA, laneDataA, orderA) =
2250 compute2DBlockIOLaneLayout(aTy.getShape(), subgroupSize,
2251 aTy.getElementType().getIntOrFloatBitWidth(),
2252 uArchInstruction->getPackedFormatBitSizeA());
2253 std::tie(laneLayoutB, laneDataB, orderB) = compute2DBlockIOLaneLayout(
2254 bTy.getShape(), subgroupSize,
2255 bTy.getElementType().getIntOrFloatBitWidth(),
2256 uArchInstruction->getPackedFormatBitSizeB(), /*vnni=*/true);
2257 std::tie(laneLayoutCD, laneDataCD, orderCD) = compute2DBlockIOLaneLayout(
2258 cdTy.getShape(), subgroupSize,
2259 cdTy.getElementType().getIntOrFloatBitWidth(),
2260 cdTy.getElementType().getIntOrFloatBitWidth());
2261 } else {
2262 assert(false && "Unsupported uArch for DPAS lane layout completion");
2263 }
2264
2265 if (!isValidLaneLayout(instDataA, laneLayoutA, laneDataA) ||
2266 !isValidLaneLayout(instDataB, laneLayoutB, laneDataB) ||
2267 !isValidLaneLayout(instDataCD, laneLayoutCD, laneDataCD))
2268 return std::nullopt;
2269 return std::make_tuple(
2270 buildInstDataLayoutWithLane(context, instDataA, laneLayoutA, laneDataA,
2271 aLayout.getOrder()),
2272 buildInstDataLayoutWithLane(context, instDataB, laneLayoutB, laneDataB,
2273 bLayout.getOrder()),
2274 buildInstDataLayoutWithLane(context, instDataCD, laneLayoutCD, laneDataCD,
2275 cdLayout.getOrder()));
2276}
2277
2278/// Like completeDpasLaneLayoutFromInstData, but for dpas_mx: also re-derives
2279/// the A_scale / B_scale layouts from the completed A / B layouts via
2280/// `createScaleLayout`, matching the default path of `setupDpasMxLayout`.
2281std::optional<
2282 std::tuple<xegpu::DistributeLayoutAttr, xegpu::DistributeLayoutAttr,
2283 xegpu::DistributeLayoutAttr, xegpu::DistributeLayoutAttr,
2284 xegpu::DistributeLayoutAttr>>
2286 xegpu::DistributeLayoutAttr aLayout, xegpu::DistributeLayoutAttr bLayout,
2287 xegpu::DistributeLayoutAttr cdLayout, VectorType aTy, VectorType bTy,
2288 VectorType cdTy, VectorType aScaleTy, VectorType bScaleTy,
2289 const xegpu::uArch::uArch *uArch) {
2290 auto completed = completeDpasLaneLayoutFromInstData(
2291 aLayout, bLayout, cdLayout, aTy, bTy, cdTy, uArch);
2292 if (!completed)
2293 return std::nullopt;
2294 auto context = aTy.getContext();
2295 auto [completedA, completedB, completedCD] = *completed;
2296
2297 auto aScaleLayout =
2298 createScaleLayout(context, aTy, aScaleTy, completedA, false, uArch);
2299 auto bScaleLayout =
2300 createScaleLayout(context, bTy, bScaleTy, completedB, true, uArch);
2301
2302 return std::make_tuple(completedA, completedB, completedCD, aScaleLayout,
2303 bScaleLayout);
2304}
2305
2306/// Sets up layout for reduction operations by creating a SliceAttr for the
2307/// result.
2308///
2309/// Algorithm Overview:
2310/// This function attempts to construct a source layout that, when sliced along
2311/// reduction dimensions, produces a result layout compatible with the
2312/// consumer layout.
2313///
2314/// For subgroup layouts, it first tries to align the source layout's subgroup
2315/// layout and data with the consumer's layout on non-reduction dimensions.
2316/// Then, it distributes remaining subgroups across reduction dimensions. This
2317/// avoids subgroup data redistribution overhead between the reduced result and
2318/// its consumer. When the consumer layout is a slice layout, it attempts to
2319/// reuse the slice layout's parent layout for the source to further minimize
2320/// potential data redistribution.
2321///
2322/// This is a best-effort alignment, not a hard constraint: the goal is only to
2323/// pick a *legal* source layout that minimizes redistribution against the
2324/// (single, first-arriving) consumer layout. There is no failure path - when
2325/// the consumer's slice layout cannot be reused as-is (example 2 below), the
2326/// function falls back to distributing all subgroups on the non-reduction
2327/// dimensions first and the remainder on the reduction dimensions, which always
2328/// yields a valid source layout. If the resulting source layout still differs
2329/// from what some consumer expects (e.g. a second, inconsistent consumer), that
2330/// mismatch is reconciled later by the layout conflict resolution process
2331/// (`ResolveLayoutConflicts`), which inserts a `convert_layout` op - this
2332/// function never has to give up.
2333///
2334/// For the InstData and Lane layout kinds only the innermost two dimensions
2335/// are distributed; all leading dimensions are assumed to be unit dimensions.
2336/// This assumption is checked via `leadingDimsAreUnit`. The lane_layout and
2337/// lane_data are computed by `computeReductionLaneLayoutAndData`, which picks
2338/// a layout that minimizes cross-lane reduction (reducing within a lane when
2339/// only one of the innermost two dims is a reduction dim). The inst_data is
2340/// simply the element-wise product lane_layout * lane_data.
2341///
2342/// The function returns the *result* layout (the SliceAttr). The *source*
2343/// layout it decides on is the parent of that slice; both are listed below so
2344/// the relationship is explicit.
2345///
2346/// Examples:
2347/// 1. Subgroup layout - Row reduction on 2D tensor:
2348/// srcShape=[32, 128], reductionDims=[1], resShape=[32], subgroupSize=16,
2349/// NumSg=32
2350/// * Consumer Layout:
2351/// #xegpu.slice<#xegpu.layout<sg_layout=[4, 8], sg_data=[8, 8]>, dims =
2352/// [1]>}
2353/// * Source Layout (decided by this function):
2354/// #xegpu.layout<sg_layout=[4, 8], sg_data=[8, 16]>
2355/// * Result Layout (returned):
2356/// #xegpu.slice<#xegpu.layout<sg_layout=[4, 8], sg_data=[8, 16]>, dims =
2357/// [1]>}
2358/// The consumer slices exactly the reduction dim, so its parent layout is
2359/// reused for the source: sg_layout is kept, but the source's sg_data on
2360/// the reduction dim is grown from 8 to 16 (= srcShape[1] / sg_layout[1] =
2361/// 128 / 8) so the source tile is evenly distributed over the reduction
2362/// dim. Slicing that source over dim 1 reproduces the consumer.
2363///
2364/// 2. Subgroup layout - Same shapes as above but consumer doesn't have a
2365/// reusable slice layout, so the algorithm distributes all subgroups on the
2366/// non-reduction dims first and the remainder on the reduction dims.
2367/// 2a. * Consumer Layout:
2368/// #xegpu.layout<sg_layout=[32], sg_data=[1]>
2369/// * Source Layout (decided by this function):
2370/// #xegpu.layout<sg_layout=[32, 1], sg_data=[1, 128]>
2371/// * Result Layout (returned):
2372/// #xegpu.slice<#xegpu.layout<sg_layout=[32, 1], sg_data=[1, 128]>,
2373/// dims = [1]>}
2374/// All 32 subgroups land on the non-reduction dim 0; the reduction dim
2375/// 1 gets the leftover (sg_layout=1, so the whole length 128 lives in
2376/// one subgroup's sg_data).
2377/// 2b. * Consumer Layout:
2378/// #xegpu.slice<#xegpu.layout<sg_layout=[8, 2, 4], sg_data=[4, 64,
2379/// 32]>, dims = [1, 2]>}
2380/// * Source Layout (decided by this function):
2381/// #xegpu.layout<sg_layout=[8, 4], sg_data=[4, 32]>
2382/// * Result Layout (returned):
2383/// #xegpu.slice<#xegpu.layout<sg_layout=[8, 4], sg_data=[4, 32]>,
2384/// dims = [1]>}
2385/// The consumer slices dims [1, 2] which do not match this op's
2386/// reductionDims, so it can't be reused as-is; subgroups are
2387/// re-distributed (non-reduction dim first, then reduction dim).
2388///
2389/// 3. Lane layout - Default (lanes on innermost dim):
2390/// srcShape=[32, 64], reductionDims=[0], subgroupSize=16
2391/// * Source Layout (decided by this function):
2392/// laneLayout=[1, 16], laneData=[1, 1] (returned sliced over dim 0).
2393/// The innermost dim is not reduced, so lanes stay on it.
2394///
2395/// 4. Lane layout - Switch (lanes moved off the reduction dim):
2396/// srcShape=[32, 64], reductionDims=[1], subgroupSize=16
2397/// * Source Layout (decided by this function):
2398/// laneLayout=[16, 1], laneData=[1, 1] (returned sliced over dim 1).
2399/// The innermost dim is the sole reduction dim, so lanes move to the
2400/// non-reduction dim to reduce within a lane. This switch only happens
2401/// when the consumer has no reduction dims to broadcast the result back
2402/// along (i.e. the consumer layout is not a slice over this reduction);
2403/// otherwise the default (example 3) is used.
2404///
2405/// 5. Lane layout - No switch when both inner dims are reduced (reduction to
2406/// scalar):
2407/// srcShape=[32, 64], reductionDims=[0, 1], subgroupSize=16
2408/// * Source Layout (decided by this function):
2409/// laneLayout=[1, 16], laneData=[1, 1] (returned sliced over dims
2410/// [0,1]).
2411/// Both dims are reduced, so this is not a *sole* innermost reduction; the
2412/// switch condition (example 4) does not apply and lanes stay on the
2413/// innermost dim. The cross-lane reduction here is unavoidable.
2414///
2415/// 6. Lane layout - No switch when the consumer slices the reduction dim:
2416/// srcShape=[32, 64], reductionDims=[1], subgroupSize=16
2417/// * Consumer Layout:
2418/// #xegpu.slice<#xegpu.layout<laneLayout=[1, 16], laneData=[1, 1]>,
2419/// dims = [1]>}
2420/// * Source Layout (decided by this function):
2421/// #xegpu.layout<laneLayout=[1, 16], laneData=[1, 1]> (the consumer
2422/// slice's parent, reused directly; returned sliced over dim 1).
2423/// Same shape/reductionDims as example 4, but here the consumer is a slice
2424/// over the reduction dim, so it can broadcast the result back along that
2425/// dim. The slice's parent layout is reused as the source (no switch, no
2426/// re-derivation); the inst_data propagation step has already inserted a
2427/// convert_layout if needed, so the lane-level layout can be reused as-is.
2428
2430 xegpu::LayoutKind layoutKind, VectorType srcVecTy,
2431 DistributeLayoutAttr consumerLayout, SmallVector<int64_t> reductionDims,
2432 int numSg, const xegpu::uArch::uArch *uArch) {
2433
2434 auto srcShape = srcVecTy.getShape();
2435 int srcRank = srcShape.size();
2436 auto context = srcVecTy.getContext();
2437
2438 const int subgroupSize = uArch->getSubgroupSize();
2439 int64_t maxReduceVectorSize = 1; // could extend to spirv vector Size
2440 xegpu::DistributeLayoutAttr srcLayout;
2441 if (layoutKind == xegpu::LayoutKind::Subgroup) {
2442 xegpu::SliceAttr consumerSliceLayout =
2443 dyn_cast_if_present<xegpu::SliceAttr>(consumerLayout);
2444 if (consumerSliceLayout &&
2445 consumerSliceLayout.getDims().asArrayRef().equals(reductionDims)) {
2446 srcLayout = consumerSliceLayout.getParent();
2447 SmallVector<int64_t> sgLayoutFromConsumer =
2448 srcLayout.getEffectiveSgLayoutAsInt();
2449 auto srcSgData = computeShapeRatio(srcShape, sgLayoutFromConsumer);
2450 if (srcSgData)
2451 for (int dim = 0; dim < srcRank; dim++) {
2452 if (llvm::is_contained(reductionDims, dim))
2453 srcLayout =
2454 srcLayout.setDimData(dim, srcSgData.value()[dim], -1, -1);
2455 }
2456 } else {
2457 SmallVector<int64_t> consumerSgLayout =
2458 consumerLayout ? consumerLayout.getEffectiveSgLayoutAsInt()
2460 SmallVector<int64_t> consumerSgData =
2461 consumerLayout ? consumerLayout.getEffectiveSgDataAsInt()
2463 SmallVector<int64_t> consumerOrder =
2464 consumerLayout ? consumerLayout.getEffectiveOrderAsInt()
2466 DenseI32ArrayAttr orderAttr =
2467 consumerLayout ? consumerLayout.getOrder() : nullptr;
2468 SmallVector<int64_t> sgLayout(srcRank), sgData(srcRank), order(srcRank);
2469 int remainingSgCount =
2470 consumerLayout ? consumerLayout.getNumSubgroups() : numSg;
2471 int consumerIdx = 0;
2472
2473 // First pass: match the consumer's layout on non-reduction dims.
2474 for (int i = 0; i < srcRank; i++) {
2475 if (!llvm::is_contained(reductionDims, i) &&
2476 consumerIdx < static_cast<int>(consumerSgLayout.size())) {
2477 sgLayout[i] = consumerSgLayout[consumerIdx];
2478 sgData[i] = consumerSgData[consumerIdx];
2479 remainingSgCount /= sgLayout[i];
2480 consumerIdx++;
2481 }
2482 }
2483
2484 // Second pass: distribute remaining subgroups across the reduction dims.
2485 // The reduction-to-scalar case is handled only by this loop.
2486 for (int i = 0; i < srcRank; i++) {
2487 if (llvm::is_contained(reductionDims, i)) {
2488 sgLayout[i] =
2489 std::min(srcShape[i], static_cast<int64_t>(remainingSgCount));
2490 assert((srcShape[i] % sgLayout[i] == 0) &&
2491 "source shape not divisible by sg_layout");
2492 sgData[i] = srcShape[i] / sgLayout[i];
2493 remainingSgCount /= sgLayout[i];
2494 }
2495 }
2496 // The reduction dims are assumed to be row-major contiguous with the
2497 // left-hand neighbor.
2498 DenseI32ArrayAttr resOrderAttr = nullptr;
2499 int numRetainedDims = srcRank - static_cast<int>(reductionDims.size());
2500 if (orderAttr && !orderAttr.empty() &&
2501 static_cast<int>(consumerOrder.size()) == numRetainedDims) {
2502 // Fill known order
2503 int retainedRank = 0;
2504 for (int dim = 0; dim < srcRank; dim++)
2505 if (!llvm::is_contained(reductionDims, dim))
2506 order[dim] = consumerOrder[retainedRank++];
2507 // Each new order is placed as close as possible to its neighbors,
2508 // row-major. Re-index the order of higher walk orders.
2509 for (int reducedDim = 0; reducedDim < srcRank; reducedDim++) {
2510 if (!llvm::is_contained(reductionDims, reducedDim))
2511 continue;
2512 int64_t insertRank = 0;
2513 // Add as an inner dim to the outer neighbor
2514 if (reducedDim)
2515 insertRank = order[reducedDim - 1];
2516 // Add as an outer dim to the inner neighbor if no left neighbor
2517 // exists.
2518 else if (srcRank > 1)
2519 insertRank = order[reducedDim + 1] + 1;
2520 for (int otherDim = 0; otherDim < srcRank; otherDim++)
2521 if (order[otherDim] >= insertRank)
2522 order[otherDim]++;
2523 order[reducedDim] = insertRank;
2524 }
2525 resOrderAttr = DenseI32ArrayAttr::get(
2526 context, SmallVector<int32_t>(order.begin(), order.end()));
2527 }
2528 assert(remainingSgCount == 1 && "not all subgroups distributed");
2529 srcLayout = buildLayout(context, sgLayout, sgData,
2530 /*instData=*/{}, /*laneLayout=*/{},
2531 /*laneData=*/{}, resOrderAttr);
2532 }
2533 } else if (layoutKind == xegpu::LayoutKind::InstData) {
2534 xegpu::SliceAttr consumerSliceLayout =
2535 dyn_cast_if_present<xegpu::SliceAttr>(consumerLayout);
2536 auto consumerReductionDims =
2537 consumerSliceLayout
2538 ? SmallVector<int64_t>(consumerSliceLayout.getDims().asArrayRef())
2540 // A[i] reduced from A[i, j] is stored out directly, use vertical Lane
2541 // layout like [16, 1]
2542 bool verticalLaneLayout = consumerReductionDims.empty() &&
2543 reductionDims.size() == 1 &&
2544 reductionDims[0] == (srcRank - 1);
2545 auto [laneLayout, laneData] = computeReductionLaneLayoutAndData(
2546 srcShape, reductionDims, subgroupSize, maxReduceVectorSize,
2547 verticalLaneLayout);
2548 // inst_data is the per-instruction data, i.e. the element-wise product of
2549 // lane_layout and lane_data.
2550 SmallVector<int64_t> instData(srcRank);
2551 for (int i = 0; i < srcRank; i++)
2552 instData[i] = laneLayout[i] * laneData[i];
2553 srcLayout =
2554 buildInstDataLayoutWithLane(context, instData, laneLayout, laneData);
2555 } else if (layoutKind == xegpu::LayoutKind::Lane) {
2556 // Only the innermost two dimensions are distributed; all leading dimensions
2557 // are assumed to be unit dimensions.
2558 assert(leadingDimsAreUnit(srcShape, /*numInnerDims=*/2) &&
2559 "Lane reduction layout assumes all leading (non-innermost-two) "
2560 "dimensions are unit dimensions");
2561 xegpu::SliceAttr consumerSliceLayout =
2562 dyn_cast_if_present<xegpu::SliceAttr>(consumerLayout);
2563 auto consumerReductionDims =
2564 consumerSliceLayout
2565 ? SmallVector<int64_t>(consumerSliceLayout.getDims().asArrayRef())
2567 if (consumerSliceLayout &&
2568 consumerSliceLayout.getDims().asArrayRef().equals(reductionDims)) {
2569 // at the lane level, the consumerSliceLayout can be directly reused
2570 // since the inst_data propagation already insert convert_layout if
2571 // the layout is not consistent
2572 srcLayout = consumerSliceLayout.getParent();
2573 } else {
2574 bool verticalLaneLayout = consumerReductionDims.empty() &&
2575 reductionDims.size() == 1 &&
2576 reductionDims[0] == (srcRank - 1);
2577 auto [laneLayout, laneData] = computeReductionLaneLayoutAndData(
2578 srcShape, reductionDims, subgroupSize, maxReduceVectorSize,
2579 verticalLaneLayout);
2580 srcLayout = buildLaneLayout(context, laneLayout, laneData);
2581 }
2582 }
2583
2584 return xegpu::SliceAttr::get(context, srcLayout,
2585 DenseI64ArrayAttr::get(context, reductionDims));
2586}
2587
2588/// Sets up layout for Reduction operations by creating a SliceAttr for the
2589/// result.
2590xegpu::SliceAttr
2592 VectorType srcVecTy,
2593 const xegpu::uArch::uArch *uArch) {
2594
2595 auto srcShape = srcVecTy.getShape();
2596 auto context = srcVecTy.getContext();
2597 auto subgroupSize = uArch->getSubgroupSize();
2598 xegpu::LayoutAttr srcLayout;
2599
2600 if (layoutKind == xegpu::LayoutKind::Subgroup) {
2601 assert(false &&
2602 "subgroup layout assignment not supported for reduction (op "
2603 "is not expected at this level).");
2604 } else if (layoutKind == xegpu::LayoutKind::InstData) {
2605 assert(false &&
2606 "instData layout assignment not supported for reduction (op "
2607 "is not expected at this level).");
2608 } else if (layoutKind == xegpu::LayoutKind::Lane) {
2609 SmallVector<int64_t> laneLayout(1), laneData(1);
2610 laneLayout[0] = std::min(static_cast<int64_t>(subgroupSize), srcShape[0]);
2611 laneData[0] = 1;
2612 srcLayout = buildLaneLayout(context, laneLayout, laneData);
2613 }
2614
2615 auto result = xegpu::SliceAttr::get(context, srcLayout,
2616 DenseI64ArrayAttr::get(context, 0));
2617 return result;
2618}
2619
2620/// Adjusts `consumerLayout`'s innermost-dim data field selected by
2621/// `layoutKind` so that the source layout can be safely inferred by dividing
2622/// that value by `ratio`. Doubles the value until the divisibility constraint
2623/// is met, bounded above by `bound` like result-shape.
2624///
2625/// Used by ops whose source relates to the result by a fixed factor along the
2626/// innermost dim (e.g., bitcast: bitwidth ratio; interleave: 2x).
2627///
2628/// Divisibility constraints per LayoutKind:
2629/// - Subgroup: sgData[innermost] % ratio == 0
2630/// - InstData: instData[innermost] % (laneLayout[innermost] * ratio) == 0
2631/// (laneLayout falls back to subgroupSize if absent)
2632/// - Lane: laneData[innermost] % ratio == 0
2633static xegpu::DistributeLayoutAttr
2634adjustInnermostDimForDivisibility(xegpu::DistributeLayoutAttr consumerLayout,
2635 xegpu::LayoutKind layoutKind,
2636 size_t innerMostDim, int ratio, int64_t bound,
2637 const xegpu::uArch::uArch *uArch) {
2638 SmallVector<int64_t> sgData = consumerLayout.getEffectiveSgDataAsInt();
2639 SmallVector<int64_t> instData = consumerLayout.getEffectiveInstDataAsInt();
2640 SmallVector<int64_t> laneData = consumerLayout.getEffectiveLaneDataAsInt();
2641 SmallVector<int64_t> laneLayout =
2642 consumerLayout.getEffectiveLaneLayoutAsInt();
2643
2644 int64_t sgDataValue = -1;
2645 int64_t instDataValue = -1;
2646 int64_t laneDataValue = -1;
2647
2648 if (layoutKind == xegpu::LayoutKind::Subgroup) {
2649 sgDataValue = sgData[innerMostDim];
2650 while ((sgDataValue <= bound) && (sgDataValue % ratio) != 0)
2651 sgDataValue *= 2;
2652 } else if (layoutKind == xegpu::LayoutKind::InstData) {
2653 instDataValue = instData[innerMostDim];
2654 const int innermostDimLaneLayout = laneLayout.empty()
2655 ? uArch->getSubgroupSize()
2656 : laneLayout[innerMostDim];
2657 while ((instDataValue <= bound) &&
2658 (instDataValue % (innermostDimLaneLayout * ratio) != 0))
2659 instDataValue *= 2;
2660 assert((bound % instDataValue) == 0 &&
2661 "bound, instData, and laneLayout for innermost must be 2^n!");
2662 } else if (layoutKind == xegpu::LayoutKind::Lane) {
2663 laneDataValue = laneData[innerMostDim];
2664 while ((laneDataValue <= bound) && (laneDataValue % ratio) != 0)
2665 laneDataValue *= 2;
2666 }
2667
2668 return consumerLayout.setDimData(innerMostDim, sgDataValue, instDataValue,
2669 laneDataValue);
2670}
2671
2672/// Sets up the result layout for a bitcast operation.
2673/// When casting to a smaller bitwidth, adjusts the layout dimensions (sgData,
2674/// instData, or laneData) by multiplying by the bitwidth ratio to ensure the
2675/// result layout can be correctly divided back to the source layout during
2676/// inference.
2677///
2678/// Examples:
2679/// 1. Casting f32 -> f16 (32-bit to 16-bit, bitWidthRatio = 2):
2680/// Consumer layout: instData=[1, 16], subgroupSize=16
2681/// Source shape: [8, 32]
2682/// Result layout: instData=[1, 32] (16 * 2)
2683/// The innermost dimension is multiplied by 2 to maintain consistency.
2684///
2685/// 2. Casting f32 -> i8 (32-bit to 8-bit, bitWidthRatio = 4):
2686/// Consumer instData=[1, 16], subgroupSize=16
2687/// Source shape: [4, 128]
2688/// adjust the instData from [1, 16] to [1, 16 * 4 = 64]
2689///
2690/// 3. Casting i8 -> i32 (8-bit to 32-bit, bitWidthRatio = 1/4):
2691/// Consumer layout: laneLayout=[1, 16], laneData=[1, 4]
2692/// No adjustment needed - returns consumer layout directly.
2693///
2694xegpu::DistributeLayoutAttr xegpu::setupBitCastResultLayout(
2695 xegpu::LayoutKind layoutKind, VectorType srcVecTy, VectorType resVecTy,
2696 DistributeLayoutAttr consumerLayout, const xegpu::uArch::uArch *uArch) {
2697
2698 int srcElemTyBitWidth = srcVecTy.getElementType().getIntOrFloatBitWidth();
2699 int resElemTyBitWidth = resVecTy.getElementType().getIntOrFloatBitWidth();
2700
2701 ArrayRef<int64_t> srcShape = srcVecTy.getShape();
2702 ArrayRef<int64_t> resShape = resVecTy.getShape();
2703
2704 assert(consumerLayout.getRank() == static_cast<int64_t>(srcShape.size()) &&
2705 "laneData must be available for all dimensions");
2706
2707 // Casting to same/larger element type: result has fewer (or equal) elements
2708 // along the innermost dim, no adjustment needed.
2709 if (srcElemTyBitWidth <= resElemTyBitWidth)
2710 return consumerLayout;
2711
2712 // Casting to smaller element type: result has more elements along innermost
2713 // dim. Adjust the innermost data field upward so the source layout can be
2714 // recovered by dividing by bitWidthRatio.
2715 size_t innerMostDim = srcShape.size() - 1;
2716 int bitWidthRatio = srcElemTyBitWidth / resElemTyBitWidth;
2717 return adjustInnermostDimForDivisibility(consumerLayout, layoutKind,
2718 innerMostDim, bitWidthRatio,
2719 resShape[innerMostDim], uArch);
2720}
2721
2722/// Sets up the result layout for an interleave operation to ensure the source
2723/// layout can be safely derived. Interleave doubles the innermost dimension,
2724/// so the result layout must ensure that laneData is a multiple
2725/// of 2, and instData must be divisible by innermostDimLaneLayout * 2.
2726///
2727/// Example:
2728/// Interleave: vector<128x256xf4> -> vector<128x512xf4>
2729/// Consumer layout: laneLayout=[1, 16], laneData=[1, 4], instData=[1, 64]
2730/// Result layout adjustment to ensure source can be safely inferred:
2731/// - laneData must be >= 2 and multiple of 2 (so source = laneData/2 is
2732/// valid)
2733/// - instData must be divisible by (16 * 2 = 32) (so source = instData/2 is
2734/// valid)
2735/// - Adjusted instData: ensure (instData % 32 == 0)
2736///
2737xegpu::DistributeLayoutAttr xegpu::setupInterleaveResultLayout(
2738 xegpu::LayoutKind layoutKind, VectorType srcVecTy, VectorType resVecTy,
2739 DistributeLayoutAttr consumerLayout, const xegpu::uArch::uArch *uArch) {
2740
2741 ArrayRef<int64_t> resShape = resVecTy.getShape();
2742 assert(consumerLayout.getRank() == static_cast<int64_t>(resShape.size()) &&
2743 "consumer layout rank must match source shape rank");
2744
2745 // Interleave doubles the innermost dimension (ratio = 2). Adjust the
2746 // innermost data field so the source layout can be recovered by dividing
2747 // by 2.
2748 const size_t innerMostDim = resShape.size() - 1;
2749 constexpr int ratio = 2;
2750 return adjustInnermostDimForDivisibility(consumerLayout, layoutKind,
2751 innerMostDim, ratio,
2752 resShape[innerMostDim], uArch);
2753}
2754
2755/// Sets up the result layout for an insert strided slice operation.
2756/// Creates a result layout based on the specified layout kind (InstData or
2757/// Lane).
2758xegpu::DistributeLayoutAttr xegpu::setupInsertStridedSliceResultLayout(
2759 xegpu::LayoutKind layoutKind, VectorType srcVectorTy,
2760 VectorType resVectorTy, xegpu::DistributeLayoutAttr consumerLayout,
2761 const xegpu::uArch::uArch *uArch) {
2762
2763 xegpu::DistributeLayoutAttr requiredResLayout;
2764 SmallVector<int64_t> consumerInstData =
2765 consumerLayout.getEffectiveInstDataAsInt();
2766 SmallVector<int64_t> consumerLaneData =
2767 consumerLayout.getEffectiveLaneDataAsInt();
2768 SmallVector<int64_t> consumerLaneLayout =
2769 consumerLayout.getEffectiveLaneLayoutAsInt();
2770 ArrayRef<int64_t> srcShape = srcVectorTy.getShape();
2771 int64_t laneDataValue = -1;
2772
2773 requiredResLayout = consumerLayout;
2774 int srcRank = srcShape.size();
2775
2776 if (layoutKind == xegpu::LayoutKind::Subgroup ||
2777 layoutKind == xegpu::LayoutKind::InstData) {
2778 assert(false && "subgroup/instData layout assignment not supported for "
2779 "insertStridedSlice.");
2780 } else if (layoutKind == xegpu::LayoutKind::Lane) {
2781 for (int dim = 0; dim < srcRank; dim++) {
2782 // A size-1 source dim is broadcast across the lanes of that dim.
2783 if (srcShape[dim] == 1) {
2784 laneDataValue = 1;
2785 } else {
2786 assert(srcShape[dim] % consumerLaneLayout[dim] == 0 &&
2787 "srcShape must be divisible by laneLayout for all dimensions");
2788 laneDataValue = std::min(srcShape[dim] / consumerLaneLayout[dim],
2789 consumerLaneData[dim]);
2790 }
2791 requiredResLayout =
2792 requiredResLayout.setDimData(dim, -1, -1, laneDataValue);
2793 }
2794 }
2795 return requiredResLayout;
2796}
2797
2798/// Back-propagates a known result layout to the layout required on `operand`
2799/// for a non-anchor (layout-propagating) vector op. Dispatches on the op kind —
2800/// broadcast, (multi)reduction, bitcast, shape/transpose, insert/extract,
2801/// interleave, etc. — applying the shape/permutation/bitwidth transform to
2802/// derive the source layout; elementwise and pass-through ops reuse resLayout
2803/// as-is. Returns nullptr for unknown ops or an absent result layout.
2804xegpu::DistributeLayoutAttr xegpu::inferSourceLayoutFromResultForNonAnchorOp(
2805 OpOperand &operand, xegpu::DistributeLayoutAttr resLayout) {
2806 if (!resLayout)
2807 return nullptr;
2808 Operation *op = operand.getOwner();
2809 unsigned idx = operand.getOperandNumber();
2810
2811 // For vector::BroadcastOp, infer the source layout from the result layout.
2812 if (auto broadcast = dyn_cast<vector::BroadcastOp>(op)) {
2813 auto srcTy = dyn_cast<VectorType>(broadcast.getSourceType());
2814 if (!srcTy)
2815 return nullptr;
2817 resLayout, broadcast.getResultVectorType().getShape(),
2818 srcTy.getShape());
2819 }
2820
2821 // For vector::MultiDimReductionOp, infer source layout from result layout
2822 // using reduction dims. Acc operand is expected to have the same layout as
2823 // the result.
2824 if (auto reduction = dyn_cast<vector::MultiDimReductionOp>(op)) {
2825 if (idx == 0) {
2826 SmallVector<int64_t> reductionDims(reduction.getReductionDims());
2827 return xegpu::inferMultiReductionSourceLayout(resLayout, reductionDims);
2828 }
2829 if (idx == 1)
2830 return resLayout;
2831 }
2832
2833 if (auto reduction = dyn_cast<vector::ReductionOp>(op))
2834 return xegpu::inferReductionSourceLayout(resLayout);
2835
2836 // For vector::BitCastOp, infer source layout from result layout using
2837 // element type bitwidths.
2838 if (auto bitcast = dyn_cast<vector::BitCastOp>(op)) {
2839 int resElemBitWidth =
2840 bitcast.getResultVectorType().getElementType().getIntOrFloatBitWidth();
2841 int srcElemBitWidth =
2842 bitcast.getSourceVectorType().getElementType().getIntOrFloatBitWidth();
2843 return xegpu::inferBitCastSourceLayout(resLayout, resElemBitWidth,
2844 srcElemBitWidth);
2845 }
2846
2847 // For vector::ShapeCastOp, infer source layout from result layout using
2848 // shapes.
2849 if (auto shapeCast = dyn_cast<vector::ShapeCastOp>(op)) {
2851 resLayout, shapeCast.getResultVectorType().getShape(),
2852 shapeCast.getSourceVectorType().getShape());
2853 }
2854
2855 // For vector::InsertStridedSliceOp, infer source layout from result
2856 // layout. Dest vector must have the same layout as the result.
2857 if (auto insertSlice = dyn_cast<vector::InsertStridedSliceOp>(op)) {
2858 if (idx == 0) {
2860 resLayout, insertSlice.getDestVectorType().getShape(),
2861 insertSlice.getSourceVectorType().getShape());
2862 }
2863 if (idx == 1)
2864 return resLayout;
2865 }
2866
2867 // For vector::Insert Op, infer source layout from result layout using
2868 // shapes.
2869 if (auto insert = dyn_cast<vector::InsertOp>(op)) {
2870 VectorType resVecTy = dyn_cast<VectorType>(insert.getResult().getType());
2871 VectorType valueToStoreTy =
2872 dyn_cast<VectorType>(insert.getValueToStore().getType());
2873
2874 if ((idx == 0) && valueToStoreTy) {
2875 return xegpu::inferInsertSourceLayout(resLayout, resVecTy.getShape(),
2876 valueToStoreTy.getShape());
2877 }
2878 if (idx == 1)
2879 return resLayout;
2880 }
2881
2882 // For vector::Extract Op, infer source layout from result layout using
2883 // shapes.
2884 if (auto extract = dyn_cast<vector::ExtractOp>(op)) {
2885 VectorType srcVecTy = dyn_cast<VectorType>(extract.getSource().getType());
2886 VectorType resVecTy = dyn_cast<VectorType>(extract.getResult().getType());
2887 if (!srcVecTy || !resVecTy)
2888 return nullptr;
2889 return xegpu::inferExtractSourceLayout(resLayout, resVecTy.getShape(),
2890 srcVecTy.getShape());
2891 }
2892
2893 // For vector::TransposeOp, infer source layout from result layout using
2894 // permutation.
2895 if (auto transpose = dyn_cast<vector::TransposeOp>(op)) {
2896 return xegpu::inferTransposeSourceLayout(resLayout,
2897 transpose.getPermutation());
2898 }
2899
2900 // For vector::BitCastOp, infer source layout from result layout using
2901 // element type bitwidths.
2902 if (auto bitcast = dyn_cast<vector::BitCastOp>(op)) {
2903 int resElemBitWidth =
2904 bitcast.getResultVectorType().getElementType().getIntOrFloatBitWidth();
2905 int srcElemBitWidth =
2906 bitcast.getSourceVectorType().getElementType().getIntOrFloatBitWidth();
2907 return xegpu::inferBitCastSourceLayout(resLayout, resElemBitWidth,
2908 srcElemBitWidth);
2909 }
2910
2911 // for vector::interleave
2912 if (auto interleave = dyn_cast<vector::InterleaveOp>(op)) {
2913 return xegpu::inferInterleaveSourceLayout(resLayout);
2914 }
2915
2916 // for vector::deinterleave
2917 if (auto deinterleave = dyn_cast<vector::DeinterleaveOp>(op)) {
2918 return xegpu::inferDeinterleaveSourceLayout(resLayout);
2919 }
2920
2921 // For vector::ExtractStridedSliceOp, simply return result layout
2922 if (dyn_cast<vector::ExtractStridedSliceOp>(op))
2923 return resLayout;
2924
2925 // For elementwise operations, all operands must have the same layout as
2926 // the result.
2928 return resLayout;
2929
2930 return nullptr;
2931}
2932
2933// For a loop terminator operand (scf.for's scf.yield, scf.while's
2934// scf.condition), returns the layout of the region iter_arg it forwards into,
2935// which is the authoritative loop-carried layout, or nullptr when that position
2936// was never assigned a layout.
2937static xegpu::DistributeLayoutAttr getLoopCarriedLayoutForYieldOperand(
2938 RegionBranchTerminatorOpInterface terminator, OpOperand &operand) {
2939 auto branch = dyn_cast<RegionBranchOpInterface>(terminator->getParentOp());
2940 if (!branch)
2941 return nullptr;
2943 branch.getSuccessorOperandInputMapping(mapping,
2944 RegionBranchPoint(terminator));
2945 auto it = mapping.find(&operand);
2946 if (it == mapping.end())
2947 return nullptr;
2948 xegpu::DistributeLayoutAttr iterArgLayout;
2949 for (auto arg : llvm::make_isa_range<BlockArgument>(it->second)) {
2950 xegpu::DistributeLayoutAttr layout = xegpu::getDistributeLayoutAttr(arg);
2951 assert((!iterArgLayout || !layout || iterArgLayout.isEqualTo(layout)) &&
2952 "region inputs fed by one terminator operand disagree on layout");
2953 if (!iterArgLayout)
2954 iterArgLayout = layout;
2955 }
2956 return iterArgLayout;
2957}
2958
2959// For the terminator of a region op that carries nothing back into its regions
2960// (scf.if), returns the layout of the parent result the operand feeds.
2961static xegpu::DistributeLayoutAttr getParentResultLayoutForYieldOperand(
2962 RegionBranchTerminatorOpInterface terminator, OpOperand &operand) {
2963 auto branch = dyn_cast<RegionBranchOpInterface>(terminator->getParentOp());
2964 if (!branch)
2965 return nullptr;
2967 branch.getSuccessorOperandInputMapping(mapping,
2968 RegionBranchPoint(terminator));
2969 auto it = mapping.find(&operand);
2970 if (it == mapping.end())
2971 return nullptr;
2972 for (Value input : it->second)
2973 if (auto result = dyn_cast<OpResult>(input))
2975 return nullptr;
2976}
2977
2978/// Returns the layout required on `operand`: anchor ops report their declared
2979/// per-operand layout directly; non-anchor ops back-derive it from their result
2980/// layout via inferSourceLayoutFromResultForNonAnchorOp.
2981xegpu::DistributeLayoutAttr xegpu::getConsumerLayoutAt(OpOperand &operand) {
2982 Operation *op = operand.getOwner();
2983 // Anchor ops declare the layout they
2984 // require on each operand. Trust that declaration directly so that
2985 // ResolveLayoutConflicts compares producer-vs-declared
2986 if (isa<xegpu::AnchorLayoutInterface>(op))
2987 return xegpu::getDistributeLayoutAttr(operand);
2988 // Region ops with forwarded operands (scf.for's and scf.while's inits) carry
2989 // the required operand layout as the layout_operand_N that
2990 // propagateRegionArgsToInits back-propagated from the region argument. Do not
2991 // re-derive it from that argument here: conflict resolution inserts
2992 // convert_layout ops as it walks, rewriting the argument's uses, so what
2993 // those uses require depends on how far the walk has progressed.
2994 if (isa<RegionBranchOpInterface>(op))
2995 return xegpu::getDistributeLayoutAttr(operand);
2996 // A region terminator requires the layout of the successor input its operand
2997 // feeds: the region iter_arg for a loop, and the parent result for a region
2998 // op with no loop-carried values (scf.if).
2999 if (auto terminator = dyn_cast<RegionBranchTerminatorOpInterface>(op)) {
3000 if (isa<LoopLikeOpInterface>(op->getParentOp()))
3001 return getLoopCarriedLayoutForYieldOperand(terminator, operand);
3002 return getParentResultLayoutForYieldOperand(terminator, operand);
3003 }
3004 // For non-anchor ops, derive the operand layout from the op's result
3005 // layout via op-specific semantics.
3006 xegpu::DistributeLayoutAttr resLayout;
3007 if (op->getNumResults() == 1 || isa<vector::DeinterleaveOp>(op))
3008 resLayout = xegpu::getDistributeLayoutAttr(op->getResult(0));
3009 return inferSourceLayoutFromResultForNonAnchorOp(operand, resLayout);
3010}
return success()
static void visit(Operation *op, DenseSet< Operation * > &visited)
Visits all the pdl.operand(s), pdl.result(s), and pdl.operation(s) connected to the given operation.
Definition PDL.cpp:62
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 xegpu::LayoutAttr buildLayout(mlir::MLIRContext *context, ArrayRef< int64_t > sgLayout, ArrayRef< int64_t > sgData, ArrayRef< int64_t > instData, ArrayRef< int64_t > laneLayout, ArrayRef< int64_t > laneData, DenseI32ArrayAttr orderAttr=nullptr)
static xegpu::DistributeLayoutAttr getStoreSubgroupLayouts(mlir::MLIRContext *context, ArrayRef< int64_t > wgShape, ArrayRef< int64_t > instData, int numSg)
Picks the subgroup layout for a scatter-style store (store_scatter / store_matrix): the most balanced...
static xegpu::DistributeLayoutAttr createScaleLayout(mlir::MLIRContext *context, VectorType matrixTy, VectorType scaleTy, xegpu::DistributeLayoutAttr matrixLayout, bool isBScale, const xegpu::uArch::uArch *uArch)
Helper to create a scale layout derived from a matrix operand layout.
static bool leadingDimsAreUnit(ArrayRef< int64_t > shape, int numInnerDims)
Returns true if every dimension of shape except the innermost numInnerDims is a unit (size-1) dimensi...
static xegpu::DistributeLayoutAttr adjustInnermostDimForDivisibility(xegpu::DistributeLayoutAttr consumerLayout, xegpu::LayoutKind layoutKind, size_t innerMostDim, int ratio, int64_t bound, const xegpu::uArch::uArch *uArch)
Adjusts consumerLayout's innermost-dim data field selected by layoutKind so that the source layout ca...
static std::pair< SmallVector< int64_t >, SmallVector< int64_t > > computeScatterIOLaneLayoutAndData(ArrayRef< int64_t > instShape, int64_t subgroupSize, int64_t maxChunkSize)
Computes lane_layout and lane_data for scatter-style store anchor layouts (store scatter,...
static xegpu::LayoutAttr buildInstDataLayoutWithLane(mlir::MLIRContext *context, ArrayRef< int64_t > instData, ArrayRef< int64_t > laneLayout, ArrayRef< int64_t > laneData, DenseI32ArrayAttr orderAttr=nullptr)
static xegpu::LayoutAttr buildSgLayout(mlir::MLIRContext *context, ArrayRef< int64_t > wgTileShape, ArrayRef< int64_t > sgLayout, int dimK=-1, DenseI32ArrayAttr orderAttr=nullptr)
static xegpu::DistributeLayoutAttr getParentResultLayoutForYieldOperand(RegionBranchTerminatorOpInterface terminator, OpOperand &operand)
static std::pair< SmallVector< int64_t >, SmallVector< int64_t > > computeReductionLaneLayoutAndData(ArrayRef< int64_t > srcShape, ArrayRef< int64_t > reductionDims, int subgroupSize, int64_t maxReduceVectorSize, bool verticalLaneLayout=false)
Computes the (lane_layout, lane_data) for a multi-reduction's source layout.
static std::optional< SmallVector< int64_t > > get2DBlockIOInstDataLayout(ArrayRef< int64_t > dataShape, Type elemTy, const xegpu::uArch::BlockIOInstructionInterface *uArchInstruction, bool transform=false, bool transpose=false)
Helper function to compute inst_data vectors for DPAS operands A, B, and C/D.
static SmallVector< LayoutRepresentation > getSgLayoutCandidates(ArrayRef< int64_t > wgShape, ArrayRef< int64_t > instData, int64_t sgCount, int64_t broadcastDim=-1)
static std::optional< std::tuple< xegpu::DistributeLayoutAttr, xegpu::DistributeLayoutAttr, xegpu::DistributeLayoutAttr > > getDpasSubgroupLayouts(mlir::MLIRContext *context, VectorType aTy, VectorType bTy, VectorType cdTy, xegpu::DistributeLayoutAttr consumerLayout, int numSg, std::tuple< SmallVector< int64_t >, SmallVector< int64_t >, SmallVector< int64_t > > instDataVecs)
Helper function to set up subgroup layouts for DPAS operands A, B, and C/D.
static SmallVector< LayoutRepresentation > enumerateFactorizations(int64_t total, int64_t rank)
Enumerates all ways to split total into rank factors whose product equals total.
static xegpu::DistributeLayoutAttr setupGenericLoadAnchorLayout(xegpu::LayoutKind layoutKind, mlir::MLIRContext *context, xegpu::DistributeLayoutAttr consumerLayout, int maxChunkSize, ArrayRef< int64_t > resShape, int subgroupSize)
Sets up the anchor layout for load gather and load matrix operation.
SmallVector< int64_t > LayoutRepresentation
static xegpu::DistributeLayoutAttr getLayoutFromUsePoints(Value result)
static xegpu::DistributeLayoutAttr setupGenericStoreAnchorLayout(xegpu::LayoutKind layoutKind, mlir::MLIRContext *context, int maxChunkSize, ArrayRef< int64_t > srcShape, int subgroupSize, int numSg)
Sets up the anchor layout for store scatter and store matrix operation, which share the same logic.
static std::tuple< SmallVector< int64_t >, SmallVector< int64_t >, SmallVector< int64_t > > compute2DBlockIOLaneLayout(ArrayRef< int64_t > instShape, int64_t subgroupSize, int64_t bitwidth, int64_t packingSize, bool transform=false, bool transpose=false)
static xegpu::DistributeLayoutAttr getLoopCarriedLayoutForYieldOperand(RegionBranchTerminatorOpInterface terminator, OpOperand &operand)
static void propagateResultsToRegularOperands(Operation *op)
static void propagateRegionResultsToYieldOperands(mlir::RegionBranchTerminatorOpInterface yieldOp)
static bool isValidLaneLayout(ArrayRef< int64_t > dataShape, ArrayRef< int64_t > laneLayout, ArrayRef< int64_t > laneData)
static void setTensorDescLayout(Value val, xegpu::DistributeLayoutAttr layout)
static void walkRegionBackward(Region &region, llvm::function_ref< void(Operation *)> visit)
static xegpu::LayoutAttr buildLaneLayout(mlir::MLIRContext *context, ArrayRef< int64_t > laneLayout, ArrayRef< int64_t > laneData, DenseI32ArrayAttr orderAttr=nullptr)
static std::optional< std::tuple< SmallVector< int64_t >, SmallVector< int64_t >, SmallVector< int64_t > > > getDpasInstDataLayouts(VectorType aTy, VectorType bTy, VectorType cdTy, const xegpu::uArch::MMAInstructionInterface *uArchInstruction)
Helper function to compute inst_data vectors for DPAS operands A, B, and C/D.
Attributes are known-constant values of operations.
Definition Attributes.h:25
Block represents an ordered list of Operations.
Definition Block.h:34
MLIRContext is the top-level object for a collection of MLIR operations.
Definition MLIRContext.h:63
This class represents an operand of an operation.
Definition Value.h:254
unsigned getOperandNumber() const
Return which operand this is in the OpOperand list of the Operation.
Definition Value.cpp:226
This is a value defined by a result of an operation.
Definition Value.h:454
This class implements the operand iterators for the Operation class.
Definition ValueRange.h:44
unsigned getBeginOperandIndex() const
Return the operand index of the first element of this range.
type_range getType() const
void walkInherentAttrs(Operation *op, InherentAttrVisitor visitor) const
Visit the inherent attributes stored in the properties of op.
Operation is the basic unit of execution within MLIR.
Definition Operation.h:87
bool hasDiscardableAttrOfType(NameT &&name)
Definition Operation.h:506
OpResult getResult(unsigned idx)
Get the 'idx'th result of this operation.
Definition Operation.h:432
unsigned getNumRegions()
Returns the number of regions held by this operation.
Definition Operation.h:726
Operation * getParentOp()
Returns the closest surrounding operation that contains this operation or nullptr if this is a top-le...
Definition Operation.h:251
MutableArrayRef< OpOperand > getOpOperands()
Definition Operation.h:408
OperationName getName()
The name of an operation is the key identifier for it.
Definition Operation.h:115
DictionaryAttr getDiscardableAttrDictionary()
Return all of the discardable attributes on this operation as a DictionaryAttr.
Definition Operation.h:553
Attribute removeDiscardableAttr(StringAttr name)
Remove the discardable attribute with the specified name if it exists.
Definition Operation.h:524
operand_range getOperands()
Returns an iterator on the underlying Value's.
Definition Operation.h:403
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
unsigned getNumResults()
Return the number of results held by this operation.
Definition Operation.h:429
This class represents a point being branched from in the methods of the RegionBranchOpInterface.
This class represents a successor of a region.
This class contains a list of basic blocks and a link to the parent operation it is attached to.
Definition Region.h:26
bool empty()
Definition Region.h:60
Instances of the Type class are uniqued, have an immutable identifier and an optional mutable compone...
Definition Types.h:74
unsigned getIntOrFloatBitWidth() const
Return the bit width of an integer or a float type, assert failure on other types.
Definition Types.cpp:124
This class provides an abstraction over the different types of ranges over Values.
Definition ValueRange.h:389
This class represents an instance of an SSA value in the MLIR system, representing a computable value...
Definition Value.h:96
void setType(Type newType)
Mutate the type of this Value to be of the specified type.
Definition Value.h:116
Type getType() const
Return the type of this value.
Definition Value.h:105
static DenseArrayAttrImpl get(MLIRContext *context, ArrayRef< int32_t > content)
Operation * getOwner() const
Return the owner of this operand.
Definition UseDefLists.h:38
bool hasElementwiseMappableTraits(Operation *op)
Together, Elementwise, Scalarizable, Vectorizable, and Tensorizable provide an easy way for scalar op...
DistributeLayoutAttr inferShapeCastSourceLayout(DistributeLayoutAttr resLayout, ArrayRef< int64_t > resShape, ArrayRef< int64_t > srcShape)
Infers the source layout attribute for a shape cast operation given the result layout attribute,...
bool matchDimCollapse(ArrayRef< int64_t > src, ArrayRef< int64_t > dst, SmallVector< SmallVector< int64_t > > &collapseDims)
DistributeLayoutAttr setupLoadNdAnchorLayout(LayoutKind layoutKind, VectorType vectorTy, DistributeLayoutAttr consumerLayout, int numSg, const uArch::uArch *uArch)
Sets up the anchor layout for a load_nd operation.
DistributeLayoutAttr inferResultLayoutFromSourceForNonAnchorOp(Operation *op, ArrayRef< DistributeLayoutAttr > operandLayouts)
Infers the result layout attribute for a non-anchor operation from the layouts of its source operands...
DistributeLayoutAttr setupLoadMatrixAnchorLayout(LayoutKind layoutKind, VectorType vectorTy, int contigChunkSize, DistributeLayoutAttr consumerLayout, const uArch::uArch *uArch)
Sets up the anchor layout for load matrix operation.
DistributeLayoutAttr setupInterleaveResultLayout(LayoutKind layoutKind, VectorType srcVectorTy, VectorType resVectorTy, DistributeLayoutAttr consumerLayout, const uArch::uArch *uArch)
Sets up the result layout for an interleave operation to ensure the source layout can be safely deriv...
DistributeLayoutAttr inferTransposeSourceLayout(DistributeLayoutAttr resLayout, ArrayRef< int64_t > permutation)
Infers the source layout attribute for a transpose operation given the result layout attribute and pe...
DistributeLayoutAttr inferInsertSourceLayout(DistributeLayoutAttr resLayout, ArrayRef< int64_t > resShape, ArrayRef< int64_t > srcShape)
Infers the source layout attribute for an insert operation.
std::optional< std::tuple< DistributeLayoutAttr, DistributeLayoutAttr, DistributeLayoutAttr, DistributeLayoutAttr, DistributeLayoutAttr > > completeDpasMxLaneLayoutFromInstData(DistributeLayoutAttr aLayout, DistributeLayoutAttr bLayout, DistributeLayoutAttr cdLayout, VectorType aTy, VectorType bTy, VectorType cdTy, VectorType aScaleTy, VectorType bScaleTy, const uArch::uArch *uArch)
Like completeDpasLaneLayoutFromInstData, but for dpas_mx: additionally re-derives the A_scale / B_sca...
DistributeLayoutAttr inferInsertStridedSliceSourceLayout(DistributeLayoutAttr resLayout, ArrayRef< int64_t > resShape, ArrayRef< int64_t > srcShape)
Infers the source layout attribute for an insert strided slice operation given the result layout attr...
DistributeLayoutAttr setupStoreMatrixAnchorLayout(LayoutKind layoutKind, VectorType vectorTy, int contigChunkSize, int numSg, const uArch::uArch *uArch)
Sets up the anchor layout for a store matrix operation.
void removeTemporaryLayoutAttrs(Operation *op)
Removes the temporary layout attributes for each OpOperand and OpResult of the given operation.
std::optional< std::tuple< DistributeLayoutAttr, DistributeLayoutAttr, DistributeLayoutAttr > > completeDpasLaneLayoutFromInstData(DistributeLayoutAttr aLayout, DistributeLayoutAttr bLayout, DistributeLayoutAttr cdLayout, VectorType aTy, VectorType bTy, VectorType cdTy, const uArch::uArch *uArch)
Completes user-provided DPAS A/B/C-D anchors that carry only inst_data by filling in lane_layout / la...
void setTemporaryLayout(const T &operandOrResult, const DistributeLayoutAttr layout)
LayoutKind
Specifies the level of a layout hierarchy for comparison or propagation.
Definition XeGPU.h:32
void setDistributeLayoutAttr(const OpResult &Result, const DistributeLayoutAttr layout)
[to-be-deprecated] Sets the DistributeLayoutAttr for a given OpResult user should use setAnchorLayout...
SmallVector< NamedAttribute > dropInstDataOnAttrs(ArrayRef< NamedAttribute > attrs)
Updates the NamedAttribute sequence by dropping inst-data information from any DistributeLayoutAttr f...
DistributeLayoutAttr inferSourceLayoutFromResultForNonAnchorOp(OpOperand &operand, DistributeLayoutAttr resLayout)
Infers the source layout attribute for an operand using result layout attribute.
DistributeLayoutAttr inferInterleaveSourceLayout(DistributeLayoutAttr resLayout)
Infers the source layout attribute for an interleave operation given the result layout attribute.
bool matchUnitDimExpansion(ArrayRef< int64_t > src, ArrayRef< int64_t > dst, SmallVector< int64_t > &expandedUnitDims)
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...
bool recoverTemporaryLayouts(Operation *rootOp)
Attach layout attributes to all vector-type operands of operations within the given operation's neste...
DistributeLayoutAttr inferBroadcastSourceLayout(DistributeLayoutAttr resLayout, ArrayRef< int64_t > resShape, ArrayRef< int64_t > srcShape)
Infers the source layout attribute for a broadcast operation given the result layout attribute,...
std::optional< std::tuple< DistributeLayoutAttr, DistributeLayoutAttr, DistributeLayoutAttr, DistributeLayoutAttr, DistributeLayoutAttr > > setupDpasMxLayout(LayoutKind layoutKind, VectorType aTy, VectorType bTy, VectorType cdTy, VectorType aScaleTy, VectorType bScaleTy, DistributeLayoutAttr consumerLayout, int numSg, const uArch::uArch *uArch)
Sets up the anchor layouts for dpas_mx operands (A, B, C/D, A_scale, and B_scale).
SliceAttr setupMultiReductionResultLayout(LayoutKind layoutKind, VectorType srcVectorTy, DistributeLayoutAttr consumerLayout, SmallVector< int64_t > reductionDims, int numSg, const uArch::uArch *uArch)
Note on the consumerLayout argument used by the consumer-driven setup* / complete* helpers below:
DistributeLayoutAttr setupLoadGatherAnchorLayout(LayoutKind layoutKind, VectorType vectorTy, int contigChunkSize, DistributeLayoutAttr consumerLayout, const uArch::uArch *uArch)
Sets up the anchor layout for a load gather operation.
llvm::function_ref< DistributeLayoutAttr(Value)> GetLayoutFnTy
Callable returning the propagated layout for a given Value, used by the layout-propagation helpers be...
std::optional< DistributeLayoutAttr > completeScatterLoadLaneLayoutFromInstData(DistributeLayoutAttr userSpecifiedLayout, DistributeLayoutAttr consumerLayout, Type elemTy, const xegpu::uArch::LoadGatherInstruction *uArchInstruction, const int subgroupSize)
If the consumer layout has only inst_data (no lane_layout/lane_data), completes it by running the cor...
bool matchSplitDimExpansion(ArrayRef< int64_t > src, ArrayRef< int64_t > dst, SmallVector< SmallVector< int64_t > > &splitDimGroups)
DistributeLayoutAttr setupStoreScatterAnchorLayout(LayoutKind layoutKind, VectorType vectorTy, int contigChunkSize, int numSg, const uArch::uArch *uArch)
Sets up the anchor layout for a store scatter operation.
DistributeLayoutAttr setupBitCastResultLayout(LayoutKind layoutKind, VectorType srcVectorTy, VectorType resVectorTy, DistributeLayoutAttr consumerLayout, const uArch::uArch *uArch)
Setup the result layout attribute for a bitcast operation based on element type bitwidths.
void removeLayoutAttr(const T &operandOrResult)
Removes the LayoutAttr for a given OpOperand or OpResult if it exists.
void dropInstDataOnInherentAttrs(Operation *op)
Drops inst-data information from DistributeLayoutAttrs stored as inherent attributes on the operation...
DistributeLayoutAttr getDistributeLayoutAttr(const Value value)
Retrieves the DistributeLayoutAttr associated with a given Value, or nullptr if none is found.
SmallVector< NamedAttribute > dropSgLayoutAndDataOnAttrs(ArrayRef< NamedAttribute > attrs)
Updates the NamedAttribute sequence by dropping sg-layout and sg-data information from any Distribute...
DistributeLayoutAttr setupPrefetchNdAnchorLayout(LayoutKind layoutKind, TensorDescType tdescTy, int numSg, const uArch::uArch *uArch)
Sets up the anchor layout for a prefetch_nd operation.
LogicalResult propagateYieldOperandsToRegionResults(RegionBranchTerminatorOpInterface terminator, GetLayoutFnTy getLayoutOfValue)
Propagate layouts from a region branch terminator's forwarded operands to the matching region results...
DistributeLayoutAttr inferShapeCastResultLayout(DistributeLayoutAttr srcLayout, ArrayRef< int64_t > srcShape, ArrayRef< int64_t > resShape)
Infers the result layout attribute for a shape cast operation given the source layout attribute,...
DistributeLayoutAttr inferExtractSourceLayout(DistributeLayoutAttr resLayout, ArrayRef< int64_t > resShape, ArrayRef< int64_t > srcShape)
Infers the source layout attribute for an extract operation.
std::string getTemporaryLayoutName(const OpOperand &operand)
Return the attribute name for the OpOperand to attach DistributeLayoutAttr.
DistributeLayoutAttr inferBitCastSourceLayout(DistributeLayoutAttr resLayout, int resElemTyBitWidth, int srcElemTyBitWidth)
Infers the source layout attribute for a bitcast operation given the result layout attribute,...
DistributeLayoutAttr setupInsertStridedSliceResultLayout(LayoutKind layoutKind, VectorType srcVectorTy, VectorType resVectorTy, DistributeLayoutAttr consumerLayout, const uArch::uArch *uArch)
Sets up the result layout for an insert strided slice operation.
DistributeLayoutAttr inferReductionSourceLayout(DistributeLayoutAttr resLayout)
Infers the source layout attribute for a reduction operation given the result layout attribute and re...
std::optional< DistributeLayoutAttr > completeScatterStoreLaneLayoutFromInstData(DistributeLayoutAttr specifiedLayout, Type elemTy, const xegpu::uArch::StoreScatterInstruction *uArchInstruction, const int subgroupSize)
Like completeScatterLoadLaneLayoutFromInstData, but for scatter stores (store_scatter / store_matrix)...
std::optional< DistributeLayoutAttr > completeBlockStoreLaneLayoutFromInstData(DistributeLayoutAttr specifiedLayout, Type elemTy, const xegpu::uArch::BlockIOInstructionInterface *uArchInstruction, const int subgroupSize)
Completes a user-provided 2D-block store_nd / prefetch_nd anchor that has only inst_data.
DistributeLayoutAttr inferDeinterleaveSourceLayout(DistributeLayoutAttr resLayout)
Infers the source layout attribute for a deinterleave operation given the result layout attribute.
DistributeLayoutAttr getConsumerLayoutAt(OpOperand &operand)
Gets the expected layout for a given consumer operand.
void removeLayoutAttrs(Operation *op)
Removes the DistributeLayoutAttr for each OpOperand and OpResult of the given operation if they exist...
DistributeLayoutAttr inferMultiReductionSourceLayout(DistributeLayoutAttr resLayout, SmallVector< int64_t > reduceDims)
Infers the source layout attribute for a reduction operation given the result layout attribute and re...
bool isTriviallyRematerializable(Operation *op)
Returns true if op is safe and cheap to clone: it has no side effects, no regions,...
DistributeLayoutAttr setupStoreNdAnchorLayout(LayoutKind layoutKind, VectorType vectorTy, int numSg, const uArch::uArch *uArch)
Sets up the anchor layout for a store_nd operation.
DistributeLayoutAttr inferTransposeResultLayout(DistributeLayoutAttr srcLayout, ArrayRef< int64_t > permutation)
Infers the result layout attribute for a transpose operation given the source layout attribute and pe...
std::optional< DistributeLayoutAttr > completeBlockLoadLaneLayoutFromInstData(DistributeLayoutAttr specifiedLayout, DistributeLayoutAttr consumerLayout, Type elemTy, const xegpu::uArch::BlockIOInstructionInterface *uArchInstruction, const int subgroupSize)
Like completeBlockStoreLaneLayoutFromInstData, but for load_nd.
LogicalResult propagateRegionArgsToInits(RegionBranchOpInterface regionOp, GetLayoutFnTy getLayoutOfValue)
Propagate layouts from a region branch op's region entry block arguments back to its init operands.
std::optional< std::tuple< DistributeLayoutAttr, DistributeLayoutAttr, DistributeLayoutAttr > > setupDpasLayout(LayoutKind layoutKind, VectorType aTy, VectorType bTy, VectorType cdTy, DistributeLayoutAttr consumerLayout, int numSg, const uArch::uArch *uArch)
Sets up the anchor layouts for a dpas operands (A, B, and C/D).
SliceAttr setupReductionResultLayout(LayoutKind layoutKind, VectorType srcVectorTy, const uArch::uArch *uArch)
Sets up layout for Reduction operations by creating a SliceAttr for the result.
Include the generated interface declarations.
DenseMap< OpOperand *, SmallVector< Value > > RegionBranchSuccessorMapping
A mapping from successor operands to successor inputs.
AffineMap inversePermutation(AffineMap map)
Returns a map of codomain to domain dimensions such that the first codomain dimension for a particula...
bool isMemoryEffectFree(Operation *op)
Returns true if the given operation is free of memory effects.
detail::DenseArrayAttrImpl< int32_t > DenseI32ArrayAttr
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.
SmallVector< int64_t > invertPermutationVector(ArrayRef< int64_t > permutation)
Helper method to apply to inverse a permutation.
virtual int32_t getPackedFormatBitSize() const =0
std::optional< BlockShapes > getBlockWidthHeightCount(Type elemTy, bool hasTransform=false, bool hasTranspose=false, bool upConv=false) const
Definition uArchBase.h:185
int32_t getMaxLaneAccessSizeBytes() const override
Definition uArchBase.h:226
virtual llvm::SmallVector< uint32_t, 8 > getSupportedN(Type type) const =0
virtual llvm::SmallVector< uint32_t, 8 > getSupportedK(Type type) const =0
virtual llvm::SmallVector< uint32_t, 8 > getSupportedM(Type type) const =0
int32_t getMaxLaneAccessSizeBytes() const override
Definition uArchBase.h:231
virtual unsigned getGeneralPackedFormatBitSize() const =0
virtual int getSubgroupSize() const =0
const Instruction * getInstruction(InstructionKind instKind) const
Definition uArchBase.h:115