[Mlir-commits] [mlir] [mlir][xegpu] Add vector layout conflict handling in XeGPU layout propagation pass. (PR #182402)
Artem Kroviakov
llvmlistbot at llvm.org
Sun Feb 22 04:38:49 PST 2026
================
@@ -1232,6 +1233,115 @@ namespace {
//===----------------------------------------------------------------------===//
// ResolveLayoutConflicts
//===----------------------------------------------------------------------===//
+
+/// Helper to get the defining CreateNdDescOp of a tensor descriptor value. This
+/// function tries to find the defining CreateNdDescOp recursively accross
+/// control-flow boundaries.
+static xegpu::CreateNdDescOp getDefiningCreateNdDescOp(Value tdescValue) {
+ // Try to get the defining CreateNdDescOp of the tensor descriptor.
+ auto definingOp = tdescValue.getDefiningOp<xegpu::CreateNdDescOp>();
+ if (definingOp)
+ return definingOp;
+ // If tdescValue is an argument, try to get the tied init value from the
+ // parent loop-like op.
+ if (auto arg = dyn_cast<BlockArgument>(tdescValue)) {
+ auto *parentOp = arg.getOwner()->getParentOp();
+ if (auto loop = dyn_cast<LoopLikeOpInterface>(parentOp)) {
+ OpOperand *tiedInit = loop.getTiedLoopInit(arg);
+ if (tiedInit)
+ return getDefiningCreateNdDescOp(tiedInit->get());
+ }
+ }
+ // If not found, return null.
+ return nullptr;
+}
+
+static xegpu::DistributeLayoutAttr
+getExpectedLayoutAt(OpOperand &operand,
+ xegpu::DistributeLayoutAttr currLayout) {
+ Operation *op = operand.getOwner();
+ unsigned idx = operand.getOperandNumber();
+
+ // For vector::BroadcastOp, infer the source layout from the result layout.
+ if (auto broadcast = dyn_cast<vector::BroadcastOp>(op)) {
+ auto resLayout = xegpu::getDistributeLayoutAttr(broadcast->getResult(0));
+ if (!resLayout)
+ return xegpu::DistributeLayoutAttr();
+ auto srcTy = dyn_cast<VectorType>(broadcast.getSourceType());
+ if (!srcTy)
+ return xegpu::DistributeLayoutAttr();
+ return xegpu::inferBroadcastSourceLayout(
+ resLayout, broadcast.getResultVectorType().getShape(),
+ srcTy.getShape());
+ }
+
+ // For vector::MultiDimReductionOp, infer source layout from result layout
+ // using reduction dims. Acc operand is expected to have the same layout as
+ // the result.
+ if (auto reduction = dyn_cast<vector::MultiDimReductionOp>(op)) {
+ auto resLayout = xegpu::getDistributeLayoutAttr(reduction->getResult(0));
+ if (!resLayout)
+ return xegpu::DistributeLayoutAttr();
+ if (idx == 0) {
+ SmallVector<int64_t> reductionDims(reduction.getReductionDims());
+ return xegpu::inferMultiReductionSourceLayout(resLayout, reductionDims);
+ }
+ if (idx == 1)
+ return resLayout;
+ }
+
+ // For vector::BitCastOp, infer source layout from result layout using
+ // element type bitwidths.
+ if (auto bitcast = dyn_cast<vector::BitCastOp>(op)) {
+ auto resLayout = xegpu::getDistributeLayoutAttr(bitcast->getResult(0));
+ if (!resLayout)
+ return xegpu::DistributeLayoutAttr();
+ int resElemBitWidth =
+ bitcast.getResultVectorType().getElementType().getIntOrFloatBitWidth();
+ int srcElemBitWidth =
+ bitcast.getSourceVectorType().getElementType().getIntOrFloatBitWidth();
+ return xegpu::inferBitCastSourceLayout(resLayout, resElemBitWidth,
+ srcElemBitWidth);
+ }
+
+ // For vector::ShapeCastOp, infer source layout from result layout using
+ // shapes.
+ if (auto shapeCast = dyn_cast<vector::ShapeCastOp>(op)) {
+ auto resLayout = xegpu::getDistributeLayoutAttr(shapeCast->getResult(0));
+ if (!resLayout)
+ return xegpu::DistributeLayoutAttr();
+ return xegpu::inferShapeCastSourceLayout(
+ resLayout, shapeCast.getResultVectorType().getShape(),
+ shapeCast.getSourceVectorType().getShape());
+ }
+
+ // For vector::InsertStridedSliceOp, infer source layout from result layout.
+ // Dest vector must have the same layout as the result.
+ if (auto insertSlice = dyn_cast<vector::InsertStridedSliceOp>(op)) {
+ auto resLayout = xegpu::getDistributeLayoutAttr(insertSlice->getResult(0));
+ if (!resLayout)
+ return xegpu::DistributeLayoutAttr();
+ if (idx == 0)
+ return xegpu::inferInsertStridedSliceSourceLayout(
+ resLayout, insertSlice.getDestVectorType().getShape(),
+ insertSlice.getSourceVectorType().getShape());
+ if (idx == 1)
+ return resLayout;
+ }
+ // For elementwise operations, all operands must have the same layout as the
+ // result.
+ if (OpTrait::hasElementwiseMappableTraits(op) && op->getNumResults() == 1) {
----------------
akroviakov wrote:
What about outlining the `resLayout` retrieval at the top and checking if an op with a vector result has a layout, and if not, `return xegpu::DistributeLayoutAttr()`? This would save a few lines from each branch
https://github.com/llvm/llvm-project/pull/182402
More information about the Mlir-commits
mailing list