MLIR 24.0.0git
Vectorization.cpp
Go to the documentation of this file.
1//===- Vectorization.cpp - Implementation of linalg Vectorization ---------===//
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 the linalg dialect Vectorization transformations.
10//
11//===----------------------------------------------------------------------===//
13
28#include "mlir/IR/AffineExpr.h"
29#include "mlir/IR/AffineMap.h"
30#include "mlir/IR/Builders.h"
35#include "mlir/IR/Value.h"
36#include "mlir/Support/LLVM.h"
38#include "llvm/ADT/STLExtras.h"
39#include "llvm/ADT/Sequence.h"
40#include "llvm/ADT/SmallPtrSet.h"
41#include "llvm/ADT/SmallVector.h"
42#include "llvm/ADT/SmallVectorExtras.h"
43#include "llvm/ADT/TypeSwitch.h"
44#include "llvm/Support/DebugLog.h"
45#include "llvm/Support/InterleavedRange.h"
46#include "llvm/Support/MathExtras.h"
47#include "llvm/Support/raw_ostream.h"
48#include <optional>
49
50using namespace mlir;
51using namespace mlir::linalg;
52
53#define DEBUG_TYPE "linalg-vectorization"
54
55/// Try to vectorize `convOp` as a convolution.
56static FailureOr<Operation *>
57vectorizeConvolution(RewriterBase &rewriter, LinalgOp convOp,
58 ArrayRef<int64_t> inputVecSizes = {},
59 ArrayRef<bool> inputVecScalableFlags = {},
60 bool flatten1DDepthwiseConv = false);
61
62/// Vectorize tensor::InsertSliceOp with:
63/// * vector::TransferReadOp + vector::TransferWriteOp
64/// The vector sizes are either:
65/// * user-provided in `inputVectorSizes`, or
66/// * inferred from the static dims in the input and output tensors.
67/// Bails out if:
68/// * vector sizes are not user-provided, and
69/// * at least one dim is dynamic (in both the input and output tensors).
70///
71/// Before:
72/// !t_in_type = tensor<1x2x3xf32>
73/// !t_out_type = tensor<9x8x7x1x2x3xf32>
74/// !v_type = vector<1x2x3xf32>
75/// %inserted_slice = tensor.insert_slice %src into %dest ... : !t_in_type
76/// into !t_out_type
77/// After:
78/// %read = vector.transfer_read %src[...], %pad ... : !t_in_type, !v_type
79/// %write = vector.transfer_write %read, %dest ... : !v_type, !t_out_type
80static LogicalResult
81vectorizeAsInsertSliceOp(RewriterBase &rewriter, tensor::InsertSliceOp sliceOp,
82 ArrayRef<int64_t> inputVectorSizes,
83 SmallVectorImpl<Value> &newResults);
84
85/// Returns the effective Pad value for the input op, provided it's a scalar.
86///
87/// Many Ops exhibit pad-like behaviour, but this isn't always explicit. If
88/// this Op performs padding, retrieve the padding value provided that it's
89/// a scalar and static/fixed for all the padded values. Returns an empty value
90/// otherwise.
92
93/// Helper function to extract the input slices after filter is unrolled along
94/// kw.
97 int64_t nSize, int64_t wSize, int64_t cSize,
98 int64_t kwSize, int strideW, int dilationW,
99 int64_t wSizeStep, bool isSingleChanneled) {
101 if (isSingleChanneled) {
102 // Extract input slice of size {wSizeStep} @ [w + kw] for non-channeled
103 // convolution.
104 SmallVector<int64_t> sizes = {wSizeStep};
105 SmallVector<int64_t> strides = {1};
106 for (int64_t kw = 0; kw < kwSize; ++kw) {
107 for (int64_t w = 0; w < wSize; w += wSizeStep) {
108 result.push_back(vector::ExtractStridedSliceOp::create(
109 rewriter, loc, input, /*offsets=*/ArrayRef<int64_t>{w + kw}, sizes,
110 strides));
111 }
112 }
113 } else {
114 // Extract lhs slice of size {n, wSizeStep, c} @ [0, sw * w + dw * kw, 0]
115 // for channeled convolution.
116 SmallVector<int64_t> sizes = {nSize, wSizeStep, cSize};
117 SmallVector<int64_t> strides = {1, 1, 1};
118 for (int64_t kw = 0; kw < kwSize; ++kw) {
119 for (int64_t w = 0; w < wSize; w += wSizeStep) {
120 result.push_back(vector::ExtractStridedSliceOp::create(
121 rewriter, loc, input,
122 /*offsets=*/ArrayRef<int64_t>{0, w * strideW + kw * dilationW, 0},
123 sizes, strides));
124 }
125 }
126 }
127 return result;
128}
129
130/// Helper function to extract the filter slices after filter is unrolled along
131/// kw.
133 Location loc, Value filter,
134 int64_t kwSize) {
136 // Extract rhs slice of size [{c, f} for channeled convolutions and {1} for
137 // non-chanelled convolution] @ [kw].
138 for (int64_t kw = 0; kw < kwSize; ++kw) {
139 result.push_back(vector::ExtractOp::create(
140 rewriter, loc, filter, /*offsets=*/ArrayRef<int64_t>{kw}));
141 }
142 return result;
143}
144
145/// Helper function to extract the result slices after filter is unrolled along
146/// kw.
149 int64_t nSize, int64_t wSize, int64_t fSize,
150 int64_t wSizeStep, bool isSingleChanneled) {
152 if (isSingleChanneled) {
153 // Extract res slice: {wSizeStep} @ [w] for non-channeled convolution.
154 SmallVector<int64_t> sizes = {wSizeStep};
155 SmallVector<int64_t> strides = {1};
156 for (int64_t w = 0; w < wSize; w += wSizeStep) {
157 result.push_back(vector::ExtractStridedSliceOp::create(
158 rewriter, loc, res, /*offsets=*/ArrayRef<int64_t>{w}, sizes,
159 strides));
160 }
161 } else {
162 // Extract res slice: {n, wSizeStep, f} @ [0, w, 0] for channeled
163 // convolution.
164 SmallVector<int64_t> sizes = {nSize, wSizeStep, fSize};
165 SmallVector<int64_t> strides = {1, 1, 1};
166 for (int64_t w = 0; w < wSize; w += wSizeStep) {
167 result.push_back(vector::ExtractStridedSliceOp::create(
168 rewriter, loc, res, /*offsets=*/ArrayRef<int64_t>{0, w, 0}, sizes,
169 strides));
170 }
171 }
172 return result;
173}
174
175/// Helper function to insert the computed result slices.
177 Value res, int64_t wSize, int64_t wSizeStep,
178 SmallVectorImpl<Value> &resVals,
179 bool isSingleChanneled) {
180
181 if (isSingleChanneled) {
182 // Write back res slice: {wSizeStep} @ [w] for non-channeled convolution.
183 // This does not depend on kw.
184 SmallVector<int64_t> strides = {1};
185 for (int64_t w = 0; w < wSize; w += wSizeStep) {
186 res = vector::InsertStridedSliceOp::create(
187 rewriter, loc, resVals[w], res, /*offsets=*/ArrayRef<int64_t>{w},
188 strides);
189 }
190 } else {
191 // Write back res slice: {n, wSizeStep, f} @ [0, w, 0] for channeled
192 // convolution. This does not depend on kw.
193 SmallVector<int64_t> strides = {1, 1, 1};
194 for (int64_t w = 0; w < wSize; w += wSizeStep) {
195 res = vector::InsertStridedSliceOp::create(
196 rewriter, loc, resVals[w], res,
197 /*offsets=*/ArrayRef<int64_t>{0, w, 0}, strides);
198 }
199 }
200 return res;
201}
202
203/// Contains the vectorization state and related methods used across the
204/// vectorization process of a given operation.
206 VectorizationState(RewriterBase &rewriter) : rewriterGuard(rewriter) {}
207
208 /// Initializes the vectorization state, including the computation of the
209 /// canonical vector shape for vectorization.
210 LogicalResult initState(RewriterBase &rewriter, LinalgOp linalgOp,
211 ArrayRef<int64_t> inputVectorSizes,
212 ArrayRef<bool> inputScalableVecDims,
213 bool assumeDynamicDimsMatchVecSizes = false);
214
215 /// Returns the canonical vector shape used to vectorize the iteration space.
216 ArrayRef<int64_t> getCanonicalVecShape() const { return canonicalVecShape; }
217
218 /// Returns the vector dimensions that are scalable in the canonical vector
219 /// shape.
220 ArrayRef<bool> getScalableVecDims() const { return scalableVecDims; }
221
222 /// Returns a vector type of the provided `elementType` with the canonical
223 /// vector shape and the corresponding fixed/scalable dimensions bit. If
224 /// `dimPermutation` is provided, the canonical vector dimensions are permuted
225 /// accordingly.
227 Type elementType,
228 std::optional<AffineMap> dimPermutation = std::nullopt) const {
230 SmallVector<bool> scalableDims;
231 if (dimPermutation.has_value()) {
233 applyPermutationMap<int64_t>(*dimPermutation, canonicalVecShape);
234 scalableDims =
235 applyPermutationMap<bool>(*dimPermutation, scalableVecDims);
236 } else {
237 vectorShape.append(canonicalVecShape.begin(), canonicalVecShape.end());
238 scalableDims.append(scalableVecDims.begin(), scalableVecDims.end());
239 }
240
241 return VectorType::get(vectorShape, elementType, scalableDims);
242 }
243
244 /// Masks an operation with the canonical vector mask if the operation needs
245 /// masking. Returns the masked operation or the original operation if masking
246 /// is not needed. If provided, the canonical mask for this operation is
247 /// permuted using `maybeIndexingMap`.
248 Operation *
249 maskOperation(RewriterBase &rewriter, Operation *opToMask, LinalgOp linalgOp,
250 std::optional<AffineMap> maybeIndexingMap = std::nullopt);
251
252private:
253 /// Initializes the iteration space static sizes using the Linalg op
254 /// information. This may become more complicated in the future.
255 void initIterSpaceStaticSizes(LinalgOp linalgOp) {
256 iterSpaceStaticSizes.append(linalgOp.getStaticLoopRanges());
257 }
258
259 /// Generates 'arith.constant' and 'tensor/memref.dim' operations for
260 /// all the static and dynamic dimensions of the iteration space to be
261 /// vectorized and store them in `iterSpaceValueSizes`.
262 LogicalResult precomputeIterSpaceValueSizes(RewriterBase &rewriter,
263 LinalgOp linalgOp);
264
265 /// Create or retrieve an existing mask value to mask `opToMask` in the
266 /// canonical vector iteration space. If `maybeMaskingMap` the mask is
267 /// permuted using that permutation map. If a new mask is created, it will be
268 /// cached for future users.
269 Value getOrCreateMaskFor(RewriterBase &rewriter, Operation *opToMask,
270 LinalgOp linalgOp,
271 std::optional<AffineMap> maybeMaskingMap);
272
273 /// Check whether this permutation map can be used for masking. At the
274 /// moment we only make sure that there are no broadcast dimensions, but this
275 /// might change if indexing maps evolve.
276 bool isValidMaskingMap(AffineMap maskingMap) {
277 return maskingMap.getBroadcastDims().empty();
278 }
279
280 /// Turn the input indexing map into a valid masking map.
281 ///
282 /// The input indexing map may contain "zero" results, e.g.:
283 /// (d0, d1, d2, d3) -> (d2, d1, d0, 0)
284 /// Applying such maps to canonical vector shapes like this one:
285 /// (1, 16, 16, 4)
286 /// would yield an invalid vector shape like this:
287 /// (16, 16, 1, 0)
288 /// Instead, drop the broadcasting dims that make no sense for masking perm.
289 /// maps:
290 /// (d0, d1, d2, d3) -> (d2, d1, d0)
291 /// This way, the corresponding vector/mask type will be:
292 /// vector<16x16x1xty>
293 /// rather than this invalid Vector type:
294 /// vector<16x16x1x0xty>
295 AffineMap getMaskingMapFromIndexingMap(AffineMap &indexingMap) {
296 return indexingMap.dropZeroResults();
297 }
298
299 // Holds the compile-time static sizes of the iteration space to vectorize.
300 // Dynamic dimensions are represented using ShapedType::kDynamic.
301 SmallVector<int64_t> iterSpaceStaticSizes;
302
303 /// Holds the value sizes of the iteration space to vectorize. Static
304 /// dimensions are represented by 'arith.constant' and dynamic
305 /// dimensions by 'tensor/memref.dim'.
306 SmallVector<Value> iterSpaceValueSizes;
307
308 /// Holds the canonical vector shape used to vectorize the iteration space.
309 SmallVector<int64_t> canonicalVecShape;
310
311 /// Holds the vector dimensions that are scalable in the canonical vector
312 /// shape.
313 SmallVector<bool> scalableVecDims;
314
315 /// Holds the active masks for permutations of the canonical vector iteration
316 /// space.
317 DenseMap<AffineMap, Value> activeMaskCache;
318
319 /// Global vectorization guard for the incoming rewriter. It's initialized
320 /// when the vectorization state is initialized.
321 OpBuilder::InsertionGuard rewriterGuard;
322
323 /// Do all dynamic dims match the corresponding vector sizes?
324 ///
325 /// When a dynamic tensor/memref dimension matches the corresponding vector
326 /// dimension, masking can be safely skipped, despite the presence of dynamic
327 /// shapes. Use this flag with care and only for cases where you are
328 /// confident the assumption holds.
329 bool assumeDynamicDimsMatchVecSizes = false;
330};
331
332LogicalResult
333VectorizationState::precomputeIterSpaceValueSizes(RewriterBase &rewriter,
334 LinalgOp linalgOp) {
335 // TODO: Support 0-d vectors.
336 for (int vecDim = 0, end = canonicalVecShape.size(); vecDim < end; ++vecDim) {
337 if (ShapedType::isStatic(iterSpaceStaticSizes[vecDim])) {
338 // Create constant index op for static dimensions.
339 iterSpaceValueSizes.push_back(arith::ConstantIndexOp::create(
340 rewriter, linalgOp.getLoc(), iterSpaceStaticSizes[vecDim]));
341 continue;
342 }
343
344 // Find an operand defined on this dimension of the iteration space to
345 // extract the runtime dimension size.
346 Value operand;
347 unsigned operandDimPos;
348 if (failed(linalgOp.mapIterationSpaceDimToOperandDim(vecDim, operand,
349 operandDimPos)))
350 return failure();
351
352 Value dynamicDim =
353 linalgOp.hasPureTensorSemantics()
354 ? (Value)tensor::DimOp::create(rewriter, linalgOp.getLoc(), operand,
355 operandDimPos)
356 : (Value)memref::DimOp::create(rewriter, linalgOp.getLoc(), operand,
357 operandDimPos);
358 iterSpaceValueSizes.push_back(dynamicDim);
359 }
360
361 return success();
362}
363
364/// Initializes the vectorization state, including the computation of the
365/// canonical vector shape for vectorization.
366// TODO: Move this to the constructor when we can remove the failure cases.
368 LinalgOp linalgOp,
369 ArrayRef<int64_t> inputVectorSizes,
370 ArrayRef<bool> inputScalableVecDims,
371 bool assumeDimsMatchVec) {
372 assumeDynamicDimsMatchVecSizes = assumeDimsMatchVec;
373 // Initialize the insertion point.
374 rewriter.setInsertionPoint(linalgOp);
375
376 if (!inputVectorSizes.empty()) {
377 // Get the canonical vector shape from the input vector sizes provided. This
378 // path should be taken to vectorize code with dynamic shapes and when using
379 // vector sizes greater than the iteration space sizes.
380 canonicalVecShape.append(inputVectorSizes.begin(), inputVectorSizes.end());
381 scalableVecDims.append(inputScalableVecDims.begin(),
382 inputScalableVecDims.end());
383 } else {
384 // Compute the canonical vector shape from the operation shape. If there are
385 // dynamic shapes, the operation won't be vectorized. We assume all the
386 // vector dimensions are fixed.
387 canonicalVecShape = linalgOp.getStaticLoopRanges();
388 scalableVecDims.append(linalgOp.getNumLoops(), false);
389 }
390
391 LDBG() << "Canonical vector shape: " << llvm::interleaved(canonicalVecShape);
392 LDBG() << "Scalable vector dims: " << llvm::interleaved(scalableVecDims);
393
394 if (ShapedType::isDynamicShape(canonicalVecShape))
395 return failure();
396
397 // Initialize iteration space static sizes.
398 initIterSpaceStaticSizes(linalgOp);
399
400 // Generate 'arith.constant' and 'tensor/memref.dim' operations for
401 // all the static and dynamic dimensions of the iteration space, needed to
402 // compute a mask during vectorization.
403 if (failed(precomputeIterSpaceValueSizes(rewriter, linalgOp)))
404 return failure();
405
406 return success();
407}
408
409/// Create or retrieve an existing mask value to mask `opToMask` in the
410/// canonical vector iteration space. If `maybeMaskingMap` the mask is permuted
411/// using that permutation map. If a new mask is created, it will be cached for
412/// future users.
413Value VectorizationState::getOrCreateMaskFor(
414 RewriterBase &rewriter, Operation *opToMask, LinalgOp linalgOp,
415 std::optional<AffineMap> maybeMaskingMap) {
416
417 assert((!maybeMaskingMap || isValidMaskingMap(*maybeMaskingMap)) &&
418 "Ill-formed masking map.");
419
420 // No mask is needed if the operation is not maskable.
421 auto maskableOp = dyn_cast<vector::MaskableOpInterface>(opToMask);
422 if (!maskableOp)
423 return Value();
424
425 assert(!maskableOp.isMasked() &&
426 "Masking an operation that is already masked");
427
428 // If no masking map was provided, use an identity map with the loop dims.
429 assert((!maybeMaskingMap || *maybeMaskingMap) &&
430 "Unexpected null mask permutation map");
431 AffineMap maskingMap =
432 maybeMaskingMap ? *maybeMaskingMap
434 linalgOp.getNumLoops(), rewriter.getContext());
435
436 LDBG() << "Masking map: " << maskingMap;
437
438 // Return the active mask for the masking map of this operation if it was
439 // already created.
440 auto activeMaskIt = activeMaskCache.find(maskingMap);
441 if (activeMaskIt != activeMaskCache.end()) {
442 Value mask = activeMaskIt->second;
443 LDBG() << "Reusing mask: " << mask;
444 return mask;
445 }
446
447 // Compute permuted projection of the iteration space to be masked and the
448 // corresponding mask shape. If the resulting iteration space dimensions are
449 // static and identical to the mask shape, masking is not needed for this
450 // operation.
451 // TODO: Improve this check. Only projected permutation indexing maps are
452 // supported.
453 SmallVector<int64_t> permutedStaticSizes =
454 applyPermutationMap<int64_t>(maskingMap, iterSpaceStaticSizes);
455 auto maskType = getCanonicalVecType(rewriter.getI1Type(), maskingMap);
456 auto maskShape = maskType.getShape();
457
458 LDBG() << "Mask shape: " << llvm::interleaved(maskShape);
459
460 if (permutedStaticSizes == maskShape) {
461 LDBG() << "Masking is not needed for masking map: " << maskingMap;
462 activeMaskCache[maskingMap] = Value();
463 return Value();
464 }
465
466 if (assumeDynamicDimsMatchVecSizes) {
467 // While for _dynamic_ dim sizes we can _assume_ that the corresponding
468 // vector sizes match, we still need to check the _static_ dim sizes. Only
469 // then we can be 100% sure that masking is not required.
470 if (llvm::all_of(llvm::zip(permutedStaticSizes, maskType.getShape()),
471 [](auto it) {
472 return std::get<0>(it) == ShapedType::kDynamic
473 ? true
474 : std::get<0>(it) == std::get<1>(it);
475 })) {
476 LDBG()
477 << "Dynamic + static dimensions match vector sizes, masking is not "
478 "required.";
479 activeMaskCache[maskingMap] = Value();
480 return Value();
481 }
482 }
483
484 // Permute the iteration space value sizes to compute the mask upper bounds.
485 SmallVector<Value> upperBounds =
486 applyPermutationMap(maskingMap, ArrayRef<Value>(iterSpaceValueSizes));
487 assert(!maskShape.empty() && !upperBounds.empty() &&
488 "Masked 0-d vectors are not supported yet");
489
490 // Create the mask based on the dimension values.
491 Value mask = vector::CreateMaskOp::create(rewriter, linalgOp.getLoc(),
492 maskType, upperBounds);
493 LDBG() << "Creating new mask: " << mask;
494 activeMaskCache[maskingMap] = mask;
495 return mask;
496}
497
498Operation *
500 LinalgOp linalgOp,
501 std::optional<AffineMap> maybeIndexingMap) {
502 LDBG() << "Trying to mask: " << *opToMask;
503
504 std::optional<AffineMap> maybeMaskingMap = std::nullopt;
505 if (maybeIndexingMap)
506 maybeMaskingMap = getMaskingMapFromIndexingMap(*maybeIndexingMap);
507
508 // Create or retrieve mask for this operation.
509 Value mask =
510 getOrCreateMaskFor(rewriter, opToMask, linalgOp, maybeMaskingMap);
511
512 if (!mask) {
513 LDBG() << "No mask required";
514 if (assumeDynamicDimsMatchVecSizes) {
516 .Case<vector::TransferReadOp, vector::TransferWriteOp>(
517 [&](auto xferOp) {
518 // For vector.transfer_read and vector.transfer_write, there is
519 // also the `in-bounds` attribute that has to be set explicitly
520 // to true. Otherwise, "out-of-bounds" access will be assumed
521 // and masks will be generated while lowering these.
522 LDBG() << "Assuming dynamic dimensions match vector sizes and "
523 "setting their in-bounds to true!";
524 SmallVector<bool> inBoundsMap = xferOp.getInBoundsValues();
525 ShapedType xferType = xferOp.getShapedType();
526 AffineMap permMap = xferOp.getPermutationMap();
527 // Only set the in-bounds values to true for dynamic dims.
528 // Different mechanisms will set these accordingly for the
529 // static dims.
530 for (unsigned i = 0; i < xferOp.getTransferRank(); i++) {
531 auto dimExpr = dyn_cast<AffineDimExpr>(permMap.getResult(i));
532 // Skip broadcast dimensions.
533 if (!dimExpr)
534 continue;
535 unsigned pos = dimExpr.getPosition();
536 if (xferType.isDynamicDim(pos))
537 inBoundsMap[i] = true;
538 }
539 rewriter.modifyOpInPlace(xferOp, [&]() {
540 xferOp.setInBoundsAttr(
541 rewriter.getBoolArrayAttr(inBoundsMap));
542 });
543 })
544 .Default([](Operation *op) {
545 // No-op if the operation is not an xfer read or write.
546 });
547 }
548 return opToMask;
549 }
550
551 // Wrap the operation with a new `vector.mask` and update D-U chain.
552 assert(opToMask && "Expected a valid operation to mask");
553 auto maskOp = cast<vector::MaskOp>(
554 mlir::vector::maskOperation(rewriter, opToMask, mask));
555 Operation *maskOpTerminator = &maskOp.getMaskRegion().front().back();
556
557 for (auto [resIdx, resVal] : llvm::enumerate(opToMask->getResults()))
558 rewriter.replaceAllUsesExcept(resVal, maskOp.getResult(resIdx),
559 maskOpTerminator);
560
561 LDBG() << "Masked operation: " << *maskOp;
562 return maskOp;
563}
564
565/// Given an indexing `map` coming from a LinalgOp indexing, restricted to a
566/// projectedPermutation, compress the unused dimensions to serve as a
567/// permutation_map for a vector transfer operation.
568/// For example, given a linalg op such as:
569///
570/// ```
571/// %0 = linalg.generic {
572/// indexing_maps = affine_map<(d0, d1, d2, d3, d4) -> (d4, d0, d2)>,
573/// indexing_maps = affine_map<(d0, d1, d2, d3, d4) -> (d1, d3)>
574/// }
575/// ins(%0 : tensor<2x3x4xf32>)
576/// outs(%1 : tensor<5x6xf32>)
577/// ```
578///
579/// the iteration domain size of the linalg op is 3x5x4x6x2. The first affine
580/// map is reindexed to `affine_map<(d0, d1, d2) -> (d2, d0, d1)>`, the second
581/// affine map is reindexed to `affine_map<(d0, d1) -> (d0, d1)>`.
583 assert(map.isProjectedPermutation(/*allowZeroInResults=*/true) &&
584 "expected projected permutation");
585 auto res = compressUnusedDims(map);
586 assert(res.getNumDims() ==
587 (res.getNumResults() - res.getNumOfZeroResults()) &&
588 "expected reindexed map with same number of dims and results");
589 return res;
590}
591
592/// Helper enum to represent conv1d input traversal order.
593enum class Conv1DOpOrder {
594 W, // Corresponds to non-channeled 1D convolution operation.
595 Ncw, // Corresponds to operation that traverses the input in (n, c, w) order.
596 Nwc // Corresponds to operation that traverses the input in (n, w, c) order.
597};
598
599/// Helper data structure to represent the result of vectorization for a single
600/// operation. In certain specific cases, like terminators, we do not want to
601/// propagate.
603 /// Op failed to vectorize.
605 /// Op vectorized and custom function took care of replacement logic
607 /// Op vectorized into a new Op whose results will replace original Op's
608 /// results.
610 // TODO: support values if Op vectorized to Many-Ops whose results we need to
611 // aggregate for replacement.
612};
613/// VectorizationHookResult contains the vectorized op returned from a
614/// CustomVectorizationHook. This is an internal implementation detail of
615/// linalg vectorization, not to be confused with VectorizationResult.
617 /// Return status from vectorizing the current op.
619 /// New vectorized operation to replace the current op.
620 /// Replacement behavior is specified by `status`.
622};
623
624std::optional<vector::CombiningKind>
626 using ::mlir::vector::CombiningKind;
627
628 if (!combinerOp)
629 return std::nullopt;
631 .Case<arith::AddIOp, arith::AddFOp>(
632 [&](auto op) { return CombiningKind::ADD; })
633 .Case([&](arith::AndIOp op) { return CombiningKind::AND; })
634 .Case([&](arith::MaxSIOp op) { return CombiningKind::MAXSI; })
635 .Case([&](arith::MaxUIOp op) { return CombiningKind::MAXUI; })
636 .Case([&](arith::MaximumFOp op) { return CombiningKind::MAXIMUMF; })
637 .Case([&](arith::MaxNumFOp op) { return CombiningKind::MAXNUMF; })
638 .Case([&](arith::MaximumNumFOp op) { return CombiningKind::MAXIMUMNUMF; })
639 .Case([&](arith::MinSIOp op) { return CombiningKind::MINSI; })
640 .Case([&](arith::MinUIOp op) { return CombiningKind::MINUI; })
641 .Case([&](arith::MinimumFOp op) { return CombiningKind::MINIMUMF; })
642 .Case([&](arith::MinNumFOp op) { return CombiningKind::MINNUMF; })
643 .Case([&](arith::MinimumNumFOp op) { return CombiningKind::MINIMUMNUMF; })
644 .Case<arith::MulIOp, arith::MulFOp>(
645 [&](auto op) { return CombiningKind::MUL; })
646 .Case([&](arith::OrIOp op) { return CombiningKind::OR; })
647 .Case([&](arith::XOrIOp op) { return CombiningKind::XOR; })
648 .Default(std::nullopt);
649}
650
651/// Check whether `outputOperand` is a reduction with a single combiner
652/// operation. Return the combiner operation of the reduction. Return
653/// nullptr otherwise. Multiple reduction operations would impose an
654/// ordering between reduction dimensions and is currently unsupported in
655/// Linalg. This limitation is motivated by the fact that e.g. min(max(X)) !=
656/// max(min(X))
657// TODO: use in LinalgOp verification, there is a circular dependency atm.
658static Operation *matchLinalgReduction(OpOperand *outputOperand) {
659 auto linalgOp = cast<LinalgOp>(outputOperand->getOwner());
660 unsigned outputPos =
661 outputOperand->getOperandNumber() - linalgOp.getNumDpsInputs();
662 // Only single combiner operations are supported for now.
663 SmallVector<Operation *, 4> combinerOps;
664 if (!matchReduction(linalgOp.getRegionOutputArgs(), outputPos, combinerOps) ||
665 combinerOps.size() != 1)
666 return nullptr;
667
668 // Return the combiner operation.
669 return combinerOps[0];
670}
671
672/// Broadcast `value` to a vector of `shape` if possible. Return value
673/// otherwise.
674static Value broadcastIfNeeded(OpBuilder &b, Value value, Type dstType) {
675 auto dstVecType = dyn_cast<VectorType>(dstType);
676 // If no shape to broadcast to, just return `value`.
677 if (dstVecType.getRank() == 0)
678 return value;
679 if (vector::isBroadcastableTo(value.getType(), dstVecType) !=
681 return value;
682 Location loc = b.getInsertionPoint()->getLoc();
683 return b.createOrFold<vector::BroadcastOp>(loc, dstVecType, value);
684}
685
686/// Create MultiDimReductionOp to compute the reduction for `reductionOp`. This
687/// assumes that `reductionOp` has two operands and one of them is the reduction
688/// initial value.buildMultiDimReduce
689// Note: this is a true builder that notifies the OpBuilder listener.
690// TODO: Consider moving as a static helper on the ReduceOp.
692 Value valueToReduce, Value acc,
693 ArrayRef<bool> dimsToMask) {
694 auto maybeKind = getCombinerOpKind(reduceOp);
695 assert(maybeKind && "Failed precondition: could not get reduction kind");
696 return vector::MultiDimReductionOp::create(
697 b, reduceOp->getLoc(), valueToReduce, acc, dimsToMask, *maybeKind);
698}
699
700static SmallVector<bool> getDimsToReduce(LinalgOp linalgOp) {
701 return llvm::map_to_vector(linalgOp.getIteratorTypesArray(),
703}
704
705/// Check if `op` is a linalg.reduce or a linalg.generic that has at least one
706/// reduction iterator.
707static bool hasReductionIterator(LinalgOp &op) {
708 return isa<linalg::ReduceOp>(op) ||
709 (isa<linalg::GenericOp>(op) &&
710 llvm::any_of(op.getIteratorTypesArray(), isReductionIterator));
711}
712
713/// Build a vector.transfer_write of `value` into `outputOperand` at indices set
714/// to all `0`; where `outputOperand` is an output operand of the LinalgOp
715/// currently being vectorized. If `dest` has null rank, build an memref.store.
716/// Return the produced value or null if no value is produced.
717// Note: this is a true builder that notifies the OpBuilder listener.
718// TODO: Consider moving as a static helper on the ReduceOp.
719static Value buildVectorWrite(RewriterBase &rewriter, Value value,
720 OpOperand *outputOperand,
721 VectorizationState &state) {
722 Location loc = value.getLoc();
723 auto linalgOp = cast<LinalgOp>(outputOperand->getOwner());
724 AffineMap opOperandMap = linalgOp.getMatchingIndexingMap(outputOperand);
725
726 // Compute the vector type of the value to store. This type should be an
727 // identity or projection of the canonical vector type without any permutation
728 // applied, given that any permutation in a transfer write happens as part of
729 // the write itself.
731 opOperandMap.getContext(), opOperandMap.getNumInputs(),
732 [&](AffineDimExpr dimExpr) -> bool {
733 return llvm::is_contained(opOperandMap.getResults(), dimExpr);
734 });
735 auto vectorType = state.getCanonicalVecType(
736 getElementTypeOrSelf(outputOperand->get().getType()), vectorTypeMap);
737
738 SmallVector<Value> indices(linalgOp.getRank(outputOperand),
739 arith::ConstantIndexOp::create(rewriter, loc, 0));
740
741 Operation *write;
742 if (vectorType.getRank() > 0) {
743 AffineMap writeMap = inversePermutation(reindexIndexingMap(opOperandMap));
744 value = broadcastIfNeeded(rewriter, value, vectorType);
745 assert(value.getType() == vectorType && "Incorrect type");
746 write = vector::TransferWriteOp::create(
747 rewriter, loc, value, outputOperand->get(), indices, writeMap);
748 } else {
749 // 0-d case is still special: do not invert the reindexing writeMap.
750 if (!isa<VectorType>(value.getType()))
751 value = vector::BroadcastOp::create(rewriter, loc, vectorType, value);
752 assert(value.getType() == vectorType && "Incorrect type");
753 write = vector::TransferWriteOp::create(rewriter, loc, value,
754 outputOperand->get(), indices);
755 }
756
757 write = state.maskOperation(rewriter, write, linalgOp, opOperandMap);
758
759 // If masked, set in-bounds to true. Masking guarantees that the access will
760 // be in-bounds.
761 if (auto maskOp = dyn_cast<vector::MaskingOpInterface>(write)) {
762 auto maskedWriteOp = cast<vector::TransferWriteOp>(maskOp.getMaskableOp());
763 SmallVector<bool> inBounds(maskedWriteOp.getVectorType().getRank(), true);
764 maskedWriteOp.setInBoundsAttr(rewriter.getBoolArrayAttr(inBounds));
765 }
766
767 LDBG() << "vectorized op: " << *write;
768 if (!write->getResults().empty())
769 return write->getResult(0);
770 return Value();
771}
772
773// Custom vectorization precondition function type. This is intented to be used
774// with CustomVectorizationHook. Returns success if the corresponding custom
775// hook can vectorize the op.
777 std::function<LogicalResult(Operation *, bool)>;
778
779// Custom vectorization function type. Produce a vector form of Operation*
780// assuming all its vectorized operands are already in the IRMapping.
781// Return nullptr if the Operation cannot be vectorized.
783 std::function<VectorizationHookResult(Operation *, const IRMapping &)>;
784
785/// Helper function to vectorize the terminator of a `linalgOp`. New result
786/// vector values are appended to `newResults`. Return
787/// VectorizationHookStatus::NoReplace to signal the vectorization algorithm
788/// that it should not try to map produced operations and instead return the
789/// results using the `newResults` vector making them available to the
790/// vectorization algorithm for RAUW. This function is meant to be used as a
791/// CustomVectorizationHook.
794 const IRMapping &bvm, VectorizationState &state,
795 LinalgOp linalgOp, SmallVectorImpl<Value> &newResults) {
796 auto yieldOp = dyn_cast<linalg::YieldOp>(op);
797 if (!yieldOp)
799 for (const auto &output : llvm::enumerate(yieldOp.getValues())) {
800 // TODO: Scan for an opportunity for reuse.
801 // TODO: use a map.
802 Value vectorValue = bvm.lookup(output.value());
803 Value newResult =
804 buildVectorWrite(rewriter, vectorValue,
805 linalgOp.getDpsInitOperand(output.index()), state);
806 if (newResult)
807 newResults.push_back(newResult);
808 }
809
811}
812
813/// Helper function to vectorize the index operations of a `linalgOp`. Return
814/// VectorizationHookStatus::NewOp to signal the vectorization algorithm that it
815/// should map the produced operations. This function is meant to be used as a
816/// CustomVectorizationHook.
818 VectorizationState &state,
819 Operation *op,
820 LinalgOp linalgOp) {
821 IndexOp indexOp = dyn_cast<linalg::IndexOp>(op);
822 if (!indexOp)
824 auto loc = indexOp.getLoc();
825 // Compute the static loop sizes of the index op.
826 ArrayRef<int64_t> targetShape = state.getCanonicalVecShape();
827 auto dim = indexOp.getDim();
828 // Compute a one-dimensional index vector for the index op dimension.
829 auto indexVectorType =
830 VectorType::get({targetShape[dim]}, rewriter.getIndexType(),
831 state.getScalableVecDims()[dim]);
832 auto indexSteps = vector::StepOp::create(rewriter, loc, indexVectorType);
833 // Return the one-dimensional index vector if it lives in the trailing
834 // dimension of the iteration space since the vectorization algorithm in this
835 // case can handle the broadcast.
836 if (dim == targetShape.size() - 1)
838 // Otherwise permute the targetShape to move the index dimension last,
839 // broadcast the one-dimensional index vector to the permuted shape, and
840 // finally transpose the broadcasted index vector to undo the permutation.
841 auto permPattern =
842 llvm::to_vector(llvm::seq<unsigned>(0, targetShape.size()));
843 std::swap(permPattern[dim], permPattern.back());
844 auto permMap =
845 AffineMap::getPermutationMap(permPattern, linalgOp.getContext());
846
847 auto broadCastOp = vector::BroadcastOp::create(
848 rewriter, loc,
849 state.getCanonicalVecType(rewriter.getIndexType(), permMap), indexSteps);
850 SmallVector<int64_t> transposition =
851 llvm::to_vector<16>(llvm::seq<int64_t>(0, linalgOp.getNumLoops()));
852 std::swap(transposition.back(), transposition[dim]);
853 auto transposeOp =
854 vector::TransposeOp::create(rewriter, loc, broadCastOp, transposition);
856}
857
858/// Helper function to check if the tensor.extract can be vectorized by the
859/// custom hook vectorizeTensorExtract.
860static LogicalResult
862 tensor::ExtractOp extractOp = dyn_cast<tensor::ExtractOp>(op);
863 if (!extractOp)
864 return failure();
865
866 if (extractOp.getIndices().size() != 1 && !vectorizeNDExtract)
867 return failure();
868
869 // Check the index type, but only for non 0-d tensors (for which we do need
870 // access indices).
871 if (not extractOp.getIndices().empty()) {
872 if (!VectorType::isValidElementType(extractOp.getIndices()[0].getType()))
873 return failure();
874 }
875
876 if (!llvm::all_of(extractOp->getResultTypes(),
877 VectorType::isValidElementType)) {
878 return failure();
879 }
880
881 return success();
882}
883
884/// Calculates the offsets (`$index_vec`) for `vector.gather` operations
885/// generated from `tensor.extract`. The offset is calculated as follows
886/// (example using scalar values):
887///
888/// offset = extractOp.indices[0]
889/// for (i = 1; i < numIndices; i++)
890/// offset = extractOp.dimSize[i] * offset + extractOp.indices[i];
891///
892/// For tensor<45 x 80 x 15 x f32> and index [1, 2, 3], this leads to:
893/// offset = ( ( 1 ) * 80 + 2 ) * 15 + 3
895 VectorizationState &state,
896 tensor::ExtractOp extractOp,
897 const IRMapping &bvm) {
898 // The vector of indices for GatherOp should be shaped as the output vector.
899 auto indexVecType = state.getCanonicalVecType(rewriter.getIndexType());
900 auto loc = extractOp.getLoc();
901
902 Value offset = broadcastIfNeeded(
903 rewriter, bvm.lookup(extractOp.getIndices()[0]), indexVecType);
904
905 const size_t numIndices = extractOp.getIndices().size();
906 for (size_t i = 1; i < numIndices; i++) {
907 Value dimIdx = arith::ConstantIndexOp::create(rewriter, loc, i);
908
909 auto dimSize = broadcastIfNeeded(
910 rewriter,
911 tensor::DimOp::create(rewriter, loc, extractOp.getTensor(), dimIdx),
912 indexVecType);
913
914 offset = arith::MulIOp::create(rewriter, loc, offset, dimSize);
915
916 auto extractOpIndex = broadcastIfNeeded(
917 rewriter, bvm.lookup(extractOp.getIndices()[i]), indexVecType);
918
919 offset = arith::AddIOp::create(rewriter, loc, extractOpIndex, offset);
920 }
921
922 return offset;
923}
924
926
927/// Find the index of the trailing non-unit dim in linalgOp. This hook is used
928/// when checking whether `tensor.extract` Op (within a `linalg.generic` Op)
929/// represents a contiguous load operation.
930///
931/// Note that when calling this hook, it is assumed that the output vector is
932/// effectively 1D. Other cases (i.e. reading n-D vectors) should've been
933/// labelled as a gather load before entering this method.
934///
935/// Following on from the above, it is assumed that:
936/// * for statically shaped loops, when no masks are used, only one dim is !=
937/// 1 (that's what the shape of the output vector is based on).
938/// * for dynamically shaped loops, there might be more non-unit dims
939/// as the output vector type is user-specified.
940///
941/// TODO: Statically shaped loops + vector masking
942static uint64_t getTrailingNonUnitLoopDimIdx(LinalgOp linalgOp) {
943 SmallVector<int64_t> loopRanges = linalgOp.getStaticLoopRanges();
944 assert(
945 (linalgOp.hasDynamicShape() ||
946 llvm::count_if(loopRanges, [](int64_t dim) { return dim != 1; }) == 1) &&
947 "For statically shaped Linalg Ops, only one "
948 "non-unit loop dim is expected");
949 assert(!loopRanges.empty() && "Empty loops, nothing to analyse.");
950
951 size_t idx = loopRanges.size() - 1;
952 for (; idx != 0; idx--)
953 if (loopRanges[idx] != 1)
954 break;
955
956 return idx;
957}
958
959/// Checks whether `val` can be used for calculating a loop invariant index.
960static bool isLoopInvariantIdx(LinalgOp &linalgOp, Value &val,
961 VectorType resType) {
962
963 assert(((llvm::count_if(resType.getShape(),
964 [](int64_t dimSize) { return dimSize > 1; }) == 1)) &&
965 "n-D vectors are not yet supported");
966
967 auto *block = linalgOp.getBlock();
968
969 // A shared DAG has exponentially many paths; a revisit adds nothing.
971 SmallVector<Value> worklist{val};
972
973 while (!worklist.empty()) {
974 Value v = worklist.pop_back_val();
975
976 // Blocks outside _this_ linalg.generic are effectively loop invariant.
977 // However, analysing block arguments for _this_ linalg.generic Op is a bit
978 // tricky. Just bail out in the latter case.
979 // TODO: We could try analysing the corresponding affine map here.
980 if (isa<BlockArgument>(v)) {
981 if (llvm::is_contained(block->getArguments(), v))
982 return false;
983 continue;
984 }
985
986 Operation *defOp = v.getDefiningOp();
987 assert(defOp && "This is neither a block argument nor an operation result");
988
989 // IndexOp is loop invariant as long as its result remains constant across
990 // iterations. Note that for dynamic shapes, the corresponding dim will also
991 // be conservatively treated as != 1.
992 if (auto indexOp = dyn_cast<linalg::IndexOp>(defOp)) {
993 if (linalgOp.getStaticLoopRanges()[indexOp.getDim()] != 1)
994 return false;
995 continue;
996 }
997
998 auto *ancestor = block->findAncestorOpInBlock(*defOp);
999
1000 // Values define outside `linalgOp` are loop invariant.
1001 if (!ancestor)
1002 continue;
1003
1004 // Values defined inside `linalgOp`, which are constant, are loop invariant.
1005 if (isa<arith::ConstantOp>(ancestor))
1006 continue;
1007
1008 if (visited.insert(ancestor).second)
1009 llvm::append_range(worklist, ancestor->getOperands());
1010 }
1011
1012 return true;
1013}
1014
1015/// Check whether `val` could be used for calculating the trailing index for a
1016/// contiguous load operation.
1017///
1018/// There are currently 3 types of values that are allowed here:
1019/// 1. loop-invariant values,
1020/// 2. values that increment by 1 with every loop iteration,
1021/// 3. results of basic arithmetic operations (linear and continuous)
1022/// involving 1., 2. and 3.
1023/// This method returns True if indeed only such values are used in calculating
1024/// `val.`
1025///
1026/// Additionally, the trailing index for a contiguous load operation should
1027/// increment by 1 with every loop iteration, i.e. be based on:
1028/// * `linalg.index <dim>` ,
1029/// where <dim> is the trailing non-unit dim of the iteration space (this way,
1030/// `linalg.index <dim>` increments by 1 with every loop iteration).
1031/// `foundIndexOp` is updated to `true` when such Op is found.
1032static bool isContiguousLoadIdx(LinalgOp &linalgOp, Value &val,
1033 bool &foundIndexOp, VectorType resType) {
1034
1035 assert(((llvm::count_if(resType.getShape(),
1036 [](int64_t dimSize) { return dimSize > 1; }) == 1)) &&
1037 "n-D vectors are not yet supported");
1038
1039 // Blocks outside _this_ linalg.generic are effectively loop invariant.
1040 // However, analysing block arguments for _this_ linalg.generic Op is a bit
1041 // tricky. Just bail out in the latter case.
1042 // TODO: We could try analysing the corresponding affine map here.
1043 auto *block = linalgOp.getBlock();
1044 if (isa<BlockArgument>(val))
1045 return !llvm::is_contained(block->getArguments(), val);
1046
1047 Operation *defOp = val.getDefiningOp();
1048 assert(defOp && "This is neither a block argument nor an operation result");
1049
1050 if (auto indexOp = dyn_cast<linalg::IndexOp>(defOp)) {
1051 auto loopDimThatIncrementsByOne = getTrailingNonUnitLoopDimIdx(linalgOp);
1052
1053 foundIndexOp = (indexOp.getDim() == loopDimThatIncrementsByOne);
1054 return true;
1055 }
1056
1057 auto *ancestor = block->findAncestorOpInBlock(*defOp);
1058
1059 if (!ancestor)
1060 return false;
1061
1062 // Conservatively reject Ops that could lead to indices with stride other
1063 // than 1.
1064 if (!isa<arith::AddIOp, arith::ConstantOp, linalg::IndexOp>(ancestor))
1065 return false;
1066
1067 bool result = false;
1068 for (auto op : ancestor->getOperands())
1069 result |= isContiguousLoadIdx(linalgOp, op, foundIndexOp, resType);
1070
1071 return result;
1072}
1073
1074/// Infer the memory access pattern for the input ExtractOp
1075///
1076/// Based on the ExtratOp result shape and the access indices, decides whether
1077/// this Op corresponds to a contiguous load (including a broadcast of a scalar)
1078/// or a gather load. When analysing the ExtractOp indices (to identify
1079/// contiguous laods), this method looks for "loop" invariant indices (e.g.
1080/// block arguments) and indices that change linearly (e.g. via `linalg.index`
1081/// Op).
1082///
1083/// Note that it is always safe to use gather load operations for contiguous
1084/// loads (albeit slow), but not vice-versa. When in doubt, bail out and assume
1085/// that `extractOp` is a gather load.
1087getTensorExtractMemoryAccessPattern(tensor::ExtractOp extractOp,
1088 LinalgOp &linalgOp, VectorType resType) {
1089
1090 auto inputShape = cast<ShapedType>(extractOp.getTensor().getType());
1091
1092 // 0. Is this a 0-D vector? If yes then this is a scalar broadcast.
1093 if (inputShape.getShape().empty())
1095
1096 // 0a. Is the result a 0-D vector? If yes, there are no iteration dimensions
1097 // so the tensor.extract is a single scalar load regardless of the index.
1098 if (resType.getRank() == 0)
1100
1101 // True for vectors that are effectively 1D, e.g. `vector<1x4x1xi32>`, false
1102 // otherwise.
1103 bool isOutput1DVector =
1104 (llvm::count_if(resType.getShape(),
1105 [](int64_t dimSize) { return dimSize > 1; }) == 1);
1106 // 1. Assume that it's a gather load when reading non-1D vector.
1107 if (!isOutput1DVector)
1109
1110 bool leadingIdxsLoopInvariant = true;
1111
1112 // 2. Analyze the leading indices of `extractOp`.
1113 // Look at the way each index is calculated and decide whether it is suitable
1114 // for a contiguous load, i.e. whether it's loop invariant. If not, it's a
1115 // gather load.
1116 auto indices = extractOp.getIndices();
1117 auto leadIndices = indices.drop_back(1);
1118
1119 for (auto [i, indexVal] : llvm::enumerate(leadIndices)) {
1120 if (inputShape.getShape()[i] == 1)
1121 continue;
1122
1123 leadingIdxsLoopInvariant &= isLoopInvariantIdx(linalgOp, indexVal, resType);
1124 }
1125
1126 if (!leadingIdxsLoopInvariant) {
1127 LDBG() << "Found gather load: " << extractOp;
1129 }
1130
1131 // 3. Analyze the trailing index for `extractOp`.
1132 // At this point we know that the leading indices are loop invariant. This
1133 // means that is potentially a scalar or a contiguous load. We can decide
1134 // based on the trailing idx.
1135 auto extractOpTrailingIdx = indices.back();
1136
1137 // 3a. Scalar broadcast load
1138 // If the trailing index is loop invariant then this is a scalar load.
1139 if (leadingIdxsLoopInvariant &&
1140 isLoopInvariantIdx(linalgOp, extractOpTrailingIdx, resType)) {
1141 LDBG() << "Found scalar broadcast load: " << extractOp;
1142
1144 }
1145
1146 // 3b. Contiguous loads
1147 // The trailing `extractOp` index should increment with every loop iteration.
1148 // This effectively means that it must be based on the trailing loop index.
1149 // This is what the following bool captures.
1150 bool foundIndexOp = false;
1151 bool isContiguousLoad = isContiguousLoadIdx(linalgOp, extractOpTrailingIdx,
1152 foundIndexOp, resType);
1153 // TODO: Support generating contiguous loads for column vectors - that will
1154 // require adding a permutation map to tranfer_read Ops.
1155 bool isRowVector = resType.getShape().back() != 1;
1156 isContiguousLoad &= (foundIndexOp && isRowVector);
1157
1158 if (isContiguousLoad) {
1159 LDBG() << "Found contigous load: " << extractOp;
1161 }
1162
1163 // 4. Fallback case - gather load.
1164 LDBG() << "Found gather load: " << extractOp;
1166}
1167
1168/// Helper function to vectorize the tensor.extract operations. Returns
1169/// VectorizationHookStatus::NewOp to signal the vectorization algorithm that it
1170/// should map the produced operations. This function is meant to be used as a
1171/// CustomVectorizationHook.
1172static VectorizationHookResult
1173vectorizeTensorExtract(RewriterBase &rewriter, VectorizationState &state,
1174 Operation *op, LinalgOp linalgOp, const IRMapping &bvm) {
1175 tensor::ExtractOp extractOp = dyn_cast<tensor::ExtractOp>(op);
1176 if (!extractOp)
1178 auto loc = extractOp.getLoc();
1179
1180 // Compute the static loop sizes of the extract op.
1181 auto resultType = state.getCanonicalVecType(extractOp.getResult().getType());
1182 auto maskConstantOp = arith::ConstantOp::create(
1183 rewriter, loc,
1184 DenseIntElementsAttr::get(state.getCanonicalVecType(rewriter.getI1Type()),
1185 /*value=*/true));
1186 auto passThruConstantOp = arith::ConstantOp::create(
1187 rewriter, loc, rewriter.getZeroAttr(resultType));
1188
1189 // Base indices are currently set to 0. We will need to re-visit if more
1190 // generic scenarios are to be supported.
1191 SmallVector<Value> baseIndices(
1192 extractOp.getIndices().size(),
1193 arith::ConstantIndexOp::create(rewriter, loc, 0));
1194
1195 VectorMemoryAccessKind memAccessKind =
1196 getTensorExtractMemoryAccessPattern(extractOp, linalgOp, resultType);
1197
1198 // 1. Handle gather access
1199 if (memAccessKind == VectorMemoryAccessKind::Gather) {
1200 Value offset = calculateGatherOffset(rewriter, state, extractOp, bvm);
1201
1202 // Generate the gather load
1203 Operation *gatherOp = vector::GatherOp::create(
1204 rewriter, loc, resultType, extractOp.getTensor(), baseIndices, offset,
1205 maskConstantOp, passThruConstantOp);
1206 gatherOp = state.maskOperation(rewriter, gatherOp, linalgOp);
1207
1208 LDBG() << "Vectorised as gather load: " << extractOp;
1210 }
1211
1212 // 2. Handle:
1213 // a. scalar loads + broadcast,
1214 // b. contiguous loads.
1215 // Both cases use vector.transfer_read.
1216
1217 // Collect indices for `vector.transfer_read`. At this point, the indices will
1218 // either be scalars or would have been broadcast to vectors matching the
1219 // result type. For indices that are vectors, there are two options:
1220 // * for non-trailing indices, all elements are identical (contiguous
1221 // loads are identified by looking for non-trailing indices that are
1222 // invariant with respect to the corresponding linalg.generic), or
1223 // * for trailing indices, the index vector will contain values with stride
1224 // one, but for `vector.transfer_read` only the first (i.e. 0th) index is
1225 // needed.
1226 // This means that
1227 // * for scalar indices - just re-use it,
1228 // * for vector indices (e.g. `vector<1x1x4xindex>`) - extract the bottom
1229 // (0th) element and use that.
1230 SmallVector<Value> transferReadIdxs;
1231 for (size_t i = 0; i < extractOp.getIndices().size(); i++) {
1232 Value idx = bvm.lookup(extractOp.getIndices()[i]);
1233 if (idx.getType().isIndex()) {
1234 transferReadIdxs.push_back(idx);
1235 continue;
1236 }
1237
1238 auto indexAs1dVector = vector::ShapeCastOp::create(
1239 rewriter, loc,
1240 VectorType::get(resultType.getShape().back(), rewriter.getIndexType(),
1241 resultType.getScalableDims().back()),
1242 idx);
1243 transferReadIdxs.push_back(
1244 vector::ExtractOp::create(rewriter, loc, indexAs1dVector, 0));
1245 }
1246
1247 // `tensor.extract_element` is always in-bounds, hence the following holds.
1248 auto dstRank = resultType.getRank();
1249 auto srcRank = extractOp.getTensor().getType().getRank();
1250 SmallVector<bool> inBounds(dstRank, true);
1251
1252 // 2a. Handle scalar broadcast access.
1253 if (memAccessKind == VectorMemoryAccessKind::ScalarBroadcast) {
1254 MLIRContext *ctx = rewriter.getContext();
1255 SmallVector<AffineExpr> exprs(dstRank, getAffineConstantExpr(0, ctx));
1256 auto permutationMap = AffineMap::get(srcRank, 0, exprs, ctx);
1257
1258 auto transferReadOp = vector::TransferReadOp::create(
1259 rewriter, loc, resultType, extractOp.getTensor(), transferReadIdxs,
1260 /*padding=*/std::nullopt, permutationMap, inBounds);
1261
1262 Operation *readOrMaskedReadOp = transferReadOp;
1263 if (dstRank > 0) {
1264 // Mask this broadcasting xfer_read here rather than relying on the
1265 // generic path (the generic path assumes identity masking map, which
1266 // wouldn't be valid here).
1267 SmallVector<int64_t> readMaskShape = {1};
1268 auto readMaskType = VectorType::get(readMaskShape, rewriter.getI1Type());
1269 auto allTrue = vector::ConstantMaskOp::create(
1270 rewriter, loc, readMaskType, vector::ConstantMaskKind::AllTrue);
1271 readOrMaskedReadOp =
1272 mlir::vector::maskOperation(rewriter, transferReadOp, allTrue);
1273 }
1274
1275 LDBG() << "Vectorised as scalar broadcast load: " << extractOp;
1277 readOrMaskedReadOp};
1278 }
1279
1280 // 2b. Handle contiguous access.
1281 auto permutationMap = AffineMap::getMinorIdentityMap(
1282 srcRank, std::min(dstRank, srcRank), rewriter.getContext());
1283
1284 int32_t rankDiff = dstRank - srcRank;
1285 // When dstRank > srcRank, broadcast the source tensor to the unitary leading
1286 // dims so that the ranks match. This is done by extending the map with 0s.
1287 // For example, for dstRank = 3, srcRank = 2, the following map created
1288 // above:
1289 // (d0, d1) --> (d0, d1)
1290 // is extended as:
1291 // (d0, d1) --> (0, d0, d1)
1292 while (rankDiff > 0) {
1293 permutationMap = permutationMap.insertResult(
1294 mlir::getAffineConstantExpr(0, rewriter.getContext()), 0);
1295 rankDiff--;
1296 }
1297
1298 auto transferReadOp = vector::TransferReadOp::create(
1299 rewriter, loc, resultType, extractOp.getTensor(), transferReadIdxs,
1300 /*padding=*/std::nullopt, permutationMap, inBounds);
1301
1302 // Mask this contiguous xfer_read here rather than relying on the generic
1303 // path (the generic path assumes an identity masking map over all the loop
1304 // dims, which wouldn't be valid here). A contiguous load only reads the
1305 // trailing `min(dstRank, srcRank)` dims of the iteration space - the leading
1306 // dims are broadcast via `permutationMap` above - so its inferred mask is
1307 // rank-reduced. Build a masking map that projects the iteration space onto
1308 // exactly those trailing dims so the created mask matches the xfer_read.
1309 int64_t numReadDims = std::min(dstRank, srcRank);
1310 auto maskingMap = AffineMap::getMinorIdentityMap(
1311 linalgOp.getNumLoops(), numReadDims, rewriter.getContext());
1312 Operation *maskedReadOp =
1313 state.maskOperation(rewriter, transferReadOp, linalgOp, maskingMap);
1314
1315 LDBG() << "Vectorised as contiguous load: " << extractOp;
1317}
1318
1319/// Emit reduction operations if the shapes of the value to reduce is different
1320/// that the result shape.
1321// Note: this is a true builder that notifies the OpBuilder listener.
1322// TODO: Consider moving as a static helper on the ReduceOp.
1323static Operation *reduceIfNeeded(OpBuilder &b, LinalgOp linalgOp, Operation *op,
1324 Value reduceValue, Value initialValue,
1325 const IRMapping &bvm) {
1326 Value reduceVec = bvm.lookup(reduceValue);
1327 Value outputVec = bvm.lookup(initialValue);
1328 auto reduceType = dyn_cast<VectorType>(reduceVec.getType());
1329 auto outputType = dyn_cast<VectorType>(outputVec.getType());
1330 // Reduce only if needed as the value may already have been reduce for
1331 // contraction vectorization.
1332 if (!reduceType ||
1333 (outputType && reduceType.getShape() == outputType.getShape()))
1334 return nullptr;
1335 SmallVector<bool> dimsToMask = getDimsToReduce(linalgOp);
1336 return buildMultiDimReduce(b, op, reduceVec, outputVec, dimsToMask);
1337}
1338
1339/// Generic vectorization for a single operation `op`, given already vectorized
1340/// operands carried by `bvm`. Vectorization occurs as follows:
1341/// 1. Try to apply any of the `customVectorizationHooks` and return its
1342/// result on success.
1343/// 2. Clone any constant in the current scope without vectorization: each
1344/// consumer of the constant will later determine the shape to which the
1345/// constant needs to be broadcast to.
1346/// 3. Fail on any remaining non `ElementwiseMappable` op. It is the purpose
1347/// of the `customVectorizationHooks` to cover such cases.
1348/// 4. Clone `op` in vector form to a vector of shape prescribed by the first
1349/// operand of maximal rank. Other operands have smaller rank and are
1350/// broadcast accordingly. It is assumed this broadcast is always legal,
1351/// otherwise, it means one of the `customVectorizationHooks` is incorrect.
1352///
1353/// This function assumes all operands of `op` have been vectorized and are in
1354/// the `bvm` mapping. As a consequence, this function is meant to be called on
1355/// a topologically-sorted list of ops.
1356/// This function does not update `bvm` but returns a VectorizationHookStatus
1357/// that instructs the caller what `bvm` update needs to occur.
1358static VectorizationHookResult
1359vectorizeOneOp(RewriterBase &rewriter, VectorizationState &state,
1360 LinalgOp linalgOp, Operation *op, const IRMapping &bvm,
1361 ArrayRef<CustomVectorizationHook> customVectorizationHooks) {
1362 LDBG() << "vectorize op " << *op;
1363
1364 // 1. Try to apply any CustomVectorizationHook.
1365 if (!customVectorizationHooks.empty()) {
1366 for (auto &customFunc : customVectorizationHooks) {
1367 VectorizationHookResult result = customFunc(op, bvm);
1369 continue;
1370 return result;
1371 }
1372 }
1373
1374 // 2. Constant ops don't get vectorized but rather broadcasted at their users.
1375 // Clone so that the constant is not confined to the linalgOp block .
1376 if (isa<arith::ConstantOp, func::ConstantOp>(op))
1378 rewriter.clone(*op)};
1379
1380 // 3. Only ElementwiseMappable are allowed in the generic vectorization.
1383
1384 // 4 . Check if the operation is a reduction.
1385 SmallVector<std::pair<Value, Value>> reductionOperands;
1386 for (Value operand : op->getOperands()) {
1387 auto blockArg = dyn_cast<BlockArgument>(operand);
1388 if (!blockArg || blockArg.getOwner() != linalgOp.getBlock() ||
1389 blockArg.getArgNumber() < linalgOp.getNumDpsInputs())
1390 continue;
1391 SmallVector<Operation *> reductionOps;
1392 Value reduceValue = matchReduction(
1393 linalgOp.getRegionOutputArgs(),
1394 blockArg.getArgNumber() - linalgOp.getNumDpsInputs(), reductionOps);
1395 if (!reduceValue)
1396 continue;
1397 reductionOperands.push_back(std::make_pair(reduceValue, operand));
1398 }
1399 if (!reductionOperands.empty()) {
1400 assert(reductionOperands.size() == 1);
1401 Operation *reduceOp =
1402 reduceIfNeeded(rewriter, linalgOp, op, reductionOperands[0].first,
1403 reductionOperands[0].second, bvm);
1404 if (reduceOp)
1406 }
1407
1408 // 5. Generic vectorization path for ElementwiseMappable ops.
1409 // a. Get the first max ranked shape.
1410 VectorType firstMaxRankedType;
1411 for (Value operand : op->getOperands()) {
1412 auto vecOperand = bvm.lookup(operand);
1413 assert(vecOperand && "Vector operand couldn't be found");
1414
1415 auto vecType = dyn_cast<VectorType>(vecOperand.getType());
1416 if (vecType && (!firstMaxRankedType ||
1417 firstMaxRankedType.getRank() < vecType.getRank()))
1418 firstMaxRankedType = vecType;
1419 }
1420 // b. Broadcast each op if needed.
1421 SmallVector<Value> vecOperands;
1422 for (Value scalarOperand : op->getOperands()) {
1423 Value vecOperand = bvm.lookup(scalarOperand);
1424 assert(vecOperand && "Vector operand couldn't be found");
1425
1426 if (firstMaxRankedType) {
1427 auto vecType = VectorType::get(firstMaxRankedType.getShape(),
1428 getElementTypeOrSelf(vecOperand.getType()),
1429 firstMaxRankedType.getScalableDims());
1430 vecOperands.push_back(broadcastIfNeeded(rewriter, vecOperand, vecType));
1431 } else {
1432 vecOperands.push_back(vecOperand);
1433 }
1434 }
1435 // c. for elementwise, the result is the vector with the firstMaxRankedShape
1436 SmallVector<Type> resultTypes;
1437 for (Type resultType : op->getResultTypes()) {
1438 resultTypes.push_back(
1439 firstMaxRankedType
1440 ? VectorType::get(firstMaxRankedType.getShape(), resultType,
1441 firstMaxRankedType.getScalableDims())
1442 : resultType);
1443 }
1444 // d. Build and return the new op.
1446 op->getLoc(), op->getName(), resultTypes, vecOperands,
1448 /*successors=*/{}, /*numRegions=*/0);
1450 rewriter.insert(newOp)};
1451}
1452
1453/// Generic vectorization function that rewrites the body of a `linalgOp` into
1454/// vector form. Generic vectorization proceeds as follows:
1455/// 1. Verify the `linalgOp` has one non-empty region.
1456/// 2. Values defined above the region are mapped to themselves and will be
1457/// broadcasted on a per-need basis by their consumers.
1458/// 3. Each region argument is vectorized into a vector.transfer_read (or 0-d
1459/// load).
1460/// TODO: Reuse opportunities for RAR dependencies.
1461/// 4a. Register CustomVectorizationHook for YieldOp to capture the results.
1462/// 4rewriter. Register CustomVectorizationHook for IndexOp to access the
1463/// iteration indices.
1464/// 5. Iteratively call vectorizeOneOp on the region operations.
1465///
1466/// When `broadcastToMaximalCommonShape` is set to true, eager broadcasting is
1467/// performed to the maximal common vector size implied by the `linalgOp`
1468/// iteration space. This eager broadcasting is introduced in the
1469/// permutation_map of the vector.transfer_read operations. The eager
1470/// broadcasting makes it trivial to determine where broadcast, transposes and
1471/// reductions should occur, without any bookkeeping. The tradeoff is that, in
1472/// the absence of good canonicalizations, the amount of work increases.
1473/// This is not deemed a problem as we expect canonicalizations and foldings to
1474/// aggressively clean up the useless work.
1475static LogicalResult
1476vectorizeAsLinalgGeneric(RewriterBase &rewriter, VectorizationState &state,
1477 LinalgOp linalgOp,
1478 SmallVectorImpl<Value> &newResults) {
1479 LDBG() << "Vectorizing operation as linalg generic/n";
1480 Block *block = linalgOp.getBlock();
1481
1482 // 2. Values defined above the region can only be broadcast for now. Make them
1483 // map to themselves.
1484 IRMapping bvm;
1485 SetVector<Value> valuesSet;
1486 mlir::getUsedValuesDefinedAbove(linalgOp->getRegion(0), valuesSet);
1487 bvm.map(valuesSet.getArrayRef(), valuesSet.getArrayRef());
1488
1489 if (linalgOp.getNumDpsInits() == 0)
1490 return failure();
1491
1492 // 3. Turn all BBArgs into vector.transfer_read / load.
1493 Location loc = linalgOp.getLoc();
1494 Value zero = arith::ConstantIndexOp::create(rewriter, loc, 0);
1495 for (OpOperand *opOperand : linalgOp.getOpOperandsMatchingBBargs()) {
1496 BlockArgument bbarg = linalgOp.getMatchingBlockArgument(opOperand);
1497 if (linalgOp.isScalar(opOperand)) {
1498 bvm.map(bbarg, opOperand->get());
1499 continue;
1500 }
1501
1502 // 3.a. Convert the indexing map for this input/output to a transfer read
1503 // permutation map and masking map.
1504 AffineMap indexingMap = linalgOp.getMatchingIndexingMap(opOperand);
1505
1506 AffineMap readMap;
1507 VectorType readType;
1508 Type elemType = getElementTypeOrSelf(opOperand->get());
1509 if (linalgOp.isDpsInput(opOperand)) {
1510 // 3.a.i. For input reads we use the canonical vector shape.
1511 readMap = inverseAndBroadcastProjectedPermutation(indexingMap);
1512 readType = state.getCanonicalVecType(elemType);
1513 } else {
1514 // 3.a.ii. For output reads (iteration-carried dependence, e.g.,
1515 // reductions), the vector shape is computed by mapping the canonical
1516 // vector shape to the output domain and back to the canonical domain.
1517 readMap = inversePermutation(reindexIndexingMap(indexingMap));
1518 readType =
1519 state.getCanonicalVecType(elemType, readMap.compose(indexingMap));
1520 }
1521
1522 SmallVector<Value> indices(linalgOp.getShape(opOperand).size(), zero);
1523
1524 Operation *read = vector::TransferReadOp::create(
1525 rewriter, loc, readType, opOperand->get(), indices,
1526 /*padding=*/std::nullopt, readMap);
1527 read = state.maskOperation(rewriter, read, linalgOp, indexingMap);
1528 Value readValue = read->getResult(0);
1529
1530 // 3.b. If masked, set in-bounds to true. Masking guarantees that the access
1531 // will be in-bounds.
1532 if (auto maskOp = dyn_cast<vector::MaskingOpInterface>(read)) {
1533 SmallVector<bool> inBounds(readType.getRank(), true);
1534 cast<vector::TransferReadOp>(maskOp.getMaskableOp())
1535 .setInBoundsAttr(rewriter.getBoolArrayAttr(inBounds));
1536 }
1537
1538 // 3.c. Not all ops support 0-d vectors, extract the scalar for now.
1539 // TODO: remove this.
1540 if (readType.getRank() == 0)
1541 readValue = vector::ExtractOp::create(rewriter, loc, readValue,
1543
1544 LDBG() << "New vectorized bbarg(" << bbarg.getArgNumber()
1545 << "): " << readValue;
1546 bvm.map(bbarg, readValue);
1547 bvm.map(opOperand->get(), readValue);
1548 }
1549
1551 // 4a. Register CustomVectorizationHook for yieldOp.
1552 CustomVectorizationHook vectorizeYield =
1553 [&](Operation *op, const IRMapping &bvm) -> VectorizationHookResult {
1554 return vectorizeLinalgYield(rewriter, op, bvm, state, linalgOp, newResults);
1555 };
1556 hooks.push_back(vectorizeYield);
1557
1558 // 4b. Register CustomVectorizationHook for indexOp.
1559 CustomVectorizationHook vectorizeIndex =
1560 [&](Operation *op, const IRMapping &bvm) -> VectorizationHookResult {
1561 return vectorizeLinalgIndex(rewriter, state, op, linalgOp);
1562 };
1563 hooks.push_back(vectorizeIndex);
1564
1565 // 4c. Register CustomVectorizationHook for extractOp.
1566 CustomVectorizationHook vectorizeExtract =
1567 [&](Operation *op, const IRMapping &bvm) -> VectorizationHookResult {
1568 return vectorizeTensorExtract(rewriter, state, op, linalgOp, bvm);
1569 };
1570 hooks.push_back(vectorizeExtract);
1571
1572 // 5. Iteratively call `vectorizeOneOp` to each op in the slice.
1573 for (Operation &op : block->getOperations()) {
1575 vectorizeOneOp(rewriter, state, linalgOp, &op, bvm, hooks);
1577 LDBG() << "failed to vectorize: " << op;
1578 return failure();
1579 }
1580 if (result.status == VectorizationHookStatus::NewOp) {
1581 Operation *maybeMaskedOp =
1582 state.maskOperation(rewriter, result.newOp, linalgOp);
1583 LDBG() << "New vector op: " << *maybeMaskedOp;
1584 bvm.map(op.getResults(), maybeMaskedOp->getResults());
1585 }
1586 }
1587
1588 return success();
1589}
1590
1591/// Given the re-associations, "collapses" the input Vector type
1592///
1593/// This is similar to CollapseShapeOp::inferCollapsedType with two notable
1594/// differences:
1595/// * We can safely assume that there are no dynamic sizes.
1596/// * Scalable flags are updated alongside regular dims.
1597///
1598/// When collapsing scalable flags, conservatively avoids cases with two
1599/// scalable dims. We could re-visit this in the future.
1600///
1601/// EXAMPLE:
1602/// type = vector<4x16x[8]x16xf32>
1603/// reassociation = [(d0, d1, d2, d3) -> (d0, d1),
1604/// (d0, d1, d2, d3) -> (d2, d3)]
1605/// Result:
1606/// vector<64x[128]xf32>
1607static VectorType getCollapsedVecType(VectorType type,
1608 ArrayRef<AffineMap> reassociation) {
1609 assert(type.getNumScalableDims() < 2 &&
1610 "Collapsing more than 1 scalable dim is not supported ATM");
1611
1612 // Use the fact that reassociation is valid to simplify the logic: only use
1613 // each map's rank.
1614 assert(isReassociationValid(reassociation) && "invalid reassociation");
1615
1616 auto shape = type.getShape();
1617 auto scalableFlags = type.getScalableDims();
1618 SmallVector<int64_t> newShape;
1619 SmallVector<bool> newScalableFlags;
1620
1621 unsigned currentDim = 0;
1622 for (AffineMap m : reassociation) {
1623 unsigned dim = m.getNumResults();
1624 int64_t size = 1;
1625 bool flag = false;
1626 for (unsigned d = 0; d < dim; ++d) {
1627 size *= shape[currentDim + d];
1628 flag |= scalableFlags[currentDim + d];
1629 }
1630 newShape.push_back(size);
1631 newScalableFlags.push_back(flag);
1632 currentDim += dim;
1633 }
1634
1635 return VectorType::get(newShape, type.getElementType(), newScalableFlags);
1636}
1637
1638/// Vectorize `linalg.pack` as:
1639/// * xfer_read -> shape_cast -> transpose -> xfer_write
1640///
1641/// The input-vector-sizes specify the _write_ vector sizes (i.e. the vector
1642/// sizes for the xfer_write operation). This is sufficient to infer the other
1643/// vector sizes required here.
1644///
1645/// If the vector sizes are not provided:
1646/// * the vector sizes are determined from the destination tensor static shape.
1647/// * the inBounds attribute is used instead of masking.
1648///
1649/// EXAMPLE (no vector sizes):
1650/// ```
1651/// %pack = tensor.pack %src
1652/// inner_dims_pos = [2, 1]
1653/// inner_tiles = [16, 2]
1654/// into %dst : tensor<32x8x16xf32> -> tensor<32x4x1x16x2xf32>
1655/// ``
1656/// is vectorizes as:
1657/// ```
1658/// %read = vector.transfer_read %src
1659/// : tensor<32x7x16xf32>, vector<32x8x16xf32>
1660/// %sc = vector.shape_cast %read
1661/// : vector<32x8x16xf32> to vector<32x4x2x1x16xf32>
1662/// %tr = vector.transpose %sc, [0, 1, 3, 4, 2]
1663/// : vector<32x4x2x1x16xf32> to vector<32x4x1x16x2xf32>
1664/// %write = vector.transfer_write %tr into %dest
1665/// : vector<32x4x1x16x2xf32>, tensor<32x4x1x16x2xf32>
1666/// ```
1667static LogicalResult
1668vectorizeAsTensorPackOp(RewriterBase &rewriter, linalg::PackOp packOp,
1669 ArrayRef<int64_t> inputVectorSizes,
1670 SmallVectorImpl<Value> &newResults) {
1671 if (!inputVectorSizes.empty()) {
1672 assert(inputVectorSizes.size() == packOp.getDestRank() &&
1673 "Invalid number of input vector sizes!");
1674 }
1675
1676 // TODO: Introduce a parent class that will handle the insertion point update.
1677 OpBuilder::InsertionGuard g(rewriter);
1678 rewriter.setInsertionPoint(packOp);
1679
1680 Location loc = packOp.getLoc();
1681 std::optional<Value> padValue = packOp.getPaddingValue()
1682 ? std::optional(packOp.getPaddingValue())
1683 : std::nullopt;
1684
1685 SmallVector<int64_t> destShape =
1686 SmallVector<int64_t>(packOp.getDestType().getShape());
1687
1688 // This is just a convenience alias to clearly communicate that the input
1689 // vector sizes determine the _write_ sizes.
1690 ArrayRef<int64_t> &writeVectorSizes = inputVectorSizes;
1691
1692 // In the absence of input-vector-sizes, use the _static_ input tensor shape.
1693 // In addition, use the inBounds attribute instead of masking.
1694 bool useInBoundsInsteadOfMasking = false;
1695 if (writeVectorSizes.empty()) {
1696 if (ShapedType::isDynamicShape(destShape))
1697 return rewriter.notifyMatchFailure(packOp,
1698 "unable to infer vector sizes");
1699
1700 writeVectorSizes = destShape;
1701 useInBoundsInsteadOfMasking = true;
1702 }
1703
1704 // Compute pre-transpose-write-vector-type, i.e. the write vector type
1705 // _before_ the transposition (i.e. before dimension permutation). This is
1706 // done by inverting the permutation/transposition that's part of the Pack
1707 // operation. This type is required to:
1708 // 1) compute the read vector type for masked-read below, and
1709 // 2) generate shape-cast Op below that expands the read vector type.
1710 PackingMetadata packMetadata;
1711 SmallVector<int64_t> preTransposeWriteVecSizses(writeVectorSizes);
1712 auto destInvPermutation = getPackInverseDestPerm(packOp, packMetadata);
1713 applyPermutationToVector(preTransposeWriteVecSizses, destInvPermutation);
1714 auto preTransposeWriteVecType =
1715 VectorType::get(preTransposeWriteVecSizses,
1716 packOp.getResult().getType().getElementType());
1717
1718 // Compute vector type for the _read_ opeartion. This is simply
1719 // pre-transpose-write-vector-type with the dimensions collapsed
1720 // as per the Pack operation.
1721 VectorType readVecType = getCollapsedVecType(
1722 preTransposeWriteVecType,
1724 rewriter.getContext(), packMetadata.reassociations)));
1725
1726 // Create masked TransferReadOp.
1727 auto maskedRead = vector::createReadOrMaskedRead(
1728 rewriter, loc, packOp.getSource(), readVecType, padValue,
1729 useInBoundsInsteadOfMasking);
1730
1731 // Create ShapeCastOp.
1732 auto shapeCastOp = vector::ShapeCastOp::create(
1733 rewriter, loc, preTransposeWriteVecType, maskedRead);
1734
1735 // Create TransposeOp.
1736 auto destPermutation = invertPermutationVector(destInvPermutation);
1737 auto transposeOp = vector::TransposeOp::create(
1738 rewriter, loc, shapeCastOp.getResult(), destPermutation);
1739
1740 // Create TransferWriteOp.
1741 Operation *write = vector::createWriteOrMaskedWrite(
1742 rewriter, loc, transposeOp.getResult(), packOp.getDest());
1743 newResults.push_back(write->getResult(0));
1744 return success();
1745}
1746
1747/// Vectorize `linalg.unpack` as:
1748/// * xfer_read -> vector.transpose -> vector.shape_cast -> xfer_write
1749///
1750/// The input-vector-sizes specify the _read_ vector sizes (i.e. the vector
1751/// sizes for the xfer_read operation). This is sufficient to infer the other
1752/// vector sizes required here.
1753///
1754/// If the vector sizes are not provided:
1755/// * the vector sizes are determined from the input tensor static shape.
1756/// * the inBounds attribute is used instead of masking.
1757///
1758/// EXAMPLE (no vector sizes):
1759/// ```
1760/// %unpack = linalg.unpack %src
1761/// inner_dims_pos = [0, 1]
1762/// inner_tiles = [8, 8]
1763/// into %dest : tensor<1x1x8x8xf32> -> tensor<8x8xf32>
1764/// ```
1765/// is vectorized as:
1766/// ```
1767/// %read = vector.transfer_read %src
1768/// : tensor<1x1x8x8xf32>, vector<1x1x8x8xf32>
1769/// %tr = vector.transpose %read, [0, 2, 1, 3]
1770/// : vector<1x1x8x8xf32> to vector<1x8x1x8xf32>
1771/// %sc = vector.shape_cast %tr
1772/// : vector<1x8x1x8xf32> to vector<8x8xf32>
1773/// %vector = vector.transfer_write %sc into %dest
1774/// : vector<8x8xf32>, tensor<8x8xf32>
1775/// ```
1776static LogicalResult
1777vectorizeAsTensorUnpackOp(RewriterBase &rewriter, linalg::UnPackOp unpackOp,
1778 ArrayRef<int64_t> inputVectorSizes,
1779 ArrayRef<bool> inputScalableVecDims,
1780 SmallVectorImpl<Value> &newResults) {
1781 if (!inputVectorSizes.empty()) {
1782 assert(inputVectorSizes.size() == unpackOp.getSourceRank() &&
1783 "Invalid number of input vector sizes!");
1784 assert(inputVectorSizes.size() == inputScalableVecDims.size() &&
1785 "Incompatible number of vector sizes and vector scalable flags!");
1786 }
1787
1788 // TODO: Introduce a parent class that will handle the insertion point update.
1789 OpBuilder::InsertionGuard g(rewriter);
1790 rewriter.setInsertionPoint(unpackOp);
1791
1792 ShapedType unpackTensorType = unpackOp.getSourceType();
1793
1794 ArrayRef<int64_t> sourceShape = unpackTensorType.getShape();
1795 bool useInBoundsInsteadOfMasking = false;
1796
1797 Location loc = unpackOp->getLoc();
1798
1799 // Obtain vector sizes for the read operation.
1800 SmallVector<int64_t> readVectorSizes(inputVectorSizes);
1801 SmallVector<bool> readScalableVectorFlags(inputScalableVecDims);
1802
1803 // In the absence of input-vector-sizes, use the _static_ input tensor shape.
1804 if (inputVectorSizes.empty()) {
1805 if (ShapedType::isDynamicShape(sourceShape))
1806 return rewriter.notifyMatchFailure(unpackOp,
1807 "Unable to infer vector sizes!");
1808
1809 readVectorSizes.assign(sourceShape.begin(), sourceShape.end());
1810 useInBoundsInsteadOfMasking = true;
1811 }
1812
1813 // -- Generate the read operation --
1814 VectorType readVecType =
1815 VectorType::get(readVectorSizes, unpackTensorType.getElementType(),
1816 readScalableVectorFlags);
1817 Value readResult = vector::createReadOrMaskedRead(
1818 rewriter, loc, unpackOp.getSource(), readVecType, std::nullopt,
1819 useInBoundsInsteadOfMasking);
1820
1821 // -- Generate the transpose operation --
1822 PackingMetadata packMetadata;
1823 SmallVector<int64_t> lastDimToInsertPosPerm =
1824 getUnPackInverseSrcPerm(unpackOp, packMetadata);
1825 vector::TransposeOp transposeOp = vector::TransposeOp::create(
1826 rewriter, loc, readResult, lastDimToInsertPosPerm);
1827
1828 // -- Generate the shape_cast operation --
1829 VectorType collapsedVecType = getCollapsedVecType(
1830 transposeOp.getType(),
1832 rewriter.getContext(), packMetadata.reassociations)));
1833 vector::ShapeCastOp shapeCastOp = vector::ShapeCastOp::create(
1834 rewriter, loc, collapsedVecType, transposeOp->getResult(0));
1835
1836 // -- Generate the write operation --
1837 Operation *write = vector::createWriteOrMaskedWrite(
1838 rewriter, loc, shapeCastOp.getResult(), unpackOp.getDest(),
1839 /*writeIndices=*/{}, useInBoundsInsteadOfMasking);
1840
1841 newResults.push_back(write->getResult(0));
1842 return success();
1843}
1844
1845/// Vectorize a `padOp` with (1) static result type, (2) constant padding value
1846/// and (3) all-zero lowPad to
1847/// `transfer_write_in_bounds(transfer_read_masked(pad_source, pad_value))`.
1848static LogicalResult
1849vectorizeAsTensorPadOp(RewriterBase &rewriter, tensor::PadOp padOp,
1850 ArrayRef<int64_t> inputVectorSizes,
1851 SmallVectorImpl<Value> &newResults) {
1852 auto padValue = padOp.getConstantPaddingValue();
1853 Location loc = padOp.getLoc();
1854
1855 // TODO: Introduce a parent class that will handle the insertion point update.
1856 OpBuilder::InsertionGuard g(rewriter);
1857 rewriter.setInsertionPoint(padOp);
1858
1859 ReifiedRankedShapedTypeDims reifiedReturnShapes;
1860 LogicalResult status =
1861 cast<ReifyRankedShapedTypeOpInterface>(padOp.getOperation())
1862 .reifyResultShapes(rewriter, reifiedReturnShapes);
1863 (void)status; // prevent unused variable warning on non-assert builds
1864 assert(succeeded(status) && "failed to reify result shapes");
1865 auto readType = VectorType::get(inputVectorSizes, padValue.getType());
1866 auto maskedRead = vector::createReadOrMaskedRead(
1867 rewriter, loc, padOp.getSource(), readType, padValue,
1868 /*useInBoundsInsteadOfMasking=*/false);
1869
1870 // Create Xfer write Op
1871 Value dest = tensor::EmptyOp::create(rewriter, loc, reifiedReturnShapes[0],
1872 padOp.getResultType().getElementType());
1873 Operation *write =
1874 vector::createWriteOrMaskedWrite(rewriter, loc, maskedRead, dest);
1875 newResults.push_back(write->getResult(0));
1876 return success();
1877}
1878
1879// TODO: probably need some extra checks for reduction followed by consumer
1880// ops that may not commute (e.g. linear reduction + non-linear instructions).
1881static LogicalResult reductionPreconditions(LinalgOp op) {
1882 if (llvm::none_of(op.getIteratorTypesArray(), isReductionIterator)) {
1883 LDBG() << "reduction precondition failed: no reduction iterator";
1884 return failure();
1885 }
1886 for (OpOperand &opOperand : op.getDpsInitsMutable()) {
1887 AffineMap indexingMap = op.getMatchingIndexingMap(&opOperand);
1888 if (indexingMap.isPermutation())
1889 continue;
1890
1891 Operation *reduceOp = matchLinalgReduction(&opOperand);
1892 if (!reduceOp || !getCombinerOpKind(reduceOp)) {
1893 LDBG() << "reduction precondition failed: reduction detection failed";
1894 return failure();
1895 }
1896 }
1897 return success();
1898}
1899
1900static LogicalResult
1901vectorizeDynamicConvOpPrecondition(linalg::LinalgOp conv,
1902 bool flatten1DDepthwiseConv) {
1903 if (flatten1DDepthwiseConv) {
1904 LDBG() << "Vectorization of flattened convs with dynamic shapes is not "
1905 "supported";
1906 return failure();
1907 }
1908
1910 LDBG() << "Not a 1D depth-wise WC conv, dynamic shapes are not supported";
1911 return failure();
1912 }
1913
1914 // Support dynamic shapes in 1D depthwise convolution, but only in the
1915 // _channel_ dimension.
1916 Value lhs = conv.getDpsInputOperand(0)->get();
1917 ArrayRef<int64_t> lhsShape = cast<ShapedType>(lhs.getType()).getShape();
1918 auto shapeWithoutCh = lhsShape.drop_back(1);
1919 if (ShapedType::isDynamicShape(shapeWithoutCh)) {
1920 LDBG() << "Dynamically-shaped op vectorization precondition failed: only "
1921 "channel dim can be dynamic";
1922 return failure();
1923 }
1924
1925 return success();
1926}
1927
1928static LogicalResult
1929vectorizeDynamicLinalgOpPrecondition(linalg::LinalgOp op,
1930 bool flatten1DDepthwiseConv) {
1932 return vectorizeDynamicConvOpPrecondition(op, flatten1DDepthwiseConv);
1933
1934 if (hasReductionIterator(op))
1935 return reductionPreconditions(op);
1936
1937 // TODO: Masking only supports dynamic element-wise ops, linalg.generic ops,
1938 // linalg.copy ops and ops that implement ContractionOpInterface for now.
1939 if (!isElementwise(op) &&
1940 !isa<linalg::GenericOp, linalg::CopyOp, linalg::ContractionOpInterface>(
1941 op.getOperation()))
1942 return failure();
1943
1944 LDBG() << "Dynamically-shaped op meets vectorization pre-conditions";
1945 return success();
1946}
1947
1948//// This hook considers two cases:
1949/// (1) If the input-vector-sizes are empty, then the vector sizes will be
1950/// infered. This is only possible when all shapes are static.
1951/// (2) If the input-vector-sizes are non-empty (i.e. user provided), then
1952/// carry out basic sanity-checking.
1953static LogicalResult
1954vectorizeUnPackOpPrecondition(linalg::UnPackOp unpackOp,
1955 ArrayRef<int64_t> inputVectorSizes) {
1956 // Pack/unpack memref transformations are unsupported. The memref forms
1957 // are mainly for bufferization and scalar lowering. Other uses are not
1958 // recommended, see #225650 for details.
1959 if (!unpackOp.hasPureTensorSemantics())
1960 return failure();
1961
1962 // If there are no input vector sizes and all shapes are static, there is
1963 // nothing left to check.
1964 if (inputVectorSizes.empty() && unpackOp.getDestType().hasStaticShape() &&
1965 unpackOp.getSourceType().hasStaticShape())
1966 return success();
1967
1968 // The number of input vector sizes must be equal to:
1969 // * read-vector-rank
1970 if (!inputVectorSizes.empty() &&
1971 (inputVectorSizes.size() != unpackOp.getSourceRank())) {
1972 LDBG() << "Incorrect number of input vector sizes";
1973 return failure();
1974 }
1975
1976 // Check the vector sizes for the read operation.
1978 unpackOp.getSourceType().getShape(), inputVectorSizes))) {
1979 LDBG() << "Invalid vector sizes for the read operation";
1980 return failure();
1981 }
1982
1983 return success();
1984}
1985
1986static LogicalResult
1987vectorizeInsertSliceOpPrecondition(tensor::InsertSliceOp sliceOp,
1988 ArrayRef<int64_t> inputVectorSizes) {
1989
1990 TypedValue<RankedTensorType> source = sliceOp.getSource();
1991 auto sourceType = source.getType();
1992 if (!VectorType::isValidElementType(sourceType.getElementType()))
1993 return failure();
1994
1995 // Get the pad value.
1996 // TransferReadOp (which is used to vectorize InsertSliceOp), requires a
1997 // scalar padding value. Note that:
1998 // * for in-bounds accesses,
1999 // the value is actually irrelevant. There are 2 cases in which xfer.read
2000 // accesses are known to be in-bounds:
2001 // 1. The source shape is static (output vector sizes would be based on
2002 // the source shape and hence all memory accesses would be in-bounds),
2003 // 2. Masking is used, i.e. the output vector sizes are user-provided. In
2004 // this case it is safe to assume that all memory accesses are in-bounds.
2005 //
2006 // When the value is not known and not needed, use 0. Otherwise, bail out.
2007 Value padValue = getStaticPadVal(sliceOp);
2008 bool isOutOfBoundsRead =
2009 !sourceType.hasStaticShape() && inputVectorSizes.empty();
2010
2011 if (!padValue && isOutOfBoundsRead) {
2012 LDBG() << "Failed to get a pad value for out-of-bounds read access";
2013 return failure();
2014 }
2015 return success();
2016}
2017
2018/// Vectorize a named linalg contraction op into:
2019/// vector::TransferReadOp - Reads vectors from the operands
2020/// vector::ContractionOp - Performs contraction
2021/// vector::TransferWriteOp - Write the result vector back to the
2022/// destination
2023/// The operands shapes are preserved and loaded directly into vectors.
2024/// Any further permutations or numerical casting remain within contraction op.
2025static LogicalResult
2026vectorizeAsLinalgContraction(RewriterBase &rewriter, VectorizationState &state,
2027 LinalgOp linalgOp,
2028 SmallVectorImpl<Value> &newResults) {
2029 Location loc = linalgOp.getLoc();
2030 MLIRContext *ctx = linalgOp.getContext();
2031
2032 // For simplicity, contraction vectorization is limited to linalg named ops.
2033 // Generic op is ignored as not every arbitrary contraction body can be
2034 // expressed by a vector.contract.
2035 if (!isa<ContractionOpInterface>(linalgOp.getOperation()))
2036 return failure();
2037
2038 OpOperand *outOperand = linalgOp.getDpsInitOperand(0);
2039 Operation *reduceOp = matchLinalgReduction(outOperand);
2040 auto maybeKind = getCombinerOpKind(reduceOp);
2041 if (!maybeKind) {
2042 LDBG() << "Failed to determine contraction combining kind.";
2043 return failure();
2044 }
2045
2046 // Check that all dimensions are present in the input operands.
2047 // Arbitrary broadcasts are not supported by the vector contraction.
2048 // Broadcasts are expected to be decomposed before vectorization.
2049 AffineMap lhsMap = linalgOp.getIndexingMapsArray()[0];
2050 AffineMap rhsMap = linalgOp.getIndexingMapsArray()[1];
2051 if (getUnusedDimsBitVector({lhsMap, rhsMap}).any()) {
2052 LDBG() << "Contractions with broadcasts are not supported.";
2053 return failure();
2054 }
2055
2056 // Load operands.
2057 SmallVector<Value> vecOperands;
2058 for (OpOperand &opOperand : linalgOp->getOpOperands()) {
2059 // The operand vector shape is computed by mapping the canonical vector
2060 // shape to the operand's domain. Further permutations are left as a part of
2061 // the contraction.
2062 AffineMap indexingMap = linalgOp.getMatchingIndexingMap(&opOperand);
2063 AffineMap readMap = AffineMap::getMultiDimIdentityMap(
2064 indexingMap.getNumResults(), rewriter.getContext());
2065 Type elemType = getElementTypeOrSelf(opOperand.get());
2066 VectorType readType =
2067 state.getCanonicalVecType(elemType, readMap.compose(indexingMap));
2068
2070 rewriter, loc, opOperand.get(), readType,
2071 /*padding=*/arith::getZeroConstant(rewriter, loc, elemType),
2072 /*useInBoundsInsteadOfMasking=*/false);
2073 vecOperands.push_back(read);
2074 }
2075
2076 // Preserve the contraction's cast semantics when converting operands to the
2077 // integer accumulator type. vector.contract provides an implicit signed
2078 // integer promotion; the cases below materialize explicit casts as needed.
2079 auto castAttr = linalgOp->getAttrOfType<TypeFnAttr>("cast");
2080 bool hasUnsignedCast =
2081 castAttr && castAttr.getValue() == TypeFn::cast_unsigned;
2082 auto accType = dyn_cast<VectorType>(vecOperands[2].getType());
2083 auto accElementType =
2084 accType ? dyn_cast<IntegerType>(accType.getElementType()) : nullptr;
2085 if (accElementType && accElementType.isSignless()) {
2086 for (Value &operand : MutableArrayRef(vecOperands).take_front(2)) {
2087 auto operandType = cast<VectorType>(operand.getType());
2088 Type operandElementType = operandType.getElementType();
2089 VectorType castType = operandType.clone(accElementType);
2090
2091 if (isa<FloatType>(operandElementType)) {
2092 operand =
2093 hasUnsignedCast
2094 ? arith::FPToUIOp::create(rewriter, loc, castType, operand)
2095 .getResult()
2096 : arith::FPToSIOp::create(rewriter, loc, castType, operand)
2097 .getResult();
2098 continue;
2099 }
2100
2101 auto operandIntegerType = dyn_cast<IntegerType>(operandElementType);
2102 if (!operandIntegerType || !operandIntegerType.isSignless())
2103 continue;
2104 if (operandIntegerType.getWidth() >= accElementType.getWidth())
2105 continue;
2106 if (!hasUnsignedCast)
2107 continue;
2108
2109 // vector.contract implicitly sign-extends integer operands. Unsigned
2110 // promotion therefore requires an explicit zero extension.
2111 operand = arith::ExtUIOp::create(rewriter, loc, castType, operand);
2112 }
2113 }
2114
2115 // Remap iterators from linalg to vector.
2116 SmallVector<Attribute> iterAttrs;
2117 auto iterators = linalgOp.getIteratorTypesArray();
2118 for (utils::IteratorType iter : iterators) {
2119 auto vecIter = iter == utils::IteratorType::parallel
2120 ? vector::IteratorType::parallel
2121 : vector::IteratorType::reduction;
2122 iterAttrs.push_back(vector::IteratorTypeAttr::get(ctx, vecIter));
2123 }
2124
2125 // Create contraction.
2126 Operation *contractOp = vector::ContractionOp::create(
2127 rewriter, loc, /*lhs=*/vecOperands[0],
2128 /*rhs=*/vecOperands[1], /*acc=*/vecOperands[2],
2129 linalgOp.getIndexingMaps(), rewriter.getArrayAttr(iterAttrs), *maybeKind);
2130 contractOp = state.maskOperation(rewriter, contractOp, linalgOp);
2131
2132 // Store result.
2133 Operation *write = vector::createWriteOrMaskedWrite(
2134 rewriter, loc, contractOp->getResult(0), outOperand->get());
2135
2136 // Finalize.
2137 if (!write->getResults().empty())
2138 newResults.push_back(write->getResult(0));
2139
2140 return success();
2141}
2142
2143namespace {
2144enum class ConvOperationKind { Conv, Pool };
2145} // namespace
2146
2147static bool isCastOfBlockArgument(Operation *op) {
2148 return isa<CastOpInterface>(op) && op->getNumOperands() == 1 &&
2149 isa<BlockArgument>(op->getOperand(0));
2150}
2151
2152// Returns the ConvOperationKind of the op using reduceOp of the generic
2153// payload. If it is neither a convolution nor a pooling, it returns
2154// std::nullopt.
2155//
2156// If (region has 2 ops (reduction + yield) or 3 ops (extension + reduction
2157// + yield) and rhs is not used) then it is the body of a pooling
2158// If conv, check for single `mul` predecessor. The `mul` operands must be
2159// block arguments or extension of block arguments.
2160// Otherwise, check for one or zero `ext` predecessor. The `ext` operands
2161// must be block arguments or extension of block arguments.
2162static std::optional<ConvOperationKind>
2163getConvOperationKind(Operation *reduceOp) {
2164 int numBlockArguments =
2165 llvm::count_if(reduceOp->getOperands(), llvm::IsaPred<BlockArgument>);
2166
2167 switch (numBlockArguments) {
2168 case 1: {
2169 // Will be convolution if feeder is a MulOp.
2170 // A strength reduced version of MulOp for i1 type is AndOp which is also
2171 // supported. Otherwise, it can be pooling. This strength reduction logic
2172 // is in `buildBinaryFn` helper in the Linalg dialect.
2173 auto feedValIt = llvm::find_if_not(reduceOp->getOperands(),
2174 llvm::IsaPred<BlockArgument>);
2175 assert(feedValIt != reduceOp->operand_end() &&
2176 "Expected a non-block argument operand");
2177 Operation *feedOp = (*feedValIt).getDefiningOp();
2178 if (isCastOfBlockArgument(feedOp)) {
2179 return ConvOperationKind::Pool;
2180 }
2181
2182 if (!((isa<arith::MulIOp, arith::MulFOp>(feedOp) ||
2183 (isa<arith::AndIOp>(feedOp) &&
2184 feedOp->getResultTypes()[0].isInteger(1))) &&
2185 llvm::all_of(feedOp->getOperands(), [](Value v) {
2186 if (isa<BlockArgument>(v))
2187 return true;
2188 if (Operation *op = v.getDefiningOp())
2189 return isCastOfBlockArgument(op);
2190 return false;
2191 }))) {
2192 return std::nullopt;
2193 }
2194
2195 return ConvOperationKind::Conv;
2196 }
2197 case 2:
2198 // Must be pooling
2199 return ConvOperationKind::Pool;
2200 default:
2201 return std::nullopt;
2202 }
2203}
2204
2205static bool isSupportedPoolKind(vector::CombiningKind kind) {
2206 switch (kind) {
2207 case vector::CombiningKind::ADD:
2208 case vector::CombiningKind::MAXNUMF:
2209 case vector::CombiningKind::MAXIMUMF:
2210 case vector::CombiningKind::MAXIMUMNUMF:
2211 case vector::CombiningKind::MAXSI:
2212 case vector::CombiningKind::MAXUI:
2213 case vector::CombiningKind::MINNUMF:
2214 case vector::CombiningKind::MINIMUMF:
2215 case vector::CombiningKind::MINIMUMNUMF:
2216 case vector::CombiningKind::MINSI:
2217 case vector::CombiningKind::MINUI:
2218 return true;
2219 default:
2220 return false;
2221 }
2222}
2223
2224static LogicalResult vectorizeConvOpPrecondition(linalg::LinalgOp convOp) {
2225 auto getOperandType = [&](auto operand) {
2226 return dyn_cast<ShapedType>((operand->get()).getType());
2227 };
2228 ShapedType lhsShapedType = getOperandType(convOp.getDpsInputOperand(0));
2229 ShapedType rhsShapedType = getOperandType(convOp.getDpsInputOperand(1));
2230 ShapedType resShapedType = getOperandType(convOp.getDpsInitOperand(0));
2231 // (LHS has dimension NCW/NWC and RES has dimension NFW/NCW/NWF/NWC) OR
2232 // (non-channeled convolution -> LHS and RHS both have single dimensions).
2233 // Note that this also ensures 2D and 3D convolutions are rejected.
2234 if ((lhsShapedType.getRank() != 3 || resShapedType.getRank() != 3) &&
2235 (lhsShapedType.getRank() != 1 || resShapedType.getRank() != 1))
2236 return failure();
2237
2238 Operation *reduceOp = matchLinalgReduction(convOp.getDpsInitOperand(0));
2239 if (!reduceOp)
2240 return failure();
2241
2242 auto maybeOper = getConvOperationKind(reduceOp);
2243 if (!maybeOper.has_value())
2244 return failure();
2245
2246 auto maybeKind = getCombinerOpKind(reduceOp);
2247 // Typically convolution will have a `Add` CombiningKind but for i1 type it
2248 // can get strength reduced to `OR` which is also supported. This strength
2249 // reduction logic is in `buildBinaryFn` helper in the Linalg dialect.
2250 if (!maybeKind || ((*maybeKind != vector::CombiningKind::ADD &&
2251 *maybeKind != vector::CombiningKind::OR) &&
2252 (*maybeOper != ConvOperationKind::Pool ||
2253 !isSupportedPoolKind(*maybeKind)))) {
2254 return failure();
2255 }
2256
2257 auto rhsRank = rhsShapedType.getRank();
2258 if (*maybeOper == ConvOperationKind::Pool) {
2259 if (rhsRank != 1)
2260 return failure();
2261 } else {
2262 if (rhsRank != 1 && rhsRank != 2 && rhsRank != 3)
2263 return failure();
2264 }
2265
2266 return success();
2267}
2268
2269static LogicalResult vectorizeLinalgOpPrecondition(
2270 LinalgOp linalgOp, ArrayRef<int64_t> inputVectorSizes,
2271 bool vectorizeNDExtract, bool flatten1DDepthwiseConv) {
2272 // tensor with dimension of 0 cannot be vectorized.
2273 if (llvm::any_of(linalgOp->getOpOperands(), [&](OpOperand &operand) {
2274 return llvm::is_contained(linalgOp.getShape(&operand), 0);
2275 }))
2276 return failure();
2277 // Check API contract for input vector sizes.
2278 if (!inputVectorSizes.empty() &&
2279 failed(vector::isValidMaskedInputVector(linalgOp.getStaticLoopRanges(),
2280 inputVectorSizes)))
2281 return failure();
2282
2283 if (linalgOp.hasDynamicShape() && failed(vectorizeDynamicLinalgOpPrecondition(
2284 linalgOp, flatten1DDepthwiseConv))) {
2285 LDBG() << "Dynamically-shaped op failed vectorization pre-conditions";
2286 return failure();
2287 }
2288
2289 SmallVector<CustomVectorizationPrecondition> customPreconditions;
2290
2291 // Register CustomVectorizationPrecondition for extractOp.
2292 customPreconditions.push_back(tensorExtractVectorizationPrecondition);
2293
2294 // All types in the body should be a supported element type for VectorType.
2295 for (Operation &innerOp : linalgOp->getRegion(0).front()) {
2296 // Check if any custom hook can vectorize the inner op.
2297 if (llvm::any_of(
2298 customPreconditions,
2299 [&](const CustomVectorizationPrecondition &customPrecondition) {
2300 return succeeded(
2301 customPrecondition(&innerOp, vectorizeNDExtract));
2302 })) {
2303 continue;
2304 }
2305 if (!llvm::all_of(innerOp.getOperandTypes(),
2306 VectorType::isValidElementType)) {
2307 return failure();
2308 }
2309 if (!llvm::all_of(innerOp.getResultTypes(),
2310 VectorType::isValidElementType)) {
2311 return failure();
2312 }
2313 }
2314 if (isElementwise(linalgOp))
2315 return success();
2316
2317 // Check for both named as well as generic convolution ops.
2318 if (isaConvolutionOpInterface(linalgOp))
2319 return vectorizeConvOpPrecondition(linalgOp);
2320
2321 // TODO: the common vector shape is equal to the static loop sizes only when
2322 // all indexing maps are projected permutations. For convs and stencils the
2323 // logic will need to evolve.
2324 if (!allIndexingsAreProjectedPermutation(linalgOp)) {
2325 LDBG() << "precondition failed: not projected permutations";
2326 return failure();
2327 }
2328 if (failed(reductionPreconditions(linalgOp))) {
2329 LDBG() << "precondition failed: reduction preconditions";
2330 return failure();
2331 }
2332 return success();
2333}
2334
2335static LogicalResult
2336vectorizePackOpPrecondition(linalg::PackOp packOp,
2337 ArrayRef<int64_t> inputVectorSizes) {
2338 // Pack/unpack memref transformations are unsupported. The memref forms
2339 // are mainly for bufferization and scalar lowering. Other uses are not
2340 // recommended, see #225650 for details.
2341 if (!packOp.hasPureTensorSemantics())
2342 return failure();
2343
2344 auto padValue = packOp.getPaddingValue();
2345 Attribute cstAttr;
2346 // TODO: Relax this condiiton
2347 if (padValue && !matchPattern(padValue, m_Constant(&cstAttr))) {
2348 LDBG() << "pad value is not constant: " << packOp;
2349 return failure();
2350 }
2351
2352 ArrayRef<int64_t> resultTensorShape = packOp.getDestType().getShape();
2353 bool satisfyEmptyCond = true;
2354 if (inputVectorSizes.empty()) {
2355 if (!packOp.getDestType().hasStaticShape() ||
2356 !packOp.getSourceType().hasStaticShape())
2357 satisfyEmptyCond = false;
2358 }
2359
2360 if (!satisfyEmptyCond &&
2362 resultTensorShape.take_front(packOp.getSourceRank()),
2363 inputVectorSizes)))
2364 return failure();
2365
2366 if (llvm::any_of(packOp.getInnerTiles(), [](OpFoldResult v) {
2367 return !getConstantIntValue(v).has_value();
2368 })) {
2369 LDBG() << "inner_tiles must be constant: " << packOp;
2370 return failure();
2371 }
2372
2373 return success();
2374}
2375
2376static LogicalResult
2377vectorizePadOpPrecondition(tensor::PadOp padOp,
2378 ArrayRef<int64_t> inputVectorSizes) {
2379 auto padValue = padOp.getConstantPaddingValue();
2380 if (!padValue) {
2381 LDBG() << "pad value is not constant: " << padOp;
2382 return failure();
2383 }
2384
2385 ArrayRef<int64_t> resultTensorShape = padOp.getResultType().getShape();
2386 if (failed(vector::isValidMaskedInputVector(resultTensorShape,
2387 inputVectorSizes)))
2388 return failure();
2389
2390 // Padding with non-zero low pad values is not supported, unless the
2391 // corresponding result dim is 1 as this would require shifting the results to
2392 // the right for the low padded dims by the required amount of low padding.
2393 // However, we do support low padding if the dims being low padded have result
2394 // sizes of 1. The reason is when we have a low pad on a unit result dim, the
2395 // input size of that dimension will be dynamically zero (as the sum of the
2396 // low pad and input dim size has to be one) and hence we will create a zero
2397 // mask as the lowering logic just makes the mask one for the input dim size -
2398 // which is zero here. Hence we will load the pad value which is what we want
2399 // in this case. If the low pad is dynamically zero then the lowering is
2400 // correct as well as no shifts are necessary.
2401 if (llvm::any_of(llvm::enumerate(padOp.getMixedLowPad()),
2402 [&](const auto &en) {
2403 OpFoldResult padValue = en.value();
2404 unsigned pos = en.index();
2405 std::optional<int64_t> pad = getConstantIntValue(padValue);
2406 return (!pad.has_value() || pad.value() != 0) &&
2407 resultTensorShape[pos] != 1;
2408 })) {
2409 LDBG() << "low pad must all be zero for all non unit dims: " << padOp;
2410 return failure();
2411 }
2412
2413 return success();
2414}
2415
2416/// Preconditions for scalable vectors.
2417///
2418/// For Ops implementing the LinalgOp interface, this is quite restrictive - it
2419/// models the fact that in practice we would only make selected dimensions
2420/// scalable. For other Ops (e.g. `linalg.unpack`), this will succeed
2421/// unconditionally - we are yet to identify meaningful conditions.
2422static LogicalResult
2423vectorizeScalableVectorPrecondition(Operation *op,
2424 ArrayRef<int64_t> inputVectorSizes,
2425 ArrayRef<bool> inputScalableVecDims) {
2426 assert(inputVectorSizes.size() == inputScalableVecDims.size() &&
2427 "Number of input vector sizes and scalable dims doesn't match");
2428
2429 size_t numOfScalableDims =
2430 llvm::count_if(inputScalableVecDims, [](bool flag) { return flag; });
2431
2432 if (numOfScalableDims == 0)
2433 return success();
2434
2435 auto linalgOp = dyn_cast<LinalgOp>(op);
2436
2437 // Cond 1: Reject Ops that don't implement the LinalgOp interface, with the
2438 // exception of UnpackOp for which there is a dedicated hook.
2439 if (!linalgOp) {
2440 return success(isa<linalg::UnPackOp>(op));
2441 }
2442
2443 // Cond 2: There's been no need for more than 2 scalable dims so far
2444 if (numOfScalableDims > 2)
2445 return failure();
2446
2447 // Cond 3: Look at the configuration in `inputScalableVecDims` and verify that
2448 // it matches one of the supported cases:
2449 // 1. Exactly 1 dim is scalable and that's the _last_ non-unit parallel dim
2450 // (*).
2451 // 2. Exactly 2 dims are scalable and those are the _last two adjacent_
2452 // parallel dims.
2453 // 3. Exactly 1 reduction dim is scalable and that's the last (innermost)
2454 // dim.
2455 // The 2nd restriction above means that only Matmul-like Ops are supported
2456 // when 2 dims are scalable, e.g. :
2457 // * iterators = [parallel, parallel, reduction]
2458 // * scalable flags = [true, true, false]
2459 //
2460 // (*) Non-unit dims get folded away in practice.
2461 // TODO: Relax these conditions as good motivating examples are identified.
2462
2463 // Find the first scalable flag.
2464 bool seenNonUnitParallel = false;
2465 auto iterators = linalgOp.getIteratorTypesArray();
2466 SmallVector<bool> scalableFlags(inputScalableVecDims);
2467 int64_t idx = scalableFlags.size() - 1;
2468 while (!scalableFlags[idx]) {
2469 bool isNonUnitDim = (inputVectorSizes[idx] != 1);
2470 seenNonUnitParallel |=
2471 (iterators[idx] == utils::IteratorType::parallel && isNonUnitDim);
2472
2473 iterators.pop_back();
2474 scalableFlags.pop_back();
2475 --idx;
2476 }
2477
2478 // Analyze the iterator corresponding to the first scalable dim.
2479 switch (iterators.back()) {
2480 case utils::IteratorType::reduction: {
2481 // Check 3. above is met.
2482 if (iterators.size() != inputVectorSizes.size()) {
2483 LDBG() << "Non-trailing reduction dim requested for scalable "
2484 "vectorization";
2485 return failure();
2486 }
2487 if (isa<linalg::MatmulOp>(op)) {
2488 LDBG()
2489 << "Scalable vectorization of the reduction dim in Matmul-like ops "
2490 "is not supported";
2491 return failure();
2492 }
2493 break;
2494 }
2495 case utils::IteratorType::parallel: {
2496 // Check 1. and 2. above are met.
2497 if (seenNonUnitParallel) {
2498 LDBG() << "Inner parallel dim not requested for scalable "
2499 "vectorization";
2500 return failure();
2501 }
2502 break;
2503 }
2504 }
2505
2506 // If present, check the 2nd scalable dim. ATM, only Matmul-like Ops are
2507 // supported for which expect the folowing config:
2508 // * iterators = [parallel, parallel, reduction]
2509 // * scalable flags = [true, true, false]
2510 if (numOfScalableDims == 2) {
2511 // Disallow below case which breaks 3. above:
2512 // * iterators = [..., parallel, reduction]
2513 // * scalable flags = [..., true, true]
2514 if (iterators.back() == utils::IteratorType::reduction) {
2515 LDBG() << "Higher dim than the trailing reduction dim requested for "
2516 "scalable "
2517 "vectorizatio";
2518 return failure();
2519 }
2520 scalableFlags.pop_back();
2521 iterators.pop_back();
2522
2523 if (!scalableFlags.back() ||
2524 (iterators.back() != utils::IteratorType::parallel))
2525 return failure();
2526 }
2527
2528 // Cond 4: Only the following ops are supported in the
2529 // presence of scalable vectors
2530 return success(
2531 isElementwise(linalgOp) || isa<linalg::MatmulOp>(op) ||
2532 isa<linalg::BatchMatmulOp>(op) ||
2534 isa<linalg::MatvecOp>(op) || isa<linalg::Mmt4DOp>(op) ||
2535 isa<linalg::BatchMmt4DOp>(op) || hasReductionIterator(linalgOp));
2536}
2537
2539 Operation *op, ArrayRef<int64_t> inputVectorSizes,
2540 ArrayRef<bool> inputScalableVecDims, bool vectorizeNDExtract,
2541 bool flatten1DDepthwiseConv) {
2542
2543 if (!hasVectorizationImpl(op))
2544 return failure();
2545
2546 if (failed(vectorizeScalableVectorPrecondition(op, inputVectorSizes,
2547 inputScalableVecDims)))
2548 return failure();
2549
2551 .Case([&](linalg::LinalgOp linalgOp) {
2552 return vectorizeLinalgOpPrecondition(linalgOp, inputVectorSizes,
2553 vectorizeNDExtract,
2554 flatten1DDepthwiseConv);
2555 })
2556 .Case([&](tensor::PadOp padOp) {
2557 return vectorizePadOpPrecondition(padOp, inputVectorSizes);
2558 })
2559 .Case([&](linalg::PackOp packOp) {
2560 return vectorizePackOpPrecondition(packOp, inputVectorSizes);
2561 })
2562 .Case([&](linalg::UnPackOp unpackOp) {
2563 return vectorizeUnPackOpPrecondition(unpackOp, inputVectorSizes);
2564 })
2565 .Case([&](tensor::InsertSliceOp sliceOp) {
2566 return vectorizeInsertSliceOpPrecondition(sliceOp, inputVectorSizes);
2567 })
2568 .Default(failure());
2569}
2570
2571/// Converts affine.apply Ops to arithmetic operations.
2572static void convertAffineApply(RewriterBase &rewriter, LinalgOp linalgOp) {
2573 OpBuilder::InsertionGuard g(rewriter);
2574 auto toReplace = linalgOp.getBlock()->getOps<affine::AffineApplyOp>();
2575
2576 for (auto op : make_early_inc_range(toReplace)) {
2577 rewriter.setInsertionPoint(op);
2578 auto expanded = affine::expandAffineExpr(
2579 rewriter, op->getLoc(), op.getAffineMap().getResult(0),
2580 op.getOperands().take_front(op.getAffineMap().getNumDims()),
2581 op.getOperands().take_back(op.getAffineMap().getNumSymbols()));
2582 rewriter.replaceOp(op, expanded);
2583 }
2584}
2585
2586bool mlir::linalg::hasVectorizationImpl(Operation *op) {
2587 return isa<linalg::LinalgOp, tensor::PadOp, linalg::PackOp, linalg::UnPackOp,
2588 tensor::InsertSliceOp>(op);
2589}
2590
2591FailureOr<VectorizationResult> mlir::linalg::vectorize(
2592 RewriterBase &rewriter, Operation *op, ArrayRef<int64_t> inputVectorSizes,
2593 ArrayRef<bool> inputScalableVecDims, bool vectorizeNDExtract,
2594 bool flatten1DDepthwiseConv, bool assumeDynamicDimsMatchVecSizes,
2595 bool createNamedContraction) {
2596 LDBG() << "Attempting to vectorize: " << *op;
2597 LDBG() << "Input vector sizes: " << llvm::interleaved(inputVectorSizes);
2598 LDBG() << "Input scalable vector dims: "
2599 << llvm::interleaved(inputScalableVecDims);
2600
2601 if (failed(vectorizeOpPrecondition(op, inputVectorSizes, inputScalableVecDims,
2602 vectorizeNDExtract,
2603 flatten1DDepthwiseConv))) {
2604 LDBG() << "Vectorization pre-conditions failed";
2605 return failure();
2606 }
2607
2608 // Initialize vectorization state.
2609 VectorizationState state(rewriter);
2610 if (auto linalgOp = dyn_cast<linalg::LinalgOp>(op)) {
2611 if (failed(state.initState(rewriter, linalgOp, inputVectorSizes,
2612 inputScalableVecDims,
2613 assumeDynamicDimsMatchVecSizes))) {
2614 LDBG() << "Vectorization state couldn't be initialized";
2615 return failure();
2616 }
2617 }
2618
2619 SmallVector<Value> results;
2620 auto vectorizeResult =
2622 .Case([&](linalg::LinalgOp linalgOp) {
2623 // Check for both named as well as generic convolution ops.
2624 if (isaConvolutionOpInterface(linalgOp)) {
2625 FailureOr<Operation *> convOr = vectorizeConvolution(
2626 rewriter, linalgOp, inputVectorSizes, inputScalableVecDims,
2627 flatten1DDepthwiseConv);
2628 if (succeeded(convOr)) {
2629 llvm::append_range(results, (*convOr)->getResults());
2630 return success();
2631 }
2632
2633 LDBG() << "Unsupported convolution can't be vectorized.";
2634 return failure();
2635 }
2636
2637 if (createNamedContraction &&
2638 isa<ContractionOpInterface>(linalgOp.getOperation()))
2639 return vectorizeAsLinalgContraction(rewriter, state, linalgOp,
2640 results);
2641
2642 LDBG()
2643 << "Vectorize generic by broadcasting to the canonical vector "
2644 "shape";
2645
2646 // Pre-process before proceeding.
2647 convertAffineApply(rewriter, linalgOp);
2648
2649 // TODO: 'vectorize' takes in a 'RewriterBase' which is up-casted
2650 // to 'OpBuilder' when it is passed over to some methods like
2651 // 'vectorizeAsLinalgGeneric'. This is highly problematic: if we
2652 // erase an op within these methods, the actual rewriter won't be
2653 // notified and we will end up with read-after-free issues!
2654 return vectorizeAsLinalgGeneric(rewriter, state, linalgOp, results);
2655 })
2656 .Case([&](tensor::PadOp padOp) {
2657 return vectorizeAsTensorPadOp(rewriter, padOp, inputVectorSizes,
2658 results);
2659 })
2660 .Case([&](linalg::PackOp packOp) {
2661 return vectorizeAsTensorPackOp(rewriter, packOp, inputVectorSizes,
2662 results);
2663 })
2664 .Case([&](linalg::UnPackOp unpackOp) {
2665 return vectorizeAsTensorUnpackOp(rewriter, unpackOp,
2666 inputVectorSizes,
2667 inputScalableVecDims, results);
2668 })
2669 .Case([&](tensor::InsertSliceOp sliceOp) {
2670 return vectorizeAsInsertSliceOp(rewriter, sliceOp, inputVectorSizes,
2671 results);
2672 })
2673 .Default(failure());
2674
2675 if (failed(vectorizeResult)) {
2676 LDBG() << "Vectorization failed";
2677 return failure();
2678 }
2679
2680 return VectorizationResult{results};
2681}
2682
2683LogicalResult mlir::linalg::vectorizeCopy(RewriterBase &rewriter,
2684 memref::CopyOp copyOp) {
2685 auto srcType = cast<MemRefType>(copyOp.getSource().getType());
2686 auto dstType = cast<MemRefType>(copyOp.getTarget().getType());
2687 if (!srcType.hasStaticShape() || !dstType.hasStaticShape())
2688 return failure();
2689
2690 auto srcElementType = getElementTypeOrSelf(srcType);
2691 auto dstElementType = getElementTypeOrSelf(dstType);
2692 if (!VectorType::isValidElementType(srcElementType) ||
2693 !VectorType::isValidElementType(dstElementType))
2694 return failure();
2695
2696 auto readType = VectorType::get(srcType.getShape(), srcElementType);
2697 auto writeType = VectorType::get(dstType.getShape(), dstElementType);
2698
2699 Location loc = copyOp->getLoc();
2700 Value zero = arith::ConstantIndexOp::create(rewriter, loc, 0);
2701 SmallVector<Value> indices(srcType.getRank(), zero);
2702
2703 Value readValue = vector::TransferReadOp::create(
2704 rewriter, loc, readType, copyOp.getSource(), indices,
2705 /*padding=*/std::nullopt,
2706 rewriter.getMultiDimIdentityMap(srcType.getRank()));
2707 if (cast<VectorType>(readValue.getType()).getRank() == 0) {
2708 readValue = vector::ExtractOp::create(rewriter, loc, readValue,
2709 ArrayRef<int64_t>());
2710 readValue =
2711 vector::BroadcastOp::create(rewriter, loc, writeType, readValue);
2712 }
2713 Operation *writeValue = vector::TransferWriteOp::create(
2714 rewriter, loc, readValue, copyOp.getTarget(), indices,
2715 rewriter.getMultiDimIdentityMap(srcType.getRank()));
2716 rewriter.replaceOp(copyOp, writeValue->getResults());
2717 return success();
2718}
2719
2720//----------------------------------------------------------------------------//
2721// Misc. vectorization patterns.
2722//----------------------------------------------------------------------------//
2723/// Base pattern for rewriting tensor::PadOps whose result is consumed by a
2724/// given operation type OpTy.
2725template <typename OpTy>
2726struct VectorizePadOpUserPattern : public OpRewritePattern<tensor::PadOp> {
2727 using OpRewritePattern<tensor::PadOp>::OpRewritePattern;
2728
2729 LogicalResult matchAndRewrite(tensor::PadOp padOp,
2730 PatternRewriter &rewriter) const final {
2731 bool changed = false;
2732 // Insert users in vector, because some users may be replaced/removed.
2733 for (auto *user : llvm::to_vector<4>(padOp->getUsers()))
2734 if (auto op = dyn_cast<OpTy>(user))
2735 changed |= rewriteUser(rewriter, padOp, op).succeeded();
2736 return success(changed);
2737 }
2738
2739protected:
2740 virtual LogicalResult rewriteUser(PatternRewriter &rewriter,
2741 tensor::PadOp padOp, OpTy op) const = 0;
2742};
2743
2744/// Rewrite use of tensor::PadOp result in TransferReadOp. E.g.:
2745/// ```
2746/// %0 = tensor.pad %src ... : tensor<?x?xf32> to tensor<17x5xf32>
2747/// %r = vector.transfer_read %0[%c0, %c0], %cst
2748/// {in_bounds = [true, true]} : tensor<17x5xf32>, vector<17x5xf32>
2749/// ```
2750/// is rewritten to:
2751/// ```
2752/// %r = vector.transfer_read %src[%c0, %c0], %padding
2753/// {in_bounds = [true, true]}
2754/// : tensor<?x?xf32>, vector<17x5xf32>
2755/// ```
2756/// Note: By restricting this pattern to in-bounds TransferReadOps, we can be
2757/// sure that the original padding value %cst was never used.
2758///
2759/// This rewrite is possible if:
2760/// - `xferOp` has no out-of-bounds dims or mask.
2761/// - Low padding is static 0.
2762/// - Single, scalar padding value.
2763struct PadOpVectorizationWithTransferReadPattern
2764 : public VectorizePadOpUserPattern<vector::TransferReadOp> {
2765 using VectorizePadOpUserPattern<
2766 vector::TransferReadOp>::VectorizePadOpUserPattern;
2767
2768 LogicalResult rewriteUser(PatternRewriter &rewriter, tensor::PadOp padOp,
2769 vector::TransferReadOp xferOp) const override {
2770 // Low padding must be static 0.
2771 if (!padOp.hasZeroLowPad())
2772 return failure();
2773 // Pad value must be a constant.
2774 auto padValue = padOp.getConstantPaddingValue();
2775 if (!padValue)
2776 return failure();
2777 // Padding value of existing `xferOp` is unused.
2778 if (xferOp.hasOutOfBoundsDim() || xferOp.getMask())
2779 return failure();
2780
2781 rewriter.modifyOpInPlace(xferOp, [&]() {
2782 SmallVector<bool> inBounds(xferOp.getVectorType().getRank(), false);
2783 xferOp->setInherentAttr(xferOp.getInBoundsAttrName(),
2784 rewriter.getBoolArrayAttr(inBounds));
2785 xferOp.getBaseMutable().assign(padOp.getSource());
2786 xferOp.getPaddingMutable().assign(padValue);
2787 });
2788
2789 return success();
2790 }
2791};
2792
2793/// Rewrite use of tensor::PadOp result in TransferWriteOp.
2794/// This pattern rewrites TransferWriteOps that write to a padded tensor
2795/// value, where the same amount of padding is immediately removed again after
2796/// the write. In such cases, the TransferWriteOp can write to the non-padded
2797/// tensor value and apply out-of-bounds masking. E.g.:
2798/// ```
2799/// %0 = tensor.extract_slice ...[...] [%s0, %s1] [1, 1]
2800/// : tensor<...> to tensor<?x?xf32>
2801/// %1 = tensor.pad %0 ... : tensor<?x?xf32> to tensor<17x5xf32>
2802/// %2 = vector.transfer_write %vec, %1[...]
2803/// : vector<17x5xf32>, tensor<17x5xf32>
2804/// %r = tensor.extract_slice %2[0, 0] [%s0, %s1] [1, 1]
2805/// : tensor<17x5xf32> to tensor<?x?xf32>
2806/// ```
2807/// is rewritten to:
2808/// ```
2809/// %0 = tensor.extract_slice ...[...] [%s0, %s1] [1, 1]
2810/// : tensor<...> to tensor<?x?xf32>
2811/// %r = vector.transfer_write %vec, %0[...] : vector<17x5xf32>,
2812/// tensor<?x?xf32>
2813/// ```
2814/// Note: It is important that the ExtractSliceOp %r resizes the result of the
2815/// TransferWriteOp to the same size as the input of the TensorPadOp (or an
2816/// even smaller size). Otherwise, %r's new (dynamic) dimensions would differ
2817/// from %r's old dimensions.
2818///
2819/// This rewrite is possible if:
2820/// - Low padding is static 0.
2821/// - `xferOp` has exactly one use, which is an ExtractSliceOp. This
2822/// ExtractSliceOp trims the same amount of padding that was added
2823/// beforehand.
2824/// - Single, scalar padding value.
2825struct PadOpVectorizationWithTransferWritePattern
2826 : public VectorizePadOpUserPattern<vector::TransferWriteOp> {
2827 using VectorizePadOpUserPattern<
2828 vector::TransferWriteOp>::VectorizePadOpUserPattern;
2829
2830 LogicalResult rewriteUser(PatternRewriter &rewriter, tensor::PadOp padOp,
2831 vector::TransferWriteOp xferOp) const override {
2832 // TODO: support 0-d corner case.
2833 if (xferOp.getTransferRank() == 0)
2834 return failure();
2835
2836 // Low padding must be static 0.
2837 if (!padOp.hasZeroLowPad())
2838 return failure();
2839 // Pad value must be a constant.
2840 auto padValue = padOp.getConstantPaddingValue();
2841 if (!padValue)
2842 return failure();
2843 // TransferWriteOp result must be directly consumed by an ExtractSliceOp.
2844 if (!xferOp->hasOneUse())
2845 return failure();
2846 auto trimPadding = dyn_cast<tensor::ExtractSliceOp>(*xferOp->user_begin());
2847 if (!trimPadding)
2848 return failure();
2849 // Only static zero offsets supported when trimming padding.
2850 if (!trimPadding.hasZeroOffset())
2851 return failure();
2852 // trimPadding must remove the amount of padding that was added earlier.
2853 if (!hasSameTensorSize(padOp.getSource(), trimPadding))
2854 return failure();
2855
2856 // Insert the new TransferWriteOp at position of the old TransferWriteOp.
2857 rewriter.setInsertionPoint(xferOp);
2858
2859 SmallVector<bool> inBounds(xferOp.getVectorType().getRank(), false);
2860 auto newXferOp = rewriter.replaceOpWithNewOp<vector::TransferWriteOp>(
2861 xferOp, padOp.getSource().getType(), xferOp.getVector(),
2862 padOp.getSource(), xferOp.getIndices(), xferOp.getPermutationMapAttr(),
2863 xferOp.getMask(), rewriter.getBoolArrayAttr(inBounds));
2864 rewriter.replaceOp(trimPadding, newXferOp->getResult(0));
2865
2866 return success();
2867 }
2868
2869 /// Check if `beforePadding` and `afterTrimming` have the same tensor size,
2870 /// i.e., same dimensions.
2871 ///
2872 /// Dimensions may be static, dynamic or mix of both. In case of dynamic
2873 /// dimensions, this function tries to infer the (static) tensor size by
2874 /// looking at the defining op and utilizing op-specific knowledge.
2875 ///
2876 /// This is a conservative analysis. In case equal tensor sizes cannot be
2877 /// proven statically, this analysis returns `false` even though the tensor
2878 /// sizes may turn out to be equal at runtime.
2879 bool hasSameTensorSize(Value beforePadding,
2880 tensor::ExtractSliceOp afterTrimming) const {
2881 // If the input to tensor::PadOp is a CastOp, try with both CastOp
2882 // result and CastOp operand.
2883 if (auto castOp = beforePadding.getDefiningOp<tensor::CastOp>())
2884 if (hasSameTensorSize(castOp.getSource(), afterTrimming))
2885 return true;
2886
2887 auto t1 = dyn_cast<RankedTensorType>(beforePadding.getType());
2888 auto t2 = dyn_cast<RankedTensorType>(afterTrimming.getType());
2889 // Only RankedTensorType supported.
2890 if (!t1 || !t2)
2891 return false;
2892 // Rank of both values must be the same.
2893 if (t1.getRank() != t2.getRank())
2894 return false;
2895
2896 // All static dimensions must be the same. Mixed cases (e.g., dimension
2897 // static in `t1` but dynamic in `t2`) are not supported.
2898 for (unsigned i = 0; i < t1.getRank(); ++i) {
2899 if (t1.isDynamicDim(i) != t2.isDynamicDim(i))
2900 return false;
2901 if (!t1.isDynamicDim(i) && t1.getDimSize(i) != t2.getDimSize(i))
2902 return false;
2903 }
2904
2905 // Nothing more to check if all dimensions are static.
2906 if (t1.getNumDynamicDims() == 0)
2907 return true;
2908
2909 // All dynamic sizes must be the same. The only supported case at the
2910 // moment is when `beforePadding` is an ExtractSliceOp (or a cast
2911 // thereof).
2912
2913 // Apart from CastOp, only ExtractSliceOp is supported.
2914 auto beforeSlice = beforePadding.getDefiningOp<tensor::ExtractSliceOp>();
2915 if (!beforeSlice)
2916 return false;
2917
2918 assert(static_cast<size_t>(t1.getRank()) ==
2919 beforeSlice.getMixedSizes().size());
2920 assert(static_cast<size_t>(t2.getRank()) ==
2921 afterTrimming.getMixedSizes().size());
2922
2923 for (unsigned i = 0; i < t1.getRank(); ++i) {
2924 // Skip static dimensions.
2925 if (!t1.isDynamicDim(i))
2926 continue;
2927 auto size1 = beforeSlice.getMixedSizes()[i];
2928 auto size2 = afterTrimming.getMixedSizes()[i];
2929
2930 // Case 1: Same value or same constant int.
2931 if (isEqualConstantIntOrValue(size1, size2))
2932 continue;
2933
2934 // Other cases: Take a deeper look at defining ops of values.
2935 auto v1 = llvm::dyn_cast_if_present<Value>(size1);
2936 auto v2 = llvm::dyn_cast_if_present<Value>(size2);
2937 if (!v1 || !v2)
2938 return false;
2939
2940 // Case 2: Both values are identical AffineMinOps. (Should not happen if
2941 // CSE is run.)
2942 auto minOp1 = v1.getDefiningOp<affine::AffineMinOp>();
2943 auto minOp2 = v2.getDefiningOp<affine::AffineMinOp>();
2944 if (minOp1 && minOp2 && minOp1.getAffineMap() == minOp2.getAffineMap() &&
2945 minOp1.getOperands() == minOp2.getOperands())
2946 continue;
2947
2948 // Add additional cases as needed.
2949 }
2950
2951 // All tests passed.
2952 return true;
2953 }
2954};
2955
2956/// Returns the effective Pad value for the input op, provided it's a scalar.
2957///
2958/// Many Ops exhibit pad-like behaviour, but this isn't always explicit. If
2959/// this Op performs padding, retrieve the padding value provided that it's
2960/// a scalar and static/fixed for all the padded values. Returns an empty value
2961/// otherwise.
2962///
2963/// TODO: This is used twice (when checking vectorization pre-conditions and
2964/// when vectorizing). Cache results instead of re-running.
2965static Value getStaticPadVal(Operation *op) {
2966 if (!op)
2967 return {};
2968
2969 // 1. vector.broadcast (f32 -> vector <...xf32>) - return the value that's
2970 // being broadcast, provided that it's a scalar.
2971 if (auto bcast = llvm::dyn_cast<vector::BroadcastOp>(op)) {
2972 auto source = bcast.getSource();
2973 if (llvm::dyn_cast<VectorType>(source.getType()))
2974 return {};
2975
2976 return source;
2977 }
2978
2979 // 2. linalg.fill - use the scalar input value that used to fill the output
2980 // tensor.
2981 if (auto fill = llvm::dyn_cast<linalg::FillOp>(op)) {
2982 return fill.getInputs()[0];
2983 }
2984
2985 // 3. tensor.generateOp - can't guarantee the value is fixed without
2986 // analysing, bail out.
2987 if (auto generate = llvm::dyn_cast<tensor::GenerateOp>(op)) {
2988 return {};
2989 }
2990
2991 // 4. vector.transfer_write - inspect the input vector that's written from. If
2992 // if contains a single value that has been broadcast (e.g. via
2993 // vector.broadcast), extract it, fail otherwise.
2994 if (auto xferWrite = llvm::dyn_cast<vector::TransferWriteOp>(op))
2995 return getStaticPadVal(xferWrite.getVector().getDefiningOp());
2996
2997 // 5. tensor.insert_slice - inspect the destination tensor. If it's larger
2998 // than the input tensor, then, provided it's constant, we'll extract the
2999 // value that was used to generate it (via e.g. linalg.fill), fail otherwise.
3000 // TODO: Clarify the semantics when the input tensor is larger than the
3001 // destination.
3002 if (auto slice = llvm::dyn_cast<tensor::InsertSliceOp>(op))
3003 return getStaticPadVal(slice.getDest().getDefiningOp());
3004
3005 return {};
3006}
3007
3008static LogicalResult
3009vectorizeAsInsertSliceOp(RewriterBase &rewriter, tensor::InsertSliceOp sliceOp,
3010 ArrayRef<int64_t> inputVectorSizes,
3011 SmallVectorImpl<Value> &newResults) {
3012 // TODO: Introduce a parent class that will handle the insertion point update.
3013 OpBuilder::InsertionGuard g(rewriter);
3014 rewriter.setInsertionPoint(sliceOp);
3015
3016 TypedValue<RankedTensorType> source = sliceOp.getSource();
3017 auto sourceType = source.getType();
3018 auto resultType = sliceOp.getResultType();
3019
3020 Value padValue = getStaticPadVal(sliceOp);
3021
3022 if (!padValue) {
3023 auto elemType = sourceType.getElementType();
3024 padValue = arith::ConstantOp::create(rewriter, sliceOp.getLoc(), elemType,
3025 rewriter.getZeroAttr(elemType));
3026 }
3027
3028 // 2. Get the vector shape
3029 // Map each source dim to its corresponding (non-dropped) result dim: for a
3030 // rank-reducing slice, dropped dims need not be the trailing ones.
3031 llvm::SmallBitVector droppedDims = sliceOp.getDroppedDims();
3032 SmallVector<int64_t> resultDimsForSourceDims;
3033 resultDimsForSourceDims.reserve(sourceType.getRank());
3034 for (int64_t resultDim = 0, end = resultType.getRank(); resultDim < end;
3035 ++resultDim)
3036 if (!droppedDims[resultDim])
3037 resultDimsForSourceDims.push_back(resultDim);
3038 assert(resultDimsForSourceDims.size() ==
3039 static_cast<size_t>(sourceType.getRank()) &&
3040 "expected one non-dropped result dim per source dim");
3041
3042 SmallVector<int64_t> vecShape;
3043 for (int64_t i = 0, end = sourceType.getRank(); i < end; ++i) {
3044 if (!inputVectorSizes.empty()) {
3045 vecShape.push_back(inputVectorSizes[i]);
3046 } else if (!sourceType.isDynamicDim(i)) {
3047 vecShape.push_back(sourceType.getDimSize(i));
3048 } else if (!resultType.isDynamicDim(resultDimsForSourceDims[i])) {
3049 // Source shape is not statically known, but result shape is.
3050 // Vectorize with size of result shape. This may be larger than the
3051 // source size.
3052 vecShape.push_back(resultType.getDimSize(resultDimsForSourceDims[i]));
3053 } else {
3054 // Neither source nor result dim of padOp is static. Cannot vectorize
3055 // the copy.
3056 return failure();
3057 }
3058 }
3059 auto vecType = VectorType::get(vecShape, sourceType.getElementType());
3060
3061 // 3. Generate TransferReadOp + TransferWriteOp
3062 auto loc = sliceOp.getLoc();
3063
3064 // Create read
3065 SmallVector<Value> readIndices(
3066 vecType.getRank(), arith::ConstantIndexOp::create(rewriter, loc, 0));
3068 rewriter, loc, source, vecType, padValue,
3069 /*useInBoundsInsteadOfMasking=*/inputVectorSizes.empty());
3070
3071 // Create write
3072 auto writeIndices =
3073 getValueOrCreateConstantIndexOp(rewriter, loc, sliceOp.getMixedOffsets());
3074 Operation *write =
3075 vector::createWriteOrMaskedWrite(rewriter, loc, read, sliceOp.getDest(),
3076 writeIndices, inputVectorSizes.empty());
3077
3078 // 4. Finalize
3079 newResults.push_back(write->getResult(0));
3080
3081 return success();
3082}
3083
3084/// Rewrite use of tensor::PadOp result in InsertSliceOp. E.g.:
3085/// ```
3086/// %0 = tensor.pad %src ... : tensor<?x?xf32> to tensor<17x5xf32>
3087/// %r = tensor.insert_slice %0
3088/// into %dest[%a, %b, 0, 0] [1, 1, 17, 5] [1, 1, 1, 1]
3089/// : tensor<17x5xf32> into tensor<?x?x17x5xf32>
3090/// ```
3091/// is rewritten to:
3092/// ```
3093/// %0 = vector.transfer_read %src[%c0, %c0], %padding
3094/// : tensor<?x?xf32>, vector<17x5xf32>
3095/// %r = vector.transfer_write %0, %dest[%a, %b, %c0, %c0]
3096/// {in_bounds = [true, true]} : vector<17x5xf32>, tensor<?x?x17x5xf32>
3097/// ```
3098///
3099/// This rewrite is possible if:
3100/// - Low padding is static 0.
3101/// - `padOp` result shape is static.
3102/// - The entire padded tensor is inserted.
3103/// (Implies that sizes of `insertOp` are all static.)
3104/// - Only unit strides in `insertOp`.
3105/// - Single, scalar padding value.
3106/// - `padOp` result not used as destination.
3107struct PadOpVectorizationWithInsertSlicePattern
3108 : public VectorizePadOpUserPattern<tensor::InsertSliceOp> {
3109 using VectorizePadOpUserPattern<
3110 tensor::InsertSliceOp>::VectorizePadOpUserPattern;
3111
3112 LogicalResult rewriteUser(PatternRewriter &rewriter, tensor::PadOp padOp,
3113 tensor::InsertSliceOp insertOp) const override {
3114 // Low padding must be static 0.
3115 if (!padOp.hasZeroLowPad())
3116 return failure();
3117 // Only unit stride supported.
3118 if (!insertOp.hasUnitStride())
3119 return failure();
3120 // Pad value must be a constant.
3121 auto padValue = padOp.getConstantPaddingValue();
3122 if (!padValue)
3123 return failure();
3124 // Dynamic shapes not supported.
3125 if (!cast<ShapedType>(padOp.getResult().getType()).hasStaticShape())
3126 return failure();
3127 // Pad result not used as destination.
3128 if (insertOp.getDest() == padOp.getResult())
3129 return failure();
3130
3131 auto vecType = VectorType::get(padOp.getType().getShape(),
3132 padOp.getType().getElementType());
3133 unsigned vecRank = vecType.getRank();
3134 unsigned tensorRank = insertOp.getType().getRank();
3135
3136 // Check if sizes match: Insert the entire tensor into most minor dims.
3137 // (No permutations allowed.)
3138 SmallVector<int64_t> expectedSizes(tensorRank - vecRank, 1);
3139 expectedSizes.append(vecType.getShape().begin(), vecType.getShape().end());
3140 if (!llvm::all_of(
3141 llvm::zip(insertOp.getMixedSizes(), expectedSizes), [](auto it) {
3142 return getConstantIntValue(std::get<0>(it)) == std::get<1>(it);
3143 }))
3144 return failure();
3145
3146 // Insert the TransferReadOp and TransferWriteOp at the position of the
3147 // InsertSliceOp.
3148 rewriter.setInsertionPoint(insertOp);
3149
3150 // Generate TransferReadOp: Read entire source tensor and add high
3151 // padding.
3152 SmallVector<Value> readIndices(
3153 vecRank, arith::ConstantIndexOp::create(rewriter, padOp.getLoc(), 0));
3154 auto read = vector::TransferReadOp::create(rewriter, padOp.getLoc(),
3155 vecType, padOp.getSource(),
3156 readIndices, padValue);
3157
3158 // Generate TransferWriteOp: Write to InsertSliceOp's dest tensor at
3159 // specified offsets. Write is fully in-bounds because a InsertSliceOp's
3160 // source must fit into the destination at the specified offsets.
3161 auto writeIndices = getValueOrCreateConstantIndexOp(
3162 rewriter, padOp.getLoc(), insertOp.getMixedOffsets());
3163 SmallVector<bool> inBounds(vecRank, true);
3164 rewriter.replaceOpWithNewOp<vector::TransferWriteOp>(
3165 insertOp, read, insertOp.getDest(), writeIndices,
3166 ArrayRef<bool>{inBounds});
3167
3168 return success();
3169 }
3170};
3171
3173 RewritePatternSet &patterns, PatternBenefit baseBenefit) {
3174 patterns.add<PadOpVectorizationWithTransferReadPattern,
3175 PadOpVectorizationWithTransferWritePattern,
3176 PadOpVectorizationWithInsertSlicePattern>(
3177 patterns.getContext(), baseBenefit.getBenefit() + 1);
3178}
3179
3180//----------------------------------------------------------------------------//
3181// Forwarding patterns
3182//----------------------------------------------------------------------------//
3183
3184/// Check whether there is any interleaved use of any `values` between
3185/// `firstOp` and `secondOp`. Conservatively return `true` if any op or value
3186/// is in a different block.
3187static bool mayExistInterleavedUses(Operation *firstOp, Operation *secondOp,
3188 ValueRange values) {
3189 if (firstOp->getBlock() != secondOp->getBlock() ||
3190 !firstOp->isBeforeInBlock(secondOp)) {
3191 LDBG() << "interleavedUses precondition failed, firstOp: " << *firstOp
3192 << ", second op: " << *secondOp;
3193 return true;
3194 }
3195 for (auto v : values) {
3196 for (auto &u : v.getUses()) {
3197 Operation *owner = u.getOwner();
3198 if (owner == firstOp || owner == secondOp)
3199 continue;
3200 // TODO: this is too conservative, use dominance info in the future.
3201 if (owner->getBlock() == firstOp->getBlock() &&
3202 (owner->isBeforeInBlock(firstOp) || secondOp->isBeforeInBlock(owner)))
3203 continue;
3204 LDBG() << " found interleaved op " << *owner << ", firstOp: " << *firstOp
3205 << ", second op: " << *secondOp;
3206 return true;
3207 }
3208 }
3209 return false;
3210}
3211
3212/// Return the unique subview use of `v` if it is indeed unique, null
3213/// otherwise.
3214static memref::SubViewOp getSubViewUseIfUnique(Value v) {
3215 memref::SubViewOp subViewOp;
3216 for (auto &u : v.getUses()) {
3217 if (auto newSubViewOp = dyn_cast<memref::SubViewOp>(u.getOwner())) {
3218 if (subViewOp)
3219 return memref::SubViewOp();
3220 subViewOp = newSubViewOp;
3221 }
3222 }
3223 return subViewOp;
3224}
3225
3226/// TODO: use interfaces, side-effects and aliasing analysis as appropriate,
3227/// when available.
3229 vector::TransferReadOp xferOp, PatternRewriter &rewriter) const {
3230
3231 // TODO: support mask.
3232 if (xferOp.getMask())
3233 return rewriter.notifyMatchFailure(xferOp, "unsupported mask");
3234
3235 // Transfer into `view`.
3236 Value viewOrAlloc = xferOp.getBase();
3237 if (!viewOrAlloc.getDefiningOp<memref::ViewOp>() &&
3238 !viewOrAlloc.getDefiningOp<memref::AllocOp>())
3239 return rewriter.notifyMatchFailure(xferOp, "source not a view or alloc");
3240
3241 // Ensure there is exactly one subview of `viewOrAlloc` defining `subView`.
3242 memref::SubViewOp subViewOp = getSubViewUseIfUnique(viewOrAlloc);
3243 if (!subViewOp)
3244 return rewriter.notifyMatchFailure(xferOp, "no subview found");
3245 Value subView = subViewOp.getResult();
3246
3247 // Find the copy into `subView` without interleaved uses.
3248 memref::CopyOp copyOp;
3249 for (auto &u : subView.getUses()) {
3250 if (auto newCopyOp = dyn_cast<memref::CopyOp>(u.getOwner())) {
3251 assert(isa<MemRefType>(newCopyOp.getTarget().getType()));
3252 if (newCopyOp.getTarget() != subView)
3253 continue;
3254 if (mayExistInterleavedUses(newCopyOp, xferOp, {viewOrAlloc, subView}))
3255 continue;
3256 copyOp = newCopyOp;
3257 break;
3258 }
3259 }
3260 if (!copyOp)
3261 return rewriter.notifyMatchFailure(xferOp, "no copy found");
3262
3263 // Find the fill into `viewOrAlloc` without interleaved uses before the
3264 // copy.
3265 FillOp maybeFillOp;
3266 for (auto &u : viewOrAlloc.getUses()) {
3267 if (auto newFillOp = dyn_cast<FillOp>(u.getOwner())) {
3268 assert(isa<MemRefType>(newFillOp.output().getType()));
3269 if (newFillOp.output() != viewOrAlloc)
3270 continue;
3271 if (mayExistInterleavedUses(newFillOp, copyOp, {viewOrAlloc, subView}))
3272 continue;
3273 maybeFillOp = newFillOp;
3274 break;
3275 }
3276 }
3277 // Ensure padding matches.
3278 if (maybeFillOp && xferOp.getPadding() != maybeFillOp.value())
3279 return rewriter.notifyMatchFailure(xferOp,
3280 "padding value does not match fill");
3281
3282 // `in` is the subview that memref.copy reads. Replace it.
3283 Value in = copyOp.getSource();
3284
3285 // memref.copy + linalg.fill can be used to create a padded local buffer.
3286 // The `masked` attribute is only valid on this padded buffer.
3287 // When forwarding to vector.transfer_read, the attribute must be reset
3288 // conservatively.
3289 auto vectorType = xferOp.getVectorType();
3290 Value res = vector::TransferReadOp::create(
3291 rewriter, xferOp.getLoc(), vectorType, in, xferOp.getIndices(),
3292 xferOp.getPermutationMapAttr(), xferOp.getPadding(), xferOp.getMask(),
3293 rewriter.getBoolArrayAttr(
3294 SmallVector<bool>(vectorType.getRank(), false)));
3295
3296 if (maybeFillOp)
3297 rewriter.eraseOp(maybeFillOp);
3298 rewriter.eraseOp(copyOp);
3299 rewriter.replaceOp(xferOp, res);
3300
3301 return success();
3302}
3303
3304/// TODO: use interfaces, side-effects and aliasing analysis as appropriate,
3305/// when available.
3307 vector::TransferWriteOp xferOp, PatternRewriter &rewriter) const {
3308 // TODO: support mask.
3309 if (xferOp.getMask())
3310 return rewriter.notifyMatchFailure(xferOp, "unsupported mask");
3311
3312 // Transfer into `viewOrAlloc`.
3313 Value viewOrAlloc = xferOp.getBase();
3314 if (!viewOrAlloc.getDefiningOp<memref::ViewOp>() &&
3315 !viewOrAlloc.getDefiningOp<memref::AllocOp>())
3316 return rewriter.notifyMatchFailure(xferOp, "source not a view or alloc");
3317
3318 // Ensure there is exactly one subview of `viewOrAlloc` defining `subView`.
3319 memref::SubViewOp subViewOp = getSubViewUseIfUnique(viewOrAlloc);
3320 if (!subViewOp)
3321 return rewriter.notifyMatchFailure(xferOp, "no subview found");
3322 Value subView = subViewOp.getResult();
3323
3324 // Find the copy from `subView` without interleaved uses.
3325 memref::CopyOp copyOp;
3326 for (auto &u : subViewOp.getResult().getUses()) {
3327 if (auto newCopyOp = dyn_cast<memref::CopyOp>(u.getOwner())) {
3328 if (newCopyOp.getSource() != subView)
3329 continue;
3330 if (mayExistInterleavedUses(xferOp, newCopyOp, {viewOrAlloc, subView}))
3331 continue;
3332 copyOp = newCopyOp;
3333 break;
3334 }
3335 }
3336 if (!copyOp)
3337 return rewriter.notifyMatchFailure(xferOp, "no copy found");
3338
3339 // `out` is the subview copied into that we replace.
3340 assert(isa<MemRefType>(copyOp.getTarget().getType()));
3341 Value out = copyOp.getTarget();
3342
3343 // Forward vector.transfer into copy.
3344 // memref.copy + linalg.fill can be used to create a padded local buffer.
3345 // The `masked` attribute is only valid on this padded buffer.
3346 // When forwarding to vector.transfer_write, the attribute must be reset
3347 // conservatively.
3348 auto vector = xferOp.getVector();
3349 vector::TransferWriteOp::create(
3350 rewriter, xferOp.getLoc(), vector, out, xferOp.getIndices(),
3351 xferOp.getPermutationMapAttr(), xferOp.getMask(),
3352 rewriter.getBoolArrayAttr(SmallVector<bool>(
3353 dyn_cast<VectorType>(vector.getType()).getRank(), false)));
3354
3355 rewriter.eraseOp(copyOp);
3356 rewriter.eraseOp(xferOp);
3357
3358 return success();
3359}
3360
3361//===----------------------------------------------------------------------===//
3362// Convolution vectorization patterns
3363//===----------------------------------------------------------------------===//
3364
3365template <int N>
3366static void bindShapeDims(ShapedType shapedType) {}
3367
3368template <int N, typename IntTy, typename... IntTy2>
3369static void bindShapeDims(ShapedType shapedType, IntTy &val, IntTy2 &...vals) {
3370 val = shapedType.getShape()[N];
3371 bindShapeDims<N + 1, IntTy2 &...>(shapedType, vals...);
3372}
3373
3374/// Bind a pack of int& to the leading dimensions of shapedType.getShape().
3375template <typename... IntTy>
3376static void bindShapeDims(ShapedType shapedType, IntTy &...vals) {
3377 bindShapeDims<0>(shapedType, vals...);
3378}
3379
3380/// Match 1D convolution or pooling operations and return their dilations and
3381/// strides. Returns std::nullopt for unrecognized ops.
3382static std::optional<DilationsAndStrides> match1DConvPoolOp(LinalgOp op) {
3383#define MATCH_1D_CONV_POOL_OP(ConvOpTy) \
3384 if (auto convParams = matchConvolutionOpOfType<ConvOpTy>(op)) \
3385 return convParams;
3386
3387 // 1D Convolution ops.
3388 MATCH_1D_CONV_POOL_OP(linalg::Conv1DOp);
3389 MATCH_1D_CONV_POOL_OP(linalg::Conv1DNwcWcfOp);
3390 MATCH_1D_CONV_POOL_OP(linalg::Conv1DNcwFcwOp);
3391 // Depthwise 1D Convolution ops.
3392 // Note: Only NWC layout without channel multiplier is supported.
3393 // DepthwiseConv1DNcwCwOp (NCW) and DepthwiseConv1DNwcWcmOp (with multiplier)
3394 // are not supported.
3395 MATCH_1D_CONV_POOL_OP(linalg::DepthwiseConv1DNwcWcOp);
3396 // 1D Pooling ops (NWC layout).
3397 MATCH_1D_CONV_POOL_OP(linalg::PoolingNwcSumOp);
3398 MATCH_1D_CONV_POOL_OP(linalg::PoolingNwcMaxOp);
3399 MATCH_1D_CONV_POOL_OP(linalg::PoolingNwcMaxUnsignedOp);
3400 MATCH_1D_CONV_POOL_OP(linalg::PoolingNwcMinOp);
3401 MATCH_1D_CONV_POOL_OP(linalg::PoolingNwcMinUnsignedOp);
3402 // 1D Pooling ops (NCW layout).
3403 MATCH_1D_CONV_POOL_OP(linalg::PoolingNcwSumOp);
3404 MATCH_1D_CONV_POOL_OP(linalg::PoolingNcwMaxOp);
3405
3406#undef MATCH_1D_CONV_POOL_OP
3407
3408 return std::nullopt;
3409}
3410
3411namespace {
3412/// Generate a vector implementation for either:
3413/// ```
3414/// Op def: ( w, kw )
3415/// Iters: ({Par(), Red()})
3416/// Layout: {{w + kw}, {kw}, {w}}
3417/// ```
3418/// kw is unrolled.
3419///
3420/// or
3421///
3422/// ```
3423/// Op def: ( n, w, c, kw, f )
3424/// Iters: ({Par(), Par(), Par(), Red(), Red()})
3425/// Layout: {{n, strideW * w + dilationW * kw, c}, {kw, c, f}, {n, w, f}}
3426/// ```
3427/// kw is unrolled, w is unrolled iff dilationW > 1.
3428///
3429/// or
3430///
3431/// ```
3432/// Op def: ( n, c, w, f, kw )
3433/// Iters: ({Par(), Par(), Par(), Red(), Red()})
3434/// Layout: {{n, c, strideW * w + dilationW * kw}, {f, c, kw}, {n, f, w}}
3435/// ```
3436/// kw is unrolled, w is unrolled iff dilationW > 1.
3437///
3438/// or
3439///
3440/// ```
3441/// Op def: ( n, w, c, kw )
3442/// Iters: ({Par(), Par(), Par(), Red()})
3443/// Layout: {{n, strideW * w + dilationW * kw, c}, {kw, c}, {n, w, c}}
3444/// ```
3445/// kw is unrolled, w is unrolled iff dilationW > 1.
3446struct Conv1DGenerator
3447 : public StructuredGenerator<LinalgOp, utils::IteratorType> {
3448 /// Factory method to create a Conv1DGenerator. Returns failure if the
3449 /// operation doesn't have valid strides/dilations.
3450 static FailureOr<Conv1DGenerator> create(RewriterBase &rewriter,
3451 LinalgOp linalgOp) {
3452 // Try to match a 1D conv/pool op using matchConvolutionOpOfType. This
3453 // works for both named ops and generic ops that match their semantics.
3454 std::optional<DilationsAndStrides> convParams = match1DConvPoolOp(linalgOp);
3455 if (!convParams)
3456 return failure();
3457
3458 int strideW = static_cast<int>(convParams->strides.front());
3459 int dilationW = static_cast<int>(convParams->dilations.front());
3460 return Conv1DGenerator(rewriter, linalgOp, strideW, dilationW);
3461 }
3462
3463private:
3464 Conv1DGenerator(RewriterBase &rewriter, LinalgOp linalgOp, int strideW,
3465 int dilationW)
3466 : StructuredGenerator<LinalgOp, utils::IteratorType>(rewriter, linalgOp),
3467 strideW(strideW), dilationW(dilationW) {
3468
3469 lhsShaped = linalgOp.getDpsInputOperand(0)->get();
3470 rhsShaped = linalgOp.getDpsInputOperand(1)->get();
3471 resShaped = linalgOp.getDpsInitOperand(0)->get();
3472 lhsShapedType = dyn_cast<ShapedType>(lhsShaped.getType());
3473 rhsShapedType = dyn_cast<ShapedType>(rhsShaped.getType());
3474 resShapedType = dyn_cast<ShapedType>(resShaped.getType());
3475
3476 Operation *reduceOp = matchLinalgReduction(linalgOp.getDpsInitOperand(0));
3477 redOp = reduceOp->getName().getIdentifier();
3478
3479 setConvOperationKind(reduceOp);
3480
3481 auto maybeKind = getCombinerOpKind(reduceOp);
3482 reductionKind = maybeKind.value();
3483 }
3484
3485public:
3486 /// Generate a vector implementation for:
3487 /// ```
3488 /// Op def: ( w, kw )
3489 /// Iters: ({Par(), Red()})
3490 /// Layout: {{w + kw}, {kw}, {w}}
3491 /// ```
3492 /// kw is always unrolled.
3493 ///
3494 /// or
3495 ///
3496 /// ```
3497 /// Op def: ( n, w, c, kw, f )
3498 /// Iters: ({Par(), Par(), Par(), Red(), Red()})
3499 /// Layout: {{n, strideW * w + dilationW * kw, c}, {kw, c, f}, {n, w, f}}
3500 /// ```
3501 /// kw is always unrolled.
3502 /// TODO: w (resp. kw) is unrolled when the strideW ( resp. dilationW) is
3503 /// > 1.
3504 FailureOr<Operation *> conv(Conv1DOpOrder conv1DOpOrder) {
3505 int64_t nSize, wSize, cSize, kwSize, fSize;
3506 SmallVector<int64_t, 3> lhsShape, rhsShape, resShape;
3507 bool isSingleChanneled = (conv1DOpOrder == Conv1DOpOrder::W);
3508 switch (conv1DOpOrder) {
3509 case Conv1DOpOrder::W:
3510 // Initialize unused dimensions
3511 nSize = fSize = cSize = 0;
3512 // out{W}
3513 bindShapeDims(resShapedType, wSize);
3514 // kernel{kw}
3515 bindShapeDims(rhsShapedType, kwSize);
3516 lhsShape = {// iw = ow + kw - 1
3517 // (i.e. 16 convolved with 3 -> 14)
3518 (wSize + kwSize - 1)};
3519 rhsShape = {kwSize};
3520 resShape = {wSize};
3521 break;
3522 case Conv1DOpOrder::Nwc:
3523 // out{n, w, f}
3524 bindShapeDims(resShapedType, nSize, wSize, fSize);
3525 switch (oper) {
3526 case ConvOperationKind::Conv:
3527 // kernel{kw, c, f}
3528 bindShapeDims(rhsShapedType, kwSize, cSize);
3529 break;
3530 case ConvOperationKind::Pool:
3531 // kernel{kw}
3532 bindShapeDims(rhsShapedType, kwSize);
3533 cSize = fSize;
3534 break;
3535 }
3536 lhsShape = {nSize,
3537 // iw = ow * sw + kw * dw - 1
3538 // (i.e. 16 convolved with 3 (@stride 1 dilation 1) -> 14)
3539 // Perform the proper inclusive -> exclusive -> inclusive.
3540 ((wSize - 1) * strideW + 1) + ((kwSize - 1) * dilationW + 1) -
3541 1,
3542 cSize};
3543 switch (oper) {
3544 case ConvOperationKind::Conv:
3545 rhsShape = {kwSize, cSize, fSize};
3546 break;
3547 case ConvOperationKind::Pool:
3548 rhsShape = {kwSize};
3549 break;
3550 }
3551 resShape = {nSize, wSize, fSize};
3552 break;
3553 case Conv1DOpOrder::Ncw:
3554 // out{n, f, w}
3555 bindShapeDims(resShapedType, nSize, fSize, wSize);
3556 switch (oper) {
3557 case ConvOperationKind::Conv:
3558 // kernel{f, c, kw}
3559 bindShapeDims(rhsShapedType, fSize, cSize, kwSize);
3560 break;
3561 case ConvOperationKind::Pool:
3562 // kernel{kw}
3563 bindShapeDims(rhsShapedType, kwSize);
3564 cSize = fSize;
3565 break;
3566 }
3567 lhsShape = {nSize, cSize,
3568 // iw = ow * sw + kw * dw - 1
3569 // (i.e. 16 convolved with 3 (@stride 1 dilation 1) -> 14)
3570 // Perform the proper inclusive -> exclusive -> inclusive.
3571 ((wSize - 1) * strideW + 1) + ((kwSize - 1) * dilationW + 1) -
3572 1};
3573 switch (oper) {
3574 case ConvOperationKind::Conv:
3575 rhsShape = {fSize, cSize, kwSize};
3576 break;
3577 case ConvOperationKind::Pool:
3578 rhsShape = {kwSize};
3579 break;
3580 }
3581 resShape = {nSize, fSize, wSize};
3582 break;
3583 }
3584
3585 vector::TransferWriteOp write;
3586 Value zero = arith::ConstantIndexOp::create(rewriter, loc, 0);
3587
3588 // w is unrolled (i.e. wSizeStep == 1) iff strideW > 1.
3589 // When strideW == 1, we can batch the contiguous loads and avoid
3590 // unrolling
3591 int64_t wSizeStep = strideW == 1 ? wSize : 1;
3592
3593 Type lhsEltType = lhsShapedType.getElementType();
3594 Type rhsEltType = rhsShapedType.getElementType();
3595 Type resEltType = resShapedType.getElementType();
3596 auto lhsType = VectorType::get(lhsShape, lhsEltType);
3597 auto rhsType = VectorType::get(rhsShape, rhsEltType);
3598 auto resType = VectorType::get(resShape, resEltType);
3599 // Zero padding with the corresponding dimensions for lhs, rhs and res.
3600 SmallVector<Value> lhsPadding(lhsShape.size(), zero);
3601 SmallVector<Value> rhsPadding(rhsShape.size(), zero);
3602 SmallVector<Value> resPadding(resShape.size(), zero);
3603
3604 // Read the whole lhs, rhs and res in one shot (with zero padding).
3605 Value lhs = vector::TransferReadOp::create(
3606 rewriter, loc, lhsType, lhsShaped, lhsPadding,
3607 /*padding=*/arith::getZeroConstant(rewriter, loc, lhsEltType));
3608 // This is needed only for Conv.
3609 Value rhs = nullptr;
3610 if (oper == ConvOperationKind::Conv)
3611 rhs = vector::TransferReadOp::create(
3612 rewriter, loc, rhsType, rhsShaped, rhsPadding,
3613 /*padding=*/arith::getZeroConstant(rewriter, loc, rhsEltType));
3614 Value res = vector::TransferReadOp::create(
3615 rewriter, loc, resType, resShaped, resPadding,
3616 /*padding=*/arith::getZeroConstant(rewriter, loc, resEltType));
3617
3618 // The base vectorization case for channeled convolution is input:
3619 // {n,w,c}, weight: {kw,c,f}, output: {n,w,f}. To reuse the base pattern
3620 // vectorization case, we do pre transpose on input, weight, and output.
3621 switch (conv1DOpOrder) {
3622 case Conv1DOpOrder::W:
3623 case Conv1DOpOrder::Nwc:
3624 // Base case, so no transposes necessary.
3625 break;
3626 case Conv1DOpOrder::Ncw: {
3627 // To match base vectorization case, we pre-transpose current case.
3628 // ncw -> nwc
3629 static constexpr std::array<int64_t, 3> permLhs = {0, 2, 1};
3630 lhs = vector::TransposeOp::create(rewriter, loc, lhs, permLhs);
3631 // fcw -> wcf
3632 static constexpr std::array<int64_t, 3> permRhs = {2, 1, 0};
3633
3634 // This is needed only for Conv.
3635 if (oper == ConvOperationKind::Conv)
3636 rhs = vector::TransposeOp::create(rewriter, loc, rhs, permRhs);
3637 // nfw -> nwf
3638 static constexpr std::array<int64_t, 3> permRes = {0, 2, 1};
3639 res = vector::TransposeOp::create(rewriter, loc, res, permRes);
3640 break;
3641 }
3642 }
3643
3644 //===------------------------------------------------------------------===//
3645 // Begin vector-only rewrite part
3646 //===------------------------------------------------------------------===//
3647 // Unroll along kw and read slices of lhs and rhs.
3648 SmallVector<Value> lhsVals, rhsVals, resVals;
3649 lhsVals = extractConvInputSlices(rewriter, loc, lhs, nSize, wSize, cSize,
3650 kwSize, strideW, dilationW, wSizeStep,
3651 isSingleChanneled);
3652 // Do not do for pooling.
3653 if (oper == ConvOperationKind::Conv)
3654 rhsVals = extractConvFilterSlices(rewriter, loc, rhs, kwSize);
3655 resVals = extractConvResultSlices(rewriter, loc, res, nSize, wSize, fSize,
3656 wSizeStep, isSingleChanneled);
3657
3658 auto linearIndex = [&](int64_t kw, int64_t w) {
3659 return kw * (wSize / wSizeStep) + w;
3660 };
3661
3662 // Compute contraction: O{n, w, f} += I{n, sw * w + dw * kw, c} * F{c, f}
3663 // or perform outerproduct for non-channeled convolution or perform simple
3664 // arith operation for pooling
3665 for (int64_t kw = 0; kw < kwSize; ++kw) {
3666 for (int64_t w = 0; w < wSize; w += wSizeStep) {
3667 switch (oper) {
3668 case ConvOperationKind::Conv:
3669 if (isSingleChanneled) {
3670 resVals[w] = conv1dSliceAsOuterProduct(rewriter, loc,
3671 lhsVals[linearIndex(kw, w)],
3672 rhsVals[kw], resVals[w]);
3673 } else {
3674 resVals[w] = conv1dSliceAsContraction(rewriter, loc,
3675 lhsVals[linearIndex(kw, w)],
3676 rhsVals[kw], resVals[w]);
3677 }
3678 break;
3679 case ConvOperationKind::Pool:
3680 resVals[w] = pool1dSlice(rewriter, loc, lhsVals[linearIndex(kw, w)],
3681 resVals[w]);
3682 break;
3683 }
3684 }
3685 }
3686
3687 res = insertConvResultSlices(rewriter, loc, res, wSize, wSizeStep, resVals,
3688 isSingleChanneled);
3689 //===------------------------------------------------------------------===//
3690 // End vector-only rewrite part
3691 //===------------------------------------------------------------------===//
3692
3693 // The base vectorization case for channeled convolution is output:
3694 // {n,w,f} To reuse the result from base pattern vectorization case, we
3695 // post transpose the base case result.
3696 switch (conv1DOpOrder) {
3697 case Conv1DOpOrder::W:
3698 case Conv1DOpOrder::Nwc:
3699 // Base case, so no transposes necessary.
3700 break;
3701 case Conv1DOpOrder::Ncw: {
3702 // nwf -> nfw
3703 static constexpr std::array<int64_t, 3> perm = {0, 2, 1};
3704 res = vector::TransposeOp::create(rewriter, loc, res, perm);
3705 break;
3706 }
3707 }
3708
3709 return vector::TransferWriteOp::create(rewriter, loc, res, resShaped,
3710 resPadding)
3711 .getOperation();
3712 }
3713
3714 // Promote `val` to the element type of `ty` using `castOp`.
3715 Value promote(RewriterBase &rewriter, Location loc, Value val, Type ty,
3716 Operation *castOp) {
3717 const Type dstElementType = getElementTypeOrSelf(ty);
3718 if (getElementTypeOrSelf(val.getType()) == dstElementType)
3719 return val;
3720
3721 assert(castOp && "expected a payload cast for promoted operand");
3722
3723 // Handle both shaped as well as scalar types.
3724 Type dstType;
3725 if (auto shapedType = dyn_cast<ShapedType>(val.getType()))
3726 dstType = shapedType.cloneWith(std::nullopt, dstElementType);
3727 else
3728 dstType = dstElementType;
3729
3730 OperationState state(loc, castOp->getName().getIdentifier(), val, dstType,
3731 castOp->getDiscardableAttrDictionary().getValue());
3732 state.propertiesAttr = castOp->getPropertiesAsAttribute();
3733 return rewriter.create(state)->getResult(0);
3734 }
3735
3736 // Create a contraction: lhs{n, w, c} * rhs{c, f} -> res{n, w, f}
3737 Value conv1dSliceAsContraction(RewriterBase &rewriter, Location loc,
3738 Value lhs, Value rhs, Value res) {
3739 vector::IteratorType par = vector::IteratorType::parallel;
3740 vector::IteratorType red = vector::IteratorType::reduction;
3741 AffineExpr n, w, f, c;
3742 bindDims(ctx, n, w, f, c);
3743 lhs = promote(rewriter, loc, lhs, res.getType(), lhsCastOp);
3744 rhs = promote(rewriter, loc, rhs, res.getType(), rhsCastOp);
3745 auto contrationOp = vector::ContractionOp::create(
3746 rewriter, loc, lhs, rhs, res,
3747 /*indexingMaps=*/MapList{{n, w, c}, {c, f}, {n, w, f}},
3748 /*iteratorTypes=*/ArrayRef<vector::IteratorType>{par, par, par, red});
3749 contrationOp.setKind(reductionKind);
3750 return contrationOp;
3751 }
3752
3753 // Create an outerproduct: lhs{w} * rhs{1} -> res{w} for single channel
3754 // convolution.
3755 Value conv1dSliceAsOuterProduct(RewriterBase &rewriter, Location loc,
3756 Value lhs, Value rhs, Value res) {
3757 lhs = promote(rewriter, loc, lhs, res.getType(), lhsCastOp);
3758 rhs = promote(rewriter, loc, rhs, res.getType(), rhsCastOp);
3759 return vector::OuterProductOp::create(rewriter, loc, res.getType(), lhs,
3760 rhs, res, vector::CombiningKind::ADD);
3761 }
3762
3763 // Create a reduction: lhs{n, w, c} -> res{n, w, c}
3764 Value pool1dSlice(RewriterBase &rewriter, Location loc, Value lhs,
3765 Value res) {
3766 if (isPoolExt)
3767 lhs = rewriter.create(loc, poolExtOp, lhs, res.getType())->getResult(0);
3768 return rewriter
3769 .create(loc, redOp, ArrayRef<Value>{lhs, res}, res.getType())
3770 ->getResult(0);
3771 }
3772
3773 /// Generate a vector implementation for:
3774 /// ```
3775 /// Op def: ( n, w, c, kw)
3776 /// Iters: ({Par(), Par(), Par(), Red()})
3777 /// Layout: {{n, strideW * w + dilationW * kw, c}, {kw, c}, {n, w, c}}
3778 /// ```
3779 /// kw is always unrolled.
3780 /// TODO: w (resp. kw) is unrolled when the strideW ( resp. dilationW) is
3781 /// > 1.
3782 FailureOr<Operation *> depthwiseConv(uint64_t channelDimVecSize,
3783 bool channelDimScalableFlag,
3784 bool flatten) {
3785 bool scalableChDim = false;
3786 bool useMasking = false;
3787 int64_t nSize, wSize, cSize, kwSize;
3788 // kernel{kw, c}
3789 bindShapeDims(rhsShapedType, kwSize, cSize);
3790 if (ShapedType::isDynamic(cSize)) {
3791 assert(channelDimVecSize != 0 && "Channel dim vec size must be > 0");
3792 cSize = channelDimVecSize;
3793 // Scalable vectors are only used when both conditions are met:
3794 // 1. channel dim is dynamic
3795 // 2. channelDimScalableFlag is set
3796 scalableChDim = channelDimScalableFlag;
3797 useMasking = true;
3798 }
3799
3800 assert(!(useMasking && flatten) &&
3801 "Unsupported flattened conv with dynamic shapes");
3802
3803 // out{n, w, c}
3804 bindShapeDims(resShapedType, nSize, wSize);
3805
3806 vector::TransferWriteOp write;
3807 Value zero = arith::ConstantIndexOp::create(rewriter, loc, 0);
3808
3809 // w is unrolled (i.e. wSizeStep == 1) iff strideW > 1.
3810 // When strideW == 1, we can batch the contiguous loads and avoid
3811 // unrolling
3812 int64_t wSizeStep = strideW == 1 ? wSize : 1;
3813
3814 Type lhsEltType = lhsShapedType.getElementType();
3815 Type rhsEltType = rhsShapedType.getElementType();
3816 Type resEltType = resShapedType.getElementType();
3817 VectorType lhsType = VectorType::get(
3818 {nSize,
3819 // iw = ow * sw + kw * dw - 1
3820 // (i.e. 16 convolved with 3 (@stride 1 dilation 1) -> 14)
3821 ((wSize - 1) * strideW + 1) + ((kwSize - 1) * dilationW + 1) - 1,
3822 cSize},
3823 lhsEltType, /*scalableDims=*/{false, false, scalableChDim});
3824 VectorType rhsType =
3825 VectorType::get({kwSize, cSize}, rhsEltType,
3826 /*scalableDims=*/{false, scalableChDim});
3827 VectorType resType =
3828 VectorType::get({nSize, wSize, cSize}, resEltType,
3829 /*scalableDims=*/{false, false, scalableChDim});
3830
3831 // Masks the input xfer Op along the channel dim, iff the corresponding
3832 // scalable flag is set.
3833 auto maybeMaskXferOp = [&](ArrayRef<int64_t> maskShape,
3834 ArrayRef<bool> scalableDims,
3835 Operation *opToMask) {
3836 if (!useMasking)
3837 return opToMask;
3838 auto maskType =
3839 VectorType::get(maskShape, rewriter.getI1Type(), scalableDims);
3840
3841 SmallVector<bool> inBounds(maskShape.size(), true);
3842 auto xferOp = cast<VectorTransferOpInterface>(opToMask);
3843 xferOp->setInherentAttr(
3844 rewriter.getStringAttr(xferOp.getInBoundsAttrName()),
3845 rewriter.getBoolArrayAttr(inBounds));
3846
3847 SmallVector<OpFoldResult> mixedDims = vector::getMixedSizesXfer(
3848 cast<LinalgOp>(op).hasPureTensorSemantics(), opToMask, rewriter);
3849
3850 Value maskOp =
3851 vector::CreateMaskOp::create(rewriter, loc, maskType, mixedDims);
3852
3853 return mlir::vector::maskOperation(rewriter, opToMask, maskOp);
3854 };
3855
3856 // Read lhs slice of size {n, w * strideW + kw * dilationW, c} @ [0, 0,
3857 // 0].
3858 Value lhs = vector::TransferReadOp::create(
3859 rewriter, loc, lhsType, lhsShaped, ValueRange{zero, zero, zero},
3860 /*padding=*/arith::getZeroConstant(rewriter, loc, lhsEltType));
3861 auto *maybeMaskedLhs = maybeMaskXferOp(
3862 lhsType.getShape(), lhsType.getScalableDims(), lhs.getDefiningOp());
3863
3864 // Read rhs slice of size {kw, c} @ [0, 0].
3865 Value rhs = vector::TransferReadOp::create(
3866 rewriter, loc, rhsType, rhsShaped, ValueRange{zero, zero},
3867 /*padding=*/arith::getZeroConstant(rewriter, loc, rhsEltType));
3868 auto *maybeMaskedRhs = maybeMaskXferOp(
3869 rhsType.getShape(), rhsType.getScalableDims(), rhs.getDefiningOp());
3870
3871 // Read res slice of size {n, w, c} @ [0, 0, 0].
3872 Value res = vector::TransferReadOp::create(
3873 rewriter, loc, resType, resShaped, ValueRange{zero, zero, zero},
3874 /*padding=*/arith::getZeroConstant(rewriter, loc, resEltType));
3875 auto *maybeMaskedRes = maybeMaskXferOp(
3876 resType.getShape(), resType.getScalableDims(), res.getDefiningOp());
3877
3878 //===------------------------------------------------------------------===//
3879 // Begin vector-only rewrite part
3880 //===------------------------------------------------------------------===//
3881 // Unroll along kw and read slices of lhs and rhs.
3882 SmallVector<Value> lhsVals, rhsVals, resVals;
3883 SmallVector<int64_t> inOutSliceSizes = {nSize, wSizeStep, cSize};
3884 SmallVector<int64_t> inOutStrides = {1, 1, 1};
3885
3886 // Extract lhs slice of size {n, wSizeStep, c}
3887 // @ [0, sw * w + dw * kw, 0].
3888 for (int64_t kw = 0; kw < kwSize; ++kw) {
3889 for (int64_t w = 0; w < wSize; w += wSizeStep) {
3890 lhsVals.push_back(vector::ExtractStridedSliceOp::create(
3891 rewriter, loc, maybeMaskedLhs->getResult(0),
3892 /*offsets=*/ArrayRef<int64_t>{0, w * strideW + kw * dilationW, 0},
3893 inOutSliceSizes, inOutStrides));
3894 }
3895 }
3896 // Extract rhs slice of size {c} @ [kw].
3897 for (int64_t kw = 0; kw < kwSize; ++kw) {
3898 rhsVals.push_back(
3899 vector::ExtractOp::create(rewriter, loc, maybeMaskedRhs->getResult(0),
3900 /*offsets=*/ArrayRef<int64_t>{kw}));
3901 }
3902 // Extract res slice: {n, wSizeStep, c} @ [0, w, 0].
3903 for (int64_t w = 0; w < wSize; w += wSizeStep) {
3904 resVals.push_back(vector::ExtractStridedSliceOp::create(
3905 rewriter, loc, maybeMaskedRes->getResult(0),
3906 /*offsets=*/ArrayRef<int64_t>{0, w, 0}, inOutSliceSizes,
3907 inOutStrides));
3908 }
3909
3910 auto linearIndex = [&](int64_t kw, int64_t w) {
3911 return kw * (wSize / wSizeStep) + w;
3912 };
3913
3914 // Note - the scalable flags are ignored as flattening combined with
3915 // scalable vectorization is not supported.
3916 SmallVector<int64_t> inOutFlattenSliceSizes = {nSize, wSizeStep * cSize};
3917 auto lhsTypeAfterFlattening =
3918 VectorType::get(inOutFlattenSliceSizes, lhsEltType);
3919 auto resTypeAfterFlattening =
3920 VectorType::get(inOutFlattenSliceSizes, resEltType);
3921
3922 // Compute contraction: O{n, w, c} += I{n, sw * w + dw * kw, c} * F{c}
3923 for (int64_t kw = 0; kw < kwSize; ++kw) {
3924 for (int64_t w = 0; w < wSize; w += wSizeStep) {
3925 Value lhsVal = lhsVals[linearIndex(kw, w)];
3926 Value resVal = resVals[w];
3927 if (flatten) {
3928 // Flatten the input and output vectors (collapse the channel
3929 // dimension)
3930 lhsVal =
3931 vector::ShapeCastOp::create(rewriter, loc, lhsTypeAfterFlattening,
3932 lhsVals[linearIndex(kw, w)]);
3933 resVal = vector::ShapeCastOp::create(
3934 rewriter, loc, resTypeAfterFlattening, resVals[w]);
3935 }
3936 resVals[w] = depthwiseConv1dSliceAsMulAcc(rewriter, loc, lhsVal,
3937 rhsVals[kw], resVal, flatten);
3938 if (flatten) {
3939 // Un-flatten the output vector (restore the channel dimension)
3940 resVals[w] = vector::ShapeCastOp::create(
3941 rewriter, loc, VectorType::get(inOutSliceSizes, resEltType),
3942 resVals[w]);
3943 }
3944 }
3945 }
3946
3947 // Its possible we failed to create the Fma.
3948 if (!llvm::all_of(resVals, [](Value v) { return v; })) {
3949 // Manually revert (in reverse order) to avoid leaving a bad IR state.
3950 for (auto &collection :
3951 {resVals, rhsVals, lhsVals, {res, rhs, lhs, zero}})
3952 for (Value v : collection)
3953 rewriter.eraseOp(v.getDefiningOp());
3954 return rewriter.notifyMatchFailure(op, "failed to create FMA");
3955 }
3956
3957 // Write back res slice: {n, wSizeStep, c} @ [0, w, 0].
3958 // This does not depend on kw.
3959 for (int64_t w = 0; w < wSize; w += wSizeStep) {
3960 maybeMaskedRes = vector::InsertStridedSliceOp::create(
3961 rewriter, loc, resVals[w], maybeMaskedRes->getResult(0),
3962 /*offsets=*/ArrayRef<int64_t>{0, w, 0},
3963 /*strides=*/ArrayRef<int64_t>{1, 1, 1});
3964 }
3965 //===------------------------------------------------------------------===//
3966 // End vector-only rewrite part
3967 //===------------------------------------------------------------------===//
3968
3969 // Write back res slice of size {n, w, c} @ [0, 0, 0].
3970 Operation *resOut = vector::TransferWriteOp::create(
3971 rewriter, loc, maybeMaskedRes->getResult(0), resShaped,
3972 ValueRange{zero, zero, zero});
3973 return maybeMaskXferOp(resType.getShape(), resType.getScalableDims(),
3974 resOut);
3975 }
3976
3977 /// Lower:
3978 /// * lhs{n, w, c} * rhs{c} -> res{n, w, c} (flatten = false)
3979 /// * lhs{n, w * c} * rhs{c} -> res{n, w * c} (flatten = true)
3980 /// to MulAcc.
3981 Value depthwiseConv1dSliceAsMulAcc(RewriterBase &rewriter, Location loc,
3982 Value lhs, Value rhs, Value res,
3983 bool flatten) {
3984 auto rhsTy = cast<ShapedType>(rhs.getType());
3985 auto resTy = cast<ShapedType>(res.getType());
3986
3987 // TODO(suderman): Change this to use a vector.ima intrinsic.
3988 lhs = promote(rewriter, loc, lhs, resTy, lhsCastOp);
3989
3990 if (flatten) {
3991 // NOTE: This following logic won't work for scalable vectors. For this
3992 // reason, "flattening" is not supported when shapes are dynamic (this
3993 // should be captured by one of the pre-conditions).
3994
3995 // There are two options for handling the filter:
3996 // * shape_cast(broadcast(filter))
3997 // * broadcast(shuffle(filter))
3998 // Opt for the option without shape_cast to simplify the codegen.
3999 auto rhsSize = cast<VectorType>(rhs.getType()).getShape()[0];
4000 auto resSize = cast<VectorType>(res.getType()).getShape()[1];
4001
4002 SmallVector<int64_t, 16> indices;
4003 for (int i = 0; i < resSize / rhsSize; ++i) {
4004 for (int j = 0; j < rhsSize; ++j)
4005 indices.push_back(j);
4006 }
4007
4008 rhs = vector::ShuffleOp::create(rewriter, loc, rhs, rhs, indices);
4009 }
4010 // Broadcast the filter to match the output vector
4011 rhs = vector::BroadcastOp::create(rewriter, loc,
4012 resTy.clone(rhsTy.getElementType()), rhs);
4013
4014 rhs = promote(rewriter, loc, rhs, resTy, rhsCastOp);
4015
4016 if (!lhs || !rhs)
4017 return nullptr;
4018
4019 if (isa<FloatType>(resTy.getElementType()))
4020 return vector::FMAOp::create(rewriter, loc, lhs, rhs, res);
4021
4022 auto mul = arith::MulIOp::create(rewriter, loc, lhs, rhs);
4023 return arith::AddIOp::create(rewriter, loc, mul, res);
4024 }
4025
4026 /// Entry point for non-channeled convolution:
4027 /// {{w + kw}, {kw}, {w}}
4028 FailureOr<Operation *> generateNonChanneledConv() {
4029 AffineExpr w, kw;
4030 bindDims(ctx, w, kw);
4031 if (!iters({Par(), Red()}))
4032 return rewriter.notifyMatchFailure(op,
4033 "failed to match conv::W 1-par 1-red");
4034
4035 // No transposition needed.
4036 if (layout({/*lhsIndex*/ {w + kw},
4037 /*rhsIndex*/ {kw},
4038 /*resIndex*/ {w}}))
4039 return conv(Conv1DOpOrder::W);
4040
4041 return rewriter.notifyMatchFailure(op, "not a conv::W layout");
4042 }
4043
4044 /// Entry point that transposes into the common form:
4045 /// {{n, strideW * w + dilationW * kw, c}, {kw, c, f}, {n, w, f}}
4046 FailureOr<Operation *> generateNwcConv() {
4047 AffineExpr n, w, f, kw, c;
4048 bindDims(ctx, n, w, f, kw, c);
4049 if (!iters({Par(), Par(), Par(), Red(), Red()}))
4050 return rewriter.notifyMatchFailure(
4051 op, "failed to match conv::Nwc 3-par 2-red");
4052
4053 // No transposition needed.
4054 if (layout({/*lhsIndex*/ {n, strideW * w + dilationW * kw, c},
4055 /*rhsIndex*/ {kw, c, f},
4056 /*resIndex*/ {n, w, f}}))
4057 return conv(Conv1DOpOrder::Nwc);
4058
4059 return rewriter.notifyMatchFailure(op, "not a conv::Nwc layout");
4060 }
4061
4062 /// Entry point that transposes into the common form:
4063 /// {{n, c, strideW * w + dilationW * kw}, {f, c, kw}, {n, f, w}}
4064 FailureOr<Operation *> generateNcwConv() {
4065 AffineExpr n, w, f, kw, c;
4066 bindDims(ctx, n, f, w, c, kw);
4067 if (!iters({Par(), Par(), Par(), Red(), Red()}))
4068 return rewriter.notifyMatchFailure(
4069 op, "failed to match conv::Ncw 3-par 2-red");
4070
4071 if (layout({/*lhsIndex*/ {n, c, strideW * w + dilationW * kw},
4072 /*rhsIndex*/ {f, c, kw},
4073 /*resIndex*/ {n, f, w}}))
4074 return conv(Conv1DOpOrder::Ncw);
4075
4076 return rewriter.notifyMatchFailure(op, "not a conv::Ncw layout");
4077 }
4078
4079 /// Entry point that transposes into the common form:
4080 /// {{n, strideW * w + dilationW * kw, c}, {kw}, {n, w, c}} for pooling
4081 FailureOr<Operation *> generateNwcPooling() {
4082 AffineExpr n, w, c, kw;
4083 bindDims(ctx, n, w, c, kw);
4084 if (!iters({Par(), Par(), Par(), Red()}))
4085 return rewriter.notifyMatchFailure(op,
4086 "failed to match pooling 3-par 1-red");
4087
4088 // No transposition needed.
4089 if (layout({/*lhsIndex*/ {n, strideW * w + dilationW * kw, c},
4090 /*rhsIndex*/ {kw},
4091 /*resIndex*/ {n, w, c}}))
4092 return conv(Conv1DOpOrder::Nwc);
4093
4094 return rewriter.notifyMatchFailure(op, "not a pooling::Nwc layout");
4095 }
4096
4097 /// Entry point that transposes into the common form:
4098 /// {{n, c, strideW * w + dilationW * kw}, {kw}, {n, c, w}} for pooling
4099 FailureOr<Operation *> generateNcwPooling() {
4100 AffineExpr n, w, c, kw;
4101 bindDims(ctx, n, c, w, kw);
4102 if (!iters({Par(), Par(), Par(), Red()}))
4103 return rewriter.notifyMatchFailure(op,
4104 "failed to match pooling 3-par 1-red");
4105
4106 if (layout({/*lhsIndex*/ {n, c, strideW * w + dilationW * kw},
4107 /*rhsIndex*/ {kw},
4108 /*resIndex*/ {n, c, w}}))
4109 return conv(Conv1DOpOrder::Ncw);
4110
4111 return rewriter.notifyMatchFailure(op, "not a pooling::Ncw layout");
4112 }
4113
4114 /// Entry point that transposes into the common form:
4115 /// {{n, strideW * w + dilationW * kw, c}, {kw, c}, {n, w, c}}
4116 FailureOr<Operation *> generateDilatedConv(uint64_t vecChDimSize = 0,
4117 bool vecChDimScalableFlag = false,
4118 bool flatten = false) {
4119 AffineExpr n, w, c, kw;
4120 bindDims(ctx, n, w, c, kw);
4121 if (!iters({Par(), Par(), Par(), Red()}))
4122 return rewriter.notifyMatchFailure(
4123 op, "failed to match depthwise::Nwc conv 3-par 1-red");
4124
4125 // No transposition needed.
4126 if (layout({/*lhsIndex*/ {n, strideW * w + dilationW * kw, c},
4127 /*rhsIndex*/ {kw, c},
4128 /*resIndex*/ {n, w, c}}))
4129 return depthwiseConv(vecChDimSize, vecChDimScalableFlag, flatten);
4130
4131 return rewriter.notifyMatchFailure(op, "not a depthwise::Nwc layout");
4132 }
4133
4134private:
4135 ConvOperationKind oper = ConvOperationKind::Conv;
4136 StringAttr redOp;
4137 StringAttr poolExtOp;
4138 bool isPoolExt = false;
4139 // Casts used to widen the convolution payload's lhs and rhs. These are null
4140 // only when the corresponding operand already has the accumulator type.
4141 Operation *lhsCastOp = nullptr;
4142 Operation *rhsCastOp = nullptr;
4143 int strideW, dilationW;
4144 Value lhsShaped, rhsShaped, resShaped;
4145 ShapedType lhsShapedType, rhsShapedType, resShapedType;
4146 vector::CombiningKind reductionKind;
4147
4148 // Sets oper, poolExtOp, isPoolExt and the conv operand casts for valid
4149 // conv/pooling ops.
4150 void setConvOperationKind(Operation *reduceOp) {
4151 int numBlockArguments =
4152 llvm::count_if(reduceOp->getOperands(), llvm::IsaPred<BlockArgument>);
4153 if (numBlockArguments == 1) {
4154 // Will be convolution if feeder is a MulOp.
4155 // A strength reduced version of MulOp for i1 type is AndOp which is also
4156 // supported. Otherwise, it can be pooling. This strength reduction logic
4157 // is in `buildBinaryFn` helper in the Linalg dialect.
4158 auto feedValIt = llvm::find_if_not(reduceOp->getOperands(),
4159 llvm::IsaPred<BlockArgument>);
4160 Operation *feedOp = (*feedValIt).getDefiningOp();
4161 if (isCastOfBlockArgument(feedOp)) {
4162 oper = ConvOperationKind::Pool;
4163 isPoolExt = true;
4164 poolExtOp = feedOp->getName().getIdentifier();
4165 return;
4166 }
4167 oper = ConvOperationKind::Conv;
4168 setConvCastOps(feedOp);
4169 return;
4170 }
4171 // numBlockArugments == 2 and this is a pooling op.
4172 oper = ConvOperationKind::Pool;
4173 isPoolExt = false;
4174 }
4175
4176 // Record the casts applied to the input and filter.
4177 void setConvCastOps(Operation *feedOp) {
4178 lhsCastOp = feedOp->getOperand(0).getDefiningOp();
4179 rhsCastOp = feedOp->getOperand(1).getDefiningOp();
4180 }
4181};
4182} // namespace
4183
4184/// Helper function to vectorize a LinalgOp with convolution semantics.
4185// TODO: extend the generic vectorization to support windows and drop this.
4186static FailureOr<Operation *> vectorizeConvolution(
4187 RewriterBase &rewriter, LinalgOp op, ArrayRef<int64_t> inputVecSizes,
4188 ArrayRef<bool> inputScalableVecDims, bool flatten1DDepthwiseConv) {
4189 FailureOr<Conv1DGenerator> conv1dGen = Conv1DGenerator::create(rewriter, op);
4190 if (failed(conv1dGen))
4191 return failure();
4192 auto res = conv1dGen->generateNonChanneledConv();
4193 if (succeeded(res))
4194 return res;
4195 res = conv1dGen->generateNwcConv();
4196 if (succeeded(res))
4197 return res;
4198 res = conv1dGen->generateNcwConv();
4199 if (succeeded(res))
4200 return res;
4201 res = conv1dGen->generateNwcPooling();
4202 if (succeeded(res))
4203 return res;
4204 res = conv1dGen->generateNcwPooling();
4205 if (succeeded(res))
4206 return res;
4207
4208 // Only depthwise 1D NWC convs are left - these can be vectorized using masks
4209 // and scalable vectors. Note that ATM the only dim that can be dynamic (i.e.
4210 // masked/scalable) is the channel dim (i.e. the trailing dim).
4211 uint64_t vecChDimSize = ShapedType::kDynamic;
4212 bool vecChDimScalableFlag = false;
4213 if (!inputVecSizes.empty()) {
4214 // Only use the input vector size corresponding to the channel dim. Other
4215 // vector dims will be inferred from the Ops.
4218 "Not a 1D depthwise conv!");
4219 size_t chDimIdx = 0;
4221 chDimIdx = 2;
4223 chDimIdx = 1;
4224
4225 vecChDimSize = inputVecSizes[chDimIdx];
4226 vecChDimScalableFlag = inputScalableVecDims[chDimIdx];
4227 }
4228 return conv1dGen->generateDilatedConv(vecChDimSize, vecChDimScalableFlag,
4229 flatten1DDepthwiseConv);
4230}
4231
4232struct VectorizeConvolution : public OpInterfaceRewritePattern<LinalgOp> {
4234
4235 LogicalResult matchAndRewrite(LinalgOp op,
4236 PatternRewriter &rewriter) const override {
4237 FailureOr<Operation *> resultOrFail = vectorizeConvolution(rewriter, op);
4238 if (failed(resultOrFail))
4239 return failure();
4240 Operation *newOp = *resultOrFail;
4241 if (newOp->getNumResults() == 0) {
4242 rewriter.eraseOp(op.getOperation());
4243 return success();
4244 }
4245 assert(newOp->getNumResults() == 1 && "expected single result");
4246 rewriter.replaceOp(op.getOperation(), newOp->getResult(0));
4247 return success();
4248 }
4249};
4250
4252 RewritePatternSet &patterns, PatternBenefit benefit) {
4253 patterns.add<VectorizeConvolution>(patterns.getContext(), benefit);
4254}
return success()
lhs
b
Return true if permutation is a valid permutation of the outer_dims_perm (case OuterOrInnerPerm::Oute...
static std::optional< VectorShape > vectorShape(Type type)
static bool isLoopInvariantIdx(LinalgOp &linalgOp, Value &val, VectorType resType)
Checks whether val can be used for calculating a loop invariant index.
static Value insertConvResultSlices(RewriterBase &rewriter, Location loc, Value res, int64_t wSize, int64_t wSizeStep, SmallVectorImpl< Value > &resVals, bool isSingleChanneled)
Helper function to insert the computed result slices.
static SmallVector< bool > getDimsToReduce(LinalgOp linalgOp)
static VectorMemoryAccessKind getTensorExtractMemoryAccessPattern(tensor::ExtractOp extractOp, LinalgOp &linalgOp, VectorType resType)
Infer the memory access pattern for the input ExtractOp.
static SmallVector< Value > extractConvInputSlices(RewriterBase &rewriter, Location loc, Value input, int64_t nSize, int64_t wSize, int64_t cSize, int64_t kwSize, int strideW, int dilationW, int64_t wSizeStep, bool isSingleChanneled)
Helper function to extract the input slices after filter is unrolled along kw.
VectorMemoryAccessKind
@ Contiguous
@ Gather
@ ScalarBroadcast
static VectorizationHookResult vectorizeTensorExtract(RewriterBase &rewriter, VectorizationState &state, Operation *op, LinalgOp linalgOp, const IRMapping &bvm)
Helper function to vectorize the tensor.extract operations.
static VectorizationHookResult vectorizeLinalgIndex(RewriterBase &rewriter, VectorizationState &state, Operation *op, LinalgOp linalgOp)
Helper function to vectorize the index operations of a linalgOp.
static LogicalResult vectorizeAsInsertSliceOp(RewriterBase &rewriter, tensor::InsertSliceOp sliceOp, ArrayRef< int64_t > inputVectorSizes, SmallVectorImpl< Value > &newResults)
Vectorize tensor::InsertSliceOp with:
static FailureOr< Operation * > vectorizeConvolution(RewriterBase &rewriter, LinalgOp convOp, ArrayRef< int64_t > inputVecSizes={}, ArrayRef< bool > inputVecScalableFlags={}, bool flatten1DDepthwiseConv=false)
Try to vectorize convOp as a convolution.
static LogicalResult vectorizeAsLinalgGeneric(RewriterBase &rewriter, VectorizationState &state, LinalgOp linalgOp, SmallVectorImpl< Value > &newResults)
Generic vectorization function that rewrites the body of a linalgOp into vector form.
#define MATCH_1D_CONV_POOL_OP(ConvOpTy)
static VectorizationHookResult vectorizeOneOp(RewriterBase &rewriter, VectorizationState &state, LinalgOp linalgOp, Operation *op, const IRMapping &bvm, ArrayRef< CustomVectorizationHook > customVectorizationHooks)
Generic vectorization for a single operation op, given already vectorized operands carried by bvm.
static Operation * matchLinalgReduction(OpOperand *outputOperand)
Check whether outputOperand is a reduction with a single combiner operation.
static Value buildVectorWrite(RewriterBase &rewriter, Value value, OpOperand *outputOperand, VectorizationState &state)
Build a vector.transfer_write of value into outputOperand at indices set to all 0; where outputOperan...
static Value getStaticPadVal(Operation *op)
Returns the effective Pad value for the input op, provided it's a scalar.
static SmallVector< Value > extractConvFilterSlices(RewriterBase &rewriter, Location loc, Value filter, int64_t kwSize)
Helper function to extract the filter slices after filter is unrolled along kw.
static bool hasReductionIterator(LinalgOp &op)
Check if op is a linalg.reduce or a linalg.generic that has at least one reduction iterator.
std::function< LogicalResult(Operation *, bool)> CustomVectorizationPrecondition
static uint64_t getTrailingNonUnitLoopDimIdx(LinalgOp linalgOp)
Find the index of the trailing non-unit dim in linalgOp.
static VectorType getCollapsedVecType(VectorType type, ArrayRef< AffineMap > reassociation)
Given the re-associations, "collapses" the input Vector type.
Conv1DOpOrder
Helper enum to represent conv1d input traversal order.
VectorizationHookStatus
Helper data structure to represent the result of vectorization for a single operation.
@ Failure
Op failed to vectorize.
@ NewOp
Op vectorized into a new Op whose results will replace original Op's results.
@ NoReplace
Op vectorized and custom function took care of replacement logic.
static Operation * reduceIfNeeded(OpBuilder &b, LinalgOp linalgOp, Operation *op, Value reduceValue, Value initialValue, const IRMapping &bvm)
Emit reduction operations if the shapes of the value to reduce is different that the result shape.
std::function< VectorizationHookResult(Operation *, const IRMapping &)> CustomVectorizationHook
static AffineMap reindexIndexingMap(AffineMap map)
Given an indexing map coming from a LinalgOp indexing, restricted to a projectedPermutation,...
static LogicalResult tensorExtractVectorizationPrecondition(Operation *op, bool vectorizeNDExtract)
Helper function to check if the tensor.extract can be vectorized by the custom hook vectorizeTensorEx...
static Value broadcastIfNeeded(OpBuilder &b, Value value, Type dstType)
Broadcast value to a vector of shape if possible.
static Value calculateGatherOffset(RewriterBase &rewriter, VectorizationState &state, tensor::ExtractOp extractOp, const IRMapping &bvm)
Calculates the offsets ($index_vec) for vector.gather operations generated from tensor....
static SmallVector< Value > extractConvResultSlices(RewriterBase &rewriter, Location loc, Value res, int64_t nSize, int64_t wSize, int64_t fSize, int64_t wSizeStep, bool isSingleChanneled)
Helper function to extract the result slices after filter is unrolled along kw.
static bool isContiguousLoadIdx(LinalgOp &linalgOp, Value &val, bool &foundIndexOp, VectorType resType)
Check whether val could be used for calculating the trailing index for a contiguous load operation.
static VectorizationHookResult vectorizeLinalgYield(RewriterBase &rewriter, Operation *op, const IRMapping &bvm, VectorizationState &state, LinalgOp linalgOp, SmallVectorImpl< Value > &newResults)
Helper function to vectorize the terminator of a linalgOp.
static Operation * buildMultiDimReduce(OpBuilder &b, Operation *reduceOp, Value valueToReduce, Value acc, ArrayRef< bool > dimsToMask)
Create MultiDimReductionOp to compute the reduction for reductionOp.
#define mul(a, b)
A dimensional identifier appearing in an affine expression.
Definition AffineExpr.h:223
A multi-dimensional affine map Affine map's are immutable like Type's, and they are uniqued.
Definition AffineMap.h:46
static AffineMap getMinorIdentityMap(unsigned dims, unsigned results, MLIRContext *context)
Returns an identity affine map (d0, ..., dn) -> (dp, ..., dn) on the most minor dimensions.
MLIRContext * getContext() const
static AffineMap getMultiDimIdentityMap(unsigned numDims, MLIRContext *context)
Returns an AffineMap with 'numDims' identity result dim exprs.
static AffineMap get(MLIRContext *context)
Returns a zero result affine map with no dimensions or symbols: () -> ().
bool isProjectedPermutation(bool allowZeroInResults=false) const
Returns true if the AffineMap represents a subset (i.e.
unsigned getNumResults() const
unsigned getNumInputs() const
AffineExpr getResult(unsigned idx) const
static AffineMap getFilteredIdentityMap(MLIRContext *ctx, unsigned numDims, llvm::function_ref< bool(AffineDimExpr)> keepDimFilter)
Returns an identity affine map with numDims input dimensions and filtered results using keepDimFilter...
AffineMap dropZeroResults()
Returns the AffineMap resulting from removing "zero" results (constant values == 0) from this map.
static AffineMap getPermutationMap(ArrayRef< unsigned > permutation, MLIRContext *context)
Returns an AffineMap representing a permutation.
SmallVector< unsigned > getBroadcastDims() const
Returns the list of broadcast dimensions (i.e.
AffineMap compose(AffineMap map) const
Returns the AffineMap resulting from composing this with map.
bool isPermutation() const
Returns true if the AffineMap represents a symbol-less permutation map.
This class represents an argument of a Block.
Definition Value.h:306
unsigned getArgNumber() const
Returns the number of this argument.
Definition Value.h:318
Block represents an ordered list of Operations.
Definition Block.h:34
OpListType & getOperations()
Definition Block.h:162
AffineMap getMultiDimIdentityMap(unsigned rank)
Definition Builders.cpp:396
StringAttr getStringAttr(const Twine &bytes)
Definition Builders.cpp:271
TypedAttr getZeroAttr(Type type)
Definition Builders.cpp:333
IntegerType getI1Type()
Definition Builders.cpp:61
ArrayAttr getArrayAttr(ArrayRef< Attribute > value)
Definition Builders.cpp:275
MLIRContext * getContext() const
Definition Builders.h:56
IndexType getIndexType()
Definition Builders.cpp:59
ArrayAttr getBoolArrayAttr(ArrayRef< bool > values)
Definition Builders.cpp:279
static DenseIntElementsAttr get(const ShapedType &type, Arg &&arg)
Get an instance of a DenseIntElementsAttr with the given arguments.
This is a utility class for mapping one set of IR entities to another.
Definition IRMapping.h:26
auto lookup(T from) const
Lookup a mapped value within the map.
Definition IRMapping.h:72
void map(Value from, Value to)
Inserts a new mapping for 'from' to 'to'.
Definition IRMapping.h:30
IRValueT get() const
Return the current value being used by this operand.
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
Operation * create(const OperationState &state)
Creates an operation given the fields represented as an OperationState.
Definition Builders.cpp:466
Operation * insert(Operation *op)
Insert the given operation at the current insertion point and return it.
Definition Builders.cpp:430
This class represents an operand of an operation.
Definition Value.h:254
unsigned getOperandNumber() const
Return which operand this is in the OpOperand list of the Operation.
Definition Value.cpp:226
StringAttr getIdentifier() const
Return the name of this operation as a StringAttr.
Operation is the basic unit of execution within MLIR.
Definition Operation.h:87
PropertyRef getPropertiesStorage()
Return a generic (but typed) reference to the property type storage.
Definition Operation.h:953
Value getOperand(unsigned idx)
Definition Operation.h:375
bool isBeforeInBlock(Operation *other)
Given an operation 'other' that is within the same parent block, return whether the current operation...
Block * getBlock()
Returns the operation block that contains this operation.
Definition Operation.h:230
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
unsigned getNumOperands()
Definition Operation.h:371
Attribute getPropertiesAsAttribute()
Return the properties converted to an attribute.
operand_iterator operand_end()
Definition Operation.h:400
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
static Operation * create(Location location, OperationName name, TypeRange resultTypes, ValueRange operands, NamedAttrList &&attributes, PropertyRef properties, BlockRange successors, unsigned numRegions)
Create a new Operation with the specific fields.
Definition Operation.cpp:65
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
unsigned short getBenefit() const
If the corresponding pattern can match, return its benefit. If the.
MLIRContext * getContext() const
RewritePatternSet & add(ConstructorArg &&arg, ConstructorArgs &&...args)
Add an instance of each of the pattern types 'Ts' to the pattern list with the given arguments.
This class coordinates the application of a rewrite on a set of IR, providing a way for clients to tr...
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.
void replaceAllUsesExcept(Value from, Value to, Operation *exceptedUser)
Find uses of from and replace them with to except if the user is exceptedUser.
std::enable_if_t<!std::is_convertible< CallbackT, Twine >::value, LogicalResult > notifyMatchFailure(Location loc, CallbackT &&reasonCallback)
Used to notify the listener that the IR failed to be rewritten because of a match failure,...
void modifyOpInPlace(Operation *root, CallableT &&callable)
This method is a utility wrapper around an in-place modification of an operation.
OpTy replaceOpWithNewOp(Operation *op, Args &&...args)
Replace the results of the given (original) op with a new op that is created without verification (re...
Instances of the Type class are uniqued, have an immutable identifier and an optional mutable compone...
Definition Types.h:74
bool isIndex() const
Definition Types.cpp:56
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
use_range getUses() const
Returns a range of all uses, which is useful for iterating over all uses.
Definition Value.h:188
Location getLoc() const
Return the location of this value.
Definition Value.cpp:24
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
Operation * getOwner() const
Return the owner of this operand.
Definition UseDefLists.h:38
bool hasElementwiseMappableTraits(Operation *op)
Together, Elementwise, Scalarizable, Vectorizable, and Tensorizable provide an easy way for scalar op...
bool hasVectorizationImpl(Operation *)
Return true if there's dedicated logic in the Linalg Vectorizer to vectorize this Op,...
SmallVector< int64_t > getUnPackInverseSrcPerm(linalg::UnPackOp, PackingMetadata &metadata)
Compute inverse permutation for the source tensor (i.e.
bool allIndexingsAreProjectedPermutation(LinalgOp op)
Check if all indexing maps are projected permutations.
Definition Utils.cpp:197
FailureOr< VectorizationResult > vectorize(RewriterBase &rewriter, Operation *op, ArrayRef< int64_t > inputVectorSizes={}, ArrayRef< bool > inputScalableVecDims={}, bool vectorizeNDExtract=false, bool flatten1DDepthwiseConv=false, bool assumeDynamicDimsMatchVecSizes=false, bool createNamedContraction=false)
Returns a VectorizationResult containing the results of the vectorized op, or failure if the transfor...
void populatePadOpVectorizationPatterns(RewritePatternSet &patterns, PatternBenefit baseBenefit=1)
Populates patterns with patterns that vectorize tensor.pad.
bool isReductionIterator(utils::IteratorType iteratorType)
Check if iterator type has "reduction" semantics.
Definition Utils.cpp:236
bool isaConvolutionOpInterface(LinalgOp linalgOp, bool allowEmptyConvolvedDims=false)
Checks whether linalgOp conforms to ConvolutionOpInterface.
void populateConvolutionVectorizationPatterns(RewritePatternSet &patterns, PatternBenefit benefit=1)
Populate patterns for vectorizing low-D convolution ops.
bool isElementwise(LinalgOp op)
Check if a LinalgOp is an element-wise operation.
Definition Utils.cpp:217
LogicalResult vectorizeCopy(RewriterBase &builder, memref::CopyOp copyOp)
Emit a suitable vector form for a Copy op with fully static shape.
LogicalResult vectorizeOpPrecondition(Operation *op, ArrayRef< int64_t > inputVectorSizes={}, ArrayRef< bool > inputScalableVecDims={}, bool vectorizeNDExtract=false, bool flatten1DDepthwiseConv=false)
Return success if the operation can be vectorized.
SmallVector< int64_t > getPackInverseDestPerm(linalg::PackOp packOp, PackingMetadata &metadata)
Compute inverse permutation for the destination tensor (i.e.
bool isaConvolutionOpOfType(LinalgOp op)
Returns true if the linalg op is a convolution op of type ConvOpTy.
Definition Utils.h:126
std::optional< vector::CombiningKind > getCombinerOpKind(Operation *combinerOp)
Return vector::CombiningKind for the given op.
detail::InFlightRemark failed(Location loc, RemarkOpts opts)
Report an optimization remark that failed.
Definition Remarks.h:734
void promote(RewriterBase &rewriter, scf::ForallOp forallOp)
Promotes the loop body of a scf::ForallOp to its containing block.
Definition SCF.cpp:753
std::enable_if_t<!is_complex< V >::value, V > readValue(char **linePtr)
Returns an element-value of non-complex type.
Definition File.h:50
Operation * maskOperation(OpBuilder &builder, Operation *maskableOp, Value mask, Value passthru=Value())
Creates a vector.mask operation around a maskable operation.
LogicalResult isValidMaskedInputVector(ArrayRef< int64_t > shape, ArrayRef< int64_t > inputVectorSizes)
Returns success if inputVectorSizes is a valid masking configuraion for given shape,...
BroadcastableToResult isBroadcastableTo(Type srcType, VectorType dstVectorType, std::pair< VectorDim, VectorDim > *mismatchingDims=nullptr)
Return whether srcType can be broadcast to dstVectorType under the semantics of the vector....
Operation * createWriteOrMaskedWrite(OpBuilder &builder, Location loc, Value vecToStore, Value dest, SmallVector< Value > writeIndices={}, bool useInBoundsInsteadOfMasking=false, AffineMap permutationMap=AffineMap())
Create a TransferWriteOp of vecToStore into dest.
Value createReadOrMaskedRead(OpBuilder &builder, Location loc, Value source, const VectorType &vecToReadTy, std::optional< Value > padValue=std::nullopt, bool useInBoundsInsteadOfMasking=false, ArrayRef< Value > indices={}, AffineMap permutationMap=AffineMap())
Creates a TransferReadOp from source.
SmallVector< OpFoldResult > getMixedSizesXfer(bool hasTensorSemantics, Operation *xfer, RewriterBase &rewriter)
A wrapper for getMixedSizes for vector.transfer_read and vector.transfer_write Ops (for source and de...
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
bool isEqualConstantIntOrValue(OpFoldResult ofr1, OpFoldResult ofr2)
Return true if ofr1 and ofr2 are the same integer constant attribute values or the same SSA value.
Type getType(OpFoldResult ofr)
Returns the int type of the integer in ofr.
Definition Utils.cpp:311
AffineMap inverseAndBroadcastProjectedPermutation(AffineMap map)
Return the reverse map of a projected permutation where the projected dimensions are transformed into...
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...
SmallVector< AffineMap, 4 > getSymbolLessAffineMaps(ArrayRef< ReassociationExprs > reassociation)
Constructs affine maps out of Array<Array<AffineExpr>>.
SmallVector< SmallVector< OpFoldResult > > ReifiedRankedShapedTypeDims
Value matchReduction(ArrayRef< BlockArgument > iterCarriedArgs, unsigned redPos, SmallVectorImpl< Operation * > &combinerOps)
Utility to match a generic reduction given a list of iteration-carried arguments, iterCarriedArgs and...
llvm::SetVector< T, Vector, Set, N > SetVector
Definition LLVM.h:125
Type getElementTypeOrSelf(Type type)
Return the element type or return the type itself.
std::conditional_t< std::is_same_v< Ty, mlir::Type >, mlir::Value, detail::TypedValue< Ty > > TypedValue
If Ty is mlir::Type this will select Value instead of having a wrapper around it.
Definition Value.h:494
void getUsedValuesDefinedAbove(Region &region, Region &limit, SetVector< Value > &values)
Fill values with a list of values defined at the ancestors of the limit region and used within region...
AffineMap compressUnusedDims(AffineMap map)
Drop the dims that are not used.
SmallVector< SmallVector< AffineExpr, 2 >, 2 > convertReassociationIndicesToExprs(MLIRContext *context, ArrayRef< ReassociationIndices > reassociationIndices)
Convert reassociation indices to affine expressions.
bool isReassociationValid(ArrayRef< AffineMap > reassociation, int *invalidIndex=nullptr)
Return true if the reassociation specification is valid, false otherwise.
llvm::TypeSwitch< T, ResultT > TypeSwitch
Definition LLVM.h:139
Value getValueOrCreateConstantIndexOp(OpBuilder &b, Location loc, OpFoldResult ofr)
Converts an OpFoldResult to a Value.
Definition Utils.cpp:114
AffineExpr getAffineConstantExpr(int64_t constant, MLIRContext *context)
auto get(MLIRContext *context, Ts &&...params)
Helper method that injects context only if needed, this helps unify some of the attribute constructio...
llvm::DenseMap< KeyT, ValueT, KeyInfoT, BucketT > DenseMap
Definition LLVM.h:120
SmallVector< T > applyPermutationMap(AffineMap map, llvm::ArrayRef< T > source)
Apply a permutation from map to source and return the result.
Definition AffineMap.h:675
llvm::SmallBitVector getUnusedDimsBitVector(ArrayRef< AffineMap > maps)
detail::constant_op_matcher m_Constant()
Matches a constant foldable operation.
Definition Matchers.h:369
void applyPermutationToVector(SmallVector< T, N > &inVec, ArrayRef< int64_t > permutation)
Apply the permutation defined by permutation to inVec.
SmallVector< int64_t > invertPermutationVector(ArrayRef< int64_t > permutation)
Helper method to apply to inverse a permutation.
VectorizationHookResult contains the vectorized op returned from a CustomVectorizationHook.
enum VectorizationHookStatus status
Return status from vectorizing the current op.
Operation * newOp
New vectorized operation to replace the current op.
ArrayRef< int64_t > getCanonicalVecShape() const
Returns the canonical vector shape used to vectorize the iteration space.
LogicalResult initState(RewriterBase &rewriter, LinalgOp linalgOp, ArrayRef< int64_t > inputVectorSizes, ArrayRef< bool > inputScalableVecDims, bool assumeDynamicDimsMatchVecSizes=false)
Initializes the vectorization state, including the computation of the canonical vector shape for vect...
Operation * maskOperation(RewriterBase &rewriter, Operation *opToMask, LinalgOp linalgOp, std::optional< AffineMap > maybeIndexingMap=std::nullopt)
Masks an operation with the canonical vector mask if the operation needs masking.
VectorType getCanonicalVecType(Type elementType, std::optional< AffineMap > dimPermutation=std::nullopt) const
Returns a vector type of the provided elementType with the canonical vector shape and the correspondi...
ArrayRef< bool > getScalableVecDims() const
Returns the vector dimensions that are scalable in the canonical vector shape.
VectorizationState(RewriterBase &rewriter)
OpInterfaceRewritePattern(MLIRContext *context, PatternBenefit benefit=1)
LogicalResult matchAndRewrite(vector::TransferReadOp xferOp, PatternRewriter &rewriter) const override
LogicalResult matchAndRewrite(vector::TransferWriteOp xferOp, PatternRewriter &rewriter) const override