18#include "llvm/ADT/SmallVector.h"
23#define GEN_PASS_DEF_LINALGBLOCKPACKMATMUL
24#include "mlir/Dialect/Linalg/Passes.h.inc"
33 if (!stride || *stride != 1)
41 return (*size - *offset);
48 if (dims.size() != tiles.size() || tiles.empty())
51 FailureOr<ContractionDimensions> contractDims =
53 if (failed(contractDims))
55 unsigned batchDimsOffset = contractDims->batch.size();
60 for (
int64_t &offsetDim : offsetDims)
61 offsetDim += batchDimsOffset;
63 auto tileOp = cast<TilingInterface>(linalgOp.getOperation());
68 for (
auto dim : llvm::enumerate(offsetDims)) {
69 if (dim.value() >=
static_cast<int64_t>(iterationDomain.size()))
73 std::optional<int64_t> rangeOnDim =
78 if (!tileSize || !rangeOnDim)
82 if (*rangeOnDim % *tileSize != 0)
90static FailureOr<PackTransposeResult>
92 linalg::PackOp packOp,
AffineMap operandMap,
94 bool transposeOuterBlocks,
bool transposeInnerBlocks) {
98 if (!packOp.hasPureTensorSemantics())
102 "expected at least 4D prepacked matmul");
103 assert(blocksStartDimPos.size() >= 2 &&
104 "expected starting outer and inner block positions");
114 bool isOuterTransposed =
115 operandMap.
getDimPosition(outerBlockPos) != blocksStartDimPos.end()[-2];
116 bool isInnerTransposed =
117 operandMap.
getDimPosition(innerBlockPos) != blocksStartDimPos.back();
122 if (isInnerTransposed != transposeInnerBlocks)
125 if (isOuterTransposed != transposeOuterBlocks)
131 for (
auto i : llvm::seq(0u, outerBlockPos))
132 offsetPerms.push_back(i);
133 for (
auto perm : outerPerm)
134 offsetPerms.push_back(perm + outerBlockPos);
135 outerPerm = offsetPerms;
137 FailureOr<PackTransposeResult> packTransposedMatmul =
139 nullptr, outerPerm, innerPerm);
141 return packTransposedMatmul;
150 if (
auto *batchMatmulOp = dyn_cast<linalg::BatchMatmulOp>(&linalgOp)) {
151 if (batchMatmulOp->hasUserDefinedMaps()) {
154 "only batch_matmul ops with non-extended semantics are supported");
158 if (linalgOp.hasPureBufferSemantics())
161 std::optional<BlockPackMatmulOptions>
options = controlPackMatmul(linalgOp);
165 if (
options->blockFactors.size() != 3)
168 bool hasScalable = !
options->scalableBlockFactors.empty();
169 if (hasScalable &&
options->scalableBlockFactors.size() != 3)
171 linalgOp,
"scalableBlockFactors must be empty or have 3 elements");
176 if (!
options->allowPadding && hasScalable)
178 linalgOp,
"scalable block factors require padding");
181 for (
auto [idx, factor] : llvm::enumerate(
options->blockFactors)) {
182 bool isScalable = hasScalable &&
options->scalableBlockFactors[idx];
189 Value vscale = vector::VectorScaleOp::create(rewriter, linalgOp.getLoc(),
192 arith::MulIOp::create(rewriter, linalgOp.getLoc(), cst, vscale)
200 "expect packing full tiles only");
212 rewriter, linalgOp, mnkTiles,
options->mnkPaddedSizesNextMultipleOf,
214 if (failed(packedMatmul))
217 assert(packedMatmul->packOps.size() == 3 &&
218 "invalid number of pack ops after matmul packing");
219 assert(packedMatmul->unPackOps.size() == 1 &&
220 "invalid number of unpack ops after matmul packing");
222 FailureOr<ContractionDimensions> contractDims =
224 if (failed(contractDims))
228 dyn_cast<linalg::GenericOp>(packedMatmul->packedLinalgOp.getOperation());
233 rewriter, packedMatmul->packedLinalgOp, packedMatmul->packOps[0], maps[0],
234 contractDims->m,
options->lhsTransposeOuterBlocks,
235 options->lhsTransposeInnerBlocks);
236 if (failed(packedLhs))
240 packedMatmul->packOps[0] = packedLhs->transposedPackOp;
241 packedMatmul->packedLinalgOp = packedLhs->transposedLinalgOp;
245 rewriter, packedMatmul->packedLinalgOp, packedMatmul->packOps[1], maps[1],
246 contractDims->k,
options->rhsTransposeOuterBlocks,
247 options->rhsTransposeInnerBlocks);
248 if (failed(packedRhs))
252 packedMatmul->packOps[1] = packedRhs->transposedPackOp;
253 packedMatmul->packedLinalgOp = packedRhs->transposedLinalgOp;
259template <
typename OpTy>
265 LogicalResult matchAndRewrite(OpTy linalgOp,
267 FailureOr<PackResult> packedMatmul =
269 if (failed(packedMatmul))
279struct BlockPackMatmul<
linalg::GenericOp>
282 PatternBenefit benefit = 1)
283 : OpRewritePattern<linalg::GenericOp>(context, benefit),
284 controlFn(std::move(fun)) {}
286 LogicalResult matchAndRewrite(linalg::GenericOp linalgOp,
287 PatternRewriter &rewriter)
const override {
293 using MapList = ArrayRef<ArrayRef<AffineExpr>>;
294 auto infer = [&](MapList m) {
299 bindDims(linalgOp->getContext(), i, j, k);
300 SmallVector<AffineMap> maps = linalgOp.getIndexingMapsArray();
303 if (!(maps == infer({{i, k}, {k, j}, {i, j}}) ||
304 maps == infer({{k, i}, {k, j}, {i, j}}) ||
305 maps == infer({{i, k}, {j, k}, {i, j}}))) {
309 FailureOr<PackResult> packedMatmul =
321struct LinalgBlockPackMatmul
322 :
public impl::LinalgBlockPackMatmulBase<LinalgBlockPackMatmul> {
323 using LinalgBlockPackMatmulBase::LinalgBlockPackMatmulBase;
325 void runOnOperation()
override {
326 Operation *op = getOperation();
330 [&](linalg::LinalgOp op) -> BlockPackMatmulOptions {
331 BlockPackMatmulOptions
options;
335 for (
const std::string &blockFactor : *blockFactors) {
336 StringRef factor(blockFactor);
337 if (factor.starts_with(
"[") && factor.ends_with(
"]")) {
339 factor.drop_front().drop_back().getAsInteger(10, val);
340 options.blockFactors.push_back(val);
341 options.scalableBlockFactors.push_back(
true);
344 factor.getAsInteger(10, val);
345 options.blockFactors.push_back(val);
346 options.scalableBlockFactors.push_back(
false);
351 if (llvm::none_of(
options.scalableBlockFactors, [](
bool b) { return b; }))
352 options.scalableBlockFactors.clear();
354 options.allowPadding = allowPadding;
355 options.mnkPaddedSizesNextMultipleOf =
356 SmallVector<int64_t>{*mnkPaddedSizesNextMultipleOf};
357 if (!mnkOrder.empty())
358 options.mnkOrder = SmallVector<int64_t>{*mnkOrder};
359 options.lhsTransposeOuterBlocks = lhsTransposeOuterBlocks;
360 options.lhsTransposeInnerBlocks = lhsTransposeInnerBlocks;
361 options.rhsTransposeOuterBlocks = rhsTransposeOuterBlocks;
362 options.rhsTransposeInnerBlocks = rhsTransposeInnerBlocks;
368 return signalPassFailure();
375 patterns.
add<BlockPackMatmul<linalg::GenericOp>,
376 BlockPackMatmul<linalg::MatmulOp>,
377 BlockPackMatmul<linalg::BatchMatmulOp>>(patterns.
getContext(),
static FailureOr< PackTransposeResult > transposePackedMatmul(RewriterBase &rewriter, linalg::LinalgOp linalgOp, linalg::PackOp packOp, AffineMap operandMap, ArrayRef< unsigned > blocksStartDimPos, bool transposeOuterBlocks, bool transposeInnerBlocks)
Return failure or packed matmul with one of its operands transposed.
static bool validateFullTilesOnDims(linalg::LinalgOp linalgOp, ArrayRef< OpFoldResult > tiles, ArrayRef< int64_t > dims)
Return true if all dimensions are fully divisible by the respective tiles.
static std::optional< int64_t > getConstantRange(const Range &range)
Return constant range span or nullopt, otherwise.
static llvm::ManagedStatic< PassManagerOptions > options
A multi-dimensional affine map Affine map's are immutable like Type's, and they are uniqued.
unsigned getDimPosition(unsigned idx) const
Extracts the position of the dimensional expression at the given result, when the caller knows it is ...
unsigned getNumDims() const
unsigned getNumResults() const
static SmallVector< AffineMap, 4 > inferFromExprList(ArrayRef< ArrayRef< AffineExpr > > exprsList, MLIRContext *context)
Returns a vector of AffineMaps; each with as many results as exprs.size(), as many dims as the larges...
IntegerAttr getIndexAttr(int64_t value)
MLIRContext is the top-level object for a collection of MLIR operations.
RAII guard to reset the insertion point of the builder when destroyed.
This class helps build Operations.
void setInsertionPointAfter(Operation *op)
Sets the insertion point to the node after the specified operation, which will cause subsequent inser...
This class represents the benefit of a pattern match in a unitless scheme that ranges from 0 (very li...
A special type of RewriterBase that coordinates the application of a rewrite pattern on the current I...
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.
This class coordinates the application of a rewrite on a set of IR, providing a way for clients to tr...
std::enable_if_t<!std::is_convertible< CallbackT, Twine >::value, LogicalResult > notifyMatchFailure(Location loc, CallbackT &&reasonCallback)
Used to notify the listener that the IR failed to be rewritten because of a match failure,...
This class represents an instance of an SSA value in the MLIR system, representing a computable value...
static ConstantIndexOp create(OpBuilder &builder, Location location, int64_t value)
void populateBlockPackMatmulPatterns(RewritePatternSet &patterns, const ControlBlockPackMatmulFn &controlFn)
Patterns to block pack Linalg matmul ops.
FailureOr< PackTransposeResult > packTranspose(RewriterBase &rewriter, linalg::PackOp packOp, linalg::LinalgOp linalgOp, linalg::UnPackOp maybeUnPackOp, ArrayRef< int64_t > outerPerm, ArrayRef< int64_t > innerPerm)
Transpose a single PackOp -> LinalgOp -> UnPackOp chain and return the transposed PackOp -> LinalgOp ...
std::function< std::optional< BlockPackMatmulOptions >(linalg::LinalgOp)> ControlBlockPackMatmulFn
Function type which is used to control matmul packing.
FailureOr< PackResult > blockPackMatmul(RewriterBase &rewriter, linalg::LinalgOp linalgOp, const ControlBlockPackMatmulFn &controlPackMatmul)
Pack a matmul operation into blocked 4D layout.
FailureOr< ContractionDimensions > inferContractionDims(LinalgOp linalgOp)
Find at least 2 parallel (m and n) and 1 reduction (k) dimension candidates that form a matmul subcom...
FailureOr< PackResult > packMatmulGreedily(RewriterBase &rewriter, LinalgOp linalgOp, ArrayRef< OpFoldResult > mnkPackedSizes, ArrayRef< int64_t > mnkPaddedSizesNextMultipleOf, ArrayRef< int64_t > mnkOrder)
Pack a LinalgOp by greedily inferring matmul dimensions (m, n, k) where m and n are proper parallel d...
bool isaContractionOpInterface(LinalgOp linalgOp)
Checks whether linalgOp conforms to ContractionOpInterface.
Include the generated interface declarations.
std::optional< int64_t > getConstantIntValue(OpFoldResult ofr)
If ofr is a constant integer or an IntegerAttr, return the integer.
void bindDims(MLIRContext *ctx, AffineExprTy &...exprs)
Bind a list of AffineExpr references to DimExpr at positions: [0 .
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...
OpRewritePattern is a wrapper around RewritePattern that allows for matching and rewriting against an...
Represents a range (offset, size, and stride) where each element of the triple may be dynamic or stat...