MLIR 24.0.0git
VectorTransforms.cpp
Go to the documentation of this file.
1//===- VectorTransforms.cpp - Conversion within the Vector dialect --------===//
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//
9// This file implements target-independent rewrites as 1->N patterns.
10//
11//===----------------------------------------------------------------------===//
12
14
25#include "mlir/IR/Location.h"
26#include "mlir/IR/Matchers.h"
29
30#include "llvm/ADT/STLExtras.h"
31#include "llvm/ADT/SmallVectorExtras.h"
32#include "llvm/Support/FormatVariadic.h"
33
34#include <cassert>
35#include <cstdint>
36#include <functional>
37#include <optional>
38
39#define DEBUG_TYPE "vector-to-vector"
40
41using namespace mlir;
42using namespace mlir::vector;
43
44template <typename IntType>
46 return llvm::to_vector<4>(llvm::map_range(
47 arrayAttr.getAsRange<IntegerAttr>(),
48 [](IntegerAttr attr) { return static_cast<IntType>(attr.getInt()); }));
49}
50
51// Helper to find an index in an affine map.
52static std::optional<int64_t> getResultIndex(AffineMap map, int64_t index) {
53 for (int64_t i = 0, e = map.getNumResults(); i < e; ++i) {
54 int64_t idx = map.getDimPosition(i);
55 if (idx == index)
56 return i;
57 }
58 return std::nullopt;
59}
60
61namespace {
62
63/// Convert MulIOp/MulFOp + MultiDimReductionOp<add> into ContractionOp.
64/// Ex:
65/// ```
66/// %0 = arith.mulf %arg0, %arg1 : vector<8x32x16xf32>
67/// %1 = vector.multi_reduction add, %0 [1]
68/// : vector<8x32x16xf32> to vector<8x16xf32>
69/// ```
70/// Gets converted to:
71/// ```
72/// %1 = vector.contract {indexing_maps = [
73/// affine_map<(d0, d1, d2) -> (d0, d1, d2)>,
74/// affine_map<(d0, d1, d2) -> (d0, d1, d2)>,
75/// affine_map<(d0, d1, d2) -> (d0, d1)>],
76/// iterator_types = ["parallel", "parallel", "reduction"],
77/// kind = add} %0, %arg1, %cst_f0
78/// : vector<8x32x16xf32>, vector<8x32x16xf32> into vector<8x32xf32>
79/// ```
80struct MultiReduceToContract
81 : public OpRewritePattern<vector::MultiDimReductionOp> {
82 using Base::Base;
83
84 LogicalResult matchAndRewrite(vector::MultiDimReductionOp reduceOp,
85 PatternRewriter &rewriter) const override {
86 if (reduceOp.getKind() != vector::CombiningKind::ADD)
87 return failure();
88 Operation *mulOp = reduceOp.getSource().getDefiningOp();
89 if (!mulOp || !isa<arith::MulIOp, arith::MulFOp>(mulOp))
90 return failure();
91 SmallVector<bool> reductionMask = reduceOp.getReductionMask();
92 auto srcMap = rewriter.getMultiDimIdentityMap(reductionMask.size());
93 SmallVector<AffineExpr> exprs;
94 SmallVector<vector::IteratorType> iteratorTypes;
95 for (const auto &isReduceDim : llvm::enumerate(reductionMask)) {
96 if (!isReduceDim.value()) {
97 iteratorTypes.push_back(vector::IteratorType::parallel);
98 exprs.push_back(rewriter.getAffineDimExpr(isReduceDim.index()));
99 } else {
100 iteratorTypes.push_back(vector::IteratorType::reduction);
101 }
102 }
103 auto dstMap =
104 AffineMap::get(/*dimCount=*/reductionMask.size(),
105 /*symbolCount=*/0, exprs, reduceOp.getContext());
106 rewriter.replaceOpWithNewOp<mlir::vector::ContractionOp>(
107 reduceOp, mulOp->getOperand(0), mulOp->getOperand(1), reduceOp.getAcc(),
108 rewriter.getAffineMapArrayAttr({srcMap, srcMap, dstMap}),
109 rewriter.getArrayAttr(llvm::map_to_vector(
110 iteratorTypes, [&](IteratorType t) -> mlir::Attribute {
111 return IteratorTypeAttr::get(rewriter.getContext(), t);
112 })));
113 return success();
114 }
115};
116
117/// Merge LHS/RHS (A/B) TransposeOp into ContractionOp user.
118/// Ex:
119/// ```
120/// %0 = vector.transpose %arg0, [2, 0, 1]
121/// : vector<32x16x8xf32> to vector<8x32x16xf32>
122/// %1 = vector.contract {indexing_maps = [
123/// affine_map<(d0, d1, d2) -> (d0, d1, d2)>,
124/// affine_map<(d0, d1, d2) -> (d0, d1, d2)>,
125/// affine_map<(d0, d1, d2) -> (d0, d1)>],
126/// iterator_types = ["parallel", "parallel", "reduction"],
127/// kind = add} %0, %arg1, %cst_f0
128/// : vector<8x32x16xf32>, vector<8x32x16xf32> into vector<8x32xf32>
129/// ```
130/// Gets converted to:
131/// ```
132/// %1 = vector.contract {indexing_maps = [
133/// affine_map<(d0, d1, d2) -> (d1, d2, d0)>,
134/// affine_map<(d0, d1, d2) -> (d0, d1, d2)>,
135/// affine_map<(d0, d1, d2) -> (d0, d1)>],
136/// iterator_types = ["parallel", "parallel", "reduction"],
137/// kind = add} %arg0, %arg1, %cst_f0
138/// : vector<8x32x16xf32>, vector<8x32x16xf32> into vector<8x32xf32>
139/// ```
140struct CombineContractABTranspose final
141 : public OpRewritePattern<vector::ContractionOp> {
142 using Base::Base;
143
144 LogicalResult matchAndRewrite(vector::ContractionOp contractOp,
145 PatternRewriter &rewriter) const override {
146 SmallVector<AffineMap> maps =
147 llvm::to_vector<4>(contractOp.getIndexingMapsArray());
148 Value lhs = contractOp.getLhs();
149 Value rhs = contractOp.getRhs();
150 size_t index = 0;
151 bool changed = false;
152 for (Value *operand : {&lhs, &rhs}) {
153 AffineMap &map = maps[index++];
154 auto transposeOp = operand->getDefiningOp<vector::TransposeOp>();
155 if (!transposeOp)
156 continue;
157 AffineMap permutationMap = AffineMap::getPermutationMap(
158 transposeOp.getPermutation(), contractOp.getContext());
159 map = inversePermutation(permutationMap).compose(map);
160 *operand = transposeOp.getVector();
161 changed = true;
162 }
163 if (!changed)
164 return failure();
165 rewriter.replaceOpWithNewOp<vector::ContractionOp>(
166 contractOp, lhs, rhs, contractOp.getAcc(),
167 rewriter.getAffineMapArrayAttr(maps), contractOp.getIteratorTypes());
168 return success();
169 }
170};
171
172/// Merges accumulator and result transposes into contract.
173///
174/// For example:
175/// ```mlir
176/// %accT = vector.transpose %acc, [0, 2, 1]
177/// : vector<2x8x4xf32> to vector<2x4x8xf32>
178/// %contract = vector.contract {
179/// indexing_maps = [
180/// affine_map<(d0, d1, d2, d3) -> (d0, d3, d1)>,
181/// affine_map<(d0, d1, d2, d3) -> (d3, d2)>,
182/// affine_map<(d0, d1, d2, d3) -> (d0, d1, d2)>
183/// ],
184/// iterator_types = ["parallel", "parallel", "parallel", "reduction"],
185/// kind = #vector.kind<add>
186/// } %lhs, %rhs, %accT
187/// : vector<2x4x4xf32>, vector<4x8xf32> into vector<2x4x8xf32>
188/// %0 = vector.transpose %contract, [0, 2, 1]
189/// : vector<2x4x8xf32> to vector<2x8x4>
190/// ```
191/// Becomes:
192/// ```mlir
193/// %0 = vector.contract {
194/// indexing_maps = [
195/// affine_map<(d0, d1, d2, d3) -> (d0, d3, d1)>,
196/// affine_map<(d0, d1, d2, d3) -> (d3, d2)>,
197/// affine_map<(d0, d1, d2, d3) -> (d0, d2, d1)>
198/// ],
199/// iterator_types = ["parallel", "parallel", "parallel", "reduction"],
200/// kind = #vector.kind<add>
201/// } %lhs, %rhs, %acc
202/// : vector<2x4x4xf32>, vector<4x8xf32> into vector<2x8x4xf32>
203/// ```
204struct CombineContractResultTranspose final
205 : public OpRewritePattern<vector::TransposeOp> {
206 using Base::Base;
207
208 LogicalResult matchAndRewrite(vector::TransposeOp resTOp,
209 PatternRewriter &rewriter) const override {
210 auto contractOp = resTOp.getVector().getDefiningOp<vector::ContractionOp>();
211 if (!contractOp || !contractOp->hasOneUse())
212 return failure();
213
214 auto accTOp = contractOp.getAcc().getDefiningOp<vector::TransposeOp>();
215 if (!accTOp)
216 return failure();
217
218 MLIRContext *context = contractOp.getContext();
219 auto maps = llvm::to_vector<3>(contractOp.getIndexingMapsArray());
220 AffineMap contractMap = maps.back();
221
222 // Accumulator transpose performs f(A) -> B. Contract performs g(C) -> B.
223 // To index into A in contract, we need revert(f)(g(C)) -> A.
224 auto accTMap =
225 AffineMap::getPermutationMap(accTOp.getPermutation(), context);
226
227 // Contract performs g(C) -> D. Result transpose performs h(D) -> E.
228 // To index into E in contract, we need h(g(C)) -> E.
229 auto resTMap =
230 AffineMap::getPermutationMap(resTOp.getPermutation(), context);
231 auto combinedResMap = resTMap.compose(contractMap);
232
233 // The accumulator and result share the same indexing map. So they should be
234 // the same to be able to merge. This means combinedResMap is the same as
235 // inversePermutation(accTMap).compose(contractMap), which means
236 if (inversePermutation(accTMap) != resTMap)
237 return failure();
238 maps.back() = combinedResMap;
239
240 rewriter.replaceOpWithNewOp<vector::ContractionOp>(
241 resTOp, contractOp.getLhs(), contractOp.getRhs(), accTOp.getVector(),
242 rewriter.getAffineMapArrayAttr(maps), contractOp.getIteratorTypes());
243 return success();
244 }
245};
246
247/// Merge BroadcastOp (and broadcast-like ShapeCastOp) into ContractionOp user.
248/// Ex:
249/// ```
250/// %0 = vector.broadcast %arg0 : vector<32x16xf32> to vector<8x32x16xf32>
251/// %1 = vector.contract {indexing_maps = [
252/// affine_map<(d0, d1, d2) -> (d0, d1, d2)>,
253/// affine_map<(d0, d1, d2) -> (d0, d1, d2)>,
254/// affine_map<(d0, d1, d2) -> (d0, d1)>],
255/// iterator_types = ["parallel", "parallel", "reduction"],
256/// kind = add} %0, %arg1, %cst_f0
257/// : vector<8x32x16xf32>, vector<8x32x16xf32> into vector<8x32xf32>
258/// ```
259/// Gets converted to:
260/// ```
261/// %1 = vector.contract {indexing_maps = [
262/// affine_map<(d0, d1, d2) -> (d1, d2)>,
263/// affine_map<(d0, d1, d2) -> (d0, d1, d2)>,
264/// affine_map<(d0, d1, d2) -> (d0, d1)>],
265/// iterator_types = ["parallel", "parallel", "reduction"],
266/// kind = add} %arg0, %arg1, %cst_f0
267/// : vector<32x16xf32>, vector<8x32x16xf32> into vector<8x32xf32>
268/// ```
269///
270/// For masked vector.contract, the mask requires updating when a dimension is
271/// dropped. In such cases, the dropped dimensions must correspond to the mask's
272/// leading unit dimensions. Supporting more generic cases (e.g. non-unit dims)
273/// is not supported.
274FailureOr<Value> combineContractAndBroadcast(vector::ContractionOp contractOp,
275 MaskingOpInterface maskingOp,
276 PatternRewriter &rewriter) {
278 llvm::to_vector<4>(contractOp.getIndexingMapsArray());
279 Value lhs = contractOp.getLhs();
280 Value rhs = contractOp.getRhs();
281 size_t index = 0;
282 bool changed = false;
283 for (Value *operand : {&lhs, &rhs}) {
284 AffineMap &map = maps[index++];
285
286 // Accept operands defined by vector.broadcast and broadcast-like
287 // vector.shape_cast.
288 auto sc = operand->getDefiningOp<vector::ShapeCastOp>();
289 auto broadcast = operand->getDefiningOp<vector::BroadcastOp>();
290 if (!broadcast && !sc)
291 continue;
292
293 if (sc && !sc.isBroadcastLike())
294 return rewriter.notifyMatchFailure(
295 contractOp, "Operand defined via vector.shape_cast that has "
296 "non-broadcast semantics");
297
298 // Get the source and the result types.
299 VectorType srcType = sc ? sc.getSourceVectorType()
300 : dyn_cast<VectorType>(broadcast.getSourceType());
301 VectorType resType =
302 sc ? sc.getResultVectorType() : broadcast.getResultVectorType();
303
304 // contractionOp can only take vector as operands.
305 // auto srcType = dyn_cast<VectorType>(broadcast.getSourceVectorType());
306 if (!srcType || srcType.getRank() >= resType.getRank())
307 continue;
308 int64_t rankDiff = resType.getRank() - srcType.getRank();
309 bool innerDimBroadcast = false;
310 SmallVector<AffineExpr> originalDims;
311 for (const auto &dim : llvm::enumerate(srcType.getShape())) {
312 if (dim.value() != resType.getDimSize(rankDiff + dim.index())) {
313 innerDimBroadcast = true;
314 break;
315 }
316 originalDims.push_back(rewriter.getAffineDimExpr(dim.index() + rankDiff));
317 }
318 // Contract doesn't support inner dimension broadcast. Once this is
319 // relaxed we can remove this case.
320 if (innerDimBroadcast)
321 continue;
322
323 // It would be incorrect to fold a broadcast onto a reduction dimension
324 // of non-unit size.
325 bool nonUnitDimReductionBroadcast = false;
326 for (int64_t i = 0; i < rankDiff; ++i) {
327 if (resType.getDimSize(i) != 1 &&
328 isReductionIterator(contractOp.getIteratorTypes()
329 .getValue()[map.getDimPosition(i)])) {
330 nonUnitDimReductionBroadcast = true;
331 break;
332 }
333 }
334 if (nonUnitDimReductionBroadcast)
335 continue;
336
337 AffineMap broadcastMap = AffineMap::get(resType.getRank(), 0, originalDims,
338 contractOp.getContext());
339 map = broadcastMap.compose(map);
340 *operand = broadcast ? broadcast.getSource() : sc.getSource();
341 changed = true;
342 }
343
344 if (!changed)
345 return failure();
346
347 // Determine which dims are usused, now that the maps have been composed
348 // with the broadcast maps.
349 llvm::SmallBitVector unusedDimsBitVector = getUnusedDimsBitVector(maps);
350 // Compress unused dims.
351 for (auto &m : maps)
352 m = compressDims(m, unusedDimsBitVector);
353 // Compute the combined iterators.
354 SmallVector<Attribute> iterators;
355 for (unsigned i = 0, e = unusedDimsBitVector.size(); i < e; ++i) {
356 if (!unusedDimsBitVector.test(i))
357 iterators.push_back(contractOp.getIteratorTypes().getValue()[i]);
358 }
359
360 // Check whether any of the unused dims is non-unit, e.g.:
361 // * vector.broadcast %arg0 : vector<8x4xi32> to vector<2x8x4xi32>
362 // This is only required when collapsing a mask. If there is no mask, skip.
363 VectorType oldMaskType;
364 bool isAnyUnusedDimNonUnit = false;
365 if (maskingOp) {
366 oldMaskType = cast<VectorType>(maskingOp.getMask().getType());
367 for (unsigned i = 0, e = unusedDimsBitVector.size(); i < e; ++i) {
368 if (unusedDimsBitVector.test(i) && oldMaskType.getShape()[i] != 1) {
369 isAnyUnusedDimNonUnit = true;
370 break;
371 }
372 }
373 }
374
375 // Check that compressing unused dims isn't removing all reduction dimension
376 // pairs. For example, if the vector.contract had only one reduction
377 // iterator and that was a unit-dimension created by a broadcast,
378 // then we should bail here, otherwise we would create a contract without
379 // a reduction dimension pair.
380 bool hasReductionIteratorApplyingOnBothSides = false;
381 for (unsigned i = 0; i < iterators.size(); ++i) {
382 if (!isReductionIterator(iterators[i]))
383 continue;
384 if (getResultIndex(maps[0], i) && getResultIndex(maps[1], i)) {
385 hasReductionIteratorApplyingOnBothSides = true;
386 break;
387 }
388 }
389 if (!hasReductionIteratorApplyingOnBothSides)
390 return failure();
391
392 // If the compressed maps have a dimension that is not used by either LHS or
393 // RHS then the ContractionOp verifier would fail.
394 if (getUnusedDimsBitVector({maps[0], maps[1]}).any())
395 return failure();
396
397 Operation *newOp = vector::ContractionOp::create(
398 rewriter, contractOp.getLoc(), lhs, rhs, contractOp.getAcc(),
399 rewriter.getAffineMapArrayAttr(maps), rewriter.getArrayAttr(iterators));
400
401 // Handle the mask.
402 if (maskingOp) {
403 if (isAnyUnusedDimNonUnit)
404 return rewriter.notifyMatchFailure(contractOp,
405 "Cannont drop non-unit mask dim.");
406 assert(unusedDimsBitVector.size() ==
407 static_cast<size_t>(oldMaskType.getRank()) &&
408 "The mask rank is incorrect!");
409
410 // If a dimension has been dropped, update the mask accordingly. Otherwise,
411 // keep it as is.
412 Value mask = maskingOp.getMask();
413 if (unusedDimsBitVector.count() != 0) {
414 // At this point, two assumptions are made:
415 // * The unused dimensions are the leading mask dimensions
416 // (vector.contract does not support inner dim broadcasting).
417 // * The unused dimensions are all unit.
418 // These conditions are effectively verified in the blocks preceeding this
419 // one.
420 auto newShape =
421 oldMaskType.getShape().drop_front(unusedDimsBitVector.count());
422 auto newShapeScalableDims =
423 oldMaskType.getScalableDims().drop_front(unusedDimsBitVector.count());
424 VectorType maskOpType =
425 VectorType::get(newShape, rewriter.getI1Type(), newShapeScalableDims);
426 mask = vector::ShapeCastOp::create(rewriter, contractOp.getLoc(),
427 maskOpType, maskingOp.getMask())
428 .getResult();
429 }
430
431 newOp = mlir::vector::maskOperation(rewriter, newOp, mask);
432 }
433 return newOp->getResult(0);
434}
435
436struct CombineContractBroadcastMask
437 : public MaskableOpRewritePattern<vector::ContractionOp> {
438 using MaskableOpRewritePattern::MaskableOpRewritePattern;
439 FailureOr<Value>
440
441 matchAndRewriteMaskableOp(vector::ContractionOp contractOp,
442 MaskingOpInterface maskingOp,
443 PatternRewriter &rewriter) const override {
444 return combineContractAndBroadcast(contractOp, maskingOp, rewriter);
445 }
446};
447
448/// Reorders cast(broadcast) to broadcast(cast). This makes broadcast ops and
449/// contraction ops closer, which kicks in CombineContractBroadcast pattern when
450/// casting ops are around these operations.
451/// Ex:
452/// ```
453/// %0 = vector.broadcast %arg0 : vector<32x16xi8> to vector<8x32x16xi8>
454/// %1 = arith.extsi %0 : vector<8x32x16xi8> to vector<8x32x16xi32>
455/// ```
456/// Gets converted to:
457/// ```
458/// %0 = arith.extsi %0 : vector<32x16xi8> to vector<32x16xi32>
459/// %1 = vector.broadcast %arg0 : vector<32x16xi32> to vector<8x32x16xi32>
460/// ```
461struct ReorderCastOpsOnBroadcast
462 : public OpInterfaceRewritePattern<CastOpInterface> {
463 using OpInterfaceRewritePattern<CastOpInterface>::OpInterfaceRewritePattern;
464
465 LogicalResult matchAndRewrite(CastOpInterface op,
466 PatternRewriter &rewriter) const override {
467 if (op->getNumOperands() != 1)
468 return failure();
469 if (!isa<VectorType>(op->getResult(0).getType()))
470 return failure();
471 auto bcastOp = op->getOperand(0).getDefiningOp<vector::BroadcastOp>();
472 if (!bcastOp)
473 return failure();
474
475 Type castResTy = getElementTypeOrSelf(op->getResult(0));
476 if (auto vecTy = dyn_cast<VectorType>(bcastOp.getSourceType()))
477 castResTy = vecTy.clone(castResTy);
478 auto *castOp =
479 rewriter.create(op->getLoc(), op->getName().getIdentifier(),
480 bcastOp.getSource(), castResTy, op->getAttrs());
481 rewriter.replaceOpWithNewOp<vector::BroadcastOp>(
482 op, op->getResult(0).getType(), castOp->getResult(0));
483 return success();
484 }
485};
486
487/// Reorders elementwise(transpose) to transpose(elementwise). This makes
488/// transpose ops and contraction ops closer, which kicks in
489/// CombineContractABTranspose pattern when elementwise ops are between these
490/// operations. Ex:
491/// ```
492/// %at = vector.transpose %a, [1, 0]: vector<4x2xf32> to vector<2x4xf32>
493/// %bt = vector.transpose %b, [1, 0]: vector<4x2xf32> to vector<2x4xf32>
494/// %r = arith.addf %at, %bt : vector<2x4xf32>
495/// ```
496/// Gets converted to:
497/// ```
498/// %0 = arith.addf %a, %b : vector<4x2xf32>
499/// %r = vector.transpose %0, [1, 0] : vector<2x4xf32>
500/// ```
501struct ReorderElementwiseOpsOnTranspose final
502 : public OpTraitRewritePattern<OpTrait::Elementwise> {
504 LogicalResult matchAndRewrite(Operation *op,
505 PatternRewriter &rewriter) const override {
506 if (op->getNumResults() != 1 || op->getNumRegions() != 0)
507 return failure();
508
509 // Make sure all operands are transpose/constant ops and collect their
510 // transposition maps.
511 SmallVector<ArrayRef<int64_t>> transposeMaps;
512 transposeMaps.reserve(op->getNumOperands());
513 // Record the initial type before transposition. We'll use its shape later.
514 // Any type will do here as we will check all transpose maps are the same.
515 VectorType srcType;
516 for (Value operand : op->getOperands()) {
517 auto transposeOp = operand.getDefiningOp<vector::TransposeOp>();
518 if (transposeOp) {
519 transposeMaps.push_back(transposeOp.getPermutation());
520 srcType = transposeOp.getSourceVectorType();
521 } else if (!matchPattern(operand, m_Constant())) {
522 return failure();
523 }
524 }
525 if (transposeMaps.empty())
526 return failure();
527 // This is an elementwise op, so all transposed operands should have the
528 // same type. We need to additionally check that all transposes uses the
529 // same map.
530 if (!llvm::all_equal(transposeMaps))
531 return rewriter.notifyMatchFailure(op, "different transpose map");
532
533 SmallVector<Value> srcValues;
534 srcValues.reserve(op->getNumOperands());
535
536 // If there are constant operands, we need to insert inverse transposes for
537 // them. Calculate the inverse order first.
538 auto order = transposeMaps.front();
539 SmallVector<int64_t> invOrder(order.size());
540 for (int i = 0, e = order.size(); i < e; ++i)
541 invOrder[order[i]] = i;
542
543 for (Value operand : op->getOperands()) {
544 auto transposeOp = operand.getDefiningOp<vector::TransposeOp>();
545 if (transposeOp) {
546 srcValues.push_back(transposeOp.getVector());
547 } else {
548 // This is a constant. Create a reverse transpose op for it.
549 auto vectorType =
550 srcType.clone(cast<VectorType>(operand.getType()).getElementType());
551 srcValues.push_back(vector::TransposeOp::create(
552 rewriter, operand.getLoc(), vectorType, operand, invOrder));
553 }
554 }
555
556 auto vectorType = srcType.clone(
557 cast<VectorType>(op->getResultTypes()[0]).getElementType());
558 Operation *elementwiseOp =
559 rewriter.create(op->getLoc(), op->getName().getIdentifier(), srcValues,
560 vectorType, op->getAttrs());
561 rewriter.replaceOpWithNewOp<vector::TransposeOp>(
562 op, op->getResultTypes()[0], elementwiseOp->getResult(0),
563 transposeMaps.front());
564 return success();
565 }
566};
567
568// Returns the values in `arrayAttr` as an integer vector.
569static SmallVector<int64_t> getIntValueVector(ArrayAttr arrayAttr) {
570 return llvm::map_to_vector<4>(arrayAttr.getAsRange<IntegerAttr>(),
571 [](IntegerAttr attr) { return attr.getInt(); });
572}
573
574// Shuffles vector.bitcast op after vector.extract op.
575//
576// This transforms IR like:
577// %0 = vector.bitcast %src : vector<4xf32> to vector<8xf16>
578// %1 = vector.extract %0[3] : f16 from vector<8xf16>
579// Into:
580// %0 = vector.extract %src[1] : f32 from vector<4xf32>
581// %1 = vector.bitcast %0: vector<1xf32> to vector<2xf16>
582// %2 = vector.extract %1[1] : f16 from vector<2xf16>
583struct BubbleDownVectorBitCastForExtract
584 : public OpRewritePattern<vector::ExtractOp> {
585 using Base::Base;
586
587 LogicalResult matchAndRewrite(vector::ExtractOp extractOp,
588 PatternRewriter &rewriter) const override {
589 // Only support extracting scalars for now.
590 if (extractOp.getSourceVectorType().getRank() != 1)
591 return failure();
592
593 auto castOp = extractOp.getSource().getDefiningOp<vector::BitCastOp>();
594 if (!castOp)
595 return failure();
596
597 VectorType castSrcType = castOp.getSourceVectorType();
598 VectorType castDstType = castOp.getResultVectorType();
599 assert(castSrcType.getRank() == castDstType.getRank());
600
601 // Fail to match if we only have one element in the cast op source.
602 // This is to avoid infinite loop given that this pattern can generate
603 // such cases.
604 if (castSrcType.getNumElements() == 1)
605 return failure();
606
607 // Only support casting to a larger number of elements or now.
608 // E.g., vector<4xf32> -> vector<8xf16>.
609 if (castSrcType.getNumElements() > castDstType.getNumElements())
610 return failure();
611
612 unsigned expandRatio =
613 castDstType.getNumElements() / castSrcType.getNumElements();
614
615 // Get the first element of the mixed position as integer.
616 auto mixedPos = extractOp.getMixedPosition();
617 if (!mixedPos.empty() && !isa<Attribute>(mixedPos[0]))
618 return failure();
619 uint64_t index = cast<IntegerAttr>(cast<Attribute>(mixedPos[0])).getInt();
620
621 // Get the single scalar (as a vector) in the source value that packs the
622 // desired scalar. E.g. extract vector<1xf32> from vector<4xf32>
623 Location loc = extractOp.getLoc();
624 Value packedValue = vector::ExtractOp::create(
625 rewriter, loc, castOp.getSource(), index / expandRatio);
626 Type packedVecType = VectorType::get(/*shape=*/{1}, packedValue.getType());
627 Value zero = arith::ConstantOp::create(rewriter, loc, packedVecType,
628 rewriter.getZeroAttr(packedVecType));
629 packedValue = vector::InsertOp::create(rewriter, loc, packedValue, zero,
630 /*position=*/0);
631
632 // Cast it to a vector with the desired scalar's type.
633 // E.g. f32 -> vector<2xf16>
634 VectorType packedType =
635 VectorType::get({expandRatio}, castDstType.getElementType());
636 Value castedValue =
637 vector::BitCastOp::create(rewriter, loc, packedType, packedValue);
638
639 // Finally extract the desired scalar.
640 rewriter.replaceOpWithNewOp<vector::ExtractOp>(extractOp, castedValue,
641 index % expandRatio);
642 return success();
643 }
644};
645
646// Shuffles vector.bitcast op after vector.extract_strided_slice op.
647//
648// This transforms IR like:
649// %cast = vector.bitcast %arg0: vector<4xf32> to vector<8xf16>
650// %0 = vector.extract_strided_slice %cast {
651// offsets = [4], sizes = [4], strides = [1]
652// } : vector<8xf16> to vector<4xf16>
653// Into:
654// %0 = vector.extract_strided_slice %src {
655// offsets = [2], sizes = [2], strides = [1]
656// } : vector<4xf32> to vector<2xf32>
657// %1 = vector.bitcast %0 : vector<2xf32> to vector<4xf16>
658struct BubbleDownBitCastForStridedSliceExtract
659 : public OpRewritePattern<vector::ExtractStridedSliceOp> {
660 using Base::Base;
661
662 LogicalResult matchAndRewrite(vector::ExtractStridedSliceOp extractOp,
663 PatternRewriter &rewriter) const override {
664 auto castOp = extractOp.getSource().getDefiningOp<vector::BitCastOp>();
665 if (!castOp)
666 return failure();
667
668 VectorType castSrcType = castOp.getSourceVectorType();
669 VectorType castDstType = castOp.getResultVectorType();
670 assert(castSrcType.getRank() == castDstType.getRank());
671
672 int64_t castSrcLastDim = castSrcType.getShape().back();
673 int64_t castDstLastDim = castDstType.getShape().back();
674 // Require casting to more elements for now; other cases to be implemented.
675 if (castSrcLastDim > castDstLastDim)
676 return failure();
677
678 // Only accept all one strides for now.
679 if (llvm::any_of(extractOp.getStrides().getAsValueRange<IntegerAttr>(),
680 [](const APInt &val) { return !val.isOne(); }))
681 return failure();
682
683 unsigned rank = extractOp.getSourceVectorType().getRank();
684 assert(castDstLastDim % castSrcLastDim == 0);
685 int64_t expandRatio = castDstLastDim / castSrcLastDim;
686
687 // If we have a less number of offsets than the rank, then implicitly we
688 // are selecting the full range for the last bitcasted dimension; other
689 // dimensions aren't affected. Otherwise, we need to scale down the last
690 // dimension's offset given we are extracting from less elements now.
691 ArrayAttr newOffsets = extractOp.getOffsets();
692 if (newOffsets.size() == rank) {
693 SmallVector<int64_t> offsets = getIntValueVector(newOffsets);
694 if (offsets.back() % expandRatio != 0)
695 return failure();
696 offsets.back() = offsets.back() / expandRatio;
697 newOffsets = rewriter.getI64ArrayAttr(offsets);
698 }
699
700 // Similarly for sizes.
701 ArrayAttr newSizes = extractOp.getSizes();
702 if (newSizes.size() == rank) {
703 SmallVector<int64_t> sizes = getIntValueVector(newSizes);
704 if (sizes.back() % expandRatio != 0)
705 return failure();
706 sizes.back() = sizes.back() / expandRatio;
707 newSizes = rewriter.getI64ArrayAttr(sizes);
708 }
709
710 SmallVector<int64_t> dims =
711 llvm::to_vector<4>(cast<VectorType>(extractOp.getType()).getShape());
712 dims.back() = dims.back() / expandRatio;
713 VectorType newExtractType =
714 VectorType::get(dims, castSrcType.getElementType());
715
716 auto newExtractOp = vector::ExtractStridedSliceOp::create(
717 rewriter, extractOp.getLoc(), newExtractType, castOp.getSource(),
718 newOffsets, newSizes, extractOp.getStrides());
719
720 rewriter.replaceOpWithNewOp<vector::BitCastOp>(
721 extractOp, extractOp.getType(), newExtractOp);
722
723 return success();
724 }
725};
726
727// Shuffles vector.bitcast op before vector.insert_strided_slice op.
728//
729// This transforms IR like:
730// %0 = vector.insert %val, %dst[4] : vector<32xi4> into vector<8x32xi4>
731// %1 = vector.bitcast %0 : vector<8x32xi4> to vector<8x16xi8>
732// Into:
733// %0 = vector.bitcast %val : vector<32xi4> to vector<16xi8>
734// %1 = vector.bitcast %dst : vector<8x32xi4> to vector<8x16xi8>
735// %2 = vector.insert %0, %1 [4] : vector<16xi8> into vector<8x16xi8>
736//
737struct BubbleUpBitCastForInsert : public OpRewritePattern<vector::BitCastOp> {
738 using Base::Base;
739
740 LogicalResult matchAndRewrite(vector::BitCastOp bitcastOp,
741 PatternRewriter &rewriter) const override {
742 VectorType castSrcType = bitcastOp.getSourceVectorType();
743 VectorType castDstType = bitcastOp.getResultVectorType();
744
745 // 0-D and scalable vectors are not supported yet.
746 if (castSrcType.getRank() == 0 || castSrcType.isScalable() ||
747 castDstType.isScalable())
748 return failure();
749
750 int64_t castSrcLastDim = castSrcType.getShape().back();
751 int64_t castDstLastDim = castDstType.getShape().back();
752 bool isNumElemsShrink = castSrcLastDim >= castDstLastDim;
753 int64_t ratio;
754 if (isNumElemsShrink) {
755 assert(castSrcLastDim % castDstLastDim == 0);
756 ratio = castSrcLastDim / castDstLastDim;
757 } else {
758 assert(castDstLastDim % castSrcLastDim == 0);
759 ratio = castDstLastDim / castSrcLastDim;
760 }
761
762 auto insertOp = bitcastOp.getSource().getDefiningOp<vector::InsertOp>();
763 if (!insertOp)
764 return failure();
765
766 // Only vector sources are supported for now.
767 auto insertSrcType = dyn_cast<VectorType>(insertOp.getValueToStoreType());
768 if (!insertSrcType)
769 return failure();
770
771 // Bitcast the source.
772 SmallVector<int64_t> srcDims(insertSrcType.getShape());
773 srcDims.back() =
774 isNumElemsShrink ? srcDims.back() / ratio : srcDims.back() * ratio;
775 VectorType newCastSrcType =
776 VectorType::get(srcDims, castDstType.getElementType());
777 auto newCastSrcOp =
778 vector::BitCastOp::create(rewriter, bitcastOp.getLoc(), newCastSrcType,
779 insertOp.getValueToStore());
780
781 SmallVector<int64_t> dstDims(insertOp.getDestVectorType().getShape());
782 dstDims.back() =
783 isNumElemsShrink ? dstDims.back() / ratio : dstDims.back() * ratio;
784 VectorType newCastDstType =
785 VectorType::get(dstDims, castDstType.getElementType());
786
787 // Bitcast the destination.
788 auto newCastDstOp = vector::BitCastOp::create(
789 rewriter, bitcastOp.getLoc(), newCastDstType, insertOp.getDest());
790
791 // Generate new insert.
792 rewriter.replaceOpWithNewOp<vector::InsertOp>(
793 bitcastOp, newCastSrcOp, newCastDstOp, insertOp.getMixedPosition());
794 return success();
795 }
796};
797
798// Shuffles vector.bitcast op before vector.insert_strided_slice op.
799//
800// This transforms IR like:
801// %0 = vector.insert_strided_slice %src, %dst {
802// offsets = [0], strides = [1]} : vector<4xf16> into vector<8xf16>
803// %1 = vector.bitcast %0: vector<8xf16> to vector<4xf32>
804// Into:
805// %0 = vector.bitcast %src : vector<4xf16> to vector<2xf32>
806// %1 = vector.bitcast %dst : vector<8xf16> to vector<4xf32>
807// %2 = vector.insert_strided_slice %src, %dst {
808// offsets = [0], strides = [1]} : vector<2xf32> into vector<4xf32>
809struct BubbleUpBitCastForStridedSliceInsert
810 : public OpRewritePattern<vector::BitCastOp> {
811 using Base::Base;
812
813 LogicalResult matchAndRewrite(vector::BitCastOp bitcastOp,
814 PatternRewriter &rewriter) const override {
815 VectorType castSrcType = bitcastOp.getSourceVectorType();
816 VectorType castDstType = bitcastOp.getResultVectorType();
817 assert(castSrcType.getRank() == castDstType.getRank());
818 // Skip 0-D vector which will not from InsertStridedSliceOp.
819 if (castSrcType.getRank() == 0)
820 return failure();
821
822 int64_t castSrcLastDim = castSrcType.getShape().back();
823 int64_t castDstLastDim = castDstType.getShape().back();
824 // Require casting to less elements for now; other cases to be implemented.
825 if (castSrcLastDim < castDstLastDim)
826 return failure();
827
828 assert(castSrcLastDim % castDstLastDim == 0);
829 int64_t shrinkRatio = castSrcLastDim / castDstLastDim;
830
831 auto insertOp =
832 bitcastOp.getSource().getDefiningOp<vector::InsertStridedSliceOp>();
833 if (!insertOp)
834 return failure();
835
836 // Only accept all one strides for now.
837 if (llvm::any_of(insertOp.getStrides().getAsValueRange<IntegerAttr>(),
838 [](const APInt &val) { return !val.isOne(); }))
839 return failure();
840
841 unsigned rank = insertOp.getSourceVectorType().getRank();
842 // Require insert op to have the same rank for the source and destination
843 // vector; other cases to be implemented.
844 if (rank != insertOp.getDestVectorType().getRank())
845 return failure();
846
847 // Requires that shape of insert op src is castable to dstType.
848 unsigned sourceWidth = castSrcType.getElementType().getIntOrFloatBitWidth();
849 unsigned destinationWidth =
850 castDstType.getElementType().getIntOrFloatBitWidth();
851 unsigned numElements = destinationWidth / sourceWidth;
852 if (insertOp.getSourceVectorType().getNumElements() % numElements != 0)
853 return failure();
854
855 ArrayAttr newOffsets = insertOp.getOffsets();
856 assert(newOffsets.size() == rank);
857 SmallVector<int64_t> offsets = getIntValueVector(newOffsets);
858 if (offsets.back() % shrinkRatio != 0)
859 return failure();
860 offsets.back() = offsets.back() / shrinkRatio;
861 newOffsets = rewriter.getI64ArrayAttr(offsets);
862
863 SmallVector<int64_t> srcDims =
864 llvm::to_vector<4>(insertOp.getSourceVectorType().getShape());
865 srcDims.back() = srcDims.back() / shrinkRatio;
866 VectorType newCastSrcType =
867 VectorType::get(srcDims, castDstType.getElementType());
868
869 auto newCastSrcOp =
870 vector::BitCastOp::create(rewriter, bitcastOp.getLoc(), newCastSrcType,
871 insertOp.getValueToStore());
872
873 SmallVector<int64_t> dstDims =
874 llvm::to_vector<4>(insertOp.getDestVectorType().getShape());
875 dstDims.back() = dstDims.back() / shrinkRatio;
876 VectorType newCastDstType =
877 VectorType::get(dstDims, castDstType.getElementType());
878
879 auto newCastDstOp = vector::BitCastOp::create(
880 rewriter, bitcastOp.getLoc(), newCastDstType, insertOp.getDest());
881
882 rewriter.replaceOpWithNewOp<vector::InsertStridedSliceOp>(
883 bitcastOp, bitcastOp.getType(), newCastSrcOp, newCastDstOp, newOffsets,
884 insertOp.getStrides());
885
886 return success();
887 }
888};
889
890// Breaks down vector.bitcast op
891//
892// This transforms IR like:
893// %1 = vector.bitcast %0: vector<8xf16> to vector<4xf32>
894// Into:
895// %cst = vector.broadcast %c0_f32 : f32 to vector<4xf32>
896// %1 = vector.extract_strided_slice %0 {
897// offsets = [0], sizes = [4], strides = [1]
898// } : vector<8xf16> to vector<4xf16>
899// %2 = vector.bitcast %1 : vector<4xf16> to vector<2xf32>
900// %4 = vector.insert_strided_slice %2, %cst {
901// offsets = [0], strides = [1]} : vector<2xf32> into vector<4xf32>
902// %5 = vector.extract_strided_slice %0 {
903// offsets = [4], sizes = [4], strides = [1]
904// } : vector<8xf16> to vector<4xf16>
905// %6 = vector.bitcast %5 : vector<4xf16> to vector<2xf32>
906// %7 = vector.insert_strided_slice %6, %cst {
907// offsets = [2], strides = [1]} : vector<2xf32> into vector<4xf32>
908struct BreakDownVectorBitCast : public OpRewritePattern<vector::BitCastOp> {
909 using Base::Base;
910
911public:
912 BreakDownVectorBitCast(MLIRContext *context,
913 std::function<bool(vector::BitCastOp)> controlFn,
914 PatternBenefit benefit)
915 : OpRewritePattern(context, benefit), controlFn(std::move(controlFn)) {}
916
917 LogicalResult matchAndRewrite(vector::BitCastOp bitcastOp,
918 PatternRewriter &rewriter) const override {
919
920 if (controlFn && !controlFn(bitcastOp))
921 return failure();
922
923 VectorType castSrcType = bitcastOp.getSourceVectorType();
924 VectorType castDstType = bitcastOp.getResultVectorType();
925 assert(castSrcType.getRank() == castDstType.getRank());
926
927 // This transformation builds on top of
928 // vector.{extract|insert}_strided_slice, which do not support
929 // extracting/inserting "scallable sub-vectors". Bail out.
930 if (castSrcType.isScalable())
931 return rewriter.notifyMatchFailure(bitcastOp,
932 "Scalable vectors are not supported");
933
934 // Only support rank 1 case for now.
935 if (castSrcType.getRank() != 1)
936 return failure();
937
938 int64_t castSrcLastDim = castSrcType.getShape().back();
939 int64_t castDstLastDim = castDstType.getShape().back();
940 // Require casting to less elements for now; other cases to be implemented.
941 if (castSrcLastDim < castDstLastDim)
942 return failure();
943
944 assert(castSrcLastDim % castDstLastDim == 0);
945 int64_t shrinkRatio = castSrcLastDim / castDstLastDim;
946 // Nothing to do if it is already bitcasting to a single element.
947 if (castSrcLastDim == shrinkRatio)
948 return failure();
949
950 Location loc = bitcastOp.getLoc();
951 Type elemType = castDstType.getElementType();
952 assert(elemType.isSignlessIntOrIndexOrFloat());
953
954 Value zero = arith::ConstantOp::create(rewriter, loc, elemType,
955 rewriter.getZeroAttr(elemType));
956 Value res = BroadcastOp::create(rewriter, loc, castDstType, zero);
957
958 SmallVector<int64_t> sliceShape = {castDstLastDim};
959 SmallVector<int64_t> strides = {1};
960 VectorType newCastDstType =
961 VectorType::get(SmallVector<int64_t>{castDstLastDim / shrinkRatio},
962 castDstType.getElementType());
963
964 for (int i = 0, e = shrinkRatio; i < e; ++i) {
965 Value extracted = ExtractStridedSliceOp::create(
966 rewriter, loc, bitcastOp.getSource(),
967 ArrayRef<int64_t>{i * castDstLastDim}, sliceShape, strides);
968 Value bitcast =
969 BitCastOp::create(rewriter, loc, newCastDstType, extracted);
970 res = InsertStridedSliceOp::create(
971 rewriter, loc, bitcast, res,
972 ArrayRef<int64_t>{i * castDstLastDim / shrinkRatio}, strides);
973 }
974 rewriter.replaceOp(bitcastOp, res);
975 return success();
976 }
977
978private:
979 std::function<bool(BitCastOp)> controlFn;
980};
981
982static bool haveSameShapeAndScaling(Type t, Type u) {
983 auto tVec = dyn_cast<VectorType>(t);
984 auto uVec = dyn_cast<VectorType>(u);
985 if (!tVec) {
986 return !uVec;
987 }
988 if (!uVec) {
989 return false;
990 }
991 return tVec.getShape() == uVec.getShape() &&
992 tVec.getScalableDims() == uVec.getScalableDims();
993}
994
995/// If `type` is shaped, clone it with `newElementType`. Otherwise,
996/// return `newElementType`.
997static Type cloneOrReplace(Type type, Type newElementType) {
998 if (auto shapedType = dyn_cast<ShapedType>(type)) {
999 return shapedType.clone(newElementType);
1000 }
1001 return newElementType;
1002}
1003
1004/// If `value` is the result of a broadcast operation, return the input
1005/// of the broadcast operation.
1006static Value getBroadcastLikeSource(Value value) {
1007
1008 Operation *op = value.getDefiningOp();
1009 if (!op)
1010 return {};
1011
1012 if (auto broadcast = dyn_cast<vector::BroadcastOp>(op))
1013 return broadcast.getSource();
1014
1015 return {};
1016}
1017
1018/// Reorders elementwise(broadcast) to broadcast(elementwise). Ex:
1019///
1020/// Example:
1021/// ```
1022/// %a = vector.broadcast %arg1 : index to vector<1x4xindex>
1023/// %b = vector.broadcast %arg2 : index to vector<1x4xindex>
1024/// %r = arith.addi %a, %b : vector<1x4xindex>
1025/// ```
1026/// Gets converted to:
1027/// ```
1028/// %r = arith.addi %arg0, %arg1 : index
1029/// %b = vector.broadcast %r : index to vector<1x4xindex>
1030/// ```
1031struct ReorderElementwiseOpsOnBroadcast final
1032 : public OpTraitRewritePattern<OpTrait::Elementwise> {
1034 LogicalResult matchAndRewrite(Operation *op,
1035 PatternRewriter &rewriter) const override {
1036 if (op->getNumResults() != 1)
1037 return failure();
1038 auto resultType = dyn_cast<VectorType>(op->getResult(0).getType());
1039 if (!resultType)
1040 return failure();
1042 return rewriter.notifyMatchFailure(
1043 op, "Op doesn't have ElementwiseMappableTraits");
1044 if (op->getNumOperands() == 0)
1045 return failure();
1046
1047 Type resultElemType = resultType.getElementType();
1048
1049 // Get the type of the first non-constant operand
1050 Value broadcastSource;
1051 for (Value operand : op->getOperands()) {
1052 Operation *definingOp = operand.getDefiningOp();
1053 if (!definingOp)
1054 return failure();
1055 if (definingOp->hasTrait<OpTrait::ConstantLike>())
1056 continue;
1057 broadcastSource = getBroadcastLikeSource(operand);
1058 break;
1059 }
1060 if (!broadcastSource)
1061 return failure();
1062 Type unbroadcastResultType =
1063 cloneOrReplace(broadcastSource.getType(), resultElemType);
1064
1065 // Some ops, e.g. `vector.fma`, only accept vector types. For such ops the
1066 // reordering is only possible when the broadcast source is a vector as
1067 // well; sinking past a broadcast from a scalar would create an invalid op.
1068 // TODO: It may be better to support scalar sources by promoting the scalar
1069 // to a single element vector.
1070 if (isa<vector::FMAOp>(op) && !isa<VectorType>(unbroadcastResultType)) {
1071 return rewriter.notifyMatchFailure(
1072 op, "Op only accepts vector types, but the broadcast source is a "
1073 "scalar");
1074 }
1075
1076 // Make sure that all operands are broadcast from identically-shaped types:
1077 // * scalar (`vector.broadcast`), or
1078 // * vector (`vector.broadcast`).
1079 // Otherwise the re-ordering wouldn't be safe.
1080 if (!llvm::all_of(op->getOperands(), [broadcastSource](Value val) {
1081 if (auto source = getBroadcastLikeSource(val))
1082 return haveSameShapeAndScaling(source.getType(),
1083 broadcastSource.getType());
1084 SplatElementsAttr splatConst;
1085 return matchPattern(val, m_Constant(&splatConst));
1086 })) {
1087 return rewriter.notifyMatchFailure(
1088 op,
1089 "not all operands are constants or broadcasts from the same type");
1090 }
1091
1092 // Collect the source values before broadcasting
1093 SmallVector<Value> srcValues;
1094 srcValues.reserve(op->getNumOperands());
1095 for (Value operand : op->getOperands()) {
1096 SplatElementsAttr splatConst;
1097 if (matchPattern(operand, m_Constant(&splatConst))) {
1098 Attribute newConst;
1099 Type elementType = getElementTypeOrSelf(operand.getType());
1100 Type newType = cloneOrReplace(unbroadcastResultType, elementType);
1101 if (auto newTypeShaped = dyn_cast<ShapedType>(newType)) {
1102 newConst = splatConst.resizeSplat(newTypeShaped);
1103 } else {
1104 newConst = splatConst.getSplatValue<Attribute>();
1105 }
1106 Operation *newConstOp =
1107 operand.getDefiningOp()->getDialect()->materializeConstant(
1108 rewriter, newConst, newType, operand.getLoc());
1109 srcValues.push_back(newConstOp->getResult(0));
1110 } else {
1111 srcValues.push_back(operand.getDefiningOp()->getOperand(0));
1112 }
1113 }
1114
1115 // Create the "elementwise" Op
1116 Operation *elementwiseOp =
1117 rewriter.create(op->getLoc(), op->getName().getIdentifier(), srcValues,
1118 unbroadcastResultType, op->getAttrs());
1119
1120 // Replace the original Op with the elementwise Op
1121 rewriter.replaceOpWithNewOp<vector::BroadcastOp>(
1122 op, resultType, elementwiseOp->getResults());
1123
1124 return success();
1125 }
1126};
1127
1128/// Pattern to rewrite a ExtractOp(Elementwise) -> Elementwise(ExtractOp).
1129/// This may result in cleaner code when extracting a single value
1130/// from multi-element vector and also to help canonicalize 1-element vectors to
1131/// scalars.
1132///
1133/// Example:
1134/// ```
1135/// %0 = arith.addf %arg0, %arg1 : vector<4xf32>
1136/// %1 = vector.extract %0[1] : f32 from vector<4xf32>
1137/// ```
1138/// Gets converted to:
1139/// ```
1140/// %0 = vector.extract %arg0[1] : f32 from vector<4xf32>
1141/// %1 = vector.extract %arg1[1] : f32 from vector<4xf32>
1142/// %2 = arith.addf %0, %1 : f32
1143/// ```
1144class ExtractOpFromElementwise final
1145 : public OpRewritePattern<vector::ExtractOp> {
1146public:
1147 using Base::Base;
1148
1149 LogicalResult matchAndRewrite(vector::ExtractOp op,
1150 PatternRewriter &rewriter) const override {
1151 Operation *eltwise = op.getSource().getDefiningOp();
1152
1153 // TODO: vector::FMAOp is not an ElemetwiseMappable even if it claims to be,
1154 // as it doesn't support scalars.
1155 if (!eltwise || !OpTrait::hasElementwiseMappableTraits(eltwise) ||
1156 isa<vector::FMAOp>(eltwise))
1157 return rewriter.notifyMatchFailure(op, "not an elementwise op");
1158
1159 if (eltwise->getNumResults() != 1)
1160 return rewriter.notifyMatchFailure(op, "expected single result");
1161
1162 if (!eltwise->hasOneUse())
1163 return rewriter.notifyMatchFailure(op, "expected single op use");
1164
1165 if (!llvm::all_equal(eltwise->getOperandTypes()))
1166 return rewriter.notifyMatchFailure(op, "operand types are different");
1167
1168 // Dynamic position can cause dominance issues, so conservatively fail for
1169 // now.
1170 if (!op.getDynamicPosition().empty())
1171 return rewriter.notifyMatchFailure(
1172 op, "dynamic position not yet implemented");
1173
1174 Type dstType = op.getType();
1175
1176 OpBuilder::InsertionGuard g(rewriter);
1177 rewriter.setInsertionPoint(eltwise);
1178
1179 IRMapping mapping;
1180 Location loc = eltwise->getLoc();
1181 SmallVector<OpFoldResult> pos = op.getMixedPosition();
1182 for (Value arg : eltwise->getOperands()) {
1183 Value newArg = vector::ExtractOp::create(rewriter, loc, arg, pos);
1184 mapping.map(arg, newArg);
1185 }
1186
1187 Operation *newEltwise = rewriter.clone(*eltwise, mapping);
1188 newEltwise->getResult(0).setType(dstType);
1189
1190 rewriter.replaceOp(op, newEltwise);
1191 rewriter.eraseOp(eltwise);
1192 return success();
1193 }
1194};
1195
1196/// Check if the element type is suitable for vector.load/store sinking.
1197/// Element type must be index or byte-aligned integer or floating-point type.
1198static bool isSupportedMemSinkElementType(Type type) {
1199 if (isa<IndexType>(type))
1200 return true;
1201
1202 return type.isIntOrFloat() && type.getIntOrFloatBitWidth() % 8 == 0;
1203}
1204
1205/// Pattern to rewrite `vector.extract(vector.load) -> vector/memref.load.
1206/// Only index and byte-aligned integer and floating-point element types are
1207/// supported for now.
1208///
1209/// Example:
1210/// ```
1211/// vector.load %arg0[%arg1] : memref<?xf32>, vector<4xf32>
1212/// vector.extract %0[1] : f32 from vector<4xf32>
1213/// ```
1214/// Gets converted to:
1215/// ```
1216/// %c1 = arith.constant 1 : index
1217/// %0 = arith.addi %arg1, %c1 overflow<nsw> : index
1218/// %1 = memref.load %arg0[%0] : memref<?xf32>
1219/// ```
1220class ExtractOpFromLoad final : public OpRewritePattern<vector::ExtractOp> {
1221public:
1222 using Base::Base;
1223
1224 LogicalResult matchAndRewrite(vector::ExtractOp op,
1225 PatternRewriter &rewriter) const override {
1226 auto loadOp = op.getSource().getDefiningOp<vector::LoadOp>();
1227 if (!loadOp)
1228 return rewriter.notifyMatchFailure(op, "expected a load op");
1229
1230 // Checking for single use so we won't duplicate load ops.
1231 if (!loadOp->hasOneUse())
1232 return rewriter.notifyMatchFailure(op, "expected single op use");
1233
1234 VectorType loadVecType = loadOp.getVectorType();
1235 if (loadVecType.isScalable())
1236 return rewriter.notifyMatchFailure(op,
1237 "scalable vectors are not supported");
1238
1239 MemRefType memType = loadOp.getMemRefType();
1240
1241 // Non-byte-aligned types are tricky and may require special handling,
1242 // ignore them for now.
1243 if (!isSupportedMemSinkElementType(memType.getElementType()))
1244 return rewriter.notifyMatchFailure(op, "unsupported element type");
1245
1246 int64_t rankOffset = memType.getRank() - loadVecType.getRank();
1247 if (rankOffset < 0)
1248 return rewriter.notifyMatchFailure(op, "unsupported ranks combination");
1249
1250 auto extractVecType = dyn_cast<VectorType>(op.getResult().getType());
1251 int64_t finalRank = 0;
1252 if (extractVecType)
1253 finalRank = extractVecType.getRank();
1254
1255 SmallVector<Value> indices = loadOp.getIndices();
1256 SmallVector<OpFoldResult> extractPos = op.getMixedPosition();
1257
1258 // There may be memory stores between the load and the extract op, so we
1259 // need to make sure that the new load op is inserted at the same place as
1260 // the original load op.
1261 OpBuilder::InsertionGuard g(rewriter);
1262 rewriter.setInsertionPoint(loadOp);
1263 Location loc = loadOp.getLoc();
1264 ArithIndexingBuilder idxBuilderf(rewriter, loc);
1265 for (auto i : llvm::seq<int64_t>(rankOffset, indices.size() - finalRank)) {
1266 OpFoldResult pos = extractPos[i - rankOffset];
1267 if (isZeroInteger(pos))
1268 continue;
1269
1270 Value offset = getValueOrCreateConstantIndexOp(rewriter, loc, pos);
1271 indices[i] = idxBuilderf.add(indices[i], offset);
1272 }
1273
1274 Value base = loadOp.getBase();
1275 if (extractVecType) {
1276 rewriter.replaceOpWithNewOp<vector::LoadOp>(op, extractVecType, base,
1277 indices);
1278 } else {
1279 rewriter.replaceOpWithNewOp<memref::LoadOp>(op, base, indices);
1280 }
1281 // We checked for single use so we can safely erase the load op.
1282 rewriter.eraseOp(loadOp);
1283 return success();
1284 }
1285};
1286
1287/// Pattern to rewrite vector.store(vector.broadcast) -> vector/memref.store.
1288///
1289/// Example:
1290/// ```
1291/// %0 = vector.broadcast %arg2 : f32 to vector<1xf32>
1292/// vector.store %0, %arg0[%arg1] : memref<?xf32>, vector<1xf32>
1293/// ```
1294/// Gets converted to:
1295/// ```
1296/// memref.store %arg2, %arg0[%arg1] : memref<?xf32>
1297/// ```
1298class StoreOpFromBroadcast final : public OpRewritePattern<vector::StoreOp> {
1299public:
1300 using Base::Base;
1301
1302 LogicalResult matchAndRewrite(vector::StoreOp op,
1303 PatternRewriter &rewriter) const override {
1304 VectorType vecType = op.getVectorType();
1305 if (vecType.isScalable())
1306 return rewriter.notifyMatchFailure(op,
1307 "scalable vectors are not supported");
1308
1309 if (isa<VectorType>(op.getMemRefType().getElementType()))
1310 return rewriter.notifyMatchFailure(
1311 op, "memrefs of vectors are not supported");
1312
1313 if (vecType.getNumElements() != 1)
1314 return rewriter.notifyMatchFailure(
1315 op, "only 1-element vectors are supported");
1316
1317 Value toStore = op.getValueToStore();
1318 Value source = getBroadcastLikeSource(toStore);
1319 if (!source)
1320 return rewriter.notifyMatchFailure(
1321 op, "value to store is not from a broadcast");
1322
1323 // Checking for single use so we can remove broadcast.
1324 Operation *broadcast = toStore.getDefiningOp();
1325 if (!broadcast->hasOneUse())
1326 return rewriter.notifyMatchFailure(op, "expected single op use");
1327
1328 Value base = op.getBase();
1329 ValueRange indices = op.getIndices();
1330
1331 if (isa<VectorType>(source.getType())) {
1332 rewriter.replaceOpWithNewOp<vector::StoreOp>(op, source, base, indices);
1333 } else {
1334 rewriter.replaceOpWithNewOp<memref::StoreOp>(op, source, base, indices);
1335 }
1336 rewriter.eraseOp(broadcast);
1337 return success();
1338 }
1339};
1340
1341// Helper that returns a vector comparison that constructs a mask:
1342// mask = [0,1,..,n-1] + [o,o,..,o] < [b,b,..,b]
1343//
1344// If `dim == 0` then the result will be a 0-D vector.
1345//
1346// NOTE: The LLVM::GetActiveLaneMaskOp intrinsic would provide an alternative,
1347// much more compact, IR for this operation, but LLVM eventually
1348// generates more elaborate instructions for this intrinsic since it
1349// is very conservative on the boundary conditions.
1350static Value buildVectorComparison(PatternRewriter &rewriter, Operation *op,
1351 bool force32BitVectorIndices, int64_t dim,
1352 Value b, Value *off = nullptr) {
1353 auto loc = op->getLoc();
1354 // If we can assume all indices fit in 32-bit, we perform the vector
1355 // comparison in 32-bit to get a higher degree of SIMD parallelism.
1356 // Otherwise we perform the vector comparison using 64-bit indices.
1357 Type idxType =
1358 force32BitVectorIndices ? rewriter.getI32Type() : rewriter.getI64Type();
1359 DenseIntElementsAttr indicesAttr;
1360 if (dim == 0 && force32BitVectorIndices) {
1361 indicesAttr = DenseIntElementsAttr::get(
1362 VectorType::get(ArrayRef<int64_t>{}, idxType), ArrayRef<int32_t>{0});
1363 } else if (dim == 0) {
1364 indicesAttr = DenseIntElementsAttr::get(
1365 VectorType::get(ArrayRef<int64_t>{}, idxType), ArrayRef<int64_t>{0});
1366 } else if (force32BitVectorIndices) {
1367 indicesAttr = rewriter.getI32VectorAttr(
1368 llvm::to_vector<4>(llvm::seq<int32_t>(0, dim)));
1369 } else {
1370 indicesAttr = rewriter.getI64VectorAttr(
1371 llvm::to_vector<4>(llvm::seq<int64_t>(0, dim)));
1372 }
1373 Value indices = arith::ConstantOp::create(rewriter, loc, indicesAttr);
1374 // Add in an offset if requested.
1375 if (off) {
1376 Value o = getValueOrCreateCastToIndexLike(rewriter, loc, idxType, *off);
1377 Value ov = vector::BroadcastOp::create(rewriter, loc, indices.getType(), o);
1378 indices = arith::AddIOp::create(rewriter, loc, ov, indices);
1379 }
1380 // Construct the vector comparison.
1381 // When using 32-bit indices, cap `b` at INT32_MAX before casting to prevent
1382 // signed overflow for large index values (e.g., 2^51 wrapping to 0 in i32).
1383 // Note: for fixed-size vectors, `dim` is a tighter bound (since any b >= dim
1384 // already implies all-true), but we use INT32_MAX for uniformity with the
1385 // scalable-vector path.
1386 if (force32BitVectorIndices) {
1387 Value maxBound =
1388 arith::ConstantIndexOp::create(rewriter, loc, (1LL << 31) - 1);
1389 b = arith::MinSIOp::create(rewriter, loc, b, maxBound);
1390 }
1391 Value bound = getValueOrCreateCastToIndexLike(rewriter, loc, idxType, b);
1392 Value bounds =
1393 vector::BroadcastOp::create(rewriter, loc, indices.getType(), bound);
1394 return arith::CmpIOp::create(rewriter, loc, arith::CmpIPredicate::slt,
1395 indices, bounds);
1396}
1397
1398template <typename ConcreteOp>
1399struct MaterializeTransferMask : public OpRewritePattern<ConcreteOp> {
1400public:
1401 explicit MaterializeTransferMask(MLIRContext *context, bool enableIndexOpt,
1402 PatternBenefit benefit = 1)
1403 : mlir::OpRewritePattern<ConcreteOp>(context, benefit),
1404 force32BitVectorIndices(enableIndexOpt) {}
1405
1406 LogicalResult matchAndRewrite(ConcreteOp xferOp,
1407 PatternRewriter &rewriter) const override {
1408 if (!xferOp.hasOutOfBoundsDim())
1409 return failure();
1410
1411 if (xferOp.getVectorType().getRank() > 1 || xferOp.getIndices().empty())
1412 return failure();
1413
1414 Location loc = xferOp->getLoc();
1415 VectorType vtp = xferOp.getVectorType();
1416
1417 // Create the in-bounds mask with all elements between [0 .. dim - offset)
1418 // set and [dim - offset .. vector_length) unset.
1419 //
1420 // TODO: when the leaf transfer rank is k > 1, we need the last `k`
1421 // dimensions here.
1422 unsigned lastIndex = llvm::size(xferOp.getIndices()) - 1;
1423 Value off = xferOp.getIndices()[lastIndex];
1424 Value dim =
1425 vector::createOrFoldDimOp(rewriter, loc, xferOp.getBase(), lastIndex);
1426 Value b = arith::SubIOp::create(rewriter, loc, dim.getType(), dim, off);
1427 Value mask = vector::CreateMaskOp::create(
1428 rewriter, loc,
1429 VectorType::get(vtp.getShape(), rewriter.getI1Type(),
1430 vtp.getScalableDims()),
1431 b);
1432 if (xferOp.getMask()) {
1433 // Intersect the in-bounds with the mask specified as an op parameter.
1434 mask = arith::AndIOp::create(rewriter, loc, mask, xferOp.getMask());
1435 }
1436
1437 rewriter.modifyOpInPlace(xferOp, [&]() {
1438 xferOp.getMaskMutable().assign(mask);
1439 xferOp.setInBoundsAttr(rewriter.getBoolArrayAttr({true}));
1440 });
1441
1442 return success();
1443 }
1444
1445private:
1446 const bool force32BitVectorIndices;
1447};
1448
1449/// Conversion pattern for a `vector.create_mask` (0-D and 1-D only).
1450class VectorCreateMaskOpConversion
1451 : public OpRewritePattern<vector::CreateMaskOp> {
1452public:
1453 explicit VectorCreateMaskOpConversion(MLIRContext *context,
1454 bool enableIndexOpt,
1455 PatternBenefit benefit = 1)
1456 : mlir::OpRewritePattern<vector::CreateMaskOp>(context, benefit),
1457 force32BitVectorIndices(enableIndexOpt) {}
1458
1459 LogicalResult matchAndRewrite(vector::CreateMaskOp op,
1460 PatternRewriter &rewriter) const override {
1461 auto dstType = op.getType();
1462 if (cast<VectorType>(dstType).isScalable())
1463 return failure();
1464 int64_t rank = dstType.getRank();
1465 if (rank > 1)
1466 return failure();
1467 rewriter.replaceOp(
1468 op, buildVectorComparison(rewriter, op, force32BitVectorIndices,
1469 rank == 0 ? 0 : dstType.getDimSize(0),
1470 op.getOperand(0)));
1471 return success();
1472 }
1473
1474private:
1475 const bool force32BitVectorIndices;
1476};
1477
1478/// Returns true if all the `i1` elements of `constantOp` are set to `value`.
1479static bool allI1ConstantValuesSetTo(arith::ConstantOp constantOp, bool value) {
1480 auto denseAttr = dyn_cast<DenseIntElementsAttr>(constantOp.getValue());
1481 // TODO: Support non-dense constant.
1482 if (!denseAttr)
1483 return false;
1484
1485 assert(denseAttr.getElementType().isInteger(1) && "Unexpected type");
1486 return denseAttr.isSplat() && denseAttr.getSplatValue<bool>() == value;
1487}
1488
1489/// Folds a select operation between an all-true and all-false vector. For now,
1490/// only single element vectors (i.e., vector<1xi1>) are supported. That is:
1491///
1492/// %true = arith.constant dense<true> : vector<1xi1>
1493/// %false = arith.constant dense<false> : vector<1xi1>
1494/// %result = arith.select %cond, %true, %false : i1, vector<1xi1>
1495/// =>
1496/// %result = vector.broadcast %cond : i1 to vector<1xi1>
1497///
1498/// InstCombine seems to handle vectors with multiple elements but not the
1499/// single element ones.
1500struct FoldI1Select : public OpRewritePattern<arith::SelectOp> {
1501 using Base::Base;
1502
1503 LogicalResult matchAndRewrite(arith::SelectOp selectOp,
1504 PatternRewriter &rewriter) const override {
1505 auto vecType = dyn_cast<VectorType>(selectOp.getType());
1506 if (!vecType || !vecType.getElementType().isInteger(1))
1507 return failure();
1508
1509 // Only scalar conditions can be folded.
1510 Value cond = selectOp.getCondition();
1511 if (isa<VectorType>(cond.getType()))
1512 return failure();
1513
1514 // TODO: Support n-D and scalable vectors.
1515 if (vecType.getRank() != 1 || vecType.isScalable())
1516 return failure();
1517
1518 // TODO: Support vectors with multiple elements.
1519 if (vecType.getShape()[0] != 1)
1520 return failure();
1521
1522 auto trueConst = selectOp.getTrueValue().getDefiningOp<arith::ConstantOp>();
1523 if (!trueConst || !allI1ConstantValuesSetTo(trueConst, true))
1524 return failure();
1525
1526 auto falseConst =
1527 selectOp.getFalseValue().getDefiningOp<arith::ConstantOp>();
1528 if (!falseConst || !allI1ConstantValuesSetTo(falseConst, false))
1529 return failure();
1530
1531 // Replace select with its condition broadcasted to single element vector.
1532 auto elemType = rewriter.getIntegerType(vecType.getNumElements());
1533 auto bcastType = VectorType::get(/*shape=*/{1}, elemType);
1534 rewriter.replaceOpWithNewOp<vector::BroadcastOp>(selectOp, bcastType, cond);
1535 return success();
1536 }
1537};
1538
1539/// Returns the number of dims can be folded away from transfer ops. It returns
1540/// a failure if it can not determine the number of dims to be folded.
1541///
1542/// Ex 1: returns "2" if `srcType` is memref<512x16x1x1xf32> and
1543/// `vectorType` is vector<16x16x1x1xf32>
1544/// (there two inner most dims can be dropped by memref.subview ops)
1545///
1546/// Ex 2: returns "1" if `srcType` is memref<512x16x1x1xf32> with
1547/// [8192, 16, 8, 1] strides and `vectorType` is vector<16x16x1x1xf32>
1548/// (only the inner most unit dim of `srcType` can be dropped)
1549///
1550/// Ex 3: return "0" if `srcType` is memref<512x16x1x1xf32> and
1551/// `vectorType` is vector<16x16x1x[1]xf32>
1552/// (the most inner dim in `vectorType` is not a unit dim (it's a "scalable
1553/// unit")
1554static FailureOr<size_t>
1555getTransferFoldableInnerUnitDims(MemRefType srcType, VectorType vectorType) {
1556 SmallVector<int64_t> srcStrides;
1557 int64_t srcOffset;
1558 if (failed(srcType.getStridesAndOffset(srcStrides, srcOffset)))
1559 return failure();
1560
1561 auto isUnitDim = [](VectorType type, int dim) {
1562 return type.getDimSize(dim) == 1 && !type.getScalableDims()[dim];
1563 };
1564
1565 // According to vector.transfer_read/write semantics, the vector can be a
1566 // slice. Thus, we have to offset the check index with `rankDiff` in
1567 // `srcStrides` and source dim sizes.
1568 size_t result = 0;
1569 int rankDiff = srcType.getRank() - vectorType.getRank();
1570 for (int64_t i = 0, e = vectorType.getRank(); i < e; ++i) {
1571 // Check that the inner dim size is 1 for both memref type and vector slice.
1572 // It can be folded only if they are 1 and the stride is 1.
1573 int dim = vectorType.getRank() - i - 1;
1574 if (srcStrides[dim + rankDiff] != 1 ||
1575 srcType.getDimSize(dim + rankDiff) != 1 || !isUnitDim(vectorType, dim))
1576 break;
1577 result++;
1578 }
1579 return result;
1580}
1581
1582/// Drop inner most contiguous unit dimensions from transfer_read operand.
1584 : public OpRewritePattern<vector::TransferReadOp> {
1585 using Base::Base;
1586
1587 LogicalResult matchAndRewrite(vector::TransferReadOp readOp,
1588 PatternRewriter &rewriter) const override {
1589 // TODO: support 0-d corner case.
1590 if (readOp.getTransferRank() == 0)
1591 return failure();
1592
1593 auto srcType = dyn_cast<MemRefType>(readOp.getBase().getType());
1594 if (!srcType)
1595 return failure();
1596
1597 if (!readOp.getPermutationMap().isMinorIdentity())
1598 return failure();
1599
1600 auto targetType = readOp.getVectorType();
1601 if (targetType.getRank() <= 1)
1602 return failure();
1603
1604 FailureOr<size_t> maybeDimsToDrop =
1605 getTransferFoldableInnerUnitDims(srcType, targetType);
1606 if (failed(maybeDimsToDrop))
1607 return failure();
1608
1609 size_t dimsToDrop = maybeDimsToDrop.value();
1610 if (dimsToDrop == 0)
1611 return failure();
1612
1613 auto inBounds = readOp.getInBoundsValues();
1614 auto droppedInBounds = ArrayRef<bool>(inBounds).take_back(dimsToDrop);
1615 if (llvm::is_contained(droppedInBounds, false))
1616 return failure();
1617
1618 auto resultTargetVecType =
1619 VectorType::get(targetType.getShape().drop_back(dimsToDrop),
1620 targetType.getElementType(),
1621 targetType.getScalableDims().drop_back(dimsToDrop));
1622
1623 auto loc = readOp.getLoc();
1625 memref::getMixedSizes(rewriter, loc, readOp.getBase());
1626 SmallVector<OpFoldResult> offsets(srcType.getRank(),
1627 rewriter.getIndexAttr(0));
1628 SmallVector<OpFoldResult> strides(srcType.getRank(),
1629 rewriter.getIndexAttr(1));
1630 MemRefType resultMemrefType = memref::SubViewOp::inferRankReducedResultType(
1631 srcType.getShape().drop_back(dimsToDrop), srcType, offsets, sizes,
1632 strides);
1633 ArrayAttr inBoundsAttr = rewriter.getArrayAttr(
1634 readOp.getInBoundsAttr().getValue().drop_back(dimsToDrop));
1635 Value rankedReducedView =
1636 memref::SubViewOp::create(rewriter, loc, resultMemrefType,
1637 readOp.getBase(), offsets, sizes, strides);
1638 auto permMap = getTransferMinorIdentityMap(
1639 cast<ShapedType>(rankedReducedView.getType()), resultTargetVecType);
1640
1641 // If there is a mask, shape_cast it to drop the same inner unit dims.
1642 Value mask = readOp.getMask();
1643 if (mask) {
1644 auto maskType = cast<VectorType>(mask.getType());
1645 auto reducedMaskType = VectorType::get(
1646 maskType.getShape().drop_back(dimsToDrop), maskType.getElementType(),
1647 maskType.getScalableDims().drop_back(dimsToDrop));
1648 mask = rewriter.createOrFold<vector::ShapeCastOp>(loc, reducedMaskType,
1649 mask);
1650 }
1651
1652 Value result = vector::TransferReadOp::create(
1653 rewriter, loc, resultTargetVecType, rankedReducedView,
1654 readOp.getIndices().drop_back(dimsToDrop), AffineMapAttr::get(permMap),
1655 readOp.getPadding(), mask, inBoundsAttr);
1656 rewriter.replaceOpWithNewOp<vector::ShapeCastOp>(readOp, targetType,
1657 result);
1658 return success();
1659 }
1660};
1661
1662/// Drop inner most contiguous unit dimensions from transfer_write operand.
1663/// E.g.,
1664/// vector.transfer_write %arg1, %arg0[%c0, %arg2, %c0, %c0, %c0]
1665/// {in_bounds = [true, true, true, true, true]}
1666/// : vector<1x16x16x1x1xf32>, memref<1x512x16x1x1xf32>
1667///
1668/// will be replaced with
1669///
1670/// %subview = memref.subview %arg0
1671/// [0, 0, 0, 0, 0] [1, 512, 16, 1, 1] [1, 1, 1, 1, 1]
1672/// : memref<1x512x16x1x1xf32> to memref<1x512x16xf32>
1673/// %0 = vector.shape_cast %arg1 : vector<1x16x16x1x1xf32>
1674/// to vector<1x16x16xf32>
1675/// vector.transfer_write %0, %subview[%c0, %arg2, %c0]
1676/// {in_bounds = [true, true, true]}
1677/// : vector<1x16x16xf32>, memref<1x512x16xf32>
1678///
1679/// Note, this pattern will not collapse "scalable unit" dims (i.e. `[1]`).
1681 : public OpRewritePattern<vector::TransferWriteOp> {
1682 using Base::Base;
1683
1684 LogicalResult matchAndRewrite(vector::TransferWriteOp writeOp,
1685 PatternRewriter &rewriter) const override {
1686 // TODO: support 0-d corner case.
1687 if (writeOp.getTransferRank() == 0)
1688 return failure();
1689
1690 auto srcType = dyn_cast<MemRefType>(writeOp.getBase().getType());
1691 if (!srcType)
1692 return failure();
1693
1694 if (!writeOp.getPermutationMap().isMinorIdentity())
1695 return failure();
1696
1697 auto targetType = writeOp.getVectorType();
1698 if (targetType.getRank() <= 1)
1699 return failure();
1700
1701 FailureOr<size_t> maybeDimsToDrop =
1702 getTransferFoldableInnerUnitDims(srcType, targetType);
1703 if (failed(maybeDimsToDrop))
1704 return failure();
1705
1706 size_t dimsToDrop = maybeDimsToDrop.value();
1707 if (dimsToDrop == 0)
1708 return failure();
1709
1710 auto inBounds = writeOp.getInBoundsValues();
1711 auto droppedInBounds = ArrayRef<bool>(inBounds).take_back(dimsToDrop);
1712 if (llvm::is_contained(droppedInBounds, false))
1713 return failure();
1714
1715 auto resultTargetVecType =
1716 VectorType::get(targetType.getShape().drop_back(dimsToDrop),
1717 targetType.getElementType(),
1718 targetType.getScalableDims().drop_back(dimsToDrop));
1719
1720 Location loc = writeOp.getLoc();
1722 memref::getMixedSizes(rewriter, loc, writeOp.getBase());
1723 SmallVector<OpFoldResult> offsets(srcType.getRank(),
1724 rewriter.getIndexAttr(0));
1725 SmallVector<OpFoldResult> strides(srcType.getRank(),
1726 rewriter.getIndexAttr(1));
1727 MemRefType resultMemrefType = memref::SubViewOp::inferRankReducedResultType(
1728 srcType.getShape().drop_back(dimsToDrop), srcType, offsets, sizes,
1729 strides);
1730 ArrayAttr inBoundsAttr = rewriter.getArrayAttr(
1731 writeOp.getInBoundsAttr().getValue().drop_back(dimsToDrop));
1732
1733 Value rankedReducedView =
1734 memref::SubViewOp::create(rewriter, loc, resultMemrefType,
1735 writeOp.getBase(), offsets, sizes, strides);
1736 auto permMap = getTransferMinorIdentityMap(
1737 cast<ShapedType>(rankedReducedView.getType()), resultTargetVecType);
1738
1739 auto shapeCast = rewriter.createOrFold<vector::ShapeCastOp>(
1740 loc, resultTargetVecType, writeOp.getVector());
1741
1742 // If there is a mask, shape_cast it to drop the same inner unit dims.
1743 Value mask = writeOp.getMask();
1744 if (mask) {
1745 auto maskType = cast<VectorType>(mask.getType());
1746 auto reducedMaskType = VectorType::get(
1747 maskType.getShape().drop_back(dimsToDrop), maskType.getElementType(),
1748 maskType.getScalableDims().drop_back(dimsToDrop));
1749 mask = rewriter.createOrFold<vector::ShapeCastOp>(loc, reducedMaskType,
1750 mask);
1751 }
1752
1753 rewriter.replaceOpWithNewOp<vector::TransferWriteOp>(
1754 writeOp, shapeCast, rankedReducedView,
1755 writeOp.getIndices().drop_back(dimsToDrop), AffineMapAttr::get(permMap),
1756 mask, inBoundsAttr);
1757 return success();
1758 }
1759};
1760
1761/// Canonicalization of a `vector.contract %a, %b, %c` with row-major matmul
1762/// semantics to a contraction suitable for MMT (matrix matrix multiplication
1763/// with the RHS transposed) lowering.
1765 : OpRewritePattern<vector::ContractionOp> {
1766 using Base::Base;
1767
1769 std::function<LogicalResult(vector::ContractionOp op)>;
1770
1772 FilterConstraintType constraint)
1773 : OpRewritePattern<vector::ContractionOp>(context, benefit),
1774 filter(std::move(constraint)) {}
1775
1776 LogicalResult matchAndRewrite(vector::ContractionOp op,
1777 PatternRewriter &rewriter) const override {
1778 if (failed(filter(op)))
1779 return failure();
1780
1781 Location loc = op.getLoc();
1782 Value lhs = op.getLhs();
1783 Value rhs = op.getRhs();
1784 Value res = op.getAcc();
1785
1786 // Set up the parallel/reduction structure in right form.
1787 using MapList = ArrayRef<ArrayRef<AffineExpr>>;
1788 auto infer = [&](MapList m) {
1789 return AffineMap::inferFromExprList(m, op.getContext());
1790 };
1791 AffineExpr m;
1792 AffineExpr n;
1793 AffineExpr k;
1794 bindDims(rewriter.getContext(), m, n, k);
1795 static constexpr std::array<int64_t, 2> perm = {1, 0};
1796 auto iteratorTypes = op.getIteratorTypes().getValue();
1797 SmallVector<AffineMap, 4> maps = op.getIndexingMapsArray();
1798 if (iteratorTypes.size() != 3 ||
1799 !vector::isParallelIterator(iteratorTypes[0]) ||
1800 !vector::isParallelIterator(iteratorTypes[1]) ||
1801 !vector::isReductionIterator(iteratorTypes[2]))
1802 return rewriter.notifyMatchFailure(op, "contraction is not a gemm");
1803
1804 // The canonical form is "TNT" = A row-major, B col-major, C row-major.
1805 const auto canonicalForm = infer({{m, k}, {n, k}, {m, n}});
1806 if (maps == canonicalForm)
1807 return rewriter.notifyMatchFailure(op, "already in the canonical form");
1808
1809 // Create a vector transpose making sure to emit zero/sign-extend at the
1810 // end.
1811 auto createTranspose = [&rewriter, loc](Value mat) -> Value {
1812 if (auto sext = mat.getDefiningOp<arith::ExtSIOp>()) {
1813 Value trans =
1814 vector::TransposeOp::create(rewriter, loc, sext.getIn(), perm);
1815 VectorType newType =
1816 cast<VectorType>(trans.getType())
1817 .clone(cast<VectorType>(mat.getType()).getElementType());
1818 return arith::ExtSIOp::create(rewriter, loc, newType, trans);
1819 }
1820 if (auto zext = mat.getDefiningOp<arith::ExtUIOp>()) {
1821 Value trans =
1822 vector::TransposeOp::create(rewriter, loc, zext.getIn(), perm);
1823 VectorType newType =
1824 VectorType::get(cast<VectorType>(trans.getType()).getShape(),
1825 cast<VectorType>(mat.getType()).getElementType());
1826 return arith::ExtUIOp::create(rewriter, loc, newType, trans);
1827 }
1828 return vector::TransposeOp::create(rewriter, loc, mat, perm);
1829 };
1830
1831 if (maps == infer({{m, k}, {k, n}, {m, n}})) {
1832 rhs = createTranspose(rhs);
1833 } else if (maps == infer({{k, m}, {n, k}, {m, n}})) {
1834 lhs = createTranspose(lhs);
1835 } else if (maps == infer({{k, m}, {k, n}, {m, n}})) {
1836 rhs = createTranspose(rhs);
1837 lhs = createTranspose(lhs);
1838 } else if (maps == infer({{k, m}, {k, n}, {n, m}})) {
1839 std::swap(rhs, lhs);
1840 rhs = createTranspose(rhs);
1841 lhs = createTranspose(lhs);
1842 } else if (maps == infer({{k, m}, {n, k}, {n, m}})) {
1843 std::swap(rhs, lhs);
1844 rhs = createTranspose(rhs);
1845 } else if (maps == infer({{m, k}, {k, n}, {n, m}})) {
1846 std::swap(lhs, rhs);
1847 lhs = createTranspose(lhs);
1848 } else if (maps == infer({{m, k}, {n, k}, {n, m}})) {
1849 std::swap(lhs, rhs);
1850 } else {
1851 return rewriter.notifyMatchFailure(op, "unhandled contraction form");
1852 }
1853 rewriter.replaceOpWithNewOp<vector::ContractionOp>(
1854 op, lhs, rhs, res, rewriter.getAffineMapArrayAttr(canonicalForm),
1855 op.getIteratorTypes());
1856 return success();
1857 };
1858
1859private:
1860 FilterConstraintType filter;
1861};
1862
1863/// Pattern to fold arithmetic extensions on floating point data types into
1864/// vector contraction operations. linalg.matmul introduces arithmetic
1865/// extensions on its operands. Please mlir snippets below for more details.
1866/// ```mlir
1867/// "linalg.matmul"(%lhs, %rhs, %acc) ({
1868/// ^bb0(%arg1: f16, %arg2: f16, %arg3: f32):
1869/// %lhs_f32 = "arith.extf"(%arg1) : (f16) -> f32
1870/// %rhs_f32 = "arith.extf"(%arg2) : (f16) -> f32
1871/// %mul = "arith.mulf"(%lhs_f32, %rhs_f32) : (f32, f32) -> f32
1872/// %acc = "arith.addf"(%arg3, %mul) : (f32, f32) -> f32
1873/// "linalg.yield"(%acc) : (f32) -> ()
1874/// })
1875/// ```
1876/// This restricts the native usage of mixed precision NVIDIA Ampere Tensor
1877/// Cores, i.e, `mma.sync.*.f32.f16.f16.f32` and `mma.sync.*.f32.bf16.bf16.f32`.
1878/// This pattern folds the arithmetic extensions into the vector contraction and
1879/// enables the usage of native mixed precision Tensor Core instructions.
1880template <typename ExtOp>
1882 : public OpRewritePattern<vector::ContractionOp> {
1883 using Base::Base;
1884
1885 LogicalResult matchAndRewrite(vector::ContractionOp contractOp,
1886 PatternRewriter &rewriter) const override {
1887
1888 auto lhsDefOp = contractOp.getLhs().getDefiningOp<ExtOp>();
1889 auto rhsDefOp = contractOp.getRhs().getDefiningOp<ExtOp>();
1890
1891 if (!lhsDefOp || !rhsDefOp) {
1892 return rewriter.notifyMatchFailure(contractOp,
1893 "no defining op on contract operands");
1894 }
1895
1896 rewriter.replaceOpWithNewOp<vector::ContractionOp>(
1897 contractOp, lhsDefOp->getOperand(0), rhsDefOp->getOperand(0),
1898 contractOp.getAcc(), contractOp.getIndexingMapsAttr(),
1899 contractOp.getIteratorTypesAttr());
1900
1901 return success();
1902 }
1903};
1904
1905/// Pattern to fold chained reduction to a series of vector additions and a
1906/// final reduction. This form should require fewer subgroup operations.
1907///
1908/// ```mlir
1909/// %a = vector.reduction <add> %x, %acc
1910/// %b = vector.reduction <add> %y, %a
1911/// ==>
1912/// %a = arith.addf %x, %y
1913/// %b = vector.reduction <add> %a, %acc
1914/// ```
1915struct ChainedReduction final : OpRewritePattern<vector::ReductionOp> {
1916 using Base::Base;
1917
1918 LogicalResult matchAndRewrite(vector::ReductionOp op,
1919 PatternRewriter &rewriter) const override {
1920 // TODO: Handle other combining kinds.
1921 if (op.getKind() != vector::CombiningKind::ADD)
1922 return failure();
1923
1924 // Accumulator is optional.
1925 Value acc = op.getAcc();
1926 if (!acc)
1927 return failure();
1928
1929 if (!acc.getType().isIntOrFloat())
1930 return failure();
1931
1932 auto parentReduction = acc.getDefiningOp<vector::ReductionOp>();
1933 if (!parentReduction)
1934 return failure();
1935
1936 Location loc = op.getLoc();
1937 Value vAdd;
1938 if (isa<IntegerType>(acc.getType())) {
1939 vAdd = rewriter.createOrFold<arith::AddIOp>(
1940 loc, parentReduction.getVector(), op.getVector());
1941 } else {
1942 vAdd = arith::AddFOp::create(rewriter, loc, parentReduction.getVector(),
1943 op.getVector());
1944 }
1945 rewriter.replaceOpWithNewOp<vector::ReductionOp>(op, op.getKind(), vAdd,
1946 parentReduction.getAcc());
1947 return success();
1948 }
1949};
1950
1951// Helper function dropping unit non-scalable dimension from a VectorType
1952// keeping at least 1 dimension to avoid generating 0-D vectors. Scalable unit
1953// dimensions are not dropped. Folding such dimensions would require "shifting"
1954// the scalable flag onto some other fixed-width dim (e.g. vector<[1]x4xf32> ->
1955// vector<[4]xf32>). This could be implemented in the future.
1956static VectorType dropNonScalableUnitDimFromType(VectorType inVecTy) {
1957 auto inVecShape = inVecTy.getShape();
1958 SmallVector<int64_t> newShape;
1959 SmallVector<bool> newScalableDims;
1960 for (auto [dim, isScalable] :
1961 llvm::zip_equal(inVecShape, inVecTy.getScalableDims())) {
1962 if (dim == 1 && !isScalable)
1963 continue;
1964
1965 newShape.push_back(dim);
1966 newScalableDims.push_back(isScalable);
1967 }
1968 // All dims have been dropped, return vector<1xeType>.
1969 if (newShape.empty()) {
1970 newShape.push_back(1);
1971 newScalableDims.push_back(false);
1972 }
1973
1974 return VectorType::get(newShape, inVecTy.getElementType(), newScalableDims);
1975}
1976
1977/// For vectors with at least one unit dim, replaces:
1978/// elementwise(a, b)
1979/// with:
1980/// sc_a = shape_cast(a)
1981/// sc_b = shape_cast(b)
1982/// res = elementwise(sc_a, sc_b)
1983/// return shape_cast(res)
1984/// The newly inserted shape_cast Ops fold (before elementwise Op) and then
1985/// restore (after elementwise Op) the unit dim. Vectors `a` and `b` are
1986/// required to be rank > 1.
1987///
1988/// Ex:
1989/// %mul = arith.mulf %B_row, %A_row : vector<1x[4]xf32>
1990/// %cast = vector.shape_cast %mul : vector<1x[4]xf32> to vector<[4]xf32>
1991///
1992/// gets converted to:
1993///
1994/// %B_row_sc = vector.shape_cast %B_row : vector<1x[4]xf32> to vector<[4]xf32>
1995/// %A_row_sc = vector.shape_cast %A_row : vector<1x[4]xf32> to vector<[4]xf32>
1996/// %mul = arith.mulf %B_row_sc, %A_row_sc : vector<[4]xf32>
1997/// %cast_new = vector.shape_cast %mul : vector<[4]xf32> to vector<1x[4]xf32>
1998/// %cast = vector.shape_cast %cast_new : vector<1x[4]xf32> to vector<[4]xf32>
1999///
2000/// Patterns for folding shape_casts should instantly eliminate `%cast_new` and
2001/// `%cast`.
2003 : public OpTraitRewritePattern<OpTrait::Elementwise> {
2005 LogicalResult matchAndRewrite(Operation *op,
2006 PatternRewriter &rewriter) const override {
2007 if (op->getNumResults() != 1 || op->getNumRegions() != 0)
2008 return failure();
2009
2010 auto resultVectorType = dyn_cast<VectorType>(op->getResult(0).getType());
2011 if (!resultVectorType)
2012 return failure();
2013
2014 // Check the operand pre-conditions. For `Elementwise` ops all operands are
2015 // guaranteed to have identical shapes (with some exceptions such as
2016 // `arith.select`) and it suffices to only check one of them.
2017 auto sourceVectorType = dyn_cast<VectorType>(op->getOperand(0).getType());
2018 if (!sourceVectorType)
2019 return failure();
2020 if (sourceVectorType.getRank() < 2)
2021 return failure();
2022
2023 SmallVector<Value> newOperands;
2024 auto loc = op->getLoc();
2025 for (auto operand : op->getOperands()) {
2026 auto opVectorType = cast<VectorType>(operand.getType());
2027 auto newVType = dropNonScalableUnitDimFromType(opVectorType);
2028 if (newVType == opVectorType)
2029 return rewriter.notifyMatchFailure(op, "No unit dimension to remove.");
2030
2031 auto opSC = vector::ShapeCastOp::create(rewriter, loc, newVType, operand);
2032 newOperands.push_back(opSC);
2033 }
2034
2035 VectorType newResultVectorType =
2036 dropNonScalableUnitDimFromType(resultVectorType);
2037 // Create an updated elementwise Op without unit dim.
2038 Operation *elementwiseOp =
2039 rewriter.create(loc, op->getName().getIdentifier(), newOperands,
2040 newResultVectorType, op->getAttrs());
2041
2042 // Restore the unit dim by applying vector.shape_cast to the result.
2043 rewriter.replaceOpWithNewOp<ShapeCastOp>(op, resultVectorType,
2044 elementwiseOp->getResult(0));
2045
2046 return success();
2047 }
2048};
2049
2050/// A pattern to drop unit dims from vector.transpose.
2051///
2052/// Example:
2053///
2054/// BEFORE:
2055/// ```mlir
2056/// %transpose = vector.transpose %vector, [3, 0, 1, 2]
2057/// : vector<1x1x4x[4]xf32> to vector<[4]x1x1x4xf32>
2058/// ```
2059///
2060/// AFTER:
2061/// ```mlir
2062/// %dropDims = vector.shape_cast %vector
2063/// : vector<1x1x4x[4]xf32> to vector<4x[4]xf32>
2064/// %transpose = vector.transpose %0, [1, 0]
2065/// : vector<4x[4]xf32> to vector<[4]x4xf32>
2066/// %restoreDims = vector.shape_cast %transpose
2067/// : vector<[4]x4xf32> to vector<[4]x1x1x4xf32>
2068/// ```
2070 : OpRewritePattern<vector::TransposeOp> {
2071 using Base::Base;
2072
2073 LogicalResult matchAndRewrite(vector::TransposeOp op,
2074 PatternRewriter &rewriter) const override {
2075 VectorType sourceType = op.getSourceVectorType();
2076 VectorType sourceTypeWithoutUnitDims =
2078
2079 if (sourceType == sourceTypeWithoutUnitDims)
2080 return failure();
2081
2082 // Construct a map from dimIdx -> number of dims dropped before dimIdx.
2083 auto sourceDims = llvm::to_vector(vector::getDims(sourceType));
2084 SmallVector<int64_t> droppedDimsBefore(sourceType.getRank());
2085 int64_t droppedDims = 0;
2086 for (auto [i, dim] : llvm::enumerate(sourceDims)) {
2087 droppedDimsBefore[i] = droppedDims;
2088 if (dim == std::make_tuple(1, false))
2089 ++droppedDims;
2090 }
2091
2092 // Drop unit dims from transpose permutation.
2093 ArrayRef<int64_t> perm = op.getPermutation();
2094 SmallVector<int64_t> newPerm;
2095 for (int64_t idx : perm) {
2096 if (sourceDims[idx] == std::make_tuple(1, false))
2097 continue;
2098 newPerm.push_back(idx - droppedDimsBefore[idx]);
2099 }
2100
2101 // Fixup for `newPerm`. The `sourceTypeWithoutUnitDims` could be vector<1xT>
2102 // type when the dimensions are unit dimensions. In this case, the newPerm
2103 // should be [0].
2104 if (newPerm.empty()) {
2105 newPerm.push_back(0);
2106 }
2107
2108 Location loc = op.getLoc();
2109 // Drop the unit dims via shape_cast.
2110 auto dropDimsShapeCast = vector::ShapeCastOp::create(
2111 rewriter, loc, sourceTypeWithoutUnitDims, op.getVector());
2112 // Create the new transpose.
2113 auto transposeWithoutUnitDims =
2114 vector::TransposeOp::create(rewriter, loc, dropDimsShapeCast, newPerm);
2115 // Restore the unit dims via shape cast.
2116 rewriter.replaceOpWithNewOp<vector::ShapeCastOp>(
2117 op, op.getResultVectorType(), transposeWithoutUnitDims);
2118
2119 return success();
2120 }
2121};
2122
2123/// A pattern to drop unit dims from the iter_args of an scf.for.
2124///
2125/// Example:
2126///
2127/// BEFORE:
2128/// ```mlir
2129/// %res = scf.for ... iter_args(%iter = %init) -> vector<[4]x1x1x4xf32> {
2130/// ...
2131/// scf.yield %
2132/// }
2133/// ```
2134///
2135/// AFTER:
2136/// ```mlir
2137/// %drop = vector.shape_cast %init
2138/// : vector<4x1x1x[4]xf32> to vector<4x[4]xf32>
2139/// %new_loop = scf.for ... iter_args(%iter = %drop) -> vector<[4]x4xf32> {
2140/// %new_iter = vector.shape_cast %iter
2141/// : vector<[4]x4xf32> to vector<[4]x1x1x4xf32>
2142/// ...
2143/// }
2144/// %res = vector.shape_cast %new_loop
2145/// : vector<[4]x4xf32> to vector<[4]x1x1x4xf32>
2146/// ```
2147struct DropUnitDimsFromScfForOp final : OpRewritePattern<scf::ForOp> {
2148 using Base::Base;
2149
2150 LogicalResult matchAndRewrite(scf::ForOp forOp,
2151 PatternRewriter &rewriter) const override {
2152 /// Find the first iter_arg with droppable unit dims. Further applications
2153 /// of this pattern will apply to later arguments.
2154 for (OpOperand &operand : forOp.getInitArgsMutable()) {
2155 auto vectorType = dyn_cast<VectorType>(operand.get().getType());
2156 if (!vectorType)
2157 continue;
2158
2159 VectorType newVectorType = dropNonScalableUnitDimFromType(vectorType);
2160 if (vectorType == newVectorType)
2161 continue;
2162
2163 // Create a new ForOp with that iter operand replaced.
2164 auto castFn = [](OpBuilder &b, Location loc, Type type, Value source) {
2165 return vector::ShapeCastOp::create(b, loc, type, source);
2166 };
2167
2169 castFn(rewriter, forOp.getLoc(), newVectorType, operand.get());
2170 rewriter.replaceOp(forOp,
2171 replaceAndCastForOpIterArg(rewriter, forOp, operand,
2172 replacement, castFn));
2173 return success();
2174 }
2175 return failure();
2176 }
2177};
2178
2179/// Pattern to eliminate redundant zero-constants added to reduction operands.
2180/// It's enough for there to be one initial zero value, so we can eliminate the
2181/// extra ones that feed into `vector.reduction <add>`. These get created by the
2182/// `ChainedReduction` pattern.
2183///
2184/// ```mlir
2185/// %a = arith.addf %x, %zero
2186/// %b = arith.addf %a, %y
2187/// %c = vector.reduction <add> %b, %acc
2188/// ==>
2189/// %b = arith.addf %a, %y
2190/// %c = vector.reduction <add> %b, %acc
2191/// ```
2192struct ReduceRedundantZero final : OpRewritePattern<vector::ReductionOp> {
2193 using Base::Base;
2194
2195 LogicalResult matchAndRewrite(vector::ReductionOp op,
2196 PatternRewriter &rewriter) const override {
2197 // TODO: Handle other reduction kinds and their identity values.
2198 if (op.getKind() != vector::CombiningKind::ADD)
2199 return failure();
2200
2201 Type elemType = op.getSourceVectorType().getElementType();
2202 // The integer case should be handled by `arith.addi` folders, only check
2203 // for floats here.
2204 if (!isa<FloatType>(elemType))
2205 return failure();
2206
2207 auto vAdd = op.getVector().getDefiningOp<arith::AddFOp>();
2208 if (!vAdd)
2209 return failure();
2210 auto addLhs = vAdd.getLhs().getDefiningOp<arith::AddFOp>();
2211 if (!addLhs)
2212 return failure();
2213
2214 if (!matchPattern(addLhs.getRhs(), m_AnyZeroFloat()))
2215 return failure();
2216
2217 auto newAdd = arith::AddFOp::create(rewriter, vAdd.getLoc(),
2218 addLhs.getLhs(), vAdd.getRhs());
2219 rewriter.replaceOpWithNewOp<vector::ReductionOp>(op, op.getKind(), newAdd,
2220 op.getAcc());
2221 return success();
2222 }
2223};
2224
2225/// Example:
2226/// ```
2227/// %a = vector.reduction <add> %x : vector<2xf32> into f32
2228/// ```
2229/// is transformed into:
2230/// ```
2231/// %y = vector.extract %x[0] : f32 from vector<2xf32>
2232/// %z = vector.extract %x[1] : f32 from vector<2xf32>
2233/// %a = arith.addf %y, %z : f32
2234/// ```
2235struct BreakDownVectorReduction final : OpRewritePattern<vector::ReductionOp> {
2237 unsigned maxNumElementsToExtract,
2238 PatternBenefit benefit)
2239 : OpRewritePattern(context, benefit),
2240 maxNumElementsToExtract(maxNumElementsToExtract) {}
2241
2242 LogicalResult matchAndRewrite(vector::ReductionOp op,
2243 PatternRewriter &rewriter) const override {
2244 VectorType type = op.getSourceVectorType();
2245 if (type.isScalable() || op.isMasked())
2246 return failure();
2247 assert(type.getRank() == 1 && "Expected a 1-d vector");
2248
2249 int64_t numElems = type.getNumElements();
2250 if (numElems > maxNumElementsToExtract) {
2251 return rewriter.notifyMatchFailure(
2252 op, llvm::formatv("has too many vector elements ({0}) to break down "
2253 "(max allowed: {1})",
2254 numElems, maxNumElementsToExtract));
2255 }
2256
2257 Location loc = op.getLoc();
2258 SmallVector<Value> extracted(numElems, nullptr);
2259 for (auto [idx, extractedElem] : llvm::enumerate(extracted))
2260 extractedElem = vector::ExtractOp::create(rewriter, loc, op.getVector(),
2261 static_cast<int64_t>(idx));
2262
2263 Value res = extracted.front();
2264 for (auto extractedElem : llvm::drop_begin(extracted))
2265 res = vector::makeArithReduction(rewriter, loc, op.getKind(), res,
2266 extractedElem, op.getFastmathAttr());
2267 if (Value acc = op.getAcc())
2268 res = vector::makeArithReduction(rewriter, loc, op.getKind(), res, acc,
2269 op.getFastmathAttr());
2270
2271 rewriter.replaceOp(op, res);
2272 return success();
2273 }
2274
2275private:
2276 unsigned maxNumElementsToExtract = 0;
2277};
2278
2279/// Fold `mulf(tr(broadcast(A)), broadcast(B))` into `vector.outerproduct(A,
2280/// B)`.
2281/// Example:
2282/// %lhsBcast = vector.broadcast %lhs : vector<4xi32> to vector<4x4xi32>
2283/// %lhsT = vector.transpose %lhsBcast, [1, 0] : vector<4x4xi32> to
2284/// vector<4x4xi32> %rhsBcast = vector.broadcast %rhs : vector<4xi32> to
2285/// vector<4x4xi32> %mul = arith.muli %lhsT, %rhsBcast : vector<4x4xi32>
2286///
2287/// Becomes :
2288///
2289/// %res = vector.outerproduct %lhs, %rhs : vector<4xi32>, vector<4xi32>
2290///
2291/// Supports only 1D-to-2D broadcasts. The following cases are not supported.
2292/// %ex1 = vector.broadcast %lhsCast : vector<1x4xf32> to vector<4x4xf32>
2293/// %ex2 = vector.broadcast %lhsCast : f32 to vector<4x4xf32>
2294/// %ex3 = vector.broadcast %lhsCast : vector<1x1xf32> to vector<4x4xf32>
2295template <typename MulOpType>
2296struct FoldArithToVectorOuterProduct : public OpRewritePattern<MulOpType> {
2297 using OpRewritePattern<MulOpType>::OpRewritePattern;
2298 // Returns whether a vector.broadcast matches requirements for an outerproduct
2299 // pattern. aka a 1D-to-2D broadcastOp without broadcasted unit dimension.
2300 bool isValidBroadcastSource(vector::BroadcastOp broadcastOp) const {
2301 // Fail if it is not a 1-to-2 dimension to broadcast to avoid generating
2302 // shape_casts/broadcasts which does not belong in this pattern.
2303 if (!broadcastOp.computeBroadcastedUnitDims().empty())
2304 return false;
2305 // Avoid broadcast like f32 or vector<f32> -> ResType
2306 auto srcType = dyn_cast<VectorType>(broadcastOp.getSourceType());
2307 return srcType && srcType.getRank() != 2;
2308 }
2309
2310 LogicalResult matchAndRewrite(MulOpType mulOp,
2311 PatternRewriter &rewriter) const override {
2312 auto resType = llvm::dyn_cast<VectorType>(mulOp.getResult().getType());
2313 if (!resType)
2314 return failure();
2315 if (resType.getRank() != 2)
2316 return failure();
2317 /// If operandA can be written as tr(broadcast(A)) and operandB as
2318 /// broadcast(B) where broadcasts are 1D-to-2D, create and return
2319 /// vector.outerproduct(A, B). Returns failure() otherwise.
2320 auto matchOuterProduct =
2321 [&](Value operandA,
2322 Value operandB) -> FailureOr<vector::OuterProductOp> {
2323 auto transposedLhs = operandA.getDefiningOp<vector::TransposeOp>();
2324 if (!transposedLhs)
2325 return failure();
2326 // Fail unless this is a true 2-D matrix transpose.
2327 ArrayRef<int64_t> permutation = transposedLhs.getPermutation();
2328 if (permutation.size() != 2 || permutation[0] != 1 || permutation[1] != 0)
2329 return failure();
2330
2331 auto broadcastedLhs =
2332 transposedLhs.getVector().getDefiningOp<vector::BroadcastOp>();
2333 if (!broadcastedLhs || !isValidBroadcastSource(broadcastedLhs))
2334 return failure();
2335
2336 auto broadcastedRhs = operandB.getDefiningOp<vector::BroadcastOp>();
2337 if (!broadcastedRhs || !isValidBroadcastSource(broadcastedRhs))
2338 return failure();
2339
2340 return vector::OuterProductOp::create(
2341 rewriter, mulOp->getLoc(), resType, broadcastedLhs.getSource(),
2342 broadcastedRhs.getSource(), Value(), vector::CombiningKind::ADD);
2343 };
2344
2345 Value lhs = mulOp->getOperand(0), rhs = mulOp->getOperand(1);
2346 auto maybeOuterP = matchOuterProduct(lhs, rhs);
2347 // Handle commutativity, the transposed op is the outerproduct LHS.
2348 if (failed(maybeOuterP))
2349 maybeOuterP = matchOuterProduct(rhs, lhs);
2350 if (failed(maybeOuterP))
2351 return failure();
2352 rewriter.replaceOp(mulOp, maybeOuterP->getResult());
2353 return success();
2354 }
2355};
2356
2357} // namespace
2358
2365
2366void mlir::vector::populateVectorMaskMaterializationPatterns(
2367 RewritePatternSet &patterns, bool force32BitVectorIndices,
2368 PatternBenefit benefit) {
2369 patterns.add<VectorCreateMaskOpConversion,
2370 MaterializeTransferMask<vector::TransferReadOp>,
2371 MaterializeTransferMask<vector::TransferWriteOp>>(
2372 patterns.getContext(), force32BitVectorIndices, benefit);
2373 patterns.add<FoldI1Select>(patterns.getContext(), benefit);
2374}
2375
2376void mlir::vector::populateDropUnitDimWithShapeCastPatterns(
2377 RewritePatternSet &patterns, PatternBenefit benefit) {
2379 DropUnitDimsFromTransposeOp>(patterns.getContext(), benefit);
2380}
2381
2382void mlir::vector::populateBubbleVectorBitCastOpPatterns(
2383 RewritePatternSet &patterns, PatternBenefit benefit) {
2384 patterns.add<BubbleDownVectorBitCastForExtract,
2385 BubbleDownBitCastForStridedSliceExtract,
2386 BubbleUpBitCastForInsert, BubbleUpBitCastForStridedSliceInsert>(
2387 patterns.getContext(), benefit);
2388}
2389
2390void mlir::vector::populateBreakDownVectorBitCastOpPatterns(
2391 RewritePatternSet &patterns,
2392 std::function<bool(vector::BitCastOp)> controlFn, PatternBenefit benefit) {
2393 patterns.add<BreakDownVectorBitCast>(patterns.getContext(),
2394 std::move(controlFn), benefit);
2395}
2396
2398 RewritePatternSet &patterns,
2399 std::function<LogicalResult(vector::ContractionOp)> constraint,
2400 PatternBenefit benefit) {
2401 patterns.add<CanonicalizeContractMatmulToMMT>(patterns.getContext(), benefit,
2402 std::move(constraint));
2403}
2404
2406 RewritePatternSet &patterns, PatternBenefit benefit) {
2407 patterns.add<MultiReduceToContract, CombineContractBroadcastMask,
2408 CombineContractABTranspose, CombineContractResultTranspose>(
2409 patterns.getContext(), benefit);
2410}
2411
2418
2420 PatternBenefit benefit) {
2421 patterns.add<ReorderElementwiseOpsOnTranspose, ReorderCastOpsOnBroadcast,
2422 ReorderElementwiseOpsOnBroadcast, ExtractOpFromElementwise>(
2423 patterns.getContext(), benefit);
2424}
2425
2426void mlir::vector::populateSinkVectorMemOpsPatterns(RewritePatternSet &patterns,
2427 PatternBenefit benefit) {
2428 // TODO: Consider converting these patterns to canonicalizations.
2429 patterns.add<ExtractOpFromLoad, StoreOpFromBroadcast>(patterns.getContext(),
2430 benefit);
2431}
2432
2433void mlir::vector::populateChainedVectorReductionFoldingPatterns(
2434 RewritePatternSet &patterns, PatternBenefit benefit) {
2435 patterns.add<ChainedReduction>(patterns.getContext(), benefit);
2436 patterns.add<ReduceRedundantZero>(patterns.getContext(),
2437 PatternBenefit(benefit.getBenefit() + 1));
2438}
2439
2440void mlir::vector::populateBreakDownVectorReductionPatterns(
2441 RewritePatternSet &patterns, unsigned maxNumElementsToExtract,
2442 PatternBenefit benefit) {
2443 patterns.add<BreakDownVectorReduction>(patterns.getContext(),
2444 maxNumElementsToExtract, benefit);
2445}
2446
2448 RewritePatternSet &patterns) {
2449 patterns.add<FoldArithToVectorOuterProduct<arith::MulFOp>,
2450 FoldArithToVectorOuterProduct<arith::MulIOp>>(
2451 patterns.getContext());
2452}
2453
2454//===----------------------------------------------------------------------===//
2455// TableGen'd enum attribute definitions
2456//===----------------------------------------------------------------------===//
2457
2458#include "mlir/Dialect/Vector/Transforms/VectorTransformsEnums.cpp.inc"
return success()
static uint64_t zext(uint32_t arg)
lhs
b
Return true if permutation is a valid permutation of the outer_dims_perm (case OuterOrInnerPerm::Oute...
ArrayAttr()
*if copies could not be generated due to yet unimplemented cases *copyInPlacementStart and copyOutPlacementStart in copyPlacementBlock *specify the insertion points where the incoming copies and outgoing should be the output argument nBegin is set to its * replacement(set to `begin` if no invalidation happens). Since outgoing *copies could have been inserted at `end`
static Value broadcast(Location loc, Value toBroadcast, unsigned numElements, const TypeConverter &typeConverter, ConversionPatternRewriter &rewriter)
Broadcasts the value to vector with numElements number of elements.
static std::optional< int64_t > getResultIndex(AffineMap map, int64_t index)
static VectorType dropNonScalableUnitDimFromType(VectorType inVecTy)
static SmallVector< IntType > extractVector(ArrayAttr arrayAttr)
static FailureOr< size_t > getTransferFoldableInnerUnitDims(MemRefType srcType, VectorType vectorType)
Returns the number of dims can be folded away from transfer ops. It returns a failure if it can not d...
Drop inner most contiguous unit dimensions from transfer_read operand.
Drop inner most contiguous unit dimensions from transfer_write operand. E.g., vector....
Base type for affine expression.
Definition AffineExpr.h:68
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 ...
static AffineMap get(MLIRContext *context)
Returns a zero result affine map with no dimensions or symbols: () -> ().
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...
static AffineMap getPermutationMap(ArrayRef< unsigned > permutation, MLIRContext *context)
Returns an AffineMap representing a permutation.
AffineMap compose(AffineMap map) const
Returns the AffineMap resulting from composing this with map.
IntegerAttr getIndexAttr(int64_t value)
Definition Builders.cpp:116
AffineMap getMultiDimIdentityMap(unsigned rank)
Definition Builders.cpp:396
IntegerType getI64Type()
Definition Builders.cpp:73
IntegerType getI32Type()
Definition Builders.cpp:71
IntegerType getIntegerType(unsigned width)
Definition Builders.cpp:75
TypedAttr getZeroAttr(Type type)
Definition Builders.cpp:333
AffineExpr getAffineDimExpr(unsigned position)
Definition Builders.cpp:373
DenseIntElementsAttr getI32VectorAttr(ArrayRef< int32_t > values)
Definition Builders.cpp:130
DenseIntElementsAttr getI64VectorAttr(ArrayRef< int64_t > values)
Definition Builders.cpp:136
IntegerType getI1Type()
Definition Builders.cpp:61
ArrayAttr getArrayAttr(ArrayRef< Attribute > value)
Definition Builders.cpp:275
MLIRContext * getContext() const
Definition Builders.h:56
ArrayAttr getI64ArrayAttr(ArrayRef< int64_t > values)
Definition Builders.cpp:290
ArrayAttr getBoolArrayAttr(ArrayRef< bool > values)
Definition Builders.cpp:279
ArrayAttr getAffineMapArrayAttr(ArrayRef< AffineMap > values)
Definition Builders.cpp:327
DenseElementsAttr resizeSplat(ShapedType newType)
Return a new DenseElementsAttr that has the same data as the current attribute, but with a different ...
std::enable_if_t<!std::is_base_of< Attribute, T >::value||std::is_same< Attribute, T >::value, T > getSplatValue() const
Return the splat value for this attribute.
An attribute that represents a reference to a dense integer vector or tensor object.
static DenseIntElementsAttr get(const ShapedType &type, Arg &&arg)
Get an instance of a DenseIntElementsAttr with the given arguments.
void map(Value from, Value to)
Inserts a new mapping for 'from' to 'to'.
Definition IRMapping.h:30
This class defines the main interface for locations in MLIR and acts as a non-nullable wrapper around...
Definition Location.h:76
MLIRContext is the top-level object for a collection of MLIR operations.
Definition MLIRContext.h:63
This class helps build Operations.
Definition Builders.h:210
Operation * clone(Operation &op, IRMapping &mapper)
Creates a deep copy of the specified operation, remapping any operands that use values outside of the...
Definition Builders.cpp:571
void setInsertionPoint(Block *block, Block::iterator insertPoint)
Set the insertion point to the specified location.
Definition Builders.h:401
void createOrFold(SmallVectorImpl< Value > &results, Location location, Args &&...args)
Create an operation of specific op type at the current insertion point, and immediately try to fold i...
Definition Builders.h:529
Operation * create(const OperationState &state)
Creates an operation given the fields represented as an OperationState.
Definition Builders.cpp:466
This class represents an operand of an operation.
Definition Value.h:254
OpTraitRewritePattern is a wrapper around RewritePattern that allows for matching and rewriting again...
OpTraitRewritePattern(MLIRContext *context, PatternBenefit benefit=1)
StringAttr getIdentifier() const
Return the name of this operation as a StringAttr.
Operation is the basic unit of execution within MLIR.
Definition Operation.h:87
Value getOperand(unsigned idx)
Definition Operation.h:375
bool hasTrait()
Returns true if the operation was registered with a particular trait, e.g.
Definition Operation.h:774
ArrayRef< NamedAttribute > getAttrs()
Return all of the attributes on this operation.
Definition Operation.h:537
bool hasOneUse()
Returns true if this operation has exactly one use.
Definition Operation.h:874
OpResult getResult(unsigned idx)
Get the 'idx'th result of this operation.
Definition Operation.h:432
unsigned getNumRegions()
Returns the number of regions held by this operation.
Definition Operation.h:699
Location getLoc()
The source location the operation was defined or derived from.
Definition Operation.h:240
unsigned getNumOperands()
Definition Operation.h:371
OperationName getName()
The name of an operation is the key identifier for it.
Definition Operation.h:115
operand_type_range getOperandTypes()
Definition Operation.h:422
result_type_range getResultTypes()
Definition Operation.h:453
operand_range getOperands()
Returns an iterator on the underlying Value's.
Definition Operation.h:403
result_range getResults()
Definition Operation.h:440
unsigned getNumResults()
Return the number of results held by this operation.
Definition Operation.h:429
This class represents the benefit of a pattern match in a unitless scheme that ranges from 0 (very li...
unsigned short getBenefit() const
If the corresponding pattern can match, return its benefit. If the.
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.
virtual void replaceOp(Operation *op, ValueRange newValues)
Replace the results of the given (original) operation with the specified list of values (replacements...
virtual void eraseOp(Operation *op)
This method erases an operation that is known to have no uses.
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,...
void modifyOpInPlace(Operation *root, CallableT &&callable)
This method is a utility wrapper around an in-place modification of an operation.
OpTy replaceOpWithNewOp(Operation *op, Args &&...args)
Replace the results of the given (original) op with a new op that is created without verification (re...
Instances of the Type class are uniqued, have an immutable identifier and an optional mutable compone...
Definition Types.h:74
bool isIntOrFloat() const
Return true if this is an integer (of any signedness) or a float type.
Definition Types.cpp:118
unsigned getIntOrFloatBitWidth() const
Return the bit width of an integer or a float type, assert failure on other types.
Definition Types.cpp:124
bool isSignlessIntOrIndexOrFloat() const
Return true if this is a signless integer, index, or float type.
Definition Types.cpp:106
This class represents an instance of an SSA value in the MLIR system, representing a computable value...
Definition Value.h:96
void setType(Type newType)
Mutate the type of this Value to be of the specified type.
Definition Value.h:116
Type getType() const
Return the type of this value.
Definition Value.h:105
bool hasOneUse() const
Returns true if this value has exactly one use.
Definition Value.h:197
Operation * getDefiningOp() const
If this value is the result of an operation, return the operation that defines it.
Definition Value.cpp:18
static ConstantIndexOp create(OpBuilder &builder, Location location, int64_t value)
Definition ArithOps.cpp:397
bool hasElementwiseMappableTraits(Operation *op)
Together, Elementwise, Scalarizable, Vectorizable, and Tensorizable provide an easy way for scalar op...
SmallVector< OpFoldResult > getMixedSizes(OpBuilder &builder, Location loc, Value value)
Return the dimensions of the given memref value.
Definition MemRefOps.cpp:79
Value makeArithReduction(OpBuilder &b, Location loc, CombiningKind kind, Value v1, Value acc, arith::FastMathFlagsAttr fastmath=nullptr, Value mask=nullptr)
Returns the result value of reducing two scalar/vector values with the corresponding arith operation.
Operation * maskOperation(OpBuilder &builder, Operation *maskableOp, Value mask, Value passthru=Value())
Creates a vector.mask operation around a maskable operation.
bool isReductionIterator(Attribute attr)
Returns true if attr has "reduction" iterator type semantics.
Definition VectorOps.h:156
auto getDims(VectorType vType)
Returns a range over the dims (size and scalability) of a VectorType.
void populateElementwiseToVectorOpsPatterns(RewritePatternSet &patterns)
Collect a set of patterns that fold elementwise op on vectors to the vector dialect.
AffineMap getTransferMinorIdentityMap(ShapedType shapedType, VectorType vectorType)
Build the default minor identity map suitable for a vector transfer.
void populateDropInnerMostUnitDimsXferOpPatterns(RewritePatternSet &patterns, PatternBenefit benefit=1)
Collect a set of patterns to collapse the most inner unit dims in xfer Ops.
bool isParallelIterator(Attribute attr)
Returns true if attr has "parallel" iterator type semantics.
Definition VectorOps.h:151
void populateFoldArithExtensionPatterns(RewritePatternSet &patterns)
Collect a set of patterns that fold arithmetic extension on floating point into vector contract for t...
void populateVectorContractCanonicalizeMatmulToMMT(RewritePatternSet &patterns, std::function< LogicalResult(vector::ContractionOp)> constraint=[](vector::ContractionOp) { return success();}, PatternBenefit=1)
Canonicalization of a vector.contract a, b, c with row-major matmul semantics to a contraction with M...
void populateSinkVectorOpsPatterns(RewritePatternSet &patterns, PatternBenefit benefit=1)
Patterns that remove redundant Vector Ops by re-ordering them with e.g.
void populateVectorReductionToContractPatterns(RewritePatternSet &patterns, PatternBenefit benefit=1)
Collect patterns to convert reduction op to vector.contract and fold transpose/broadcast ops into the...
Value createOrFoldDimOp(OpBuilder &b, Location loc, Value source, int64_t dim)
Helper function that creates a memref::DimOp or tensor::DimOp depending on the type of source.
Include the generated interface declarations.
bool matchPattern(Value value, const Pattern &pattern)
Entry point for matching a pattern over a Value.
Definition Matchers.h:490
Type getType(OpFoldResult ofr)
Returns the int type of the integer in ofr.
Definition Utils.cpp:307
void bindDims(MLIRContext *ctx, AffineExprTy &...exprs)
Bind a list of AffineExpr references to DimExpr at positions: [0 .
Definition AffineExpr.h:311
AffineMap inversePermutation(AffineMap map)
Returns a map of codomain to domain dimensions such that the first codomain dimension for a particula...
Value getValueOrCreateCastToIndexLike(OpBuilder &b, Location loc, Type targetType, Value value)
Create a cast from an index-like value (index or integer) to another index-like value.
Definition Utils.cpp:122
Type getElementTypeOrSelf(Type type)
Return the element type or return the type itself.
bool isZeroInteger(OpFoldResult v)
Return "true" if v is an integer value/attribute with constant value 0.
detail::constant_float_predicate_matcher m_AnyZeroFloat()
Matches a constant scalar / vector splat / tensor splat float (both positive and negative) zero.
Definition Matchers.h:399
Value getValueOrCreateConstantIndexOp(OpBuilder &b, Location loc, OpFoldResult ofr)
Converts an OpFoldResult to a Value.
Definition Utils.cpp:114
AffineMap compressDims(AffineMap map, const llvm::SmallBitVector &unusedDims)
Drop the dims that are listed in unusedDims.
llvm::SmallBitVector getUnusedDimsBitVector(ArrayRef< AffineMap > maps)
detail::constant_op_matcher m_Constant()
Matches a constant foldable operation.
Definition Matchers.h:369
LogicalResult matchAndRewrite(vector::ReductionOp op, PatternRewriter &rewriter) const override
BreakDownVectorReduction(MLIRContext *context, unsigned maxNumElementsToExtract, PatternBenefit benefit)
Canonicalization of a vector.contract a, b, c with row-major matmul semantics to a contraction suitab...
LogicalResult matchAndRewrite(vector::ContractionOp op, PatternRewriter &rewriter) const override
std::function< LogicalResult(vector::ContractionOp op)> FilterConstraintType
CanonicalizeContractMatmulToMMT(MLIRContext *context, PatternBenefit benefit, FilterConstraintType constraint)
Pattern to fold chained reduction to a series of vector additions and a final reduction....
LogicalResult matchAndRewrite(vector::ReductionOp op, PatternRewriter &rewriter) const override
For vectors with at least one unit dim, replaces: elementwise(a, b) with: sc_a = shape_cast(a) sc_b =...
OpTraitRewritePattern(MLIRContext *context, PatternBenefit benefit=1)
LogicalResult matchAndRewrite(Operation *op, PatternRewriter &rewriter) const override
Attempt to match against code rooted at the specified operation, which is the same operation code as ...
A pattern to drop unit dims from the iter_args of an scf.for.
LogicalResult matchAndRewrite(scf::ForOp forOp, PatternRewriter &rewriter) const override
A pattern to drop unit dims from vector.transpose.
LogicalResult matchAndRewrite(vector::TransposeOp op, PatternRewriter &rewriter) const override
Pattern to fold arithmetic extensions on floating point data types into vector contraction operations...
LogicalResult matchAndRewrite(vector::ContractionOp contractOp, PatternRewriter &rewriter) const override
Pattern to eliminate redundant zero-constants added to reduction operands. It's enough for there to b...
LogicalResult matchAndRewrite(vector::ReductionOp op, PatternRewriter &rewriter) const override
OpInterfaceRewritePattern is a wrapper around RewritePattern that allows for matching and rewriting a...
OpRewritePattern is a wrapper around RewritePattern that allows for matching and rewriting against an...
OpRewritePattern Base
Type alias to allow derived classes to inherit constructors with using Base::Base;.
OpRewritePattern(MLIRContext *context, PatternBenefit benefit=1, ArrayRef< StringRef > generatedNames={})
LogicalResult matchAndRewrite(Operation *op, PatternRewriter &rewriter) const final
Wrapper around the RewritePattern method that passes the derived op type.
A pattern for ops that implement MaskableOpInterface and that might be masked (i.e.