MLIR 24.0.0git
XeGPUWgToSgDistribute.cpp
Go to the documentation of this file.
1//===- XeGPUWgToSgDistribute.cpp - XeGPU Workgroup to Subgroup Pass -------===//
2//
3// Part of the LLVM Project, under the Apache License v2.0 with LLVM Exceptions.
4// See https://llvm.org/LICENSE.txt for license information.
5// SPDX-License-Identifier: Apache-2.0 WITH LLVM-exception
6//
7//===----------------------------------------------------------------------===//
9
25#include "llvm/ADT/SetVector.h"
26#include <optional>
27
28namespace mlir {
29namespace xegpu {
30#define GEN_PASS_DEF_XEGPUWGTOSGDISTRIBUTE
31#include "mlir/Dialect/XeGPU/Transforms/Passes.h.inc"
32} // namespace xegpu
33} // namespace mlir
34
35using namespace mlir;
36
37namespace {
38
39// Retrieve the RangeAttr if it is specified.
40static xegpu::RangeAttr getRangeSpecAttr(Operation *op) {
41 Operation *parent = op->getParentOfType<scf::IfOp>();
42 while (parent) {
43 if (auto attr = llvm::dyn_cast_if_present<xegpu::RangeAttr>(
44 parent->getDiscardableAttr("sg_id_range")))
45 return attr;
46 parent = parent->getParentOfType<scf::IfOp>();
47 }
48 return {};
49}
50
51static std::pair<SmallVector<int64_t>, int>
52getSgShapeAndCount(ArrayRef<int64_t> shape,
53 xegpu::DistributeLayoutAttr layout) {
54 int count = 1;
56 auto distributedShape = layout.computeDistributedShape(
57 SmallVector<int64_t>(shape.begin(), shape.end()));
58 if (failed(distributedShape))
59 return std::make_pair(sgShape, count);
60 auto sgData = layout.getEffectiveSgDataAsInt();
61 count = computeProduct(distributedShape.value()) / computeProduct(sgData);
62 return std::make_pair(sgData, count);
63}
64
65/// Utility helper for deriving a list of offsets for each sub-TensorDescs
66/// or sub-MemDescs to be accessed by current subgroup (sgId) based on the
67/// associated distribute layout attribute, the shape, subgroup id and the
68/// original offsets of the op
69template <typename OpType,
70 typename = std::enable_if_t<llvm::is_one_of<
71 OpType, xegpu::LoadNdOp, xegpu::StoreNdOp, xegpu::PrefetchNdOp,
72 xegpu::LoadMatrixOp, xegpu::StoreMatrixOp>::value>>
73static LogicalResult
74genOffsetsList(ConversionPatternRewriter &rewriter, OpType op,
76 Location loc = op.getLoc();
77 SmallVector<OpFoldResult> origOffsets = op.getMixedOffsets();
78 // not applicable to ops without offsets operands.
79 if (origOffsets.empty())
80 return failure();
81
82 // if op is xegpu::CreateNdDescOp, call op.getDescLayoutAttr()
83 xegpu::DistributeLayoutAttr layout;
84 if constexpr (std::is_same_v<OpType, xegpu::LoadMatrixOp> ||
85 std::is_same_v<OpType, xegpu::StoreMatrixOp>) {
86 layout = op.getLayoutAttr();
87 } else {
88 layout = op.getDescLayoutAttr();
89 }
90
91 // not applicable to ops without workgroup layout attributes
92 if (!layout || !layout.isForWorkgroup())
93 return failure();
94
95 Value sgId =
96 gpu::SubgroupIdOp::create(rewriter, loc, /*upper_bound=*/nullptr);
97
98 // verify and adjust the sgId if the range specifier is present
99 xegpu::RangeAttr sgIdRange = getRangeSpecAttr(op);
100 if (sgIdRange) {
101 int64_t startOfRange = sgIdRange.getStart().getInt();
102 int64_t endOfRange = sgIdRange.getEnd().getInt();
103 // verify the RangeAttr against the layout attribute
104 if (layout.getNumSubgroups() != endOfRange - startOfRange)
105 return rewriter.notifyMatchFailure(
106 op, "sg_layout size must match the sg_id_range");
107 // adjust the sgId if necessary
108 if (startOfRange > 0) {
109 Value startOfRangeVal =
110 arith::ConstantIndexOp::create(rewriter, loc, startOfRange);
111 sgId = index::SubOp::create(rewriter, loc, sgId, startOfRangeVal);
112 }
113 }
114
115 // Compute the list of subgroup-relative offsets for sub-tensors or sub-memory
116 // descriptors to be accessed, based on the layout information.
117 ArrayRef<int64_t> wgShape = op.getDataShape();
118 auto maybeDescOffsets =
119 layout.computeDistributedCoords(rewriter, loc, sgId, wgShape);
120 if (failed(maybeDescOffsets))
121 return failure();
122
123 // Compute the final global offsets for each accessed sub-tensor
124 // or sub-memory descriptor.
125 for (const auto &sgOffsets : *maybeDescOffsets) {
127 rewriter, loc, getAsOpFoldResult(sgOffsets), origOffsets);
128 offsetsList.push_back(std::move(newOffsets));
129 }
130
131 // callback(offsetsList);
132 return success();
133}
134
135/// This pattern transforms the CreateNdDescOp to create a subgroup descriptor
136/// from a workgroup descriptor. It replaces the offsets and sizes with
137/// appropriate values for the subgroup.
138/// It uses round-robin assignment to distribute the work to the subgroups.
139/// Following create_nd_desc operation:
140/// %tdesc = xegpu.create_nd_tdesc %src : memref<24x24xf32>
141/// -> !xegpu.tensor_desc<24x24xf32, #xegpu.layout<sg_layout = [4, 4],
142/// sg_data = [2, 2], lane_layout = [2, 2], lane_data = [1, 1]>>
143/// is converted to 9 subgroup level operations based on the sg_layout &
144/// sg_data:
145/// %tdesc = xegpu.create_nd_tdesc %src : memref<24x24xf32> ->
146/// !xegpu.tensor_desc<2x2xf32, #xegpu.layout<lane_layout = [2, 2],
147/// lane_data = [1, 1]>>
148///
149/// The sg_layout and sg_data attributes are dropped after the pass as they are
150/// no longer needed.
151///
152/// 24x24 matrix distribution example:
153/// sg_layout = [4, 4], sg_data = [2, 2]
154/// Each 8x8 matrix within the 24x24 matrix is called a distribution unit.
155/// dist_unit_shape = [8, 8] --> sg_layout[i] * sg_data[i]
156///
157/// +------------------------+
158/// | 8x8 | 8x8 | 8x8 | <- 3 tiles across
159/// |-----+-----+-----|
160/// | 8x8 | 8x8 | 8x8 | <- 3 tiles down
161/// |-----+-----+-----|
162/// | 8x8 | 8x8 | 8x8 |
163/// +------------------------+
164///
165/// Each 8x8 tile is further subdivided among subgroups:
166/// +------------------------+
167/// | 2x2 2x2 2x2 2x2 | <- 4 subgroups across (each handles 2 columns)
168/// | 2x2 2x2 2x2 2x2 | <- 4 subgroups down (each handles 2 rows)
169/// | 2x2 2x2 2x2 2x2 |
170/// | 2x2 2x2 2x2 2x2 |
171/// +------------------------+
172///
173/// Since the 24x24 matrix is divided into 8x8 distribution units, there will be
174/// 9 distribution units (3x3) in total. Hence the 9 subgroup level operations.
175
176/// The pass currently has entire distribution logic in the WgToSgCreateNdOp
177/// pattern and all the other ops just follow.
178/// TODO: Decouple the distribution logic from WgToSgCreateNdOp for all the
179/// ops in the pass.
180// This pattern transforms the CreateNdDescOp to create a
181// subgroup descriptor from a workgroup descriptor.
182struct WgToSgCreateNdOp : public OpConversionPattern<xegpu::CreateNdDescOp> {
183 using OpConversionPattern<xegpu::CreateNdDescOp>::OpConversionPattern;
184
185 LogicalResult
186 matchAndRewrite(xegpu::CreateNdDescOp op, OneToNOpAdaptor adaptor,
187 ConversionPatternRewriter &rewriter) const override {
188
189 Location loc = op.getLoc();
190 MLIRContext *ctx = op.getContext();
191 xegpu::TensorDescType tdescTy = op.getType();
192 auto layout = dyn_cast<xegpu::DistributeLayoutAttr>(tdescTy.getLayout());
193 if (!layout || !layout.isForWorkgroup())
194 return failure();
195
196 Type elemTy = tdescTy.getElementType();
197 ArrayRef<int64_t> wgShape = tdescTy.getShape();
198
199 SmallVector<int64_t> sgShape;
200 int count;
201 std::tie(sgShape, count) = getSgShapeAndCount(wgShape, layout);
202 xegpu::TensorDescType newTdescTy =
203 xegpu::TensorDescType::get(ctx, sgShape, elemTy, tdescTy.getEncoding(),
204 layout.dropSgLayoutAndData());
205
206 Value src = op.getSource();
207 SmallVector<Value> newCreateNdOps(count);
208 std::generate(newCreateNdOps.begin(), newCreateNdOps.end(), [&]() -> Value {
209 if (isa<MemRefType>(src.getType()))
210 return xegpu::CreateNdDescOp::create(rewriter, loc, newTdescTy,
211 cast<TypedValue<MemRefType>>(src));
212 return xegpu::CreateNdDescOp::create(rewriter, loc, newTdescTy, src,
213 op.getMixedSizes(),
214 op.getMixedStrides());
215 });
216
217 rewriter.replaceOpWithMultiple(op, {newCreateNdOps});
218 return success();
219 }
220};
221
222/// This pattern transforms the LoadNdOp to load subgroup data.
223struct WgToSgLoadNdOp : public OpConversionPattern<xegpu::LoadNdOp> {
224 using OpConversionPattern<xegpu::LoadNdOp>::OpConversionPattern;
225 LogicalResult
226 matchAndRewrite(xegpu::LoadNdOp op, OneToNOpAdaptor adaptor,
227 ConversionPatternRewriter &rewriter) const override {
228
229 SmallVector<SmallVector<OpFoldResult>> offsetsList;
230 if (failed(genOffsetsList(rewriter, op, offsetsList)))
231 return failure();
232
233 xegpu::DistributeLayoutAttr layout = op.getLayoutAttr();
234 if (layout)
235 layout = layout.dropSgLayoutAndData();
236 SmallVector<Value> newOps;
237 for (auto [tdesc, offsets] :
238 llvm::zip(adaptor.getTensorDesc(), offsetsList)) {
239 auto tdescTy = dyn_cast<xegpu::TensorDescType>(tdesc.getType());
240 VectorType newResTy =
241 VectorType::get(tdescTy.getShape(), tdescTy.getElementType());
242 auto newOp = xegpu::LoadNdOp::create(
243 rewriter, op.getLoc(), newResTy, tdesc, offsets,
244 /*packed = */ nullptr, /*transpose = */ nullptr, op.getL1HintAttr(),
245 op.getL2HintAttr(), op.getL3HintAttr(), layout);
246 newOps.push_back(newOp);
247 }
248 rewriter.replaceOpWithMultiple(op, {newOps});
249
250 return success();
251 }
252};
253
254/// This pattern transforms the StoreNdOp to store subgroup data.
255struct WgToSgStoreNdOp : public OpConversionPattern<xegpu::StoreNdOp> {
256 using OpConversionPattern<xegpu::StoreNdOp>::OpConversionPattern;
257 LogicalResult
258 matchAndRewrite(xegpu::StoreNdOp op, OneToNOpAdaptor adaptor,
259 ConversionPatternRewriter &rewriter) const override {
260 SmallVector<SmallVector<OpFoldResult>> offsetsList;
261 if (failed(genOffsetsList(rewriter, op, offsetsList)))
262 return failure();
263
264 xegpu::DistributeLayoutAttr layout = op.getLayoutAttr();
265 if (layout)
266 layout = layout.dropSgLayoutAndData();
267 for (auto [v, tdesc, offsets] :
268 llvm::zip(adaptor.getValue(), adaptor.getTensorDesc(), offsetsList)) {
269 xegpu::StoreNdOp::create(rewriter, op.getLoc(), v, tdesc, offsets,
270 op.getL1HintAttr(), op.getL2HintAttr(),
271 op.getL3HintAttr(), layout);
272 }
273 rewriter.eraseOp(op);
274
275 return success();
276 }
277};
278
279/// This pattern transforms the PrefetchNdOp to prefetch subgroup data.
280struct WgToSgPrefetchNdOp : public OpConversionPattern<xegpu::PrefetchNdOp> {
281 using OpConversionPattern<xegpu::PrefetchNdOp>::OpConversionPattern;
282 LogicalResult
283 matchAndRewrite(xegpu::PrefetchNdOp op, OneToNOpAdaptor adaptor,
284 ConversionPatternRewriter &rewriter) const override {
285 SmallVector<SmallVector<OpFoldResult>> offsetsList;
286 if (failed(genOffsetsList(rewriter, op, offsetsList)))
287 return failure();
288
289 xegpu::DistributeLayoutAttr layout = op.getLayoutAttr();
290 if (layout)
291 layout = layout.dropSgLayoutAndData();
292 for (auto [tdesc, offsets] :
293 llvm::zip(adaptor.getTensorDesc(), offsetsList)) {
294 xegpu::PrefetchNdOp::create(rewriter, op.getLoc(), tdesc, offsets,
295 op.getL1HintAttr(), op.getL2HintAttr(),
296 op.getL3HintAttr(), layout);
297 }
298 rewriter.eraseOp(op);
299
300 return success();
301 }
302};
303
304/// This pattern transforms the DpasOp to work at subgroup level.
305struct WgToSgDpasOp : public OpConversionPattern<xegpu::DpasOp> {
306 using OpConversionPattern<xegpu::DpasOp>::OpConversionPattern;
307 LogicalResult
308 matchAndRewrite(xegpu::DpasOp op, OneToNOpAdaptor adaptor,
309 ConversionPatternRewriter &rewriter) const override {
310 Location loc = op.getLoc();
311 VectorType resultTy = op.getResult().getType();
312 if (resultTy.getRank() < 2)
313 return failure();
314
315 auto layoutCd = op.getLayoutCdAttr();
316 auto layoutA = op.getLayoutAAttr();
317 auto layoutB = op.getLayoutBAttr();
318 if (!layoutCd || !layoutA || !layoutB)
319 return failure();
320 size_t i = 0;
321 SmallVector<Value> newDpasOps;
322 for (auto aVec : adaptor.getLhs()) {
323 for (auto bVec : adaptor.getRhs()) {
324
325 Value tmpC;
326 if (op.getAcc())
327 tmpC = adaptor.getAcc()[i++];
328
329 ArrayRef<int64_t> aVecShape =
330 cast<VectorType>(aVec.getType()).getShape();
331 ArrayRef<int64_t> bVecShape =
332 cast<VectorType>(bVec.getType()).getShape();
333 // Build result shape: batch dims from A + [M, N] from last dims of
334 // A and B.
335 SmallVector<int64_t> resShape(aVecShape.drop_back(2));
336 resShape.push_back(aVecShape[aVecShape.size() - 2]);
337 resShape.push_back(bVecShape[bVecShape.size() - 1]);
338 VectorType resTy = VectorType::get(resShape, resultTy.getElementType());
339 auto newDpasOp = xegpu::DpasOp::create(
340 rewriter, loc, resTy, aVec, bVec, tmpC,
341 /*layout_a=*/nullptr, /*layout_b=*/nullptr, /*layout_cd=*/nullptr);
342 newDpasOp.setLayoutCdAttr(layoutCd.dropSgLayoutAndData());
343 newDpasOp.setLayoutAAttr(layoutA.dropSgLayoutAndData());
344 newDpasOp.setLayoutBAttr(layoutB.dropSgLayoutAndData());
345
346 newDpasOps.push_back(newDpasOp);
347 }
348 }
349 rewriter.replaceOpWithMultiple(op, {newDpasOps});
350 return success();
351 }
352};
353
354/// This pattern transforms the DpasMxOp to work at subgroup level.
355struct WgToSgDpasMxOp : public OpConversionPattern<xegpu::DpasMxOp> {
356 using OpConversionPattern<xegpu::DpasMxOp>::OpConversionPattern;
357 LogicalResult
358 matchAndRewrite(xegpu::DpasMxOp op, OneToNOpAdaptor adaptor,
359 ConversionPatternRewriter &rewriter) const override {
360
361 Location loc = op.getLoc();
362 VectorType resultTy = op.getResult().getType();
363
364 if (resultTy.getRank() < 2)
365 return failure();
366
367 auto layoutCd = op.getLayoutCdAttr();
368 auto layoutA = op.getLayoutAAttr();
369 auto layoutB = op.getLayoutBAttr();
370 auto layoutAScale = op.getLayoutAScaleAttr();
371 auto layoutBScale = op.getLayoutBScaleAttr();
372
373 if (!layoutCd || !layoutA || !layoutB || !layoutAScale || !layoutBScale)
374 return failure();
375
376 size_t index_c = 0;
377 SmallVector<Value> newDpasMxOps;
378 for (auto [index_a, aVec] : llvm::enumerate(adaptor.getA())) {
379 for (auto [index_b, bVec] : llvm::enumerate(adaptor.getB())) {
380 Value accVal = (op.getAcc()) ? adaptor.getAcc()[index_c++] : Value();
381 Value scaleAVal =
382 (op.getScaleA()) ? adaptor.getScaleA()[index_a] : Value();
383 Value scaleBVal =
384 (op.getScaleB()) ? adaptor.getScaleB()[index_b] : Value();
385
386 ArrayRef<int64_t> aVecShape =
387 cast<VectorType>(aVec.getType()).getShape();
388 ArrayRef<int64_t> bVecShape =
389 cast<VectorType>(bVec.getType()).getShape();
390 // Build result shape: batch dims from A + [M, N]
391 SmallVector<int64_t> resShape(aVecShape.drop_back(2));
392 resShape.push_back(aVecShape[aVecShape.size() - 2]);
393 resShape.push_back(bVecShape[bVecShape.size() - 1]);
394 VectorType resTy = VectorType::get(resShape, resultTy.getElementType());
395 auto newDpasMxOp = xegpu::DpasMxOp::create(
396 rewriter, loc, resTy, aVec, bVec, accVal, scaleAVal, scaleBVal,
397 layoutA.dropSgLayoutAndData(), layoutB.dropSgLayoutAndData(),
398 layoutCd.dropSgLayoutAndData(), layoutAScale.dropSgLayoutAndData(),
399 layoutBScale.dropSgLayoutAndData());
400
401 newDpasMxOps.push_back(newDpasMxOp);
402 }
403 }
404 rewriter.replaceOpWithMultiple(op, {newDpasMxOps});
405 return success();
406 }
407};
408
409/// This pattern transforms vector.broadcast ops to work at subgroup level.
410struct WgToSgVectorBroadcastOp
411 : public OpConversionPattern<vector::BroadcastOp> {
412 using OpConversionPattern<vector::BroadcastOp>::OpConversionPattern;
413
414 LogicalResult
415 matchAndRewrite(vector::BroadcastOp op, OneToNOpAdaptor adaptor,
416 ConversionPatternRewriter &rewriter) const override {
417
418 VectorType resultType = op.getResult().getType();
419 ArrayRef<int64_t> wgShape = resultType.getShape();
420
421 xegpu::DistributeLayoutAttr layout =
422 xegpu::getTemporaryLayout(llvm::cast<OpResult>(op.getResult()));
423 if (!layout || !layout.isForWorkgroup())
424 return failure();
425
426 SmallVector<int64_t> sgShape;
427 int count;
428 std::tie(sgShape, count) = getSgShapeAndCount(wgShape, layout);
429 VectorType newResultType =
430 VectorType::get(sgShape, resultType.getElementType());
431
432 SmallVector<Value> newBroadcastOps;
433 auto distSource = adaptor.getOperands().front();
434 int numDistributions = count / distSource.size();
435 for (int i = 0; i < numDistributions; ++i) {
436 for (auto operand : distSource) {
437 auto newBroadcast = vector::BroadcastOp::create(rewriter, op.getLoc(),
438 newResultType, operand);
439
440 newBroadcastOps.push_back(newBroadcast.getResult());
441 }
442 }
443 rewriter.replaceOpWithMultiple(op, {newBroadcastOps});
444 return success();
445 }
446};
447
448// This pattern transforms elementwise ops to work at subgroup level.
449struct WgToSgElementwiseOp : public ConversionPattern {
450 WgToSgElementwiseOp(MLIRContext *ctx)
451 : ConversionPattern(MatchAnyOpTypeTag(), /*benefit=*/1, ctx) {}
452
453 LogicalResult
454 matchAndRewrite(Operation *op, ArrayRef<ValueRange> operands,
455 ConversionPatternRewriter &rewriter) const override {
456 // Only match ops with elementwise trait and single result.
458 return failure();
459
460 auto resultType = dyn_cast<VectorType>(op->getResult(0).getType());
461 assert(resultType && "Expected result to be a VectorType");
462
463 ArrayRef<int64_t> wgShape = resultType.getShape();
464
465 xegpu::DistributeLayoutAttr layout =
466 xegpu::getTemporaryLayout(llvm::cast<OpResult>(op->getResult(0)));
467 if (!layout || !layout.isForWorkgroup())
468 return failure();
469
470 SmallVector<int64_t> sgShape = getSgShapeAndCount(wgShape, layout).first;
471
472 size_t numVariants = operands.empty() ? 0 : operands.front().size();
473
474 if (llvm::any_of(operands, [&](const ValueRange &operandVec) {
475 return operandVec.size() != numVariants;
476 }))
477 return failure();
478
479 SmallVector<Value> newResults;
480 VectorType newResultType =
481 VectorType::get(sgShape, resultType.getElementType());
482
483 for (size_t i = 0; i < numVariants; ++i) {
484 SmallVector<Value> opOperands;
485 for (auto &operandVec : operands)
486 opOperands.push_back(operandVec[i]);
487
488 OperationState state(op->getLoc(), op->getName());
489 state.addOperands(opOperands);
490 state.addTypes(newResultType);
491 state.addAttributes(op->getDiscardableAttrDictionary().getValue());
492 state.propertiesAttr = op->getPropertiesAsAttribute();
493 Operation *newOp = rewriter.create(state);
495 newResults.push_back(newOp->getResult(0));
496 }
497
498 rewriter.replaceOpWithMultiple(op, {newResults});
499 return success();
500 }
501};
502
503// clang-format off
504// Pattern for lowering ConvertLayoutOp based on sg_layout and sg_data.
505// If input_layout and target_layout have identical sg_layout and sg_data,
506// the op is rewritten to a subgroup-level ConvertLayoutOp with these fields
507// dropped. For example:
508// #a = #xegpu.layout<sg_layout = [2, 2], sg_data = [16, 16], inst_data = [16, 16]>
509// #b = #xegpu.layout<sg_layout = [2, 2], sg_data = [16, 16], inst_data = [8, 16]>
510// xegpu.convert_layout %1 <{input_layout = #a, target_layout = #b}> : vector<32x64xf32>
511// becomes:
512// #a = #xegpu.layout<inst_data = [16, 16]>
513// #b = #xegpu.layout<inst_data = [8, 16]>
514// xegpu.convert_layout %1 <{input_layout = #a, target_layout = #b}> : vector<16x16xf32>
515// (vector<16x16xf32> is determined by sg_data = [16, 16])
516//
517// If sg_layout or sg_data differ, SLM is used to redistribute data across subgroups.
518// For example:
519// #a = #xegpu.layout<sg_layout = [1, 4], sg_data = [32, 16], inst_data = [16, 16]>
520// #b = #xegpu.layout<sg_layout = [2, 2], sg_data = [16, 32], inst_data = [8, 16]>
521// xegpu.convert_layout %1 <{input_layout = #a, target_layout = #b}> : vector<32x64xf32>
522// is lowered to:
523// #a = #xegpu.layout<inst_data = [16, 16]>
524// #b = #xegpu.layout<inst_data = [8, 16]>
525// store_matrix %1, %slm <{layout_input_0 = #a}> : vector<32x16>, mem_desc<32x64xf32>
526// %d = load_matrix %slm <{layout_result_0 = #a}> : mem_desc<32x64xf32> -> vector<16x32xf32>
527// xegpu.convert_layout %d <{input_layout = #a, target_layout = #b}> : vector<16x32xf32>
528// clang-format on
529struct WgToSgConvertLayoutOp
530 : public OpConversionPattern<xegpu::ConvertLayoutOp> {
531 using OpConversionPattern<xegpu::ConvertLayoutOp>::OpConversionPattern;
532
533 LogicalResult
534 matchAndRewrite(xegpu::ConvertLayoutOp op, OneToNOpAdaptor adaptor,
535 ConversionPatternRewriter &rewriter) const override {
536 Location loc = op.getLoc();
537 auto inputLayout = op.getEffectiveInputLayout();
538 auto targetLayout = op.getTargetLayout();
539
540 if (!inputLayout || !targetLayout || !inputLayout.isForWorkgroup() ||
541 !targetLayout.isForWorkgroup())
542 return rewriter.notifyMatchFailure(
543 op, "Input and target layouts must have subgroup layout");
544
545 Type resultType = op.getResult().getType();
546 if (resultType.isIntOrFloat()) {
547 rewriter.replaceOp(op, op.getSource());
548 assert(!inputLayout.dropSgLayoutAndData() &&
549 !targetLayout.dropSgLayoutAndData() &&
550 "unexpected layout attributes for scalar type");
551 return success();
552 }
553
554 ArrayRef<int64_t> wgShape = cast<VectorType>(resultType).getShape();
555 SmallVector<int64_t> inputSgLayout =
556 inputLayout.getEffectiveSgLayoutAsInt();
557 SmallVector<int64_t> inputSgData = inputLayout.getEffectiveSgDataAsInt();
558 SmallVector<int64_t> targetSgLayout =
559 targetLayout.getEffectiveSgLayoutAsInt();
560 SmallVector<int64_t> targetSgData = targetLayout.getEffectiveSgDataAsInt();
561
562 // Fast path: if sg_layout and sg_data are identical, no SLM needed
563 SmallVector<int64_t> wgShapeVec(wgShape.begin(), wgShape.end());
564 if (inputLayout.isCompatibleWith(targetLayout, wgShapeVec,
565 xegpu::LayoutKind::Subgroup)) {
566 inputLayout = inputLayout.dropSgLayoutAndData();
567 targetLayout = targetLayout.dropSgLayoutAndData();
568
569 SmallVector<Value> newOps(adaptor.getSource());
570 if (inputLayout && targetLayout) {
571 for (auto [i, src] : llvm::enumerate(adaptor.getSource())) {
572 auto newOp = xegpu::ConvertLayoutOp::create(
573 rewriter, loc, src.getType(), src, inputLayout, targetLayout);
574 newOps[i] = newOp;
575 }
576 }
577 rewriter.replaceOpWithMultiple(op, {newOps});
578 return success();
579 }
580
581 // SLM path: layouts differ, need cross-subgroup data redistribution
582 Type elemTy = cast<VectorType>(op.getSource().getType()).getElementType();
583
584 SmallVector<int64_t> slmShape = llvm::to_vector(wgShape);
585
586 // Calculate SLM size requirements
587 auto bitWidth = elemTy.getIntOrFloatBitWidth();
588 auto bytesPerElement = bitWidth / 8;
589 auto slmSize = computeProduct(slmShape) * bytesPerElement;
590
591 // Allocate SLM
592 auto slmTy = MemRefType::get({slmSize}, rewriter.getI8Type(), {}, 3);
593 auto slm = memref::AllocaOp::create(rewriter, loc, slmTy);
594
595 auto memDescType = xegpu::MemDescType::get(rewriter.getContext(), slmShape,
596 elemTy, nullptr);
597 auto memDesc =
598 xegpu::CreateMemDescOp::create(rewriter, loc, memDescType, slm);
599
600 auto sgId = gpu::SubgroupIdOp::create(rewriter, loc,
601 rewriter.getIndexType(), nullptr);
602
603 // STORE PHASE: Each subgroup stores in SLM using input layout
604 auto storeCoords = inputLayout.computeDistributedCoords(
605 rewriter, loc, sgId.getResult(), wgShape);
606 if (failed(storeCoords))
607 return failure();
608
609 // Store to SLM
610 for (auto [src, coords] : llvm::zip(adaptor.getSource(), *storeCoords)) {
611 SmallVector<OpFoldResult> storeMatrixOffsets;
612 for (Value coord : coords) {
613 storeMatrixOffsets.push_back(coord);
614 }
615 xegpu::StoreMatrixOp::create(rewriter, loc, src, memDesc.getResult(),
616 storeMatrixOffsets, nullptr /*layout*/);
617 }
618
619 gpu::BarrierOp::create(rewriter, loc);
620
621 // LOAD PHASE: Each target subgroup loads from SLM using target layout
622 auto loadCoords = targetLayout.computeDistributedCoords(
623 rewriter, loc, sgId.getResult(), wgShape);
624 if (failed(loadCoords))
625 return failure();
626
627 VectorType loadType = VectorType::get(targetSgData, elemTy);
628
629 // Load vectors from SLM
630 SmallVector<Value> finalResults;
631 for (auto coords : *loadCoords) {
632 SmallVector<OpFoldResult> loadMatrixOffsets;
633 for (Value coord : coords) {
634 loadMatrixOffsets.push_back(coord);
635 }
636 auto loadOp = xegpu::LoadMatrixOp::create(
637 rewriter, loc, loadType, memDesc.getResult(), loadMatrixOffsets,
638 targetLayout.dropSgLayoutAndData());
639
640 finalResults.push_back(loadOp.getResult());
641 }
642
643 rewriter.replaceOpWithMultiple(op, {finalResults});
644 return success();
645 }
646};
647
648// This pattern distributes arith.constant op into subgroup-level constants
649struct WgToSgArithConstantOp : public OpConversionPattern<arith::ConstantOp> {
650 using OpConversionPattern<arith::ConstantOp>::OpConversionPattern;
651
652 LogicalResult
653 matchAndRewrite(arith::ConstantOp op, OneToNOpAdaptor adaptor,
654 ConversionPatternRewriter &rewriter) const override {
655 auto vecAttr = dyn_cast<DenseElementsAttr>(op.getValue());
656 auto vecType = dyn_cast<VectorType>(op.getType());
657 if (!vecAttr || !vecType)
658 return failure();
659
660 xegpu::DistributeLayoutAttr layout =
661 xegpu::getTemporaryLayout(dyn_cast<OpResult>(op.getResult()));
662 if (!layout || !layout.isForWorkgroup())
663 return failure();
664
665 ArrayRef<int64_t> wgShape = vecType.getShape();
666 SmallVector<int64_t> sgShape;
667 int count;
668 std::tie(sgShape, count) = getSgShapeAndCount(wgShape, layout);
669
670 auto newType = VectorType::get(sgShape, vecType.getElementType());
671 Location loc = op.getLoc();
672 auto eltType = vecType.getElementType();
673
674 if (vecAttr.isSplat()) {
675 // Splat: single value for all subgroups
676 Attribute singleVal = vecAttr.getSplatValue<Attribute>();
677 auto sgAttr = DenseElementsAttr::get(newType, singleVal);
678 SmallVector<Value> newConstOps;
679 for (int i = 0; i < count; ++i) {
680 auto cstOp = arith::ConstantOp::create(rewriter, loc, newType, sgAttr);
681 newConstOps.push_back(cstOp);
682 }
683 rewriter.replaceOpWithMultiple(op, {newConstOps});
684 return success();
685 } else if (sgShape == wgShape) { // if the entire vector is shared by all
686 // subgroups, don't distribute
687 auto newConstOp =
688 arith::ConstantOp::create(rewriter, op.getLoc(), vecType, vecAttr);
689 rewriter.replaceOp(op, newConstOp);
690 return success();
691 } else {
692 // Non-splat constant
693 // Only supports 1D & 2D
694 // TODO: support other cases that require SLM access
695 if (!eltType.isIndex())
696 return rewriter.notifyMatchFailure(
697 op, "Unsupported element type for non-splat constant op.");
698
699 if (wgShape.size() > 2)
700 return rewriter.notifyMatchFailure(
701 op, "Only 1D & 2D vector constant supported");
702
703 SmallVector<Attribute> values(vecAttr.getValues<Attribute>());
704 int64_t rowStride = 0, colStride = 0;
705 int64_t rows = wgShape.size() == 1 ? 1 : wgShape[0];
706 int64_t cols = wgShape.size() == 1 ? wgShape[0] : wgShape[1];
707
708 // Compute colStride and rowStride, and check for constant strides.
709 if (cols > 1) {
710 colStride = cast<IntegerAttr>(values[1]).getInt() -
711 cast<IntegerAttr>(values[0]).getInt();
712 }
713 if (rows > 1) {
714 rowStride = cast<IntegerAttr>(values[cols]).getInt() -
715 cast<IntegerAttr>(values[0]).getInt();
716 }
717
718 for (int64_t r = 0; r < rows; ++r) {
719 for (int64_t c = 0; c < cols; ++c) {
720 int64_t idx = r * cols + c;
721 // Check column stride
722 if (c > 0 && cols > 1) {
723 int64_t prevIdx = r * cols + (c - 1);
724 int64_t diff = cast<IntegerAttr>(values[idx]).getInt() -
725 cast<IntegerAttr>(values[prevIdx]).getInt();
726 if (diff != colStride)
727 return rewriter.notifyMatchFailure(
728 op, "Non-constant column stride in constant op.");
729 }
730 // Check row stride
731 if (r > 0 && rows > 1) {
732 int64_t prevIdx = (r - 1) * cols + c;
733 int64_t diff = cast<IntegerAttr>(values[idx]).getInt() -
734 cast<IntegerAttr>(values[prevIdx]).getInt();
735 if (diff != rowStride)
736 return rewriter.notifyMatchFailure(
737 op, "Non-constant row stride in constant op.");
738 }
739 }
740 }
741
742 // Create a constant for the base tile.
743 // For 2D case, extract the top-left sgShape[0] x sgShape[1] submatrix.
744 // For 1D case, extract the first sgShape[0] elements.
745 SmallVector<Attribute> baseTileValues;
746 int baseTileCols = sgShape[sgShape.size() - 1];
747 int64_t baseTileRows = sgShape.size() == 1 ? 1 : sgShape[0];
748 for (int64_t r = 0; r < baseTileRows; ++r) {
749 for (int64_t c = 0; c < baseTileCols; ++c) {
750 baseTileValues.push_back(values[r * cols + c]);
751 }
752 }
753
754 auto tileAttr = DenseElementsAttr::get(VectorType::get(sgShape, eltType),
755 baseTileValues);
756 auto baseConstVec = arith::ConstantOp::create(rewriter, loc, tileAttr);
757
758 // Get subgroup id
759 Value sgId =
760 gpu::SubgroupIdOp::create(rewriter, loc, /*upper_bound=*/nullptr);
761 auto sgOffsets =
762 layout.computeDistributedCoords(rewriter, loc, sgId, wgShape);
763 if (failed(sgOffsets))
764 return failure();
765
766 SmallVector<Value, 2> strideConsts;
767 strideConsts.push_back(
768 arith::ConstantIndexOp::create(rewriter, loc, colStride));
769 if (rows > 1)
770 strideConsts.insert(
771 strideConsts.begin(),
772 arith::ConstantIndexOp::create(rewriter, loc, rowStride));
773
774 SmallVector<Value> newConstOps;
775 for (auto offsets : *sgOffsets) {
776 // Multiply offset with stride, broadcast it and add to baseConstVec
777 Value mulOffset = arith::ConstantIndexOp::create(rewriter, loc, 0);
778 for (size_t i = 0; i < strideConsts.size(); ++i) {
779 Value mul =
780 arith::MulIOp::create(rewriter, loc, rewriter.getIndexType(),
781 offsets[i], strideConsts[i]);
782 mulOffset = arith::AddIOp::create(
783 rewriter, loc, rewriter.getIndexType(), mulOffset, mul);
784 }
785 // Broadcast to baseConstVec size
786 auto bcastOffset = vector::BroadcastOp::create(
787 rewriter, loc, baseConstVec.getType(), mulOffset);
788 auto finalConst =
789 arith::AddIOp::create(rewriter, loc, baseConstVec, bcastOffset);
790 newConstOps.push_back(finalConst);
791 }
792 rewriter.replaceOpWithMultiple(op, {newConstOps});
793 return success();
794 }
795 }
796};
797
798// This pattern transforms the LoadGatherOp with explicit offsets to load
799// subgroup data
800struct WgToSgLoadGatherOp : public OpConversionPattern<xegpu::LoadGatherOp> {
801 using OpConversionPattern<xegpu::LoadGatherOp>::OpConversionPattern;
802 LogicalResult
803 matchAndRewrite(xegpu::LoadGatherOp op, OneToNOpAdaptor adaptor,
804 ConversionPatternRewriter &rewriter) const override {
805
806 Location loc = op.getLoc();
807 VectorType resultType = dyn_cast<VectorType>(op.getResult().getType());
808 if (!resultType)
809 return failure();
810 ArrayRef<int64_t> wgShape = resultType.getShape();
811
812 xegpu::DistributeLayoutAttr layout = op.getLayoutAttr();
813
814 if (!layout || !layout.isForWorkgroup())
815 return failure();
816
817 SmallVector<int64_t> sgShape = getSgShapeAndCount(wgShape, layout).first;
818
819 // The offsets need to be distributed
820 auto offsetsVecType =
821 dyn_cast<VectorType>(adaptor.getOffsets().front().getType());
822 auto maskVecType =
823 dyn_cast<VectorType>(adaptor.getMask().front().getType());
824 if (!offsetsVecType || !maskVecType ||
825 offsetsVecType.getShape() != maskVecType.getShape()) {
826 return rewriter.notifyMatchFailure(op,
827 "offsets have not been distributed");
828 }
829
830 SmallVector<Value> newLoadOps;
831 VectorType newTy = VectorType::get(sgShape, resultType.getElementType());
832 for (auto [offsets, mask] :
833 llvm::zip(adaptor.getOffsets(), adaptor.getMask())) {
834 auto newLayout = layout.dropSgLayoutAndData();
835 auto newLoadOp = xegpu::LoadGatherOp::create(
836 rewriter, loc, newTy, op.getSource(), offsets, mask,
837 op.getL1HintAttr(), op.getL2HintAttr(), op.getL3HintAttr(), newLayout,
838 /*contiguity=*/nullptr);
839 newLoadOps.push_back(newLoadOp);
840 }
841 rewriter.replaceOpWithMultiple(op, {newLoadOps});
842 return success();
843 }
844};
845
846// This pattern transforms the StoreScatterOp with explicit offsets to store
847// subgroup data
848struct WgToSgStoreScatterOp
849 : public OpConversionPattern<xegpu::StoreScatterOp> {
850 using OpConversionPattern<xegpu::StoreScatterOp>::OpConversionPattern;
851 LogicalResult
852 matchAndRewrite(xegpu::StoreScatterOp op, OneToNOpAdaptor adaptor,
853 ConversionPatternRewriter &rewriter) const override {
854
855 Location loc = op.getLoc();
856 VectorType valueType = dyn_cast<VectorType>(op.getValue().getType());
857 if (!valueType)
858 return failure();
859
860 xegpu::DistributeLayoutAttr layout = op.getLayoutAttr();
861
862 if (!layout || !layout.isForWorkgroup())
863 return failure();
864
865 // The offsets need to be distributed
866 auto offsetsVecType =
867 dyn_cast<VectorType>(adaptor.getOffsets().front().getType());
868 auto maskVecType =
869 dyn_cast<VectorType>(adaptor.getMask().front().getType());
870 if (!offsetsVecType || !maskVecType ||
871 offsetsVecType.getShape() != maskVecType.getShape()) {
872 return rewriter.notifyMatchFailure(op,
873 "offsets have not been distributed");
874 }
875
876 for (auto [val, offs, mask] : llvm::zip(
877 adaptor.getValue(), adaptor.getOffsets(), adaptor.getMask())) {
878 xegpu::StoreScatterOp::create(
879 rewriter, loc, val, op.getDest(), offs, mask, op.getL1HintAttr(),
880 op.getL2HintAttr(), op.getL3HintAttr(), layout.dropSgLayoutAndData(),
881 /*contiguity=*/nullptr);
882 }
883 rewriter.eraseOp(op);
884 return success();
885 }
886};
887
888struct WgToSgLoadMatrixOp : public OpConversionPattern<xegpu::LoadMatrixOp> {
889 using OpConversionPattern<xegpu::LoadMatrixOp>::OpConversionPattern;
890 LogicalResult
891 matchAndRewrite(xegpu::LoadMatrixOp op, OneToNOpAdaptor adaptor,
892 ConversionPatternRewriter &rewriter) const override {
893
894 SmallVector<SmallVector<OpFoldResult>> offsetsList;
895 if (failed(genOffsetsList(rewriter, op, offsetsList)))
896 return failure();
897
898 ArrayRef<int64_t> wgShape = op.getDataShape();
899 VectorType valueTy = llvm::dyn_cast<VectorType>(op.getRes().getType());
900 assert(valueTy && "the value type must be vector type!");
901 Type elemTy = valueTy.getElementType();
902
903 xegpu::DistributeLayoutAttr layout = op.getLayoutAttr();
904 SmallVector<int64_t> sgShape = getSgShapeAndCount(wgShape, layout).first;
905 VectorType newResTy = VectorType::get(sgShape, elemTy);
906 SmallVector<Value> newOps;
907 for (auto offsets : offsetsList) {
908 auto newOp = xegpu::LoadMatrixOp::create(rewriter, op.getLoc(), newResTy,
909 op.getMemDesc(), offsets,
910 layout.dropSgLayoutAndData());
911 newOps.push_back(newOp);
912 }
913 rewriter.replaceOpWithMultiple(op, {newOps});
914
915 return success();
916 }
917};
918
919struct WgToSgStoreMatrixOp : public OpConversionPattern<xegpu::StoreMatrixOp> {
920 using OpConversionPattern<xegpu::StoreMatrixOp>::OpConversionPattern;
921 LogicalResult
922 matchAndRewrite(xegpu::StoreMatrixOp op, OneToNOpAdaptor adaptor,
923 ConversionPatternRewriter &rewriter) const override {
924
925 SmallVector<SmallVector<OpFoldResult>> offsetsList;
926 if (failed(genOffsetsList(rewriter, op, offsetsList)))
927 return failure();
928
929 xegpu::DistributeLayoutAttr layout = op.getLayoutAttr();
930 for (auto [v, offsets] : llvm::zip(adaptor.getData(), offsetsList))
931 xegpu::StoreMatrixOp::create(rewriter, op.getLoc(), v, op.getMemDesc(),
932 offsets, layout.dropSgLayoutAndData());
933 rewriter.eraseOp(op);
934 return success();
935 }
936};
937
938// This pattern distributes the vector.step ops to work at subgroup level
939struct WgToSgVectorStepOp : public OpConversionPattern<vector::StepOp> {
940 using OpConversionPattern<vector::StepOp>::OpConversionPattern;
941 LogicalResult
942 matchAndRewrite(vector::StepOp op, OneToNOpAdaptor adaptor,
943 ConversionPatternRewriter &rewriter) const override {
944 xegpu::DistributeLayoutAttr layout =
945 xegpu::getTemporaryLayout(dyn_cast<OpResult>(op.getResult()));
946 if (!layout || !layout.isForWorkgroup())
947 return failure();
948
949 Location loc = op.getLoc();
950 VectorType type = op.getResult().getType();
951 auto wgShape = type.getShape();
952 std::optional<SmallVector<int64_t>> sgShape =
953 getSgShapeAndCount(wgShape, layout).first;
954 if (!sgShape)
955 return failure();
956
957 Value sgId =
958 gpu::SubgroupIdOp::create(rewriter, loc, /*upper_bound=*/nullptr);
959 auto sgOffsets =
960 layout.computeDistributedCoords(rewriter, loc, sgId, wgShape);
961 if (failed(sgOffsets))
962 return failure();
963
964 VectorType newTy = type.cloneWith(*sgShape, type.getElementType());
965 auto steps = vector::StepOp::create(rewriter, loc, newTy);
966 SmallVector<Value> newOps;
967 for (auto offsets : *sgOffsets) {
968 // Broadcast the offset scalar to a vector & add to the base steps
969 auto bcastOffset =
970 vector::BroadcastOp::create(rewriter, loc, newTy, offsets[0]);
971 auto finalSteps =
972 arith::AddIOp::create(rewriter, loc, steps, bcastOffset);
973 newOps.push_back(finalSteps);
974 }
975
976 rewriter.replaceOpWithMultiple(op, {newOps});
977 return success();
978 }
979};
980
981// This pattern transforms vector.shape_cast ops to work at subgroup level.
982struct WgToSgVectorShapeCastOp
983 : public OpConversionPattern<vector::ShapeCastOp> {
984 using OpConversionPattern<vector::ShapeCastOp>::OpConversionPattern;
985
986 LogicalResult
987 matchAndRewrite(vector::ShapeCastOp op, OneToNOpAdaptor adaptor,
988 ConversionPatternRewriter &rewriter) const override {
989
990 VectorType resultType = dyn_cast<VectorType>(op.getResult().getType());
991 if (!resultType)
992 return failure();
993
994 ArrayRef<int64_t> wgShape = resultType.getShape();
995 xegpu::DistributeLayoutAttr layout =
996 xegpu::getTemporaryLayout(dyn_cast<OpResult>(op.getResult()));
997 if (!layout || !layout.isForWorkgroup())
998 return failure();
999
1000 // Check that srcShape and destShape, if they differ, only differ by
1001 // expand of unit dimensions.
1002 auto srcType = dyn_cast<VectorType>(op.getSource().getType());
1003 if (!srcType)
1004 return failure();
1005
1006 ArrayRef<int64_t> srcShape = srcType.getShape();
1007
1008 xegpu::DistributeLayoutAttr layoutToDistribute = layout;
1009 SmallVector<int64_t> expandedUnitDims;
1010 if (xegpu::matchUnitDimExpansion(srcShape, wgShape, expandedUnitDims)) {
1011 xegpu::DistributeLayoutAttr sourceLayout =
1012 xegpu::getTemporaryLayout(op->getOpOperand(0));
1013
1014 if (!sourceLayout.isSliceOf(layout))
1015 return rewriter.notifyMatchFailure(
1016 op, "The ShapeCast op only expands dimensions, the input layout "
1017 "must be a slice of the result layout.");
1018
1019 assert(layoutToDistribute.isEqualTo(
1020 layoutToDistribute.setUnitDimData(expandedUnitDims)) &&
1021 "The sg_data for unit dimensions should be set as 1");
1022 }
1023
1024 SmallVector<int64_t> sgShape =
1025 getSgShapeAndCount(wgShape, layoutToDistribute).first;
1026 VectorType newResultType =
1027 VectorType::get(sgShape, resultType.getElementType());
1028
1029 SmallVector<Value> newShapeCastOps;
1030 for (auto src : adaptor.getSource()) {
1031 auto newShapeCast = vector::ShapeCastOp::create(rewriter, op.getLoc(),
1032 newResultType, src);
1033 newShapeCastOps.push_back(newShapeCast.getResult());
1034 }
1035
1036 rewriter.replaceOpWithMultiple(op, {newShapeCastOps});
1037 return success();
1038 }
1039};
1040
1041/// This pattern transforms vector.multi_dim_reduction operations from
1042/// workgroup-level to subgroup-level execution with support for multiple
1043/// reduction dimensions.
1044///
1045/// Steps include:
1046/// 1. LOCAL REDUCTION :
1047/// - Each subgroup performs local reduction on its data slice
1048/// - Uses ZERO accumulator to avoid double-counting during cross-subgroup
1049/// phase
1050///
1051/// 2. CROSS-SUBGROUP :
1052/// - Determines if cross-subgroup reduction is needed (when sg_layout > 1 in
1053/// reduction dims & sgData[reduction dims] < wgData[reduction dims])
1054/// - If not needed, adds original accumulator and returns local results
1055///
1056/// 3. SHARED LOCAL MEMORY (SLM) PHASE (when cross-subgroup reduction needed):
1057/// a) SLM Layout Design:
1058/// - Rows: subgroups participating in reduction (product of sg_layout in
1059/// reduction dims)
1060/// - Cols: total result elements across non-reduction dimensions
1061///
1062/// b) Store Phase:
1063/// - Each subgroup stores its local reduction result to SLM
1064/// - Row offset: linearized index of subgroup in reduction dimensions
1065/// - Col offset: linearized index of subgroup in non-reduction dimensions
1066///
1067/// c) Load and Final Reduction Phase:
1068/// - Each subgroup loads a column of data (all reduction participants for
1069/// its position)
1070/// - Performs final reduction along the loaded dimension
1071/// - Adds original accumulator to get final result
1072///
1073struct WgToSgMultiDimReductionOp
1074 : public OpConversionPattern<vector::MultiDimReductionOp> {
1075 using OpConversionPattern<vector::MultiDimReductionOp>::OpConversionPattern;
1076
1077 LogicalResult
1078 matchAndRewrite(vector::MultiDimReductionOp op, OneToNOpAdaptor adaptor,
1079 ConversionPatternRewriter &rewriter) const override {
1080 Location loc = op.getLoc();
1081
1082 VectorType srcType = op.getSourceVectorType();
1083 Type resultTy = op.getResult().getType();
1084 VectorType dstVecType = dyn_cast<VectorType>(resultTy);
1085 bool isScalarResult = !dstVecType;
1086
1087 auto originalSrcShape = srcType.getShape();
1088 Type elemTy = srcType.getElementType();
1089
1090 xegpu::DistributeLayoutAttr layout =
1091 xegpu::getTemporaryLayout(dyn_cast<OpResult>(op.getResult()));
1092 if (!layout || !layout.isForWorkgroup())
1093 return failure();
1094
1095 auto reductionDims = llvm::to_vector(op.getReductionDims());
1096
1097 // Get sg_layout and sg_data from the parent layout
1098 SmallVector<int64_t> sgLayout;
1099 SmallVector<int64_t> sgData;
1100 xegpu::DistributeLayoutAttr parentLayout;
1101 if (auto sliceAttr = dyn_cast<xegpu::SliceAttr>(layout)) {
1102 parentLayout = sliceAttr.getParent();
1103 sgLayout = parentLayout.getEffectiveSgLayoutAsInt();
1104 sgData = parentLayout.getEffectiveSgDataAsInt();
1105 } else
1106 return rewriter.notifyMatchFailure(
1107 op, "Reduction should have SliceAttr layout");
1108
1109 // Step 1: perform local subgroup reductions with neutral accumulator
1110 SmallVector<Value> localReductions;
1111 auto sgSrcs = adaptor.getSource();
1112 auto sgSrcType = dyn_cast<VectorType>(sgSrcs.front().getType());
1113 SmallVector<int64_t> sgSrcShape(sgSrcType.getShape().begin(),
1114 sgSrcType.getShape().end());
1115
1116 // Determine the SG-level destination type.
1117 // For scalar results (all dims reduced), the sg result is also scalar.
1118 // For vector results, compute the sg destination shape from layout.
1119 Type sgDstType;
1120 if (dstVecType) {
1121 auto originalDstShape = dstVecType.getShape();
1122 SmallVector<int64_t> sgDstShape =
1123 getSgShapeAndCount(originalDstShape, layout).first;
1124 sgDstType = VectorType::get(sgDstShape, elemTy);
1125 } else {
1126 sgDstType = elemTy;
1127 }
1128
1129 for (auto sgSrc : sgSrcs) {
1130 // Create neutral accumulator for local reduction
1131 Value neutralLocalAcc = xegpu::createReductionNeutralValue(
1132 rewriter, loc, sgDstType, op.getKind());
1133 // Local reduction with neutral accumulator
1134 auto localReduce = vector::MultiDimReductionOp::create(
1135 rewriter, loc, sgDstType, op.getKind(), sgSrc, neutralLocalAcc,
1136 reductionDims);
1137 localReductions.push_back(localReduce.getResult());
1138 }
1139
1140 // Check if cross-subgroup reduction is needed for any reduction dimension
1141 SmallVector<int64_t> crossSgReductionDims;
1142 for (int64_t reductionDim : reductionDims) {
1143 bool needsCrossSubgroupReduction =
1144 (sgLayout[reductionDim] > 1) &&
1145 (sgData[reductionDim] < originalSrcShape[reductionDim]);
1146
1147 if (needsCrossSubgroupReduction) {
1148 crossSgReductionDims.push_back(reductionDim);
1149 }
1150 }
1151
1152 // If no cross-subgroup reduction needed, add accumulator and return
1153 if (crossSgReductionDims.empty()) {
1154 SmallVector<Value> results;
1155 for (auto localResult : localReductions) {
1156 auto finalResult = vector::makeArithReduction(
1157 rewriter, loc, op.getKind(), localResult, adaptor.getAcc()[0]);
1158 results.push_back(finalResult);
1159 }
1160 rewriter.replaceOpWithMultiple(op, {results});
1161 return success();
1162 }
1163
1164 // Step 2: cross-subgroup reduction using SLM - allocating slm memory
1165 auto slmStoreDataShape = sgSrcShape;
1166 for (int64_t dim : reductionDims)
1167 slmStoreDataShape[dim] = 1;
1168 VectorType slmStoreDataType = VectorType::get(slmStoreDataShape, elemTy);
1169 SmallVector<Value> slmStoreData;
1170 for (auto localResult : localReductions) {
1171 if (isScalarResult) {
1172 // Scalar result: broadcast scalar to vector<1x...x1> for SLM store
1173 slmStoreData.push_back(vector::BroadcastOp::create(
1174 rewriter, loc, slmStoreDataType, localResult));
1175 } else {
1176 slmStoreData.push_back(vector::ShapeCastOp::create(
1177 rewriter, loc, slmStoreDataType, localResult));
1178 }
1179 }
1180 // for reduction dimension, SLM stores partial results from each subgroup
1181 SmallVector<int64_t> slmShape(originalSrcShape.begin(),
1182 originalSrcShape.end());
1183 SmallVector<int> slmSgData(sgData.begin(), sgData.end());
1184 SmallVector<int> slmSgLayout(sgLayout.begin(), sgLayout.end());
1185 for (int dim : reductionDims) {
1186 slmShape[dim] = sgLayout[dim];
1187 slmSgData[dim] = 1;
1188 }
1189 xegpu::LayoutAttr slmStoreLayout =
1190 xegpu::LayoutAttr::get(rewriter.getContext(), slmSgLayout, slmSgData);
1191
1192 // Allocate SLM
1193 auto bitWidth = elemTy.getIntOrFloatBitWidth();
1194 auto bytesPerElement = bitWidth / 8;
1195 auto slmSize = computeProduct(slmShape) * bytesPerElement;
1196 auto slmTy = MemRefType::get({slmSize}, rewriter.getI8Type(), {}, 3);
1197 auto slm = memref::AllocaOp::create(rewriter, loc, slmTy);
1198
1199 auto memDescType = xegpu::MemDescType::get(rewriter.getContext(), slmShape,
1200 elemTy, nullptr);
1201 auto memDesc =
1202 xegpu::CreateMemDescOp::create(rewriter, loc, memDescType, slm);
1203
1204 // Step 3: Store local results to SLM
1205 auto sgId = gpu::SubgroupIdOp::create(rewriter, loc,
1206 rewriter.getIndexType(), nullptr);
1207
1208 auto slmStoreCoords =
1209 slmStoreLayout.computeDistributedCoords(rewriter, loc, sgId, slmShape);
1210 if (failed(slmStoreCoords))
1211 return failure();
1212 for (auto [data, coord] : llvm::zip(slmStoreData, *slmStoreCoords)) {
1213 SmallVector<OpFoldResult> coordOfr(coord.begin(), coord.end());
1214 xegpu::StoreMatrixOp::create(rewriter, loc, data, memDesc.getResult(),
1215 coordOfr,
1216 /*layout=*/nullptr);
1217 }
1218
1219 gpu::BarrierOp::create(rewriter, loc);
1220
1221 // Step 4: Load from SLM for final reduction
1222 SmallVector<int64_t> slmLoadDataShape(sgSrcShape.begin(), sgSrcShape.end());
1223 for (int64_t dim : reductionDims) {
1224 slmLoadDataShape[dim] = slmShape[dim];
1225 slmSgData[dim] = slmShape[dim];
1226 }
1227 xegpu::LayoutAttr slmLoadLayout =
1228 xegpu::LayoutAttr::get(rewriter.getContext(), slmSgLayout, slmSgData);
1229 auto slmLoadCoords =
1230 slmLoadLayout.computeDistributedCoords(rewriter, loc, sgId, slmShape);
1231 if (failed(slmLoadCoords))
1232 return failure();
1233
1234 VectorType slmLoadType = VectorType::get(slmLoadDataShape, elemTy);
1235 SmallVector<Value> slmLoadData;
1236 for (auto coord : *slmLoadCoords) {
1237 SmallVector<OpFoldResult> coordOfr(coord.begin(), coord.end());
1238 slmLoadData.push_back(xegpu::LoadMatrixOp::create(
1239 rewriter, loc, slmLoadType, memDesc.getResult(), coordOfr,
1240 /*layout=*/nullptr));
1241 }
1242
1243 // Step 5: Perform final reduction with neutral accumulator and add the
1244 // original accumulator at the end
1245 Value neutralFinalAcc = xegpu::createReductionNeutralValue(
1246 rewriter, loc, sgDstType, op.getKind());
1247
1248 SmallVector<Value> finalResults;
1249 for (size_t i = 0; i < slmLoadData.size(); ++i) {
1250 auto loaded = slmLoadData[i];
1251 auto finalReduce = vector::MultiDimReductionOp::create(
1252 rewriter, loc, sgDstType, op.getKind(), loaded, neutralFinalAcc,
1253 reductionDims);
1254 finalResults.push_back(vector::makeArithReduction(
1255 rewriter, loc, op.getKind(), finalReduce.getResult(),
1256 adaptor.getAcc()[i]));
1257 }
1258 rewriter.replaceOpWithMultiple(op, {finalResults});
1259 return success();
1260 }
1261};
1262
1263// This pattern transforms vector.transpose ops to work at subgroup level.
1264struct WgToSgVectorTransposeOp
1265 : public OpConversionPattern<vector::TransposeOp> {
1266 using OpConversionPattern<vector::TransposeOp>::OpConversionPattern;
1267
1268 LogicalResult
1269 matchAndRewrite(vector::TransposeOp op, OneToNOpAdaptor adaptor,
1270 ConversionPatternRewriter &rewriter) const override {
1271 VectorType resultType = op.getResultVectorType();
1272
1273 ArrayRef<int64_t> wgShape = resultType.getShape();
1274 xegpu::DistributeLayoutAttr layout =
1275 xegpu::getTemporaryLayout(dyn_cast<OpResult>(op.getResult()));
1276 if (!layout || !layout.isForWorkgroup())
1277 return failure();
1278 xegpu::DistributeLayoutAttr sourceLayout =
1279 xegpu::getTemporaryLayout(op->getOpOperand(0));
1280 if (!sourceLayout || !sourceLayout.isForWorkgroup())
1281 return failure();
1282
1283 SmallVector<int64_t> sourceSgLayout =
1284 sourceLayout.getEffectiveSgLayoutAsInt();
1285 SmallVector<int64_t> resultSgLayout = layout.getEffectiveSgLayoutAsInt();
1286
1287 ArrayRef<int64_t> permutation = op.getPermutation();
1288 size_t permutationSize = permutation.size();
1289 if (sourceSgLayout.size() != permutationSize ||
1290 resultSgLayout.size() != permutationSize) {
1291 return rewriter.notifyMatchFailure(
1292 op, "Layouts and permutation must have the same rank");
1293 }
1294
1295 // Check that sgLayout, sgData & order are properly transposed for source
1296 // and result
1297 if (!layout.isTransposeOf(sourceLayout, permutation,
1298 xegpu::LayoutKind::Subgroup))
1299 return rewriter.notifyMatchFailure(
1300 op, "Result layout is not a valid transpose of source layout "
1301 "according to permutation");
1302
1303 SmallVector<int64_t> sgShape = getSgShapeAndCount(wgShape, layout).first;
1304 VectorType newResultType =
1305 VectorType::get(sgShape, resultType.getElementType());
1306
1307 SmallVector<Value> newTransposeOps;
1308 for (auto src : adaptor.getVector()) {
1309 auto newTranspose = vector::TransposeOp::create(
1310 rewriter, op.getLoc(), newResultType, src, permutation);
1311 newTransposeOps.push_back(newTranspose.getResult());
1312 }
1313 rewriter.replaceOpWithMultiple(op, {newTransposeOps});
1314 return success();
1315 }
1316};
1317
1318// Distribute vector mask ops to work at subgroup level.
1319template <typename MaskOpType>
1320struct WgToSgVectorMaskOp : public OpConversionPattern<MaskOpType> {
1321 using OpConversionPattern<MaskOpType>::OpConversionPattern;
1322
1323 LogicalResult matchAndRewrite(
1324 MaskOpType op,
1325 typename OpConversionPattern<MaskOpType>::OneToNOpAdaptor adaptor,
1326 ConversionPatternRewriter &rewriter) const override {
1327 xegpu::DistributeLayoutAttr layout =
1328 xegpu::getTemporaryLayout(dyn_cast<OpResult>(op.getResult()));
1329 if (!layout || !layout.isForWorkgroup())
1330 return failure();
1331
1332 Location loc = op.getLoc();
1333 VectorType type = op.getResult().getType();
1334 auto wgShape = type.getShape();
1335
1336 SmallVector<Value> wgMaskDimSizes;
1337 if constexpr (std::is_same_v<MaskOpType, vector::ConstantMaskOp>) {
1338 for (int64_t maskSize : op.getMaskDimSizes()) {
1339 wgMaskDimSizes.push_back(
1340 arith::ConstantIndexOp::create(rewriter, loc, maskSize));
1341 }
1342 } else if constexpr (std::is_same_v<MaskOpType, vector::CreateMaskOp>) {
1343 wgMaskDimSizes = llvm::to_vector(op.getOperands());
1344 }
1345
1346 Value sgId =
1347 gpu::SubgroupIdOp::create(rewriter, loc, /*upper_bound=*/nullptr);
1348 auto sgOffsets =
1349 layout.computeDistributedCoords(rewriter, loc, sgId, wgShape);
1350 if (failed(sgOffsets))
1351 return failure();
1352
1353 SmallVector<int64_t> sgShape = getSgShapeAndCount(wgShape, layout).first;
1354 VectorType resultType = VectorType::get(sgShape, type.getElementType());
1355
1356 // In each dimension, each subgroup computes its local mask size as:
1357 // min(max(wgMaskDimSize[d] - offset[d], 0), sgDimSize[d])
1358 SmallVector<Value> newCreateMaskOps;
1359 for (auto offsetSet : *sgOffsets) {
1360 SmallVector<Value> maskOperands;
1361
1362 for (auto [i, wgMaskDimSize] : llvm::enumerate(wgMaskDimSizes)) {
1363 Value dimSizeVal =
1364 arith::ConstantIndexOp::create(rewriter, loc, sgShape[i]);
1365 Value offset = offsetSet[i];
1366 Value adjustedMaskSize =
1367 arith::SubIOp::create(rewriter, loc, wgMaskDimSize, offset);
1368 Value zero = arith::ConstantIndexOp::create(rewriter, loc, 0);
1369 Value nonNegative =
1370 arith::MaxSIOp::create(rewriter, loc, adjustedMaskSize, zero);
1371 Value sgMaskSize =
1372 arith::MinSIOp::create(rewriter, loc, nonNegative, dimSizeVal);
1373 maskOperands.push_back(sgMaskSize);
1374 }
1375
1376 auto newCreateMaskOp =
1377 vector::CreateMaskOp::create(rewriter, loc, resultType, maskOperands);
1378 newCreateMaskOps.push_back(newCreateMaskOp.getResult());
1379 }
1380
1381 rewriter.replaceOpWithMultiple(op, {newCreateMaskOps});
1382 return success();
1383 }
1384};
1385
1386using WgToSgVectorConstantMaskOp = WgToSgVectorMaskOp<vector::ConstantMaskOp>;
1387using WgToSgVectorCreateMaskOp = WgToSgVectorMaskOp<vector::CreateMaskOp>;
1388
1389// This pattern transforms vector.bitcast ops to work at subgroup level.
1390struct WgToSgVectorBitCastOp : public OpConversionPattern<vector::BitCastOp> {
1391 using OpConversionPattern<vector::BitCastOp>::OpConversionPattern;
1392
1393 LogicalResult
1394 matchAndRewrite(vector::BitCastOp op, OneToNOpAdaptor adaptor,
1395 ConversionPatternRewriter &rewriter) const override {
1396 VectorType resultType = op.getResultVectorType();
1397
1398 ArrayRef<int64_t> wgShape = resultType.getShape();
1399 xegpu::DistributeLayoutAttr layout =
1400 xegpu::getTemporaryLayout(dyn_cast<OpResult>(op.getResult()));
1401 if (!layout || !layout.isForWorkgroup())
1402 return failure();
1403
1404 SmallVector<int64_t> sgShape = getSgShapeAndCount(wgShape, layout).first;
1405 VectorType newResultType =
1406 VectorType::get(sgShape, resultType.getElementType());
1407
1408 SmallVector<Value> newBitCastOps;
1409 for (auto src : adaptor.getSource()) {
1410 auto newBitCast =
1411 vector::BitCastOp::create(rewriter, op.getLoc(), newResultType, src);
1412 newBitCastOps.push_back(newBitCast.getResult());
1413 }
1414
1415 rewriter.replaceOpWithMultiple(op, {newBitCastOps});
1416 return success();
1417 }
1418};
1419
1420// This pattern transforms vector.interleave ops to work at subgroup level.
1421struct WgToSgVectorInterleaveOp
1422 : public OpConversionPattern<vector::InterleaveOp> {
1423 using OpConversionPattern<vector::InterleaveOp>::OpConversionPattern;
1424
1425 LogicalResult
1426 matchAndRewrite(vector::InterleaveOp op, OneToNOpAdaptor adaptor,
1427 ConversionPatternRewriter &rewriter) const override {
1428 VectorType resultType = op.getResultVectorType();
1429
1430 ArrayRef<int64_t> wgShape = resultType.getShape();
1431 xegpu::DistributeLayoutAttr layout =
1432 xegpu::getTemporaryLayout(dyn_cast<OpResult>(op.getResult()));
1433 if (!layout || !layout.isForWorkgroup())
1434 return failure();
1435
1436 SmallVector<int64_t> sgShape = getSgShapeAndCount(wgShape, layout).first;
1437 VectorType newResultType =
1438 VectorType::get(sgShape, resultType.getElementType());
1439
1440 SmallVector<Value> newInterleaveOps;
1441 // Interleave operates pairwise: each lhs value is interleaved with
1442 // corresponding rhs value
1443 for (auto [lhs, rhs] : llvm::zip(adaptor.getLhs(), adaptor.getRhs())) {
1444 auto newInterleave = vector::InterleaveOp::create(
1445 rewriter, op.getLoc(), newResultType, lhs, rhs);
1446 newInterleaveOps.push_back(newInterleave.getResult());
1447 }
1448
1449 rewriter.replaceOpWithMultiple(op, {newInterleaveOps});
1450 return success();
1451 }
1452};
1453
1454// This pattern transforms vector.deinterleave ops to work at subgroup level.
1455struct WgToSgVectorDeinterleaveOp
1456 : public OpConversionPattern<vector::DeinterleaveOp> {
1457 using OpConversionPattern<vector::DeinterleaveOp>::OpConversionPattern;
1458
1459 LogicalResult
1460 matchAndRewrite(vector::DeinterleaveOp op, OneToNOpAdaptor adaptor,
1461 ConversionPatternRewriter &rewriter) const override {
1462 SmallVector<Value> newRes1Ops;
1463 SmallVector<Value> newRes2Ops;
1464
1465 for (auto src : adaptor.getSource()) {
1466 auto newDeinterleave =
1467 vector::DeinterleaveOp::create(rewriter, op.getLoc(), src);
1468 newRes1Ops.push_back(newDeinterleave.getRes1());
1469 newRes2Ops.push_back(newDeinterleave.getRes2());
1470 }
1471
1472 SmallVector<SmallVector<Value>> results = {newRes1Ops, newRes2Ops};
1473 rewriter.replaceOpWithMultiple(op, results);
1474 return success();
1475 }
1476};
1477
1478} // namespace
1479
1480namespace mlir {
1481namespace xegpu {
1483 Operation *topLevelOp) {
1484 // Pass through all types by default.
1485 converter.addConversion([](Type type) -> Type { return type; });
1486
1487 // For TensorDescType, convert WG-level tensor descs to N SG-level descs.
1488 converter.addConversion(
1489 [](xegpu::TensorDescType type,
1490 SmallVectorImpl<Type> &result) -> std::optional<LogicalResult> {
1491 xegpu::DistributeLayoutAttr layout = type.getLayoutAttr();
1492 if (!layout || !layout.isForWorkgroup())
1493 return std::nullopt;
1494
1495 Type elemTy = type.getElementType();
1496 ArrayRef<int64_t> shape = type.getShape();
1497
1498 int count;
1499 SmallVector<int64_t> subShape;
1500 std::tie(subShape, count) = getSgShapeAndCount(shape, layout);
1501
1502 layout = layout.dropSgLayoutAndData();
1503
1504 auto newTy = xegpu::TensorDescType::get(
1505 type.getContext(), subShape, elemTy, type.getEncoding(), layout);
1506 result.append(count, newTy);
1507 return success();
1508 });
1509
1510 // Context-aware VectorType conversion based on sg_layout/sg_data
1511 // (1:1 shape-changing or 1:N).
1512 auto getSubShapeAndCount = [](VectorType vecTy,
1513 xegpu::DistributeLayoutAttr layout)
1514 -> std::pair<SmallVector<int64_t>, int> {
1515 if (!layout.isForWorkgroup())
1516 return {{}, 0};
1517 return getSgShapeAndCount(vecTy.getShape(), layout);
1518 };
1519 auto loopArgTypes =
1520 xegpu::precomputeLoopBlockArgTypes(topLevelOp, getSubShapeAndCount);
1521 xegpu::addVectorTypeConversion(converter, getSubShapeAndCount,
1522 std::move(loopArgTypes));
1523}
1524
1526 patterns.add<WgToSgCreateNdOp, WgToSgLoadNdOp, WgToSgStoreNdOp, WgToSgDpasOp,
1527 WgToSgDpasMxOp, WgToSgPrefetchNdOp, WgToSgElementwiseOp,
1528 WgToSgVectorBroadcastOp, WgToSgConvertLayoutOp,
1529 WgToSgArithConstantOp, WgToSgLoadGatherOp, WgToSgStoreScatterOp,
1530 WgToSgLoadMatrixOp, WgToSgStoreMatrixOp, WgToSgVectorStepOp,
1531 WgToSgVectorShapeCastOp, WgToSgMultiDimReductionOp,
1532 WgToSgVectorTransposeOp, WgToSgVectorConstantMaskOp,
1533 WgToSgVectorCreateMaskOp, WgToSgVectorBitCastOp,
1534 WgToSgVectorInterleaveOp, WgToSgVectorDeinterleaveOp>(
1535 patterns.getContext());
1536}
1537} // namespace xegpu
1538} // namespace mlir
1539
1540namespace {
1541struct XeGPUWgToSgDistributePass
1542 : public xegpu::impl::XeGPUWgToSgDistributeBase<XeGPUWgToSgDistributePass> {
1543 void runOnOperation() override;
1544};
1545} // namespace
1546
1547void XeGPUWgToSgDistributePass::runOnOperation() {
1548
1549 Operation *op = getOperation();
1551 signalPassFailure();
1552 return;
1553 }
1554
1555 // Collect existing UnrealizedConversionCastOps. These must be preserved.
1556 llvm::SmallSetVector<UnrealizedConversionCastOp, 8> existingCasts;
1557 getOperation()->walk(
1558 [&](UnrealizedConversionCastOp castOp) { existingCasts.insert(castOp); });
1559
1560 // Perform workgroup to subgroup distribution for TensorDesc and Vector
1561 // values, as well as XeGPU, Arith, and Vector operations. Uses a
1562 // context-aware type converter that inspects Values to retrieve the
1563 // distribute layout attribute for 1:N type conversion.
1564 MLIRContext *ctx = &getContext();
1565 RewritePatternSet patterns(ctx);
1566 ConversionTarget target(*ctx);
1567 TypeConverter converter;
1568 // Source (N:1) and target (1:1) materializations using
1569 // UnrealizedConversionCastOp.
1570 auto materializeCast = [](OpBuilder &builder, Type type, ValueRange inputs,
1571 Location loc) -> Value {
1572 return UnrealizedConversionCastOp::create(builder, loc, type, inputs)
1573 .getResult(0);
1574 };
1575 converter.addSourceMaterialization(materializeCast);
1576 converter.addTargetMaterialization(materializeCast);
1578 getOperation());
1579
1580 auto getTensorDescType = [](Operation *op) -> xegpu::TensorDescType {
1581 if (auto createOp = dyn_cast<xegpu::CreateNdDescOp>(op))
1582 return createOp.getType();
1583 if (auto loadOp = dyn_cast<xegpu::LoadNdOp>(op))
1584 return loadOp.getTensorDescType();
1585 if (auto storeOp = dyn_cast<xegpu::StoreNdOp>(op))
1586 return storeOp.getTensorDescType();
1587 if (auto prefetchOp = dyn_cast<xegpu::PrefetchNdOp>(op))
1588 return prefetchOp.getTensorDescType();
1589 return xegpu::TensorDescType();
1590 };
1591
1592 auto isLegal = [&](xegpu::DistributeLayoutAttr layout) -> bool {
1593 return !layout || !layout.isForWorkgroup();
1594 };
1595
1596 target.addDynamicallyLegalOp<xegpu::CreateNdDescOp, xegpu::LoadNdOp,
1597 xegpu::StoreNdOp, xegpu::PrefetchNdOp>(
1598 [=](Operation *op) -> bool {
1599 auto tdescTy = getTensorDescType(op);
1600 auto layout = dyn_cast_if_present<xegpu::DistributeLayoutAttr>(
1601 tdescTy.getLayout());
1602 return isLegal(layout);
1603 });
1604
1605 target.addDynamicallyLegalOp<xegpu::DpasOp>([=](xegpu::DpasOp op) -> bool {
1606 auto layout = op.getLayoutCdAttr();
1607 return isLegal(layout);
1608 });
1609
1610 target.addDynamicallyLegalOp<xegpu::DpasMxOp>(
1611 [=](xegpu::DpasMxOp op) -> bool {
1612 auto layout = op.getLayoutCdAttr();
1613 return isLegal(layout);
1614 });
1615
1616 target.addDynamicallyLegalOp<xegpu::LoadMatrixOp>(
1617 [=](xegpu::LoadMatrixOp op) -> bool {
1618 return isLegal(op.getLayoutAttr());
1619 });
1620
1621 target.addDynamicallyLegalOp<xegpu::StoreMatrixOp>(
1622 [=](xegpu::StoreMatrixOp op) -> bool {
1623 return isLegal(op.getLayoutAttr());
1624 });
1625
1626 target.addDynamicallyLegalOp<arith::ConstantOp>(
1627 [=](arith::ConstantOp op) -> bool {
1628 auto vecType = dyn_cast<VectorType>(op.getType());
1629 if (!vecType)
1630 return true;
1631
1632 auto layout =
1633 xegpu::getTemporaryLayout(dyn_cast<OpResult>(op.getResult()));
1634 return isLegal(layout);
1635 });
1636
1637 target.addDynamicallyLegalOp<
1638 vector::ShapeCastOp, vector::StepOp, vector::TransposeOp,
1639 vector::BroadcastOp, vector::MultiDimReductionOp, vector::ConstantMaskOp,
1640 vector::CreateMaskOp, vector::BitCastOp, vector::InterleaveOp,
1641 vector::DeinterleaveOp>([=](Operation *op) -> bool {
1642 // Check for either a SliceAttr or LayoutAttr on the result.
1643 auto layout =
1644 xegpu::getTemporaryLayout(dyn_cast<OpResult>(op->getResult(0)));
1645 return isLegal(layout);
1646 });
1647
1648 target.addDynamicallyLegalOp<xegpu::LoadGatherOp>(
1649 [=](xegpu::LoadGatherOp op) -> bool {
1650 auto layout = op.getLayoutAttr();
1651 return isLegal(layout);
1652 });
1653
1654 target.addDynamicallyLegalOp<xegpu::StoreScatterOp>(
1655 [=](xegpu::StoreScatterOp op) -> bool {
1656 auto layout = op.getLayoutAttr();
1657 return isLegal(layout);
1658 });
1659
1660 target.addDynamicallyLegalOp<xegpu::ConvertLayoutOp>(
1661 [=](xegpu::ConvertLayoutOp op) -> bool {
1662 return isLegal(op.getEffectiveInputLayout()) &&
1663 isLegal(op.getTargetLayout());
1664 });
1665
1666 target.addDynamicallyLegalDialect<math::MathDialect, arith::ArithDialect>(
1667 [=](Operation *op) -> std::optional<bool> {
1668 // Only handle elementwise mappable ops
1670 return true;
1671
1672 VectorType resultType =
1673 dyn_cast<VectorType>(op->getResult(0).getType());
1674 if (!resultType)
1675 return true;
1676
1677 // Check if all operands are vectors of the same shape
1678 // TODO: Support other types.
1679 for (Value operand : op->getOperands()) {
1680 VectorType operandType = dyn_cast<VectorType>(operand.getType());
1681 if (!operandType || operandType.getShape() != resultType.getShape()) {
1682 return true;
1683 }
1684 }
1685
1686 xegpu::DistributeLayoutAttr layout =
1688 return isLegal(layout);
1689 });
1690
1691 target.addLegalOp<UnrealizedConversionCastOp>();
1692
1693 target.markUnknownOpDynamicallyLegal([](Operation *) { return true; });
1694
1696 target);
1698 if (failed(
1699 applyPartialConversion(getOperation(), target, std::move(patterns))))
1700 return signalPassFailure();
1701
1702 // Fold cancelling cast chains and erase dead casts.
1703 xegpu::cleanupUnrealizedConversionCasts(getOperation(), existingCasts);
1704 xegpu::removeTemporaryLayoutAttrs(getOperation());
1705}
return success()
lhs
b getContext())
#define mul(a, b)
static DenseElementsAttr get(ShapedType type, ArrayRef< Attribute > values)
Constructs a dense elements attribute from an array of element values.
This class defines the main interface for locations in MLIR and acts as a non-nullable wrapper around...
Definition Location.h:76
MLIRContext * getContext() const
Return the context this location is uniqued in.
Definition Location.h:86
Operation is the basic unit of execution within MLIR.
Definition Operation.h:87
Attribute getDiscardableAttr(StringRef name)
Access a discardable attribute by name, returns a null Attribute if the discardable attribute does no...
Definition Operation.h:485
OpResult getResult(unsigned idx)
Get the 'idx'th result of this operation.
Definition Operation.h:432
Location getLoc()
The source location the operation was defined or derived from.
Definition Operation.h:240
Attribute getPropertiesAsAttribute()
Return the properties converted to an attribute.
OpTy getParentOfType()
Return the closest surrounding parent operation that is of type 'OpTy'.
Definition Operation.h:255
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
operand_range getOperands()
Returns an iterator on the underlying Value's.
Definition Operation.h:403
unsigned getNumResults()
Return the number of results held by this operation.
Definition Operation.h:429
MLIRContext * getContext() const
RewritePatternSet & add(ConstructorArg &&arg, ConstructorArgs &&...args)
Add an instance of each of the pattern types 'Ts' to the pattern list with the given arguments.
Instances of the Type class are uniqued, have an immutable identifier and an optional mutable compone...
Definition Types.h:74
MLIRContext * getContext() const
Return the MLIRContext in which this type was uniqued.
Definition Types.cpp:35
bool isIntOrFloat() const
Return true if this is an integer (of any signedness) or a float type.
Definition Types.cpp:118
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 represents an instance of an SSA value in the MLIR system, representing a computable value...
Definition Value.h:96
Type getType() const
Return the type of this value.
Definition Value.h:105
static ConstantIndexOp create(OpBuilder &builder, Location location, int64_t value)
Definition ArithOps.cpp:398
bool hasElementwiseMappableTraits(Operation *op)
Together, Elementwise, Scalarizable, Vectorizable, and Tensorizable provide an easy way for scalar op...
detail::InFlightRemark failed(Location loc, RemarkOpts opts)
Report an optimization remark that failed.
Definition Remarks.h:734
void populateSCFStructuralTypeConversionsAndLegality(const TypeConverter &typeConverter, RewritePatternSet &patterns, ConversionTarget &target, PatternBenefit benefit=1)
Populates patterns for SCF structural type conversions and sets up the provided ConversionTarget with...
Value makeArithReduction(OpBuilder &b, Location loc, CombiningKind kind, Value v1, Value acc, arith::FastMathFlagsAttr fastmath=nullptr, Value mask=nullptr)
Returns the result value of reducing two scalar/vector values with the corresponding arith operation.
void removeTemporaryLayoutAttrs(Operation *op)
Removes the temporary layout attributes for each OpOperand and OpResult of the given operation.
void populateXeGPUWgToSgDistributeTypeConversions(TypeConverter &converter, Operation *topLevelOp)
Define the type conversions needed for XeGPU workgroup to subgroup distribution.
Value createReductionNeutralValue(OpBuilder &builder, Location loc, Type type, vector::CombiningKind kind)
Creates a constant filled with the neutral (identity) value for the given reduction kind.
bool matchUnitDimExpansion(ArrayRef< int64_t > src, ArrayRef< int64_t > dst, SmallVector< int64_t > &expandedUnitDims)
bool recoverTemporaryLayouts(Operation *rootOp)
Attach layout attributes to all vector-type operands of operations within the given operation's neste...
DenseMap< Value, SmallVector< Type > > precomputeLoopBlockArgTypes(Operation *topLevelOp, SubShapeAndCountFn getSubShapeAndCount)
Pre-computes distributed VectorType mappings for every value carried through an SCF loop under topLev...
void populateXeGPUWgToSgDistributePatterns(RewritePatternSet &patterns)
Appends patterns for XeGPU workgroup to subgroup distribution into patterns.
void addVectorTypeConversion(TypeConverter &converter, SubShapeAndCountFn getSubShapeAndCount, DenseMap< Value, SmallVector< Type > > loopArgTypes)
Adds a context-aware VectorType conversion to converter (1:1 shape-changing or 1:N,...
DistributeLayoutAttr getTemporaryLayout(const T &operandOrResult)
get and set distribute layout attribute for non-anchor operations (and offsets/masks of load/store op...
void removeLayoutAttrs(Operation *op)
Removes the DistributeLayoutAttr for each OpOperand and OpResult of the given operation if they exist...
void cleanupUnrealizedConversionCasts(Operation *root, const llvm::SmallSetVector< UnrealizedConversionCastOp, 8 > &existingCasts)
Cleans up UnrealizedConversionCastOps inserted during SCF structural type conversion and/or XeGPU unr...
SmallVector< OpFoldResult > addWithRightAligned(OpBuilder &builder, Location loc, ArrayRef< OpFoldResult > lhs, ArrayRef< OpFoldResult > rhs)
Generates element-wise addition ops of two arrays with automatic alignment.
Include the generated interface declarations.
int64_t computeProduct(ArrayRef< int64_t > basis)
Self-explicit.
OpFoldResult getAsOpFoldResult(Value val)
Given a value, try to extract a constant Attribute.