MLIR 24.0.0git
VectorDropLeadUnitDim.cpp
Go to the documentation of this file.
1//===- VectorDropLeadUnitDim.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#include <numeric>
10
16#include "mlir/IR/Builders.h"
18#include "llvm/ADT/STLExtras.h"
19
20#define DEBUG_TYPE "vector-drop-unit-dim"
21
22using namespace mlir;
23using namespace mlir::vector;
24
25// Trims leading one dimensions (fixed-width) from `oldType` and returns the
26// result type. Returns `vector<1xT>` if `oldType` only has one element.
27static VectorType trimLeadingUnitDims(VectorType oldType,
28 bool trimOnlyOneDim = false,
29 bool allowRank0 = false) {
30 ArrayRef<int64_t> oldShape = oldType.getShape();
31 ArrayRef<int64_t> newShape = oldShape;
32
33 ArrayRef<bool> oldScalableDims = oldType.getScalableDims();
34 ArrayRef<bool> newScalableDims = oldScalableDims;
35
36 while (!newShape.empty() && newShape.front() == 1 &&
37 !newScalableDims.front()) {
38 newShape = newShape.drop_front(1);
39 newScalableDims = newScalableDims.drop_front(1);
40
41 if (trimOnlyOneDim)
42 break;
43 }
44
45 // Make sure we have at least 1 dimension per vector type requirements.
46 if (newShape.empty() && !allowRank0) {
47 newShape = oldShape.take_back();
48 newScalableDims = oldType.getScalableDims().take_back();
49 }
50 return VectorType::get(newShape, oldType.getElementType(), newScalableDims);
51}
52
53/// Return a smallVector of size `rank` containing all zeros.
55 return SmallVector<int64_t>(rank, 0);
56}
57
59 ValueRange operands,
60 TypeRange resultTypes) {
61 OperationState state(op->getLoc(), op->getName(), operands, resultTypes,
62 op->getDiscardableAttrDictionary().getValue());
64 return builder.create(state);
65}
66namespace {
67
68// Casts away leading one dimensions in vector.extract_strided_slice's vector
69// input by inserting vector.broadcast.
70struct CastAwayExtractStridedSliceLeadingOneDim
71 : public OpRewritePattern<vector::ExtractStridedSliceOp> {
72 using Base::Base;
73
74 LogicalResult matchAndRewrite(vector::ExtractStridedSliceOp extractOp,
75 PatternRewriter &rewriter) const override {
76 // vector.extract_strided_slice requires the input and output vector to have
77 // the same rank. Here we drop leading one dimensions from the input vector
78 // type to make sure we don't cause mismatch.
79 VectorType oldSrcType = extractOp.getSourceVectorType();
80 VectorType newSrcType = trimLeadingUnitDims(oldSrcType);
81
82 if (newSrcType.getRank() == oldSrcType.getRank())
83 return failure();
84
85 int64_t dropCount = oldSrcType.getRank() - newSrcType.getRank();
86
87 VectorType oldDstType = extractOp.getType();
88 VectorType newDstType =
89 VectorType::get(oldDstType.getShape().drop_front(dropCount),
90 oldDstType.getElementType(),
91 oldDstType.getScalableDims().drop_front(dropCount));
92
93 Location loc = extractOp.getLoc();
94
95 Value newSrcVector = rewriter.createOrFold<ShapeCastOp>(
96 loc, newSrcType, extractOp.getSource());
97
98 // The offsets/sizes/strides attribute can have a less number of elements
99 // than the input vector's rank: it is meant for the leading dimensions.
100 auto newOffsets = rewriter.getArrayAttr(
101 extractOp.getOffsets().getValue().drop_front(dropCount));
102 auto newSizes = rewriter.getArrayAttr(
103 extractOp.getSizes().getValue().drop_front(dropCount));
104 auto newStrides = rewriter.getArrayAttr(
105 extractOp.getStrides().getValue().drop_front(dropCount));
106
107 auto newExtractOp = vector::ExtractStridedSliceOp::create(
108 rewriter, loc, newDstType, newSrcVector, newOffsets, newSizes,
109 newStrides);
110
111 rewriter.replaceOpWithNewOp<vector::ShapeCastOp>(extractOp, oldDstType,
112 newExtractOp);
113
114 return success();
115 }
116};
117
118// Casts away leading one dimensions in vector.insert_strided_slice's vector
119// inputs by inserting vector.broadcast.
120struct CastAwayInsertStridedSliceLeadingOneDim
121 : public OpRewritePattern<vector::InsertStridedSliceOp> {
122 using Base::Base;
123
124 LogicalResult matchAndRewrite(vector::InsertStridedSliceOp insertOp,
125 PatternRewriter &rewriter) const override {
126 VectorType oldSrcType = insertOp.getSourceVectorType();
127 VectorType newSrcType = trimLeadingUnitDims(oldSrcType);
128 VectorType oldDstType = insertOp.getDestVectorType();
129 VectorType newDstType = trimLeadingUnitDims(oldDstType);
130
131 int64_t srcDropCount = oldSrcType.getRank() - newSrcType.getRank();
132 int64_t dstDropCount = oldDstType.getRank() - newDstType.getRank();
133 if (srcDropCount == 0 && dstDropCount == 0)
134 return failure();
135
136 // Trim leading one dimensions from both operands.
137 Location loc = insertOp.getLoc();
138
139 Value newSrcVector = rewriter.createOrFold<vector::ShapeCastOp>(
140 loc, newSrcType, insertOp.getValueToStore());
141 Value newDstVector = rewriter.createOrFold<vector::ShapeCastOp>(
142 loc, newDstType, insertOp.getDest());
143
144 auto newOffsets = rewriter.getArrayAttr(
145 insertOp.getOffsets().getValue().take_back(newDstType.getRank()));
146 auto newStrides = rewriter.getArrayAttr(
147 insertOp.getStrides().getValue().take_back(newSrcType.getRank()));
148
149 auto newInsertOp = vector::InsertStridedSliceOp::create(
150 rewriter, loc, newDstType, newSrcVector, newDstVector, newOffsets,
151 newStrides);
152
153 rewriter.replaceOpWithNewOp<vector::ShapeCastOp>(insertOp, oldDstType,
154 newInsertOp);
155
156 return success();
157 }
158};
159
160// Casts away leading one dimensions in vector.insert's vector inputs by
161// inserting vector.shape_cast.
162struct CastAwayInsertLeadingOneDim : public OpRewritePattern<vector::InsertOp> {
163 using Base::Base;
164
165 LogicalResult matchAndRewrite(vector::InsertOp insertOp,
166 PatternRewriter &rewriter) const override {
167 Type oldSrcType = insertOp.getValueToStoreType();
168 Type newSrcType = oldSrcType;
169 int64_t oldSrcRank = 0, newSrcRank = 0;
170 if (auto type = dyn_cast<VectorType>(oldSrcType)) {
171 newSrcType = trimLeadingUnitDims(type);
172 oldSrcRank = type.getRank();
173 newSrcRank = cast<VectorType>(newSrcType).getRank();
174 }
175
176 VectorType oldDstType = insertOp.getDestVectorType();
177 VectorType newDstType = trimLeadingUnitDims(oldDstType);
178
179 int64_t srcDropCount = oldSrcRank - newSrcRank;
180 int64_t dstDropCount = oldDstType.getRank() - newDstType.getRank();
181 if (srcDropCount == 0 && dstDropCount == 0)
182 return failure();
183
184 // Trim leading one dimensions from both operands.
185 Location loc = insertOp.getLoc();
186
187 Value newSrcVector = insertOp.getValueToStore();
188 if (oldSrcRank != 0) {
189 newSrcVector = rewriter.createOrFold<vector::ShapeCastOp>(
190 loc, cast<VectorType>(newSrcType), insertOp.getValueToStore());
191 }
192 Value newDstVector = rewriter.createOrFold<vector::ShapeCastOp>(
193 loc, newDstType, insertOp.getDest());
194
195 // New position rank needs to be computed in two steps: (1) if destination
196 // type has leading unit dims, we also trim the position array accordingly,
197 // then (2) if source type also has leading unit dims, we need to append
198 // zeroes to the position array accordingly.
199 unsigned oldPosRank = insertOp.getNumIndices();
200 unsigned newPosRank = std::max<int64_t>(0, oldPosRank - dstDropCount);
201 SmallVector<OpFoldResult> oldPosition = insertOp.getMixedPosition();
202 SmallVector<OpFoldResult> newPosition =
203 llvm::to_vector(ArrayRef(oldPosition).take_back(newPosRank));
204 newPosition.resize(newDstType.getRank() - newSrcRank,
205 rewriter.getI64IntegerAttr(0));
206
207 auto newInsertOp = vector::InsertOp::create(rewriter, loc, newSrcVector,
208 newDstVector, newPosition);
209
210 rewriter.replaceOpWithNewOp<vector::ShapeCastOp>(insertOp, oldDstType,
211 newInsertOp);
212
213 return success();
214 }
215};
216
217static Value dropUnitDimsFromMask(OpBuilder &b, Location loc, Value mask,
218 VectorType newType, AffineMap newMap) {
219 VectorType newMaskType = inferTransferOpMaskType(newType, newMap);
220
221 return vector::ShapeCastOp::create(b, loc, newMaskType, mask);
222}
223
224// Turns vector.transfer_read on vector with leading 1 dimensions into
225// vector.shape_cast followed by vector.transfer_read on vector without leading
226// 1 dimensions.
227struct CastAwayTransferReadLeadingOneDim
228 : public OpRewritePattern<vector::TransferReadOp> {
229 using Base::Base;
230
231 LogicalResult matchAndRewrite(vector::TransferReadOp read,
232 PatternRewriter &rewriter) const override {
233 // TODO(#78787): Not supported masked op yet.
234 if (cast<MaskableOpInterface>(read.getOperation()).isMasked())
235 return failure();
236
237 if (read.getTransferRank() == 0)
238 return rewriter.notifyMatchFailure(
239 read, "Nothing to trim - the transfer itself has rank zero");
240
241 auto shapedType = cast<ShapedType>(read.getBase().getType());
242 if (shapedType.getElementType() != read.getVectorType().getElementType())
243 return failure();
244
245 VectorType oldType = read.getVectorType();
246 VectorType newType = trimLeadingUnitDims(oldType);
247
248 if (newType == oldType)
249 return failure();
250
251 AffineMap oldMap = read.getPermutationMap();
252 ArrayRef<AffineExpr> newResults =
253 oldMap.getResults().take_back(newType.getRank());
254 AffineMap newMap =
255 AffineMap::get(oldMap.getNumDims(), oldMap.getNumSymbols(), newResults,
256 rewriter.getContext());
257
258 ArrayAttr inBoundsAttr;
259 if (read.getInBounds())
260 inBoundsAttr = rewriter.getArrayAttr(
261 read.getInBoundsAttr().getValue().take_back(newType.getRank()));
262
263 Value mask = Value();
264 if (read.getMask())
265 mask = dropUnitDimsFromMask(rewriter, read.getLoc(), read.getMask(),
266 newType, newMap);
267
268 auto newRead = vector::TransferReadOp::create(
269 rewriter, read.getLoc(), newType, read.getBase(), read.getIndices(),
270 AffineMapAttr::get(newMap), read.getPadding(), mask, inBoundsAttr);
271 rewriter.replaceOpWithNewOp<vector::BroadcastOp>(read, oldType, newRead);
272
273 return success();
274 }
275};
276
277// Turns vector.transfer_write on vector with leading 1 dimensions into
278// vector.shape_cast followed by vector.transfer_write on vector without leading
279// 1 dimensions.
280struct CastAwayTransferWriteLeadingOneDim
281 : public OpRewritePattern<vector::TransferWriteOp> {
282 using Base::Base;
283
284 LogicalResult matchAndRewrite(vector::TransferWriteOp write,
285 PatternRewriter &rewriter) const override {
286 // TODO(#78787): Not supported masked op yet.
287 if (cast<MaskableOpInterface>(write.getOperation()).isMasked())
288 return failure();
289
290 if (write.getTransferRank() == 0)
291 return rewriter.notifyMatchFailure(
292 write, "Nothing to trim - the transfer itself has rank zero");
293
294 auto shapedType = dyn_cast<ShapedType>(write.getBase().getType());
295 if (shapedType.getElementType() != write.getVectorType().getElementType())
296 return failure();
297
298 VectorType oldType = write.getVectorType();
299 VectorType newType = trimLeadingUnitDims(oldType);
300 if (newType == oldType)
301 return failure();
302
303 AffineMap oldMap = write.getPermutationMap();
304 ArrayRef<AffineExpr> newResults =
305 oldMap.getResults().take_back(newType.getRank());
306 AffineMap newMap =
307 AffineMap::get(oldMap.getNumDims(), oldMap.getNumSymbols(), newResults,
308 rewriter.getContext());
309
310 ArrayAttr inBoundsAttr;
311 if (write.getInBounds())
312 inBoundsAttr = rewriter.getArrayAttr(
313 write.getInBoundsAttr().getValue().take_back(newType.getRank()));
314
315 auto newVector = rewriter.createOrFold<vector::ShapeCastOp>(
316 write.getLoc(), newType, write.getVector());
317
318 if (write.getMask()) {
319 Value newMask = dropUnitDimsFromMask(rewriter, write.getLoc(),
320 write.getMask(), newType, newMap);
321 rewriter.replaceOpWithNewOp<vector::TransferWriteOp>(
322 write, newVector, write.getBase(), write.getIndices(),
323 AffineMapAttr::get(newMap), newMask, inBoundsAttr);
324 return success();
325 }
326
327 rewriter.replaceOpWithNewOp<vector::TransferWriteOp>(
328 write, newVector, write.getBase(), write.getIndices(),
329 AffineMapAttr::get(newMap), inBoundsAttr);
330 return success();
331 }
332};
333
334} // namespace
335
336// Takes `oldVal` and "drops" the leading unit dim with either ShapeCastOp or
337// ExtractOp. The latter is used for rank-1 vectors to make sure that a scalar
338// (as opposed to rank-0 vector) is generated. This is a requirement of e.g.
339// ContractOp.
341 Location loc,
342 mlir::Value oldVal) {
343 auto oldValTy = cast<VectorType>(oldVal.getType());
344 if (oldValTy.getRank() == 1) {
345 return rewriter.createOrFold<ExtractOp>(loc, oldVal, 0);
346 }
347
348 return rewriter.createOrFold<ShapeCastOp>(
349 loc,
350 trimLeadingUnitDims(oldValTy,
351 /*trimOnlyOneDim=*/true,
352 /*allowRank0=*/true),
353 oldVal);
354}
355
356// Takes `oldVal` and "adds" leading unit dim with either ShapeCastOp or
357// BroadcastOp. The latter is used for scalar inputs as ShapeCastOp cannot
358// "broadcast" from a scalar. Scalars are used (instead of rank-0 vectors) as
359// ContractOp operands.
361 Location loc,
362 mlir::Value oldVal,
363 mlir::Type newTy) {
364 if (!isa<VectorType>(oldVal.getType())) {
365 return rewriter.createOrFold<BroadcastOp>(loc, newTy, oldVal);
366 }
367
368 return rewriter.createOrFold<ShapeCastOp>(loc, newTy, oldVal);
369}
370
371FailureOr<Value>
372mlir::vector::castAwayContractionLeadingOneDim(vector::ContractionOp contractOp,
373 MaskingOpInterface maskingOp,
374 RewriterBase &rewriter) {
375 VectorType oldAccType = dyn_cast<VectorType>(contractOp.getAccType());
376 if (oldAccType == nullptr)
377 return failure();
378 if (oldAccType.getRank() < 1)
379 return failure();
380 if (oldAccType.getShape()[0] != 1 || oldAccType.getScalableDims()[0])
381 return failure();
382
383 auto oldIndexingMaps = contractOp.getIndexingMapsArray();
384 SmallVector<AffineMap> newIndexingMaps;
385
386 auto oldIteratorTypes = contractOp.getIteratorTypes();
387 SmallVector<Attribute> newIteratorTypes;
388
389 // 0-th dim from the accumulator
390 int64_t dimToDrop = oldIndexingMaps[2].getDimPosition(0);
391
392 if (!isParallelIterator(oldIteratorTypes[dimToDrop]))
393 // only parallel type iterators can be dropped.
394 return failure();
395
396 for (const auto &it : llvm::enumerate(oldIteratorTypes)) {
397 int64_t currDim = it.index();
398 if (currDim == dimToDrop)
399 continue;
400 newIteratorTypes.push_back(it.value());
401 }
402
403 SmallVector<Value> operands = {contractOp.getLhs(), contractOp.getRhs(),
404 contractOp.getAcc()};
405 SmallVector<Value> newOperands;
406 auto loc = contractOp.getLoc();
407
408 for (const auto &it : llvm::enumerate(oldIndexingMaps)) {
409 // Check if the dim to be dropped exists as a leading dim in the operand
410 // if it does then we use vector.extract to drop it.
411 bool validExtract = false;
413 auto map = it.value();
414 int64_t orginalZeroDim = it.value().getDimPosition(0);
415 if (orginalZeroDim != dimToDrop) {
416 // There are two reasons to be in this path, 1. We need to
417 // transpose the operand to make the dim to be dropped
418 // leading. 2. The dim to be dropped does not exist and in
419 // that case we dont want to add a unit transpose but we must
420 // check all the indices to make sure this is the case.
421 bool transposeNeeded = false;
423 SmallVector<AffineExpr> transposeResults;
424
425 for (int64_t i = 0, e = map.getNumResults(); i < e; ++i) {
426 int64_t currDim = map.getDimPosition(i);
427 if (currDim == dimToDrop) {
428 transposeNeeded = true;
429 perm.insert(perm.begin(), i);
430 auto targetExpr = rewriter.getAffineDimExpr(currDim);
431 transposeResults.insert(transposeResults.begin(), targetExpr);
432 } else {
433 perm.push_back(i);
434 auto targetExpr = rewriter.getAffineDimExpr(currDim);
435 transposeResults.push_back(targetExpr);
436 }
437 }
438
439 // Checks if only the outer, unit dimensions (of size 1) are permuted.
440 // Such transposes do not materially effect the underlying vector and can
441 // be omitted. EG: perm [1, 0, 2] applied to vector<1x1x8xi32>
442 bool transposeNonOuterUnitDims = false;
443 auto operandShape = cast<ShapedType>(operands[it.index()].getType());
444 for (auto [index, dim] :
445 llvm::enumerate(ArrayRef<int64_t>(perm).drop_back(1))) {
446 if (dim != static_cast<int64_t>(index) &&
447 operandShape.getDimSize(index) != 1) {
448 transposeNonOuterUnitDims = true;
449 break;
450 }
451 }
452
453 // Do the transpose now if needed so that we can drop the
454 // correct dim using extract later.
455 if (transposeNeeded) {
456 map = AffineMap::get(map.getNumDims(), 0, transposeResults,
457 contractOp.getContext());
458 if (transposeNonOuterUnitDims) {
459 // TODO: While the discussion on the validity of folding
460 // TransposeOp into ShapeCastOp continues, see e.g.
461 // * https://github.com/llvm/llvm-project/pull/219611,
462 // keep the explicit TransposeOp here. Note that existing TransposeOp
463 // folders already turn it into a ShapeCastOp, as demonstrated by the
464 // tests. Revisit this and consider inserting ShapeCastOp directly
465 // once the discussion progresses.
466 operands[it.index()] = rewriter.createOrFold<vector::TransposeOp>(
467 loc, operands[it.index()], perm);
468 }
469 }
470 }
471 // We have taken care to have the dim to be dropped be
472 // the leading dim. If its still not leading that means it
473 // does not exist in this operand and hence we do not need
474 // an extract.
475 if (map.getDimPosition(0) == dimToDrop)
476 validExtract = true;
477
478 for (int64_t i = 0, e = map.getNumResults(); i < e; ++i) {
479 int64_t currDim = map.getDimPosition(i);
480 if (currDim == dimToDrop)
481 // This is the dim we are dropping.
482 continue;
483 auto targetExpr = rewriter.getAffineDimExpr(
484 currDim < dimToDrop ? currDim : currDim - 1);
485 results.push_back(targetExpr);
486 }
487 newIndexingMaps.push_back(AffineMap::get(map.getNumDims() - 1, 0, results,
488 contractOp.getContext()));
489 // Extract if its a valid extraction, otherwise use the operand
490 // without extraction.
491 auto oldVal = operands[it.index()];
492 newOperands.push_back(
493 validExtract
494 ? dropLeadingUnitDimViaShapeCastOrExtract(rewriter, loc, oldVal)
495 : oldVal);
496 }
497
498 // Depending on whether this vector.contract is masked, the replacing Op
499 // should either be a new vector.contract Op or vector.mask Op.
500 Operation *newOp = vector::ContractionOp::create(
501 rewriter, loc, newOperands[0], newOperands[1], newOperands[2],
502 rewriter.getAffineMapArrayAttr(newIndexingMaps),
503 rewriter.getArrayAttr(newIteratorTypes), contractOp.getKind());
504
505 if (maskingOp) {
506 auto newMask = rewriter.createOrFold<ShapeCastOp>(
507 loc,
508 trimLeadingUnitDims(cast<VectorType>(maskingOp.getMask().getType()),
509 /*trimOnlyOneDim=*/true, /*allowRank0=*/true),
510 maskingOp.getMask());
511
512 newOp = mlir::vector::maskOperation(rewriter, newOp, newMask);
513 }
514
516 rewriter, loc, newOp->getResult(0), contractOp->getResultTypes()[0]);
517}
518
519namespace {
520
521/// Turns vector.contract on vector with leading 1 dimensions into
522/// vector.shape_cast followed by vector.contract on vector without leading
523/// 1 dimensions. Also performs transpose of lhs and rhs operands if required.
524///
525/// TODO: While the discussion on the validity of folding TransposeOp into
526/// ShapeCastOp continues, see e.g.
527/// * https://github.com/llvm/llvm-project/pull/219611,
528/// keep the explicit TransposeOp here. Once the discussion settles, revisit and
529/// consider replacing TransposeOp with ShapeCastOp.
530struct CastAwayContractionLeadingOneDim
531 : public MaskableOpRewritePattern<vector::ContractionOp> {
532 using MaskableOpRewritePattern::MaskableOpRewritePattern;
533
534 FailureOr<Value>
535 matchAndRewriteMaskableOp(vector::ContractionOp contractOp,
536 MaskingOpInterface maskingOp,
537 PatternRewriter &rewriter) const override {
538 return castAwayContractionLeadingOneDim(contractOp, maskingOp, rewriter);
539 }
540};
541
542/// Looks at elementwise operations on vectors with at least one leading
543/// dimension equal 1, e.g. vector<1x[4]x1xf32> (but not vector<2x[4]x1xf32>),
544/// and cast aways the leading one dimensions (_plural_) and then broadcasts
545/// the results.
546///
547/// Example before:
548/// %1 = arith.mulf %arg0, %arg1 : vector<1x4x1xf32>
549/// Example after:
550/// %2 = arith.mulf %0, %1 : vector<4x1xf32>
551/// %3 = vector.broadcast %2 : vector<4x1xf32> to vector<1x4x1xf32>
552///
553/// Does support scalable vectors.
554class CastAwayElementwiseLeadingOneDim : public RewritePattern {
555public:
556 CastAwayElementwiseLeadingOneDim(MLIRContext *context,
557 PatternBenefit benefit = 1)
558 : RewritePattern(MatchAnyOpTypeTag(), benefit, context) {}
559
560 LogicalResult matchAndRewrite(Operation *op,
561 PatternRewriter &rewriter) const override {
563 return failure();
564 auto vecType = dyn_cast<VectorType>(op->getResultTypes()[0]);
565 if (!vecType)
566 return failure();
567 VectorType newVecType = trimLeadingUnitDims(vecType);
568 if (newVecType == vecType)
569 return failure();
570 int64_t dropDim = vecType.getRank() - newVecType.getRank();
571 SmallVector<Value, 4> newOperands;
572 for (Value operand : op->getOperands()) {
573 if (auto opVecType = dyn_cast<VectorType>(operand.getType())) {
574 newOperands.push_back(vector::ExtractOp::create(
575 rewriter, op->getLoc(), operand, splatZero(dropDim)));
576 } else {
577 newOperands.push_back(operand);
578 }
579 }
580 Operation *newOp =
581 createWithProperties(rewriter, op, newOperands, TypeRange{newVecType});
582 rewriter.replaceOpWithNewOp<vector::BroadcastOp>(op, vecType,
583 newOp->getResult(0));
584 return success();
585 }
586};
587} // namespace
588
589// Drops `dropDim` leading dimensions from `operand` using vector.extract when
590// those dims are all non-scalable units (the cheap, structural rewrite); falls
591// back to vector.shape_cast otherwise.
593 Value operand, int64_t nDropped) {
594 auto oldType = cast<VectorType>(operand.getType());
595 ArrayRef<int64_t> leadingShape = oldType.getShape().take_front(nDropped);
596 ArrayRef<bool> leadingScalable =
597 oldType.getScalableDims().take_front(nDropped);
598 bool extractable =
599 llvm::all_of(leadingShape, [](int64_t d) { return d == 1; }) &&
600 llvm::none_of(leadingScalable, [](bool s) { return s; });
601 if (extractable)
602 return vector::ExtractOp::create(b, loc, operand, splatZero(nDropped));
603 VectorType newType = VectorType::get(
604 oldType.getShape().drop_front(nDropped), oldType.getElementType(),
605 oldType.getScalableDims().drop_front(nDropped));
606 return vector::ShapeCastOp::create(b, loc, newType, operand);
607}
608
609namespace {
610
611// Drops leading 1 dimensions from load-like memory operaitons. REmoves leading
612// unit dimensions from the result types and then broadcasts back in those 1s,
613// while also extracting (or shape_cast-ing) any leading unit dimensions on
614// the input operands.
615template <typename OpTy>
616struct CastAwayLoadLikeLeadingOneDim : public OpRewritePattern<OpTy> {
617 using OpRewritePattern<OpTy>::OpRewritePattern;
618
619 LogicalResult matchAndRewrite(OpTy op,
620 PatternRewriter &rewriter) const override {
621 VectorType oldResultType = op.getVectorType();
622 VectorType newResultType = trimLeadingUnitDims(oldResultType);
623 if (newResultType == oldResultType)
624 return failure();
625 int64_t nDropped = oldResultType.getRank() - newResultType.getRank();
626
627 Location loc = op.getLoc();
628 SmallVector<Value> newOperands;
629 newOperands.reserve(op->getNumOperands());
630 for (Value operand : op->getOperands()) {
631 if (isa<VectorType>(operand.getType())) {
632 newOperands.push_back(
633 dropLeadingOneDimsFromOperand(rewriter, loc, operand, nDropped));
634 } else {
635 newOperands.push_back(operand);
636 }
637 }
638
639 Operation *newOp = createWithProperties(rewriter, op, newOperands,
640 TypeRange{newResultType});
641 rewriter.replaceOpWithNewOp<vector::BroadcastOp>(op, oldResultType,
642 newOp->getResult(0));
643 return success();
644 }
645};
646
647// Drops leading 1 dimensions from store-like memory ops. Extracts or
648// `shape_cast`s away those leading unit dimensions and leaves any scalar
649// operands alone.
650template <typename OpTy>
651struct CastAwayStoreLikeLeadingOneDim : public OpRewritePattern<OpTy> {
652 using OpRewritePattern<OpTy>::OpRewritePattern;
653
654 LogicalResult matchAndRewrite(OpTy op,
655 PatternRewriter &rewriter) const override {
656 VectorType oldVecType = op.getVectorType();
657 VectorType newVecType = trimLeadingUnitDims(oldVecType);
658 if (newVecType == oldVecType)
659 return failure();
660 int64_t nDropped = oldVecType.getRank() - newVecType.getRank();
661
662 Location loc = op.getLoc();
663 SmallVector<Value> newOperands;
664 newOperands.reserve(op->getNumOperands());
665 for (Value operand : op->getOperands()) {
666 if (isa<VectorType>(operand.getType())) {
667 newOperands.push_back(
668 dropLeadingOneDimsFromOperand(rewriter, loc, operand, nDropped));
669 } else {
670 newOperands.push_back(operand);
671 }
672 }
673
674 Operation *newOp =
675 createWithProperties(rewriter, op, newOperands, op->getResultTypes());
676 rewriter.replaceOp(op, newOp->getResults());
677 return success();
678 }
679};
680
681// Drops leading 1 dimensions from vector.constant_mask and inserts a
682// vector.broadcast back to the original shape.
683struct CastAwayConstantMaskLeadingOneDim
684 : public OpRewritePattern<vector::ConstantMaskOp> {
685 using Base::Base;
686
687 LogicalResult matchAndRewrite(vector::ConstantMaskOp mask,
688 PatternRewriter &rewriter) const override {
689 VectorType oldType = mask.getType();
690 VectorType newType = trimLeadingUnitDims(oldType);
691
692 if (newType == oldType)
693 return failure();
694
695 int64_t dropDim = oldType.getRank() - newType.getRank();
696 ArrayRef<int64_t> dimSizes = mask.getMaskDimSizes();
697
698 // If any of the dropped unit dims has a size of `0`, the entire mask is a
699 // zero mask, else the unit dim has no effect on the mask.
700 int64_t flatLeadingSize =
701 llvm::product_of(dimSizes.take_front(dropDim + 1));
702 SmallVector<int64_t> newDimSizes = {flatLeadingSize};
703 newDimSizes.append(dimSizes.begin() + dropDim + 1, dimSizes.end());
704
705 auto newMask = vector::ConstantMaskOp::create(rewriter, mask.getLoc(),
706 newType, newDimSizes);
707 rewriter.replaceOpWithNewOp<vector::BroadcastOp>(mask, oldType, newMask);
708 return success();
709 }
710};
711
712} // namespace
713
714void mlir::vector::populateCastAwayVectorLeadingOneDimPatterns(
715 RewritePatternSet &patterns, PatternBenefit benefit) {
716 patterns
717 .add<CastAwayExtractStridedSliceLeadingOneDim,
718 CastAwayInsertStridedSliceLeadingOneDim, CastAwayInsertLeadingOneDim,
719 CastAwayConstantMaskLeadingOneDim, CastAwayTransferReadLeadingOneDim,
720 CastAwayTransferWriteLeadingOneDim, CastAwayElementwiseLeadingOneDim,
721 CastAwayContractionLeadingOneDim,
722 CastAwayLoadLikeLeadingOneDim<vector::LoadOp>,
723 CastAwayLoadLikeLeadingOneDim<vector::MaskedLoadOp>,
724 CastAwayLoadLikeLeadingOneDim<vector::ExpandLoadOp>,
725 CastAwayLoadLikeLeadingOneDim<vector::GatherOp>,
726 CastAwayStoreLikeLeadingOneDim<vector::StoreOp>,
727 CastAwayStoreLikeLeadingOneDim<vector::MaskedStoreOp>,
728 CastAwayStoreLikeLeadingOneDim<vector::CompressStoreOp>,
729 CastAwayStoreLikeLeadingOneDim<vector::ScatterOp>>(
730 patterns.getContext(), benefit);
731}
return success()
b
Return true if permutation is a valid permutation of the outer_dims_perm (case OuterOrInnerPerm::Oute...
ArrayAttr()
static VectorType trimLeadingUnitDims(VectorType oldType, bool trimOnlyOneDim=false, bool allowRank0=false)
static SmallVector< int64_t > splatZero(int64_t rank)
Return a smallVector of size rank containing all zeros.
static Value restoreLeadingUnitDimViaShapeCastOrBcast(RewriterBase &rewriter, Location loc, mlir::Value oldVal, mlir::Type newTy)
static Value dropLeadingUnitDimViaShapeCastOrExtract(RewriterBase &rewriter, Location loc, mlir::Value oldVal)
static Value dropLeadingOneDimsFromOperand(OpBuilder &b, Location loc, Value operand, int64_t nDropped)
static Operation * createWithProperties(OpBuilder &builder, Operation *op, ValueRange operands, TypeRange resultTypes)
A multi-dimensional affine map Affine map's are immutable like Type's, and they are uniqued.
Definition AffineMap.h:46
static AffineMap get(MLIRContext *context)
Returns a zero result affine map with no dimensions or symbols: () -> ().
unsigned getNumSymbols() const
unsigned getNumDims() const
ArrayRef< AffineExpr > getResults() const
static AffineMap getPermutationMap(ArrayRef< unsigned > permutation, MLIRContext *context)
Returns an AffineMap representing a permutation.
IntegerAttr getI64IntegerAttr(int64_t value)
Definition Builders.cpp:120
AffineExpr getAffineDimExpr(unsigned position)
Definition Builders.cpp:373
ArrayAttr getArrayAttr(ArrayRef< Attribute > value)
Definition Builders.cpp:275
MLIRContext * getContext() const
Definition Builders.h:56
ArrayAttr getAffineMapArrayAttr(ArrayRef< AffineMap > values)
Definition Builders.cpp:327
This class defines the main interface for locations in MLIR and acts as a non-nullable wrapper around...
Definition Location.h:76
This class helps build Operations.
Definition Builders.h:210
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
Operation is the basic unit of execution within MLIR.
Definition Operation.h:87
OpResult getResult(unsigned idx)
Get the 'idx'th result of this operation.
Definition Operation.h:432
Location getLoc()
The source location the operation was defined or derived from.
Definition Operation.h:240
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
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...
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.
RewritePattern is the common base class for all DAG to DAG replacements.
This class coordinates the application of a rewrite on a set of IR, providing a way for clients to tr...
virtual void replaceOp(Operation *op, ValueRange newValues)
Replace the results of the given (original) operation with the specified list of values (replacements...
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,...
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
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
Type getType() const
Return the type of this value.
Definition Value.h:105
Location getLoc() const
Return the location of this value.
Definition Value.cpp:24
bool hasElementwiseMappableTraits(Operation *op)
Together, Elementwise, Scalarizable, Vectorizable, and Tensorizable provide an easy way for scalar op...
Operation * maskOperation(OpBuilder &builder, Operation *maskableOp, Value mask, Value passthru=Value())
Creates a vector.mask operation around a maskable operation.
VectorType inferTransferOpMaskType(VectorType vecType, AffineMap permMap)
Infers the mask type for a transfer op given its vector type and permutation map.
bool isParallelIterator(Attribute attr)
Returns true if attr has "parallel" iterator type semantics.
Definition VectorOps.h:151
Include the generated interface declarations.
OpRewritePattern is a wrapper around RewritePattern that allows for matching and rewriting against an...
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.
A pattern for ops that implement MaskableOpInterface and that might be masked (i.e.