31 ScalableMaskedAddIIntrOp>;
34 ScalableMaskedAddFIntrOp>;
37 ScalableMaskedSubIIntrOp>;
40 ScalableMaskedSubFIntrOp>;
43 ScalableMaskedMulIIntrOp>;
46 ScalableMaskedMulFIntrOp>;
49 ScalableMaskedSDivIIntrOp>;
52 ScalableMaskedUDivIIntrOp>;
55 ScalableMaskedDivFIntrOp>;
78template <
typename Op,
typename IntrOp>
83 matchAndRewrite(
Op convertOp,
typename Op::Adaptor,
84 ConversionPatternRewriter &rewriter)
const override {
85 auto loc = convertOp.
getLoc();
87 auto source = convertOp.getSource();
88 VectorType sourceType = source.getType();
89 VectorType resultType = convertOp.getResult().getType();
91 Value result = arith::ConstantOp::create(rewriter, loc, resultType,
92 rewriter.getZeroAttr(resultType));
98 tileShape.back() = sourceType.getShape().back();
104 auto sourceVector = vector::ExtractOp::create(rewriter, loc, source,
105 extractOrInsertPosition);
106 VectorType convertedType =
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);
115 rewriter.replaceOp(convertOp,
result);
120using ConvertToSvboolOpLowering =
121 SvboolConversionOpLowering<ConvertToSvboolOp, ConvertToSvboolIntrOp>;
123using ConvertFromSvboolOpLowering =
124 SvboolConversionOpLowering<ConvertFromSvboolOp, ConvertFromSvboolIntrOp>;
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,
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);
158struct CreateMaskOpLowering
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");
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");
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]);
193 patterns.
add<ConvertFromSvboolOpLowering,
194 ConvertToSvboolOpLowering,
215 patterns.
add<CreateMaskOpLowering>(converter, 4096);
222 target.addLegalOp<BfmmlaOp,
223 ConvertFromSvboolIntrOp,
224 ConvertToSvboolIntrOp,
227 ScalableMaskedAddFIntrOp,
228 ScalableMaskedAddIIntrOp,
229 ScalableMaskedDivFIntrOp,
230 ScalableMaskedMulFIntrOp,
231 ScalableMaskedMulIIntrOp,
232 ScalableMaskedSDivIIntrOp,
233 ScalableMaskedSubFIntrOp,
234 ScalableMaskedSubIIntrOp,
235 ScalableMaskedUDivIIntrOp,
244 target.addIllegalOp<ConvertFromSvboolOp,
248 ScalableMaskedAddFOp,
249 ScalableMaskedAddIOp,
250 ScalableMaskedDivFOp,
251 ScalableMaskedMulFOp,
252 ScalableMaskedMulIOp,
253 ScalableMaskedSDivIOp,
254 ScalableMaskedSubFOp,
255 ScalableMaskedSubIOp,
256 ScalableMaskedUDivIOp,
Utility class for operation conversions targeting the LLVM dialect that match exactly one source oper...
ConvertOpToLLVMPattern(const LLVMTypeConverter &typeConverter, PatternBenefit benefit=1)
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...
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.
This class represents an instance of an SSA value in the MLIR system, representing a computable value...
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.
void populateArmSVELegalizeForLLVMExportPatterns(const LLVMTypeConverter &converter, RewritePatternSet &patterns)
Collect a set of patterns to lower ArmSVE ops to ops that map to LLVM intrinsics.