MLIR 24.0.0git
LowerVectorContract.cpp
Go to the documentation of this file.
1//===- LowerVectorContract.cpp - Lower 'vector.contract' operation --------===//
2//
3// Part of the LLVM Project, under the Apache License v2.0 with LLVM Exceptions.
4// See https://llvm.org/LICENSE.txt for license information.
5// SPDX-License-Identifier: Apache-2.0 WITH LLVM-exception
6//
7//===----------------------------------------------------------------------===//
8//
9// This file implements target-independent rewrites and utilities to lower the
10// 'vector.contract' operation.
11//
12//===----------------------------------------------------------------------===//
13
22#include "mlir/IR/Location.h"
25
26#define DEBUG_TYPE "vector-contract-lowering"
27
28using namespace mlir;
29using namespace mlir::vector;
30
31//===----------------------------------------------------------------------===//
32// Helper functions
33//===----------------------------------------------------------------------===//
34// Helper to find an index in an affine map.
35static std::optional<int64_t> getResultIndex(AffineMap map, int64_t index) {
36 for (int64_t i = 0, e = map.getNumResults(); i < e; ++i) {
37 int64_t idx = map.getDimPosition(i);
38 if (idx == index)
39 return i;
40 }
41 return std::nullopt;
42}
43
44// Helper to construct iterator types with one index removed.
46 int64_t index) {
48 for (const auto &it : llvm::enumerate(iteratorTypes)) {
49 int64_t idx = it.index();
50 if (idx == index)
51 continue;
52 results.push_back(it.value());
53 }
54 return results;
55}
56
57// Helper to construct an affine map with one index removed.
59 PatternRewriter &rewriter) {
60 auto *ctx = rewriter.getContext();
62 for (int64_t i = 0, e = map.getNumResults(); i < e; ++i) {
63 int64_t idx = map.getDimPosition(i);
64 if (idx == index)
65 continue;
66 // Re-insert remaining indices, but renamed when occurring
67 // after the removed index.
68 auto targetExpr = getAffineDimExpr(idx < index ? idx : idx - 1, ctx);
69 results.push_back(targetExpr);
70 }
71 return AffineMap::get(map.getNumDims() - 1, 0, results, ctx);
72}
73
74/// Returns `val` with the dimension at position `index` dropped by indexing
75/// that dimension with `pos`.
76///
77/// If `index == -1`, returns `val` unchanged. If `index == 0`, the result is
78/// a single `vector.extract %val[pos]`.
79///
80/// Example (`index == 0`): extract the sub-vector at `pos` along the leading
81/// dimension.
82/// // val : vector<4x8xf32>, pos = 2
83/// %res = vector.extract %val[2] : vector<8xf32> from vector<4x8xf32>
84///
85/// For `index > 0`, recursively applies the same drop to each sub-vector of
86/// the leading dimension and reassembles the result.
88 PatternRewriter &rewriter) {
89 if (index == -1)
90 return val;
91
92 // At extraction dimension?
93 if (index == 0)
94 return vector::ExtractOp::create(rewriter, loc, val, pos);
95
96 // Unroll leading dimensions.
97 VectorType type = cast<VectorType>(val.getType());
98 VectorType resType = VectorType::Builder(type).dropDim(index);
99 Value result = arith::ConstantOp::create(rewriter, loc, resType,
100 rewriter.getZeroAttr(resType));
101 for (int64_t d = 0, e = resType.getDimSize(0); d < e; d++) {
102 Value ext = vector::ExtractOp::create(rewriter, loc, val, d);
103 Value load = reshapeLoad(loc, ext, index - 1, pos, rewriter);
104 result = vector::InsertOp::create(rewriter, loc, load, result, d);
105 }
106 return result;
107}
108
109/// Inserts `val` into `result` at position `pos` along dimension `index`.
110///
111/// This is the inverse of `reshapeLoad`. If `index == -1`, returns `val`. If
112/// `index == 0`, the result is a single `vector.insert %val, %result [pos]`.
113///
114/// Example (`index == 0`): insert `val` at `pos` along the leading dimension.
115/// // val : vector<4xf32>, acc : vector<2x4xf32>, pos = 1
116/// %res = vector.insert %val, %acc [1] : vector<4xf32> into vector<2x4xf32>
117///
118/// For `index > 0`, recursively applies the same insertion to each sub-vector
119/// of the leading dimension and reassembles the result.
121 int64_t pos, PatternRewriter &rewriter) {
122 // Unmodified?
123 if (index == -1)
124 return val;
125 // At insertion dimension?
126 if (index == 0)
127 return vector::InsertOp::create(rewriter, loc, val, result, pos);
128
129 // Unroll leading dimensions.
130 VectorType type = cast<VectorType>(result.getType());
131 for (int64_t d = 0, e = type.getDimSize(0); d < e; d++) {
132 Value ext = vector::ExtractOp::create(rewriter, loc, result, d);
133 Value ins = vector::ExtractOp::create(rewriter, loc, val, d);
134 Value sto = reshapeStore(loc, ins, ext, index - 1, pos, rewriter);
135 result = vector::InsertOp::create(rewriter, loc, sto, result, d);
136 }
137 return result;
138}
139
140/// Helper to create arithmetic operation associated with a kind of contraction.
141static std::optional<Value>
143 vector::CombiningKind kind, PatternRewriter &rewriter,
144 bool isInt, Value mask = Value(),
145 arith::FastMathFlagsAttr fmf = {}) {
146 using vector::CombiningKind;
147 Value mul;
148
149 if (isInt) {
150 if (kind == CombiningKind::MINNUMF || kind == CombiningKind::MAXNUMF ||
151 kind == CombiningKind::MINIMUMF || kind == CombiningKind::MAXIMUMF ||
152 kind == CombiningKind::MINIMUMNUMF ||
153 kind == CombiningKind::MAXIMUMNUMF)
154 // Only valid for floating point types.
155 return std::nullopt;
156 mul = arith::MulIOp::create(rewriter, loc, x, y);
157 } else {
158 // Float case.
159 if (kind == CombiningKind::AND || kind == CombiningKind::MINUI ||
160 kind == CombiningKind::MINSI || kind == CombiningKind::MAXUI ||
161 kind == CombiningKind::MAXSI || kind == CombiningKind::OR ||
162 kind == CombiningKind::XOR)
163 // Only valid for integer types.
164 return std::nullopt;
165 // Special case for fused multiply-add.
166 if (acc && isa<VectorType>(acc.getType()) && kind == CombiningKind::ADD) {
167 Value fma = vector::FMAOp::create(rewriter, loc, x, y, acc);
168 if (mask)
169 // The fma op doesn't need explicit masking. However, fma ops used in
170 // reductions must preserve previous 'acc' values for masked-out lanes.
171 fma = selectPassthru(rewriter, mask, fma, acc);
172 return fma;
173 }
174 mul = arith::MulFOp::create(rewriter, loc, x, y, fmf);
175 }
176
177 if (!acc)
178 return std::optional<Value>(mul);
179
180 return makeArithReduction(rewriter, loc, kind, mul, acc, fmf, mask);
181}
182
183/// Return the positions of the reductions in the given map.
185 ArrayAttr iteratorTypes) {
186 SmallVector<int64_t> dimsIdx;
187 for (unsigned i = 0, e = map.getNumResults(); i < e; i++) {
188 if (isReductionIterator(iteratorTypes[map.getDimPosition(i)]))
189 dimsIdx.push_back(i);
190 }
191 return dimsIdx;
192}
193
194/// Look for a given dimension in an affine map and return its position. Return
195/// std::nullopt if the dimension is not in the map results.
196static std::optional<unsigned> getDimPosition(AffineMap map, unsigned dim) {
197 for (unsigned i = 0, e = map.getNumResults(); i < e; i++) {
198 if (map.getDimPosition(i) == dim)
199 return i;
200 }
201 return std::nullopt;
202}
203
204/// Creates an AddIOp if `isInt` is true otherwise create an arith::AddFOp using
205/// operands `x` and `y`.
206static Value createAdd(Location loc, Value x, Value y, bool isInt,
207 PatternRewriter &rewriter,
208 arith::FastMathFlagsAttr fmf = {}) {
209 if (isInt)
210 return arith::AddIOp::create(rewriter, loc, x, y);
211 return arith::AddFOp::create(rewriter, loc, x, y, fmf);
212}
213
214/// Creates a MulIOp if `isInt` is true otherwise create an MulFOp using
215/// operands `x and `y`.
216static Value createMul(Location loc, Value x, Value y, bool isInt,
217 PatternRewriter &rewriter,
218 arith::FastMathFlagsAttr fmf = {}) {
219 if (isInt)
220 return arith::MulIOp::create(rewriter, loc, x, y);
221 return arith::MulFOp::create(rewriter, loc, x, y, fmf);
222}
223
224namespace {
225
226/// Progressive lowering of a `vector.contract %a, %b, %c` with row-major matmul
227/// semantics to a reduction_size-unrolled sequence:
228/// ```
229/// %at = vector.transpose %a, [1, 0]
230/// %bRow0 = vector.extract %b[0]
231/// %atRow0 = vector.extract %at[0]
232/// %c0 = vector.outerproduct %atRow0, %bRow0, %c
233/// ...
234/// %bRowK = vector.extract %b[K]
235/// %atRowK = vector.extract %at[K]
236/// %cK = vector.outerproduct %atRowK, %bRowK, %cK-1
237/// ```
238///
239/// This only kicks in when vectorContractLowering is set to OuterProduct and
240/// the vector.contract op is a row-major matrix multiply.
241class ContractionOpToOuterProductOpLowering
242 : public MaskableOpRewritePattern<vector::ContractionOp> {
243public:
244 using MaskableOpRewritePattern::MaskableOpRewritePattern;
245
246 using FilterConstraintType =
247 std::function<LogicalResult(vector::ContractionOp op)>;
248
249 static LogicalResult defaultFilter(vector::ContractionOp op) {
250 return success();
251 }
252
253 ContractionOpToOuterProductOpLowering(
254 vector::VectorContractLowering vectorContractLowering,
255 MLIRContext *context, PatternBenefit benefit = 1,
256 FilterConstraintType constraint = defaultFilter)
257 : MaskableOpRewritePattern<vector::ContractionOp>(context, benefit),
258 vectorContractLowering(vectorContractLowering),
259 filter(std::move(constraint)) {}
260
261 FailureOr<Value>
262 matchAndRewriteMaskableOp(vector::ContractionOp op, MaskingOpInterface maskOp,
263 PatternRewriter &rewriter) const override;
264
265private:
266 /// Options to control the vector patterns.
267 vector::VectorContractLowering vectorContractLowering;
268 FilterConstraintType filter;
269};
270
271/// Progressive lowering of a `vector.contract %a, %b, %c` with row-major matmul
272/// semantics to an output-size-unrolled sequence:
273/// ```
274/// %out = arith.constant ... : vector<MxNxelt_type>
275/// %bt = vector.transpose %b, [1, 0]
276/// %aRow0 = vector.extract %a[0]
277/// %btRow0 = vector.extract %bt[0]
278/// %c00 = vector.reduction %atRow0, %bRow0
279/// %out00 = vector.insert %c00, %out[0, 0]
280/// ...
281/// %aRowLast = vector.extract %at[M-1]
282/// %btRowLast = vector.extract %b[N-1]
283/// %cLastLast = vector.reduction %atRowLast, %bRowLast
284/// %outcLastLast = vector.insert %cLastLast, %out[M-1, N-1]
285/// ```
286///
287/// This only kicks in when VectorTransformsOptions is set to Dot and
288/// the vector.contract op is a row-major matmul or matvec.
289class ContractionOpToDotLowering
290 : public MaskableOpRewritePattern<vector::ContractionOp> {
291public:
292 using MaskableOpRewritePattern::MaskableOpRewritePattern;
293
294 using FilterConstraintType =
295 std::function<LogicalResult(vector::ContractionOp op)>;
296
297 static LogicalResult defaultFilter(vector::ContractionOp op) {
298 return success();
299 }
300
301 ContractionOpToDotLowering(
302 vector::VectorContractLowering vectorContractLowering,
303 MLIRContext *context, PatternBenefit benefit = 1,
304 const FilterConstraintType &constraint = defaultFilter)
305 : MaskableOpRewritePattern<vector::ContractionOp>(context, benefit),
306 vectorContractLowering(vectorContractLowering), filter(defaultFilter) {}
307
308 FailureOr<Value>
309 matchAndRewriteMaskableOp(vector::ContractionOp op, MaskingOpInterface maskOp,
310 PatternRewriter &rewriter) const override;
311
312private:
313 /// Options to control the vector patterns.
314 vector::VectorContractLowering vectorContractLowering;
315 FilterConstraintType filter;
316};
317
318/// Progressive lowering of ContractionOp.
319///
320/// One:
321/// %x = vector.contract with at least one free/batch dimension
322/// is replaced by:
323/// %a = vector.contract with one less free/batch dimension
324/// %b = vector.contract with one less free/batch dimension
325/// ..
326/// %x = combine %a %b ..
327/// until a pure contraction is reached (no free/batch dimensions),
328/// which is replaced by a dot-product.
329///
330/// This only kicks in when either VectorTransformsOptions is set
331/// to Dot or when other contraction patterns fail.
332class ContractionOpLowering
333 : public MaskableOpRewritePattern<vector::ContractionOp> {
334public:
335 using MaskableOpRewritePattern::MaskableOpRewritePattern;
336 using FilterConstraintType =
337 std::function<LogicalResult(vector::ContractionOp op)>;
338
339 static LogicalResult defaultFilter(vector::ContractionOp op) {
340 return success();
341 }
342
343 ContractionOpLowering(
344 vector::VectorContractLowering vectorContractLoweringOption,
345 MLIRContext *context, PatternBenefit benefit = 1,
346 FilterConstraintType constraint = defaultFilter)
347 : MaskableOpRewritePattern<vector::ContractionOp>(context, benefit),
348 vectorContractLoweringOption(vectorContractLoweringOption),
349 filter(std::move(constraint)) {}
350
351 FailureOr<Value>
352 matchAndRewriteMaskableOp(vector::ContractionOp op, MaskingOpInterface maskOp,
353 PatternRewriter &rewriter) const override;
354
355private:
356 /// Options to control the vector patterns.
357 vector::VectorContractLowering vectorContractLoweringOption;
358 FilterConstraintType filter;
359 // Lower one parallel dimension.
360 FailureOr<Value> lowerParallel(PatternRewriter &rewriter,
361 vector::ContractionOp op, int64_t lhsIndex,
362 int64_t rhsIndex, Value mask) const;
363 // Lower one reduction dimension.
364 FailureOr<Value> lowerReduction(PatternRewriter &rewriter,
365 vector::ContractionOp op, Value mask) const;
366};
367
368/// Generate a vector implementation for matmat, matvec and tmatvec.
369/// This unrolls outer-products along the reduction dimension.
370struct UnrolledOuterProductGenerator
371 : public StructuredGenerator<vector::ContractionOp, vector::IteratorType> {
372 UnrolledOuterProductGenerator(RewriterBase &b, vector::ContractionOp op)
373 : StructuredGenerator<vector::ContractionOp, vector::IteratorType>(b, op),
374 kind(op.getKind()), lhs(op.getLhs()), rhs(op.getRhs()),
375 res(op.getAcc()), lhsType(op.getLhsType()) {
376 auto maskableOp = cast<MaskableOpInterface>(op.getOperation());
377 if (maskableOp.isMasked())
378 mask = maskableOp.getMaskingOp().getMask();
379 }
380
381 Value t(Value v, ArrayRef<int64_t> perm = {1, 0}) {
382 if (!v)
383 return v;
384 return vector::TransposeOp::create(rewriter, loc, v, perm);
385 }
386
387 Value promote(Value v, Type dstElementType) {
388 Type elementType = v.getType();
389 auto vecType = dyn_cast<VectorType>(elementType);
390 if (vecType)
391 elementType = vecType.getElementType();
392 if (elementType == dstElementType)
393 return v;
394 Type promotedType = dstElementType;
395 if (vecType)
396 promotedType = vecType.clone(promotedType);
397 if (isa<FloatType>(dstElementType))
398 return arith::ExtFOp::create(rewriter, loc, promotedType, v,
399 /*fastmath=*/{});
400 return arith::ExtSIOp::create(rewriter, loc, promotedType, v);
401 }
402
403 FailureOr<Value> outerProd(Value lhs, Value rhs, Value res,
404 VectorType lhsType, int reductionSize,
405 std::optional<Value> maybeMask = std::nullopt) {
406 // Incremental support for masking.
407 if (mask && !maybeMask.has_value())
408 return failure();
409
410 Type resElementType = cast<VectorType>(res.getType()).getElementType();
411 for (int64_t k = 0; k < reductionSize; ++k) {
412 Value extractA = vector::ExtractOp::create(rewriter, loc, lhs, k);
413 Value extractB = vector::ExtractOp::create(rewriter, loc, rhs, k);
414 extractA = promote(extractA, resElementType);
415 extractB = promote(extractB, resElementType);
416 Value extractMask;
417 if (maybeMask.has_value() && maybeMask.value())
418 extractMask =
419 vector::ExtractOp::create(rewriter, loc, maybeMask.value(), k);
420
421 Operation *outerProdOp = vector::OuterProductOp::create(
422 rewriter, loc, res.getType(), extractA, extractB, res, kind);
423 res = maskOperation(rewriter, outerProdOp, extractMask)->getResult(0);
424 }
425 return res;
426 }
427
428 /// Helper function for `matmat`, `matvec`, `tmatvec`. Returns the size of
429 /// dimension `reductionDim`. If the dimension is a scalable dimension,
430 /// returns "nullopt".
431 std::optional<int64_t> getReductionSize(VectorType vecType,
432 int64_t reductionDim) {
433 // Cannot unroll scalable dimension.
434 if (vecType.getScalableDims()[reductionDim])
435 return std::nullopt;
436 int64_t reductionSize = vecType.getDimSize(reductionDim);
437 assert(reductionSize > 0 &&
438 "Reduction dim must be a known static size to allow unrolling");
439 return reductionSize;
440 }
441
442 /// Two outer parallel, one inner reduction (matmat flavor).
443 FailureOr<Value> matmat() {
444 if (!iters({Par(), Par(), Red()}))
445 return failure();
446 // Set up the parallel/reduction structure in the right form.
447 AffineExpr m, n, k;
448 bindDims(rewriter.getContext(), m, n, k);
449
450 // Classical row-major matmul: Just permute the lhs.
451 if (layout({{m, k}, {k, n}, {m, n}})) {
452 if (auto reductionSize = getReductionSize(lhsType, 1)) {
453 // Note: `t` creates new IR. It must be nested within this `if` check
454 // so that no IR is created when then pattern returns "failure".
455 Value tLhs = t(lhs);
456 Value tMask = t(mask, {2, 0, 1});
457 return outerProd(tLhs, rhs, res, lhsType, *reductionSize, tMask);
458 }
459 }
460 // TODO: may be better to fail and use some vector<k> -> scalar reduction.
461 if (layout({{m, k}, {n, k}, {m, n}})) {
462 if (auto reductionSize = getReductionSize(lhsType, 1)) {
463 Value tLhs = t(lhs);
464 Value tRhs = t(rhs);
465 Value tMask = t(mask, {2, 0, 1});
466 return outerProd(tLhs, tRhs, res, lhsType, *reductionSize, tMask);
467 }
468 }
469 // No need to permute anything.
470 if (layout({{k, m}, {k, n}, {m, n}})) {
471 if (auto reductionSize = getReductionSize(lhsType, 0)) {
472 Value tMask = t(mask, {2, 0, 1});
473 return outerProd(lhs, rhs, res, lhsType, *reductionSize, tMask);
474 }
475 }
476 // Just permute the rhs.
477 if (layout({{k, m}, {n, k}, {m, n}})) {
478 if (auto reductionSize = getReductionSize(lhsType, 0)) {
479 Value tRhs = t(rhs);
480 Value tMask = t(mask, {2, 0, 1});
481 return outerProd(lhs, tRhs, res, lhsType, *reductionSize, tMask);
482 }
483 }
484 // Transposed output: swap RHS and LHS.
485 // Classical row-major matmul: permute the lhs.
486 if (layout({{m, k}, {k, n}, {n, m}})) {
487 if (auto reductionSize = getReductionSize(lhsType, 1)) {
488 Value tLhs = t(lhs);
489 Value tMask = t(mask, {2, 0, 1});
490 return outerProd(rhs, tLhs, res, lhsType, *reductionSize, tMask);
491 }
492 }
493 // TODO: may be better to fail and use some vector<k> -> scalar reduction.
494 if (layout({{m, k}, {n, k}, {n, m}})) {
495 if (auto reductionSize = getReductionSize(lhsType, 1)) {
496 Value tRhs = t(rhs);
497 Value tLhs = t(lhs);
498 Value tMask = t(mask, {2, 0, 1});
499 return outerProd(tRhs, tLhs, res, lhsType, *reductionSize, tMask);
500 }
501 }
502 if (layout({{k, m}, {k, n}, {n, m}})) {
503 if (auto reductionSize = getReductionSize(lhsType, 0)) {
504 Value tMask = t(mask, {2, 0, 1});
505 return outerProd(rhs, lhs, res, lhsType, *reductionSize, tMask);
506 }
507 }
508 if (layout({{k, m}, {n, k}, {n, m}})) {
509 if (auto reductionSize = getReductionSize(lhsType, 0)) {
510 Value tRhs = t(rhs);
511 Value tMask = t(mask, {2, 0, 1});
512 return outerProd(tRhs, lhs, res, lhsType, *reductionSize, tMask);
513 }
514 }
515 return failure();
516 }
517
518 //
519 // One outer parallel, one inner reduction (matvec flavor).
520 // Mask needs to be transposed everywhere to turn the reduction dimension
521 // outermost as required by outerproduct.
522 //
523 FailureOr<Value> matvec() {
524 if (!iters({Par(), Red()}))
525 return failure();
526 AffineExpr m, k;
527 bindDims(rewriter.getContext(), m, k);
528
529 // Case mat-vec: transpose.
530 if (layout({{m, k}, {k}, {m}})) {
531 if (auto reductionSize = getReductionSize(lhsType, 1)) {
532 Value tLhs = t(lhs);
533 Value tMask = t(mask);
534 return outerProd(tLhs, rhs, res, lhsType, *reductionSize, tMask);
535 }
536 }
537 // Case mat-trans-vec: ready to go.
538 if (layout({{k, m}, {k}, {m}})) {
539 if (auto reductionSize = getReductionSize(lhsType, 0)) {
540 Value tMask = t(mask);
541 return outerProd(lhs, rhs, res, lhsType, *reductionSize, tMask);
542 }
543 }
544 // Case vec-mat: swap and transpose.
545 if (layout({{k}, {m, k}, {m}})) {
546 if (auto reductionSize = getReductionSize(lhsType, 0)) {
547 Value tRhs = t(rhs);
548 Value tMask = t(mask);
549 return outerProd(tRhs, lhs, res, lhsType, *reductionSize, tMask);
550 }
551 }
552 // Case vec-mat-trans: swap and ready to go.
553 if (layout({{k}, {k, m}, {m}})) {
554 if (auto reductionSize = getReductionSize(lhsType, 0)) {
555 Value tMask = t(mask);
556 return outerProd(rhs, lhs, res, lhsType, *reductionSize, tMask);
557 }
558 }
559 return failure();
560 }
561
562 //
563 // One outer reduction, one inner parallel (tmatvec flavor).
564 // Mask already has the shape of the outer product.
565 //
566 FailureOr<Value> tmatvec() {
567 if (!iters({Red(), Par()}))
568 return failure();
569 AffineExpr k, m;
570 bindDims(rewriter.getContext(), k, m);
571
572 // Case mat-vec: transpose.
573 if (layout({{m, k}, {k}, {m}}))
574 if (auto reductionSize = getReductionSize(lhsType, 1))
575 return outerProd(t(lhs), rhs, res, lhsType, *reductionSize, mask);
576 // Case mat-trans-vec: ready to go.
577 if (layout({{k, m}, {k}, {m}}))
578 if (auto reductionSize = getReductionSize(lhsType, 0))
579 return outerProd(lhs, rhs, res, lhsType, *reductionSize, mask);
580 // Case vec-mat: swap and transpose.
581 if (layout({{k}, {m, k}, {m}}))
582 if (auto reductionSize = getReductionSize(lhsType, 0))
583 return outerProd(t(rhs), lhs, res, lhsType, *reductionSize, mask);
584 // Case vec-mat-trans: swap and ready to go.
585 if (layout({{k}, {k, m}, {m}}))
586 if (auto reductionSize = getReductionSize(lhsType, 0))
587 return outerProd(rhs, lhs, res, lhsType, *reductionSize, mask);
588 return failure();
589 }
590
591private:
592 vector::CombiningKind kind;
593 Value lhs, rhs, res, mask;
594 VectorType lhsType;
595};
596
597/// Progressively lower a `vector.contract %a, %b, %c` with row-major matmul
598/// semantics to a reduction_size-unrolled sequence:
599/// ```
600/// %at = vector.transpose %a, [1, 0]
601/// %bRow0 = vector.extract %b[0]
602/// %atRow0 = vector.extract %at[0]
603/// %c0 = vector.outerproduct %atRow0, %bRow0, %c
604/// ...
605/// %bRowK = vector.extract %b[K]
606/// %atRowK = vector.extract %at[K]
607/// %cK = vector.outerproduct %atRowK, %bRowK, %cK-1
608/// ```
609///
610/// This only kicks in when vectorContractLowering is set to OuterProduct but
611/// otherwise supports any layout permutation of the matrix-multiply.
612FailureOr<Value>
613ContractionOpToOuterProductOpLowering::matchAndRewriteMaskableOp(
614 vector::ContractionOp op, MaskingOpInterface maskOp,
615 PatternRewriter &rewriter) const {
616 if (vectorContractLowering != vector::VectorContractLowering::OuterProduct)
617 return failure();
618
619 if (failed(filter(op)))
620 return failure();
621
622 UnrolledOuterProductGenerator e(rewriter, op);
623 FailureOr<Value> matmatRes = e.matmat();
624 if (succeeded(matmatRes)) {
625 return matmatRes;
626 }
627 FailureOr<Value> matvecRes = e.matvec();
628 if (succeeded(matvecRes)) {
629 return matvecRes;
630 }
631
632 FailureOr<Value> tmatvecRes = e.tmatvec();
633 return tmatvecRes;
634}
635
636FailureOr<Value> ContractionOpToDotLowering::matchAndRewriteMaskableOp(
637 vector::ContractionOp op, MaskingOpInterface maskOp,
638 PatternRewriter &rewriter) const {
639 // TODO: Support vector.mask.
640 if (maskOp)
641 return failure();
642
643 if (failed(filter(op)))
644 return failure();
645
646 if (vectorContractLowering != vector::VectorContractLowering::Dot)
647 return failure();
648
649 auto iteratorTypes = op.getIteratorTypes().getValue();
650 static constexpr std::array<int64_t, 2> perm = {1, 0};
651 Location loc = op.getLoc();
652 Value lhs = op.getLhs(), rhs = op.getRhs();
653
654 using MapList = ArrayRef<ArrayRef<AffineExpr>>;
655 auto infer = [&](MapList m) {
656 return AffineMap::inferFromExprList(m, op.getContext());
657 };
658 AffineExpr m, n, k;
659 bindDims(rewriter.getContext(), m, n, k);
660 SmallVector<AffineMap> maps = op.getIndexingMapsArray();
661 //
662 // In the following we wish to make the reduction dimension innermost so we
663 // can load vectors and just fmul + reduce into a scalar.
664 //
665 if (isParallelIterator(iteratorTypes[0]) &&
666 isParallelIterator(iteratorTypes[1]) &&
667 isReductionIterator(iteratorTypes[2])) {
668 //
669 // Two outer parallel, one inner reduction (matmat flavor).
670 //
671 if (maps == infer({{m, k}, {k, n}, {m, n}})) {
672 rhs = vector::TransposeOp::create(rewriter, loc, rhs, perm);
673 } else if (maps == infer({{m, k}, {n, k}, {m, n}})) {
674 // No need to permute anything.
675 } else if (maps == infer({{k, m}, {k, n}, {m, n}})) {
676 lhs = vector::TransposeOp::create(rewriter, loc, lhs, perm);
677 rhs = vector::TransposeOp::create(rewriter, loc, rhs, perm);
678 } else if (maps == infer({{k, m}, {n, k}, {m, n}})) {
679 lhs = vector::TransposeOp::create(rewriter, loc, lhs, perm);
680 } else if (maps == infer({{m, k}, {k, n}, {n, m}})) {
681 // This is the classical row-major matmul. Just permute the lhs.
682 Value tmp = lhs;
683 lhs = vector::TransposeOp::create(rewriter, loc, rhs, perm);
684 rhs = tmp;
685 } else if (maps == infer({{m, k}, {n, k}, {n, m}})) {
686 std::swap(lhs, rhs);
687 } else if (maps == infer({{k, m}, {k, n}, {n, m}})) {
688 Value tmp = lhs;
689 lhs = vector::TransposeOp::create(rewriter, loc, rhs, perm);
690 rhs = vector::TransposeOp::create(rewriter, loc, tmp, perm);
691 } else if (maps == infer({{k, m}, {n, k}, {n, m}})) {
692 Value tmp = rhs;
693 rhs = vector::TransposeOp::create(rewriter, loc, lhs, perm);
694 lhs = tmp;
695 } else {
696 return failure();
697 }
698 } else if (isParallelIterator(iteratorTypes[0]) &&
699 isReductionIterator(iteratorTypes[1])) {
700 //
701 // One outer parallel, one inner reduction (matvec flavor)
702 //
703 if (maps == infer({{m, n}, {n}, {m}})) {
704 // No need to permute anything.
705 } else if (maps == infer({{n, m}, {n}, {m}})) {
706 lhs = vector::TransposeOp::create(rewriter, loc, lhs, perm);
707 } else if (maps == infer({{n}, {m, n}, {m}})) {
708 std::swap(lhs, rhs);
709 } else if (maps == infer({{n}, {n, m}, {m}})) {
710 std::swap(lhs, rhs);
711 lhs = vector::TransposeOp::create(rewriter, loc, lhs, perm);
712 } else {
713 return failure();
714 }
715 } else {
716 return failure();
717 }
718
719 VectorType dstType = cast<VectorType>(op.getResultType());
720 assert(dstType.getRank() >= 1 && dstType.getRank() <= 2 &&
721 "Expected dst type of rank 1 or 2");
722
723 unsigned rank = dstType.getRank();
724 unsigned dstRows = dstType.getShape()[0];
725 unsigned dstColumns = rank == 1 ? 1 : dstType.getShape()[1];
726
727 // ExtractOp does not allow dynamic indexing, we must unroll explicitly.
728 Value res = arith::ConstantOp::create(rewriter, loc, dstType,
729 rewriter.getZeroAttr(dstType));
730 bool isInt = isa<IntegerType>(dstType.getElementType());
731 arith::FastMathFlagsAttr fmf = op.getFastmathAttr();
732 llvm::SmallVector<Value> extractedCols;
733 extractedCols.reserve(dstColumns);
734 for (unsigned r = 0; r < dstRows; ++r) {
735 Value rowLhs = vector::ExtractOp::create(rewriter, op.getLoc(), lhs, r);
736 for (unsigned c = 0; c < dstColumns; ++c) {
737 // Extract each respective row and column of the LHS and RHS once to
738 // avoid having duplicate SSA values pointing to the same rows/columns.
739 if (r == 0) {
740 Value colRhs =
741 rank == 1
742 ? rhs
743 : vector::ExtractOp::create(rewriter, op.getLoc(), rhs, c);
744 extractedCols.push_back(colRhs);
745 }
746 Value extractedColRhs = extractedCols[c];
747 Value product =
748 createMul(op.getLoc(), rowLhs, extractedColRhs, isInt, rewriter, fmf);
749 Value sum = vector::ReductionOp::create(rewriter, op.getLoc(),
750 vector::CombiningKind::ADD,
751 product, op.getFastmath());
752
755 res = vector::InsertOp::create(rewriter, op.getLoc(), sum, res, pos);
756 }
757 }
758 if (auto acc = op.getAcc())
759 res = createAdd(op.getLoc(), res, acc, isInt, rewriter, fmf);
760 return res;
761}
762
763/// Lower vector.contract with all size one reduction dimensions to
764/// elementwise ops when possible.
765struct ContractOpToElementwise
766 : public MaskableOpRewritePattern<vector::ContractionOp> {
767 using MaskableOpRewritePattern::MaskableOpRewritePattern;
768 using FilterConstraintType =
769 std::function<LogicalResult(vector::ContractionOp op)>;
770 static LogicalResult defaultFilter(vector::ContractionOp op) {
771 return success();
772 }
773 ContractOpToElementwise(
774 vector::VectorContractLowering vectorContractLowering,
775 MLIRContext *context, PatternBenefit benefit = 1,
776 const FilterConstraintType &constraint = defaultFilter)
777 : MaskableOpRewritePattern<vector::ContractionOp>(context, benefit),
778 vectorContractLowering(vectorContractLowering), filter(defaultFilter) {}
779
780 FailureOr<Value>
781 matchAndRewriteMaskableOp(vector::ContractionOp contractOp,
782 MaskingOpInterface maskOp,
783 PatternRewriter &rewriter) const override {
784 // TODO: Support vector.mask.
785 if (maskOp)
786 return failure();
787
788 if (failed(filter(contractOp)))
789 return failure();
790
791 if (vectorContractLowering != vector::VectorContractLowering::ParallelArith)
792 return failure();
793
794 ArrayRef<int64_t> lhsShape = contractOp.getLhsType().getShape();
795 ArrayRef<int64_t> rhsShape = contractOp.getRhsType().getShape();
796 AffineMap lhsMap = contractOp.getIndexingMapsArray()[0];
797 AffineMap rhsMap = contractOp.getIndexingMapsArray()[1];
798 SmallVector<int64_t> lhsReductionDims =
799 getReductionIndex(lhsMap, contractOp.getIteratorTypes());
800 SmallVector<int64_t> rhsReductionDims =
801 getReductionIndex(rhsMap, contractOp.getIteratorTypes());
802 // All the reduction dimensions must be a size 1.
803 for (int64_t dim : lhsReductionDims) {
804 if (lhsShape[dim] != 1)
805 return failure();
806 }
807 for (int64_t dim : rhsReductionDims) {
808 if (rhsShape[dim] != 1)
809 return failure();
810 }
811 AffineMap accMap = contractOp.getIndexingMapsArray()[2];
812 unsigned numParallelDims = accMap.getNumResults();
813 unsigned numLhsDimToBroadcast =
814 numParallelDims - (lhsMap.getNumResults() - lhsReductionDims.size());
815 unsigned numRhsDimToBroadcast =
816 numParallelDims - (rhsMap.getNumResults() - rhsReductionDims.size());
817 SmallVector<int64_t> lhsDims;
818 SmallVector<int64_t> lhsTranspose;
819 SmallVector<int64_t> rhsDims;
820 SmallVector<int64_t> rhsTranspose;
821 for (int64_t dim : lhsReductionDims)
822 lhsTranspose.push_back(numLhsDimToBroadcast + dim);
823 for (int64_t dim : rhsReductionDims)
824 rhsTranspose.push_back(numRhsDimToBroadcast + dim);
825 // Loop through the parallel dimensions to calculate the dimensions to
826 // broadcast and to permute in order to extract only parallel dimensions.
827 for (unsigned i = 0; i < numParallelDims; i++) {
828 std::optional<unsigned> lhsDim =
829 getDimPosition(lhsMap, accMap.getDimPosition(i));
830 if (lhsDim) {
831 lhsTranspose.push_back(numLhsDimToBroadcast + *lhsDim);
832 } else {
833 // If the parallel dimension doesn't exist we will have to broadcast it.
834 lhsDims.push_back(
835 cast<VectorType>(contractOp.getResultType()).getDimSize(i));
836 lhsTranspose.push_back(lhsDims.size() - 1);
837 }
838 std::optional<unsigned> rhsDim =
839 getDimPosition(rhsMap, accMap.getDimPosition(i));
840 if (rhsDim) {
841 rhsTranspose.push_back(numRhsDimToBroadcast + *rhsDim);
842 } else {
843 // If the parallel dimension doesn't exist we will have to broadcast it.
844 rhsDims.push_back(
845 cast<VectorType>(contractOp.getResultType()).getDimSize(i));
846 rhsTranspose.push_back(rhsDims.size() - 1);
847 }
848 }
849 Value newLhs = contractOp.getLhs();
850 Value newRhs = contractOp.getRhs();
851 Location loc = contractOp.getLoc();
852 if (!lhsDims.empty()) {
853 lhsDims.append(lhsShape.begin(), lhsShape.end());
854 auto expandedType =
855 VectorType::get(lhsDims, contractOp.getLhsType().getElementType());
856 newLhs = vector::BroadcastOp::create(rewriter, loc, expandedType, newLhs);
857 }
858 if (!rhsDims.empty()) {
859 rhsDims.append(rhsShape.begin(), rhsShape.end());
860 auto expandedType =
861 VectorType::get(rhsDims, contractOp.getRhsType().getElementType());
862 newRhs = vector::BroadcastOp::create(rewriter, loc, expandedType, newRhs);
863 }
864 bool isInt = contractOp.getLhsType().getElementType().isIntOrIndex();
865 newLhs = vector::TransposeOp::create(rewriter, loc, newLhs, lhsTranspose);
866 newRhs = vector::TransposeOp::create(rewriter, loc, newRhs, rhsTranspose);
867 SmallVector<int64_t> lhsOffsets(lhsReductionDims.size(), 0);
868 SmallVector<int64_t> rhsOffsets(rhsReductionDims.size(), 0);
869 newLhs = vector::ExtractOp::create(rewriter, loc, newLhs, lhsOffsets);
870 newRhs = vector::ExtractOp::create(rewriter, loc, newRhs, rhsOffsets);
871 std::optional<Value> result =
872 createContractArithOp(loc, newLhs, newRhs, contractOp.getAcc(),
873 contractOp.getKind(), rewriter, isInt,
874 /*mask=*/Value(), contractOp.getFastmathAttr());
875 if (result)
876 return *result;
877
878 return failure();
879 }
880
881private:
882 /// Options to control the vector patterns.
883 vector::VectorContractLowering vectorContractLowering;
884 FilterConstraintType filter;
885};
886
887/// Progressive lowering of ContractionOp.
888/// One:
889/// %x = vector.contract with at least one free/batch dimension
890/// is replaced by:
891/// %a = vector.contract with one less free/batch dimension
892/// %b = vector.contract with one less free/batch dimension
893/// ..
894/// %x = combine %a %b ..
895/// until a pure contraction is reached (no free/batch dimensions),
896/// which is replaced by a dot-product.
897///
898/// This only kicks in when either vectorContractLoweringOption is set
899/// to DOT or when other contraction patterns fail.
900//
901// TODO: break down into transpose/reshape/cast ops
902// when they become available to avoid code dup
903// TODO: investigate lowering order impact on performance
904FailureOr<Value> ContractionOpLowering::matchAndRewriteMaskableOp(
905 vector::ContractionOp op, MaskingOpInterface maskOp,
906 PatternRewriter &rewriter) const {
907 if (failed(filter(op)))
908 return failure();
909
910 // TODO: support mixed mode contract lowering.
911 if (op.getLhsType().getElementType() !=
912 getElementTypeOrSelf(op.getAccType()) ||
913 op.getRhsType().getElementType() != getElementTypeOrSelf(op.getAccType()))
914 return failure();
915
916 // TODO: the code below assumes the default contraction, make sure it supports
917 // other kinds before enabling this lowering.
918 if (op.getKind() != vector::CombiningKind::ADD) {
919 return rewriter.notifyMatchFailure(
920 op, "contractions other than 'add' not supported");
921 }
922
923 // TODO: implement benefits, cost models.
924 MLIRContext *ctx = op.getContext();
925
926 ContractionOpToOuterProductOpLowering pat1(vectorContractLoweringOption, ctx);
927 FailureOr<Value> newVal1 =
928 pat1.matchAndRewriteMaskableOp(op, maskOp, rewriter);
929 if (!failed(newVal1))
930 return newVal1;
931
932 ContractionOpToDotLowering pat2(vectorContractLoweringOption, ctx);
933 FailureOr<Value> newVal2 =
934 pat2.matchAndRewriteMaskableOp(op, maskOp, rewriter);
935 if (!failed(newVal2))
936 return newVal2;
937
938 ContractOpToElementwise pat4(vectorContractLoweringOption, ctx);
939 FailureOr<Value> newVal4 =
940 pat4.matchAndRewriteMaskableOp(op, maskOp, rewriter);
941 if (!failed(newVal4))
942 return newVal4;
943
944 // Vector mask setup.
945
946 Value mask;
947 if (maskOp)
948 mask = maskOp.getMask();
949 // Find first batch dimension in LHS/RHS, and lower when found.
950 std::vector<std::pair<int64_t, int64_t>> batchDimMap = op.getBatchDimMap();
951 if (!batchDimMap.empty()) {
952 int64_t lhsIndex = batchDimMap[0].first;
953 int64_t rhsIndex = batchDimMap[0].second;
954 auto newOp = lowerParallel(rewriter, op, lhsIndex, rhsIndex, mask);
955 if (failed(newOp))
956 return failure();
957 return newOp;
958 }
959
960 // Collect contracting dimensions.
961 std::vector<std::pair<int64_t, int64_t>> contractingDimMap =
962 op.getContractingDimMap();
963 DenseSet<int64_t> lhsContractingDimSet;
964 DenseSet<int64_t> rhsContractingDimSet;
965 for (auto &dimPair : contractingDimMap) {
966 lhsContractingDimSet.insert(dimPair.first);
967 rhsContractingDimSet.insert(dimPair.second);
968 }
969
970 // Find first free dimension in LHS, and lower when found.
971 VectorType lhsType = op.getLhsType();
972 for (int64_t lhsIndex = 0, e = lhsType.getRank(); lhsIndex < e; ++lhsIndex) {
973 if (lhsContractingDimSet.count(lhsIndex) == 0) {
974 auto newOp = lowerParallel(rewriter, op, lhsIndex, /*rhsIndex=*/-1, mask);
975 if (failed(newOp))
976 return failure();
977 return newOp;
978 }
979 }
980
981 // Find first free dimension in RHS, and lower when found.
982 VectorType rhsType = op.getRhsType();
983 for (int64_t rhsIndex = 0, e = rhsType.getRank(); rhsIndex < e; ++rhsIndex) {
984 if (rhsContractingDimSet.count(rhsIndex) == 0) {
985 auto newOp = lowerParallel(rewriter, op, /*lhsIndex=*/-1, rhsIndex, mask);
986 if (failed(newOp))
987 return failure();
988 return newOp;
989 }
990 }
991
992 // Lower the first remaining reduction dimension.
993 if (!contractingDimMap.empty()) {
994 auto newOp = lowerReduction(rewriter, op, mask);
995 if (failed(newOp))
996 return failure();
997 return newOp;
998 }
999
1000 return failure();
1001}
1002
1003// Lower one parallel dimension.
1004// Incidentally also tolerates unit-size (hence trivial) reduction dimensions.
1005// TODO: consider reusing existing contract unrolling
1006FailureOr<Value> ContractionOpLowering::lowerParallel(PatternRewriter &rewriter,
1007 vector::ContractionOp op,
1008 int64_t lhsIndex,
1009 int64_t rhsIndex,
1010 Value mask) const {
1011 VectorType lhsType = op.getLhsType();
1012 VectorType rhsType = op.getRhsType();
1013 VectorType resType = cast<VectorType>(op.getResultType());
1014 // Find the iterator type index and result index.
1015 SmallVector<AffineMap> iMap = op.getIndexingMapsArray();
1016 int64_t iterIndex = -1;
1017 int64_t dimSize = -1;
1018 if (lhsIndex >= 0) {
1019 iterIndex = iMap[0].getDimPosition(lhsIndex);
1020 if (rhsIndex >= 0 && iterIndex != iMap[1].getDimPosition(rhsIndex))
1021 return rewriter.notifyMatchFailure(op, [&](Diagnostic &diag) {
1022 diag << "expected lhsIndex=" << lhsIndex << " and rhsIndex=" << rhsIndex
1023 << " to map to the same dimension";
1024 });
1025 if (lhsType.getScalableDims()[lhsIndex])
1026 return rewriter.notifyMatchFailure(op, [&](Diagnostic &diag) {
1027 diag << "Unrolling scalable dimension (lhsIndex=" << lhsIndex
1028 << ") is not supported yet";
1029 });
1030 dimSize = lhsType.getDimSize(lhsIndex);
1031 } else if (rhsIndex >= 0) {
1032 iterIndex = iMap[1].getDimPosition(rhsIndex);
1033 if (rhsType.getScalableDims()[rhsIndex])
1034 return rewriter.notifyMatchFailure(op, [&](Diagnostic &diag) {
1035 diag << "Unrolling scalable dimension (rhsIndex=" << rhsIndex
1036 << ") is not supported yet";
1037 });
1038 dimSize = rhsType.getDimSize(rhsIndex);
1039 }
1040 if (iterIndex < 0)
1041 return rewriter.notifyMatchFailure(op, [&](Diagnostic &diag) {
1042 diag << "expected either lhsIndex=" << lhsIndex
1043 << " or rhsIndex=" << rhsIndex << " to be nonnegative";
1044 });
1045 // value_or(-1) means that we tolerate a dimension not appearing
1046 // in the result map. That can't happen for actual parallel iterators, but
1047 // the caller ContractionOpLowering::matchAndRewrite is currently calling
1048 // lowerParallel also for the case of unit-size reduction dims appearing only
1049 // on one of LHS or RHS, not both. At the moment, such cases are created by
1050 // CastAwayContractionLeadingOneDim, so we need to either support that or
1051 // modify that pattern.
1052 int64_t resIndex = getResultIndex(iMap[2], iterIndex).value_or(-1);
1053 if (resIndex == -1 && dimSize != 1)
1054 return rewriter.notifyMatchFailure(op, [&](Diagnostic &diag) {
1055 diag << "expected the dimension for iterIndex=" << iterIndex
1056 << " to either appear in the result map, or to be a unit dimension";
1057 });
1058
1059 // Construct new iterator types and affine map array attribute.
1060 std::array<AffineMap, 3> lowIndexingMaps = {
1061 adjustMap(iMap[0], iterIndex, rewriter),
1062 adjustMap(iMap[1], iterIndex, rewriter),
1063 adjustMap(iMap[2], iterIndex, rewriter)};
1064 auto lowAffine = rewriter.getAffineMapArrayAttr(lowIndexingMaps);
1065 auto lowIter =
1066 rewriter.getArrayAttr(adjustIter(op.getIteratorTypes(), iterIndex));
1067 // Unroll into a series of lower dimensional vector.contract ops.
1068 Location loc = op.getLoc();
1069 Value result = arith::ConstantOp::create(rewriter, loc, resType,
1070 rewriter.getZeroAttr(resType));
1071
1072 for (int64_t d = 0; d < dimSize; ++d) {
1073 auto lhs = reshapeLoad(loc, op.getLhs(), lhsIndex, d, rewriter);
1074 auto rhs = reshapeLoad(loc, op.getRhs(), rhsIndex, d, rewriter);
1075 auto acc = reshapeLoad(loc, op.getAcc(), resIndex, d, rewriter);
1076
1077 Value lowMask;
1078 if (mask)
1079 lowMask = reshapeLoad(loc, mask, iterIndex, d, rewriter);
1080
1081 Operation *lowContract =
1082 vector::ContractionOp::create(rewriter, loc, lhs, rhs, acc, lowAffine,
1083 lowIter, op.getKind(), op.getFastmath());
1084 lowContract = maskOperation(rewriter, lowContract, lowMask);
1085 result = reshapeStore(loc, lowContract->getResult(0), result, resIndex, d,
1086 rewriter);
1087 }
1088 return result;
1089}
1090
1091// Lower one reduction dimension.
1092FailureOr<Value> ContractionOpLowering::lowerReduction(
1093 PatternRewriter &rewriter, vector::ContractionOp op, Value mask) const {
1094 auto loc = op.getLoc();
1095 VectorType lhsType = op.getLhsType();
1096 VectorType rhsType = op.getRhsType();
1097 Type resType = op.getResultType();
1098 if (isa<VectorType>(resType))
1099 return rewriter.notifyMatchFailure(op,
1100 "did not expect a VectorType result");
1101 bool isInt = isa<IntegerType>(resType);
1102 // Use iterator index 0.
1103 int64_t iterIndex = 0;
1104 SmallVector<AffineMap> iMap = op.getIndexingMapsArray();
1105 std::optional<int64_t> lookupLhs = getResultIndex(iMap[0], iterIndex);
1106 std::optional<int64_t> lookupRhs = getResultIndex(iMap[1], iterIndex);
1107 if (!lookupLhs.has_value())
1108 return rewriter.notifyMatchFailure(op, [&](Diagnostic &diag) {
1109 diag << "expected iterIndex=" << iterIndex << "to map to a LHS dimension";
1110 });
1111 if (!lookupRhs.has_value())
1112 return rewriter.notifyMatchFailure(op, [&](Diagnostic &diag) {
1113 diag << "expected iterIndex=" << iterIndex << "to map to a RHS dimension";
1114 });
1115 int64_t lhsIndex = *lookupLhs;
1116 int64_t rhsIndex = *lookupRhs;
1117 int64_t dimSize = lhsType.getDimSize(lhsIndex);
1118 if (dimSize != rhsType.getDimSize(rhsIndex))
1119 return rewriter.notifyMatchFailure(op, [&](Diagnostic &diag) {
1120 diag << "expect LHS dimension " << lhsIndex
1121 << " to have the same size as RHS dimension " << rhsIndex;
1122 });
1123 // Base case.
1124 if (lhsType.getRank() == 1) {
1125 if (rhsType.getRank() != 1)
1126 return rewriter.notifyMatchFailure(
1127 op, "When LHS has rank 1, expected also RHS to have rank 1");
1128 arith::FastMathFlagsAttr fmf = op.getFastmathAttr();
1129 Value m = createMul(loc, op.getLhs(), op.getRhs(), isInt, rewriter, fmf);
1130 auto kind = vector::CombiningKind::ADD;
1131
1132 Value acc = op.getAcc();
1133 Operation *reductionOp =
1134 acc ? vector::ReductionOp::create(rewriter, loc, kind, m, acc,
1135 op.getFastmath())
1136 : vector::ReductionOp::create(rewriter, loc, kind, m,
1137 op.getFastmath());
1138 return maskOperation(rewriter, reductionOp, mask)->getResult(0);
1139 }
1140 // Construct new iterator types and affine map array attribute.
1141 std::array<AffineMap, 3> lowIndexingMaps = {
1142 adjustMap(iMap[0], iterIndex, rewriter),
1143 adjustMap(iMap[1], iterIndex, rewriter),
1144 adjustMap(iMap[2], iterIndex, rewriter)};
1145 auto lowAffine = rewriter.getAffineMapArrayAttr(lowIndexingMaps);
1146 auto lowIter =
1147 rewriter.getArrayAttr(adjustIter(op.getIteratorTypes(), iterIndex));
1148 // Unroll into a series of lower dimensional vector.contract ops.
1149 // By feeding the initial accumulator into the first contraction,
1150 // and the result of each contraction into the next, eventually
1151 // the sum of all reductions is computed.
1152 Value result = op.getAcc();
1153 for (int64_t d = 0; d < dimSize; ++d) {
1154 auto lhs = reshapeLoad(loc, op.getLhs(), lhsIndex, d, rewriter);
1155 auto rhs = reshapeLoad(loc, op.getRhs(), rhsIndex, d, rewriter);
1156 Value newMask;
1157 if (mask)
1158 newMask = reshapeLoad(loc, mask, iterIndex, d, rewriter);
1159
1160 Operation *newContract = vector::ContractionOp::create(
1161 rewriter, loc, lhs, rhs, result, lowAffine, lowIter, op.getKind(),
1162 op.getFastmath());
1163 result = maskOperation(rewriter, newContract, newMask)->getResult(0);
1164 }
1165 return result;
1166}
1167
1168/// Progressive lowering of OuterProductOp.
1169/// One:
1170/// %x = vector.outerproduct %lhs, %rhs, %acc
1171/// is replaced by:
1172/// %z = zero-result
1173/// %0 = vector.extract %lhs[0]
1174/// %1 = vector.broadcast %0
1175/// %2 = vector.extract %acc[0]
1176/// %3 = vector.fma %1, %rhs, %2
1177/// %4 = vector.insert %3, %z[0]
1178/// ..
1179/// %x = vector.insert %.., %..[N-1]
1180///
1181class OuterProductOpLowering : public OpRewritePattern<vector::OuterProductOp> {
1182public:
1183 using Base::Base;
1184
1185 LogicalResult matchAndRewrite(vector::OuterProductOp op,
1186 PatternRewriter &rewriter) const override {
1187 VectorType resType = op.getResultVectorType();
1188 if ((resType.getShape().size() >= 2) && resType.allDimsScalable())
1189 return failure();
1190
1191 auto loc = op.getLoc();
1192
1193 VectorType lhsType = op.getOperandVectorTypeLHS();
1194 VectorType rhsType = dyn_cast<VectorType>(op.getOperandTypeRHS());
1195 Type eltType = resType.getElementType();
1196 bool isInt = isa<IntegerType, IndexType>(eltType);
1197 Value acc = op.getAcc();
1198 vector::CombiningKind kind = op.getKind();
1199
1200 // Vector mask setup.
1201 OpBuilder::InsertionGuard guard(rewriter);
1202 auto maskableOp = cast<vector::MaskableOpInterface>(op.getOperation());
1203 Operation *rootOp;
1204 Value mask;
1205 if (maskableOp.isMasked()) {
1206 rewriter.setInsertionPoint(maskableOp.getMaskingOp());
1207 rootOp = maskableOp.getMaskingOp();
1208 mask = maskableOp.getMaskingOp().getMask();
1209 } else {
1210 rootOp = op;
1211 }
1212
1213 if (!rhsType) {
1214 // Special case: AXPY operation.
1215 Value b =
1216 vector::BroadcastOp::create(rewriter, loc, lhsType, op.getRhs());
1217 std::optional<Value> mult = createContractArithOp(
1218 loc, op.getLhs(), b, acc, kind, rewriter, isInt, mask);
1219 if (!mult.has_value())
1220 return failure();
1221 rewriter.replaceOp(rootOp, *mult);
1222 return success();
1223 }
1224
1225 Value result = arith::ConstantOp::create(rewriter, loc, resType,
1226 rewriter.getZeroAttr(resType));
1227 for (int64_t d = 0, e = resType.getDimSize(0); d < e; ++d) {
1228 Value x = vector::ExtractOp::create(rewriter, loc, op.getLhs(), d);
1229 Value a = vector::BroadcastOp::create(rewriter, loc, rhsType, x);
1230 Value r = nullptr;
1231 if (acc)
1232 r = vector::ExtractOp::create(rewriter, loc, acc, d);
1233 Value extrMask;
1234 if (mask)
1235 extrMask = vector::ExtractOp::create(rewriter, loc, mask, d);
1236
1237 std::optional<Value> m = createContractArithOp(
1238 loc, a, op.getRhs(), r, kind, rewriter, isInt, extrMask);
1239 if (!m.has_value())
1240 return failure();
1241 result = vector::InsertOp::create(rewriter, loc, *m, result, d);
1242 }
1243
1244 rewriter.replaceOp(rootOp, result);
1245 return success();
1246 }
1247};
1248
1249} // namespace
1250
1252 RewritePatternSet &patterns,
1253 VectorContractLowering vectorContractLoweringOption, PatternBenefit benefit,
1254 bool disableOuterProductLowering) {
1255 if (!disableOuterProductLowering)
1256 patterns.add<OuterProductOpLowering>(patterns.getContext(), benefit);
1257 patterns.add<ContractionOpLowering, ContractionOpToOuterProductOpLowering>(
1258 vectorContractLoweringOption, patterns.getContext(), benefit);
1259}
1260
1262 RewritePatternSet &patterns, PatternBenefit benefit) {
1263 patterns.add<OuterProductOpLowering>(patterns.getContext(), benefit);
1264}
return success()
static int64_t product(ArrayRef< int64_t > vals)
lhs
b
Return true if permutation is a valid permutation of the outer_dims_perm (case OuterOrInnerPerm::Oute...
ArrayAttr()
auto load
static std::optional< int64_t > getResultIndex(AffineMap map, int64_t index)
static SmallVector< int64_t > getReductionIndex(AffineMap map, ArrayAttr iteratorTypes)
Return the positions of the reductions in the given map.
static std::optional< unsigned > getDimPosition(AffineMap map, unsigned dim)
Look for a given dimension in an affine map and return its position.
static Value reshapeStore(Location loc, Value val, Value result, int64_t index, int64_t pos, PatternRewriter &rewriter)
Inserts val into result at position pos along dimension index.
static SmallVector< Attribute > adjustIter(ArrayAttr iteratorTypes, int64_t index)
FailureOr< Value > tmatvec()
static Value createAdd(Location loc, Value x, Value y, bool isInt, PatternRewriter &rewriter, arith::FastMathFlagsAttr fmf={})
Creates an AddIOp if isInt is true otherwise create an arith::AddFOp using operands x and y.
FailureOr< Value > outerProd(Value lhs, Value rhs, Value res, VectorType lhsType, int reductionSize, std::optional< Value > maybeMask=std::nullopt)
FailureOr< Value > matvec()
static AffineMap adjustMap(AffineMap map, int64_t index, PatternRewriter &rewriter)
static Value reshapeLoad(Location loc, Value val, int64_t index, int64_t pos, PatternRewriter &rewriter)
Returns val with the dimension at position index dropped by indexing that dimension with pos.
FailureOr< Value > matmat()
Two outer parallel, one inner reduction (matmat flavor).
static std::optional< Value > createContractArithOp(Location loc, Value x, Value y, Value acc, vector::CombiningKind kind, PatternRewriter &rewriter, bool isInt, Value mask=Value(), arith::FastMathFlagsAttr fmf={})
Helper to create arithmetic operation associated with a kind of contraction.
std::optional< int64_t > getReductionSize(VectorType vecType, int64_t reductionDim)
Helper function for matmat, matvec, tmatvec. Returns the size of dimension reductionDim....
#define mul(a, b)
Base type for affine expression.
Definition AffineExpr.h:68
A multi-dimensional affine map Affine map's are immutable like Type's, and they are uniqued.
Definition AffineMap.h:46
unsigned getDimPosition(unsigned idx) const
Extracts the position of the dimensional expression at the given result, when the caller knows it is ...
static AffineMap get(MLIRContext *context)
Returns a zero result affine map with no dimensions or symbols: () -> ().
unsigned getNumDims() const
unsigned getNumResults() const
static SmallVector< AffineMap, 4 > inferFromExprList(ArrayRef< ArrayRef< AffineExpr > > exprsList, MLIRContext *context)
Returns a vector of AffineMaps; each with as many results as exprs.size(), as many dims as the larges...
TypedAttr getZeroAttr(Type type)
Definition Builders.cpp:333
ArrayAttr getArrayAttr(ArrayRef< Attribute > value)
Definition Builders.cpp:275
MLIRContext * getContext() const
Definition Builders.h:56
ArrayAttr getAffineMapArrayAttr(ArrayRef< AffineMap > values)
Definition Builders.cpp:327
This class contains all of the information necessary to report a diagnostic to the DiagnosticEngine.
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
void setInsertionPoint(Block *block, Block::iterator insertPoint)
Set the insertion point to the specified location.
Definition Builders.h:401
Operation is the basic unit of execution within MLIR.
Definition Operation.h:87
OpResult getResult(unsigned idx)
Get the 'idx'th result of this operation.
Definition Operation.h:432
This class represents the benefit of a pattern match in a unitless scheme that ranges from 0 (very li...
A special type of RewriterBase that coordinates the application of a rewrite pattern on the current I...
MLIRContext * getContext() const
RewritePatternSet & add(ConstructorArg &&arg, ConstructorArgs &&...args)
Add an instance of each of the pattern types 'Ts' to the pattern list with the given arguments.
virtual void replaceOp(Operation *op, ValueRange newValues)
Replace the results of the given (original) operation with the specified list of values (replacements...
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,...
Helper StructuredGenerator class to manipulate and rewrite ops with StructuredOpInterface.
bool iters(ArrayRef< IteratorType > its)
Instances of the Type class are uniqued, have an immutable identifier and an optional mutable compone...
Definition Types.h:74
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
This is a builder type that keeps local references to arguments.
Builder & dropDim(unsigned pos)
Erase a dim from shape @pos.
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
Value makeArithReduction(OpBuilder &b, Location loc, CombiningKind kind, Value v1, Value acc, arith::FastMathFlagsAttr fastmath=nullptr, Value mask=nullptr)
Returns the result value of reducing two scalar/vector values with the corresponding arith operation.
Operation * maskOperation(OpBuilder &builder, Operation *maskableOp, Value mask, Value passthru=Value())
Creates a vector.mask operation around a maskable operation.
bool isReductionIterator(Attribute attr)
Returns true if attr has "reduction" iterator type semantics.
Definition VectorOps.h:156
Value selectPassthru(OpBuilder &builder, Value mask, Value newValue, Value passthru)
Creates a vector select operation that picks values from newValue or passthru for each result vector ...
bool isParallelIterator(Attribute attr)
Returns true if attr has "parallel" iterator type semantics.
Definition VectorOps.h:151
void populateVectorOuterProductLoweringPatterns(RewritePatternSet &patterns, PatternBenefit benefit=1)
Populate the pattern set with the following patterns:
void populateVectorContractLoweringPatterns(RewritePatternSet &patterns, VectorContractLowering vectorContractLoweringOption, PatternBenefit benefit=1, bool disableOuterProductLowering=false)
Populate the pattern set with the following patterns:
Include the generated interface declarations.
void bindDims(MLIRContext *ctx, AffineExprTy &...exprs)
Bind a list of AffineExpr references to DimExpr at positions: [0 .
Definition AffineExpr.h:311
llvm::DenseSet< ValueT, ValueInfoT > DenseSet
Definition LLVM.h:122
Type getElementTypeOrSelf(Type type)
Return the element type or return the type itself.
AffineExpr getAffineDimExpr(unsigned position, MLIRContext *context)
These free functions allow clients of the API to not use classes in detail.
OpRewritePattern is a wrapper around RewritePattern that allows for matching and rewriting against an...
OpRewritePattern Base
Type alias to allow derived classes to inherit constructors with using Base::Base;.
A pattern for ops that implement MaskableOpInterface and that might be masked (i.e.
virtual FailureOr< Value > matchAndRewriteMaskableOp(SourceOp sourceOp, MaskingOpInterface maskingOp, PatternRewriter &rewriter) const =0