[Mlir-commits] [mlir] [MLIR][XeGPU] Add distribution pattern for xegpu.load & store for sg to wi pass (PR #181917)
Charitha Saumya
llvmlistbot at llvm.org
Fri Feb 20 13:25:11 PST 2026
================
@@ -522,6 +594,80 @@ struct LowerVectorMultiReductionPattern
}
};
+/// Distributes a subgroup-level StoreScatter (xegpu.store) op to
+/// workitem-level.
+struct SgToWiStoreScatter : public OpConversionPattern<xegpu::StoreScatterOp> {
+ using OpConversionPattern<xegpu::StoreScatterOp>::OpConversionPattern;
+
+ LogicalResult
+ matchAndRewrite(xegpu::StoreScatterOp op, OpAdaptor adaptor,
+ ConversionPatternRewriter &rewriter) const override {
+ xegpu::DistributeLayoutAttr layout = op.getAnchorLayout();
+ if (!layout)
+ return failure();
+
+ VectorType valueTy = op.getValueType();
+ if (!valueTy)
+ return failure();
+
+ // Check that all leading dimensions are unit dimensions.
+ int chunkSize = op.getChunkSize().value_or(1);
+ int effectiveVecRank = (chunkSize == 1) ? 1 : 2;
+ for (int i = 0; i < valueTy.getRank() - effectiveVecRank; i++) {
+ if (valueTy.getShape()[i] != 1)
+ return rewriter.notifyMatchFailure(
+ op, "Only unit dimensions allowed for the leading "
+ "dimensions of the store vector!");
+ }
+
+ auto expectedWiValueTyOrFailure =
+ xegpu::getDistVecTypeBasedOnLaneLayout(layout, valueTy);
+ if (failed(expectedWiValueTyOrFailure))
+ return rewriter.notifyMatchFailure(
+ op,
+ "unable to compute expected workitem vector type from lane layout");
+
+ VectorType expectedWiValueTy = expectedWiValueTyOrFailure.value();
+ VectorType supportedWiValueTy =
+ VectorType::get({expectedWiValueTy.getNumElements()},
+ expectedWiValueTy.getElementType());
+
+ Value adaptedValue = adaptor.getValue();
+ if (adaptedValue.getType() != supportedWiValueTy)
+ adaptedValue =
+ vector::ShapeCastOp::create(rewriter, op.getLoc(), supportedWiValueTy,
+ adaptedValue)
+ .getResult();
----------------
charithaintc wrote:
use `castValueTo`
https://github.com/llvm/llvm-project/pull/181917
More information about the Mlir-commits
mailing list