MLIR 24.0.0git
LegalizeForLLVMExport.cpp
Go to the documentation of this file.
1//===- LegalizeForLLVMExport.cpp - Prepare ArmSVE for LLVM translation ----===//
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
18
19using namespace mlir;
20using namespace mlir::arm_sve;
21
30 OneToOneConvertToLLVMPattern<ScalableMaskedAddIOp,
31 ScalableMaskedAddIIntrOp>;
33 OneToOneConvertToLLVMPattern<ScalableMaskedAddFOp,
34 ScalableMaskedAddFIntrOp>;
36 OneToOneConvertToLLVMPattern<ScalableMaskedSubIOp,
37 ScalableMaskedSubIIntrOp>;
39 OneToOneConvertToLLVMPattern<ScalableMaskedSubFOp,
40 ScalableMaskedSubFIntrOp>;
42 OneToOneConvertToLLVMPattern<ScalableMaskedMulIOp,
43 ScalableMaskedMulIIntrOp>;
45 OneToOneConvertToLLVMPattern<ScalableMaskedMulFOp,
46 ScalableMaskedMulFIntrOp>;
48 OneToOneConvertToLLVMPattern<ScalableMaskedSDivIOp,
49 ScalableMaskedSDivIIntrOp>;
51 OneToOneConvertToLLVMPattern<ScalableMaskedUDivIOp,
52 ScalableMaskedUDivIIntrOp>;
54 OneToOneConvertToLLVMPattern<ScalableMaskedDivFOp,
55 ScalableMaskedDivFIntrOp>;
56
57namespace {
58
59/// Unrolls a conversion to/from equivalent vector types, to allow using a
60/// conversion intrinsic that only supports 1-D vector types.
61///
62/// Example:
63/// ```
64/// %result = arm_sve.convert_to_svbool %source : vector<2x[4]xi1>
65/// ```
66/// is rewritten into:
67/// ```
68/// %cst = arith.constant dense<false> : vector<2x[16]xi1>
69/// %1 = vector.extract %source[0] : vector<[4]xi1> from vector<2x[4]xi1>
70/// %2 = "arm_sve.intr.convert.to.svbool"(%1)
71/// : (vector<[4]xi1>) -> vector<[16]xi1>
72/// %3 = vector.insert %2, %cst[0] : vector<[16]xi1> into vector<2x[16]xi1>
73/// %4 = vector.extract %source[1] : vector<[4]xi1> from vector<2x[4]xi1>
74/// %5 = "arm_sve.intr.convert.to.svbool"(%4)
75/// : (vector<[4]xi1>) -> vector<[16]xi1>
76/// %result = vector.insert %5, %3[1] : vector<[16]xi1> into vector<2x[16]xi1>
77/// ```
78template <typename Op, typename IntrOp>
79struct SvboolConversionOpLowering : public ConvertOpToLLVMPattern<Op> {
81
82 LogicalResult
83 matchAndRewrite(Op convertOp, typename Op::Adaptor,
84 ConversionPatternRewriter &rewriter) const override {
85 auto loc = convertOp.getLoc();
86
87 auto source = convertOp.getSource();
88 VectorType sourceType = source.getType();
89 VectorType resultType = convertOp.getResult().getType();
90
91 Value result = arith::ConstantOp::create(rewriter, loc, resultType,
92 rewriter.getZeroAttr(resultType));
93
94 // We want to iterate over the input vector in steps of the trailing
95 // dimension. So this creates tile shape where all leading dimensions are 1,
96 // and the trailing dimension step is the size of the dimension.
97 SmallVector<int64_t> tileShape(sourceType.getRank(), 1);
98 tileShape.back() = sourceType.getShape().back();
99
100 // Iterate over all scalable mask/predicate slices of the source vector.
102 StaticTileOffsetRange(sourceType.getShape(), tileShape)) {
103 auto extractOrInsertPosition = ArrayRef(index).drop_back();
104 auto sourceVector = vector::ExtractOp::create(rewriter, loc, source,
105 extractOrInsertPosition);
106 VectorType convertedType =
107 VectorType::Builder(llvm::cast<VectorType>(sourceVector.getType()))
108 .setDim(0, resultType.getShape().back());
109 auto convertedVector =
110 IntrOp::create(rewriter, loc, TypeRange{convertedType}, sourceVector);
111 result = vector::InsertOp::create(rewriter, loc, convertedVector, result,
112 extractOrInsertPosition);
113 }
114
115 rewriter.replaceOp(convertOp, result);
116 return success();
117 }
118};
119
120using ConvertToSvboolOpLowering =
121 SvboolConversionOpLowering<ConvertToSvboolOp, ConvertToSvboolIntrOp>;
122
123using ConvertFromSvboolOpLowering =
124 SvboolConversionOpLowering<ConvertFromSvboolOp, ConvertFromSvboolIntrOp>;
125
128
129/// Lower `arm_sve.psel` to LLVM intrinsics. This is almost a 1-to-1 conversion
130/// but first input (P1) and result predicates need conversion to/from svbool.
131struct PselOpLowering : public ConvertOpToLLVMPattern<PselOp> {
133
134 LogicalResult
135 matchAndRewrite(PselOp pselOp, PselOp::Adaptor adaptor,
136 ConversionPatternRewriter &rewriter) const override {
137 auto svboolType = VectorType::get(16, rewriter.getI1Type(), true);
138 auto loc = pselOp.getLoc();
139 auto svboolP1 = ConvertToSvboolIntrOp::create(rewriter, loc, svboolType,
140 adaptor.getP1());
141 auto indexI32 = arith::IndexCastOp::create(
142 rewriter, loc, rewriter.getI32Type(), pselOp.getIndex());
143 auto pselIntr = PselIntrOp::create(rewriter, loc, svboolType, svboolP1,
144 pselOp.getP2(), indexI32);
145 rewriter.replaceOpWithNewOp<ConvertFromSvboolIntrOp>(
146 pselOp, adaptor.getP1().getType(), pselIntr);
147 return success();
148 }
149};
150
151/// Converts `vector.create_mask` ops that match the size of an SVE predicate
152/// to the `whilelt` intrinsic. This produces more canonical codegen than the
153/// generic LLVM lowering, see https://github.com/llvm/llvm-project/issues/81840
154/// for more details. Note that we can't use (the more general) active.lane.mask
155/// as its semantics don't neatly map on to `vector.create_mask`, as it does an
156/// unsigned comparison (whereas `create_mask` is signed), and is UB/posion if
157/// `n` is zero (whereas `create_mask` just returns an all-false mask).
158struct CreateMaskOpLowering
159 : public ConvertOpToLLVMPattern<vector::CreateMaskOp> {
161
162 LogicalResult
163 matchAndRewrite(vector::CreateMaskOp createMaskOp,
164 vector::CreateMaskOp::Adaptor adaptor,
165 ConversionPatternRewriter &rewriter) const override {
166 auto maskType = createMaskOp.getVectorType();
167 if (maskType.getRank() != 1 || !maskType.isScalable())
168 return rewriter.notifyMatchFailure(createMaskOp, "not 1-D and scalable");
169
170 // TODO: Support masks which are multiples of SVE predicates.
171 auto maskBaseSize = maskType.getDimSize(0);
172 if (maskBaseSize < 2 || maskBaseSize > 16 ||
173 !llvm::isPowerOf2_32(uint32_t(maskBaseSize)))
174 return rewriter.notifyMatchFailure(createMaskOp,
175 "not SVE predicate-sized");
176
177 auto loc = createMaskOp.getLoc();
178 auto zero = LLVM::ZeroOp::create(rewriter, loc, rewriter.getI64Type());
179 rewriter.replaceOpWithNewOp<WhileLTIntrOp>(createMaskOp, maskType, zero,
180 adaptor.getOperands()[0]);
181 return success();
182 }
183};
184
185} // namespace
186
187/// Populate the given list with patterns that convert from ArmSVE to LLVM.
189 const LLVMTypeConverter &converter, RewritePatternSet &patterns) {
190 // Populate conversion patterns
191
192 // clang-format off
193 patterns.add<ConvertFromSvboolOpLowering,
194 ConvertToSvboolOpLowering,
196 PselOpLowering,
210 ZipX2OpLowering,
211 ZipX4OpLowering,
212 SdotOpLowering>(converter);
213 // Add vector.create_mask conversion with a high benefit as it produces much
214 // nicer code than the generic lowering.
215 patterns.add<CreateMaskOpLowering>(converter, /*benefit=*/4096);
216 // clang-format on
217}
218
221 // clang-format off
222 target.addLegalOp<BfmmlaOp,
223 ConvertFromSvboolIntrOp,
224 ConvertToSvboolIntrOp,
225 DupQLaneIntrOp,
226 PselIntrOp,
227 ScalableMaskedAddFIntrOp,
228 ScalableMaskedAddIIntrOp,
229 ScalableMaskedDivFIntrOp,
230 ScalableMaskedMulFIntrOp,
231 ScalableMaskedMulIIntrOp,
232 ScalableMaskedSDivIIntrOp,
233 ScalableMaskedSubFIntrOp,
234 ScalableMaskedSubIIntrOp,
235 ScalableMaskedUDivIIntrOp,
236 SmmlaIntrOp,
237 UdotIntrOp,
238 UmmlaIntrOp,
239 UsmmlaIntrOp,
240 WhileLTIntrOp,
241 ZipX2IntrOp,
242 ZipX4IntrOp,
243 SdotIntrOp>();
244 target.addIllegalOp<ConvertFromSvboolOp,
245 ConvertToSvboolOp,
246 DupQLaneOp,
247 PselOp,
248 ScalableMaskedAddFOp,
249 ScalableMaskedAddIOp,
250 ScalableMaskedDivFOp,
251 ScalableMaskedMulFOp,
252 ScalableMaskedMulIOp,
253 ScalableMaskedSDivIOp,
254 ScalableMaskedSubFOp,
255 ScalableMaskedSubIOp,
256 ScalableMaskedUDivIOp,
257 SmmlaOp,
258 UdotOp,
259 UmmlaOp,
260 UsmmlaOp,
261 ZipX2Op,
262 ZipX4Op,
263 SdotOp>();
264 // clang-format on
265}
return success()
OneToOneConvertToLLVMPattern< ScalableMaskedMulIOp, ScalableMaskedMulIIntrOp > ScalableMaskedMulIOpLowering
OneToOneConvertToLLVMPattern< ScalableMaskedAddFOp, ScalableMaskedAddFIntrOp > ScalableMaskedAddFOpLowering
OneToOneConvertToLLVMPattern< UmmlaOp, UmmlaIntrOp > UmmlaOpLowering
OneToOneConvertToLLVMPattern< ScalableMaskedMulFOp, ScalableMaskedMulFIntrOp > ScalableMaskedMulFOpLowering
OneToOneConvertToLLVMPattern< SdotOp, SdotIntrOp > SdotOpLowering
OneToOneConvertToLLVMPattern< ScalableMaskedUDivIOp, ScalableMaskedUDivIIntrOp > ScalableMaskedUDivIOpLowering
OneToOneConvertToLLVMPattern< UsmmlaOp, UsmmlaIntrOp > UsmmlaOpLowering
OneToOneConvertToLLVMPattern< SmmlaOp, SmmlaIntrOp > SmmlaOpLowering
OneToOneConvertToLLVMPattern< ScalableMaskedDivFOp, ScalableMaskedDivFIntrOp > ScalableMaskedDivFOpLowering
OneToOneConvertToLLVMPattern< DupQLaneOp, DupQLaneIntrOp > DupQLaneLowering
OneToOneConvertToLLVMPattern< ScalableMaskedSDivIOp, ScalableMaskedSDivIIntrOp > ScalableMaskedSDivIOpLowering
OneToOneConvertToLLVMPattern< ScalableMaskedAddIOp, ScalableMaskedAddIIntrOp > ScalableMaskedAddIOpLowering
OneToOneConvertToLLVMPattern< UdotOp, UdotIntrOp > UdotOpLowering
OneToOneConvertToLLVMPattern< ScalableMaskedSubIOp, ScalableMaskedSubIIntrOp > ScalableMaskedSubIOpLowering
OneToOneConvertToLLVMPattern< ScalableMaskedSubFOp, ScalableMaskedSubFIntrOp > ScalableMaskedSubFOpLowering
Utility class for operation conversions targeting the LLVM dialect that match exactly one source oper...
Definition Pattern.h:233
ConvertOpToLLVMPattern(const LLVMTypeConverter &typeConverter, PatternBenefit benefit=1)
Definition Pattern.h:239
Derived class that automatically populates legalization information for different LLVM ops.
Conversion from types to the LLVM IR dialect.
Generic implementation of one-to-one conversion from "SourceOp" to "TargetOp" where the latter belong...
Definition Pattern.h:336
Location getLoc()
The source location the operation was defined or derived from.
This provides public APIs that all operations should have.
RewritePatternSet & add(ConstructorArg &&arg, ConstructorArgs &&...args)
Add an instance of each of the pattern types 'Ts' to the pattern list with the given arguments.
A range-style iterator that allows for iterating over the offsets of all potential tiles of size tile...
This class provides an abstraction over the various different ranges of value types.
Definition TypeRange.h:40
This class represents an instance of an SSA value in the MLIR system, representing a computable value...
Definition Value.h:96
This is a builder type that keeps local references to arguments.
Builder & setDim(unsigned pos, int64_t val)
Set a dim in shape @pos to val.
Include the generated interface declarations.
void configureArmSVELegalizeForExportTarget(LLVMConversionTarget &target)
Configure the target to support lowering ArmSVE ops to ops that map to LLVM intrinsics.
Type getType(OpFoldResult ofr)
Returns the int type of the integer in ofr.
Definition Utils.cpp:311
void populateArmSVELegalizeForLLVMExportPatterns(const LLVMTypeConverter &converter, RewritePatternSet &patterns)
Collect a set of patterns to lower ArmSVE ops to ops that map to LLVM intrinsics.