23#include "llvm/ADT/STLExtras.h"
24#include "llvm/ADT/SetVector.h"
25#include "llvm/Support/DebugLog.h"
29#define GEN_PASS_DEF_XEGPUBLOCKING
30#include "mlir/Dialect/XeGPU/Transforms/Passes.h.inc"
34#define DEBUG_TYPE "xegpu-blocking"
48class XeGPUBlockingPass final
49 :
public xegpu::impl::XeGPUBlockingBase<XeGPUBlockingPass> {
51 void runOnOperation()
override;
58 typename = std::enable_if_t<std::is_same_v<T, OpOperand> ||
59 std::is_same_v<T, OpResult>>>
60 std::optional<SmallVector<int64_t>>
73template <
typename T,
typename>
74std::optional<SmallVector<int64_t>>
75XeGPUBlockingPass::getTileShape(
const T &operandOrResult)
const {
77 if constexpr (std::is_same_v<T, OpOperand>) {
78 value = operandOrResult.get();
80 value = (Value)operandOrResult;
83 xegpu::DistributeLayoutAttr layout =
85 if (layout && layout.isForSubgroup()) {
86 if (!layout.getEffectiveInstDataAsInt().empty()) {
87 SmallVector<int64_t> instData = layout.getEffectiveInstDataAsInt();
90 if (
auto type = dyn_cast<ShapedType>(value.
getType()))
91 return llvm::to_vector(type.getShape());
93 LDBG() <<
"failed to getTileShape for: " << value;
97std::optional<SmallVector<int64_t>>
98XeGPUBlockingPass::getTileShape(Operation *op)
const {
99 if (isa<xegpu::CreateNdDescOp, xegpu::LoadMatrixOp>(op))
101 if (isa<xegpu::PrefetchNdOp, xegpu::LoadNdOp, xegpu::PrefetchOp,
102 xegpu::StoreMatrixOp>(op))
104 if (isa<xegpu::StoreNdOp>(op))
107 if (isa<xegpu::LoadGatherOp>(op))
110 if (
auto convertLayoutOp = dyn_cast<xegpu::ConvertLayoutOp>(op)) {
112 convertLayoutOp.getEffectiveInputLayout().getEffectiveInstDataAsInt();
113 auto targetInstData =
114 convertLayoutOp.getTargetLayout().getEffectiveInstDataAsInt();
115 assert(inputInstData.size() == targetInstData.size() &&
116 "convert_layout layouts must both carry inst_data of the same rank");
117 SmallVector<int64_t>
tile(inputInstData.size());
118 for (
size_t i = 0; i <
tile.size(); ++i)
119 tile[i] = std::max(inputInstData[i], targetInstData[i]);
123 if (isa<xegpu::StoreScatterOp>(op))
127 auto validateABTiles = [&](Operation *op)
128 -> std::optional<std::pair<SmallVector<int64_t>, SmallVector<int64_t>>> {
129 std::optional<SmallVector<int64_t>> aTile =
131 std::optional<SmallVector<int64_t>> bTile =
134 if (!aTile || aTile->size() < 2 || !bTile || bTile->size() < 2)
138 int64_t aBatchRank = aTile->size() - 2;
139 int64_t bBatchRank = bTile->size() - 2;
140 if (aBatchRank != bBatchRank)
144 for (int64_t i = 0; i < aBatchRank; ++i) {
145 if ((*aTile)[i] != (*bTile)[i])
151 if ((*aTile).back() != (*bTile)[bBatchRank])
154 return std::make_pair(*aTile, *bTile);
158 auto validateCTile = [&](Operation *op,
unsigned cOperandIdx,
159 const SmallVector<int64_t> &aTile,
160 const SmallVector<int64_t> &bTile) ->
bool {
164 std::optional<SmallVector<int64_t>> cTile =
169 int64_t aBatchRank = aTile.size() - 2;
170 SmallVector<int64_t> expectedCTile(aTile.begin(),
171 aTile.begin() + aBatchRank);
172 expectedCTile.push_back(aTile[aBatchRank]);
173 expectedCTile.push_back(bTile.back());
174 if (!llvm::equal(*cTile, expectedCTile))
180 auto validateScaleATile =
181 [&](Operation *op,
unsigned scaleAOperandIdx,
182 const SmallVector<int64_t> &aTile) -> std::optional<int64_t> {
183 std::optional<SmallVector<int64_t>> aScaleTile =
186 if (!aScaleTile || aScaleTile->size() < 2)
191 int64_t scaleRank = aScaleTile->size();
192 int64_t aBatchRank = aTile.size() - 2;
193 if ((*aScaleTile)[scaleRank - 2] != aTile[aBatchRank])
197 return aScaleTile->back();
201 auto validateScaleBTile =
202 [&](Operation *op,
unsigned scaleBOperandIdx,
203 const SmallVector<int64_t> &bTile) -> std::optional<int64_t> {
204 std::optional<SmallVector<int64_t>> bScaleTile =
207 if (!bScaleTile || bScaleTile->size() < 2)
212 if (bScaleTile->back() != bTile.back())
216 int64_t scaleRank = bScaleTile->size();
217 return (*bScaleTile)[scaleRank - 2];
220 if (isa<xegpu::DpasOp>(op)) {
221 auto abTiles = validateABTiles(op);
225 auto [aTile, bTile] = *abTiles;
228 if (!validateCTile(op, 2, aTile, bTile))
232 int64_t aBatchRank = aTile.size() - 2;
233 SmallVector<int64_t> tileShape(aTile.begin(), aTile.begin() + aBatchRank);
234 tileShape.push_back(aTile[aBatchRank]);
235 tileShape.push_back(aTile[aBatchRank + 1]);
236 tileShape.push_back(bTile.back());
240 if (
auto dpasMxOp = dyn_cast<xegpu::DpasMxOp>(op)) {
241 auto abTiles = validateABTiles(op);
245 auto [aTile, bTile] = *abTiles;
248 if (dpasMxOp.getAcc()) {
249 unsigned accOperandIdx = 2;
250 if (!validateCTile(op, accOperandIdx, aTile, bTile))
255 int64_t kScaleFactor = 1;
256 std::optional<int64_t> scaleAFactor;
257 std::optional<int64_t> scaleBFactor;
259 if (dpasMxOp.getScaleA()) {
260 unsigned scaleAOperandIdx = 2 + (dpasMxOp.getAcc() ? 1 : 0);
261 scaleAFactor = validateScaleATile(op, scaleAOperandIdx, aTile);
266 if (dpasMxOp.getScaleB()) {
267 unsigned scaleBOperandIdx =
268 2 + (dpasMxOp.getAcc() ? 1 : 0) + (dpasMxOp.getScaleA() ? 1 : 0);
269 scaleBFactor = validateScaleBTile(op, scaleBOperandIdx, bTile);
275 if (scaleAFactor && scaleBFactor) {
276 if (*scaleAFactor != *scaleBFactor)
278 kScaleFactor = *scaleAFactor;
279 }
else if (scaleAFactor) {
280 kScaleFactor = *scaleAFactor;
281 }
else if (scaleBFactor) {
282 kScaleFactor = *scaleBFactor;
286 int64_t aBatchRank = aTile.size() - 2;
287 SmallVector<int64_t> tileShape(aTile.begin(), aTile.begin() + aBatchRank);
288 tileShape.push_back(aTile[aBatchRank]);
289 tileShape.push_back(aTile[aBatchRank + 1]);
290 tileShape.push_back(bTile.back());
291 tileShape.push_back(kScaleFactor);
298 if (isa<vector::MultiDimReductionOp>(op))
301 if (isa<vector::TransposeOp, vector::BroadcastOp, vector::StepOp,
302 vector::ShapeCastOp, vector::ConstantMaskOp, vector::CreateMaskOp,
303 vector::BitCastOp, vector::InterleaveOp, vector::DeinterleaveOp>(op))
309bool XeGPUBlockingPass::needsUnroll(Operation *op)
const {
311 bool hasWgLayoutOperands =
313 xegpu::DistributeLayoutAttr layout =
314 xegpu::getDistributeLayoutAttr(opr);
315 return layout && layout.isForWorkgroup();
317 bool hasWgLayoutResults =
319 xegpu::DistributeLayoutAttr layout =
320 xegpu::getDistributeLayoutAttr(result);
321 return layout && layout.isForWorkgroup();
323 if (hasWgLayoutOperands || hasWgLayoutResults) {
324 LDBG() <<
"skip unrolling for op with workgroup level layout: " << *op;
328 auto isUnrollable = [](Value value, ArrayRef<int64_t> tileShape) {
330 if (
auto tdescTy = dyn_cast<xegpu::TensorDescType>(valTy)) {
331 xegpu::DistributeLayoutAttr layout = tdescTy.getLayoutAttr();
332 return layout && !layout.getEffectiveInstDataAsInt().empty();
334 auto shapedType = dyn_cast<ShapedType>(valTy);
335 return shapedType && !llvm::equal(tileShape, shapedType.getShape());
338 bool hasUnrollableOperands =
340 std::optional<SmallVector<int64_t>> tileShape = getTileShape(opr);
341 return tileShape.has_value() && isUnrollable(opr.get(), *tileShape);
343 bool hasUnrollableResults =
345 std::optional<SmallVector<int64_t>> tileShape = getTileShape(result);
346 return tileShape.has_value() && isUnrollable(result, *tileShape);
349 bool isConvertLayoutWithInstData =
false;
350 if (
auto convertLayoutOp = dyn_cast<xegpu::ConvertLayoutOp>(op)) {
351 auto targettLayout = convertLayoutOp.getTargetLayout();
352 if (targettLayout && !targettLayout.getEffectiveInstDataAsInt().empty()) {
353 isConvertLayoutWithInstData =
true;
356 return hasUnrollableOperands || hasUnrollableResults ||
357 isConvertLayoutWithInstData;
360void XeGPUBlockingPass::runOnOperation() {
362 Operation *op = getOperation();
369 auto getTileShapeAndCount = [](llvm::ArrayRef<int64_t> shape,
370 xegpu::DistributeLayoutAttr layout) {
372 SmallVector<int64_t> tileShape(shape);
373 if (layout && !layout.getEffectiveInstDataAsInt().empty()) {
374 tileShape = layout.getEffectiveInstDataAsInt();
377 assert(count >= 1 &&
"count must be at least 1");
378 return std::make_pair(tileShape, count);
383 llvm::SmallSetVector<UnrealizedConversionCastOp, 8> existingCasts;
385 [&](UnrealizedConversionCastOp castOp) { existingCasts.insert(castOp); });
388 TypeConverter converter;
389 converter.addConversion([](Type type) -> Type {
return type; });
392 converter.addConversion(
393 [&](xegpu::TensorDescType type,
394 SmallVectorImpl<Type> &
result) -> std::optional<LogicalResult> {
395 Type elemTy = type.getElementType();
396 ArrayRef<int64_t> shape = type.getShape();
398 xegpu::DistributeLayoutAttr layout = type.getLayoutAttr();
399 if (layout && layout.isForWorkgroup())
403 SmallVector<int64_t> subShape;
404 std::tie(subShape, count) = getTileShapeAndCount(shape, layout);
407 layout = layout.dropInstData();
409 auto newTy = xegpu::TensorDescType::get(
410 type.getContext(), subShape, elemTy, type.getEncoding(), layout);
411 result.append(count, newTy);
417 auto getSubShapeAndCount = [&](VectorType vecTy,
418 xegpu::DistributeLayoutAttr layout)
419 -> std::pair<SmallVector<int64_t>,
int> {
420 return getTileShapeAndCount(vecTy.getShape(), layout);
425 std::move(loopArgTypes));
433 op->
walk([](Operation *loopOp) {
434 if (!isa<scf::ForOp, scf::WhileOp, scf::ConditionOp, scf::IfOp>(loopOp))
436 SmallVector<StringRef> toRemove;
437 for (
const NamedAttribute &attr :
439 StringRef name = attr.getName().strref();
440 if (name.starts_with(
"layout_operand_") ||
441 name.starts_with(
"layout_result_"))
442 toRemove.push_back(name);
444 for (StringRef name : toRemove)
450 auto materializeCast = [](OpBuilder &builder, Type type,
ValueRange inputs,
451 Location loc) -> Value {
452 return UnrealizedConversionCastOp::create(builder, loc, type, inputs)
455 converter.addSourceMaterialization(materializeCast);
456 converter.addTargetMaterialization(materializeCast);
459 converter.addTargetMaterialization(
460 [](mlir::OpBuilder &builder, mlir::TypeRange types,
461 mlir::ValueRange inputs, mlir::Location loc) -> SmallVector<Value> {
463 UnrealizedConversionCastOp::create(builder, loc, types, inputs);
464 return SmallVector<Value>(castOp.getResults());
467 ConversionTarget
target(*ctx);
468 target.addLegalOp<UnrealizedConversionCastOp>();
469 target.markUnknownOpDynamicallyLegal([](Operation *) {
return true; });
471 RewritePatternSet scfPatterns(ctx);
474 if (
failed(applyPartialConversion(op,
target, std::move(scfPatterns))))
475 return signalPassFailure();
483 [&](Operation *op) -> LogicalResult {
return success(needsUnroll(op)); });
487 options.setUnrolledTypesFn([&](ShapedType type, ArrayRef<int64_t> tileShape) {
488 Type elemTy = type.getElementType();
490 if (
auto tdescTy = dyn_cast<xegpu::TensorDescType>(type)) {
492 Attribute encoding = tdescTy.getEncoding();
494 xegpu::TensorDescType newTy =
495 xegpu::TensorDescType::get(ctx, tileShape, elemTy, encoding,
496 tdescTy.getLayoutAttr().dropInstData());
498 ArrayRef<int64_t> shape = type.getShape();
501 return SmallVector<Type>(batchCount, newTy);
503 Type newTy = VectorType::get(tileShape, elemTy);
505 std::optional<SmallVector<int64_t>> ratio =
507 assert(ratio &&
"The shape of the type must be a multiple of tileShape.");
511 RewritePatternSet patterns(ctx);
512 vector::UnrollVectorOptions vectorOptions;
516 vector::populateVectorUnrollPatterns(patterns, vectorOptions);
525 op->
walk([](Operation *op) {
539 if (!isa<LoopLikeOpInterface>(op))
559 RewritePatternSet emptyPatterns(ctx);
static std::array< int64_t, 2 > getTileShape(ArrayRef< int64_t > operandShape, Type elementType, int64_t lineSizeBits)
Returns the number of 8 x [128|256|512] bit tiles that compose the given operand shape.
static llvm::ManagedStatic< PassManagerOptions > options
Operation is the basic unit of execution within MLIR.
bool hasDiscardableAttrOfType(NameT &&name)
OpResult getOpResult(unsigned idx)
MutableArrayRef< OpOperand > getOpOperands()
unsigned getNumOperands()
DictionaryAttr getDiscardableAttrDictionary()
Return all of the discardable attributes on this operation as a DictionaryAttr.
Attribute removeDiscardableAttr(StringAttr name)
Remove the discardable attribute with the specified name if it exists.
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),...
AttrClass getDiscardableAttrOfType(StringRef name)
Access a discardable attribute by name and cast it to AttrClass.
result_range getOpResults()
OpOperand & getOpOperand(unsigned idx)
void setDiscardableAttrs(DictionaryAttr newAttrs)
Set the discardable attribute dictionary on this operation.
unsigned getNumResults()
Return the number of results held by this operation.
Type getType() const
Return the type of this value.
bool hasElementwiseMappableTraits(Operation *op)
Together, Elementwise, Scalarizable, Vectorizable, and Tensorizable provide an easy way for scalar op...
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...
void populateXeGPUUnrollPatterns(RewritePatternSet &patterns, const UnrollOptions &options)
Collect a set of patterns to unroll xegpu operations to a smaller shapes.
void setDistributeLayoutAttr(const OpResult &Result, const DistributeLayoutAttr layout)
[to-be-deprecated] Sets the DistributeLayoutAttr for a given OpResult user should use setAnchorLayout...
SmallVector< NamedAttribute > dropInstDataOnAttrs(ArrayRef< NamedAttribute > attrs)
Updates the NamedAttribute sequence by dropping inst-data information from any DistributeLayoutAttr f...
bool recoverTemporaryLayouts(Operation *rootOp)
Attach layout attributes to all vector-type operands of operations within the given operation's neste...
void dropInstDataOnInherentAttrs(Operation *op)
Drops inst-data information from DistributeLayoutAttrs stored as inherent attributes on the operation...
DistributeLayoutAttr getDistributeLayoutAttr(const Value value)
Retrieves the DistributeLayoutAttr associated with a given Value, or nullptr if none is found.
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::string getTemporaryLayoutName(const OpOperand &operand)
Return the attribute name for the OpOperand to attach DistributeLayoutAttr.
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,...
void cleanupUnrealizedConversionCasts(Operation *root, const llvm::SmallSetVector< UnrealizedConversionCastOp, 8 > &existingCasts)
Cleans up UnrealizedConversionCastOps inserted during SCF structural type conversion and/or XeGPU unr...
Include the generated interface declarations.
LogicalResult applyPatternsGreedily(Region ®ion, const FrozenRewritePatternSet &patterns, GreedyRewriteConfig config=GreedyRewriteConfig(), bool *changed=nullptr)
Rewrite ops in the given region, which must be isolated from above, by repeatedly applying the highes...
int64_t computeProduct(ArrayRef< int64_t > basis)
Self-explicit.
SmallVector< Loops, 8 > tile(ArrayRef< scf::ForOp > forOps, ArrayRef< Value > sizes, ArrayRef< scf::ForOp > targets)
Performs tiling fo imperfectly nested loops (with interchange) by strip-mining the forOps by sizes an...
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.
UnrollVectorOptions & setNativeShapeFn(NativeShapeFnType fn)