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