MLIR 24.0.0git
LowerVectorScan.cpp
Go to the documentation of this file.
1//===- LowerVectorScam.cpp - Lower 'vector.scan' 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.scan' operation.
11//
12//===----------------------------------------------------------------------===//
13
21#include "mlir/IR/Location.h"
24
25#define DEBUG_TYPE "vector-broadcast-lowering"
26
27using namespace mlir;
28using namespace mlir::vector;
29
30/// This function checks to see if the vector combining kind
31/// is consistent with the integer or float element type.
32static bool isValidKind(bool isInt, vector::CombiningKind kind) {
33 using vector::CombiningKind;
34 enum class KindType { FLOAT, INT, INVALID };
35 KindType type{KindType::INVALID};
36 switch (kind) {
37 case CombiningKind::MINNUMF:
38 case CombiningKind::MINIMUMF:
39 case CombiningKind::MAXNUMF:
40 case CombiningKind::MAXIMUMF:
41 case CombiningKind::MAXIMUMNUMF:
42 case CombiningKind::MINIMUMNUMF:
43 type = KindType::FLOAT;
44 break;
45 case CombiningKind::MINUI:
46 case CombiningKind::MINSI:
47 case CombiningKind::MAXUI:
48 case CombiningKind::MAXSI:
49 case CombiningKind::AND:
50 case CombiningKind::OR:
51 case CombiningKind::XOR:
52 type = KindType::INT;
53 break;
54 case CombiningKind::ADD:
55 case CombiningKind::MUL:
56 type = isInt ? KindType::INT : KindType::FLOAT;
57 break;
58 }
59 bool isValidIntKind = (type == KindType::INT) && isInt;
60 bool isValidFloatKind = (type == KindType::FLOAT) && (!isInt);
61 return (isValidIntKind || isValidFloatKind);
62}
63
64namespace {
65/// Convert vector.scan op into arith ops and vector.insert_strided_slice /
66/// vector.extract_strided_slice.
67///
68/// Example:
69///
70/// ```
71/// %0:2 = vector.scan <add>, %arg0, %arg1
72/// {inclusive = true, reduction_dim = 1} :
73/// (vector<2x3xi32>, vector<2xi32>) to (vector<2x3xi32>, vector<2xi32>)
74/// ```
75///
76/// is converted to:
77///
78/// ```
79/// %cst = arith.constant dense<0> : vector<2x3xi32>
80/// %0 = vector.extract_strided_slice %arg0
81/// {offsets = [0, 0], sizes = [2, 1], strides = [1, 1]}
82/// : vector<2x3xi32> to vector<2x1xi32>
83/// %1 = vector.insert_strided_slice %0, %cst
84/// {offsets = [0, 0], strides = [1, 1]}
85/// : vector<2x1xi32> into vector<2x3xi32>
86/// %2 = vector.extract_strided_slice %arg0
87/// {offsets = [0, 1], sizes = [2, 1], strides = [1, 1]}
88/// : vector<2x3xi32> to vector<2x1xi32>
89/// %3 = arith.muli %0, %2 : vector<2x1xi32>
90/// %4 = vector.insert_strided_slice %3, %1
91/// {offsets = [0, 1], strides = [1, 1]}
92/// : vector<2x1xi32> into vector<2x3xi32>
93/// %5 = vector.extract_strided_slice %arg0
94/// {offsets = [0, 2], sizes = [2, 1], strides = [1, 1]}
95/// : vector<2x3xi32> to vector<2x1xi32>
96/// %6 = arith.muli %3, %5 : vector<2x1xi32>
97/// %7 = vector.insert_strided_slice %6, %4
98/// {offsets = [0, 2], strides = [1, 1]}
99/// : vector<2x1xi32> into vector<2x3xi32>
100/// %8 = vector.shape_cast %6 : vector<2x1xi32> to vector<2xi32>
101/// return %7, %8 : vector<2x3xi32>, vector<2xi32>
102/// ```
103struct ScanToArithOps : public OpRewritePattern<vector::ScanOp> {
104 using Base::Base;
105
106 LogicalResult matchAndRewrite(vector::ScanOp scanOp,
107 PatternRewriter &rewriter) const override {
108 auto loc = scanOp.getLoc();
109 VectorType destType = scanOp.getDestType();
110 ArrayRef<int64_t> destShape = destType.getShape();
111 auto elType = destType.getElementType();
112 bool isInt = elType.isIntOrIndex();
113 if (!isValidKind(isInt, scanOp.getKind()))
114 return failure();
115
116 int64_t reductionDim = scanOp.getReductionDim();
117 bool inclusive = scanOp.getInclusive();
118 int64_t destRank = destType.getRank();
119 VectorType initialValueType = scanOp.getInitialValueType();
120 int64_t initialValueRank = initialValueType.getRank();
121
122 SmallVector<int64_t> reductionShape(destShape);
123 SmallVector<bool> reductionScalableDims(destType.getScalableDims());
124
125 // Check before creating any IR so that returning failure() does not
126 // violate the pattern API contract.
127 if (reductionScalableDims[reductionDim])
128 return rewriter.notifyMatchFailure(
129 scanOp, "Trying to reduce scalable dimension - not yet supported!");
130
131 VectorType resType = destType;
132 Value result = arith::ConstantOp::create(rewriter, loc, resType,
133 rewriter.getZeroAttr(resType));
134
135 // The reduction dimension, after reducing, becomes 1. It's a fixed-width
136 // dimension - no need to touch the scalability flag.
137 reductionShape[reductionDim] = 1;
138 VectorType reductionType =
139 VectorType::get(reductionShape, elType, reductionScalableDims);
140
141 SmallVector<int64_t> offsets(destRank, 0);
142 SmallVector<int64_t> strides(destRank, 1);
143 SmallVector<int64_t> sizes(destShape);
144 sizes[reductionDim] = 1;
145 ArrayAttr scanSizes = rewriter.getI64ArrayAttr(sizes);
146 ArrayAttr scanStrides = rewriter.getI64ArrayAttr(strides);
147
148 Value lastOutput, lastInput;
149 for (int i = 0; i < destShape[reductionDim]; i++) {
150 offsets[reductionDim] = i;
151 ArrayAttr scanOffsets = rewriter.getI64ArrayAttr(offsets);
152 Value input = vector::ExtractStridedSliceOp::create(
153 rewriter, loc, reductionType, scanOp.getSource(), scanOffsets,
154 scanSizes, scanStrides);
155 Value output;
156 if (i == 0) {
157 if (inclusive) {
158 output = input;
159 } else {
160 if (initialValueRank == 0) {
161 // ShapeCastOp cannot handle 0-D vectors
162 output = vector::BroadcastOp::create(rewriter, loc, input.getType(),
163 scanOp.getInitialValue());
164 } else {
165 output = vector::ShapeCastOp::create(rewriter, loc, input.getType(),
166 scanOp.getInitialValue());
167 }
168 }
169 } else {
170 Value y = inclusive ? input : lastInput;
171 output = vector::makeArithReduction(rewriter, loc, scanOp.getKind(),
172 lastOutput, y);
173 }
174 result = vector::InsertStridedSliceOp::create(rewriter, loc, output,
175 result, offsets, strides);
176 lastOutput = output;
177 lastInput = input;
178 }
179
180 Value reduction;
181 if (initialValueRank == 0) {
182 Value v = vector::ExtractOp::create(rewriter, loc, lastOutput, 0);
183 reduction =
184 vector::BroadcastOp::create(rewriter, loc, initialValueType, v);
185 } else {
186 reduction = vector::ShapeCastOp::create(rewriter, loc, initialValueType,
187 lastOutput);
188 }
189
190 rewriter.replaceOp(scanOp, {result, reduction});
191 return success();
192 }
193};
194} // namespace
195
197 RewritePatternSet &patterns, PatternBenefit benefit) {
198 patterns.add<ScanToArithOps>(patterns.getContext(), benefit);
199}
return success()
ArrayAttr()
static bool isValidKind(bool isInt, vector::CombiningKind kind)
This function checks to see if the vector combining kind is consistent with the integer or float elem...
TypedAttr getZeroAttr(Type type)
Definition Builders.cpp:333
ArrayAttr getI64ArrayAttr(ArrayRef< int64_t > values)
Definition Builders.cpp:290
This class represents the benefit of a pattern match in a unitless scheme that ranges from 0 (very li...
MLIRContext * getContext() const
RewritePatternSet & add(ConstructorArg &&arg, ConstructorArgs &&...args)
Add an instance of each of the pattern types 'Ts' to the pattern list with the given arguments.
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,...
Type getType() const
Return the type of this value.
Definition Value.h:105
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.
void populateVectorScanLoweringPatterns(RewritePatternSet &patterns, PatternBenefit benefit=1)
Populate the pattern set with the following patterns:
Include the generated interface declarations.
OpRewritePattern is a wrapper around RewritePattern that allows for matching and rewriting against an...