MLIR 24.0.0git
XeGPUSgToLaneDistribute.cpp
Go to the documentation of this file.
1//===- XeGPUSgToLaneDistribute.cpp - XeGPU SG to Lane 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//===----------------------------------------------------------------------===//
21#include "mlir/IR/Builders.h"
23#include "mlir/IR/BuiltinOps.h"
25#include "mlir/IR/MLIRContext.h"
26#include "mlir/IR/Operation.h"
27#include "mlir/IR/Value.h"
28#include "mlir/IR/ValueRange.h"
30#include "llvm/ADT/SetVector.h"
31#include "llvm/Support/LogicalResult.h"
32#include "llvm/Support/raw_ostream.h"
33#include <optional>
34
35namespace mlir {
36namespace xegpu {
37#define GEN_PASS_DEF_XEGPUSGTOLANEDISTRIBUTE
38#include "mlir/Dialect/XeGPU/Transforms/Passes.h.inc"
39} // namespace xegpu
40} // namespace mlir
41
42using namespace mlir;
43
44#define DEBUG_TYPE "xegpu-sg-to-lane-distribute"
45#define DBGS() (llvm::dbgs() << "[" DEBUG_TYPE "]: ")
46
47namespace {
48
49/// Casts the given vector value `v` to the expected vector type `expectedTy`.
50static Value castValueTo(ConversionPatternRewriter &rewriter,
51 TypedValue<VectorType> v, VectorType expectedTy) {
52 // If the type matches, simply return the value itself.
53 if (v.getType() == expectedTy)
54 return v;
55 // If only shape differs, use shape cast.
56 if (isa<VectorType>(v.getType()) &&
57 v.getType().getNumElements() == expectedTy.getNumElements())
58 return vector::ShapeCastOp::create(rewriter, v.getLoc(), expectedTy, v);
59
60 // Else create an unrealized cast.
61 auto newOp = UnrealizedConversionCastOp::create(rewriter, v.getLoc(),
62 expectedTy, ValueRange{v});
63 return newOp.getResult(0);
64}
65
66/// A vector::MultiDimReductionOp at subgroup level in expected form if, it has
67/// exactly 1 reduction dimension, it had valid result layout attribute, and
68/// result type can be distributed to lanes using the layout.
69static bool isValidSubgroupMultiReductionOp(vector::MultiDimReductionOp op) {
70 auto resLayout = xegpu::getTemporaryLayout(op->getOpResult(0));
71 // If no layout, not valid.
72 if (!resLayout || !resLayout.isForSubgroup())
73 return false;
74 // Scalar result (e.g., vector<32xf32> to f32) is valid.
75 if (op.getType().isIntOrFloat())
76 return op.getReductionDims().size() == 1;
77 VectorType resTy = dyn_cast<VectorType>(op.getType());
78 if (!resTy)
79 return false;
80 // Compute the distributed result vector type based on the layout.
81 FailureOr<VectorType> resDistTypeOrFailure =
82 getDistVecTypeBasedOnLaneLayout(resLayout, resTy);
83 if (failed(resDistTypeOrFailure))
84 return false;
85 return op.getReductionDims().size() == 1;
86}
87
88/// A vector::MultiDimReductionOp reduces lane-locally when no data is combined
89/// across lanes, which holds exactly when the source's lane_layout along the
90/// reduction dimension is 1: every reduced element then lives in one lane.
91static bool isReductionLaneLocal(vector::MultiDimReductionOp op) {
92 // Must be valid MultiDimReductionOp.
93 assert(isValidSubgroupMultiReductionOp(op) && "Expecting a valid subgroup "
94 "MultiDimReductionOp");
95 auto srcLayout = xegpu::getTemporaryLayout(op->getOpOperand(0));
96 ArrayRef<int64_t> reductionDims = op.getReductionDims();
97 assert(reductionDims.size() == 1 &&
98 "Expecting single reduction dimension for subgroup multi "
99 "reduction op");
100 int64_t reductionDim = reductionDims[0];
101 SmallVector<int64_t> srcLaneLayout = srcLayout.getEffectiveLaneLayoutAsInt();
102 assert(reductionDim < static_cast<int64_t>(srcLaneLayout.size()) &&
103 "Expecting a source lane_layout covering the reduction dimension");
104 return srcLaneLayout[reductionDim] == 1;
105}
106
107/// Given a vector type and its distributed vector type, return the list of
108/// dimensions that are distributed.
109static SmallVector<int64_t> getDistributedDims(VectorType originalType,
110 VectorType distributedType) {
111 assert(originalType.getRank() == distributedType.getRank() &&
112 "original and distributed vector types must have the same rank");
113 SmallVector<int64_t> distributedDims;
114 for (int64_t i = 0; i < originalType.getRank(); ++i) {
115 if (distributedType.getDimSize(i) != originalType.getDimSize(i))
116 distributedDims.push_back(i);
117 }
118 return distributedDims;
119}
120
121/// Distributes a subgroup-level CreateNdDesc op to lane-level CreateNdDesc
122/// op. This simply drops the layout attribute from the tensor descriptor type.
123struct SgToLaneCreateNdDesc
124 : public OpConversionPattern<xegpu::CreateNdDescOp> {
125 using OpConversionPattern<xegpu::CreateNdDescOp>::OpConversionPattern;
126
127 LogicalResult
128 matchAndRewrite(xegpu::CreateNdDescOp op, OpAdaptor adaptor,
129 ConversionPatternRewriter &rewriter) const override {
130 xegpu::TensorDescType resultType = op.getType();
131 // If no layout, nothing to do.
132 if (!resultType.getLayout())
133 return failure();
134
135 auto newOp = xegpu::CreateNdDescOp::create(
136 rewriter, op.getLoc(), TypeRange{resultType.dropLayouts()},
137 op.getOperands(), op.getProperties(),
138 op->getDiscardableAttrDictionary().getValue());
139 rewriter.replaceOp(op, newOp.getResult());
140 return success();
141 }
142};
143
144/// Distributes a subgroup-level LoadNd op to lane-level LoadNd op. Output
145/// of lane-level LoadNd op is 1D. ShapeCast is added to restore the
146/// original rank.
147struct SgToLaneLoadNd : public OpConversionPattern<xegpu::LoadNdOp> {
148 using OpConversionPattern<xegpu::LoadNdOp>::OpConversionPattern;
149
150 LogicalResult
151 matchAndRewrite(xegpu::LoadNdOp op, OpAdaptor adaptor,
152 ConversionPatternRewriter &rewriter) const override {
153 xegpu::DistributeLayoutAttr layout = op.getAnchorLayout();
154 // If no layout, nothing to do.
155 if (!layout)
156 return failure();
157 // Check if the layout attached to the tensor descriptor is same as the
158 // anchor layout. Otherwise, this is a conflict.
159 if (op.getTensorDescType().getLayout() != layout)
160 return rewriter.notifyMatchFailure(
161 op, "conflicting layout attributes on tensor descriptor and anchor");
162 const auto *uArch =
164 if (!uArch)
165 return rewriter.notifyMatchFailure(
166 op, "xegpu::LoadNdOp require target attribute attached to "
167 "determine transpose "
168 "requirement");
169 auto supportedLaneResultTyOrFailure =
170 xegpu::getDistributedVectorType(op.getTensorDescType());
171 auto expectedLaneResultTyOrFailure =
172 xegpu::getDistVecTypeBasedOnLaneLayout(layout, op.getType());
173 if (failed(supportedLaneResultTyOrFailure))
174 return rewriter.notifyMatchFailure(
175 op, "unable to compute the lane vector type for LoadNdOp");
176 if (failed(expectedLaneResultTyOrFailure))
177 return rewriter.notifyMatchFailure(
178 op, "unable to compute expected lane vector type from lane layout");
179 auto newOp = xegpu::LoadNdOp::create(
180 rewriter, op.getLoc(), supportedLaneResultTyOrFailure.value(),
181 adaptor.getTensorDesc(), op.getMixedOffsets(), op.getPackedAttr(),
182 op.getTransposeAttr(), op.getL1HintAttr(), op.getL2HintAttr(),
183 op.getL3HintAttr(), /**layout**/ nullptr);
184 // Set the packed attribute if the layout requires it.
185 newOp.setPacked(xegpu::requirePacked(cast<xegpu::LayoutAttr>(layout)));
186 // Set the transpose attribute if the layout requires it.
187 if (xegpu::requireTranspose(cast<xegpu::LayoutAttr>(layout), uArch))
188 newOp.setTranspose(DenseI64ArrayAttr::get(rewriter.getContext(), {1, 0}));
189 rewriter.replaceOp(op, castValueTo(rewriter, newOp.getResult(),
190 expectedLaneResultTyOrFailure.value()));
191 return success();
192 }
193};
194
195/// Distributes a subgroup-level StoreNd op to lane-level StoreNd op. Stored
196/// value in lane-level StoreNd op is 1D. ShapeCast is added to cast the
197/// incoming value to 1D.
198struct SgToLaneStoreNd : public OpConversionPattern<xegpu::StoreNdOp> {
199 using OpConversionPattern<xegpu::StoreNdOp>::OpConversionPattern;
200
201 LogicalResult
202 matchAndRewrite(xegpu::StoreNdOp op, OpAdaptor adaptor,
203 ConversionPatternRewriter &rewriter) const override {
204 xegpu::DistributeLayoutAttr layout = op.getAnchorLayout();
205 // If no layout, nothing to do.
206 if (!layout)
207 return failure();
208 // Check if the layout attached to the tensor descriptor and value layout is
209 // same as the anchor layout. Otherwise, this is a conflict.
210 if (op.getTensorDescType().getLayout() != layout)
211 return rewriter.notifyMatchFailure(
212 op, "conflicting layout attributes on tensor descriptor and anchor");
213 auto valueLayout = xegpu::getDistributeLayoutAttr(op->getOpOperand(0));
214 if (valueLayout != layout)
215 return rewriter.notifyMatchFailure(
216 op, "conflicting layout attributes on value and anchor");
217 auto supportedLaneValueTyOrFailure =
218 xegpu::getDistributedVectorType(op.getTensorDescType());
219 if (failed(supportedLaneValueTyOrFailure))
220 return rewriter.notifyMatchFailure(
221 op,
222 "unable to compute lane vector type for StoreNdOp value from tensor "
223 "descriptor");
224
225 xegpu::StoreNdOp::create(
226 rewriter, op.getLoc(),
227 castValueTo(rewriter, cast<TypedValue<VectorType>>(adaptor.getValue()),
228 supportedLaneValueTyOrFailure.value()),
229 adaptor.getTensorDesc(), op.getMixedOffsets(), op.getL1HintAttr(),
230 op.getL2HintAttr(), op.getL3HintAttr(), /**layout**/ nullptr);
231 rewriter.eraseOp(op);
232 return success();
233 }
234};
235
236/// Distributes a subgroup-level Dpas op to lane-level Dpas op. All inpputs
237/// and output of lane-level Dpas op are 1D. Necessary casts are added to
238/// convert the inputs and output to/from 1D.
239struct SgToLaneDpas : public OpConversionPattern<xegpu::DpasOp> {
240 using OpConversionPattern<xegpu::DpasOp>::OpConversionPattern;
241
242 LogicalResult
243 matchAndRewrite(xegpu::DpasOp op, OpAdaptor adaptor,
244 ConversionPatternRewriter &rewriter) const override {
245 // Check if the op has A, B and CD layouts attached.
246 auto layoutA = cast<xegpu::LayoutAttr>(op.getLayoutAAttr());
247 auto layoutB = cast<xegpu::LayoutAttr>(op.getLayoutBAttr());
248 auto layoutCd = cast<xegpu::LayoutAttr>(op.getLayoutCdAttr());
249 if (!layoutA || !layoutB || !layoutCd)
250 return failure();
251 auto laneResultTyOrFailure =
252 xegpu::getDistributedVectorType(op.getType(), layoutCd);
253 auto laneATypeOrFailure =
254 xegpu::getDistributedVectorType(op.getLhs().getType(), layoutA);
255 auto laneBTypeOrFailure =
256 xegpu::getDistributedVectorType(op.getRhs().getType(), layoutB);
257 auto expectedLaneResultTyOrFailure =
258 xegpu::getDistVecTypeBasedOnLaneLayout(layoutCd, op.getType());
259 if (failed(laneResultTyOrFailure) || failed(laneATypeOrFailure) ||
260 failed(laneBTypeOrFailure))
261 return rewriter.notifyMatchFailure(
262 op, "failed to calculate supported lane vector types for DpasOp "
263 "from layouts");
264 if (failed(expectedLaneResultTyOrFailure))
265 return rewriter.notifyMatchFailure(
266 op, "unable to compute expected lane vector type for DpasOp from "
267 "lane layout");
268
269 // Validate bit widths match uArch packed format requirements
270 const auto *uArch =
272 if (uArch) {
273 const auto *uArchInstruction =
274 dyn_cast<xegpu::uArch::SubgroupMatrixMultiplyAcc>(
275 uArch->getInstruction(
277 if (uArchInstruction) {
278 auto laneAType = laneATypeOrFailure.value();
279 auto laneBType = laneBTypeOrFailure.value();
280 // Calculate total packed bit width = element bit width * vector size
281 unsigned aPackedBitWidth =
282 laneAType.getElementTypeBitWidth() * laneAType.getNumElements();
283 unsigned bPackedBitWidth =
284 laneBType.getElementTypeBitWidth() * laneBType.getNumElements();
285 unsigned expectedABitSize = uArchInstruction->getPackedFormatBitSizeA();
286 unsigned expectedBBitSize = uArchInstruction->getPackedFormatBitSizeB();
287
288 if (aPackedBitWidth % expectedABitSize != 0)
289 return rewriter.notifyMatchFailure(
290 op,
291 "A operand packed bit width must be a multiple of uArch packed "
292 "format requirement");
293 if (bPackedBitWidth % expectedBBitSize != 0)
294 return rewriter.notifyMatchFailure(
295 op,
296 "B operand packed bit width must be a multiple of uArch packed "
297 "format requirement");
298 }
299 }
300
301 auto newOp = xegpu::DpasOp::create(
302 rewriter, op->getLoc(), laneResultTyOrFailure.value(),
303 castValueTo(rewriter, cast<TypedValue<VectorType>>(adaptor.getLhs()),
304 laneATypeOrFailure.value()),
305 castValueTo(rewriter, cast<TypedValue<VectorType>>(adaptor.getRhs()),
306 laneBTypeOrFailure.value()),
307 castValueTo(rewriter, cast<TypedValue<VectorType>>(adaptor.getAcc()),
308 laneResultTyOrFailure.value()),
309 /** layoutA**/ nullptr,
310 /** layoutB**/ nullptr, /** layoutCd**/ nullptr);
311 // Explicitly set the new types to enable correct type materializations.
312 rewriter.replaceOp(op, castValueTo(rewriter, newOp.getResult(),
313 expectedLaneResultTyOrFailure.value()));
314 return success();
315 }
316};
317
318/// Distributes elementwise ops to lane-level elementwise ops. This
319/// currently handles elementwise ops with single result only.
320struct SgToLaneElementWise : public ConversionPattern {
321 SgToLaneElementWise(TypeConverter &typeConverter, MLIRContext *ctx)
322 : ConversionPattern(MatchAnyOpTypeTag(), /*benefit=*/1, ctx) {}
323
324 LogicalResult
325 matchAndRewrite(Operation *op, ArrayRef<Value> operands,
326 ConversionPatternRewriter &rewriter) const override {
327 // Only match ops with elementwise trait and single result.
329 return failure();
330
331 auto resultType = dyn_cast<VectorType>(op->getResult(0).getType());
332 if (!resultType)
333 return rewriter.notifyMatchFailure(
334 op, "operation result is not a vector type");
335
336 xegpu::DistributeLayoutAttr layout =
337 xegpu::getTemporaryLayout(llvm::cast<OpResult>(op->getResult(0)));
338 if (!layout || !layout.isForSubgroup())
339 return rewriter.notifyMatchFailure(
340 op, "operation result does not have subgroup distribute layout");
341
342 auto laneShapeOrFailure =
343 xegpu::getDistVecTypeBasedOnLaneLayout(layout, resultType);
344
345 if (failed(laneShapeOrFailure))
346 return rewriter.notifyMatchFailure(
347 op, "unable to compute lane vector type from the layout");
348
349 VectorType newResultType = laneShapeOrFailure.value();
350 OperationState state(op->getLoc(), op->getName());
351 state.addOperands(operands);
352 state.addTypes(newResultType);
353 // Copy all attributes except for DistributeLayoutAttr.
354 for (auto attr : op->getDiscardableAttrDictionary().getValue()) {
355 if (!isa<xegpu::DistributeLayoutAttr>(attr.getValue()))
356 state.addAttribute(attr.getName(), attr.getValue());
357 }
359 Operation *newOp = rewriter.create(state);
360
361 rewriter.replaceOp(op, newOp->getResult(0));
362 return success();
363 }
364};
365
366/// Distributes a subgroup-level arith ConstantOp to lane-level arith
367/// ConstantOp.
368///
369/// Splat constants are rebuilt with the lane-local vector type. Non-splat
370/// constants are distributed by extracting each lane_data-sized block from
371/// the full constant and inserting it at the correct position in the
372/// distributed vector using insert_strided_slice.
373struct SgToLaneArithConstant : public OpConversionPattern<arith::ConstantOp> {
374 using OpConversionPattern<arith::ConstantOp>::OpConversionPattern;
375
376 LogicalResult
377 matchAndRewrite(arith::ConstantOp op, OpAdaptor adaptor,
378 ConversionPatternRewriter &rewriter) const override {
379 auto resultType = dyn_cast<VectorType>(op.getType());
380 if (!resultType)
381 return failure();
382
383 // Only handle dense vector constants.
384 auto denseAttr = dyn_cast<DenseElementsAttr>(op.getValue());
385 if (!denseAttr)
386 return rewriter.notifyMatchFailure(
387 op, "only dense vector constants are supported");
388
389 xegpu::DistributeLayoutAttr layout =
390 xegpu::getTemporaryLayout(llvm::cast<OpResult>(op.getResult()));
391 if (!layout || !layout.isForSubgroup())
392 return rewriter.notifyMatchFailure(
393 op, "operation result does not have subgroup distribute layout");
394
395 auto laneShapeOrFailure =
396 xegpu::getDistVecTypeBasedOnLaneLayout(layout, resultType);
397
398 if (failed(laneShapeOrFailure))
399 return rewriter.notifyMatchFailure(
400 op, "unable to compute lane vector type from the layout");
401
402 VectorType newResultType = laneShapeOrFailure.value();
403 Location loc = op.getLoc();
404
405 // Splat constants: every lane gets the same value, so just rebuild the
406 // splat with the distributed type.
407 if (denseAttr.isSplat()) {
408 auto scalarValue = denseAttr.getSplatValue<Attribute>();
409 auto newDenseAttr = DenseElementsAttr::get(newResultType, scalarValue);
410 auto newOp =
411 arith::ConstantOp::create(rewriter, loc, newResultType, newDenseAttr);
412 rewriter.replaceOp(op, newOp.getResult());
413 return success();
414 }
415
416 // Non-splat constants: each lane extracts the elements it owns from the
417 // full constant using the distributed coordinates from the layout.
418 auto fullConst =
419 arith::ConstantOp::create(rewriter, loc, resultType, denseAttr);
420
421 Value laneId = gpu::LaneIdOp::create(rewriter, loc, rewriter.getIndexType(),
422 /*upperBound=*/mlir::IntegerAttr());
423 auto maybeCoordsVec = layout.computeDistributedCoords(
424 rewriter, loc, laneId, resultType.getShape());
425 if (failed(maybeCoordsVec))
426 return rewriter.notifyMatchFailure(
427 op, "failed to compute distributed coordinates from layout");
428
429 SmallVector<SmallVector<Value>> coordsVec = maybeCoordsVec.value();
430 SmallVector<int64_t> laneData = layout.getEffectiveLaneDataAsInt();
431 ArrayRef<int64_t> distShape = newResultType.getShape();
432 int64_t rank = newResultType.getRank();
433
434 // Each lane owns one lane_data-sized block per distribution unit.
435 // computeDistributedCoords returns those block starts in row-major order
436 // over the block grid (distShape / laneData).
437 SmallVector<int64_t> blockGridShape(rank);
438 for (int64_t d = 0; d < rank; d++)
439 blockGridShape[d] = distShape[d] / laneData[d];
440 SmallVector<int64_t> blockGridStrides = computeStrides(blockGridShape);
441
442 auto blockType = VectorType::get(laneData, newResultType.getElementType());
443 SmallVector<int64_t> unitTile(rank, 1);
444 SmallVector<int64_t> strides(rank, 1);
445
446 Value result = arith::ConstantOp::create(
447 rewriter, loc, newResultType, rewriter.getZeroAttr(newResultType));
448
449 for (auto [blockIdx, blockStart] : llvm::enumerate(coordsVec)) {
450 // Gather the block's elements from the full constant. The block start is
451 // lane-dynamic, so extract element-by-element (row-major over lane_data)
452 // instead.
453 SmallVector<Value> blockElems;
454 for (SmallVector<int64_t> off :
455 StaticTileOffsetRange(laneData, unitTile)) {
457 for (int64_t d = 0; d < rank; d++)
458 pos[d] = getAsOpFoldResult(arith::AddIOp::create(
459 rewriter, loc, blockStart[d],
460 arith::ConstantIndexOp::create(rewriter, loc, off[d])));
461 blockElems.push_back(vector::ExtractOp::create(
462 rewriter, loc, fullConst.getResult(), pos));
463 }
464
465 // Rebuild the block keeping its lane_data shape, then place it with
466 // insert_strided_slice so the block keeps its orientation in the
467 // distributed vector (e.g. a [2, 1] block stays a vertical 2x1 slice).
468 Value block =
469 vector::FromElementsOp::create(rewriter, loc, blockType, blockElems);
470 SmallVector<int64_t> blockGridPos =
471 delinearize(blockIdx, blockGridStrides);
472 SmallVector<int64_t> offsets(rank);
473 for (int64_t d = 0; d < rank; d++)
474 offsets[d] = blockGridPos[d] * laneData[d];
475 result = vector::InsertStridedSliceOp::create(rewriter, loc, block,
476 result, offsets, strides);
477 }
478
479 rewriter.replaceOp(op, result);
480 return success();
481 }
482};
483
484/// Distributes a subgroup-level PrefetchNd op to lane-level PrefetchNd op.
485struct SgToLanePrefetchNd : public OpConversionPattern<xegpu::PrefetchNdOp> {
486 using OpConversionPattern<xegpu::PrefetchNdOp>::OpConversionPattern;
487
488 LogicalResult
489 matchAndRewrite(xegpu::PrefetchNdOp op, OpAdaptor adaptor,
490 ConversionPatternRewriter &rewriter) const override {
491 xegpu::DistributeLayoutAttr layout = op.getAnchorLayout();
492 // If no layout, nothing to do.
493 if (!layout)
494 return failure();
495
496 xegpu::PrefetchNdOp::create(rewriter, op.getLoc(), adaptor.getTensorDesc(),
497 op.getMixedOffsets(), op.getL1HintAttr(),
498 op.getL2HintAttr(), op.getL3HintAttr(),
499 /**layout**/ nullptr);
500 rewriter.eraseOp(op);
501 return success();
502 }
503};
504
505/// Distributes a subgroup-level LoadGather (xegpu.load) op to lane-level.
506///
507/// Example 1 (1D):
508/// layout = #xegpu.layout<lane_layout = [16], lane_data = [1]>
509/// %mask = producer_op : vector<16xi1>
510/// %offset = producer_op : vector<16xindex>
511/// %0 = xegpu.load %src[%offset], %mask : memref<256xf16>,
512/// vector<16xindex>, vector<16xi1> -> vector<16xf16>
513/// Distributed to:
514/// %mask = producer_op : vector<1xi1>
515/// %offset = producer_op : vector<1xindex>
516/// %0 = xegpu.load %src[%offset], %mask : memref<256xf16>,
517/// vector<1xindex>, vector<1xi1> -> vector<1xf16>
518///
519/// Example 2 (3D with leading unit dims):
520/// layout = #xegpu.layout<lane_layout = [1, 1, 16], lane_data = [1, 1, 1]>
521/// %mask = producer_op : vector<1x1x16xi1>
522/// %offset = producer_op : vector<1x1x16xindex>
523/// %0 = xegpu.load %src[%offset], %mask : memref<256xf16>,
524/// vector<1x1x16xindex>, vector<1x1x16xi1> -> vector<1x1x16xf16>
525/// Distributed to:
526/// %mask = producer_op : vector<1x1x1xi1>
527/// %offset = producer_op : vector<1x1x1xindex>
528/// %0 = xegpu.load %src[%offset], %mask : memref<256xf16>,
529/// vector<1xindex>, vector<1xi1> -> vector<1xf16>
530struct SgToLaneLoadGather : public OpConversionPattern<xegpu::LoadGatherOp> {
531 using OpConversionPattern<xegpu::LoadGatherOp>::OpConversionPattern;
532
533 LogicalResult
534 matchAndRewrite(xegpu::LoadGatherOp op, OpAdaptor adaptor,
535 ConversionPatternRewriter &rewriter) const override {
536 xegpu::DistributeLayoutAttr layout = op.getAnchorLayout();
537 if (!layout)
538 return failure();
539
540 VectorType origResultTy = op.getValueType();
541 if (!origResultTy)
542 return failure();
543
544 // Check that leading dimensions are unit.
545 int effectiveVecRank = 1;
546 ArrayRef<int64_t> shape = origResultTy.getShape();
547 if (llvm::any_of(
548 shape.take_front(origResultTy.getRank() - effectiveVecRank),
549 [](int64_t d) { return d != 1; }))
550 return rewriter.notifyMatchFailure(
551 op, "Only unit dimensions allowed for the leading "
552 "dimensions of the load vector!");
553
554 auto distResultTyOrFailure =
555 xegpu::getDistVecTypeBasedOnLaneLayout(layout, origResultTy);
556 if (failed(distResultTyOrFailure))
557 return rewriter.notifyMatchFailure(
558 op, "unable to compute expected lane vector type from lane layout");
559
560 VectorType distResultTy = distResultTyOrFailure.value();
561 VectorType distResultTy1D = VectorType::get({distResultTy.getNumElements()},
562 distResultTy.getElementType());
563
564 // Flatten offsets and mask to 1D to match the 1D result type.
565 Value distOffsets = adaptor.getOffsets();
566 auto distOffsetsTy = cast<VectorType>(distOffsets.getType());
567 VectorType offsetsTy1D = VectorType::get({distOffsetsTy.getNumElements()},
568 distOffsetsTy.getElementType());
569 distOffsets = castValueTo(
570 rewriter, cast<TypedValue<VectorType>>(distOffsets), offsetsTy1D);
571
572 Value distMask = adaptor.getMask();
573 auto distMaskTy = cast<VectorType>(distMask.getType());
574 VectorType maskTy1D = VectorType::get({distMaskTy.getNumElements()},
575 distMaskTy.getElementType());
576 distMask =
577 castValueTo(rewriter, cast<TypedValue<VectorType>>(distMask), maskTy1D);
578
579 Value distSource = adaptor.getSource();
580 auto newOp = xegpu::LoadGatherOp::create(
581 rewriter, op.getLoc(), distResultTy1D, distSource, distOffsets,
582 distMask, op.getL1HintAttr(), op.getL2HintAttr(), op.getL3HintAttr(),
583 /*layout=*/nullptr, /*contiguity=*/nullptr);
584
585 Value result = newOp->getResult(0);
586 if (distResultTy1D != distResultTy)
587 result = castValueTo(rewriter, cast<TypedValue<VectorType>>(result),
588 distResultTy);
589 rewriter.replaceOp(op, result);
590 return success();
591 }
592};
593
594/// This pattern distributes a subgroup-level vector.reduction op to
595/// lane-level. This require shuffling the data across the lanes (using
596/// gpu::ShuffleOp) and reducing in stages until all lanes have the final
597/// result.
598struct SgToLaneVectorReduction
599 : public OpConversionPattern<vector::ReductionOp> {
600 using OpConversionPattern<vector::ReductionOp>::OpConversionPattern;
601
602 LogicalResult
603 matchAndRewrite(vector::ReductionOp op, OpAdaptor adaptor,
604 ConversionPatternRewriter &rewriter) const override {
605 auto layout = xegpu::getDistributeLayoutAttr(op.getVector());
606
607 // If no layout, nothing to do.
608 if (!layout || !layout.isForSubgroup())
609 return failure();
610
611 VectorType srcVecType = op.getSourceVectorType();
612 // Only rank 1 vectors supported.
613 if (srcVecType.getRank() != 1)
614 return rewriter.notifyMatchFailure(
615 op, "Only rank 1 reductions can be distributed.");
616 // Lane layout must have the same rank as the vector.
617 if (layout.getRank() != srcVecType.getRank())
618 return rewriter.notifyMatchFailure(
619 op, "Layout rank does not match vector rank.");
620
621 // Get the subgroup size from the layout.
622 int64_t sgSize = layout.getEffectiveLaneLayoutAsInt()[0];
623 const auto *uArch =
625 if (!uArch)
626 return rewriter.notifyMatchFailure(
627 op, "xegpu::ReductionOp require target attribute attached to "
628 "determine subgroup size");
629
630 // Only subgroup-sized vectors supported.
631 if (sgSize != uArch->getSubgroupSize() ||
632 srcVecType.getShape()[0] % sgSize != 0)
633 return rewriter.notifyMatchFailure(op,
634 "Invalid layout or reduction vector "
635 "dimension must match subgroup size.");
636
637 if (!op.getType().isIntOrFloat())
638 return rewriter.notifyMatchFailure(
639 op, "Reduction distribution currently only supports floats and "
640 "integer types.");
641
642 // Get the distributed vector (per lane portion).
643 Value laneValVec = adaptor.getVector();
644
645 // Distribute and reduce across lanes in the subgroup.
646 Value fullReduce = xegpu::subgroupReduction(
647 op.getLoc(), rewriter, laneValVec, op.getKind(), sgSize);
648
649 // If there's an accumulator, combine it with the reduced value.
650 if (adaptor.getAcc())
651 fullReduce = vector::makeArithReduction(
652 rewriter, op.getLoc(), op.getKind(), fullReduce, adaptor.getAcc());
653
654 rewriter.replaceOp(op, fullReduce);
655 return success();
656 }
657};
658
659/// This pattern distributes a subgroup-level vector.multi_reduction op to
660/// lane-level only if the reduction is lane-local. This means that
661/// reduction dimension is not distributed to lanes and each lane does its own
662/// local reduction.
663struct SgToLaneMultiDimReduction
664 : public OpConversionPattern<vector::MultiDimReductionOp> {
665 using OpConversionPattern<vector::MultiDimReductionOp>::OpConversionPattern;
666
667 LogicalResult
668 matchAndRewrite(vector::MultiDimReductionOp op, OpAdaptor adaptor,
669 ConversionPatternRewriter &rewriter) const override {
671 ArrayRef<int64_t> reductionDims = op.getReductionDims();
672 assert(reductionDims.size() == 1 &&
673 "Expecting single reduction dimension for subgroup multi "
674 "reduction op");
675 // For rank > 2, ensure leading dimensions are unit.
676 VectorType sourceType = op.getSourceVectorType();
677 int64_t rank = sourceType.getRank();
678 if (rank > 2) {
679 ArrayRef<int64_t> shape = sourceType.getShape();
680 if (llvm::any_of(shape.take_front(rank - 2),
681 [](int64_t d) { return d != 1; }))
682 return rewriter.notifyMatchFailure(
683 op, "only unit leading dimensions are supported for "
684 "multi_reduction with rank > 2");
685 }
686 // Handle scalar result: full reduction of a distributed vector to a
687 // scalar. First do a local vector reduction, then cross-lane shuffles.
688 if (op.getType().isIntOrFloat()) {
689 auto reductionDim = reductionDims[0];
690 VectorType origSourceType = op.getSourceVectorType();
691 int64_t reductionDimSize = origSourceType.getShape()[reductionDim];
692 // Local reduction to scalar, then cross-lane butterfly shuffles.
693 result =
694 xegpu::subgroupReduction(op.getLoc(), rewriter, adaptor.getSource(),
695 op.getKind(), reductionDimSize);
696 // Combine with accumulator if present.
697 if (adaptor.getAcc())
698 result = vector::makeArithReduction(rewriter, op.getLoc(), op.getKind(),
699 result, adaptor.getAcc());
700 } else if (isReductionLaneLocal(op)) {
701 // For lane-local reduction, lower to a sequence of vector.reduction ops
702 // over 1D slices extracted from the distributed source vector. This is
703 // required so we dont have 2D source vectors at xegpu-linearize.
704 auto reductionDim = reductionDims[0];
706 cast<TypedValue<VectorType>>(adaptor.getSource()),
707 cast<TypedValue<VectorType>>(adaptor.getAcc()), op.getKind(),
708 reductionDim, op.getLoc(), rewriter);
709 } else {
710 auto reductionDim = reductionDims[0];
711 VectorType sourceType = op.getSourceVectorType();
712 int64_t reductionDimSize = sourceType.getShape()[reductionDim];
714 cast<TypedValue<VectorType>>(adaptor.getSource()),
715 cast<TypedValue<VectorType>>(adaptor.getAcc()), op.getKind(),
716 reductionDim, reductionDimSize, op.getLoc(), rewriter);
717 }
718 rewriter.replaceOp(op, result);
719 return success();
720 }
721};
722
723/// Helper to compute distributed coordinates for matrix ops.
724/// When not using subgroup_block_io, each lane computes its own
725/// coordinates based on the layout and lane ID.
726static SmallVector<Value> computeDistributedCoordsForMatrixOp(
727 ConversionPatternRewriter &rewriter, Location loc,
728 xegpu::DistributeLayoutAttr layout, ArrayRef<int64_t> payloadShape,
729 ValueRange origOffsets) {
730 Value laneId = gpu::LaneIdOp::create(rewriter, loc, rewriter.getIndexType(),
731 /*upperBound=*/mlir::IntegerAttr());
732 auto maybeCoords =
733 layout.computeDistributedCoords(rewriter, loc, laneId, payloadShape);
734 if (failed(maybeCoords))
735 return {};
736 assert(maybeCoords.value().size() == 1 &&
737 "Expected one set of distributed offsets");
739 rewriter, loc, getAsOpFoldResult(maybeCoords.value()[0]),
740 getAsOpFoldResult(origOffsets));
741 return llvm::map_to_vector(ofrVec, llvm::CastTo<Value>);
742}
743
744/// This pattern distributes a subgroup-level LoadMatrix op to lane-level.
745struct SgToLaneLoadMatrix : public OpConversionPattern<xegpu::LoadMatrixOp> {
746 using OpConversionPattern<xegpu::LoadMatrixOp>::OpConversionPattern;
747
748 LogicalResult
749 matchAndRewrite(xegpu::LoadMatrixOp op, OpAdaptor adaptor,
750 ConversionPatternRewriter &rewriter) const override {
751 auto layout = op.getLayoutAttr();
752 // If no layout, nothing to do.
753 if (!layout)
754 return failure();
755
756 VectorType sgPayloadTy = dyn_cast<VectorType>(op.getResult().getType());
757 if (!sgPayloadTy)
758 return rewriter.notifyMatchFailure(
759 op, "the matrix op payload must be a vector type");
760
761 auto loc = op.getLoc();
762 auto offsets = op.getMixedOffsets();
763 if (offsets.empty())
764 return rewriter.notifyMatchFailure(op, "the load op must have offsets");
765
766 FailureOr<VectorType> distPayloadTyOrFailure =
767 getDistVecTypeBasedOnLaneLayout(layout, sgPayloadTy);
768 if (failed(distPayloadTyOrFailure))
769 return rewriter.notifyMatchFailure(
770 op, "Failed to distribute matrix op payload based on layout.");
771
772 SmallVector<Value> offsetsAsValues =
773 vector::getAsValues(rewriter, loc, offsets);
774
775 SmallVector<Value> newCoords = offsetsAsValues;
776 if (!op.getSubgroupBlockIoAttr()) {
777 newCoords = computeDistributedCoordsForMatrixOp(
778 rewriter, loc, layout, sgPayloadTy.getShape(), offsetsAsValues);
779 if (newCoords.empty())
780 return rewriter.notifyMatchFailure(
781 op, "Failed to compute distributed coordinates.");
782 }
783
784 SmallVector<int64_t> newConstOffsets(op.getConstOffsets().size(),
785 ShapedType::kDynamic);
786 DenseI64ArrayAttr newConstOffsetsAttr =
787 rewriter.getDenseI64ArrayAttr(newConstOffsets);
788
789 auto newOp = xegpu::LoadMatrixOp::create(
790 rewriter, loc, *distPayloadTyOrFailure, adaptor.getMemDesc(),
791 ValueRange(newCoords), newConstOffsetsAttr, op.getSubgroupBlockIoAttr(),
792 xegpu::DistributeLayoutAttr{});
793 rewriter.replaceOp(op, newOp.getResult());
794 return success();
795 }
796};
797
798/// Distributes a subgroup-level vector.transpose op to lane-level.
799struct SgToLaneVectorTranspose
800 : public OpConversionPattern<vector::TransposeOp> {
801 using OpConversionPattern<vector::TransposeOp>::OpConversionPattern;
802
803 LogicalResult
804 matchAndRewrite(vector::TransposeOp op, OpAdaptor adaptor,
805 ConversionPatternRewriter &rewriter) const override {
806 xegpu::DistributeLayoutAttr sourceLayout =
807 xegpu::getTemporaryLayout(op->getOpOperand(0));
808 xegpu::DistributeLayoutAttr resultLayout =
809 xegpu::getTemporaryLayout(op->getOpResult(0));
810 if (!sourceLayout || !resultLayout)
811 return rewriter.notifyMatchFailure(
812 op, "the source or result vector of the transpose op lacks layout "
813 "attribute");
814 ArrayRef<int64_t> perm = op.getPermutation();
815 // Result layout must be a transpose of source layout.
816 if (!resultLayout.isTransposeOf(sourceLayout, perm,
818 return rewriter.notifyMatchFailure(
819 op, "the source or result vector layouts must be transposes of "
820 "each other");
821 FailureOr<VectorType> distributedResultTypeOrFailure =
822 getDistVecTypeBasedOnLaneLayout(resultLayout, op.getResultVectorType());
823 if (failed(distributedResultTypeOrFailure))
824 return rewriter.notifyMatchFailure(
825 op, "Failed to distribute the result vector type in "
826 "vector::Transpose op");
827 auto newOp = vector::TransposeOp::create(rewriter, op.getLoc(),
828 adaptor.getVector(), perm);
829 rewriter.replaceOp(op, castValueTo(rewriter, newOp.getResult(),
830 distributedResultTypeOrFailure.value()));
831 return success();
832 }
833};
834
835/// Distributes a subgroup-level vector.bitcast op to lane-level.
836/// Bitcast only impacts the innermost dimension of the source/result vectors.
837struct SgToLaneVectorBitcast : public OpConversionPattern<vector::BitCastOp> {
838 using OpConversionPattern<vector::BitCastOp>::OpConversionPattern;
839
840 LogicalResult
841 matchAndRewrite(vector::BitCastOp op, OpAdaptor adaptor,
842 ConversionPatternRewriter &rewriter) const override {
843 xegpu::DistributeLayoutAttr resultLayout =
844 xegpu::getTemporaryLayout(op->getOpResult(0));
845 if (!resultLayout)
846 return rewriter.notifyMatchFailure(
847 op, "result vector of the bitcast op lacks layout attribute");
848 FailureOr<VectorType> distributedResultTypeOrFailure =
849 getDistVecTypeBasedOnLaneLayout(resultLayout, op.getResultVectorType());
850 if (failed(distributedResultTypeOrFailure))
851 return rewriter.notifyMatchFailure(
852 op, "Failed to distribute the result vector type in "
853 "vector::BitCast op");
854 auto newOp = vector::BitCastOp::create(
855 rewriter, op.getLoc(), distributedResultTypeOrFailure.value(),
856 adaptor.getSource());
857 rewriter.replaceOp(op, newOp.getResult());
858 return success();
859 }
860};
861
862/// Distributes a subgroup-level vector.create_mask or vector.constant_mask op
863/// to lane-level.
864/// The pattern constructs a mask based on the following bounds check:
865/// ```
866/// for d in [0, ..., maskRank):
867/// mask &= (staticOffset[d] < (bound[d] - base[d]))
868/// ```
869/// where
870/// - `base` is the coordinate vector of the first distributed unit.
871/// - `staticOffset` is the offsets vector per element.
872/// - `bound` is the original mask bound for the corresponding dimension.
873/// The mask vector contains *all* elements (i.e., non-unit `lane_data` is
874/// flattened). For example,
875/// ```
876/// %mask = vector.create_mask %bound_0, %bound_1 {
877/// lane_layout = [1, 16], lane_data = [2, 1]
878/// }: vector<8x32xi1>
879/// ```
880/// Has 8 dist units (of shape [2, 1]) with offsets for lane 0:
881/// {
882/// [0, 0], [0, 16],
883/// [2, 0], [2, 16],
884/// [4, 0], [4, 16],
885/// [6, 0], [6, 16]
886/// }
887/// We check the distance to the bound from the first dist. unit:
888/// %base_0 = genCoords(gpu.lane_id, layout)[0][0]
889/// %base_1 = genCoords(gpu.lane_id, layout)[0][1]
890/// %baseDistanceToBound_0 = %bound_0 - %base_0
891/// %baseDistanceToBound_1 = %bound_1 - %base_1
892/// The pattern flattens the dimension d coordinates of each element:
893/// %staticOffset_0 = [0, 0, 1, 1, 2, 2, 3, 3, 4, 4, 5, 5, 6, 6, 7, 7]
894/// %staticOffset_1 = [0, 16, 0, 16, 0, 16, 0, 16, 0, 16, 0, 16, 0, 16, 0, 16]
895/// and compares each coordinate against the corresponding mask bound dim:
896/// %mask_0 = arith.cmpi slt, %staticOffset_0, bcast(%baseDistanceToBound_0)
897/// %mask_1 = arith.cmpi slt, %staticOffset_1, bcast(%baseDistanceToBound_1)
898/// %mask = shape_cast(%mask_0 & %mask_1) : vector<16xi1> to vector<8x2xi1>
899///
900template <typename OpType,
901 typename = std::enable_if_t<llvm::is_one_of<
902 OpType, vector::CreateMaskOp, vector::ConstantMaskOp>::value>>
903struct SgToLaneCreateMask : public OpConversionPattern<OpType> {
904 using OpConversionPattern<OpType>::OpConversionPattern;
905
906 LogicalResult
907 matchAndRewrite(OpType op, typename OpType::Adaptor adaptor,
908 ConversionPatternRewriter &rewriter) const override {
909 xegpu::DistributeLayoutAttr layout =
910 xegpu::getTemporaryLayout(op->getOpResult(0));
911 if (!layout || !layout.isForSubgroup())
912 return rewriter.notifyMatchFailure(
913 op, "operation result does not have subgroup distribute layout");
914
915 VectorType origType = op.getType();
916 FailureOr<VectorType> distTypeOrFailure =
917 getDistVecTypeBasedOnLaneLayout(layout, origType);
918 if (failed(distTypeOrFailure))
919 return rewriter.notifyMatchFailure(
920 op, "unable to compute lane vector type from the layout");
921
922 VectorType distType = distTypeOrFailure.value();
923 Location loc = op.getLoc();
924
925 // Materialize the original mask bounds as Values.
926 SmallVector<Value> origBounds;
927 if constexpr (std::is_same_v<OpType, vector::CreateMaskOp>) {
928 origBounds.append(op.getOperands().begin(), op.getOperands().end());
929 } else {
930 auto dimSizes = op.getMaskDimSizesAttr().asArrayRef();
931 for (auto dimSize : dimSizes)
932 origBounds.push_back(
933 arith::ConstantIndexOp::create(rewriter, loc, dimSize).getResult());
934 }
935
936 ArrayRef<int64_t> origShape = origType.getShape();
937
938 // Use computeDistributedCoords to get the coordinates each WI owns.
939 Value laneId = gpu::LaneIdOp::create(rewriter, loc, rewriter.getIndexType(),
940 /*upperBound=*/mlir::IntegerAttr());
941 auto maybeCoordsVec =
942 layout.computeDistributedCoords(rewriter, loc, laneId, origShape);
943 if (failed(maybeCoordsVec))
944 return rewriter.notifyMatchFailure(
945 op, "failed to compute distributed coordinates from layout");
946
947 SmallVector<SmallVector<Value>> laneDataCoords = maybeCoordsVec.value();
948 SmallVector<int64_t> laneData = layout.getEffectiveLaneDataAsInt();
949 ArrayRef<int64_t> distShape = distType.getShape();
950 int64_t rank = distType.getRank();
951 int64_t numElements = distType.getNumElements();
952
953 if (static_cast<int64_t>(laneData.size()) != rank ||
954 !computeShapeRatio(distShape, laneData))
955 return rewriter.notifyMatchFailure(
956 op, "lane_data does not tile the distributed vector");
957
958 SmallVector<int64_t> distStrides = computeStrides(distShape);
959 assert(static_cast<int64_t>(laneDataCoords.size()) *
960 computeProduct(laneData) ==
961 numElements &&
962 "number of coordinate sets must match number of lane_data blocks");
963
964 // Static offsets to be applied to a dynamic lane's dist units coordinates.
965 SmallVector<SmallVector<int64_t>> staticLaneDataOffsetsOrig =
966 layout.computeStaticDistributedCoords(/*linearId=*/0, origShape);
967 if (staticLaneDataOffsetsOrig.empty() ||
968 staticLaneDataOffsetsOrig.size() != laneDataCoords.size())
969 return rewriter.notifyMatchFailure(
970 op, "static and dynamic coordinates disagree on the distribution "
971 "unit count");
972
973 SmallVector<int64_t> unitTile(rank, 1);
974 int64_t linearLaneDataIdx = 0;
975 // Layout and shape information is static, compute static offset of lane's
976 // elements in the source dimensions. Each dim is a flat vec of element
977 // offsets. We store dim-major to have an easy materialization as one const
978 // vector.
979 SmallVector<SmallVector<int64_t>> staticElemOffset(
980 rank, SmallVector<int64_t>(numElements));
981 // For each lane_data block in the lane's dist shape
982 for (SmallVector<int64_t> laneDataOffsetDist :
983 StaticTileOffsetRange(distShape, laneData)) {
984 ArrayRef<int64_t> staticLaneDataOffsetOrig =
985 staticLaneDataOffsetsOrig[linearLaneDataIdx++];
986 // For each element in the lane_data block
987 for (SmallVector<int64_t> elemOffsetInLaneData :
988 StaticTileOffsetRange(laneData, unitTile)) {
989 // For each dim of indexing space
990 SmallVector<int64_t> elementOffsetDist(rank);
991 for (int64_t d = 0; d < rank; d++)
992 elementOffsetDist[d] =
993 laneDataOffsetDist[d] + elemOffsetInLaneData[d];
994 int64_t elemLinearizedIdxDist =
995 linearize(elementOffsetDist, distStrides);
996 for (int64_t d = 0; d < rank; d++)
997 staticElemOffset[d][elemLinearizedIdxDist] =
998 staticLaneDataOffsetOrig[d] + elemOffsetInLaneData[d];
999 }
1000 }
1001
1002 // Check, whether ALL elements of the distributed mask are within the valid
1003 // extent of EACH source dimension.
1004 // Expressed as dyn_offset[d] + static_offset[d] < dyn_mask_bound[d],
1005 // or equivalently,
1006 // static_offset[d] < dyn_mask_bound[d] - dyn_offset[d]
1007 auto flatIndexType = VectorType::get(numElements, rewriter.getIndexType());
1008 Value inBounds;
1009 for (int64_t d = 0; d < rank; d++) {
1010 std::optional<int64_t> constBound = getConstantIntValue(origBounds[d]);
1011 if (constBound && *constBound >= origShape[d])
1012 continue;
1013
1014 Value materializedStaticOffset = arith::ConstantOp::create(
1015 rewriter, loc, rewriter.getIndexVectorAttr(staticElemOffset[d]));
1016 // Consider only the first unit, all others are compile-time multiple
1017 // offsets of the first one and are already encoded in staticOffsets.
1018 Value validDimExtent = arith::SubIOp::create(rewriter, loc, origBounds[d],
1019 laneDataCoords[0][d]);
1020 // Dim extent applies to all elements.
1021 Value validDimExtentPerElement = vector::BroadcastOp::create(
1022 rewriter, loc, flatIndexType, validDimExtent);
1023 Value elementMaskInDim = arith::CmpIOp::create(
1024 rewriter, loc, arith::CmpIPredicate::slt, materializedStaticOffset,
1025 validDimExtentPerElement);
1026 // An element is valid iff it is within the valid extent of ALL
1027 // dimensions.
1028 inBounds = inBounds ? arith::AndIOp::create(rewriter, loc, inBounds,
1029 elementMaskInDim)
1030 .getResult()
1031 : elementMaskInDim;
1032 }
1033 // Every dim's bound saturates the mask extent, so no comparison was
1034 // emitted at all: every element of every lane is in bounds.
1035 if (!inBounds) {
1036 rewriter.replaceOp(op, arith::ConstantOp::create(
1037 rewriter, loc, distType,
1038 DenseElementsAttr::get(distType, true)));
1039 return success();
1040 }
1041 auto resMask =
1042 rewriter.createOrFold<vector::ShapeCastOp>(loc, distType, inBounds);
1043 rewriter.replaceOp(op, resMask);
1044 return success();
1045 }
1046};
1047
1048/// This pattern distributes a subgroup-level StoreMatrix op to lane-level.
1049struct SgToLaneStoreMatrix : public OpConversionPattern<xegpu::StoreMatrixOp> {
1050 using OpConversionPattern<xegpu::StoreMatrixOp>::OpConversionPattern;
1051
1052 LogicalResult
1053 matchAndRewrite(xegpu::StoreMatrixOp op, OpAdaptor adaptor,
1054 ConversionPatternRewriter &rewriter) const override {
1055 auto layout = op.getLayoutAttr();
1056 // If no layout, nothing to do.
1057 if (!layout)
1058 return failure();
1059
1060 VectorType sgPayloadTy = dyn_cast<VectorType>(op.getData().getType());
1061 if (!sgPayloadTy)
1062 return rewriter.notifyMatchFailure(
1063 op, "the matrix op payload must be a vector type");
1064
1065 auto loc = op.getLoc();
1066 auto offsets = op.getMixedOffsets();
1067 if (offsets.empty())
1068 return rewriter.notifyMatchFailure(op, "the store op must have offsets");
1069
1070 FailureOr<VectorType> distPayloadTyOrFailure =
1071 getDistVecTypeBasedOnLaneLayout(layout, sgPayloadTy);
1072 if (failed(distPayloadTyOrFailure))
1073 return rewriter.notifyMatchFailure(
1074 op, "Failed to distribute matrix op payload based on layout.");
1075
1076 SmallVector<Value> offsetsAsValues =
1077 vector::getAsValues(rewriter, loc, offsets);
1078
1079 SmallVector<Value> newCoords = offsetsAsValues;
1080 if (!op.getSubgroupBlockIoAttr()) {
1081 newCoords = computeDistributedCoordsForMatrixOp(
1082 rewriter, loc, layout, sgPayloadTy.getShape(), offsetsAsValues);
1083 if (newCoords.empty())
1084 return rewriter.notifyMatchFailure(
1085 op, "Failed to compute distributed coordinates.");
1086 }
1087
1088 SmallVector<int64_t> newConstOffsets(op.getConstOffsets().size(),
1089 ShapedType::kDynamic);
1090 DenseI64ArrayAttr newConstOffsetsAttr =
1091 rewriter.getDenseI64ArrayAttr(newConstOffsets);
1092
1093 xegpu::StoreMatrixOp::create(
1094 rewriter, loc, TypeRange{},
1095 castValueTo(rewriter, cast<TypedValue<VectorType>>(adaptor.getData()),
1096 distPayloadTyOrFailure.value()),
1097 adaptor.getMemDesc(), ValueRange(newCoords), newConstOffsetsAttr,
1098 op.getSubgroupBlockIoAttr(), xegpu::DistributeLayoutAttr{});
1099 rewriter.eraseOp(op);
1100 return success();
1101 }
1102};
1103
1104/// Distributes a subgroup-level StoreScatter (xegpu.store) op to
1105/// lane-level.
1106///
1107/// Example 1 (1D):
1108/// layout = #xegpu.layout<lane_layout = [16], lane_data = [1]>
1109/// %mask = producer_op : vector<16xi1>
1110/// %offset = producer_op : vector<16xindex>
1111/// xegpu.store %payload, %src[%offset], %mask : vector<16xf16>,
1112/// memref<256xf16>, vector<16xindex>, vector<16xi1>
1113/// Distributed to:
1114/// %mask = producer_op : vector<1xi1>
1115/// %offset = producer_op : vector<1xindex>
1116/// xegpu.store %payload, %src[%offset], %mask : vector<1xf16>,
1117/// memref<256xf16>, vector<1xindex>, vector<1xi1>
1118///
1119/// Example 2 (3D with leading unit dims):
1120/// layout = #xegpu.layout<lane_layout = [1, 1, 16], lane_data = [1, 1, 1]>
1121/// %mask = producer_op : vector<1x1x16xi1>
1122/// %offset = producer_op : vector<1x1x16xindex>
1123/// xegpu.store %payload, %src[%offset], %mask : vector<1x1x16xf16>,
1124/// memref<256xf16>, vector<1x1x16xindex>, vector<1x1x16xi1>
1125/// Distributed to:
1126/// %mask = producer_op : vector<1x1x1xi1>
1127/// %offset = producer_op : vector<1x1x1xindex>
1128/// xegpu.store %payload, %src[%offset], %mask : vector<1xf16>,
1129/// memref<256xf16>, vector<1xindex>, vector<1xi1>
1130struct SgToLaneStoreScatter
1131 : public OpConversionPattern<xegpu::StoreScatterOp> {
1132 using OpConversionPattern<xegpu::StoreScatterOp>::OpConversionPattern;
1133
1134 LogicalResult
1135 matchAndRewrite(xegpu::StoreScatterOp op, OpAdaptor adaptor,
1136 ConversionPatternRewriter &rewriter) const override {
1137 xegpu::DistributeLayoutAttr layout = op.getAnchorLayout();
1138 if (!layout)
1139 return failure();
1140
1141 VectorType origValueTy = op.getValueType();
1142 if (!origValueTy)
1143 return failure();
1144
1145 // Check that all leading dimensions are unit dimensions.
1146 int effectiveVecRank = 1;
1147 ArrayRef<int64_t> shape = origValueTy.getShape();
1148 if (llvm::any_of(shape.take_front(origValueTy.getRank() - effectiveVecRank),
1149 [](int64_t d) { return d != 1; }))
1150 return rewriter.notifyMatchFailure(
1151 op, "Only unit dimensions allowed for the leading "
1152 "dimensions of the store vector!");
1153
1154 auto distValueTyOrFailure =
1155 xegpu::getDistVecTypeBasedOnLaneLayout(layout, origValueTy);
1156 if (failed(distValueTyOrFailure))
1157 return rewriter.notifyMatchFailure(
1158 op, "unable to compute expected lane vector type from lane layout");
1159
1160 VectorType distValueTy = distValueTyOrFailure.value();
1161 VectorType distValueTy1D = VectorType::get({distValueTy.getNumElements()},
1162 distValueTy.getElementType());
1163
1164 Value distValue = adaptor.getValue();
1165 if (distValue.getType() != distValueTy1D)
1166 distValue = castValueTo(rewriter, cast<TypedValue<VectorType>>(distValue),
1167 distValueTy1D);
1168
1169 // Flatten offsets and mask to 1D to match the 1D value type.
1170 Value distOffsets = adaptor.getOffsets();
1171 auto distOffsetsTy = cast<VectorType>(distOffsets.getType());
1172 VectorType offsetsTy1D = VectorType::get({distOffsetsTy.getNumElements()},
1173 distOffsetsTy.getElementType());
1174 distOffsets = castValueTo(
1175 rewriter, cast<TypedValue<VectorType>>(distOffsets), offsetsTy1D);
1176
1177 Value distMask = adaptor.getMask();
1178 auto distMaskTy = cast<VectorType>(distMask.getType());
1179 VectorType maskTy1D = VectorType::get({distMaskTy.getNumElements()},
1180 distMaskTy.getElementType());
1181 distMask =
1182 castValueTo(rewriter, cast<TypedValue<VectorType>>(distMask), maskTy1D);
1183
1184 Value distDest = adaptor.getDest();
1185 xegpu::StoreScatterOp::create(rewriter, op.getLoc(), distValue, distDest,
1186 distOffsets, distMask, op.getL1HintAttr(),
1187 op.getL2HintAttr(), op.getL3HintAttr(),
1188 /*layout=*/nullptr,
1189 /*contiguity=*/nullptr);
1190 rewriter.eraseOp(op);
1191 return success();
1192 }
1193};
1194
1195/// Distribute a vector::StepOp to lane-level.
1196/// The layout must have exactly 1 effective lane dimension.
1197/// We completely resolve the vector::StepOp by computing the lane_data-sized
1198/// subranges.
1199struct SgToLaneVectorStep : public OpConversionPattern<vector::StepOp> {
1200 using OpConversionPattern<vector::StepOp>::OpConversionPattern;
1201
1202 LogicalResult
1203 matchAndRewrite(vector::StepOp op, OpAdaptor adaptor,
1204 ConversionPatternRewriter &rewriter) const override {
1205 xegpu::DistributeLayoutAttr resultLayout =
1206 xegpu::getTemporaryLayout(op->getResult(0));
1207 if (!resultLayout || !resultLayout.isForSubgroup())
1208 return rewriter.notifyMatchFailure(
1209 op, "the result vector of the step op lacks subgroup layout");
1210
1211 auto loc = op.getLoc();
1212 auto stepResultVecTy = op.getResult().getType();
1213 auto laneShapeOrFailure =
1214 xegpu::getDistVecTypeBasedOnLaneLayout(resultLayout, stepResultVecTy);
1215 if (failed(laneShapeOrFailure))
1216 return rewriter.notifyMatchFailure(
1217 op, "unable to compute lane vector type from the layout");
1218 VectorType newVecTy = laneShapeOrFailure.value();
1219
1220 Value laneId = gpu::LaneIdOp::create(rewriter, loc, rewriter.getIndexType(),
1221 /*upperBound=*/mlir::IntegerAttr());
1222 auto laneDataBlockCoords = resultLayout.computeDistributedCoords(
1223 rewriter, loc, laneId, stepResultVecTy.getShape());
1224 if (failed(laneDataBlockCoords))
1225 return rewriter.notifyMatchFailure(
1226 op, "failed to compute lane data block coordinates");
1227
1228 auto laneDataBlockCoordsVec = laneDataBlockCoords.value();
1229 auto laneDataBlockLength = resultLayout.getEffectiveLaneDataAsInt()[0];
1230 assert(static_cast<int64_t>(laneDataBlockCoordsVec.size()) ==
1231 newVecTy.getNumElements() / laneDataBlockLength);
1232 SmallVector<Value> stepVals;
1233 // For each lane_data block, reconstruct its sub-range
1234 // from the range of SG-level vector.step.Example: vector.step
1235 // {slice<layout<lane_layout=[2,4,2], lane_data=[1,2,1]>, dims=[0,2]>} :
1236 // vector<16xindex>
1237 // Each logical lane holds 4 elements as 2 blocks of 2 elements each.
1238 // The blocks are round-robin distributed, so logical lane id 0
1239 // holds values [0,1, 8,9].
1240 for (auto &laneDataBlockCoords : laneDataBlockCoordsVec) {
1241 auto laneDataBlockStartCoord = laneDataBlockCoords[0];
1242 stepVals.push_back(laneDataBlockStartCoord);
1243 for (int i = 1; i < laneDataBlockLength; ++i) {
1244 auto offset = arith::ConstantIndexOp::create(rewriter, loc, i);
1245 stepVals.push_back(arith::AddIOp::create(
1246 rewriter, loc, laneDataBlockStartCoord, offset));
1247 }
1248 }
1249 assert(static_cast<int64_t>(stepVals.size()) == newVecTy.getNumElements() &&
1250 "Expecting the number of step values to match the number of "
1251 "elements in the vector");
1252 auto stepOpVal =
1253 vector::FromElementsOp::create(rewriter, loc, newVecTy, stepVals);
1254 rewriter.replaceOp(op, stepOpVal);
1255 return success();
1256 }
1257};
1258
1259/// Distributes a subgroup-level vector.extract op to lane-level. Only
1260/// handles sub-vector extraction (result is VectorType, not scalar).
1261struct SgToLaneVectorExtract : public OpConversionPattern<vector::ExtractOp> {
1262 using OpConversionPattern<vector::ExtractOp>::OpConversionPattern;
1263
1264 LogicalResult
1265 matchAndRewrite(vector::ExtractOp op, OpAdaptor adaptor,
1266 ConversionPatternRewriter &rewriter) const override {
1267 // Only handle vector results (not scalar extraction).
1268 auto resultType = dyn_cast<VectorType>(op.getType());
1269 if (!resultType)
1270 return rewriter.notifyMatchFailure(op, "scalar extract not supported");
1271
1272 xegpu::DistributeLayoutAttr layout =
1273 xegpu::getTemporaryLayout(op->getOpResult(0));
1274 if (!layout || !layout.isForSubgroup())
1275 return failure();
1276
1277 // This implementation assumes distribution only happens on the innermost
1278 // dimension. Verify that lane_layout[0...n-2] are all unit.
1279 auto laneLayout = layout.getEffectiveLaneLayoutAsInt();
1280 if (llvm::any_of(ArrayRef<int64_t>(laneLayout).drop_back(1),
1281 [](int64_t v) { return v != 1; }))
1282 return rewriter.notifyMatchFailure(
1283 op, "only innermost dimension distribution is supported for "
1284 "vector.extract");
1285
1286 auto newOp = vector::ExtractOp::create(
1287 rewriter, op.getLoc(), adaptor.getSource(), op.getMixedPosition());
1288 rewriter.replaceOp(op, newOp.getResult());
1289 return success();
1290 }
1291};
1292
1293/// This pattern distributes a subgroup-level ShapeCast op to lane-level.
1294struct SgToLaneVectorShapeCast
1295 : public OpConversionPattern<vector::ShapeCastOp> {
1296 using OpConversionPattern<vector::ShapeCastOp>::OpConversionPattern;
1297
1298 LogicalResult
1299 matchAndRewrite(vector::ShapeCastOp op, OpAdaptor adaptor,
1300 ConversionPatternRewriter &rewriter) const override {
1301 xegpu::DistributeLayoutAttr resultLayout =
1302 xegpu::getTemporaryLayout(op->getOpResult(0));
1303 if (!resultLayout || !resultLayout.isForSubgroup())
1304 return rewriter.notifyMatchFailure(
1305 op, "the result vector of the shape_cast op lacks subgroup layout");
1306
1307 auto resultDistTypeOrFailure = xegpu::getDistVecTypeBasedOnLaneLayout(
1308 resultLayout, op.getResultVectorType());
1309 if (failed(resultDistTypeOrFailure))
1310 return rewriter.notifyMatchFailure(
1311 op, "failed to get distributed vector type for result");
1312
1313 Value source = adaptor.getSource();
1314 auto newShapeCast = vector::ShapeCastOp::create(
1315 rewriter, op.getLoc(), resultDistTypeOrFailure.value(), source);
1316 rewriter.replaceOp(op, newShapeCast);
1317 return success();
1318 }
1319};
1320
1321/// Distributes a subgroup-level vector.extract_strided_slice op to
1322/// lane-level. If the result is distributed, the offsets and sizes are
1323/// adjusted to match the distributed types.
1324struct SgToLaneVectorExtractStridedSlice
1325 : public OpConversionPattern<vector::ExtractStridedSliceOp> {
1326 using OpConversionPattern<vector::ExtractStridedSliceOp>::OpConversionPattern;
1327
1328 LogicalResult
1329 matchAndRewrite(vector::ExtractStridedSliceOp op, OpAdaptor adaptor,
1330 ConversionPatternRewriter &rewriter) const override {
1331 xegpu::DistributeLayoutAttr resultLayout =
1332 xegpu::getTemporaryLayout(op->getOpResult(0));
1333 if (!resultLayout || !resultLayout.isForSubgroup())
1334 return failure();
1335
1336 VectorType resultType = op.getType();
1337 auto distResultTyOrFailure =
1338 xegpu::getDistVecTypeBasedOnLaneLayout(resultLayout, resultType);
1339 if (failed(distResultTyOrFailure))
1340 return rewriter.notifyMatchFailure(
1341 op, "unable to compute distributed vector type from lane layout");
1342 VectorType distResultTy = *distResultTyOrFailure;
1343
1344 SmallVector<int64_t> distributedDims =
1345 getDistributedDims(resultType, distResultTy);
1346
1347 // Collect updated sizes, offsets, strides. Pad to full source rank.
1348 int64_t sourceRank = op.getSourceVectorType().getRank();
1349 SmallVector<Attribute> updatedSizes =
1350 llvm::map_to_vector(op.getSizes(), [](Attribute attr) { return attr; });
1351 SmallVector<Attribute> updatedOffsets = llvm::map_to_vector(
1352 op.getOffsets(), [](Attribute attr) { return attr; });
1353 SmallVector<Attribute> updatedStrides = llvm::map_to_vector(
1354 op.getStrides(), [](Attribute attr) { return attr; });
1355 for (int64_t i = op.getSizes().size(); i < sourceRank; ++i) {
1356 updatedSizes.push_back(
1357 rewriter.getI64IntegerAttr(op.getSourceVectorType().getDimSize(i)));
1358 updatedOffsets.push_back(rewriter.getI64IntegerAttr(0));
1359 updatedStrides.push_back(rewriter.getI64IntegerAttr(1));
1360 }
1361
1362 // Each distributed dim shrinks by its own lane count, so its size and
1363 // offset are rescaled by that count.
1364 if (!distributedDims.empty()) {
1365 auto sourceLayout = xegpu::getTemporaryLayout(op->getOpOperand(0));
1366 if (!sourceLayout || sourceLayout.getEffectiveLaneLayoutAsInt().empty())
1367 return rewriter.notifyMatchFailure(
1368 op, "source of extract_strided_slice lacks distribution layout");
1369 SmallVector<int64_t> laneLayout =
1370 sourceLayout.getEffectiveLaneLayoutAsInt();
1371 SmallVector<int64_t> laneData = sourceLayout.getEffectiveLaneDataAsInt();
1372 ArrayRef<int64_t> sourceShape = op.getSourceVectorType().getShape();
1373 for (int64_t distDim : distributedDims) {
1374 int64_t lanes = laneLayout[distDim];
1375 if (lanes == 0 || sourceShape[distDim] % lanes != 0)
1376 return rewriter.notifyMatchFailure(
1377 op, "source size along a distributed dim is not a multiple of "
1378 "its lane count");
1379 int64_t distrDimOffset =
1380 cast<IntegerAttr>(updatedOffsets[distDim]).getInt();
1381 if (distrDimOffset % (lanes * laneData[distDim]) != 0)
1382 return rewriter.notifyMatchFailure(
1383 op, "offset along a distributed dim is not a multiple of its "
1384 "lane tile");
1385 updatedSizes[distDim] =
1386 rewriter.getI64IntegerAttr(distResultTy.getDimSize(distDim));
1387 updatedOffsets[distDim] =
1388 rewriter.getI64IntegerAttr(distrDimOffset / lanes);
1389 }
1390 }
1391
1392 auto newOp = vector::ExtractStridedSliceOp::create(
1393 rewriter, op.getLoc(), distResultTy, adaptor.getSource(),
1394 ArrayAttr::get(rewriter.getContext(), updatedOffsets),
1395 ArrayAttr::get(rewriter.getContext(), updatedSizes),
1396 ArrayAttr::get(rewriter.getContext(), updatedStrides));
1397 rewriter.replaceOp(op, newOp.getResult());
1398 return success();
1399 }
1400};
1401
1402/// This pattern distributes a subgroup-level `vector.broadcast` op to
1403/// lane-level. The pattern supports three cases:
1404///
1405/// 1) Broadcast a low-rank vector to high-rank vector: The low-rank input
1406/// vector must have a slice layout of the result. If the distributed source
1407/// and target vector types are identical, this lowers to a no-op; otherwise,
1408/// it remains a broadcast but operates on distributed vectors.
1409///
1410/// 2) Broadcast a same-rank vector with identical layouts for source and
1411/// target: The source vector must have unit dimensions, and lane_data must
1412/// be unit size for those unit dims. This always lowers to a no-op.
1413///
1414/// 3) Broadcast a scalar with no layout: This always lowers to a broadcast
1415/// from scalar to distributed result type.
1416///
1417/// Example 1 (low-rank to high-rank broadcast):
1418/// ```
1419/// %0 = "some_op"() {layout_result_0 =
1420/// #xegpu.slice<#xegpu.layout<lane_layout = [1, 16], lane_data = [1, 1]>,
1421/// dims = [0]>} : () -> vector<16xf16>
1422/// %1 = vector.broadcast %0 {layout_result_0 =
1423/// #xegpu.layout<lane_layout = [1, 16], lane_data = [1, 1]>}
1424/// : vector<16xf16> to vector<16x16xf16>
1425/// ```
1426/// is distributed to:
1427/// ```
1428/// %0 = "some_op"() : () -> vector<1xf16>
1429/// %1 = vector.broadcast %0 : vector<1xf16> to vector<16x1xf16>
1430/// ```
1431///
1432/// Example 2 (same-rank broadcast, no-op):
1433/// ```
1434/// %0 = "some_op"() {layout_result_0 =
1435/// #xegpu.layout<lane_layout = [1, 16], lane_data = [1, 1]>}
1436/// : () -> vector<16x1xf16>
1437/// %1 = vector.broadcast %0 {layout_result_0 =
1438/// #xegpu.layout<lane_layout = [1, 16], lane_data = [1, 1]>}
1439/// : vector<16x1xf16> to vector<16x16xf16>
1440/// ```
1441/// is distributed to (no-op, source already matches distributed result type):
1442/// ```
1443/// %0 = "some_op"() : () -> vector<16x1xf16>
1444/// // broadcast is eliminated, %0 is used directly
1445/// ```
1446///
1447/// Example 3 (scalar to vector broadcast):
1448/// ```
1449/// %0 = "some_op"() : () -> f16
1450/// %1 = vector.broadcast %0 {layout_result_0 =
1451/// #xegpu.layout<lane_layout = [1, 16], lane_data = [1, 1]>}
1452/// : f16 to vector<16x16xf16>
1453/// ```
1454/// is distributed to:
1455/// ```
1456/// %0 = "some_op"() : f16
1457/// %1 = vector.broadcast %0 : f16 to vector<16x1xf16>
1458/// ```
1459struct SgToLaneBroadcast : public OpConversionPattern<vector::BroadcastOp> {
1460 using OpConversionPattern<vector::BroadcastOp>::OpConversionPattern;
1461
1462 LogicalResult
1463 matchAndRewrite(vector::BroadcastOp op, OpAdaptor adaptor,
1464 ConversionPatternRewriter &rewriter) const override {
1465 xegpu::DistributeLayoutAttr resultLayout =
1466 xegpu::getTemporaryLayout(cast<OpResult>(op.getResult()));
1467 if (!resultLayout || !resultLayout.isForSubgroup())
1468 return rewriter.notifyMatchFailure(
1469 op, "result does not have subgroup distribute layout");
1470
1471 VectorType destType = op.getResultVectorType();
1472 VectorType sourceType = dyn_cast<VectorType>(op.getSourceType());
1473
1474 xegpu::DistributeLayoutAttr sourceLayout =
1475 xegpu::getTemporaryLayout(op->getOpOperand(0));
1476
1477 if (sourceType) {
1478 int64_t rankDiff = destType.getRank() - sourceType.getRank();
1479 if (rankDiff > 0) {
1480 // Case 1: Low-rank to high-rank broadcast.
1481 if (!sourceLayout || !sourceLayout.isSliceOf(resultLayout))
1482 op.emitWarning(
1483 "broadcast source layout must be a slice of result layout");
1484 } else if (rankDiff == 0) {
1485 // Case 2: Same-rank broadcast.
1486 auto broadcastUnitDimsSet = op.computeBroadcastedUnitDims();
1487 SmallVector<int64_t> broadcastUnitDims(broadcastUnitDimsSet.begin(),
1488 broadcastUnitDimsSet.end());
1489 assert(sourceLayout.isEqualTo(
1490 sourceLayout.setUnitDimData(broadcastUnitDims)) &&
1491 "The sg_data for unit dimensions should be set as 1");
1492 sourceLayout = sourceLayout.setUnitDimLayout(broadcastUnitDims);
1493 }
1494 } else {
1495 // Case 3: Scalar to vector broadcast.
1496 if (sourceLayout)
1497 return rewriter.notifyMatchFailure(
1498 op, "broadcast from scalar must not have a layout attribute");
1499 }
1500
1501 auto destDistType =
1502 xegpu::getDistVecTypeBasedOnLaneLayout(resultLayout, destType);
1503 if (failed(destDistType))
1504 return rewriter.notifyMatchFailure(
1505 op, "failed to distribute the result vector type");
1506
1507 Value source = adaptor.getSource();
1508 // If the adapted source already matches the dest dist type, it's a no-op.
1509 if (source.getType() == destDistType.value()) {
1510 rewriter.replaceOp(op, source);
1511 return success();
1512 }
1513
1514 auto newOp = vector::BroadcastOp::create(rewriter, op.getLoc(),
1515 destDistType.value(), source);
1516 rewriter.replaceOp(op, newOp);
1517 return success();
1518 }
1519};
1520
1521/// Distributes a subgroup-level vector.insert_strided_slice op to
1522/// lane-level. If the dest is distributed, the offsets are adjusted to
1523/// match the distributed types.
1524struct SgToLaneVectorInsertStridedSlice
1525 : public OpConversionPattern<vector::InsertStridedSliceOp> {
1526 using OpConversionPattern<vector::InsertStridedSliceOp>::OpConversionPattern;
1527
1528 LogicalResult
1529 matchAndRewrite(vector::InsertStridedSliceOp op, OpAdaptor adaptor,
1530 ConversionPatternRewriter &rewriter) const override {
1531 xegpu::DistributeLayoutAttr resultLayout =
1532 xegpu::getTemporaryLayout(op->getOpResult(0));
1533 if (!resultLayout || !resultLayout.isForSubgroup())
1534 return failure();
1535
1536 VectorType destType = op.getDestVectorType();
1537 auto distDestTyOrFailure =
1538 xegpu::getDistVecTypeBasedOnLaneLayout(resultLayout, destType);
1539 if (failed(distDestTyOrFailure))
1540 return rewriter.notifyMatchFailure(
1541 op, "unable to compute distributed vector type from lane layout");
1542 VectorType distDestTy = *distDestTyOrFailure;
1543
1544 SmallVector<int64_t> destDistributedDims =
1545 getDistributedDims(destType, distDestTy);
1546
1547 SmallVector<Attribute> updatedOffsets = llvm::map_to_vector(
1548 op.getOffsets(), [](Attribute attr) { return attr; });
1549
1550 if (!destDistributedDims.empty()) {
1551 if (destDistributedDims.size() != 1)
1552 return rewriter.notifyMatchFailure(
1553 op, "only single dimension distribution is supported");
1554 int64_t destDistDim = destDistributedDims[0];
1555
1556 VectorType srcType = op.getSourceVectorType();
1557 // The distributed dim must be in the last k (source rank) dims of dest.
1558 int64_t sourceDistDim =
1559 destDistDim - (destType.getRank() - srcType.getRank());
1560 if (sourceDistDim < 0)
1561 return rewriter.notifyMatchFailure(
1562 op, "distributed dimension must be in the last k dims of dest");
1563
1564 auto destLayout = xegpu::getTemporaryLayout(op->getOpOperand(1));
1565 auto sourceLayout = xegpu::getTemporaryLayout(op->getOpOperand(0));
1566 if (!destLayout || !sourceLayout ||
1567 destLayout.getEffectiveLaneLayoutAsInt().empty() ||
1568 sourceLayout.getEffectiveLaneLayoutAsInt().empty())
1569 return rewriter.notifyMatchFailure(
1570 op, "source or dest of insert_strided_slice lacks distribution "
1571 "layout");
1572
1573 auto destLaneData = destLayout.getEffectiveLaneDataAsInt();
1574 auto sourceLaneData = sourceLayout.getEffectiveLaneDataAsInt();
1575 // Only check lane_data for the distributed dimension. Non-distributed
1576 // dimensions may have non-unit lane_data (e.g., packed layouts).
1577 if ((destDistDim < static_cast<int64_t>(destLaneData.size()) &&
1578 destLaneData[destDistDim] != 1) ||
1579 (sourceDistDim < static_cast<int64_t>(sourceLaneData.size()) &&
1580 sourceLaneData[sourceDistDim] != 1))
1581 return rewriter.notifyMatchFailure(
1582 op, "expecting unit lane data along the distributed dimension");
1583
1584 // The distributed dimension may span only a subset of the subgroup, so
1585 // divide its size and offset by the lanes that actually cover it.
1586 auto destLaneLayout = destLayout.getEffectiveLaneLayoutAsInt();
1587 int64_t numLanesAlongDim =
1588 destDistDim < static_cast<int64_t>(destLaneLayout.size())
1589 ? destLaneLayout[destDistDim]
1590 : 1;
1591
1592 int64_t srcDistrDimSize = srcType.getDimSize(sourceDistDim);
1593 if (srcDistrDimSize % numLanesAlongDim != 0)
1594 return rewriter.notifyMatchFailure(
1595 op, "source distributed dim size is not a multiple of "
1596 "the number of lanes along that dimension");
1597
1598 int64_t destDistrDimOffset =
1599 cast<IntegerAttr>(op.getOffsets()[destDistDim]).getInt();
1600 if (destDistrDimOffset % numLanesAlongDim != 0)
1601 return rewriter.notifyMatchFailure(
1602 op, "offset along distributed dim is not a multiple of "
1603 "the number of lanes along that dimension");
1604 // Adjust offset for the distributed dimension.
1605 updatedOffsets[destDistDim] =
1606 rewriter.getI64IntegerAttr(destDistrDimOffset / numLanesAlongDim);
1607 }
1608
1609 auto newOp = vector::InsertStridedSliceOp::create(
1610 rewriter, op.getLoc(), distDestTy, adaptor.getValueToStore(),
1611 adaptor.getDest(),
1612 ArrayAttr::get(rewriter.getContext(), updatedOffsets), op.getStrides());
1613 rewriter.replaceOp(op, newOp.getResult());
1614 return success();
1615 }
1616};
1617
1618/// Distributes a subgroup-level vector.insert op to lane-level. Only
1619/// handles sub-vector insertion (value to store is VectorType, not scalar).
1620struct SgToLaneVectorInsert : public OpConversionPattern<vector::InsertOp> {
1621 using OpConversionPattern<vector::InsertOp>::OpConversionPattern;
1622
1623 LogicalResult
1624 matchAndRewrite(vector::InsertOp op, OpAdaptor adaptor,
1625 ConversionPatternRewriter &rewriter) const override {
1626 // Only handle vector value-to-store (not scalar insertion).
1627 auto valueType = dyn_cast<VectorType>(op.getValueToStoreType());
1628 if (!valueType)
1629 return rewriter.notifyMatchFailure(op, "scalar insert not supported");
1630
1631 xegpu::DistributeLayoutAttr layout =
1632 xegpu::getTemporaryLayout(op->getOpResult(0));
1633 if (!layout || !layout.isForSubgroup())
1634 return failure();
1635
1636 // verify that the outer k dimensions (for offsets)
1637 // don't have non-unit lane_layout.
1638 auto laneLayout = layout.getEffectiveLaneLayoutAsInt();
1639 if (llvm::any_of(ArrayRef<int64_t>(laneLayout).drop_back(1),
1640 [](int64_t v) { return v != 1; }))
1641 return rewriter.notifyMatchFailure(
1642 op, "only innermost dimension distribution is supported for "
1643 "vector.insert");
1644
1645 auto newOp = vector::InsertOp::create(
1646 rewriter, op.getLoc(), adaptor.getValueToStore(), adaptor.getDest(),
1647 op.getMixedPosition());
1648 rewriter.replaceOp(op, newOp.getResult());
1649 return success();
1650 }
1651};
1652
1653/// Redistributes `src` for a `convert_layout` that changes only the
1654/// `lane_layout` along the outer (distributed) dimension, shrinking it from
1655/// `currentLaneNum` to `targetLaneNum` lanes (a partial-subgroup
1656/// distribution). Because the data is no longer replicated across all lanes,
1657/// each surviving lane must gather the values that previously lived in the
1658/// lanes that are dropped. The values are gathered with `gpu.shuffle` and
1659/// concatenated with the lane-local data using `vector.shuffle`, which doubles
1660/// the distributed outer dimension when the lane count is halved.
1661///
1662/// Only halving the lane count (a factor of two) is currently supported.
1663/// Returns the redistributed value on success, or failure if `src` cannot be
1664/// shuffled (e.g. it is not a rank-2 vector or its bit width is not a multiple
1665/// of 32).
1666static FailureOr<Value>
1667shuffleDataAsLaneLayoutChange(ConversionPatternRewriter &rewriter, Location loc,
1668 Value src, int64_t currentLaneNum,
1669 int64_t targetLaneNum) {
1670 VectorType srcTy = dyn_cast<VectorType>(src.getType());
1671 if (!srcTy || srcTy.getRank() != 2)
1672 return failure();
1673 // Only halving the lane count (factor of two) is supported for now.
1674 if (targetLaneNum <= 0 || currentLaneNum != targetLaneNum * 2)
1675 return failure();
1676 // gpu.shuffle operates on i32, so the data must be a multiple of 32 bits.
1677 int64_t vectorBitWidth =
1678 srcTy.getNumElements() * srcTy.getElementTypeBitWidth();
1679 if (vectorBitWidth % 32 != 0)
1680 return failure();
1681
1682 // A vector cannot be shuffled across lanes directly:
1683 // -- cast the source to a 1D vector of i32
1684 // -- create a temp 1D vector of i32 initialized to zero
1685 // -- for each i32 element:
1686 // ---- extract it from the source bundle
1687 // ---- gpu.shuffle to gather the value from the partner lane
1688 // ---- insert it into the temp bundle
1689 // -- cast the temp back to the source vector type
1690 // -- vector.shuffle the source and temp to concatenate along the outer dim
1691 Type shuffleElemTy = rewriter.getI32Type();
1692 int64_t numShuffles = vectorBitWidth / 32;
1693 VectorType shuffleBundleTy = VectorType::get({numShuffles}, shuffleElemTy);
1694 // Initialize temp to zero.
1695 Value temp = arith::ConstantOp::create(
1696 rewriter, loc,
1697 DenseElementsAttr::get(shuffleBundleTy,
1698 IntegerAttr::get(shuffleElemTy, 0)));
1699 VectorType flatSrcTy =
1700 VectorType::get({srcTy.getNumElements()}, srcTy.getElementType());
1701 Value flatSrc = vector::ShapeCastOp::create(rewriter, loc, flatSrcTy, src);
1702 Value shuffleBundle =
1703 vector::BitCastOp::create(rewriter, loc, shuffleBundleTy, flatSrc);
1704 for (int64_t i = 0; i < numShuffles; i++) {
1705 Value shuffleElem =
1706 vector::ExtractOp::create(rewriter, loc, shuffleBundle, i);
1707 shuffleElem = gpu::ShuffleOp::create(rewriter, loc, shuffleElem, 0,
1708 targetLaneNum, gpu::ShuffleMode::UP)
1709 .getResult(0);
1710 temp = vector::InsertOp::create(rewriter, loc, shuffleElem, temp, i);
1711 }
1712 temp = vector::BitCastOp::create(rewriter, loc, flatSrcTy, temp);
1713 temp = vector::ShapeCastOp::create(rewriter, loc, srcTy, temp);
1714
1715 // Concatenate the lane-local and gathered data along the outer dimension.
1716 SmallVector<int64_t> indices(srcTy.getShape()[0] * 2);
1717 std::iota(indices.begin(), indices.end(), 0);
1718 Value res = vector::ShuffleOp::create(rewriter, loc, src, temp, indices);
1719 return res;
1720}
1721
1722/// Repacks `src`'s `lane_data` along `repackDim` between round-robin and
1723/// contiguous form with an `xegpu.lane_shuffle`, which moves each lane's run of
1724/// `k` elements across lanes while preserving the element type.
1725///
1726/// `inputData`/`targetData` are the `repackDim` `lane_data` of the input and
1727/// target layouts; exactly one must be 1 (round-robin) and the other `k`
1728/// (contiguous). Returns failure if that does not hold.
1729static FailureOr<Value> repackLaneData(ConversionPatternRewriter &rewriter,
1730 Location loc, Value src,
1731 int64_t repackDim, int64_t inputData,
1732 int64_t targetData) {
1733 auto srcTy = dyn_cast<VectorType>(src.getType());
1734 if (!srcTy)
1735 return failure();
1736 int64_t rank = srcTy.getRank();
1737 Type elemTy = srcTy.getElementType();
1738 int64_t k = srcTy.getShape()[repackDim];
1739
1740 bool roundRobinToContig = inputData == 1 && targetData == k;
1741 bool contigToRoundRobin = inputData == k && targetData == 1;
1742 if (!roundRobinToContig && !contigToRoundRobin)
1743 return failure();
1744
1745 // Round-robin -> contiguous gathers a lane's strided elements into
1746 // consecutive positions (pack); the reverse scatters them back (unpack).
1747 xegpu::LaneShuffleMode mode = roundRobinToContig
1748 ? xegpu::LaneShuffleMode::Pack
1749 : xegpu::LaneShuffleMode::Unpack;
1750 VectorType runTy = VectorType::get({k}, elemTy);
1751
1752 // Common case: the lane fragment is a single run (every dimension other than
1753 // `repackDim` is unit), so collapse it to 1D, shuffle once, and restore it.
1754 if (srcTy.getNumElements() == k) {
1755 if (rank == 1)
1756 return Value(
1757 xegpu::LaneShuffleOp::create(rewriter, loc, runTy, src, mode));
1758 Value flat = vector::ShapeCastOp::create(rewriter, loc, runTy, src);
1759 Value shuffled =
1760 xegpu::LaneShuffleOp::create(rewriter, loc, runTy, flat, mode);
1761 return Value(vector::ShapeCastOp::create(rewriter, loc, srcTy, shuffled));
1762 }
1763
1764 // When `repackDim` is innermost each run is a contiguous sub-vector, so it is
1765 // extracted and re-inserted as a whole.
1766 if (repackDim == rank - 1) {
1767 SmallVector<int64_t> outerShape(srcTy.getShape().drop_back());
1768 int64_t numRuns = computeProduct(outerShape);
1769 SmallVector<int64_t> outerStrides = computeStrides(outerShape);
1770 Value result = arith::ConstantOp::create(rewriter, loc, srcTy,
1771 rewriter.getZeroAttr(srcTy));
1772 for (int64_t i = 0; i < numRuns; ++i) {
1773 SmallVector<int64_t> pos = delinearize(i, outerStrides);
1774 Value run = vector::ExtractOp::create(rewriter, loc, src, pos);
1775 Value shuffled =
1776 xegpu::LaneShuffleOp::create(rewriter, loc, runTy, run, mode);
1777 result = vector::InsertOp::create(rewriter, loc, shuffled, result, pos);
1778 }
1779 return result;
1780 }
1781
1782 // Otherwise each run is strided along `repackDim`: extract the `k`-long slice
1783 // (a sub-vector that is unit along every other dim), flatten it to 1D,
1784 // shuffle, and insert it back.
1785 SmallVector<int64_t> keptShape;
1786 SmallVector<int64_t> keptDims;
1787 for (int64_t d = 0; d < rank; ++d)
1788 if (d != repackDim) {
1789 keptShape.push_back(srcTy.getShape()[d]);
1790 keptDims.push_back(d);
1791 }
1792 int64_t numRuns = computeProduct(keptShape);
1793 SmallVector<int64_t> keptStrides = computeStrides(keptShape);
1794 SmallVector<int64_t> sliceSizes(rank, 1);
1795 sliceSizes[repackDim] = k;
1796 SmallVector<int64_t> sliceStrides(rank, 1);
1797 VectorType sliceTy = VectorType::get(sliceSizes, elemTy);
1798 Value result = arith::ConstantOp::create(rewriter, loc, srcTy,
1799 rewriter.getZeroAttr(srcTy));
1800 for (int64_t i = 0; i < numRuns; ++i) {
1801 SmallVector<int64_t> keptPos = delinearize(i, keptStrides);
1802 SmallVector<int64_t> offsets(rank, 0);
1803 for (auto [dim, coord] : llvm::zip_equal(keptDims, keptPos))
1804 offsets[dim] = coord;
1805 Value slice = vector::ExtractStridedSliceOp::create(
1806 rewriter, loc, src, offsets, sliceSizes, sliceStrides);
1807 Value run = vector::ShapeCastOp::create(rewriter, loc, runTy, slice);
1808 Value repacked =
1809 xegpu::LaneShuffleOp::create(rewriter, loc, runTy, run, mode);
1810 Value repackedSlice =
1811 vector::ShapeCastOp::create(rewriter, loc, sliceTy, repacked);
1812 result = vector::InsertStridedSliceOp::create(
1813 rewriter, loc, repackedSlice, result, offsets, sliceStrides);
1814 }
1815 return result;
1816}
1817
1818/// Folds a subgroup-level ConvertLayout op with compatible lane layouts.
1819struct SgToLaneConvertLayout
1820 : public OpConversionPattern<xegpu::ConvertLayoutOp> {
1821 using OpConversionPattern<xegpu::ConvertLayoutOp>::OpConversionPattern;
1822
1823 LogicalResult
1824 matchAndRewrite(xegpu::ConvertLayoutOp op, OpAdaptor adaptor,
1825 ConversionPatternRewriter &rewriter) const override {
1826 auto inputLayout = op.getEffectiveInputLayout();
1827 auto targetLayout = op.getTargetLayoutAttr();
1828 Type valType = op.getResult().getType();
1829
1830 if (valType.isIntOrFloat()) {
1831 rewriter.replaceOp(op, op.getSource());
1832 return success();
1833 }
1834
1835 auto resShape = cast<VectorType>(valType).getShape();
1836 SmallVector<int64_t> resShapeVec(resShape.begin(), resShape.end());
1837
1838 // Equivalent layouts: the convert_layout is a no-op and folds to its
1839 // source.
1840 if (inputLayout.isCompatibleWith(targetLayout, resShapeVec,
1842 rewriter.replaceOp(op, adaptor.getSource());
1843 return success();
1844 }
1845
1846 // Handle the special case where the conversion redistributes a value
1847 // across a fraction of the subgroup: the lane_layout shrinks along the
1848 // outer (distributed) dimension while lane_data stays the same. Only a
1849 // pure outer-dimension lane_layout change is supported, so the inner
1850 // lane_layout must be unit (making the outer dim the only distributed one)
1851 // and the outer lane_layout must be genuinely distributed (> 1), which
1852 // also rules out the degenerate [1, 1] layout.
1853 if (inputLayout.getEffectiveOrderAsInt() ==
1854 targetLayout.getEffectiveOrderAsInt() &&
1855 inputLayout.getRank() == 2 && targetLayout.getRank() == 2) {
1856 auto laneLayout = inputLayout.getEffectiveLaneLayoutAsInt();
1857 auto targetLaneLayout = targetLayout.getEffectiveLaneLayoutAsInt();
1858 auto laneData = inputLayout.getEffectiveLaneDataAsInt();
1859 auto targetLaneData = targetLayout.getEffectiveLaneDataAsInt();
1860 if (laneLayout.size() == 2 && targetLaneLayout.size() == 2 &&
1861 laneData == targetLaneData && laneLayout[1] == 1 &&
1862 targetLaneLayout[1] == 1 && laneLayout[0] > 1 &&
1863 laneLayout[0] != targetLaneLayout[0]) {
1864 FailureOr<Value> res = shuffleDataAsLaneLayoutChange(
1865 rewriter, op.getLoc(), adaptor.getSource(), laneLayout[0],
1866 targetLaneLayout[0]);
1867 if (succeeded(res)) {
1868 rewriter.replaceOp(op, *res);
1869 return success();
1870 }
1871 }
1872 }
1873
1874 // Handle a pure `lane_data` repack: `lane_layout` and `order` are unchanged
1875 // and exactly one dimension's `lane_data` switches between round-robin
1876 // (lane_data 1) and contiguous (lane_data == run length). The elements per
1877 // lane are unchanged, but their assignment to lanes is not, so the data is
1878 // moved across lanes with `xegpu.lane_shuffle`. The changed dimension must
1879 // be one of the two innermost ones, since sg-to-lane distribution is 2D.
1880 if (inputLayout.getEffectiveOrderAsInt() ==
1881 targetLayout.getEffectiveOrderAsInt() &&
1882 inputLayout.getEffectiveLaneLayoutAsInt() ==
1883 targetLayout.getEffectiveLaneLayoutAsInt()) {
1884 auto laneLayout = inputLayout.getEffectiveLaneLayoutAsInt();
1885 auto laneData = inputLayout.getEffectiveLaneDataAsInt();
1886 auto targetLaneData = targetLayout.getEffectiveLaneDataAsInt();
1887 // Find the single dimension whose lane_data changed; bail out if more
1888 // than one differs.
1889 int64_t rank = laneData.size();
1890 int64_t repackDim = -1;
1891 bool multipleChanged = false;
1892 for (int64_t d = 0; d < rank; ++d)
1893 if (laneData[d] != targetLaneData[d]) {
1894 if (repackDim != -1)
1895 multipleChanged = true;
1896 repackDim = d;
1897 }
1898
1899 // `repackDim` must be the distributed dim (lane_layout != 1) and the
1900 // other innermost dim non-distributed (lane_layout == 1).
1901 int64_t otherDim = repackDim == rank - 1 ? rank - 2 : rank - 1;
1902 bool laneLayoutOk = repackDim != -1 && laneLayout[repackDim] != 1 &&
1903 (rank < 2 || laneLayout[otherDim] == 1);
1904
1905 // Exactly one dimension must change, and it must be one of the two
1906 // innermost (>= rank - 2).
1907 if (repackDim != -1 && repackDim >= rank - 2 && !multipleChanged &&
1908 laneLayoutOk) {
1909 FailureOr<Value> res = repackLaneData(
1910 rewriter, op.getLoc(), adaptor.getSource(), repackDim,
1911 laneData[repackDim], targetLaneData[repackDim]);
1912 if (succeeded(res)) {
1913 rewriter.replaceOp(op, *res);
1914 return success();
1915 }
1916 }
1917 }
1918
1919 return rewriter.notifyMatchFailure(
1920 op, "lowering incompatible convert_layout not yet supported");
1921 }
1922};
1923
1924/// The offsets and shape of one `lane_data`-sized piece of a lane's fragment.
1925/// `computeDistributedCoords` returns one coordinate per distribution unit,
1926/// enumerated row major over `shape / (lane_layout * lane_data)`, and unit `u`
1927/// owns the piece of the fragment at `delinearize(u) * lane_data`.
1928struct LaneFragmentPiece {
1929 SmallVector<int64_t> offsets;
1931};
1932
1933static LaneFragmentPiece getLaneFragmentPiece(int64_t unit,
1935 ArrayRef<int64_t> laneLayout,
1936 ArrayRef<int64_t> laneData) {
1937 int64_t rank = shape.size();
1938 SmallVector<int64_t> units(rank);
1939 for (int64_t d = 0; d < rank; ++d)
1940 units[d] = shape[d] / std::min(shape[d], laneLayout[d] * laneData[d]);
1941 SmallVector<int64_t> unitCoords = delinearize(unit, computeStrides(units));
1942 LaneFragmentPiece piece;
1943 for (int64_t d = 0; d < rank; ++d) {
1944 piece.offsets.push_back(unitCoords[d] * laneData[d]);
1945 piece.sizes.push_back(laneData[d]);
1946 }
1947 return piece;
1948}
1949
1950/// Last-resort lowering for a `convert_layout` that none of the patterns above
1951/// handle: round-trip the value through shared local memory, writing each
1952/// lane's fragment to the coordinates `input_layout` gives it and reading it
1953/// back from the ones `target_layout` gives it. Any pair of layouts that can
1954/// both distribute the value works, at the cost of an SLM write and read.
1955///
1956/// The op is subgroup level, but every subgroup of the workgroup executes it on
1957/// its own data, so the scratch holds one tile per subgroup and each subgroup
1958/// addresses its own through a `gpu.subgroup_id` offset on the outermost
1959/// dimension. The subgroup count comes from the kernel's `known_block_size`.
1960///
1961/// A lane's fragment is one `lane_data`-sized piece per distribution unit, so
1962/// the transfer is one `store_matrix`/`load_matrix` per unit; layouts with
1963/// a single unit, which is the common case, give a single pair.
1964struct SgToLaneConvertLayoutViaSLM
1965 : public OpConversionPattern<xegpu::ConvertLayoutOp> {
1966 using OpConversionPattern<xegpu::ConvertLayoutOp>::OpConversionPattern;
1967
1968 LogicalResult
1969 matchAndRewrite(xegpu::ConvertLayoutOp op, OpAdaptor adaptor,
1970 ConversionPatternRewriter &rewriter) const override {
1971 xegpu::DistributeLayoutAttr inputLayout = op.getEffectiveInputLayout();
1972 xegpu::DistributeLayoutAttr targetLayout = op.getTargetLayoutAttr();
1973 if (!inputLayout || !targetLayout || !inputLayout.isForSubgroup() ||
1974 !targetLayout.isForSubgroup())
1975 return rewriter.notifyMatchFailure(op, "both layouts must be lane level");
1976
1977 auto valueTy = dyn_cast<VectorType>(op.getResult().getType());
1978 if (!valueTy)
1979 return rewriter.notifyMatchFailure(op, "value type must be a vector");
1980 Type elemTy = valueTy.getElementType();
1981 // The scratch is sized in bytes, so sub-byte elements would not get a
1982 // well-defined address.
1983 if (!elemTy.isIntOrFloat() || elemTy.getIntOrFloatBitWidth() % 8 != 0)
1984 return rewriter.notifyMatchFailure(
1985 op, "element type must be a whole number of bytes");
1986
1987 ArrayRef<int64_t> sgShape = valueTy.getShape();
1988 int64_t rank = sgShape.size();
1989 if (rank != inputLayout.getRank() || rank != targetLayout.getRank())
1990 return rewriter.notifyMatchFailure(
1991 op, "both layouts must have the rank of the value");
1992
1993 FailureOr<VectorType> distInputTy =
1994 xegpu::getDistVecTypeBasedOnLaneLayout(inputLayout, valueTy);
1995 FailureOr<VectorType> distTargetTy =
1996 xegpu::getDistVecTypeBasedOnLaneLayout(targetLayout, valueTy);
1997 if (failed(distInputTy) || failed(distTargetTy))
1998 return rewriter.notifyMatchFailure(
1999 op, "value type must be distributable by both layouts");
2000
2001 const auto *uArch =
2002 xegpu::uArch::getUArch(xegpu::getChipStr(op).value_or(""));
2003 if (!uArch)
2004 return rewriter.notifyMatchFailure(
2005 op, "target attribute is required to determine the subgroup size");
2006 FailureOr<int64_t> numSubgroups =
2007 xegpu::getNumSubgroupsFromBlockSize(op, uArch->getSubgroupSize());
2008 if (failed(numSubgroups))
2009 return rewriter.notifyMatchFailure(
2010 op, "the scratch holds one tile per subgroup, so @known_block_size "
2011 "must be attached to the kernel, with power-of-two dimensions "
2012 "covering at least one subgroup");
2013
2014 Location loc = op.getLoc();
2015
2016 // One tile per subgroup, stacked along the outermost dimension.
2017 SmallVector<int64_t> slmShape(sgShape);
2018 slmShape[0] *= *numSubgroups;
2019 int64_t slmBytes =
2020 computeProduct(slmShape) * elemTy.getIntOrFloatBitWidth() / 8;
2021 auto slmTy = MemRefType::get({slmBytes}, rewriter.getI8Type(), {}, 3);
2022 Value slm = memref::AllocaOp::create(rewriter, loc, slmTy);
2023 Value memDesc = xegpu::CreateMemDescOp::create(
2024 rewriter, loc,
2025 xegpu::MemDescType::get(rewriter.getContext(), slmShape, elemTy,
2026 nullptr),
2027 slm);
2028
2029 // This subgroup's tile starts at `subgroup_id` tiles into the scratch.
2030 Value sgId = gpu::SubgroupIdOp::create(rewriter, loc,
2031 rewriter.getIndexType(), nullptr);
2032 SmallVector<Value> base(rank);
2033 base[0] = arith::MulIOp::create(
2034 rewriter, loc, sgId,
2035 arith::ConstantIndexOp::create(rewriter, loc, sgShape[0]));
2036 for (int64_t d = 1; d < rank; ++d)
2037 base[d] = arith::ConstantIndexOp::create(rewriter, loc, 0);
2038
2039 Value laneId = gpu::LaneIdOp::create(rewriter, loc, rewriter.getIndexType(),
2040 /*upperBound=*/mlir::IntegerAttr());
2041 auto dynamicOffsets = rewriter.getDenseI64ArrayAttr(
2042 SmallVector<int64_t>(rank, ShapedType::kDynamic));
2043
2044 // Adds this subgroup's base to one distribution unit's coordinates.
2045 auto withBase = [&](ArrayRef<Value> coords) {
2046 SmallVector<Value> offsets;
2047 for (auto [coord, b] : llvm::zip_equal(coords, base))
2048 offsets.push_back(arith::AddIOp::create(rewriter, loc, coord, b));
2049 return offsets;
2050 };
2051
2052 // Write phase: every lane stores the pieces `input_layout` gives it.
2053 auto storeCoords =
2054 inputLayout.computeDistributedCoords(rewriter, loc, laneId, sgShape);
2055 if (failed(storeCoords))
2056 return rewriter.notifyMatchFailure(
2057 op, "failed to compute the input_layout coordinates");
2058 SmallVector<int64_t> inputLaneLayout =
2059 inputLayout.getEffectiveLaneLayoutAsInt();
2060 SmallVector<int64_t> inputLaneData =
2061 inputLayout.getEffectiveLaneDataAsInt();
2062 Value fragment =
2063 castValueTo(rewriter, cast<TypedValue<VectorType>>(adaptor.getSource()),
2064 *distInputTy);
2065 for (auto [unit, coords] : llvm::enumerate(*storeCoords)) {
2066 LaneFragmentPiece piece =
2067 getLaneFragmentPiece(unit, sgShape, inputLaneLayout, inputLaneData);
2068 Value data = fragment;
2069 if (piece.sizes != SmallVector<int64_t>(distInputTy->getShape()))
2070 data = vector::ExtractStridedSliceOp::create(
2071 rewriter, loc, fragment, piece.offsets, piece.sizes,
2072 SmallVector<int64_t>(rank, 1));
2073 xegpu::StoreMatrixOp::create(rewriter, loc, TypeRange{}, data, memDesc,
2074 ValueRange(withBase(coords)), dynamicOffsets,
2075 nullptr, xegpu::DistributeLayoutAttr{});
2076 }
2077
2078 // The scratch is only exchanged between the lanes of one subgroup, which is
2079 // a single hardware thread, so ordering the write before the read is enough
2080 // and no workgroup barrier is needed.
2081 xegpu::FenceOp::create(rewriter, loc, xegpu::MemorySpace::SLM,
2082 xegpu::FenceScope::Workgroup);
2083
2084 // Read phase: every lane loads the pieces `target_layout` gives it.
2085 auto loadCoords =
2086 targetLayout.computeDistributedCoords(rewriter, loc, laneId, sgShape);
2087 if (failed(loadCoords))
2088 return rewriter.notifyMatchFailure(
2089 op, "failed to compute the target_layout coordinates");
2090 SmallVector<int64_t> targetLaneLayout =
2091 targetLayout.getEffectiveLaneLayoutAsInt();
2092 SmallVector<int64_t> targetLaneData =
2093 targetLayout.getEffectiveLaneDataAsInt();
2094 Value result;
2095 if (loadCoords->size() > 1)
2096 result = arith::ConstantOp::create(rewriter, loc, *distTargetTy,
2097 rewriter.getZeroAttr(*distTargetTy));
2098 for (auto [unit, coords] : llvm::enumerate(*loadCoords)) {
2099 LaneFragmentPiece piece =
2100 getLaneFragmentPiece(unit, sgShape, targetLaneLayout, targetLaneData);
2101 auto pieceTy = VectorType::get(piece.sizes, elemTy);
2102 Value loaded = xegpu::LoadMatrixOp::create(
2103 rewriter, loc, pieceTy, memDesc, ValueRange(withBase(coords)),
2104 dynamicOffsets, nullptr, xegpu::DistributeLayoutAttr{});
2105 if (!result) {
2106 result = loaded;
2107 break;
2108 }
2109 result = vector::InsertStridedSliceOp::create(
2110 rewriter, loc, loaded, result, piece.offsets,
2111 SmallVector<int64_t>(rank, 1));
2112 }
2113
2114 rewriter.replaceOp(op, castValueTo(rewriter,
2116 *distTargetTy));
2117 return success();
2118 }
2119};
2120
2121/// `getEffectiveLaneDataAsInt` is empty when `lane_data` is unset, so the unit
2122/// check also rejects layouts that are not lane level.
2123static bool hasDefaultOrderAndUnitLaneData(xegpu::DistributeLayoutAttr layout) {
2124 if (layout.getRank() != 2)
2125 return false;
2126 return layout.getEffectiveLaneDataAsInt() == SmallVector<int64_t>{1, 1} &&
2127 layout.getEffectiveOrderAsInt() == SmallVector<int64_t>{1, 0};
2128}
2129
2130/// The quantities every element-to-lane redistribution needs, see
2131/// `matchElementLaneRedistribution`.
2132struct ElementLaneRedistribution {
2133 /// Subgroup-level type of the converted value.
2134 VectorType valueType;
2135 /// `valueType` as distributed by `input_layout` and by `target_layout`.
2136 VectorType distributedInput;
2137 VectorType distributedTarget;
2138 /// Effective `lane_layout` of `input_layout` and of `target_layout`. These
2139 /// always differ, otherwise the conversion would already have folded.
2140 SmallVector<int64_t> inputLaneLayout;
2141 SmallVector<int64_t> targetLaneLayout;
2142 int64_t subgroupSize;
2143};
2144
2145/// Recognizes an `xegpu.convert_layout` that redistributes individual elements
2146/// between the lanes of a subgroup, which is the common condition the three
2147/// slice-attributed lowerings below have.
2148/// Both input and target layouts must satisfy `hasDefaultOrderAndUnitLaneData`
2149/// and must be able to distribute the value, the two can only differ in which
2150/// lane owns which element.
2151static FailureOr<ElementLaneRedistribution>
2152matchElementLaneRedistribution(xegpu::ConvertLayoutOp op,
2153 ConversionPatternRewriter &rewriter) {
2154 xegpu::DistributeLayoutAttr inputLayout = op.getEffectiveInputLayout();
2155 xegpu::DistributeLayoutAttr targetLayout = op.getTargetLayoutAttr();
2156 if (!hasDefaultOrderAndUnitLaneData(inputLayout) ||
2157 !hasDefaultOrderAndUnitLaneData(targetLayout))
2158 return rewriter.notifyMatchFailure(
2159 op, "both layouts must be rank 2 with effective lane_data [1, 1] and "
2160 "effective order [1, 0]");
2161
2162 auto valueType = dyn_cast<VectorType>(op.getResult().getType());
2163 if (!valueType)
2164 return rewriter.notifyMatchFailure(op, "value type must be a vector");
2165
2166 FailureOr<VectorType> distributedInput =
2167 xegpu::getDistVecTypeBasedOnLaneLayout(inputLayout, valueType);
2168 FailureOr<VectorType> distributedTarget =
2169 xegpu::getDistVecTypeBasedOnLaneLayout(targetLayout, valueType);
2170 if (failed(distributedInput) || failed(distributedTarget))
2171 return rewriter.notifyMatchFailure(
2172 op, "value type must be distributable by both layouts");
2173
2174 const auto *uArch =
2175 xegpu::uArch::getUArch(xegpu::getChipStr(op).value_or(""));
2176 if (!uArch)
2177 return rewriter.notifyMatchFailure(
2178 op, "target attribute is required to determine the subgroup size");
2179
2180 ElementLaneRedistribution redistribution;
2181 redistribution.valueType = valueType;
2182 redistribution.distributedInput = *distributedInput;
2183 redistribution.distributedTarget = *distributedTarget;
2184 redistribution.inputLaneLayout = inputLayout.getEffectiveLaneLayoutAsInt();
2185 redistribution.targetLaneLayout = targetLayout.getEffectiveLaneLayoutAsInt();
2186 redistribution.subgroupSize = uArch->getSubgroupSize();
2187 return redistribution;
2188}
2189
2190/// Returns the lane stride of the single distributed dimension of `slice`, i.e.
2191/// the distance in lane ids between two adjacent coordinates along that
2192/// dimension. `slice` must have exactly one non-sliced dimension whose parent
2193/// `lane_layout` extent is greater than one; the stride is the product of the
2194/// parent `lane_layout` extents that precede it in the parent `order`.
2195///
2196/// Example: for
2197/// #xegpu.slice<#xegpu.layout<lane_layout = [8, 1, 2], lane_data = [4, 1, 1],
2198/// order = [0, 2, 1]>, dims = [0]>
2199/// the only such dimension is parent dim 2, of extent 2. `order` makes dim 0
2200/// the fastest varying, with extent 8, so lanes 0..7 hold coordinate 0 and
2201/// lanes 8..15 hold coordinate 1: the stride is 8.
2202static FailureOr<int64_t> getDistributedDimLaneStride(xegpu::SliceAttr slice) {
2203 xegpu::SliceAttr flattened = slice.flatten();
2204 auto parent = dyn_cast<xegpu::LayoutAttr>(flattened.getParent());
2205 if (!parent)
2206 return failure();
2207
2208 SmallVector<int64_t> parentLaneLayout = parent.getEffectiveLaneLayoutAsInt();
2209 SmallVector<int64_t> parentOrder = parent.getEffectiveOrderAsInt();
2210 if (parentLaneLayout.size() != parentOrder.size())
2211 return failure();
2212
2213 ArrayRef<int64_t> slicedDims = flattened.getDims().asArrayRef();
2214 std::optional<int64_t> distributedDim;
2215 for (int64_t dim = 0, rank = static_cast<int64_t>(parentLaneLayout.size());
2216 dim < rank; ++dim) {
2217 if (llvm::is_contained(slicedDims, dim) || parentLaneLayout[dim] == 1)
2218 continue;
2219 if (distributedDim)
2220 return failure();
2221 distributedDim = dim;
2222 }
2223 if (!distributedDim)
2224 return failure();
2225
2226 int64_t stride = 1;
2227 for (int64_t dim : parentOrder) {
2228 if (dim == *distributedDim)
2229 return stride;
2230 stride *= parentLaneLayout[dim];
2231 }
2232 return failure();
2233}
2234
2235/// Distributes the slice-attributed `xegpu.convert_layout` whose source is
2236/// fully broadcast, i.e. every lane holds the whole value, onto the rows of a
2237/// subset of the lanes. Each lane keeps a single element, so no data crosses
2238/// lanes and one `vector.extract` suffices.
2239///
2240/// The input layout has effective `lane_layout` [1, 1], so the distributed
2241/// source is the whole value; the target has [n, 1], so lane `l` keeps row
2242/// `l % n` and the distributed result is `vector<1x1>`.
2243///
2244/// The source is flattened first because `xegpu-vector-linearize` cannot
2245/// linearize a `vector.extract` with a dynamic position out of a rank-2 value.
2246struct SgToLaneConvertLayoutBroadcastExtract
2247 : public OpConversionPattern<xegpu::ConvertLayoutOp> {
2248 using OpConversionPattern<xegpu::ConvertLayoutOp>::OpConversionPattern;
2249
2250 LogicalResult
2251 matchAndRewrite(xegpu::ConvertLayoutOp op, OpAdaptor adaptor,
2252 ConversionPatternRewriter &rewriter) const override {
2253 xegpu::DistributeLayoutAttr inputLayout = op.getEffectiveInputLayout();
2254 xegpu::DistributeLayoutAttr targetLayout = op.getTargetLayoutAttr();
2255
2256 if (!isa<xegpu::SliceAttr>(inputLayout))
2257 return rewriter.notifyMatchFailure(op,
2258 "input_layout must be #xegpu.slice");
2259 if (!isa<xegpu::LayoutAttr>(targetLayout))
2260 return rewriter.notifyMatchFailure(op,
2261 "target_layout must be #xegpu.layout");
2262
2263 FailureOr<ElementLaneRedistribution> redistribution =
2264 matchElementLaneRedistribution(op, rewriter);
2265 if (failed(redistribution))
2266 return failure();
2267
2268 if (redistribution->inputLaneLayout != SmallVector<int64_t>{1, 1})
2269 return rewriter.notifyMatchFailure(
2270 op, "input_layout effective lane_layout must be [1, 1]");
2271 SmallVector<int64_t> targetLaneLayout = redistribution->targetLaneLayout;
2272 if (targetLaneLayout[0] <= 1 || targetLaneLayout[1] != 1)
2273 return rewriter.notifyMatchFailure(
2274 op, "target_layout effective lane_layout must be [n, 1] with n > 1");
2275 if (redistribution->distributedInput != redistribution->valueType)
2276 return rewriter.notifyMatchFailure(
2277 op, "distributed input_layout type must equal the value type");
2278 if (redistribution->distributedTarget.getShape() != ArrayRef<int64_t>{1, 1})
2279 return rewriter.notifyMatchFailure(
2280 op, "distributed target_layout type must be vector<1x1>");
2281 if (targetLaneLayout[0] > redistribution->subgroupSize)
2282 return rewriter.notifyMatchFailure(
2283 op, "target_layout effective lane_layout[0] must not exceed the "
2284 "subgroup size");
2285
2286 VectorType valueType = redistribution->valueType;
2287 Location loc = op.getLoc();
2288 Value src =
2289 castValueTo(rewriter, cast<TypedValue<VectorType>>(adaptor.getSource()),
2290 redistribution->distributedInput);
2291 auto flatType = VectorType::get({valueType.getNumElements()},
2292 valueType.getElementType());
2293 Value flat = vector::ShapeCastOp::create(rewriter, loc, flatType, src);
2294 Value laneId = gpu::LaneIdOp::create(rewriter, loc, rewriter.getIndexType(),
2295 /*upperBound=*/mlir::IntegerAttr());
2296 Value rowCount =
2297 arith::ConstantIndexOp::create(rewriter, loc, targetLaneLayout[0]);
2298 Value row = arith::RemUIOp::create(rewriter, loc, laneId, rowCount);
2299 Value element = vector::ExtractOp::create(rewriter, loc, flat,
2301 rewriter.replaceOpWithNewOp<vector::FromElementsOp>(
2302 op, redistribution->distributedTarget, element);
2303 return success();
2304 }
2305};
2306
2307/// Distributes the slice-attributed `xegpu.convert_layout` whose source is
2308/// broadcast over two lane groups onto the rows of a subset of the lanes.
2309/// Column `c` of row `r` is owned by lane `r + c * stride`, where `stride` is
2310/// the lane stride of the input's distributed dimension (see
2311/// `getDistributedDimLaneStride`), so each lane extracts its own element and
2312/// gathers the columns of its row with one `gpu.shuffle idx` per column.
2313///
2314/// The input layout has effective `lane_layout` [1, 2], so the distributed
2315/// source holds one column per lane group; the target has [n, 1], so lane `l`
2316/// keeps row `l % n` of both columns and the distributed result is
2317/// `vector<1x2>`.
2318///
2319/// xegpu.convert_layout %src
2320/// <{input_layout = #xegpu.slice<#xegpu.layout<lane_layout = [8, 1, 2],
2321/// lane_data = [4, 1, 1],
2322/// order = [0, 2, 1]>,
2323/// dims = [0]>,
2324/// target_layout = #xegpu.layout<lane_layout = [8, 1],
2325/// lane_data = [1, 1]>}>
2326/// : vector<8x2xf8E8M0FNU>
2327///
2328/// becomes, with lane `l` holding all 8 rows of column `l / 8`:
2329///
2330/// %flat = vector.shape_cast %src : vector<8x1xf8E8M0FNU> to
2331/// vector<8xf8E8M0FNU>
2332/// %lane = gpu.lane_id
2333/// %row = arith.remui %lane, %c8 : index
2334/// %own = vector.extract %flat[%row] : f8E8M0FNU from
2335/// vector<8xf8E8M0FNU>
2336/// %rowI32 = arith.index_cast %row : index to i32
2337/// %owner1 = arith.addi %rowI32, %c8_i32 : i32
2338/// %col0, %v0 = gpu.shuffle idx %own, %rowI32, %c16_i32 : f8E8M0FNU
2339/// %col1, %v1 = gpu.shuffle idx %own, %owner1, %c16_i32 : f8E8M0FNU
2340/// %res = vector.from_elements %col0, %col1 : vector<1x2xf8E8M0FNU>
2341///
2342/// The source is flattened first because `xegpu-vector-linearize` cannot
2343/// linearize a `vector.extract` with a dynamic position out of a rank-2 value.
2344///
2345/// All lanes gather every column, including the one they already hold, so that
2346/// lanes outside the target `lane_layout` hold replicas as partial lane layouts
2347/// require.
2348struct SgToLaneConvertLayoutPartialBroadcastExtractShuffle
2349 : public OpConversionPattern<xegpu::ConvertLayoutOp> {
2350 using OpConversionPattern<xegpu::ConvertLayoutOp>::OpConversionPattern;
2351
2352 LogicalResult
2353 matchAndRewrite(xegpu::ConvertLayoutOp op, OpAdaptor adaptor,
2354 ConversionPatternRewriter &rewriter) const override {
2355 xegpu::DistributeLayoutAttr inputLayout = op.getEffectiveInputLayout();
2356 xegpu::DistributeLayoutAttr targetLayout = op.getTargetLayoutAttr();
2357
2358 auto inputSlice = dyn_cast<xegpu::SliceAttr>(inputLayout);
2359 if (!inputSlice)
2360 return rewriter.notifyMatchFailure(op,
2361 "input_layout must be #xegpu.slice");
2362 if (!isa<xegpu::LayoutAttr>(targetLayout))
2363 return rewriter.notifyMatchFailure(op,
2364 "target_layout must be #xegpu.layout");
2365
2366 FailureOr<ElementLaneRedistribution> redistribution =
2367 matchElementLaneRedistribution(op, rewriter);
2368 if (failed(redistribution))
2369 return failure();
2370
2371 VectorType valueType = redistribution->valueType;
2372 Type elementType = valueType.getElementType();
2373 if (!elementType.isIntOrFloat() ||
2374 !llvm::is_contained({8u, 16u, 32u, 64u},
2375 elementType.getIntOrFloatBitWidth()))
2376 return rewriter.notifyMatchFailure(
2377 op, "element type must be an int or float of bit width 8, 16, 32 or "
2378 "64 to be carried by gpu.shuffle");
2379
2380 if (redistribution->inputLaneLayout != SmallVector<int64_t>{1, 2})
2381 return rewriter.notifyMatchFailure(
2382 op, "input_layout effective lane_layout must be [1, 2]");
2383 SmallVector<int64_t> targetLaneLayout = redistribution->targetLaneLayout;
2384 if (targetLaneLayout[0] <= 1 || targetLaneLayout[1] != 1)
2385 return rewriter.notifyMatchFailure(
2386 op, "target_layout effective lane_layout must be [n, 1] with n > 1");
2387 if (redistribution->distributedInput.getShape() !=
2388 ArrayRef<int64_t>{valueType.getDimSize(0), 1})
2389 return rewriter.notifyMatchFailure(
2390 op, "distributed input_layout type must be vector<shape[0]x1>");
2391 if (redistribution->distributedTarget.getShape() != ArrayRef<int64_t>{1, 2})
2392 return rewriter.notifyMatchFailure(
2393 op, "distributed target_layout type must be vector<1x2>");
2394
2395 FailureOr<int64_t> stride = getDistributedDimLaneStride(inputSlice);
2396 if (failed(stride))
2397 return rewriter.notifyMatchFailure(
2398 op, "input_layout parent must have exactly one non-sliced dimension "
2399 "with lane_layout extent greater than one");
2400 // The two broadcast groups must together cover the whole subgroup, so that
2401 // lane `r + c * stride` is the lane holding column `c` of row `r`.
2402 if (*stride * redistribution->inputLaneLayout[1] !=
2403 redistribution->subgroupSize)
2404 return rewriter.notifyMatchFailure(
2405 op, "input_layout distributed dimension lane stride times its extent "
2406 "must equal the subgroup size");
2407 // Lane `r + c * stride` must hold row `r`, i.e. (r + stride) %
2408 // lane_layout[0] must be r, which requires the row count to divide the
2409 // stride.
2410 if (*stride % targetLaneLayout[0] != 0)
2411 return rewriter.notifyMatchFailure(
2412 op, "target_layout effective lane_layout[0] must divide the "
2413 "input_layout distributed dimension lane stride");
2414
2415 Location loc = op.getLoc();
2416 Value src =
2417 castValueTo(rewriter, cast<TypedValue<VectorType>>(adaptor.getSource()),
2418 redistribution->distributedInput);
2419 auto flatType =
2420 VectorType::get({redistribution->distributedInput.getNumElements()},
2421 valueType.getElementType());
2422 Value flat = vector::ShapeCastOp::create(rewriter, loc, flatType, src);
2423 Value laneId = gpu::LaneIdOp::create(rewriter, loc, rewriter.getIndexType(),
2424 /*upperBound=*/mlir::IntegerAttr());
2425 Value rowCount =
2426 arith::ConstantIndexOp::create(rewriter, loc, targetLaneLayout[0]);
2427 Value row = arith::RemUIOp::create(rewriter, loc, laneId, rowCount);
2428 Value own = vector::ExtractOp::create(rewriter, loc, flat,
2430 // Gather column `c` of `row` from lane `row + c * stride`, which owns it.
2431 Type i32Type = rewriter.getI32Type();
2432 Value width = arith::ConstantIntOp::create(rewriter, loc, i32Type,
2433 redistribution->subgroupSize);
2434 Value rowI32 = arith::IndexCastOp::create(rewriter, loc, i32Type, row);
2435 Value strideI32 =
2436 arith::ConstantIntOp::create(rewriter, loc, i32Type, *stride);
2437 Value otherOwner = arith::AddIOp::create(rewriter, loc, rowI32, strideI32);
2438 Value column0 = gpu::ShuffleOp::create(rewriter, loc, own, rowI32, width,
2439 gpu::ShuffleMode::IDX)
2440 .getShuffleResult();
2441 Value column1 = gpu::ShuffleOp::create(rewriter, loc, own, otherOwner,
2442 width, gpu::ShuffleMode::IDX)
2443 .getShuffleResult();
2444 rewriter.replaceOpWithNewOp<vector::FromElementsOp>(
2445 op, redistribution->distributedTarget, ValueRange{column0, column1});
2446 return success();
2447 }
2448};
2449
2450/// Distributes the slice-attributed `xegpu.convert_layout` whose source is
2451/// fully broadcast and whose target splits the two columns of the value over
2452/// two lane groups. Each lane keeps one whole column, which is a stride-2
2453/// subset of the row-major source, so the two columns are separated with one
2454/// `vector.deinterleave` and the lane's group selects between them.
2455///
2456/// The input layout has effective `lane_layout` [1, 1], so the distributed
2457/// source is the whole value; the target has [1, 2] and its two groups tile the
2458/// subgroup, so the first half of the lanes keeps column 0 and the second half
2459/// column 1, one column per lane.
2460struct SgToLaneConvertLayoutDeinterleaveSelect
2461 : public OpConversionPattern<xegpu::ConvertLayoutOp> {
2462 using OpConversionPattern<xegpu::ConvertLayoutOp>::OpConversionPattern;
2463
2464 LogicalResult
2465 matchAndRewrite(xegpu::ConvertLayoutOp op, OpAdaptor adaptor,
2466 ConversionPatternRewriter &rewriter) const override {
2467 xegpu::DistributeLayoutAttr inputLayout = op.getEffectiveInputLayout();
2468 xegpu::DistributeLayoutAttr targetLayout = op.getTargetLayoutAttr();
2469
2470 if (!isa<xegpu::SliceAttr>(inputLayout))
2471 return rewriter.notifyMatchFailure(op,
2472 "input_layout must be #xegpu.slice");
2473 auto targetSlice = dyn_cast<xegpu::SliceAttr>(targetLayout);
2474 if (!targetSlice)
2475 return rewriter.notifyMatchFailure(op,
2476 "target_layout must be #xegpu.slice");
2477
2478 FailureOr<ElementLaneRedistribution> redistribution =
2479 matchElementLaneRedistribution(op, rewriter);
2480 if (failed(redistribution))
2481 return failure();
2482
2483 VectorType valueType = redistribution->valueType;
2484 if (redistribution->inputLaneLayout != SmallVector<int64_t>{1, 1})
2485 return rewriter.notifyMatchFailure(
2486 op, "input_layout effective lane_layout must be [1, 1]");
2487 SmallVector<int64_t> targetLaneLayout = redistribution->targetLaneLayout;
2488 if (targetLaneLayout != SmallVector<int64_t>{1, 2})
2489 return rewriter.notifyMatchFailure(
2490 op, "target_layout effective lane_layout must be [1, 2]");
2491 if (redistribution->distributedInput != valueType)
2492 return rewriter.notifyMatchFailure(
2493 op, "distributed input_layout type must equal the value type");
2494 if (redistribution->distributedTarget.getShape() !=
2495 ArrayRef<int64_t>{valueType.getDimSize(0), 1})
2496 return rewriter.notifyMatchFailure(
2497 op, "distributed target_layout type must be vector<shape[0]x1>");
2498
2499 FailureOr<int64_t> stride = getDistributedDimLaneStride(targetSlice);
2500 if (failed(stride))
2501 return rewriter.notifyMatchFailure(
2502 op, "target_layout parent must have exactly one non-sliced dimension "
2503 "with lane_layout extent greater than one");
2504 // The two lane groups must together cover the whole subgroup, so that
2505 // `lane_id / stride` is the index of the column the lane keeps.
2506 if (*stride * targetLaneLayout[1] != redistribution->subgroupSize)
2507 return rewriter.notifyMatchFailure(
2508 op, "target_layout distributed dimension lane stride times its "
2509 "extent must equal the subgroup size");
2510
2511 Location loc = op.getLoc();
2512 Value src =
2513 castValueTo(rewriter, cast<TypedValue<VectorType>>(adaptor.getSource()),
2514 redistribution->distributedInput);
2515 auto flatType = VectorType::get({valueType.getNumElements()},
2516 valueType.getElementType());
2517 Value flat = vector::ShapeCastOp::create(rewriter, loc, flatType, src);
2518 auto deinterleaved = vector::DeinterleaveOp::create(rewriter, loc, flat);
2519 Value laneId = gpu::LaneIdOp::create(rewriter, loc, rewriter.getIndexType(),
2520 /*upperBound=*/mlir::IntegerAttr());
2521 Value strideVal = arith::ConstantIndexOp::create(rewriter, loc, *stride);
2522 Value zero = arith::ConstantIndexOp::create(rewriter, loc, 0);
2523 Value group = arith::DivUIOp::create(rewriter, loc, laneId, strideVal);
2524 Value isFirstGroup = arith::CmpIOp::create(
2525 rewriter, loc, arith::CmpIPredicate::eq, group, zero);
2526 Value selected = arith::SelectOp::create(rewriter, loc, isFirstGroup,
2527 deinterleaved.getRes1(),
2528 deinterleaved.getRes2());
2529 rewriter.replaceOpWithNewOp<vector::ShapeCastOp>(
2530 op, redistribution->distributedTarget, selected);
2531 return success();
2532 }
2533};
2534
2535// Trivially distribute `vector.interleave`
2536struct SgToLaneVectorInterleave
2537 : public OpConversionPattern<vector::InterleaveOp> {
2538 using OpConversionPattern<vector::InterleaveOp>::OpConversionPattern;
2539
2540 LogicalResult
2541 matchAndRewrite(vector::InterleaveOp op, OpAdaptor adaptor,
2542 ConversionPatternRewriter &rewriter) const override {
2543
2544 auto newOp = vector::InterleaveOp::create(
2545 rewriter, op.getLoc(), adaptor.getLhs(), adaptor.getRhs());
2546 rewriter.replaceOp(op, newOp.getResult());
2547 return success();
2548 }
2549};
2550
2551// Trivially distribute `vector.deinterleave`
2552struct SgToLaneVectorDeinterleave
2553 : public OpConversionPattern<vector::DeinterleaveOp> {
2554 using OpConversionPattern<vector::DeinterleaveOp>::OpConversionPattern;
2555
2556 LogicalResult
2557 matchAndRewrite(vector::DeinterleaveOp op, OpAdaptor adaptor,
2558 ConversionPatternRewriter &rewriter) const override {
2559
2560 auto newOp = vector::DeinterleaveOp::create(rewriter, op.getLoc(),
2561 adaptor.getSource());
2562 rewriter.replaceOp(op, newOp.getResults());
2563 return success();
2564 }
2565};
2566
2567struct SgToLaneDpasMx : public OpConversionPattern<xegpu::DpasMxOp> {
2568 using OpConversionPattern<xegpu::DpasMxOp>::OpConversionPattern;
2569
2570 LogicalResult
2571 matchAndRewrite(xegpu::DpasMxOp op, OpAdaptor adaptor,
2572 ConversionPatternRewriter &rewriter) const override {
2573 const auto *uArch =
2574 xegpu::uArch::getUArch(xegpu::getChipStr(op).value_or(""));
2575 if (!uArch)
2576 return failure();
2577 if (!uArch->isSupportedInstruction(
2579 return rewriter.notifyMatchFailure(
2580 op, "target uArch does not support scaled subgroup mma");
2581 // Check if the op has A, B and CD layouts attached.
2582 auto layoutA = cast<xegpu::LayoutAttr>(op.getLayoutAAttr());
2583 auto layoutB = cast<xegpu::LayoutAttr>(op.getLayoutBAttr());
2584 auto layoutCd = cast<xegpu::LayoutAttr>(op.getLayoutCdAttr());
2585 if (!layoutA || !layoutB || !layoutCd)
2586 return rewriter.notifyMatchFailure(
2587 op, "missing required layout attributes for DpasMxOp distribution");
2588
2589 // Retrieve expected types, according to anchor layouts.
2590 auto expected1DTypeResult =
2591 xegpu::getDistributedVectorType(op.getType(), layoutCd);
2592 auto expected1DTypeA =
2593 xegpu::getDistributedVectorType(op.getA().getType(), layoutA);
2594 auto expected1DTypeB =
2595 xegpu::getDistributedVectorType(op.getB().getType(), layoutB);
2596
2597 VectorType expected1DTypeScaleA, expected1DTypeScaleB;
2598 if (op.getScaleA()) {
2599 auto layoutScaleA = cast<xegpu::LayoutAttr>(op.getLayoutAScaleAttr());
2600 auto expected1DTypeScaleAOrFailure = xegpu::getDistributedVectorType(
2601 cast<VectorType>(op.getScaleA().getType()), layoutScaleA);
2602 if (failed(expected1DTypeScaleAOrFailure))
2603 return rewriter.notifyMatchFailure(
2604 op, "failed to calculate expected 1D vector type for scale A");
2605 expected1DTypeScaleA = expected1DTypeScaleAOrFailure.value();
2606 }
2607 if (op.getScaleB()) {
2608 auto layoutScaleB = cast<xegpu::LayoutAttr>(op.getLayoutBScaleAttr());
2609 auto expected1DTypeScaleBOrFailure = xegpu::getDistributedVectorType(
2610 cast<VectorType>(op.getScaleB().getType()), layoutScaleB);
2611 if (failed(expected1DTypeScaleBOrFailure))
2612 return rewriter.notifyMatchFailure(
2613 op, "failed to calculate expected 1D vector type for scale B");
2614 expected1DTypeScaleB = expected1DTypeScaleBOrFailure.value();
2615 }
2616
2617 auto expectedNDTypeResult =
2618 xegpu::getDistVecTypeBasedOnLaneLayout(layoutCd, op.getType());
2619 if (failed(expected1DTypeResult) || failed(expected1DTypeA) ||
2620 failed(expected1DTypeB))
2621 return rewriter.notifyMatchFailure(
2622 op,
2623 "failed to calculate supported workitem 1D vector types for DpasOp "
2624 "from layouts");
2625 if (failed(expectedNDTypeResult))
2626 return rewriter.notifyMatchFailure(
2627 op, "unable to compute expected workitem vector type for DpasOp from "
2628 "lane layout");
2629
2630 // Validate bit widths match uArch packed format requirements
2631 const auto *uArchInstruction = dyn_cast<
2632 xegpu::uArch::SubgroupScaledMatrixMultiplyAcc>(uArch->getInstruction(
2634 assert(uArchInstruction);
2635 auto wiAType = expected1DTypeA.value();
2636 auto wiBType = expected1DTypeB.value();
2637 // Calculate total packed bit width = element bit width * vector size
2638 unsigned aPackedBitWidth =
2639 wiAType.getElementTypeBitWidth() * wiAType.getNumElements();
2640 unsigned bPackedBitWidth =
2641 wiBType.getElementTypeBitWidth() * wiBType.getNumElements();
2642 if (aPackedBitWidth % uArchInstruction->getPackedFormatBitSizeA())
2643 return rewriter.notifyMatchFailure(
2644 op, "A operand packed bit width must be a multiple of uArch packed "
2645 "format requirement");
2646 if (bPackedBitWidth % uArchInstruction->getPackedFormatBitSizeB())
2647 return rewriter.notifyMatchFailure(
2648 op, "B operand packed bit width must be a multiple of uArch packed "
2649 "format requirement");
2650
2651 auto newOp = xegpu::DpasMxOp::create(
2652 rewriter, op->getLoc(), expected1DTypeResult.value(),
2653 castValueTo(rewriter, cast<TypedValue<VectorType>>(adaptor.getA()),
2654 expected1DTypeA.value()),
2655 castValueTo(rewriter, cast<TypedValue<VectorType>>(adaptor.getB()),
2656 expected1DTypeB.value()),
2657 op.getAcc()
2658 ? castValueTo(rewriter,
2659 cast<TypedValue<VectorType>>(adaptor.getAcc()),
2660 expected1DTypeResult.value())
2661 : nullptr,
2662
2663 op.getScaleA()
2664 ? castValueTo(rewriter,
2665 cast<TypedValue<VectorType>>(adaptor.getScaleA()),
2666 expected1DTypeScaleA)
2667 : nullptr,
2668 op.getScaleB()
2669 ? castValueTo(rewriter,
2670 cast<TypedValue<VectorType>>(adaptor.getScaleB()),
2671 expected1DTypeScaleB)
2672 : nullptr,
2673 /** layoutA**/ nullptr,
2674 /** layoutB**/ nullptr, /** layoutCd**/ nullptr,
2675 /** layoutAScale**/ nullptr, /** layoutBScale**/ nullptr);
2676 // Explicitly set the new types to enable correct type materializations.
2677 rewriter.replaceOp(op, castValueTo(rewriter, newOp.getResult(),
2678 expectedNDTypeResult.value()));
2679 return success();
2680 }
2681};
2682
2683struct XeGPUSgToLaneDistributePass
2684 : public xegpu::impl::XeGPUSgToLaneDistributeBase<
2685 XeGPUSgToLaneDistributePass> {
2686 void runOnOperation() override;
2687};
2688
2689} // namespace
2690
2691void XeGPUSgToLaneDistributePass::runOnOperation() {
2692
2693 // Recover temporary operand layouts for usage in patterns.
2694 Operation *root = getOperation();
2695 if (!xegpu::recoverTemporaryLayouts(root)) {
2696 signalPassFailure();
2697 return;
2698 }
2699
2700 // Collect existing UnrealizedConversionCastOps. These must be preserved.
2701 llvm::SmallSetVector<UnrealizedConversionCastOp, 8> existingCasts;
2702 root->walk(
2703 [&](UnrealizedConversionCastOp castOp) { existingCasts.insert(castOp); });
2704 // Perform a structural type conversion to convert structural ops to have WI
2705 // types. This will insert UnrealizedConversionCastOps to make the IR
2706 // valid.
2707 {
2708 ConversionTarget target(getContext());
2709 TypeConverter typeConverter;
2710 RewritePatternSet patterns(&getContext());
2711 // Source (N:1) and target (1:1) materializations using
2712 // UnrealizedConversionCastOp.
2713 auto materializeCast = [](OpBuilder &builder, Type type, ValueRange inputs,
2714 Location loc) -> Value {
2715 return UnrealizedConversionCastOp::create(builder, loc, type, inputs)
2716 .getResult(0);
2717 };
2718 typeConverter.addSourceMaterialization(materializeCast);
2719 typeConverter.addTargetMaterialization(materializeCast);
2722 patterns, target);
2724 typeConverter, patterns, target, root);
2725 target.addLegalOp<UnrealizedConversionCastOp>();
2726 if (failed(applyPartialConversion(root, target, std::move(patterns))))
2727 return signalPassFailure();
2728 }
2729 // Fold cancelling cast chains and erase dead casts.
2730 xegpu::cleanupUnrealizedConversionCasts(root, existingCasts);
2731 xegpu::removeTemporaryLayoutAttrs(getOperation());
2732}
2733
2735 TypeConverter &typeConverter, Operation *topLevelOp) {
2736 // Pass through any type by default; more specific conversions registered
2737 // below override this for TensorDescType and (distributing) VectorType.
2738 typeConverter.addConversion([](Type type) -> Type { return type; });
2739 // For TensorDescType, drop the layout attribute if any.
2740 typeConverter.addConversion([](TensorDescType type) -> Type {
2741 if (type.getLayoutAttr()) {
2742 return type.dropLayouts();
2743 }
2744 return type;
2745 });
2746 // For VectorType, distribute based on the lane layout (1:1 shape-changing
2747 // conversion). Uses xegpu::addVectorTypeConversion with a pre-computed
2748 // map for SCF loop block args (see precomputeLoopBlockArgTypes for the
2749 // rationale).
2750 auto getSubShapeAndCount = [](VectorType vecTy,
2751 xegpu::DistributeLayoutAttr layout)
2752 -> std::pair<SmallVector<int64_t>, int> {
2753 auto distTyOrFailure = getDistVecTypeBasedOnLaneLayout(layout, vecTy);
2754 if (failed(distTyOrFailure))
2755 return {{}, 0};
2756 return {SmallVector<int64_t>(distTyOrFailure->getShape()), 1};
2757 };
2758 auto loopArgTypes =
2759 xegpu::precomputeLoopBlockArgTypes(topLevelOp, getSubShapeAndCount);
2760 xegpu::addVectorTypeConversion(typeConverter, getSubShapeAndCount,
2761 std::move(loopArgTypes));
2762}
2763
2765 TypeConverter &typeConverter, RewritePatternSet &patterns,
2766 ConversionTarget &target, Operation *topLevelOp) {
2767 populateXeGPUSgToLaneDistributeTypeConversions(typeConverter, topLevelOp);
2768 // CreateNdDescOp is legal only if its result type has no layout attribute.
2769 target.addDynamicallyLegalOp<xegpu::CreateNdDescOp>(
2770 [&](xegpu::CreateNdDescOp op) { return !op.getType().getLayoutAttr(); });
2771 // Any anchor XeGPU op is legal only if it has no anchor layout.
2772 target.addDynamicallyLegalDialect<xegpu::XeGPUDialect>([](Operation *op) {
2773 if (isa<xegpu::ConvertLayoutOp>(op))
2774 return false;
2775 auto anchorOp = dyn_cast<AnchorLayoutInterface>(op);
2776 if (!anchorOp)
2777 return true;
2778 return !anchorOp.getAnchorLayout();
2779 });
2780 // Arith constants are legal only if they have no temporary layout attribute.
2781 target.addDynamicallyLegalOp<arith::ConstantOp>(
2782 [=](arith::ConstantOp op) -> bool {
2783 // If the result type is not a vector, it's legal.
2784 if (!isa<VectorType>(op.getResult().getType()))
2785 return true;
2786 return !xegpu::getTemporaryLayout(dyn_cast<OpResult>(op.getResult()));
2787 });
2788 // In math and arith dialects, only handle elementwise ops with a single
2789 // result and with a result layout attribute.
2790 target.addDynamicallyLegalDialect<math::MathDialect, arith::ArithDialect>(
2791 [=](Operation *op) -> std::optional<bool> {
2792 // Only handle elementwise mappable ops
2794 return true;
2795 // Only handle ops with single vector result
2796 if (op->getNumResults() != 1)
2797 return true;
2798
2799 VectorType resultType =
2800 dyn_cast<VectorType>(op->getResult(0).getType());
2801 if (!resultType)
2802 return true;
2803
2804 // Check if all operands are vectors of the same shape
2805 for (Value operand : op->getOperands()) {
2806 VectorType operandType = dyn_cast<VectorType>(operand.getType());
2807 if (!operandType || operandType.getShape() != resultType.getShape()) {
2808 return true;
2809 }
2810 }
2811 return !xegpu::getTemporaryLayout(dyn_cast<OpResult>(op->getResult(0)));
2812 });
2813 // vector::ReductionOp is legal only if its source has no distribute layout
2814 // attribute.
2815 target.addDynamicallyLegalOp<vector::ReductionOp>(
2816 [=](vector::ReductionOp op) -> bool {
2817 auto layout = xegpu::getDistributeLayoutAttr(op.getVector());
2818 return !layout;
2819 });
2820 // vector::MultiDimReductionOp op legality.
2821 target.addDynamicallyLegalOp<vector::MultiDimReductionOp>(
2822 [=](vector::MultiDimReductionOp op) -> bool {
2823 return !isValidSubgroupMultiReductionOp(op);
2824 });
2825 target.addDynamicallyLegalOp<vector::CreateMaskOp, vector::ConstantMaskOp,
2826 vector::TransposeOp, vector::BitCastOp,
2827 vector::ShapeCastOp, vector::StepOp,
2828 vector::BroadcastOp>([=](Operation *op) -> bool {
2829 return !xegpu::getTemporaryLayout(op->getOpResult(0));
2830 });
2831 target.addDynamicallyLegalOp<vector::ExtractOp>(
2832 [=](vector::ExtractOp op) -> bool {
2833 if (!isa<VectorType>(op.getType()))
2834 return true;
2835 return !xegpu::getTemporaryLayout(op->getOpResult(0));
2836 });
2837 target.addDynamicallyLegalOp<vector::InsertOp>(
2838 [=](vector::InsertOp op) -> bool {
2839 return !xegpu::getTemporaryLayout(op->getOpResult(0));
2840 });
2841 target.addDynamicallyLegalOp<vector::ExtractStridedSliceOp>(
2842 [=](vector::ExtractStridedSliceOp op) -> bool {
2843 return !xegpu::getTemporaryLayout(op->getOpResult(0));
2844 });
2845 target.addDynamicallyLegalOp<vector::InsertStridedSliceOp>(
2846 [=](vector::InsertStridedSliceOp op) -> bool {
2847 return !xegpu::getTemporaryLayout(op->getOpResult(0));
2848 });
2849 target.addDynamicallyLegalOp<vector::InterleaveOp, vector::DeinterleaveOp>(
2850 [=](Operation *op) -> bool {
2851 return !xegpu::getTemporaryLayout(op->getOpResult(0));
2852 });
2853 target.markUnknownOpDynamicallyLegal([](Operation *op) { return true; });
2854 patterns
2855 .add<SgToLaneCreateNdDesc, SgToLaneLoadNd, SgToLaneStoreNd, SgToLaneDpas,
2856 SgToLaneElementWise, SgToLaneArithConstant, SgToLanePrefetchNd,
2857 SgToLaneLoadGather, SgToLaneStoreScatter, SgToLaneVectorReduction,
2858 SgToLaneMultiDimReduction, SgToLaneVectorExtract,
2859 SgToLaneVectorInsert, SgToLaneVectorExtractStridedSlice,
2860 SgToLaneVectorInsertStridedSlice, SgToLaneLoadMatrix,
2861 SgToLaneStoreMatrix, SgToLaneConvertLayout, SgToLaneVectorTranspose,
2862 SgToLaneVectorBitcast, SgToLaneVectorStep, SgToLaneVectorShapeCast,
2863 SgToLaneBroadcast, SgToLaneCreateMask<vector::CreateMaskOp>,
2864 SgToLaneCreateMask<vector::ConstantMaskOp>,
2865 SgToLaneVectorDeinterleave, SgToLaneVectorInterleave, SgToLaneDpasMx,
2866 SgToLaneConvertLayoutBroadcastExtract,
2867 SgToLaneConvertLayoutPartialBroadcastExtractShuffle,
2868 SgToLaneConvertLayoutDeinterleaveSelect>(typeConverter,
2869 patterns.getContext());
2870 // The SLM round-trip handles any pair of layouts but pays for a memory
2871 // round-trip, so it only runs when every pattern above has declined.
2872 patterns.add<SgToLaneConvertLayoutViaSLM>(typeConverter,
2873 patterns.getContext(),
2874 /*benefit=*/PatternBenefit(0));
2875}
return success()
b
Return true if permutation is a valid permutation of the outer_dims_perm (case OuterOrInnerPerm::Oute...
b getContext())
Attributes are known-constant values of operations.
Definition Attributes.h:25
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 is the top-level object for a collection of MLIR operations.
Definition MLIRContext.h:63
Operation is the basic unit of execution within MLIR.
Definition Operation.h:87
OpResult getResult(unsigned idx)
Get the 'idx'th result of this operation.
Definition Operation.h:432
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.
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
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 the benefit of a pattern match in a unitless scheme that ranges from 0 (very li...
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.
A range-style iterator that allows for iterating over the offsets of all potential tiles of size tile...
This class provides an abstraction over the various different ranges of value types.
Definition TypeRange.h:40
Instances of the Type class are uniqued, have an immutable identifier and an optional mutable compone...
Definition Types.h:74
bool isIntOrFloat() const
Return true if this is an integer (of any signedness) or a float type.
Definition Types.cpp:118
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
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
static ConstantIntOp create(OpBuilder &builder, Location location, int64_t value, unsigned width)
Definition ArithOps.cpp:297
static DenseArrayAttrImpl get(MLIRContext *context, ArrayRef< int64_t > content)
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.
SmallVector< Value > getAsValues(OpBuilder &builder, Location loc, ArrayRef< OpFoldResult > foldResults)
Convert foldResults into Values.
const uArch * getUArch(llvm::StringRef archName)
Definition uArchCommon.h:24
void populateXeGPUSgToLaneDistributeTypeConversionAndLegality(TypeConverter &typeConverter, RewritePatternSet &patterns, ConversionTarget &target, Operation *topLevelOp)
Defines type conversions and legality for XeGPU subgroup to lane distribution and appends the require...
bool requirePacked(const DistributeLayoutAttr layout)
Helper function to check if the layout is packed.
void removeTemporaryLayoutAttrs(Operation *op)
Removes the temporary layout attributes for each OpOperand and OpResult of the given operation.
Value subgroupReduction(Location loc, OpBuilder &builder, Value input, vector::CombiningKind kind, uint32_t size)
Given an input value representing per-lane data, this function returns the result after performing a ...
bool recoverTemporaryLayouts(Operation *rootOp)
Attach layout attributes to all vector-type operands of operations within the given operation's neste...
FailureOr< int64_t > getNumSubgroupsFromBlockSize(Operation *op, int64_t subgroupSize)
Returns the number of subgroups the kernel enclosing op runs, derived from the known_block_size of it...
FailureOr< VectorType > getDistVecTypeBasedOnLaneLayout(DistributeLayoutAttr layout, VectorType originalType)
Helper function to get distributed vector type for a source vector type according to the lane_layout.
Value lowerToVectorReductions(TypedValue< VectorType > src, TypedValue< VectorType > acc, vector::CombiningKind kind, int64_t reductionDim, Location loc, PatternRewriter &rewriter)
Given a src and an acc argumments from a vector::MultiDimReductionOp, lower to a set of vector::Reduc...
bool requireTranspose(const DistributeLayoutAttr layout, const uArch::uArch *uArch)
Helper function to check if the layout requires a transpose effect.
DistributeLayoutAttr getDistributeLayoutAttr(const Value value)
Retrieves the DistributeLayoutAttr associated with a given Value, or nullptr if none is found.
DenseMap< Value, SmallVector< Type > > precomputeLoopBlockArgTypes(Operation *topLevelOp, SubShapeAndCountFn getSubShapeAndCount)
Pre-computes distributed VectorType mappings for every value carried through an SCF loop under topLev...
std::optional< std::string > getChipStr(Operation *op)
Retrieves the chip string from the XeVM target attribute of the parent GPU module operation.
void addVectorTypeConversion(TypeConverter &converter, SubShapeAndCountFn getSubShapeAndCount, DenseMap< Value, SmallVector< Type > > loopArgTypes)
Adds a context-aware VectorType conversion to converter (1:1 shape-changing or 1:N,...
DistributeLayoutAttr getTemporaryLayout(const T &operandOrResult)
get and set distribute layout attribute for non-anchor operations (and offsets/masks of load/store op...
void populateXeGPUSgToLaneDistributeTypeConversions(TypeConverter &typeConverter, Operation *topLevelOp)
Define only the type conversions needed for XeGPU subgroup to lane distribution.
Value lowerCrossLaneReductionToShuffles(TypedValue< VectorType > src, TypedValue< VectorType > acc, vector::CombiningKind kind, int64_t reductionDim, int64_t reductionSize, Location loc, PatternRewriter &rewriter)
Lowers cross-lane reductions to shuffle operations on a 2D vector.
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.
FailureOr< VectorType > getDistributedVectorType(xegpu::TensorDescType tdescTy)
If tensor descriptor has a layout attribute it is used in SIMT mode.
Include the generated interface declarations.
detail::DenseArrayAttrImpl< int64_t > DenseI64ArrayAttr
std::optional< int64_t > getConstantIntValue(OpFoldResult ofr)
If ofr is a constant integer or an IntegerAttr, return the integer.
SmallVector< int64_t > computeStrides(ArrayRef< int64_t > sizes)
SmallVector< int64_t > delinearize(int64_t linearIndex, ArrayRef< int64_t > strides)
Given the strides together with a linear index in the dimension space, return the vector-space offset...
int64_t computeProduct(ArrayRef< int64_t > basis)
Self-explicit.
std::conditional_t< std::is_same_v< Ty, mlir::Type >, mlir::Value, detail::TypedValue< Ty > > TypedValue
If Ty is mlir::Type this will select Value instead of having a wrapper around it.
Definition Value.h:494
OpFoldResult getAsOpFoldResult(Value val)
Given a value, try to extract a constant Attribute.
int64_t linearize(ArrayRef< int64_t > offsets, ArrayRef< int64_t > basis)
Return the linearized index of 'offsets' w.r.t.
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.
This represents an operation in an abstracted form, suitable for use with the builder APIs.
void addOperands(ValueRange newOperands)
void addAttribute(StringRef name, Attribute attr)
Add an attribute with the specified name.
void addTypes(ArrayRef< Type > newTypes)
Attribute propertiesAttr
This Attribute is used to opaquely construct the properties of the operation.