MLIR 24.0.0git
VectorToSCF.cpp
Go to the documentation of this file.
1//===- VectorToSCF.cpp - Convert vector to SCF dialect ----------*- C++ -*-===//
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 lowering of vector transfer operations to SCF.
10//
11//===----------------------------------------------------------------------===//
12
13#include <numeric>
14#include <optional>
15
17
26#include "mlir/IR/Builders.h"
27#include "mlir/Pass/Pass.h"
29#include "llvm/ADT/STLExtras.h"
30
31namespace mlir {
32#define GEN_PASS_DEF_CONVERTVECTORTOSCF
33#include "mlir/Conversion/Passes.h.inc"
34} // namespace mlir
35
36using namespace mlir;
37using vector::TransferReadOp;
38using vector::TransferWriteOp;
39
40namespace {
41
42/// Attribute name used for labeling transfer ops during progressive lowering.
43static const char kPassLabel[] = "__vector_to_scf_lowering__";
44
45/// Return true if this transfer op operates on a source tensor.
46static bool isTensorOp(VectorTransferOpInterface xferOp) {
47 if (isa<RankedTensorType>(xferOp.getShapedType())) {
48 if (isa<vector::TransferWriteOp>(xferOp)) {
49 // TransferWriteOps on tensors have a result.
50 assert(xferOp->getNumResults() > 0);
51 }
52 return true;
53 }
54 return false;
55}
56
57/// Patterns that inherit from this struct have access to
58/// VectorTransferToSCFOptions.
59template <typename OpTy>
60struct VectorToSCFPattern : public OpRewritePattern<OpTy> {
61 explicit VectorToSCFPattern(MLIRContext *context,
62 VectorTransferToSCFOptions opt)
63 : OpRewritePattern<OpTy>(context), options(opt) {}
64
65 LogicalResult checkLowerTensors(VectorTransferOpInterface xferOp,
66 PatternRewriter &rewriter) const {
67 if (isTensorOp(xferOp) && !options.lowerTensors) {
68 return rewriter.notifyMatchFailure(
69 xferOp, "lowering tensor transfers is disabled");
70 }
71 return success();
72 }
73
74 VectorTransferToSCFOptions options;
75};
76
77/// Given a vector transfer op, calculate which dimension of the `source`
78/// memref should be unpacked in the next application of TransferOpConversion.
79/// A return value of std::nullopt indicates a broadcast.
80template <typename OpTy>
81static std::optional<int64_t> unpackedDim(OpTy xferOp) {
82 // TODO: support 0-d corner case.
83 assert(xferOp.getTransferRank() > 0 && "unexpected 0-d transfer");
84 auto map = xferOp.getPermutationMap();
85 if (auto expr = dyn_cast<AffineDimExpr>(map.getResult(0))) {
86 return expr.getPosition();
87 }
88 assert(xferOp.isBroadcastDim(0) &&
89 "Expected AffineDimExpr or AffineConstantExpr");
90 return std::nullopt;
91}
92
93/// Compute the permutation map for the new (N-1)-D vector transfer op. This
94/// map is identical to the current permutation map, but the first result is
95/// omitted.
96template <typename OpTy>
97static AffineMap unpackedPermutationMap(OpBuilder &b, OpTy xferOp) {
98 // TODO: support 0-d corner case.
99 assert(xferOp.getTransferRank() > 0 && "unexpected 0-d transfer");
100 auto map = xferOp.getPermutationMap();
101 return AffineMap::get(map.getNumDims(), 0, map.getResults().drop_front(),
102 b.getContext());
103}
104
105/// Calculate the indices for the new vector transfer op.
106///
107/// E.g.: transfer_read %A[%a, %b, %c, %d] ... : vector<5x4x3xf32> ...
108/// --> transfer_read %A[%a, %b + iv, %c, %d] ... vector<4x3f32>
109/// ^^^^^^
110/// `iv` is the iteration variable of the (new) surrounding loop.
111template <typename OpTy>
112static void getXferIndices(OpBuilder &b, OpTy xferOp, Value iv,
114 typename OpTy::Adaptor adaptor(xferOp);
115 // Corresponding memref dim of the vector dim that is unpacked.
116 auto dim = unpackedDim(xferOp);
117 auto prevIndices = adaptor.getIndices();
118 indices.append(prevIndices.begin(), prevIndices.end());
119
120 Location loc = xferOp.getLoc();
121 bool isBroadcast = !dim.has_value();
122 if (!isBroadcast) {
123 AffineExpr d0, d1;
124 bindDims(xferOp.getContext(), d0, d1);
125 Value offset = adaptor.getIndices()[*dim];
126 indices[*dim] =
127 affine::makeComposedAffineApply(b, loc, d0 + d1, {offset, iv});
128 }
129}
130
131static void maybeYieldValue(OpBuilder &b, Location loc, bool hasRetVal,
132 Value value) {
133 if (hasRetVal) {
134 assert(value && "Expected non-empty value");
135 scf::YieldOp::create(b, loc, value);
136 } else {
137 scf::YieldOp::create(b, loc);
138 }
139}
140
141/// Generates a boolean Value that is true if the iv-th bit in xferOp's mask
142/// is set to true. No such check is generated under following circumstances:
143/// * xferOp does not have a mask.
144/// * xferOp's mask is not 1D. (In case of (N>1)-D, a subvector of the mask is
145/// computed and attached to the new transfer op in the pattern.)
146/// * The to-be-unpacked dim of xferOp is a broadcast.
147template <typename OpTy>
148static Value generateMaskCheck(OpBuilder &b, OpTy xferOp, Value iv) {
149 if (!xferOp.getMask())
150 return Value();
151 if (xferOp.getMaskType().getRank() != 1)
152 return Value();
153 if (xferOp.isBroadcastDim(0))
154 return Value();
155
156 Location loc = xferOp.getLoc();
157 return vector::ExtractOp::create(b, loc, xferOp.getMask(), iv);
158}
159
160/// Helper function TransferOpConversion and TransferOp1dConversion.
161/// Generate an in-bounds check if the transfer op may go out-of-bounds on the
162/// specified dimension `dim` with the loop iteration variable `iv`.
163/// E.g., when unpacking dimension 0 from:
164/// ```
165/// %vec = vector.transfer_read %A[%a, %b] %cst
166/// : vector<5x4xf32>, memref<?x?xf32>
167/// ```
168/// An if check similar to this will be generated inside the loop:
169/// ```
170/// %d = memref.dim %A, %c0 : memref<?x?xf32>
171/// if (%a + iv < %d) {
172/// (in-bounds case)
173/// } else {
174/// (out-of-bounds case)
175/// }
176/// ```
177///
178/// If the transfer is 1D and has a mask, this function generates a more complex
179/// check also accounts for potentially masked out elements.
180///
181/// This function variant returns the value returned by `inBoundsCase` or
182/// `outOfBoundsCase`. The MLIR type of the return value must be specified in
183/// `resultTypes`.
184template <typename OpTy>
185static Value generateInBoundsCheck(
186 OpBuilder &b, OpTy xferOp, Value iv, std::optional<int64_t> dim,
187 TypeRange resultTypes,
188 function_ref<Value(OpBuilder &, Location)> inBoundsCase,
189 function_ref<Value(OpBuilder &, Location)> outOfBoundsCase = nullptr) {
190 bool hasRetVal = !resultTypes.empty();
191 Value cond; // Condition to be built...
192
193 // Condition check 1: Access in-bounds?
194 bool isBroadcast = !dim; // No in-bounds check for broadcasts.
195 Location loc = xferOp.getLoc();
196 ImplicitLocOpBuilder lb(xferOp.getLoc(), b);
197 if (!xferOp.isDimInBounds(0) && !isBroadcast) {
198 Value memrefDim = vector::createOrFoldDimOp(b, loc, xferOp.getBase(), *dim);
199 AffineExpr d0, d1;
200 bindDims(xferOp.getContext(), d0, d1);
201 Value base = xferOp.getIndices()[*dim];
202 Value memrefIdx =
203 affine::makeComposedAffineApply(b, loc, d0 + d1, {base, iv});
205 Value nonNegative =
206 arith::CmpIOp::create(lb, arith::CmpIPredicate::sge, memrefIdx, zero);
207 Value inRange = arith::CmpIOp::create(lb, arith::CmpIPredicate::slt,
208 memrefIdx, memrefDim);
209 cond = arith::AndIOp::create(lb, nonNegative, inRange);
210 }
211
212 // Condition check 2: Masked in?
213 if (auto maskCond = generateMaskCheck(b, xferOp, iv)) {
214 if (cond)
215 cond = arith::AndIOp::create(lb, cond, maskCond);
216 else
217 cond = maskCond;
218 }
219
220 // If the condition is non-empty, generate an SCF::IfOp.
221 if (cond) {
222 auto check = scf::IfOp::create(
223 lb, cond,
224 /*thenBuilder=*/
225 [&](OpBuilder &b, Location loc) {
226 maybeYieldValue(b, loc, hasRetVal, inBoundsCase(b, loc));
227 },
228 /*elseBuilder=*/
229 [&](OpBuilder &b, Location loc) {
230 if (outOfBoundsCase) {
231 maybeYieldValue(b, loc, hasRetVal, outOfBoundsCase(b, loc));
232 } else {
233 scf::YieldOp::create(b, loc);
234 }
235 });
236
237 return hasRetVal ? check.getResult(0) : Value();
238 }
239
240 // Condition is empty, no need for an SCF::IfOp.
241 return inBoundsCase(b, loc);
242}
243
244/// In this function variant, `inBoundsCase` and `outOfBoundsCase` do not have
245/// a return value. Consequently, this function does not have a return value.
246template <typename OpTy>
247static void generateInBoundsCheck(
248 OpBuilder &b, OpTy xferOp, Value iv, std::optional<int64_t> dim,
249 function_ref<void(OpBuilder &, Location)> inBoundsCase,
250 function_ref<void(OpBuilder &, Location)> outOfBoundsCase = nullptr) {
251 generateInBoundsCheck(
252 b, xferOp, iv, dim, /*resultTypes=*/TypeRange(),
253 /*inBoundsCase=*/
254 [&](OpBuilder &b, Location loc) {
255 inBoundsCase(b, loc);
256 return Value();
257 },
258 /*outOfBoundsCase=*/
259 [&](OpBuilder &b, Location loc) {
260 if (outOfBoundsCase)
261 outOfBoundsCase(b, loc);
262 return Value();
263 });
264}
265
266/// Given an ArrayAttr, return a copy where the first element is dropped.
267static ArrayAttr dropFirstElem(OpBuilder &b, ArrayAttr attr) {
268 if (!attr)
269 return attr;
270 return ArrayAttr::get(b.getContext(), attr.getValue().drop_front());
271}
272
273/// Add the pass label to a vector transfer op if its rank is not the target
274/// rank.
275template <typename OpTy>
276static void maybeApplyPassLabel(OpBuilder &b, OpTy newXferOp,
277 unsigned targetRank) {
278 if (newXferOp.getVectorType().getRank() > targetRank)
279 newXferOp->setDiscardableAttr(kPassLabel, b.getUnitAttr());
280}
281
282namespace lowering_n_d {
283
284/// Helper data structure for data and mask buffers.
285struct BufferAllocs {
286 Value dataBuffer;
287 Value maskBuffer;
288};
289
290// TODO: Parallelism and threadlocal considerations with a ParallelScope trait.
291static Operation *getAutomaticAllocationScope(Operation *op) {
292 Operation *scope =
294 assert(scope && "Expected op to be inside automatic allocation scope");
295 return scope;
296}
297
298/// Allocate temporary buffers for data (vector) and mask (if present).
299template <typename OpTy>
300static BufferAllocs allocBuffers(OpBuilder &b, OpTy xferOp) {
301 Location loc = xferOp.getLoc();
303 Operation *scope = getAutomaticAllocationScope(xferOp);
304 assert(scope->getNumRegions() == 1 &&
305 "AutomaticAllocationScope with >1 regions");
306 b.setInsertionPointToStart(&scope->getRegion(0).front());
307
308 BufferAllocs result;
309 auto bufferType = MemRefType::get({}, xferOp.getVectorType());
310 result.dataBuffer = memref::AllocaOp::create(b, loc, bufferType);
311
312 if (xferOp.getMask()) {
313 auto maskType = MemRefType::get({}, xferOp.getMask().getType());
314 auto maskBuffer = memref::AllocaOp::create(b, loc, maskType);
315 b.setInsertionPoint(xferOp);
316 memref::StoreOp::create(b, loc, xferOp.getMask(), maskBuffer);
317 result.maskBuffer =
318 memref::LoadOp::create(b, loc, maskBuffer, ValueRange());
319 }
320
321 return result;
322}
323
324/// Given a MemRefType with VectorType element type, unpack one dimension from
325/// the VectorType into the MemRefType.
326///
327/// E.g.: memref<9xvector<5x6xf32>> --> memref<9x5xvector<6xf32>>
328static FailureOr<MemRefType> unpackOneDim(MemRefType type) {
329 auto vectorType = dyn_cast<VectorType>(type.getElementType());
330 // Vectors with leading scalable dims are not supported.
331 // It may be possible to support these in future by using dynamic memref dims.
332 if (vectorType.getScalableDims().front())
333 return failure();
334 auto memrefShape = type.getShape();
335 SmallVector<int64_t, 8> newMemrefShape;
336 newMemrefShape.append(memrefShape.begin(), memrefShape.end());
337 newMemrefShape.push_back(vectorType.getDimSize(0));
338 return MemRefType::get(newMemrefShape,
339 VectorType::Builder(vectorType).dropDim(0));
340}
341
342/// Given a transfer op, find the memref from which the mask is loaded. This
343/// is similar to Strategy<TransferWriteOp>::getBuffer.
344template <typename OpTy>
345static Value getMaskBuffer(OpTy xferOp) {
346 assert(xferOp.getMask() && "Expected that transfer op has mask");
347 auto loadOp = xferOp.getMask().template getDefiningOp<memref::LoadOp>();
348 assert(loadOp && "Expected transfer op mask produced by LoadOp");
349 return loadOp.getMemRef();
350}
351
352/// Codegen strategy, depending on the operation.
353template <typename OpTy>
354struct Strategy;
355
356/// Code strategy for vector TransferReadOp.
357template <>
358struct Strategy<TransferReadOp> {
359 /// Find the StoreOp that is used for writing the current TransferReadOp's
360 /// result to the temporary buffer allocation.
361 static memref::StoreOp getStoreOp(TransferReadOp xferOp) {
362 assert(xferOp->hasOneUse() && "Expected exactly one use of TransferReadOp");
363 auto storeOp = dyn_cast<memref::StoreOp>((*xferOp->use_begin()).getOwner());
364 assert(storeOp && "Expected TransferReadOp result used by StoreOp");
365 return storeOp;
366 }
367
368 /// Find the temporary buffer allocation. All labeled TransferReadOps are
369 /// used like this, where %buf is either the buffer allocation or a type cast
370 /// of the buffer allocation:
371 /// ```
372 /// %vec = vector.transfer_read ... { __vector_to_scf_lowering__ } ...
373 /// memref.store %vec, %buf[...] ...
374 /// ```
375 static Value getBuffer(TransferReadOp xferOp) {
376 return getStoreOp(xferOp).getMemRef();
377 }
378
379 /// Retrieve the indices of the current StoreOp that stores into the buffer.
380 static void getBufferIndices(TransferReadOp xferOp,
382 auto storeOp = getStoreOp(xferOp);
383 auto prevIndices = memref::StoreOpAdaptor(storeOp).getIndices();
384 indices.append(prevIndices.begin(), prevIndices.end());
385 }
386
387 /// Rewrite the TransferReadOp, assuming that there are no out-of-bounds
388 /// accesses on the to-be-unpacked dimension.
389 ///
390 /// 1. Generate a new (N-1)-d TransferReadOp using the loop iteration
391 /// variable `iv`.
392 /// 2. Store the result into the (already `vector.type_cast`ed) buffer.
393 ///
394 /// E.g.:
395 /// ```
396 /// %vec = vector.transfer_read %A[%a+%i, %b, %c], %cst
397 /// : memref<?x?x?xf32>, vector<4x3xf32>
398 /// memref.store %vec, %buf[%i] : memref<5xvector<4x3xf32>>
399 /// ```
400 /// Is rewritten to:
401 /// ```
402 /// %casted = vector.type_cast %buf
403 /// : memref<5xvector<4x3xf32>> to memref<5x4xvector<3xf32>>
404 /// for %j = 0 to 4 {
405 /// %vec = vector.transfer_read %A[%a+%i, %b+%j, %c], %cst
406 /// : memref<?x?x?xf32>, vector<3xf32>
407 /// memref.store %vec, %casted[%i, %j] : memref<5x4xvector<3xf32>>
408 /// }
409 /// ```
410 ///
411 /// Note: The loop and type cast are generated in TransferOpConversion.
412 /// The original TransferReadOp and store op are deleted in `cleanup`.
413 /// Note: The `mask` operand is set in TransferOpConversion.
414 static TransferReadOp rewriteOp(OpBuilder &b,
416 TransferReadOp xferOp, Value buffer, Value iv,
417 ValueRange /*loopState*/) {
418 SmallVector<Value, 8> storeIndices;
419 getBufferIndices(xferOp, storeIndices);
420 storeIndices.push_back(iv);
421
422 SmallVector<Value, 8> xferIndices;
423 getXferIndices(b, xferOp, iv, xferIndices);
424
425 Location loc = xferOp.getLoc();
426 auto bufferType = dyn_cast<ShapedType>(buffer.getType());
427 auto vecType = dyn_cast<VectorType>(bufferType.getElementType());
428 auto inBoundsAttr = dropFirstElem(b, xferOp.getInBoundsAttr());
429 auto newXferOp = vector::TransferReadOp::create(
430 b, loc, vecType, xferOp.getBase(), xferIndices,
431 AffineMapAttr::get(unpackedPermutationMap(b, xferOp)),
432 xferOp.getPadding(), Value(), inBoundsAttr);
433
434 maybeApplyPassLabel(b, newXferOp, options.targetRank);
435
436 memref::StoreOp::create(b, loc, newXferOp.getVector(), buffer,
437 storeIndices);
438 return newXferOp;
439 }
440
441 /// Handle out-of-bounds accesses on the to-be-unpacked dimension: Write
442 /// padding value to the temporary buffer.
443 static Value handleOutOfBoundsDim(OpBuilder &b, TransferReadOp xferOp,
444 Value buffer, Value iv,
445 ValueRange /*loopState*/) {
446 SmallVector<Value, 8> storeIndices;
447 getBufferIndices(xferOp, storeIndices);
448 storeIndices.push_back(iv);
449
450 Location loc = xferOp.getLoc();
451 auto bufferType = dyn_cast<ShapedType>(buffer.getType());
452 auto vecType = dyn_cast<VectorType>(bufferType.getElementType());
453 auto vec =
454 vector::BroadcastOp::create(b, loc, vecType, xferOp.getPadding());
455 memref::StoreOp::create(b, loc, vec, buffer, storeIndices);
456
457 return Value();
458 }
459
460 /// Cleanup after rewriting the op.
461 static void cleanup(PatternRewriter &rewriter, TransferReadOp xferOp,
462 scf::ForOp /*forOp*/) {
463 rewriter.eraseOp(getStoreOp(xferOp));
464 rewriter.eraseOp(xferOp);
465 }
466
467 /// Return the initial loop state for the generated scf.for loop.
468 static Value initialLoopState(TransferReadOp xferOp) { return Value(); }
469};
470
471/// Codegen strategy for vector TransferWriteOp.
472template <>
473struct Strategy<TransferWriteOp> {
474 /// Find the temporary buffer allocation. All labeled TransferWriteOps are
475 /// used like this, where %buf is either the buffer allocation or a type cast
476 /// of the buffer allocation:
477 /// ```
478 /// %vec = memref.load %buf[...] ...
479 /// vector.transfer_write %vec ... { __vector_to_scf_lowering__ } ...
480 /// ```
481 static Value getBuffer(TransferWriteOp xferOp) {
482 auto loadOp = xferOp.getVector().getDefiningOp<memref::LoadOp>();
483 assert(loadOp && "Expected transfer op vector produced by LoadOp");
484 return loadOp.getMemRef();
485 }
486
487 /// Retrieve the indices of the current LoadOp that loads from the buffer.
488 static void getBufferIndices(TransferWriteOp xferOp,
490 auto loadOp = xferOp.getVector().getDefiningOp<memref::LoadOp>();
491 auto prevIndices = memref::LoadOpAdaptor(loadOp).getIndices();
492 indices.append(prevIndices.begin(), prevIndices.end());
493 }
494
495 /// Rewrite the TransferWriteOp, assuming that there are no out-of-bounds
496 /// accesses on the to-be-unpacked dimension.
497 ///
498 /// 1. Load an (N-1)-d vector from the (already `vector.type_cast`ed) buffer,
499 /// using the loop iteration variable `iv`.
500 /// 2. Generate a new (N-1)-d TransferWriteOp, writing the loaded vector back
501 /// to memory.
502 ///
503 /// Note: For more details, see comments on Strategy<TransferReadOp>.
504 static TransferWriteOp rewriteOp(OpBuilder &b,
506 TransferWriteOp xferOp, Value buffer,
507 Value iv, ValueRange loopState) {
508 SmallVector<Value, 8> loadIndices;
509 getBufferIndices(xferOp, loadIndices);
510 loadIndices.push_back(iv);
511
512 SmallVector<Value, 8> xferIndices;
513 getXferIndices(b, xferOp, iv, xferIndices);
514
515 Location loc = xferOp.getLoc();
516 auto vec = memref::LoadOp::create(b, loc, buffer, loadIndices);
517 auto inBoundsAttr = dropFirstElem(b, xferOp.getInBoundsAttr());
518 auto source = loopState.empty() ? xferOp.getBase() : loopState[0];
519 Type type = isTensorOp(xferOp) ? xferOp.getShapedType() : Type();
520 auto newXferOp = vector::TransferWriteOp::create(
521 b, loc, type, vec, source, xferIndices,
522 AffineMapAttr::get(unpackedPermutationMap(b, xferOp)), Value(),
523 inBoundsAttr);
524
525 maybeApplyPassLabel(b, newXferOp, options.targetRank);
526
527 return newXferOp;
528 }
529
530 /// Handle out-of-bounds accesses on the to-be-unpacked dimension.
531 static Value handleOutOfBoundsDim(OpBuilder &b, TransferWriteOp xferOp,
532 Value buffer, Value iv,
533 ValueRange loopState) {
534 return isTensorOp(xferOp) ? loopState[0] : Value();
535 }
536
537 /// Cleanup after rewriting the op.
538 static void cleanup(PatternRewriter &rewriter, TransferWriteOp xferOp,
539 scf::ForOp forOp) {
540 if (isTensorOp(xferOp)) {
541 assert(forOp->getNumResults() == 1 && "Expected one for loop result");
542 rewriter.replaceOp(xferOp, forOp->getResult(0));
543 } else {
544 rewriter.eraseOp(xferOp);
545 }
546 }
547
548 /// Return the initial loop state for the generated scf.for loop.
549 static Value initialLoopState(TransferWriteOp xferOp) {
550 return isTensorOp(xferOp) ? xferOp.getBase() : Value();
551 }
552};
553
554template <typename OpTy>
555static LogicalResult checkPrepareXferOp(OpTy xferOp, PatternRewriter &rewriter,
557 if (xferOp->hasDiscardableAttr(kPassLabel))
558 return rewriter.notifyMatchFailure(
559 xferOp, "kPassLabel is present (vector-to-scf lowering in progress)");
560 if (xferOp.getVectorType().getRank() <= options.targetRank)
561 return rewriter.notifyMatchFailure(
562 xferOp, "xferOp vector rank <= transformation target rank");
563 if (xferOp.getVectorType().getScalableDims().front())
564 return rewriter.notifyMatchFailure(
565 xferOp, "Unpacking of the leading dimension into the memref is not yet "
566 "supported for scalable dims");
567 if (isTensorOp(xferOp) && !options.lowerTensors)
568 return rewriter.notifyMatchFailure(
569 xferOp, "Unpacking for tensors has been disabled.");
570 if (xferOp.getVectorType().getElementType() !=
571 xferOp.getShapedType().getElementType())
572 return rewriter.notifyMatchFailure(
573 xferOp, "Mismatching source and destination element types.");
574 Operation *op = xferOp.getOperation();
576 return rewriter.notifyMatchFailure(
577 xferOp, "xferOp is not inside an automatic allocation scope");
578
579 return success();
580}
581
582/// Prepare a TransferReadOp for progressive lowering.
583///
584/// 1. Allocate a temporary buffer.
585/// 2. Label the TransferReadOp, marking it eligible for progressive lowering.
586/// 3. Store the result of the TransferReadOp into the temporary buffer.
587/// 4. Load the result from the temporary buffer and replace all uses of the
588/// original TransferReadOp with this load.
589///
590/// E.g.:
591/// ```
592/// %vec = vector.transfer_read %A[%a, %b, %c], %cst
593/// : vector<5x4xf32>, memref<?x?x?xf32>
594/// ```
595/// is rewritten to:
596/// ```
597/// %0 = memref.alloca() : memref<vector<5x4xf32>>
598/// %1 = vector.transfer_read %A[%a, %b, %c], %cst
599/// { __vector_to_scf_lowering__ } : vector<5x4xf32>, memref<?x?x?xf32>
600/// memref.store %1, %0[] : memref<vector<5x4xf32>>
601/// %vec = memref.load %0[] : memref<vector<5x4xf32>>
602/// ```
603///
604/// Note: A second temporary buffer may be allocated for the `mask` operand.
605struct PrepareTransferReadConversion
606 : public VectorToSCFPattern<TransferReadOp> {
607 using VectorToSCFPattern<TransferReadOp>::VectorToSCFPattern;
608
609 LogicalResult matchAndRewrite(TransferReadOp xferOp,
610 PatternRewriter &rewriter) const override {
611 if (checkPrepareXferOp(xferOp, rewriter, options).failed())
612 return rewriter.notifyMatchFailure(
613 xferOp, "checkPrepareXferOp conditions not met!");
614
615 auto buffers = allocBuffers(rewriter, xferOp);
616 auto *newXfer = rewriter.clone(*xferOp.getOperation());
617 newXfer->setDiscardableAttr(kPassLabel, rewriter.getUnitAttr());
618 if (xferOp.getMask()) {
619 dyn_cast<TransferReadOp>(newXfer).getMaskMutable().assign(
620 buffers.maskBuffer);
621 }
622
623 Location loc = xferOp.getLoc();
624 memref::StoreOp::create(rewriter, loc, newXfer->getResult(0),
625 buffers.dataBuffer);
626 rewriter.replaceOpWithNewOp<memref::LoadOp>(xferOp, buffers.dataBuffer,
627 ValueRange{});
628
629 return success();
630 }
631};
632
633/// Prepare a TransferWriteOp for progressive lowering.
634///
635/// 1. Allocate a temporary buffer.
636/// 2. Store the vector into the buffer.
637/// 3. Load the vector from the buffer again.
638/// 4. Use the loaded vector as a TransferWriteOp operand and label the op,
639/// marking it eligible for progressive lowering via TransferOpConversion.
640///
641/// E.g.:
642/// ```
643/// vector.transfer_write %vec, %A[%a, %b, %c]
644/// : vector<5x4xf32>, memref<?x?x?xf32>
645/// ```
646/// is rewritten to:
647/// ```
648/// %0 = memref.alloca() : memref<vector<5x4xf32>>
649/// memref.store %vec, %0[] : memref<vector<5x4xf32>>
650/// %1 = memref.load %0[] : memref<vector<5x4xf32>>
651/// vector.transfer_write %1, %A[%a, %b, %c] { __vector_to_scf_lowering__ }
652/// : vector<5x4xf32>, memref<?x?x?xf32>
653/// ```
654///
655/// Note: A second temporary buffer may be allocated for the `mask` operand.
656struct PrepareTransferWriteConversion
657 : public VectorToSCFPattern<TransferWriteOp> {
658 using VectorToSCFPattern<TransferWriteOp>::VectorToSCFPattern;
659
660 LogicalResult matchAndRewrite(TransferWriteOp xferOp,
661 PatternRewriter &rewriter) const override {
662 if (checkPrepareXferOp(xferOp, rewriter, options).failed())
663 return rewriter.notifyMatchFailure(
664 xferOp, "checkPrepareXferOp conditions not met!");
665
666 Location loc = xferOp.getLoc();
667 auto buffers = allocBuffers(rewriter, xferOp);
668 memref::StoreOp::create(rewriter, loc, xferOp.getVector(),
669 buffers.dataBuffer);
670 auto loadedVec =
671 memref::LoadOp::create(rewriter, loc, buffers.dataBuffer, ValueRange{});
672 rewriter.modifyOpInPlace(xferOp, [&]() {
673 xferOp.getValueToStoreMutable().assign(loadedVec);
674 xferOp->setDiscardableAttr(kPassLabel, rewriter.getUnitAttr());
675 });
676
677 if (xferOp.getMask()) {
678 rewriter.modifyOpInPlace(xferOp, [&]() {
679 xferOp.getMaskMutable().assign(buffers.maskBuffer);
680 });
681 }
682
683 return success();
684 }
685};
686
687/// Decompose a n-D PrintOp into a loop of elementary/scalar prints. This allows
688/// printing both 1D scalable vectors and n-D fixed size vectors.
689///
690/// E.g.:
691/// ```
692/// vector.print %v : vector<[4]xi32>
693/// ```
694/// is rewritten to:
695/// ```
696/// %c0 = arith.constant 0 : index
697/// %c4 = arith.constant 4 : index
698/// %c1 = arith.constant 1 : index
699/// %vscale = vector.vscale
700/// %length = arith.muli %vscale, %c4 : index
701/// %lastIndex = arith.subi %length, %c1 : index
702/// vector.print punctuation <open>
703/// scf.for %i = %c0 to %length step %c1 {
704/// %el = vector.extract %v[%i] : i32 from vector<[4]xi32>
705/// vector.print %el : i32 punctuation <no_punctuation>
706/// %notLastIndex = arith.cmpi ult, %i, %lastIndex : index
707/// scf.if %notLastIndex {
708/// vector.print punctuation <comma>
709/// }
710/// }
711/// vector.print punctuation <close>
712/// vector.print
713/// ```
714struct DecomposePrintOpConversion : public VectorToSCFPattern<vector::PrintOp> {
715 using VectorToSCFPattern<vector::PrintOp>::VectorToSCFPattern;
716 LogicalResult matchAndRewrite(vector::PrintOp printOp,
717 PatternRewriter &rewriter) const override {
718 if (!printOp.getSource())
719 return failure();
720
721 VectorType vectorType = dyn_cast<VectorType>(printOp.getPrintType());
722 if (!vectorType)
723 return failure();
724
725 // Currently >= 2D scalable vectors are not supported.
726 // These can't be lowered to LLVM (as LLVM does not support scalable vectors
727 // of scalable vectors), and due to limitations of current ops can't be
728 // indexed with SSA values or flattened. This may change after
729 // https://reviews.llvm.org/D155034, though there still needs to be a path
730 // for lowering to LLVM.
731 if (vectorType.getRank() > 1 && vectorType.isScalable())
732 return failure();
733
734 auto loc = printOp.getLoc();
735 auto value = printOp.getSource();
736
737 if (auto intTy = dyn_cast<IntegerType>(vectorType.getElementType())) {
738 // Oddly sized integers are (somewhat) buggy on a lot of backends, so to
739 // avoid issues extend them to a more standard size.
740 // https://github.com/llvm/llvm-project/issues/30613
741 auto width = intTy.getWidth();
742 auto legalWidth = llvm::NextPowerOf2(std::max(8u, width) - 1);
743 auto legalIntTy = IntegerType::get(rewriter.getContext(), legalWidth,
744 intTy.getSignedness());
745 // arith can only take signless integers, so we must cast back and forth.
746 auto signlessSourceVectorType =
747 vectorType.cloneWith({}, getIntTypeWithSignlessSemantics(intTy));
748 auto signlessTargetVectorType =
749 vectorType.cloneWith({}, getIntTypeWithSignlessSemantics(legalIntTy));
750 auto targetVectorType = vectorType.cloneWith({}, legalIntTy);
751 value = vector::BitCastOp::create(rewriter, loc, signlessSourceVectorType,
752 value);
753 if (value.getType() != signlessTargetVectorType) {
754 if (width == 1 || intTy.isUnsigned())
755 value = arith::ExtUIOp::create(rewriter, loc,
756 signlessTargetVectorType, value);
757 else
758 value = arith::ExtSIOp::create(rewriter, loc,
759 signlessTargetVectorType, value);
760 }
761 value = vector::BitCastOp::create(rewriter, loc, targetVectorType, value);
762 vectorType = targetVectorType;
763 }
764
765 auto scalableDimensions = vectorType.getScalableDims();
766 auto shape = vectorType.getShape();
767 constexpr int64_t singletonShape[] = {1};
768 if (vectorType.getRank() == 0)
769 shape = singletonShape;
770
771 if (vectorType.getRank() != 1) {
772 // Flatten n-D vectors to 1D. This is done to allow indexing with a
773 // non-constant value.
774 int64_t flatLength = llvm::product_of(shape);
775 auto flatVectorType =
776 VectorType::get({flatLength}, vectorType.getElementType());
777 value = vector::ShapeCastOp::create(rewriter, loc, flatVectorType, value);
778 }
779
780 vector::PrintOp firstClose;
781 SmallVector<Value, 8> loopIndices;
782 for (unsigned d = 0; d < shape.size(); d++) {
783 // Setup loop bounds and step.
784 Value lowerBound = arith::ConstantIndexOp::create(rewriter, loc, 0);
785 Value upperBound =
786 arith::ConstantIndexOp::create(rewriter, loc, shape[d]);
787 Value step = arith::ConstantIndexOp::create(rewriter, loc, 1);
788 if (!scalableDimensions.empty() && scalableDimensions[d]) {
789 auto vscale = vector::VectorScaleOp::create(rewriter, loc,
790 rewriter.getIndexType());
791 upperBound = arith::MulIOp::create(rewriter, loc, upperBound, vscale);
792 }
793 auto lastIndex = arith::SubIOp::create(rewriter, loc, upperBound, step);
794
795 // Create a loop to print the elements surrounded by parentheses.
796 vector::PrintOp::create(rewriter, loc, vector::PrintPunctuation::Open);
797 auto loop =
798 scf::ForOp::create(rewriter, loc, lowerBound, upperBound, step);
799 auto printClose = vector::PrintOp::create(
800 rewriter, loc, vector::PrintPunctuation::Close);
801 if (!firstClose)
802 firstClose = printClose;
803
804 auto loopIdx = loop.getInductionVar();
805 loopIndices.push_back(loopIdx);
806
807 // Print a comma after all but the last element.
808 rewriter.setInsertionPointToStart(loop.getBody());
809 auto notLastIndex = arith::CmpIOp::create(
810 rewriter, loc, arith::CmpIPredicate::ult, loopIdx, lastIndex);
811 scf::IfOp::create(rewriter, loc, notLastIndex,
812 [&](OpBuilder &builder, Location loc) {
813 vector::PrintOp::create(
814 builder, loc, vector::PrintPunctuation::Comma);
815 scf::YieldOp::create(builder, loc);
816 });
817
818 rewriter.setInsertionPointToStart(loop.getBody());
819 }
820
821 // Compute the flattened index.
822 // Note: For the > rank 1 vectors this assumes non-scalable.
823 Value flatIndex;
824 auto currentStride = 1;
825 for (int d = shape.size() - 1; d >= 0; d--) {
826 auto stride =
827 arith::ConstantIndexOp::create(rewriter, loc, currentStride);
828 auto index = arith::MulIOp::create(rewriter, loc, stride, loopIndices[d]);
829 if (flatIndex)
830 flatIndex = arith::AddIOp::create(rewriter, loc, flatIndex, index);
831 else
832 flatIndex = index;
833 currentStride *= shape[d];
834 }
835
836 // Print the scalar elements in the inner most loop.
837 auto element = vector::ExtractOp::create(rewriter, loc, value, flatIndex);
838 vector::PrintOp::create(rewriter, loc, element,
839 vector::PrintPunctuation::NoPunctuation);
840
841 rewriter.setInsertionPointAfter(firstClose);
842 vector::PrintOp::create(rewriter, loc, printOp.getPunctuation());
843 rewriter.eraseOp(printOp);
844 return success();
845 }
846
847 static IntegerType getIntTypeWithSignlessSemantics(IntegerType intTy) {
848 return IntegerType::get(intTy.getContext(), intTy.getWidth(),
849 IntegerType::Signless);
850 };
851};
852
853/// Progressive lowering of vector transfer ops: Unpack one dimension.
854///
855/// 1. Unpack one dimension from the current buffer type and cast the buffer
856/// to that new type. E.g.:
857/// ```
858/// %vec = memref.load %0[%1] : memref<5xvector<4x3xf32>>
859/// vector.transfer_write %vec ...
860/// ```
861/// The following cast is generated:
862/// ```
863/// %casted = vector.type_cast %0
864/// : memref<5xvector<4x3xf32>> to memref<5x4xvector<3xf32>>
865/// ```
866/// 2. Generate a for loop and rewrite the transfer op according to the
867/// corresponding Strategy<OpTy>. If the to-be-unpacked dimension can be
868/// out-of-bounds, generate an if-check and handle both cases separately.
869/// 3. Clean up according to the corresponding Strategy<OpTy>.
870///
871/// Note: If the transfer op is a TransferWriteOp and operates on a tensor
872/// source (as opposed to a memref source), then each iteration of the generated
873/// scf.for loop yields the new tensor value. E.g.:
874/// ```
875/// %result = scf.for i = 0 to 5 {
876/// %0 = memref.load %buffer[i] : memref<5xvector<4x3xf32>>
877/// %1 = vector.transfer_write %0, %source[...]
878/// : vector<4x3xf32>, tensor<5x4x3xf32>
879/// scf.yield %1 : tensor<5x4x3xf32>
880/// }
881/// ```
882template <typename OpTy>
883struct TransferOpConversion : public VectorToSCFPattern<OpTy> {
884 using VectorToSCFPattern<OpTy>::VectorToSCFPattern;
885
886 void initialize() {
887 // This pattern recursively unpacks one dimension at a time. The recursion
888 // bounded as the rank is strictly decreasing.
889 this->setHasBoundedRewriteRecursion();
890 }
891
892 static void getMaskBufferLoadIndices(OpTy xferOp, Value castedMaskBuffer,
893 SmallVectorImpl<Value> &loadIndices,
894 Value iv) {
895 assert(xferOp.getMask() && "Expected transfer op to have mask");
896
897 // Add load indices from the previous iteration.
898 // The mask buffer depends on the permutation map, which makes determining
899 // the indices quite complex, so this is why we need to "look back" to the
900 // previous iteration to find the right indices.
901 Value maskBuffer = getMaskBuffer(xferOp);
902 for (Operation *user : maskBuffer.getUsers()) {
903 // If there is no previous load op, then the indices are empty.
904 if (auto loadOp = dyn_cast<memref::LoadOp>(user)) {
905 Operation::operand_range prevIndices = loadOp.getIndices();
906 loadIndices.append(prevIndices.begin(), prevIndices.end());
907 break;
908 }
909 }
910
911 // In case of broadcast: Use same indices to load from memref
912 // as before.
913 if (!xferOp.isBroadcastDim(0))
914 loadIndices.push_back(iv);
915 }
916
917 LogicalResult matchAndRewrite(OpTy xferOp,
918 PatternRewriter &rewriter) const override {
919 if (!xferOp->hasDiscardableAttr(kPassLabel))
920 return rewriter.notifyMatchFailure(
921 xferOp, "kPassLabel is present (progressing lowering in progress)");
922
923 // Find and cast data buffer. How the buffer can be found depends on OpTy.
924 ImplicitLocOpBuilder locB(xferOp.getLoc(), rewriter);
925 Value dataBuffer = Strategy<OpTy>::getBuffer(xferOp);
926 auto dataBufferType = dyn_cast<MemRefType>(dataBuffer.getType());
927 FailureOr<MemRefType> castedDataType = unpackOneDim(dataBufferType);
928 if (failed(castedDataType))
929 return rewriter.notifyMatchFailure(xferOp,
930 "Failed to unpack one vector dim.");
931
932 auto castedDataBuffer =
933 vector::TypeCastOp::create(locB, *castedDataType, dataBuffer);
934
935 // If the xferOp has a mask: Find and cast mask buffer.
936 Value castedMaskBuffer;
937 if (xferOp.getMask()) {
938 Value maskBuffer = getMaskBuffer(xferOp);
939 if (xferOp.isBroadcastDim(0) || xferOp.getMaskType().getRank() == 1) {
940 // Do not unpack a dimension of the mask, if:
941 // * To-be-unpacked transfer op dimension is a broadcast.
942 // * Mask is 1D, i.e., the mask cannot be further unpacked.
943 // (That means that all remaining dimensions of the transfer op must
944 // be broadcasted.)
945 castedMaskBuffer = maskBuffer;
946 } else {
947 // It's safe to assume the mask buffer can be unpacked if the data
948 // buffer was unpacked.
949 auto maskBufferType = cast<MemRefType>(maskBuffer.getType());
950 MemRefType castedMaskType = *unpackOneDim(maskBufferType);
951 castedMaskBuffer =
952 vector::TypeCastOp::create(locB, castedMaskType, maskBuffer);
953 }
954 }
955
956 // Loop bounds and step.
957 auto lb = arith::ConstantIndexOp::create(locB, 0);
959 locB, castedDataType->getDimSize(castedDataType->getRank() - 1));
960 auto step = arith::ConstantIndexOp::create(locB, 1);
961 // TransferWriteOps that operate on tensors return the modified tensor and
962 // require a loop state.
963 auto loopState = Strategy<OpTy>::initialLoopState(xferOp);
964
965 // Generate for loop.
966 auto result = scf::ForOp::create(
967 locB, lb, ub, step, loopState ? ValueRange(loopState) : ValueRange(),
968 [&](OpBuilder &b, Location loc, Value iv, ValueRange loopState) {
969 Type stateType = loopState.empty() ? Type() : loopState[0].getType();
970
971 auto result = generateInBoundsCheck(
972 b, xferOp, iv, unpackedDim(xferOp),
973 stateType ? TypeRange(stateType) : TypeRange(),
974 /*inBoundsCase=*/
975 [&](OpBuilder &b, Location loc) {
976 // Create new transfer op.
977 OpTy newXfer = Strategy<OpTy>::rewriteOp(
978 b, this->options, xferOp, castedDataBuffer, iv, loopState);
979
980 // If old transfer op has a mask: Set mask on new transfer op.
981 // Special case: If the mask of the old transfer op is 1D and
982 // the unpacked dim is not a broadcast, no mask is needed on
983 // the new transfer op.
984 if (xferOp.getMask() && (xferOp.isBroadcastDim(0) ||
985 xferOp.getMaskType().getRank() > 1)) {
987 b.setInsertionPoint(newXfer); // Insert load before newXfer.
988
989 SmallVector<Value, 8> loadIndices;
990 getMaskBufferLoadIndices(xferOp, castedMaskBuffer,
991 loadIndices, iv);
992 auto mask = memref::LoadOp::create(b, loc, castedMaskBuffer,
993 loadIndices);
994 rewriter.modifyOpInPlace(newXfer, [&]() {
995 newXfer.getMaskMutable().assign(mask);
996 });
997 }
998
999 return loopState.empty() ? Value() : newXfer->getResult(0);
1000 },
1001 /*outOfBoundsCase=*/
1002 [&](OpBuilder &b, Location /*loc*/) {
1003 return Strategy<OpTy>::handleOutOfBoundsDim(
1004 b, xferOp, castedDataBuffer, iv, loopState);
1005 });
1006
1007 maybeYieldValue(b, loc, !loopState.empty(), result);
1008 });
1009
1010 Strategy<OpTy>::cleanup(rewriter, xferOp, result);
1011 return success();
1012 }
1013};
1014
1015/// Retrieves the dimensions sizes of a mask. Currently supports CreateMaskOp
1016/// and ConstantMaskOp.
1017template <typename VscaleConstantBuilder>
1018static FailureOr<SmallVector<OpFoldResult>>
1019getMaskDimSizes(Value mask, VscaleConstantBuilder &createVscaleMultiple) {
1020 if (!mask)
1022 if (auto createMaskOp = mask.getDefiningOp<vector::CreateMaskOp>()) {
1023 return llvm::map_to_vector(createMaskOp.getOperands(), [](Value dimSize) {
1024 return OpFoldResult(dimSize);
1025 });
1026 }
1027 if (auto constantMask = mask.getDefiningOp<vector::ConstantMaskOp>()) {
1028 int dimIdx = 0;
1029 VectorType maskType = constantMask.getVectorType();
1030 auto indexType = IndexType::get(mask.getContext());
1031 return llvm::map_to_vector(
1032 constantMask.getMaskDimSizes(), [&](int64_t dimSize) {
1033 // A scalable dim in a constant_mask means vscale x dimSize.
1034 if (maskType.getScalableDims()[dimIdx++])
1035 return OpFoldResult(createVscaleMultiple(dimSize));
1036 return OpFoldResult(IntegerAttr::get(indexType, dimSize));
1037 });
1038 }
1039 return failure();
1040}
1041
1042/// Scalable vector lowering of transfer_write(transpose). This lowering only
1043/// supports rank 2 (scalable) vectors, but can be used in conjunction with
1044/// `UnrollTransferWriteConversion` to support n-D cases. The unroll conversion
1045/// unrolls until the first scalable dimension.
1046///
1047/// Example:
1048///
1049/// BEFORE:
1050/// ```mlir
1051/// %transpose = vector.transpose %vec, [1, 0]
1052/// : vector<4x[4]xf32> to vector<[4]x4xf32>
1053/// vector.transfer_write %transpose, %dest[%i, %j] {in_bounds = [true, true]}
1054/// : vector<[4]x4xf32>, memref<?x?xf32>
1055/// ```
1056///
1057/// AFTER:
1058/// ```mlir
1059/// %c1 = arith.constant 1 : index
1060/// %c4 = arith.constant 4 : index
1061/// %c0 = arith.constant 0 : index
1062/// %0 = vector.extract %arg0[0] : vector<[4]xf32> from vector<4x[4]xf32>
1063/// %1 = vector.extract %arg0[1] : vector<[4]xf32> from vector<4x[4]xf32>
1064/// %2 = vector.extract %arg0[2] : vector<[4]xf32> from vector<4x[4]xf32>
1065/// %3 = vector.extract %arg0[3] : vector<[4]xf32> from vector<4x[4]xf32>
1066/// %vscale = vector.vscale
1067/// %c4_vscale = arith.muli %vscale, %c4 : index
1068/// scf.for %idx = %c0 to %c4_vscale step %c1 {
1069/// %4 = vector.extract %0[%idx] : f32 from vector<[4]xf32>
1070/// %5 = vector.extract %1[%idx] : f32 from vector<[4]xf32>
1071/// %6 = vector.extract %2[%idx] : f32 from vector<[4]xf32>
1072/// %7 = vector.extract %3[%idx] : f32 from vector<[4]xf32>
1073/// %slice_i = affine.apply #map(%idx)[%i]
1074/// %slice = vector.from_elements %4, %5, %6, %7 : vector<4xf32>
1075/// vector.transfer_write %slice, %arg1[%slice_i, %j] {in_bounds = [true]}
1076/// : vector<4xf32>, memref<?x?xf32>
1077/// }
1078/// ```
1079struct ScalableTransposeTransferWriteConversion
1080 : VectorToSCFPattern<vector::TransferWriteOp> {
1081 using VectorToSCFPattern::VectorToSCFPattern;
1082
1083 LogicalResult matchAndRewrite(TransferWriteOp writeOp,
1084 PatternRewriter &rewriter) const override {
1085 if (failed(checkLowerTensors(writeOp, rewriter)))
1086 return failure();
1087
1088 VectorType vectorType = writeOp.getVectorType();
1089
1090 // Note: By comparing the scalable dims to an ArrayRef of length two this
1091 // implicitly checks the rank (is also two).
1092 ArrayRef<bool> scalableFlags = vectorType.getScalableDims();
1093 if (scalableFlags != ArrayRef<bool>{true, false}) {
1094 return rewriter.notifyMatchFailure(
1095 writeOp, "expected vector of the form vector<[N]xMxty>");
1096 }
1097
1098 auto permutationMap = writeOp.getPermutationMap();
1099 if (!permutationMap.isIdentity()) {
1100 return rewriter.notifyMatchFailure(
1101 writeOp, "non-identity permutations are unsupported (lower first)");
1102 }
1103
1104 // Note: This pattern is only lowering the leading dimension (to a loop),
1105 // so we only check if the leading dimension is in bounds. The in-bounds
1106 // attribute for the trailing dimension will be propagated.
1107 if (!writeOp.isDimInBounds(0)) {
1108 return rewriter.notifyMatchFailure(
1109 writeOp, "out-of-bounds dims are unsupported (use masking)");
1110 }
1111
1112 Value vector = writeOp.getVector();
1113 auto transposeOp = vector.getDefiningOp<vector::TransposeOp>();
1114 if (!transposeOp ||
1115 transposeOp.getPermutation() != ArrayRef<int64_t>{1, 0}) {
1116 return rewriter.notifyMatchFailure(writeOp, "source not transpose");
1117 }
1118
1119 auto loc = writeOp.getLoc();
1120 auto createVscaleMultiple =
1121 vector::makeVscaleConstantBuilder(rewriter, loc);
1122
1123 auto maskDims = getMaskDimSizes(writeOp.getMask(), createVscaleMultiple);
1124 if (failed(maskDims)) {
1125 return rewriter.notifyMatchFailure(writeOp,
1126 "failed to resolve mask dims");
1127 }
1128
1129 int64_t fixedDimSize = vectorType.getDimSize(1);
1130 auto fixedDimOffsets = llvm::seq(fixedDimSize);
1131
1132 // Extract all slices from the source of the transpose.
1133 auto transposeSource = transposeOp.getVector();
1134 SmallVector<Value> transposeSourceSlices =
1135 llvm::map_to_vector(fixedDimOffsets, [&](int64_t idx) -> Value {
1136 return vector::ExtractOp::create(rewriter, loc, transposeSource, idx);
1137 });
1138
1139 // Loop bounds and step.
1140 auto lb = arith::ConstantIndexOp::create(rewriter, loc, 0);
1141 auto ub =
1142 maskDims->empty()
1143 ? Value(createVscaleMultiple(vectorType.getDimSize(0)))
1144 : vector::getAsValues(rewriter, loc, maskDims->front()).front();
1145 auto step = arith::ConstantIndexOp::create(rewriter, loc, 1);
1146
1147 // Generate a new mask for the slice.
1148 VectorType sliceType = VectorType::Builder(vectorType).dropDim(0);
1149 Value sliceMask = nullptr;
1150 if (!maskDims->empty()) {
1151 sliceMask = vector::CreateMaskOp::create(
1152 rewriter, loc, sliceType.clone(rewriter.getI1Type()),
1153 ArrayRef<OpFoldResult>(*maskDims).drop_front());
1154 }
1155
1156 Value initDest = isTensorOp(writeOp) ? writeOp.getBase() : Value{};
1157 ValueRange initLoopArgs = initDest ? initDest : ValueRange{};
1158 auto result = scf::ForOp::create(
1159 rewriter, loc, lb, ub, step, initLoopArgs,
1160 [&](OpBuilder &b, Location loc, Value iv, ValueRange loopIterArgs) {
1161 // Indices for the new transfer op.
1162 SmallVector<Value, 8> xferIndices;
1163 getXferIndices(b, writeOp, iv, xferIndices);
1164
1165 // Extract a transposed slice from the source vector.
1166 SmallVector<Value> transposeElements =
1167 llvm::map_to_vector(fixedDimOffsets, [&](int64_t idx) -> Value {
1168 return vector::ExtractOp::create(
1169 b, loc, transposeSourceSlices[idx], iv);
1170 });
1171 auto sliceVec = vector::FromElementsOp::create(b, loc, sliceType,
1172 transposeElements);
1173
1174 // Create the transfer_write for the slice.
1175 Value dest =
1176 loopIterArgs.empty() ? writeOp.getBase() : loopIterArgs.front();
1177 auto newWriteOp = vector::TransferWriteOp::create(
1178 b, loc, sliceVec, dest, xferIndices,
1179 ArrayRef<bool>(writeOp.getInBoundsValues()).drop_front());
1180 if (sliceMask)
1181 newWriteOp.getMaskMutable().assign(sliceMask);
1182
1183 // Yield from the loop.
1184 scf::YieldOp::create(b, loc,
1185 loopIterArgs.empty() ? ValueRange{}
1186 : newWriteOp.getResult());
1187 });
1188
1189 if (isTensorOp(writeOp))
1190 rewriter.replaceOp(writeOp, result);
1191 else
1192 rewriter.eraseOp(writeOp);
1193
1194 return success();
1195 }
1196};
1197
1198} // namespace lowering_n_d
1199
1201
1202/// If the original transfer op has a mask, compute the mask of the new transfer
1203/// op (for the current iteration `i`) and assign it.
1204template <typename OpTy>
1205static void maybeAssignMask(OpBuilder &b, OpTy xferOp, OpTy newXferOp,
1206 int64_t i) {
1207 if (!xferOp.getMask())
1208 return;
1209
1210 if (xferOp.isBroadcastDim(0)) {
1211 // To-be-unpacked dimension is a broadcast, which does not have a
1212 // corresponding mask dimension. Mask attribute remains unchanged.
1213 newXferOp.getMaskMutable().assign(xferOp.getMask());
1214 return;
1215 }
1216
1217 if (xferOp.getMaskType().getRank() > 1) {
1218 // Unpack one dimension of the mask.
1220 b.setInsertionPoint(newXferOp); // Insert load before newXfer.
1221
1223 Location loc = xferOp.getLoc();
1224 auto newMask = vector::ExtractOp::create(b, loc, xferOp.getMask(), indices);
1225 newXferOp.getMaskMutable().assign(newMask);
1226 }
1227
1228 // If we end up here: The mask of the old transfer op is 1D and the unpacked
1229 // dim is not a broadcast, so no mask is needed on the new transfer op.
1230 // `generateInBoundsCheck` will have evaluated the mask already.
1231}
1232
1233/// Progressive lowering of vector TransferReadOp with unrolling: Unpack one
1234/// dimension. This is similar to TransferOpConversion<TransferReadOp>, but no
1235/// memref buffer is allocated and the SCF loop is fully unrolled.
1236///
1237/// ```
1238/// E.g.:
1239/// ```
1240/// %vec = vector.transfer_read %A[%a, %b, %c], %padding
1241/// : memref<?x?x?xf32>, vector<5x4xf32>
1242/// ```
1243/// is rewritten to IR such as (simplified):
1244/// ```
1245/// %v_init = splat %padding : vector<5x4xf32>
1246/// %tmp0 = vector.transfer_read %A[%a, %b, %c], %padding
1247/// : memref<?x?x?xf32>, vector<4xf32>
1248/// %v0 = vector.insert %tmp0, %v_init[0] : vector<4xf32> into vector<5x4xf32>
1249/// %tmp1 = vector.transfer_read %A[%a, %b + 1, %c], %padding
1250/// : memref<?x?x?xf32>, vector<4xf32>
1251/// %v1 = vector.insert %tmp1, %v0[1] : vector<4xf32> into vector<5x4xf32>
1252/// ...
1253/// %tmp4 = vector.transfer_read %A[%a, %b + 4, %c], %padding
1254/// : memref<?x?x?xf32>, vector<4xf32>
1255/// %vec = vector.insert %tmp1, %v3[4] : vector<4xf32> into vector<5x4xf32>
1256/// ```
1257///
1258/// Note: As an optimization, if the result of the original TransferReadOp
1259/// was directly inserted into another vector, no new %v_init vector is created.
1260/// Instead, the new TransferReadOp results are inserted into that vector.
1261struct UnrollTransferReadConversion
1262 : public VectorToSCFPattern<TransferReadOp> {
1263 using VectorToSCFPattern<TransferReadOp>::VectorToSCFPattern;
1264
1265 void initialize() {
1266 // This pattern recursively unpacks one dimension at a time. The recursion
1267 // bounded as the rank is strictly decreasing.
1268 setHasBoundedRewriteRecursion();
1269 }
1270
1271 /// Get or build the vector into which the newly created TransferReadOp
1272 /// results are inserted.
1273 Value buildResultVector(PatternRewriter &rewriter,
1274 TransferReadOp xferOp) const {
1275 if (auto insertOp = getInsertOp(xferOp))
1276 return insertOp.getDest();
1277 Location loc = xferOp.getLoc();
1278 return vector::BroadcastOp::create(rewriter, loc, xferOp.getVectorType(),
1279 xferOp.getPadding());
1280 }
1281
1282 /// If the result of the TransferReadOp has exactly one user, which is a
1283 /// vector::InsertOp, return that operation.
1284 vector::InsertOp getInsertOp(TransferReadOp xferOp) const {
1285 if (xferOp->hasOneUse()) {
1286 Operation *xferOpUser = *xferOp->getUsers().begin();
1287 if (auto insertOp = dyn_cast<vector::InsertOp>(xferOpUser))
1288 return insertOp;
1289 }
1290
1291 return vector::InsertOp();
1292 }
1293
1294 /// If the result of the TransferReadOp has exactly one user, which is a
1295 /// vector::InsertOp, return that operation's indices.
1296 void getInsertionIndices(TransferReadOp xferOp,
1298 if (auto insertOp = getInsertOp(xferOp)) {
1299 auto pos = insertOp.getMixedPosition();
1300 indices.append(pos.begin(), pos.end());
1301 }
1302 }
1303
1304 /// Rewrite the op: Unpack one dimension. Can handle masks, out-of-bounds
1305 /// accesses, and broadcasts and transposes in permutation maps.
1306 LogicalResult matchAndRewrite(TransferReadOp xferOp,
1307 PatternRewriter &rewriter) const override {
1308 if (xferOp.getVectorType().getRank() <= options.targetRank)
1309 return rewriter.notifyMatchFailure(
1310 xferOp, "vector rank is less or equal to target rank");
1311 if (failed(checkLowerTensors(xferOp, rewriter)))
1312 return failure();
1313 if (xferOp.getVectorType().getElementType() !=
1314 xferOp.getShapedType().getElementType())
1315 return rewriter.notifyMatchFailure(
1316 xferOp, "not yet supported: element type mismatch");
1317 auto xferVecType = xferOp.getVectorType();
1318 if (xferVecType.getScalableDims()[0]) {
1319 return rewriter.notifyMatchFailure(
1320 xferOp, "scalable dimensions cannot be unrolled at compile time");
1321 }
1322
1323 auto insertOp = getInsertOp(xferOp);
1324 auto vec = buildResultVector(rewriter, xferOp);
1325 auto vecType = dyn_cast<VectorType>(vec.getType());
1326
1327 VectorType newXferVecType = VectorType::Builder(xferVecType).dropDim(0);
1328
1329 int64_t dimSize = xferVecType.getShape()[0];
1330
1331 // Generate fully unrolled loop of transfer ops.
1332 Location loc = xferOp.getLoc();
1333 for (int64_t i = 0; i < dimSize; ++i) {
1334 Value iv = arith::ConstantIndexOp::create(rewriter, loc, i);
1335
1336 // FIXME: Rename this lambda - it does much more than just
1337 // in-bounds-check generation.
1338 vec = generateInBoundsCheck(
1339 rewriter, xferOp, iv, unpackedDim(xferOp), TypeRange(vecType),
1340 /*inBoundsCase=*/
1341 [&](OpBuilder &b, Location loc) {
1342 // Indices for the new transfer op.
1343 SmallVector<Value, 8> xferIndices;
1344 getXferIndices(b, xferOp, iv, xferIndices);
1345
1346 // Indices for the new vector.insert op.
1347 SmallVector<OpFoldResult, 8> insertionIndices;
1348 getInsertionIndices(xferOp, insertionIndices);
1349 insertionIndices.push_back(rewriter.getIndexAttr(i));
1350
1351 auto inBoundsAttr = dropFirstElem(b, xferOp.getInBoundsAttr());
1352
1353 auto newXferOp = vector::TransferReadOp::create(
1354 b, loc, newXferVecType, xferOp.getBase(), xferIndices,
1355 AffineMapAttr::get(unpackedPermutationMap(b, xferOp)),
1356 xferOp.getPadding(), Value(), inBoundsAttr);
1357 maybeAssignMask(b, xferOp, newXferOp, i);
1358
1359 Value valToInser = newXferOp.getResult();
1360 if (newXferVecType.getRank() == 0) {
1361 // vector.insert does not accept rank-0 as the non-indexed
1362 // argument. Extract the scalar before inserting.
1363 valToInser = vector::ExtractOp::create(b, loc, valToInser,
1365 }
1366 return vector::InsertOp::create(b, loc, valToInser, vec,
1367 insertionIndices);
1368 },
1369 /*outOfBoundsCase=*/
1370 [&](OpBuilder &b, Location loc) {
1371 // Loop through original (unmodified) vector.
1372 return vec;
1373 });
1374 }
1375
1376 if (insertOp) {
1377 // Rewrite single user of the old TransferReadOp, which was an InsertOp.
1378 rewriter.replaceOp(insertOp, vec);
1379 rewriter.eraseOp(xferOp);
1380 } else {
1381 rewriter.replaceOp(xferOp, vec);
1382 }
1383
1384 return success();
1385 }
1386};
1387
1388/// Progressive lowering of vector TransferWriteOp with unrolling: Unpack one
1389/// dimension. This is similar to TransferOpConversion<TransferWriteOp>, but no
1390/// memref buffer is allocated and the SCF loop is fully unrolled.
1391///
1392/// ```
1393/// E.g.:
1394/// ```
1395/// vector.transfer_write %vec, %A[%a, %b, %c]
1396/// : vector<5x4xf32>, memref<?x?x?xf32>
1397/// ```
1398/// is rewritten to IR such as (simplified):
1399/// ```
1400/// %v0 = vector.extract %vec[0] : vector<4xf32> from vector<5x4xf32>
1401/// vector.transfer_write %v0, %A[%a, %b, %c] : vector<4xf32>, memref<...>
1402/// %v1 = vector.extract %vec[1] : vector<4xf32> from vector<5x4xf32>
1403/// vector.transfer_write %v1, %A[%a, %b + 1, %c] : vector<4xf32>, memref<...>
1404/// ...
1405/// %v4 = vector.extract %vec[4] : vector<4xf32> from vector<5x4xf32>
1406/// vector.transfer_write %v4, %A[%a, %b + 4, %c] : vector<4xf32>, memref<...>
1407/// ```
1408///
1409/// Note: As an optimization, if the vector of the original TransferWriteOp
1410/// was directly extracted from another vector via an ExtractOp `a`, extract
1411/// the vectors for the newly generated TransferWriteOps from `a`'s input. By
1412/// doing so, `a` may become dead, and the number of ExtractOps generated during
1413/// recursive application of this pattern will be minimal.
1414struct UnrollTransferWriteConversion
1415 : public VectorToSCFPattern<TransferWriteOp> {
1416 using VectorToSCFPattern<TransferWriteOp>::VectorToSCFPattern;
1417
1418 void initialize() {
1419 // This pattern recursively unpacks one dimension at a time. The recursion
1420 // bounded as the rank is strictly decreasing.
1421 setHasBoundedRewriteRecursion();
1422 }
1423
1424 /// Return the vector from which newly generated ExtracOps will extract.
1425 Value getDataVector(TransferWriteOp xferOp) const {
1426 if (auto extractOp = getExtractOp(xferOp))
1427 return extractOp.getSource();
1428 return xferOp.getVector();
1429 }
1430
1431 /// If the input of the given TransferWriteOp is an ExtractOp, return it.
1432 vector::ExtractOp getExtractOp(TransferWriteOp xferOp) const {
1433 if (auto *op = xferOp.getVector().getDefiningOp())
1434 return dyn_cast<vector::ExtractOp>(op);
1435 return vector::ExtractOp();
1436 }
1437
1438 /// If the input of the given TransferWriteOp is an ExtractOp, return its
1439 /// indices.
1440 void getExtractionIndices(TransferWriteOp xferOp,
1442 if (auto extractOp = getExtractOp(xferOp)) {
1443 auto pos = extractOp.getMixedPosition();
1444 indices.append(pos.begin(), pos.end());
1445 }
1446 }
1447
1448 /// Rewrite the op: Unpack one dimension. Can handle masks, out-of-bounds
1449 /// accesses, and broadcasts and transposes in permutation maps.
1450 LogicalResult matchAndRewrite(TransferWriteOp xferOp,
1451 PatternRewriter &rewriter) const override {
1452 VectorType inputVectorTy = xferOp.getVectorType();
1453
1454 if (inputVectorTy.getRank() <= options.targetRank)
1455 return failure();
1456
1457 if (failed(checkLowerTensors(xferOp, rewriter)))
1458 return failure();
1459 // Transfer ops that modify the element type are not supported atm.
1460 if (inputVectorTy.getElementType() !=
1461 xferOp.getShapedType().getElementType())
1462 return failure();
1463
1464 auto vec = getDataVector(xferOp);
1465 if (inputVectorTy.getScalableDims()[0]) {
1466 // Cannot unroll a scalable dimension at compile time.
1467 return failure();
1468 }
1469
1470 int64_t dimSize = inputVectorTy.getShape()[0];
1471 Value source = xferOp.getBase(); // memref or tensor to be written to.
1472 auto sourceType = isTensorOp(xferOp) ? xferOp.getShapedType() : Type();
1473
1474 // Generate fully unrolled loop of transfer ops.
1475 Location loc = xferOp.getLoc();
1476 for (int64_t i = 0; i < dimSize; ++i) {
1477 Value iv = arith::ConstantIndexOp::create(rewriter, loc, i);
1478
1479 auto updatedSource = generateInBoundsCheck(
1480 rewriter, xferOp, iv, unpackedDim(xferOp),
1481 isTensorOp(xferOp) ? TypeRange(sourceType) : TypeRange(),
1482 /*inBoundsCase=*/
1483 [&](OpBuilder &b, Location loc) {
1484 // Indices for the new transfer op.
1485 SmallVector<Value, 8> xferIndices;
1486 getXferIndices(b, xferOp, iv, xferIndices);
1487
1488 // Indices for the new vector.extract op.
1489 SmallVector<OpFoldResult, 8> extractionIndices;
1490 getExtractionIndices(xferOp, extractionIndices);
1491 extractionIndices.push_back(b.getI64IntegerAttr(i));
1492
1493 auto extracted =
1494 vector::ExtractOp::create(b, loc, vec, extractionIndices);
1495 auto inBoundsAttr = dropFirstElem(b, xferOp.getInBoundsAttr());
1496 Value xferVec;
1497 if (inputVectorTy.getRank() == 1) {
1498 // When target-rank=0, unrolling would causes the vector input
1499 // argument into `transfer_write` to become a scalar. We solve
1500 // this by broadcasting the scalar to a 0D vector.
1501 xferVec = vector::BroadcastOp::create(
1502 b, loc, VectorType::get({}, extracted.getType()), extracted);
1503 } else {
1504 xferVec = extracted;
1505 }
1506 auto newXferOp = vector::TransferWriteOp::create(
1507 b, loc, sourceType, xferVec, source, xferIndices,
1508 AffineMapAttr::get(unpackedPermutationMap(b, xferOp)), Value(),
1509 inBoundsAttr);
1510
1511 maybeAssignMask(b, xferOp, newXferOp, i);
1512
1513 return isTensorOp(xferOp) ? newXferOp->getResult(0) : Value();
1514 },
1515 /*outOfBoundsCase=*/
1516 [&](OpBuilder &b, Location loc) {
1517 return isTensorOp(xferOp) ? source : Value();
1518 });
1519
1520 if (isTensorOp(xferOp))
1521 source = updatedSource;
1522 }
1523
1524 if (isTensorOp(xferOp))
1525 rewriter.replaceOp(xferOp, source);
1526 else
1527 rewriter.eraseOp(xferOp);
1528
1529 return success();
1530 }
1531};
1532
1533} // namespace lowering_n_d_unrolled
1534
1535namespace lowering_1_d {
1536
1537/// Compute the indices into the memref for the LoadOp/StoreOp generated as
1538/// part of TransferOp1dConversion. Return the memref dimension on which
1539/// the transfer is operating. A return value of std::nullopt indicates a
1540/// broadcast.
1541template <typename OpTy>
1542static std::optional<int64_t>
1543get1dMemrefIndices(OpBuilder &b, OpTy xferOp, Value iv,
1544 SmallVector<Value, 8> &memrefIndices) {
1545 auto indices = xferOp.getIndices();
1546 auto map = xferOp.getPermutationMap();
1547 assert(xferOp.getTransferRank() > 0 && "unexpected 0-d transfer");
1548
1549 memrefIndices.append(indices.begin(), indices.end());
1550 assert(map.getNumResults() == 1 &&
1551 "Expected 1 permutation map result for 1D transfer");
1552 if (auto expr = dyn_cast<AffineDimExpr>(map.getResult(0))) {
1553 Location loc = xferOp.getLoc();
1554 auto dim = expr.getPosition();
1555 AffineExpr d0, d1;
1556 bindDims(xferOp.getContext(), d0, d1);
1557 Value offset = memrefIndices[dim];
1558 memrefIndices[dim] =
1559 affine::makeComposedAffineApply(b, loc, d0 + d1, {offset, iv});
1560 return dim;
1561 }
1562
1563 assert(xferOp.isBroadcastDim(0) &&
1564 "Expected AffineDimExpr or AffineConstantExpr");
1565 return std::nullopt;
1566}
1567
1568/// Codegen strategy for TransferOp1dConversion, depending on the
1569/// operation.
1570template <typename OpTy>
1571struct Strategy1d;
1572
1573/// Codegen strategy for TransferReadOp.
1574template <>
1575struct Strategy1d<TransferReadOp> {
1576 static void generateForLoopBody(OpBuilder &b, Location loc,
1577 TransferReadOp xferOp, Value iv,
1578 ValueRange loopState) {
1580 auto dim = get1dMemrefIndices(b, xferOp, iv, indices);
1581 auto vec = loopState[0];
1582
1583 // In case of out-of-bounds access, leave `vec` as is (was initialized with
1584 // padding value).
1585 auto nextVec = generateInBoundsCheck(
1586 b, xferOp, iv, dim, TypeRange(xferOp.getVectorType()),
1587 /*inBoundsCase=*/
1588 [&](OpBuilder &b, Location loc) {
1589 Value val = memref::LoadOp::create(b, loc, xferOp.getBase(), indices);
1590 return vector::InsertOp::create(b, loc, val, vec, iv);
1591 },
1592 /*outOfBoundsCase=*/
1593 [&](OpBuilder & /*b*/, Location loc) { return vec; });
1594 scf::YieldOp::create(b, loc, nextVec);
1595 }
1596
1597 static Value initialLoopState(OpBuilder &b, TransferReadOp xferOp) {
1598 // Inititalize vector with padding value.
1599 Location loc = xferOp.getLoc();
1600 return vector::BroadcastOp::create(b, loc, xferOp.getVectorType(),
1601 xferOp.getPadding());
1602 }
1603};
1604
1605/// Codegen strategy for TransferWriteOp.
1606template <>
1607struct Strategy1d<TransferWriteOp> {
1608 static void generateForLoopBody(OpBuilder &b, Location loc,
1609 TransferWriteOp xferOp, Value iv,
1610 ValueRange /*loopState*/) {
1612 auto dim = get1dMemrefIndices(b, xferOp, iv, indices);
1613
1614 // Nothing to do in case of out-of-bounds access.
1615 generateInBoundsCheck(
1616 b, xferOp, iv, dim,
1617 /*inBoundsCase=*/[&](OpBuilder &b, Location loc) {
1618 auto val = vector::ExtractOp::create(b, loc, xferOp.getVector(), iv);
1619 memref::StoreOp::create(b, loc, val, xferOp.getBase(), indices);
1620 });
1621 scf::YieldOp::create(b, loc);
1622 }
1623
1624 static Value initialLoopState(OpBuilder &b, TransferWriteOp xferOp) {
1625 return Value();
1626 }
1627};
1628
1629/// Lower a 1D vector transfer op to SCF using scalar loads/stores. This is
1630/// necessary in cases where a 1D vector transfer op cannot be lowered into
1631/// vector load/stores due to non-unit strides or broadcasts:
1632///
1633/// * Transfer dimension is not the last memref dimension
1634/// * Transfer dimension is a broadcast (i.e., scalar load + broadcast)
1635/// * Memref has a layout map with non-unit stride on the last dimension
1636///
1637/// This pattern generates IR as follows:
1638///
1639/// 1. Generate a for loop iterating over each vector element.
1640/// 2. Inside the loop, generate a InsertElementOp or ExtractElementOp,
1641/// depending on OpTy.
1642///
1643/// TODO: In some cases (no masking, etc.), LLVM::MatrixColumnMajorLoadOp
1644/// can be generated instead of TransferOp1dConversion. Add such a pattern
1645/// to ConvertVectorToLLVM.
1646///
1647/// E.g.:
1648/// ```
1649/// vector.transfer_write %vec, %A[%a, %b]
1650/// {permutation_map = affine_map<(d0, d1) -> (d0)>, in_bounds = [true]}
1651/// : vector<9xf32>, memref<?x?xf32>
1652/// ```
1653/// Is rewritten to approximately the following pseudo-IR:
1654/// ```
1655/// for i = 0 to 9 {
1656/// %t = vector.extract %vec[i] : f32 from vector<9xf32>
1657/// memref.store %t, %arg0[%a + i, %b] : memref<?x?xf32>
1658/// }
1659/// ```
1660template <typename OpTy>
1661struct TransferOp1dConversion : public VectorToSCFPattern<OpTy> {
1662 using VectorToSCFPattern<OpTy>::VectorToSCFPattern;
1663
1664 LogicalResult matchAndRewrite(OpTy xferOp,
1665 PatternRewriter &rewriter) const override {
1666 // TODO: support 0-d corner case.
1667 if (xferOp.getTransferRank() == 0)
1668 return failure();
1669 auto map = xferOp.getPermutationMap();
1670 auto memRefType = dyn_cast<MemRefType>(xferOp.getShapedType());
1671
1672 if (!memRefType)
1673 return failure();
1674 if (xferOp.getVectorType().getRank() != 1)
1675 return failure();
1676 if (map.isMinorIdentity() && memRefType.isLastDimUnitStride())
1677 return failure(); // Handled by ConvertVectorToLLVM
1678
1679 // Loop bounds, step, state...
1680 Location loc = xferOp.getLoc();
1681 auto vecType = xferOp.getVectorType();
1682 auto lb = arith::ConstantIndexOp::create(rewriter, loc, 0);
1683 Value ub =
1684 arith::ConstantIndexOp::create(rewriter, loc, vecType.getDimSize(0));
1685 if (vecType.isScalable()) {
1686 Value vscale =
1687 vector::VectorScaleOp::create(rewriter, loc, rewriter.getIndexType());
1688 ub = arith::MulIOp::create(rewriter, loc, ub, vscale);
1689 }
1690 auto step = arith::ConstantIndexOp::create(rewriter, loc, 1);
1691 auto loopState = Strategy1d<OpTy>::initialLoopState(rewriter, xferOp);
1692
1693 // Generate for loop.
1694 rewriter.replaceOpWithNewOp<scf::ForOp>(
1695 xferOp, lb, ub, step, loopState ? ValueRange(loopState) : ValueRange(),
1696 [&](OpBuilder &b, Location loc, Value iv, ValueRange loopState) {
1697 Strategy1d<OpTy>::generateForLoopBody(b, loc, xferOp, iv, loopState);
1698 });
1699
1700 return success();
1701 }
1702};
1703
1704} // namespace lowering_1_d
1705} // namespace
1706
1709 if (options.unroll) {
1710 patterns.add<lowering_n_d_unrolled::UnrollTransferReadConversion,
1711 lowering_n_d_unrolled::UnrollTransferWriteConversion>(
1712 patterns.getContext(), options);
1713 } else {
1714 patterns.add<lowering_n_d::PrepareTransferReadConversion,
1715 lowering_n_d::PrepareTransferWriteConversion,
1716 lowering_n_d::TransferOpConversion<TransferReadOp>,
1717 lowering_n_d::TransferOpConversion<TransferWriteOp>>(
1718 patterns.getContext(), options);
1719 }
1720 if (options.lowerScalable) {
1721 patterns.add<lowering_n_d::ScalableTransposeTransferWriteConversion>(
1722 patterns.getContext(), options);
1723 }
1724 if (options.targetRank == 1) {
1725 patterns.add<lowering_1_d::TransferOp1dConversion<TransferReadOp>,
1726 lowering_1_d::TransferOp1dConversion<TransferWriteOp>>(
1727 patterns.getContext(), options);
1728 }
1729 patterns.add<lowering_n_d::DecomposePrintOpConversion>(patterns.getContext(),
1730 options);
1731}
1732
1733namespace {
1734
1735struct ConvertVectorToSCFPass
1736 : public impl::ConvertVectorToSCFBase<ConvertVectorToSCFPass> {
1737 ConvertVectorToSCFPass() = default;
1738 ConvertVectorToSCFPass(const VectorTransferToSCFOptions &options) {
1739 this->fullUnroll = options.unroll;
1740 this->targetRank = options.targetRank;
1741 this->lowerTensors = options.lowerTensors;
1742 this->lowerScalable = options.lowerScalable;
1743 }
1744
1745 void runOnOperation() override {
1746 VectorTransferToSCFOptions options;
1747 options.unroll = fullUnroll;
1748 options.targetRank = targetRank;
1749 options.lowerTensors = lowerTensors;
1750 options.lowerScalable = lowerScalable;
1751
1752 // Lower permutation maps first.
1753 RewritePatternSet lowerTransferPatterns(&getContext());
1755 lowerTransferPatterns);
1756 (void)applyPatternsGreedily(getOperation(),
1757 std::move(lowerTransferPatterns));
1758
1759 RewritePatternSet patterns(&getContext());
1761 (void)applyPatternsGreedily(getOperation(), std::move(patterns));
1762 }
1763};
1764
1765} // namespace
1766
1767std::unique_ptr<Pass>
1769 return std::make_unique<ConvertVectorToSCFPass>(options);
1770}
return success()
MLIR_CRUNNERUTILS_EXPORT void printClose()
LogicalResult initialize(unsigned origNumLoops, ArrayRef< ReassociationIndices > foldedIterationDims)
b
Return true if permutation is a valid permutation of the outer_dims_perm (case OuterOrInnerPerm::Oute...
ArrayAttr()
b getContext())
static llvm::ManagedStatic< PassManagerOptions > options
static void printOp(llvm::raw_ostream &os, Operation *op, OpPrintingFlags &flags)
Definition Unit.cpp:18
static void getXferIndices(RewriterBase &rewriter, TransferOpType xferOp, AffineMap offsetMap, ArrayRef< Value > dimValues, SmallVector< Value, 4 > &indices)
For a vector TransferOpType xferOp, an empty indices vector, and an AffineMap representing offsets to...
Base type for affine expression.
Definition AffineExpr.h:68
A multi-dimensional affine map Affine map's are immutable like Type's, and they are uniqued.
Definition AffineMap.h:46
static AffineMap get(MLIRContext *context)
Returns a zero result affine map with no dimensions or symbols: () -> ().
IntegerAttr getIndexAttr(int64_t value)
Definition Builders.cpp:116
UnitAttr getUnitAttr()
Definition Builders.cpp:106
IntegerType getI1Type()
Definition Builders.cpp:61
MLIRContext * getContext() const
Definition Builders.h:56
IndexType getIndexType()
Definition Builders.cpp:59
ImplicitLocOpBuilder maintains a 'current location', allowing use of the create<> method without spec...
Definition Builders.h:632
This class defines the main interface for locations in MLIR and acts as a non-nullable wrapper around...
Definition Location.h:76
RAII guard to reset the insertion point of the builder when destroyed.
Definition Builders.h:351
This class helps build Operations.
Definition Builders.h:210
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 setInsertionPointToStart(Block *block)
Sets the insertion point to the start of the specified block.
Definition Builders.h:434
void setInsertionPointAfter(Operation *op)
Sets the insertion point to the node after the specified operation, which will cause subsequent inser...
Definition Builders.h:415
A trait of region holding operations that define a new scope for automatic allocations,...
Operation is the basic unit of execution within MLIR.
Definition Operation.h:87
Region & getRegion(unsigned index)
Returns the region held by this operation at position 'index'.
Definition Operation.h:738
void setDiscardableAttr(StringAttr name, Attribute value)
Set a discardable attribute by name.
Definition Operation.h:512
Operation * getParentWithTrait()
Returns the closest surrounding parent operation with trait Trait.
Definition Operation.h:273
unsigned getNumRegions()
Returns the number of regions held by this operation.
Definition Operation.h:726
OperandRange operand_range
Definition Operation.h:396
user_range getUsers()
Returns a range of all users.
Definition Operation.h:925
A special type of RewriterBase that coordinates the application of a rewrite pattern on the current I...
Block & front()
Definition Region.h:65
MLIRContext * getContext() const
RewritePatternSet & add(ConstructorArg &&arg, ConstructorArgs &&...args)
Add an instance of each of the pattern types 'Ts' to the pattern list with the given arguments.
virtual void replaceOp(Operation *op, ValueRange newValues)
Replace the results of the given (original) operation with the specified list of values (replacements...
virtual void eraseOp(Operation *op)
This method erases an operation that is known to have no uses.
std::enable_if_t<!std::is_convertible< CallbackT, Twine >::value, LogicalResult > notifyMatchFailure(Location loc, CallbackT &&reasonCallback)
Used to notify the listener that the IR failed to be rewritten because of a match failure,...
void modifyOpInPlace(Operation *root, CallableT &&callable)
This method is a utility wrapper around an in-place modification of an operation.
OpTy replaceOpWithNewOp(Operation *op, Args &&...args)
Replace the results of the given (original) op with a new op that is created without verification (re...
This class provides an abstraction over the various different ranges of value types.
Definition TypeRange.h:40
Instances of the Type class are uniqued, have an immutable identifier and an optional mutable compone...
Definition Types.h:74
This class provides an abstraction over the different types of ranges over Values.
Definition ValueRange.h:389
type_range getType() const
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
user_range getUsers() const
Definition Value.h:218
Operation * getDefiningOp() const
If this value is the result of an operation, return the operation that defines it.
Definition Value.cpp:18
This is a builder type that keeps local references to arguments.
Builder & dropDim(unsigned pos)
Erase a dim from shape @pos.
static ConstantIndexOp create(OpBuilder &builder, Location location, int64_t value)
Definition ArithOps.cpp:398
AffineApplyOp makeComposedAffineApply(OpBuilder &b, Location loc, AffineMap map, ArrayRef< OpFoldResult > operands, bool composeAffineMin=false)
Returns a composed AffineApplyOp by composing map and operands with other AffineApplyOps supplying th...
void populateVectorTransferPermutationMapLoweringPatterns(RewritePatternSet &patterns, PatternBenefit benefit=1)
Collect a set of transfer read/write lowering patterns that simplify the permutation map (e....
Value createOrFoldDimOp(OpBuilder &b, Location loc, Value source, int64_t dim)
Helper function that creates a memref::DimOp or tensor::DimOp depending on the type of source.
SmallVector< Value > getAsValues(OpBuilder &builder, Location loc, ArrayRef< OpFoldResult > foldResults)
Convert foldResults into Values.
auto makeVscaleConstantBuilder(PatternRewriter &rewriter, Location loc)
Returns a functor (int64_t -> Value) which returns a constant vscale multiple.
Include the generated interface declarations.
Type getType(OpFoldResult ofr)
Returns the int type of the integer in ofr.
Definition Utils.cpp:311
void bindDims(MLIRContext *ctx, AffineExprTy &...exprs)
Bind a list of AffineExpr references to DimExpr at positions: [0 .
Definition AffineExpr.h:311
LogicalResult applyPatternsGreedily(Region &region, const FrozenRewritePatternSet &patterns, GreedyRewriteConfig config=GreedyRewriteConfig(), bool *changed=nullptr)
Rewrite ops in the given region, which must be isolated from above, by repeatedly applying the highes...
void populateVectorToSCFConversionPatterns(RewritePatternSet &patterns, const VectorTransferToSCFOptions &options=VectorTransferToSCFOptions())
Collect a set of patterns to convert from the Vector dialect to SCF + func.
std::unique_ptr< Pass > createConvertVectorToSCFPass(const VectorTransferToSCFOptions &options=VectorTransferToSCFOptions())
Create a pass to convert a subset of vector ops to SCF.
llvm::function_ref< Fn > function_ref
Definition LLVM.h:147
OpRewritePattern is a wrapper around RewritePattern that allows for matching and rewriting against an...
When lowering an N-d vector transfer op to an (N-1)-d vector transfer op, a temporary buffer is creat...
Definition VectorToSCF.h:52