MLIR 24.0.0git
BlockPackMatmul.cpp
Go to the documentation of this file.
1//===- BlockPackMatmul.cpp - Linalg matmul block packing ------------------===//
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//===----------------------------------------------------------------------===//
8
10
18#include "llvm/ADT/SmallVector.h"
19
20#include <optional>
21
22namespace mlir {
23#define GEN_PASS_DEF_LINALGBLOCKPACKMATMUL
24#include "mlir/Dialect/Linalg/Passes.h.inc"
25} // namespace mlir
26
27using namespace mlir;
28using namespace mlir::linalg;
29
30/// Return constant range span or nullopt, otherwise.
31static std::optional<int64_t> getConstantRange(const Range &range) {
32 std::optional<int64_t> stride = getConstantIntValue(range.stride);
33 if (!stride || *stride != 1)
34 return std::nullopt;
35 std::optional<int64_t> offset = getConstantIntValue(range.offset);
36 if (!offset)
37 return std::nullopt;
38 std::optional<int64_t> size = getConstantIntValue(range.size);
39 if (!size)
40 return std::nullopt;
41 return (*size - *offset);
42}
43
44/// Return true if all dimensions are fully divisible by the respective tiles.
45static bool validateFullTilesOnDims(linalg::LinalgOp linalgOp,
47 ArrayRef<int64_t> dims) {
48 if (dims.size() != tiles.size() || tiles.empty())
49 return false;
50
51 FailureOr<ContractionDimensions> contractDims =
52 inferContractionDims(linalgOp);
53 if (failed(contractDims))
54 return false;
55 unsigned batchDimsOffset = contractDims->batch.size();
56
57 // Skip the batch dimension if present.
58 // Offset all dimensions accordingly.
59 SmallVector<int64_t, 3> offsetDims(dims);
60 for (int64_t &offsetDim : offsetDims)
61 offsetDim += batchDimsOffset;
62
63 auto tileOp = cast<TilingInterface>(linalgOp.getOperation());
64 OpBuilder builder(tileOp);
65 OpBuilder::InsertionGuard guard(builder);
66 SmallVector<Range> iterationDomain = tileOp.getIterationDomain(builder);
67
68 for (auto dim : llvm::enumerate(offsetDims)) {
69 if (dim.value() >= static_cast<int64_t>(iterationDomain.size()))
70 return false;
71
72 std::optional<int64_t> tileSize = getConstantIntValue(tiles[dim.index()]);
73 std::optional<int64_t> rangeOnDim =
74 getConstantRange(iterationDomain[dim.value()]);
75
76 // If the tile factor or the range are non-constant, the tile size is
77 // considered to be invalid.
78 if (!tileSize || !rangeOnDim)
79 return false;
80
81 // The dimension must be fully divisible by the tile.
82 if (*rangeOnDim % *tileSize != 0)
83 return false;
84 }
85
86 return true;
87}
88
89/// Return failure or packed matmul with one of its operands transposed.
90static FailureOr<PackTransposeResult>
91transposePackedMatmul(RewriterBase &rewriter, linalg::LinalgOp linalgOp,
92 linalg::PackOp packOp, AffineMap operandMap,
93 ArrayRef<unsigned> blocksStartDimPos,
94 bool transposeOuterBlocks, bool transposeInnerBlocks) {
95 // Pack/unpack memref transformations are unsupported. The memref forms
96 // are mainly for bufferization and scalar lowering. Other uses are not
97 // recommended, see #225650 for details.
98 if (!packOp.hasPureTensorSemantics())
99 return failure();
100
101 assert(operandMap.getNumDims() >= 4 &&
102 "expected at least 4D prepacked matmul");
103 assert(blocksStartDimPos.size() >= 2 &&
104 "expected starting outer and inner block positions");
105
106 // Bias toward innermost dimensions.
107 unsigned outerBlockPos = operandMap.getNumResults() - 4;
108 unsigned innerBlockPos = operandMap.getNumResults() - 2;
109
110 // Transpose control options define the desired block and element layout.
111 // Block transposition (outer dimensions) or element transposition (inner
112 // dimensions) may not be necessary depending on the original matmul data
113 // layout.
114 bool isOuterTransposed =
115 operandMap.getDimPosition(outerBlockPos) != blocksStartDimPos.end()[-2];
116 bool isInnerTransposed =
117 operandMap.getDimPosition(innerBlockPos) != blocksStartDimPos.back();
118
119 // Transpose only the dimensions that need that to conform to the provided
120 // transpotion settings.
121 SmallVector<int64_t> innerPerm = {0, 1};
122 if (isInnerTransposed != transposeInnerBlocks)
123 innerPerm = {1, 0};
124 SmallVector<int64_t> outerPerm = {0, 1};
125 if (isOuterTransposed != transposeOuterBlocks)
126 outerPerm = {1, 0};
127
128 // Leave the outer dimensions, like batch, unchanged by offsetting all
129 // outer dimensions permutations.
130 SmallVector<int64_t> offsetPerms;
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;
136
137 FailureOr<PackTransposeResult> packTransposedMatmul =
138 packTranspose(rewriter, packOp, linalgOp,
139 /*maybeUnPackOp=*/nullptr, outerPerm, innerPerm);
140
141 return packTransposedMatmul;
142}
143
144/// Pack a matmul operation into blocked 4D layout.
145FailureOr<PackResult>
146linalg::blockPackMatmul(RewriterBase &rewriter, linalg::LinalgOp linalgOp,
147 const ControlBlockPackMatmulFn &controlPackMatmul) {
148 // Check to not let go the batch_matmul with extended semantic, through this
149 // transform.
150 if (auto *batchMatmulOp = dyn_cast<linalg::BatchMatmulOp>(&linalgOp)) {
151 if (batchMatmulOp->hasUserDefinedMaps()) {
152 return rewriter.notifyMatchFailure(
153 *batchMatmulOp,
154 "only batch_matmul ops with non-extended semantics are supported");
155 }
156 }
157
158 if (linalgOp.hasPureBufferSemantics())
159 return rewriter.notifyMatchFailure(linalgOp, "require tensor semantics");
160
161 std::optional<BlockPackMatmulOptions> options = controlPackMatmul(linalgOp);
162 if (!options)
163 return rewriter.notifyMatchFailure(linalgOp, "invalid packing options");
164
165 if (options->blockFactors.size() != 3)
166 return rewriter.notifyMatchFailure(linalgOp, "require 3 tile factors");
167
168 bool hasScalable = !options->scalableBlockFactors.empty();
169 if (hasScalable && options->scalableBlockFactors.size() != 3)
170 return rewriter.notifyMatchFailure(
171 linalgOp, "scalableBlockFactors must be empty or have 3 elements");
172
173 // Scalable tile sizes are non-constant at compile time, so they can never
174 // satisfy the full-tile divisibility check. Reject early before creating
175 // any ops to avoid modifying IR before returning notifyMatchFailure.
176 if (!options->allowPadding && hasScalable)
177 return rewriter.notifyMatchFailure(
178 linalgOp, "scalable block factors require padding");
179
181 for (auto [idx, factor] : llvm::enumerate(options->blockFactors)) {
182 bool isScalable = hasScalable && options->scalableBlockFactors[idx];
183 if (!isScalable) {
184 mnkTiles.push_back(rewriter.getIndexAttr(factor));
185 continue;
186 }
187 Value cst =
188 arith::ConstantIndexOp::create(rewriter, linalgOp.getLoc(), factor);
189 Value vscale = vector::VectorScaleOp::create(rewriter, linalgOp.getLoc(),
190 rewriter.getIndexType());
191 mnkTiles.push_back(
192 arith::MulIOp::create(rewriter, linalgOp.getLoc(), cst, vscale)
193 .getResult());
194 }
195
196 // If padding is disabled, make sure that dimensions can be packed cleanly.
197 if (!options->allowPadding &&
198 !validateFullTilesOnDims(linalgOp, mnkTiles, options->mnkOrder)) {
199 return rewriter.notifyMatchFailure(linalgOp,
200 "expect packing full tiles only");
201 }
202
203 OpBuilder::InsertionGuard guard(rewriter);
204 // The op is replaced, we need to set the insertion point after it.
205 rewriter.setInsertionPointAfter(linalgOp);
206
207 // Pack the matmul operation into blocked layout with two levels of
208 // subdivision:
209 // - major 2D blocks - outer dimensions, consist of minor blocks
210 // - minor 2D blocks - inner dimensions, consist of scalar elements
211 FailureOr<PackResult> packedMatmul = packMatmulGreedily(
212 rewriter, linalgOp, mnkTiles, options->mnkPaddedSizesNextMultipleOf,
213 options->mnkOrder);
214 if (failed(packedMatmul))
215 return failure();
216
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");
221
222 FailureOr<ContractionDimensions> contractDims =
223 inferContractionDims(packedMatmul->packedLinalgOp);
224 if (failed(contractDims))
225 return failure();
226
227 auto genericOp =
228 dyn_cast<linalg::GenericOp>(packedMatmul->packedLinalgOp.getOperation());
229 SmallVector<AffineMap> maps = genericOp.getIndexingMapsArray();
230
231 // Transpose LHS matrix according to the options.
232 FailureOr<PackTransposeResult> packedLhs = transposePackedMatmul(
233 rewriter, packedMatmul->packedLinalgOp, packedMatmul->packOps[0], maps[0],
234 contractDims->m, options->lhsTransposeOuterBlocks,
235 options->lhsTransposeInnerBlocks);
236 if (failed(packedLhs))
237 return failure();
238
239 // Update results.
240 packedMatmul->packOps[0] = packedLhs->transposedPackOp;
241 packedMatmul->packedLinalgOp = packedLhs->transposedLinalgOp;
242
243 // Transpose RHS matrix according to the options.
244 FailureOr<PackTransposeResult> packedRhs = transposePackedMatmul(
245 rewriter, packedMatmul->packedLinalgOp, packedMatmul->packOps[1], maps[1],
246 contractDims->k, options->rhsTransposeOuterBlocks,
247 options->rhsTransposeInnerBlocks);
248 if (failed(packedRhs))
249 return failure();
250
251 // Update results.
252 packedMatmul->packOps[1] = packedRhs->transposedPackOp;
253 packedMatmul->packedLinalgOp = packedRhs->transposedLinalgOp;
254
255 return packedMatmul;
256}
257
258namespace {
259template <typename OpTy>
260struct BlockPackMatmul : public OpRewritePattern<OpTy> {
261 BlockPackMatmul(MLIRContext *context, ControlBlockPackMatmulFn fun,
262 PatternBenefit benefit = 1)
263 : OpRewritePattern<OpTy>(context, benefit), controlFn(std::move(fun)) {}
264
265 LogicalResult matchAndRewrite(OpTy linalgOp,
266 PatternRewriter &rewriter) const override {
267 FailureOr<PackResult> packedMatmul =
268 blockPackMatmul(rewriter, linalgOp, controlFn);
269 if (failed(packedMatmul))
270 return failure();
271 return success();
272 }
273
274private:
275 ControlBlockPackMatmulFn controlFn;
276};
277
278template <>
279struct BlockPackMatmul<linalg::GenericOp>
280 : public OpRewritePattern<linalg::GenericOp> {
281 BlockPackMatmul(MLIRContext *context, ControlBlockPackMatmulFn fun,
282 PatternBenefit benefit = 1)
283 : OpRewritePattern<linalg::GenericOp>(context, benefit),
284 controlFn(std::move(fun)) {}
285
286 LogicalResult matchAndRewrite(linalg::GenericOp linalgOp,
287 PatternRewriter &rewriter) const override {
288 // Match suitable generics.
289 if (!linalg::isaContractionOpInterface(linalgOp)) {
290 return rewriter.notifyMatchFailure(linalgOp, "not a contraction");
291 }
292
293 using MapList = ArrayRef<ArrayRef<AffineExpr>>;
294 auto infer = [&](MapList m) {
295 return AffineMap::inferFromExprList(m, linalgOp.getContext());
296 };
297
298 AffineExpr i, j, k;
299 bindDims(linalgOp->getContext(), i, j, k);
300 SmallVector<AffineMap> maps = linalgOp.getIndexingMapsArray();
301
302 // For now, only match simple matmuls.
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}}))) {
306 return rewriter.notifyMatchFailure(linalgOp, "not a suitable matmul");
307 }
308
309 FailureOr<PackResult> packedMatmul =
310 blockPackMatmul(rewriter, linalgOp, controlFn);
311 if (failed(packedMatmul))
312 return failure();
313 return success();
314 }
315
316private:
317 ControlBlockPackMatmulFn controlFn;
318};
319
320/// Convert linalg matmul ops to block layout and back.
321struct LinalgBlockPackMatmul
322 : public impl::LinalgBlockPackMatmulBase<LinalgBlockPackMatmul> {
323 using LinalgBlockPackMatmulBase::LinalgBlockPackMatmulBase;
324
325 void runOnOperation() override {
326 Operation *op = getOperation();
327 RewritePatternSet patterns(&getContext());
328
329 ControlBlockPackMatmulFn controlFn =
330 [&](linalg::LinalgOp op) -> BlockPackMatmulOptions {
331 BlockPackMatmulOptions options;
332
333 // Parse block-factors strings. Each element is either "N" (static) or
334 // "[N]" (scalable, i.e. N * vscale at runtime).
335 for (const std::string &blockFactor : *blockFactors) {
336 StringRef factor(blockFactor);
337 if (factor.starts_with("[") && factor.ends_with("]")) {
338 int64_t val = 0;
339 factor.drop_front().drop_back().getAsInteger(10, val);
340 options.blockFactors.push_back(val);
341 options.scalableBlockFactors.push_back(true);
342 } else {
343 int64_t val = 0;
344 factor.getAsInteger(10, val);
345 options.blockFactors.push_back(val);
346 options.scalableBlockFactors.push_back(false);
347 }
348 }
349 // If all flags are false, clear the vector so blockPackMatmul can take
350 // the cheaper static path.
351 if (llvm::none_of(options.scalableBlockFactors, [](bool b) { return b; }))
352 options.scalableBlockFactors.clear();
353
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;
363 return options;
364 };
365
366 linalg::populateBlockPackMatmulPatterns(patterns, controlFn);
367 if (failed(applyPatternsGreedily(op, std::move(patterns))))
368 return signalPassFailure();
369 }
370};
371} // namespace
372
374 RewritePatternSet &patterns, const ControlBlockPackMatmulFn &controlFn) {
375 patterns.add<BlockPackMatmul<linalg::GenericOp>,
376 BlockPackMatmul<linalg::MatmulOp>,
377 BlockPackMatmul<linalg::BatchMatmulOp>>(patterns.getContext(),
378 controlFn);
379}
return success()
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.
b
Return true if permutation is a valid permutation of the outer_dims_perm (case OuterOrInnerPerm::Oute...
b getContext())
static llvm::ManagedStatic< PassManagerOptions > options
A multi-dimensional affine map Affine map's are immutable like Type's, and they are uniqued.
Definition AffineMap.h:46
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)
Definition Builders.cpp:116
IndexType getIndexType()
Definition Builders.cpp:59
MLIRContext is the top-level object for a collection of MLIR operations.
Definition MLIRContext.h:63
RAII guard to reset the insertion point of the builder when destroyed.
Definition Builders.h:351
This class helps build Operations.
Definition Builders.h:210
void setInsertionPointAfter(Operation *op)
Sets the insertion point to the node after the specified operation, which will cause subsequent inser...
Definition Builders.h:415
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...
Definition Value.h:96
static ConstantIndexOp create(OpBuilder &builder, Location location, int64_t value)
Definition ArithOps.cpp:398
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.
detail::InFlightRemark failed(Location loc, RemarkOpts opts)
Report an optimization remark that failed.
Definition Remarks.h:734
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 .
Definition AffineExpr.h:311
LogicalResult applyPatternsGreedily(Region &region, 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...
OpFoldResult stride
OpFoldResult size
OpFoldResult offset