MLIR 24.0.0git
ConvertVectorToLLVM.cpp
Go to the documentation of this file.
1//===- VectorToLLVM.cpp - Conversion from Vector to the LLVM dialect ------===//
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
10
32#include "llvm/ADT/APFloat.h"
33#include "llvm/IR/LLVMContext.h"
34#include "llvm/Support/Casting.h"
35
36#include <optional>
37
38using namespace mlir;
39using namespace mlir::vector;
40
41// Helper that picks the proper sequence for inserting.
42static Value insertOne(ConversionPatternRewriter &rewriter,
43 const LLVMTypeConverter &typeConverter, Location loc,
44 Value val1, Value val2, Type llvmType, int64_t rank,
45 int64_t pos) {
46 assert(rank > 0 && "0-D vector corner case should have been handled already");
47 if (rank == 1) {
48 Type idxType = typeConverter.convertType(rewriter.getIndexType());
49 auto constant = LLVM::ConstantOp::create(
50 rewriter, loc, idxType, rewriter.getIntegerAttr(idxType, pos));
51 return LLVM::InsertElementOp::create(rewriter, loc, llvmType, val1, val2,
52 constant);
53 }
54 return LLVM::InsertValueOp::create(rewriter, loc, val1, val2, pos);
55}
56
57// Helper that picks the proper sequence for extracting.
58static Value extractOne(ConversionPatternRewriter &rewriter,
59 const LLVMTypeConverter &typeConverter, Location loc,
60 Value val, Type llvmType, int64_t rank, int64_t pos) {
61 if (rank <= 1) {
62 Type idxType = typeConverter.convertType(rewriter.getIndexType());
63 auto constant = LLVM::ConstantOp::create(
64 rewriter, loc, idxType, rewriter.getIntegerAttr(idxType, pos));
65 return LLVM::ExtractElementOp::create(rewriter, loc, llvmType, val,
66 constant);
67 }
68 return LLVM::ExtractValueOp::create(rewriter, loc, val, pos);
69}
70
71// Helper that returns data layout alignment of a vector.
72LogicalResult getVectorAlignment(const LLVMTypeConverter &typeConverter,
73 VectorType vectorType, unsigned &align) {
74 Type convertedVectorTy = typeConverter.convertType(vectorType);
75 if (!convertedVectorTy)
76 return failure();
77
78 llvm::LLVMContext llvmContext;
79 align = LLVM::TypeToLLVMIRTranslator(llvmContext)
80 .getPreferredAlignment(convertedVectorTy,
81 typeConverter.getDataLayout());
82
83 return success();
84}
85
86// Helper that returns data layout alignment of a memref.
87LogicalResult getMemRefAlignment(const LLVMTypeConverter &typeConverter,
88 MemRefType memrefType, unsigned &align) {
89 Type elementTy = typeConverter.convertType(memrefType.getElementType());
90 if (!elementTy)
91 return failure();
92
93 // TODO: this should use the MLIR data layout when it becomes available and
94 // stop depending on translation.
95 llvm::LLVMContext llvmContext;
96 align = LLVM::TypeToLLVMIRTranslator(llvmContext)
97 .getPreferredAlignment(elementTy, typeConverter.getDataLayout());
98 return success();
99}
100
101// Helper to resolve the alignment for vector load/store, gather and scatter
102// ops. If useVectorAlignment is true, get the preferred alignment for the
103// vector type in the operation. This option is used for hardware backends with
104// vectorization. Otherwise, use the preferred alignment of the element type of
105// the memref. Note that if you choose to use vector alignment, the shape of the
106// vector type must be resolved before the ConvertVectorToLLVM pass is run.
107LogicalResult getVectorToLLVMAlignment(const LLVMTypeConverter &typeConverter,
108 VectorType vectorType,
109 MemRefType memrefType, unsigned &align,
110 bool useVectorAlignment) {
111 if (useVectorAlignment) {
112 if (failed(getVectorAlignment(typeConverter, vectorType, align))) {
113 return failure();
114 }
115 } else {
116 if (failed(getMemRefAlignment(typeConverter, memrefType, align))) {
117 return failure();
118 }
119 }
120 return success();
121}
122
123// Check if the last stride is non-unit and has a valid memory space.
124static LogicalResult isMemRefTypeSupported(MemRefType memRefType,
125 const LLVMTypeConverter &converter) {
126 if (!memRefType.isLastDimUnitStride())
127 return failure();
128 if (failed(converter.getMemRefAddressSpace(memRefType)))
129 return failure();
130 return success();
131}
132
133// Add an index vector component to a base pointer.
134static Value getIndexedPtrs(ConversionPatternRewriter &rewriter, Location loc,
135 const LLVMTypeConverter &typeConverter,
136 MemRefType memRefType, Value llvmMemref, Value base,
137 Value index, VectorType vectorType) {
138 assert(succeeded(isMemRefTypeSupported(memRefType, typeConverter)) &&
139 "unsupported memref type");
140 assert(vectorType.getRank() == 1 && "expected a 1-d vector type");
141 auto pType = MemRefDescriptor(llvmMemref).getElementPtrType();
142 auto ptrsType =
143 LLVM::getVectorType(pType, vectorType.getDimSize(0),
144 /*isScalable=*/vectorType.getScalableDims()[0]);
145 return LLVM::GEPOp::create(
146 rewriter, loc, ptrsType,
147 typeConverter.convertType(memRefType.getElementType()), base, index);
148}
149
150/// Convert `foldResult` into a Value. Integer attribute is converted to
151/// an LLVM constant op.
153 OpFoldResult foldResult) {
154 if (auto attr = dyn_cast<Attribute>(foldResult)) {
155 auto intAttr = cast<IntegerAttr>(attr);
156 return LLVM::ConstantOp::create(builder, loc, intAttr).getResult();
157 }
158
159 return cast<Value>(foldResult);
160}
161
162namespace {
163
164/// Trivial Vector to LLVM conversions
165using VectorScaleOpConversion =
167
168/// Conversion pattern for a vector.bitcast.
169class VectorBitCastOpConversion
170 : public ConvertOpToLLVMPattern<vector::BitCastOp> {
171public:
172 using ConvertOpToLLVMPattern<vector::BitCastOp>::ConvertOpToLLVMPattern;
173
174 LogicalResult
175 matchAndRewrite(vector::BitCastOp bitCastOp, OpAdaptor adaptor,
176 ConversionPatternRewriter &rewriter) const override {
177 // Only 0-D and 1-D vectors can be lowered to LLVM.
178 VectorType resultTy = bitCastOp.getResultVectorType();
179 if (resultTy.getRank() > 1)
180 return failure();
181 Type newResultTy = typeConverter->convertType(resultTy);
182 rewriter.replaceOpWithNewOp<LLVM::BitcastOp>(bitCastOp, newResultTy,
183 adaptor.getOperands()[0]);
184 return success();
185 }
186};
187
188/// Overloaded utility that replaces a vector.load, vector.store,
189/// vector.maskedload and vector.maskedstore with their respective LLVM
190/// couterparts.
191static void replaceLoadOrStoreOp(vector::LoadOp loadOp,
192 vector::LoadOpAdaptor adaptor,
193 VectorType vectorTy, Value ptr, unsigned align,
194 ConversionPatternRewriter &rewriter) {
195 rewriter.replaceOpWithNewOp<LLVM::LoadOp>(loadOp, vectorTy, ptr, align,
196 /*volatile_=*/false,
197 loadOp.getNontemporal());
198}
199
200static void replaceLoadOrStoreOp(vector::MaskedLoadOp loadOp,
201 vector::MaskedLoadOpAdaptor adaptor,
202 VectorType vectorTy, Value ptr, unsigned align,
203 ConversionPatternRewriter &rewriter) {
204 rewriter.replaceOpWithNewOp<LLVM::MaskedLoadOp>(
205 loadOp, vectorTy, ptr, adaptor.getMask(), adaptor.getPassThru(), align);
206}
207
208static void replaceLoadOrStoreOp(vector::StoreOp storeOp,
209 vector::StoreOpAdaptor adaptor,
210 VectorType vectorTy, Value ptr, unsigned align,
211 ConversionPatternRewriter &rewriter) {
212 rewriter.replaceOpWithNewOp<LLVM::StoreOp>(storeOp, adaptor.getValueToStore(),
213 ptr, align, /*volatile_=*/false,
214 storeOp.getNontemporal());
215}
216
217static void replaceLoadOrStoreOp(vector::MaskedStoreOp storeOp,
218 vector::MaskedStoreOpAdaptor adaptor,
219 VectorType vectorTy, Value ptr, unsigned align,
220 ConversionPatternRewriter &rewriter) {
221 rewriter.replaceOpWithNewOp<LLVM::MaskedStoreOp>(
222 storeOp, adaptor.getValueToStore(), ptr, adaptor.getMask(), align);
223}
224
225/// Conversion pattern for a vector.load, vector.store, vector.maskedload, and
226/// vector.maskedstore.
227template <class LoadOrStoreOp>
228class VectorLoadStoreConversion : public ConvertOpToLLVMPattern<LoadOrStoreOp> {
229public:
230 explicit VectorLoadStoreConversion(const LLVMTypeConverter &typeConv,
231 bool useVectorAlign,
232 bool enableGEPInboundsNuw)
233 : ConvertOpToLLVMPattern<LoadOrStoreOp>(typeConv),
234 useVectorAlignment(useVectorAlign),
235 enableGEPInboundsNuw(enableGEPInboundsNuw) {}
236
237 LogicalResult
238 matchAndRewrite(LoadOrStoreOp loadOrStoreOp,
239 typename LoadOrStoreOp::Adaptor adaptor,
240 ConversionPatternRewriter &rewriter) const override {
241 // Only 1-D vectors can be lowered to LLVM.
242 VectorType vectorTy = loadOrStoreOp.getVectorType();
243 if (vectorTy.getRank() > 1)
244 return failure();
245
246 auto loc = loadOrStoreOp->getLoc();
247 MemRefType memRefTy = loadOrStoreOp.getMemRefType();
248
249 // Resolve alignment.
250 // Explicit alignment takes priority over use-vector-alignment.
251 unsigned align = loadOrStoreOp.getAlignment().value_or(0);
252 if (!align &&
253 failed(getVectorToLLVMAlignment(*this->getTypeConverter(), vectorTy,
254 memRefTy, align, useVectorAlignment)))
255 return rewriter.notifyMatchFailure(loadOrStoreOp,
256 "could not resolve alignment");
257
258 // Resolve address.
259 // When --enable-gep-inbounds-nuw is set, emit inbounds|nuw on the GEP so
260 // LLVM can apply no-wrap optimizations on the index arithmetic. This
261 // assumes 0 <= idx < dim_size and non-negative strides; the caller is
262 // responsible for ensuring those conditions hold. Masked variants are
263 // designed for near-boundary access and never receive these flags.
264 LLVM::GEPNoWrapFlags noWrapFlags = LLVM::GEPNoWrapFlags::none;
265 if constexpr (std::is_same_v<LoadOrStoreOp, vector::LoadOp> ||
266 std::is_same_v<LoadOrStoreOp, vector::StoreOp>) {
267 // The verifier (verifyLoadStoreMemRefLayout) guarantees that the
268 // trailing (most minor) stride of the memref is 1. Assert to make
269 // the invariant explicit in the lowering code.
270 auto [strides, offset] = memRefTy.getStridesAndOffset();
271 assert((strides.empty() || strides.back() == 1) &&
272 "vector.load/store requires unit trailing memref stride");
273 if (enableGEPInboundsNuw) {
274 noWrapFlags = noWrapFlags | LLVM::GEPNoWrapFlags::inbounds;
275
276 // `nuw` additionally requires non-negative strides.
277 assert(
278 !(memref::hasNegativeStaticStride(memRefTy)) &&
279 "Invalid MemRef type - should have been rejected by Op verifier.");
280 noWrapFlags = noWrapFlags | LLVM::GEPNoWrapFlags::nuw;
281 }
282 }
283 auto vtype = cast<VectorType>(
284 this->typeConverter->convertType(loadOrStoreOp.getVectorType()));
285 Value dataPtr =
286 this->getStridedElementPtr(rewriter, loc, memRefTy, adaptor.getBase(),
287 adaptor.getIndices(), noWrapFlags);
288 replaceLoadOrStoreOp(loadOrStoreOp, adaptor, vtype, dataPtr, align,
289 rewriter);
290 return success();
291 }
292
293private:
294 // If true, use the preferred alignment of the vector type.
295 // If false, use the preferred alignment of the element type
296 // of the memref. This flag is intended for use with hardware
297 // backends that require alignment of vector operations.
298 const bool useVectorAlignment;
299 const bool enableGEPInboundsNuw;
300};
301
302/// Conversion pattern for a vector.gather.
303class VectorGatherOpConversion
304 : public ConvertOpToLLVMPattern<vector::GatherOp> {
305public:
306 explicit VectorGatherOpConversion(const LLVMTypeConverter &typeConv,
307 bool useVectorAlign)
308 : ConvertOpToLLVMPattern<vector::GatherOp>(typeConv),
309 useVectorAlignment(useVectorAlign) {}
310 using ConvertOpToLLVMPattern<vector::GatherOp>::ConvertOpToLLVMPattern;
311
312 LogicalResult
313 matchAndRewrite(vector::GatherOp gather, OpAdaptor adaptor,
314 ConversionPatternRewriter &rewriter) const override {
315 Location loc = gather->getLoc();
316 MemRefType memRefType = dyn_cast<MemRefType>(gather.getBaseType());
317 assert(memRefType && "The base should be bufferized");
318
319 // TODO: Add support for strided MemRef.
320 if (failed(isMemRefTypeSupported(memRefType, *this->getTypeConverter())))
321 return rewriter.notifyMatchFailure(gather, "memref type not supported");
322
323 VectorType vType = gather.getVectorType();
324 if (vType.getRank() > 1) {
325 return rewriter.notifyMatchFailure(
326 gather, "only 1-D vectors can be lowered to LLVM");
327 }
328
329 // Resolve alignment.
330 // Explicit alignment takes priority over use-vector-alignment.
331 unsigned align = gather.getAlignment().value_or(0);
332 if (!align &&
333 failed(getVectorToLLVMAlignment(*this->getTypeConverter(), vType,
334 memRefType, align, useVectorAlignment)))
335 return rewriter.notifyMatchFailure(gather, "could not resolve alignment");
336
337 // Resolve address.
338 Value ptr = getStridedElementPtr(rewriter, loc, memRefType,
339 adaptor.getBase(), adaptor.getOffsets());
340 Value base = adaptor.getBase();
341 Value ptrs =
342 getIndexedPtrs(rewriter, loc, *this->getTypeConverter(), memRefType,
343 base, ptr, adaptor.getIndices(), vType);
344
345 // Replace with the gather intrinsic.
346 rewriter.replaceOpWithNewOp<LLVM::masked_gather>(
347 gather, typeConverter->convertType(vType), ptrs, adaptor.getMask(),
348 adaptor.getPassThru(), align);
349 return success();
350 }
351
352private:
353 // If true, use the preferred alignment of the vector type.
354 // If false, use the preferred alignment of the element type
355 // of the memref. This flag is intended for use with hardware
356 // backends that require alignment of vector operations.
357 const bool useVectorAlignment;
358};
359
360/// Conversion pattern for a vector.scatter.
361class VectorScatterOpConversion
362 : public ConvertOpToLLVMPattern<vector::ScatterOp> {
363public:
364 explicit VectorScatterOpConversion(const LLVMTypeConverter &typeConv,
365 bool useVectorAlign)
366 : ConvertOpToLLVMPattern<vector::ScatterOp>(typeConv),
367 useVectorAlignment(useVectorAlign) {}
368
369 using ConvertOpToLLVMPattern<vector::ScatterOp>::ConvertOpToLLVMPattern;
370
371 LogicalResult
372 matchAndRewrite(vector::ScatterOp scatter, OpAdaptor adaptor,
373 ConversionPatternRewriter &rewriter) const override {
374 auto loc = scatter->getLoc();
375 auto memRefType = dyn_cast<MemRefType>(scatter.getBaseType());
376 assert(memRefType && "The base should be bufferized");
377
378 // TODO: Add support for strided MemRef.
379 if (failed(isMemRefTypeSupported(memRefType, *this->getTypeConverter())))
380 return rewriter.notifyMatchFailure(scatter, "memref type not supported");
381
382 VectorType vType = scatter.getVectorType();
383 if (vType.getRank() > 1) {
384 return rewriter.notifyMatchFailure(
385 scatter, "only 1-D vectors can be lowered to LLVM");
386 }
387
388 // Resolve alignment.
389 // Explicit alignment takes priority over use-vector-alignment.
390 unsigned align = scatter.getAlignment().value_or(0);
391 if (!align &&
392 failed(getVectorToLLVMAlignment(*this->getTypeConverter(), vType,
393 memRefType, align, useVectorAlignment)))
394 return rewriter.notifyMatchFailure(scatter,
395 "could not resolve alignment");
396
397 // Resolve address.
398 Value ptr = getStridedElementPtr(rewriter, loc, memRefType,
399 adaptor.getBase(), adaptor.getOffsets());
400 Value ptrs =
401 getIndexedPtrs(rewriter, loc, *this->getTypeConverter(), memRefType,
402 adaptor.getBase(), ptr, adaptor.getIndices(), vType);
403
404 // Replace with the scatter intrinsic.
405 rewriter.replaceOpWithNewOp<LLVM::masked_scatter>(
406 scatter, adaptor.getValueToStore(), ptrs, adaptor.getMask(), align);
407 return success();
408 }
409
410private:
411 // If true, use the preferred alignment of the vector type.
412 // If false, use the preferred alignment of the element type
413 // of the memref. This flag is intended for use with hardware
414 // backends that require alignment of vector operations.
415 const bool useVectorAlignment;
416};
417
418/// Conversion pattern for a vector.expandload.
419class VectorExpandLoadOpConversion
420 : public ConvertOpToLLVMPattern<vector::ExpandLoadOp> {
421public:
422 using ConvertOpToLLVMPattern<vector::ExpandLoadOp>::ConvertOpToLLVMPattern;
423
424 LogicalResult
425 matchAndRewrite(vector::ExpandLoadOp expand, OpAdaptor adaptor,
426 ConversionPatternRewriter &rewriter) const override {
427 auto loc = expand->getLoc();
428 MemRefType memRefType = expand.getMemRefType();
429
430 // Resolve address.
431 auto vtype = typeConverter->convertType(expand.getVectorType());
432 Value ptr = getStridedElementPtr(rewriter, loc, memRefType,
433 adaptor.getBase(), adaptor.getIndices());
434
435 // From:
436 // https://llvm.org/docs/LangRef.html#llvm-masked-expandload-intrinsics
437 // The pointer alignment defaults to 1.
438 uint64_t alignment = expand.getAlignment().value_or(1);
439
440 rewriter.replaceOpWithNewOp<LLVM::masked_expandload>(
441 expand, vtype, ptr, adaptor.getMask(), adaptor.getPassThru(),
442 alignment);
443 return success();
444 }
445};
446
447/// Conversion pattern for a vector.compressstore.
448class VectorCompressStoreOpConversion
449 : public ConvertOpToLLVMPattern<vector::CompressStoreOp> {
450public:
451 using ConvertOpToLLVMPattern<vector::CompressStoreOp>::ConvertOpToLLVMPattern;
452
453 LogicalResult
454 matchAndRewrite(vector::CompressStoreOp compress, OpAdaptor adaptor,
455 ConversionPatternRewriter &rewriter) const override {
456 auto loc = compress->getLoc();
457 MemRefType memRefType = compress.getMemRefType();
458
459 // Resolve address.
460 Value ptr = getStridedElementPtr(rewriter, loc, memRefType,
461 adaptor.getBase(), adaptor.getIndices());
462
463 // From:
464 // https://llvm.org/docs/LangRef.html#llvm-masked-compressstore-intrinsics
465 // The pointer alignment defaults to 1.
466 uint64_t alignment = compress.getAlignment().value_or(1);
467
468 rewriter.replaceOpWithNewOp<LLVM::masked_compressstore>(
469 compress, adaptor.getValueToStore(), ptr, adaptor.getMask(), alignment);
470 return success();
471 }
472};
473
474/// Reduction neutral classes for overloading.
475class ReductionNeutralZero {};
476class ReductionNeutralIntOne {};
477class ReductionNeutralFPOne {};
478class ReductionNeutralAllOnes {};
479class ReductionNeutralSIntMin {};
480class ReductionNeutralUIntMin {};
481class ReductionNeutralSIntMax {};
482class ReductionNeutralUIntMax {};
483class ReductionNeutralFPQNaN {};
484class ReductionNeutralFPNegQNaN {};
485class ReductionNeutralFPNegInf {};
486class ReductionNeutralFPPosInf {};
487class ReductionNeutralFPLowestFinite {};
488class ReductionNeutralFPLargestFinite {};
489
490/// Create the reduction neutral zero value.
491static Value createReductionNeutralValue(ReductionNeutralZero neutral,
492 ConversionPatternRewriter &rewriter,
493 Location loc, Type llvmType) {
494 return LLVM::ConstantOp::create(rewriter, loc, llvmType,
495 rewriter.getZeroAttr(llvmType));
496}
497
498/// Create the reduction neutral integer one value.
499static Value createReductionNeutralValue(ReductionNeutralIntOne neutral,
500 ConversionPatternRewriter &rewriter,
501 Location loc, Type llvmType) {
502 return LLVM::ConstantOp::create(rewriter, loc, llvmType,
503 rewriter.getIntegerAttr(llvmType, 1));
504}
505
506/// Create the reduction neutral fp one value.
507static Value createReductionNeutralValue(ReductionNeutralFPOne neutral,
508 ConversionPatternRewriter &rewriter,
509 Location loc, Type llvmType) {
510 return LLVM::ConstantOp::create(rewriter, loc, llvmType,
511 rewriter.getFloatAttr(llvmType, 1.0));
512}
513
514/// Create the reduction neutral all-ones value.
515static Value createReductionNeutralValue(ReductionNeutralAllOnes neutral,
516 ConversionPatternRewriter &rewriter,
517 Location loc, Type llvmType) {
518 return LLVM::ConstantOp::create(
519 rewriter, loc, llvmType,
520 rewriter.getIntegerAttr(
521 llvmType, llvm::APInt::getAllOnes(llvmType.getIntOrFloatBitWidth())));
522}
523
524/// Create the reduction neutral signed int minimum value.
525static Value createReductionNeutralValue(ReductionNeutralSIntMin neutral,
526 ConversionPatternRewriter &rewriter,
527 Location loc, Type llvmType) {
528 return LLVM::ConstantOp::create(
529 rewriter, loc, llvmType,
530 rewriter.getIntegerAttr(llvmType, llvm::APInt::getSignedMinValue(
531 llvmType.getIntOrFloatBitWidth())));
532}
533
534/// Create the reduction neutral unsigned int minimum value.
535static Value createReductionNeutralValue(ReductionNeutralUIntMin neutral,
536 ConversionPatternRewriter &rewriter,
537 Location loc, Type llvmType) {
538 return LLVM::ConstantOp::create(
539 rewriter, loc, llvmType,
540 rewriter.getIntegerAttr(llvmType, llvm::APInt::getMinValue(
541 llvmType.getIntOrFloatBitWidth())));
542}
543
544/// Create the reduction neutral signed int maximum value.
545static Value createReductionNeutralValue(ReductionNeutralSIntMax neutral,
546 ConversionPatternRewriter &rewriter,
547 Location loc, Type llvmType) {
548 return LLVM::ConstantOp::create(
549 rewriter, loc, llvmType,
550 rewriter.getIntegerAttr(llvmType, llvm::APInt::getSignedMaxValue(
551 llvmType.getIntOrFloatBitWidth())));
552}
553
554/// Create the reduction neutral unsigned int maximum value.
555static Value createReductionNeutralValue(ReductionNeutralUIntMax neutral,
556 ConversionPatternRewriter &rewriter,
557 Location loc, Type llvmType) {
558 return LLVM::ConstantOp::create(
559 rewriter, loc, llvmType,
560 rewriter.getIntegerAttr(llvmType, llvm::APInt::getMaxValue(
561 llvmType.getIntOrFloatBitWidth())));
562}
563
564/// Create the reduction neutral quiet NaN value.
565static Value createReductionNeutralValue(ReductionNeutralFPQNaN neutral,
566 ConversionPatternRewriter &rewriter,
567 Location loc, Type llvmType) {
568 auto floatType = cast<FloatType>(llvmType);
569 return LLVM::ConstantOp::create(
570 rewriter, loc, llvmType,
571 rewriter.getFloatAttr(
572 llvmType, llvm::APFloat::getQNaN(floatType.getFloatSemantics(),
573 /*Negative=*/false)));
574}
575
576/// Create the reduction neutral negative quiet NaN value.
577static Value createReductionNeutralValue(ReductionNeutralFPNegQNaN neutral,
578 ConversionPatternRewriter &rewriter,
579 Location loc, Type llvmType) {
580 auto floatType = cast<FloatType>(llvmType);
581 return LLVM::ConstantOp::create(
582 rewriter, loc, llvmType,
583 rewriter.getFloatAttr(
584 llvmType, llvm::APFloat::getQNaN(floatType.getFloatSemantics(),
585 /*Negative=*/true)));
586}
587
588/// Create the reduction neutral negative infinity value.
589static Value createReductionNeutralValue(ReductionNeutralFPNegInf neutral,
590 ConversionPatternRewriter &rewriter,
591 Location loc, Type llvmType) {
592 auto floatType = cast<FloatType>(llvmType);
593 return LLVM::ConstantOp::create(
594 rewriter, loc, llvmType,
595 rewriter.getFloatAttr(llvmType,
596 llvm::APFloat::getInf(floatType.getFloatSemantics(),
597 /*Negative=*/true)));
598}
599
600/// Create the reduction neutral positive infinity value.
601static Value createReductionNeutralValue(ReductionNeutralFPPosInf neutral,
602 ConversionPatternRewriter &rewriter,
603 Location loc, Type llvmType) {
604 auto floatType = cast<FloatType>(llvmType);
605 return LLVM::ConstantOp::create(
606 rewriter, loc, llvmType,
607 rewriter.getFloatAttr(llvmType,
608 llvm::APFloat::getInf(floatType.getFloatSemantics(),
609 /*Negative=*/false)));
610}
611
612/// Create the reduction neutral lowest finite value.
613static Value createReductionNeutralValue(ReductionNeutralFPLowestFinite neutral,
614 ConversionPatternRewriter &rewriter,
615 Location loc, Type llvmType) {
616 auto floatType = cast<FloatType>(llvmType);
617 return LLVM::ConstantOp::create(
618 rewriter, loc, llvmType,
619 rewriter.getFloatAttr(
620 llvmType, llvm::APFloat::getLargest(floatType.getFloatSemantics(),
621 /*Negative=*/true)));
622}
623
624/// Create the reduction neutral largest finite value.
625static Value
626createReductionNeutralValue(ReductionNeutralFPLargestFinite neutral,
627 ConversionPatternRewriter &rewriter, Location loc,
628 Type llvmType) {
629 auto floatType = cast<FloatType>(llvmType);
630 return LLVM::ConstantOp::create(
631 rewriter, loc, llvmType,
632 rewriter.getFloatAttr(
633 llvmType, llvm::APFloat::getLargest(floatType.getFloatSemantics(),
634 /*Negative=*/false)));
635}
636
637/// Returns `accumulator` if it has a valid value. Otherwise, creates and
638/// returns a new accumulator value using `ReductionNeutral`.
639template <class ReductionNeutral>
640static Value getOrCreateAccumulator(ConversionPatternRewriter &rewriter,
641 Location loc, Type llvmType,
642 Value accumulator) {
643 if (accumulator)
644 return accumulator;
645
646 return createReductionNeutralValue(ReductionNeutral(), rewriter, loc,
647 llvmType);
648}
649
650/// Creates a value with the 1-D vector shape provided in `llvmType`.
651/// This is used as effective vector length by some intrinsics supporting
652/// dynamic vector lengths at runtime.
653static Value createVectorLengthValue(ConversionPatternRewriter &rewriter,
654 Location loc, Type llvmType) {
655 VectorType vType = cast<VectorType>(llvmType);
656 auto vShape = vType.getShape();
657 assert(vShape.size() == 1 && "Unexpected multi-dim vector type");
658
659 Value baseVecLength = LLVM::ConstantOp::create(
660 rewriter, loc, rewriter.getI32Type(),
661 rewriter.getIntegerAttr(rewriter.getI32Type(), vShape[0]));
662
663 if (!vType.getScalableDims()[0])
664 return baseVecLength;
665
666 // For a scalable vector type, create and return `vScale * baseVecLength`.
667 Value vScale = vector::VectorScaleOp::create(rewriter, loc);
668 vScale =
669 arith::IndexCastOp::create(rewriter, loc, rewriter.getI32Type(), vScale);
670 Value scalableVecLength =
671 arith::MulIOp::create(rewriter, loc, baseVecLength, vScale);
672 return scalableVecLength;
673}
674
675/// Helper method to lower a `vector.reduction` op that performs an arithmetic
676/// operation like add,mul, etc.. `VectorOp` is the LLVM vector intrinsic to use
677/// and `ScalarOp` is the scalar operation used to add the accumulation value if
678/// non-null.
679template <class LLVMRedIntrinOp, class ScalarOp>
680static Value createIntegerReductionArithmeticOpLowering(
681 ConversionPatternRewriter &rewriter, Location loc, Type llvmType,
682 Value vectorOperand, Value accumulator) {
683
684 Value result =
685 LLVMRedIntrinOp::create(rewriter, loc, llvmType, vectorOperand);
686
687 if (accumulator)
688 result = ScalarOp::create(rewriter, loc, accumulator, result);
689 return result;
690}
691
692/// Helper method to lower a `vector.reduction` operation that performs
693/// a comparison operation like `min`/`max`. `VectorOp` is the LLVM vector
694/// intrinsic to use and `predicate` is the predicate to use to compare+combine
695/// the accumulator value if non-null.
696template <class LLVMRedIntrinOp>
697static Value createIntegerReductionComparisonOpLowering(
698 ConversionPatternRewriter &rewriter, Location loc, Type llvmType,
699 Value vectorOperand, Value accumulator, LLVM::ICmpPredicate predicate) {
700 Value result =
701 LLVMRedIntrinOp::create(rewriter, loc, llvmType, vectorOperand);
702 if (accumulator) {
703 Value cmp =
704 LLVM::ICmpOp::create(rewriter, loc, predicate, accumulator, result);
705 result = LLVM::SelectOp::create(rewriter, loc, cmp, accumulator, result);
706 }
707 return result;
708}
709
710namespace {
711template <typename Source>
712struct VectorToScalarMapper;
713template <>
714struct VectorToScalarMapper<LLVM::vector_reduce_fmaximum> {
715 using Type = LLVM::MaximumOp;
716};
717template <>
718struct VectorToScalarMapper<LLVM::vector_reduce_fminimum> {
719 using Type = LLVM::MinimumOp;
720};
721template <>
722struct VectorToScalarMapper<LLVM::vector_reduce_fmax> {
723 using Type = LLVM::MaxNumOp;
724};
725template <>
726struct VectorToScalarMapper<LLVM::vector_reduce_fmin> {
727 using Type = LLVM::MinNumOp;
728};
729template <>
730struct VectorToScalarMapper<LLVM::vector_reduce_fmaximumnum> {
731 using Type = LLVM::MaximumNumOp;
732};
733template <>
734struct VectorToScalarMapper<LLVM::vector_reduce_fminimumnum> {
735 using Type = LLVM::MinimumNumOp;
736};
737} // namespace
738
739template <class LLVMRedIntrinOp>
740static Value createFPReductionComparisonOpLowering(
741 ConversionPatternRewriter &rewriter, Location loc, Type llvmType,
742 Value vectorOperand, Value accumulator, LLVM::FastmathFlagsAttr fmf) {
743 Value result =
744 LLVMRedIntrinOp::create(rewriter, loc, llvmType, vectorOperand, fmf);
745
746 if (accumulator) {
747 result = VectorToScalarMapper<LLVMRedIntrinOp>::Type::create(
748 rewriter, loc, result, accumulator);
749 }
750
751 return result;
752}
753
754/// Mask neutral classes for overloading.
755class MaskNeutralFMaximumNum {};
756class MaskNeutralFMinimumNum {};
757
758/// Get the mask neutral value for a `fmaximumnum` reduction. `maximumnum`
759/// ignores NaN operands, making a quiet NaN the identity element. When `nnan`
760/// promises no NaN reaches the reduction, fall back to the smallest
761/// representable value, which `ninf` further constrains to be finite.
762static llvm::APFloat getMaskNeutralValue(MaskNeutralFMaximumNum,
763 const llvm::fltSemantics &semantics,
764 bool noNaNs, bool noInfs) {
765 if (!noNaNs)
766 return llvm::APFloat::getQNaN(semantics);
767 if (noInfs)
768 return llvm::APFloat::getLargest(semantics, /*Negative=*/true);
769 return llvm::APFloat::getInf(semantics, /*Negative=*/true);
770}
771
772/// Get the mask neutral value for a `fminimumnum` reduction. See the
773/// `fmaximumnum` overload above for the rationale.
774static llvm::APFloat getMaskNeutralValue(MaskNeutralFMinimumNum,
775 const llvm::fltSemantics &semantics,
776 bool noNaNs, bool noInfs) {
777 if (!noNaNs)
778 return llvm::APFloat::getQNaN(semantics);
779 if (noInfs)
780 return llvm::APFloat::getLargest(semantics);
781 return llvm::APFloat::getInf(semantics);
782}
783
784/// Lowers masked `fmaximumnum` and `fminimumnum` reductions using the
785/// non-masked intrinsics, since LLVM has no predicated counterpart for them.
786/// Inactive lanes are replaced by a mask neutral value before reducing.
787/// TODO: Switch to `lowerPredicatedReductionWithStartValue` once
788/// `llvm.vp.reduce.fmaximumnum`/`fminimumnum` are added to LLVM IR.
789template <class LLVMRedIntrinOp, class MaskNeutral>
790static Value
791lowerMaskedReductionWithRegular(ConversionPatternRewriter &rewriter,
792 Location loc, Type llvmType,
793 Value vectorOperand, Value accumulator,
794 Value mask, LLVM::FastmathFlagsAttr fmf) {
795 const auto &floatSemantics = cast<FloatType>(llvmType).getFloatSemantics();
796 auto value = getMaskNeutralValue(
797 MaskNeutral{}, floatSemantics,
798 LLVM::bitEnumContainsAny(fmf.getValue(), LLVM::FastmathFlags::nnan),
799 LLVM::bitEnumContainsAny(fmf.getValue(), LLVM::FastmathFlags::ninf));
800 Type vectorType = vectorOperand.getType();
801 auto denseValue = DenseElementsAttr::get(cast<ShapedType>(vectorType), value);
802 const Value vectorMaskNeutral =
803 LLVM::ConstantOp::create(rewriter, loc, vectorType, denseValue);
804 const Value selectedVectorByMask = LLVM::SelectOp::create(
805 rewriter, loc, mask, vectorOperand, vectorMaskNeutral);
806 return createFPReductionComparisonOpLowering<LLVMRedIntrinOp>(
807 rewriter, loc, llvmType, selectedVectorByMask, accumulator, fmf);
808}
809
810template <class LLVMRedIntrinOp, class ReductionNeutral>
811static Value
812lowerReductionWithStartValue(ConversionPatternRewriter &rewriter, Location loc,
813 Type llvmType, Value vectorOperand,
814 Value accumulator, LLVM::FastmathFlagsAttr fmf) {
815 accumulator = getOrCreateAccumulator<ReductionNeutral>(rewriter, loc,
816 llvmType, accumulator);
817 return LLVMRedIntrinOp::create(rewriter, loc, llvmType,
818 /*start_value=*/accumulator, vectorOperand,
819 fmf);
820}
821
822template <class LLVMVPRedIntrinOp, class ReductionNeutral>
823static Value lowerPredicatedReductionWithStartValue(
824 ConversionPatternRewriter &rewriter, Location loc, Type llvmType,
825 Value vectorOperand, Value accumulator, Value mask) {
826 accumulator = getOrCreateAccumulator<ReductionNeutral>(rewriter, loc,
827 llvmType, accumulator);
828 Value vectorLength =
829 createVectorLengthValue(rewriter, loc, vectorOperand.getType());
830 return LLVMVPRedIntrinOp::create(rewriter, loc, llvmType,
831 /*satrt_value=*/accumulator, vectorOperand,
832 mask, vectorLength);
833}
834
835template <class LLVMIntVPRedIntrinOp, class IntReductionNeutral,
836 class LLVMFPVPRedIntrinOp, class FPReductionNeutral>
837static Value lowerPredicatedReductionWithStartValue(
838 ConversionPatternRewriter &rewriter, Location loc, Type llvmType,
839 Value vectorOperand, Value accumulator, Value mask) {
840 if (llvmType.isIntOrIndex())
841 return lowerPredicatedReductionWithStartValue<LLVMIntVPRedIntrinOp,
842 IntReductionNeutral>(
843 rewriter, loc, llvmType, vectorOperand, accumulator, mask);
844
845 // FP dispatch.
846 return lowerPredicatedReductionWithStartValue<LLVMFPVPRedIntrinOp,
847 FPReductionNeutral>(
848 rewriter, loc, llvmType, vectorOperand, accumulator, mask);
849}
850
851/// Conversion pattern for all vector reductions.
852class VectorReductionOpConversion
853 : public ConvertOpToLLVMPattern<vector::ReductionOp> {
854public:
855 explicit VectorReductionOpConversion(const LLVMTypeConverter &typeConv,
856 bool reassociateFPRed)
857 : ConvertOpToLLVMPattern<vector::ReductionOp>(typeConv),
858 reassociateFPReductions(reassociateFPRed) {}
859
860 LogicalResult
861 matchAndRewrite(vector::ReductionOp reductionOp, OpAdaptor adaptor,
862 ConversionPatternRewriter &rewriter) const override {
863 auto kind = reductionOp.getKind();
864 Type eltType = reductionOp.getDest().getType();
865 Type llvmType = typeConverter->convertType(eltType);
866 Value operand = adaptor.getVector();
867 Value acc = adaptor.getAcc();
868 Location loc = reductionOp.getLoc();
869
870 if (eltType.isIntOrIndex()) {
871 // Integer reductions: add/mul/min/max/and/or/xor.
872 Value result;
873 switch (kind) {
874 case vector::CombiningKind::ADD:
875 result =
876 createIntegerReductionArithmeticOpLowering<LLVM::vector_reduce_add,
877 LLVM::AddOp>(
878 rewriter, loc, llvmType, operand, acc);
879 break;
880 case vector::CombiningKind::MUL:
881 result =
882 createIntegerReductionArithmeticOpLowering<LLVM::vector_reduce_mul,
883 LLVM::MulOp>(
884 rewriter, loc, llvmType, operand, acc);
885 break;
886 case vector::CombiningKind::MINUI:
887 result = createIntegerReductionComparisonOpLowering<
888 LLVM::vector_reduce_umin>(rewriter, loc, llvmType, operand, acc,
889 LLVM::ICmpPredicate::ule);
890 break;
891 case vector::CombiningKind::MINSI:
892 result = createIntegerReductionComparisonOpLowering<
893 LLVM::vector_reduce_smin>(rewriter, loc, llvmType, operand, acc,
894 LLVM::ICmpPredicate::sle);
895 break;
896 case vector::CombiningKind::MAXUI:
897 result = createIntegerReductionComparisonOpLowering<
898 LLVM::vector_reduce_umax>(rewriter, loc, llvmType, operand, acc,
899 LLVM::ICmpPredicate::uge);
900 break;
901 case vector::CombiningKind::MAXSI:
902 result = createIntegerReductionComparisonOpLowering<
903 LLVM::vector_reduce_smax>(rewriter, loc, llvmType, operand, acc,
904 LLVM::ICmpPredicate::sge);
905 break;
906 case vector::CombiningKind::AND:
907 result =
908 createIntegerReductionArithmeticOpLowering<LLVM::vector_reduce_and,
909 LLVM::AndOp>(
910 rewriter, loc, llvmType, operand, acc);
911 break;
912 case vector::CombiningKind::OR:
913 result =
914 createIntegerReductionArithmeticOpLowering<LLVM::vector_reduce_or,
915 LLVM::OrOp>(
916 rewriter, loc, llvmType, operand, acc);
917 break;
918 case vector::CombiningKind::XOR:
919 result =
920 createIntegerReductionArithmeticOpLowering<LLVM::vector_reduce_xor,
921 LLVM::XOrOp>(
922 rewriter, loc, llvmType, operand, acc);
923 break;
924 default:
925 return failure();
926 }
927 rewriter.replaceOp(reductionOp, result);
928
929 return success();
930 }
931
932 if (!isa<FloatType>(eltType))
933 return failure();
934
935 arith::FastMathFlagsAttr fMFAttr = reductionOp.getFastMathFlagsAttr();
936 LLVM::FastmathFlagsAttr fmf = LLVM::FastmathFlagsAttr::get(
937 reductionOp.getContext(),
938 convertArithFastMathFlagsToLLVM(fMFAttr.getValue()));
939 fmf = LLVM::FastmathFlagsAttr::get(
940 reductionOp.getContext(),
941 fmf.getValue() | (reassociateFPReductions ? LLVM::FastmathFlags::reassoc
942 : LLVM::FastmathFlags::none));
943
944 // Floating-point reductions: add/mul/min/max
945 Value result;
946 if (kind == vector::CombiningKind::ADD) {
947 result = lowerReductionWithStartValue<LLVM::vector_reduce_fadd,
948 ReductionNeutralZero>(
949 rewriter, loc, llvmType, operand, acc, fmf);
950 } else if (kind == vector::CombiningKind::MUL) {
951 result = lowerReductionWithStartValue<LLVM::vector_reduce_fmul,
952 ReductionNeutralFPOne>(
953 rewriter, loc, llvmType, operand, acc, fmf);
954 } else if (kind == vector::CombiningKind::MINIMUMF) {
955 result =
956 createFPReductionComparisonOpLowering<LLVM::vector_reduce_fminimum>(
957 rewriter, loc, llvmType, operand, acc, fmf);
958 } else if (kind == vector::CombiningKind::MAXIMUMF) {
959 result =
960 createFPReductionComparisonOpLowering<LLVM::vector_reduce_fmaximum>(
961 rewriter, loc, llvmType, operand, acc, fmf);
962 } else if (kind == vector::CombiningKind::MINNUMF) {
963 result = createFPReductionComparisonOpLowering<LLVM::vector_reduce_fmin>(
964 rewriter, loc, llvmType, operand, acc, fmf);
965 } else if (kind == vector::CombiningKind::MAXNUMF) {
966 result = createFPReductionComparisonOpLowering<LLVM::vector_reduce_fmax>(
967 rewriter, loc, llvmType, operand, acc, fmf);
968 } else if (kind == vector::CombiningKind::MAXIMUMNUMF) {
969 result = createFPReductionComparisonOpLowering<
970 LLVM::vector_reduce_fmaximumnum>(rewriter, loc, llvmType, operand,
971 acc, fmf);
972 } else if (kind == vector::CombiningKind::MINIMUMNUMF) {
973 result = createFPReductionComparisonOpLowering<
974 LLVM::vector_reduce_fminimumnum>(rewriter, loc, llvmType, operand,
975 acc, fmf);
976 } else {
977 return failure();
978 }
979
980 rewriter.replaceOp(reductionOp, result);
981 return success();
982 }
983
984private:
985 const bool reassociateFPReductions;
986};
987
988/// Base class to convert a `vector.mask` operation while matching traits
989/// of the maskable operation nested inside. A `VectorMaskOpConversionBase`
990/// instance matches against a `vector.mask` operation. The `matchAndRewrite`
991/// method performs a second match against the maskable operation `MaskedOp`.
992/// Finally, it invokes the virtual method `matchAndRewriteMaskableOp` to be
993/// implemented by the concrete conversion classes. This method can match
994/// against specific traits of the `vector.mask` and the maskable operation. It
995/// must replace the `vector.mask` operation.
996template <class MaskedOp>
997class VectorMaskOpConversionBase
998 : public ConvertOpToLLVMPattern<vector::MaskOp> {
999public:
1000 using ConvertOpToLLVMPattern<vector::MaskOp>::ConvertOpToLLVMPattern;
1001
1002 LogicalResult
1003 matchAndRewrite(vector::MaskOp maskOp, OpAdaptor adaptor,
1004 ConversionPatternRewriter &rewriter) const final {
1005 // Match against the maskable operation kind.
1006 auto maskedOp = llvm::dyn_cast_or_null<MaskedOp>(maskOp.getMaskableOp());
1007 if (!maskedOp)
1008 return failure();
1009 return matchAndRewriteMaskableOp(maskOp, maskedOp, rewriter);
1010 }
1011
1012protected:
1013 virtual LogicalResult
1014 matchAndRewriteMaskableOp(vector::MaskOp maskOp,
1015 vector::MaskableOpInterface maskableOp,
1016 ConversionPatternRewriter &rewriter) const = 0;
1017};
1018
1019class MaskedReductionOpConversion
1020 : public VectorMaskOpConversionBase<vector::ReductionOp> {
1021
1022public:
1023 using VectorMaskOpConversionBase<
1024 vector::ReductionOp>::VectorMaskOpConversionBase;
1025
1026 LogicalResult matchAndRewriteMaskableOp(
1027 vector::MaskOp maskOp, MaskableOpInterface maskableOp,
1028 ConversionPatternRewriter &rewriter) const override {
1029 auto reductionOp = cast<ReductionOp>(maskableOp.getOperation());
1030 auto kind = reductionOp.getKind();
1031 Type eltType = reductionOp.getDest().getType();
1032 Type llvmType = typeConverter->convertType(eltType);
1033 Value operand = reductionOp.getVector();
1034 Value acc = reductionOp.getAcc();
1035 Location loc = reductionOp.getLoc();
1036
1037 arith::FastMathFlagsAttr fMFAttr = reductionOp.getFastMathFlagsAttr();
1038 LLVM::FastmathFlagsAttr fmf = LLVM::FastmathFlagsAttr::get(
1039 reductionOp.getContext(),
1040 convertArithFastMathFlagsToLLVM(fMFAttr.getValue()));
1041 const bool noInfs =
1042 LLVM::bitEnumContainsAny(fmf.getValue(), LLVM::FastmathFlags::ninf);
1043
1044 Value result;
1045 switch (kind) {
1046 case vector::CombiningKind::ADD:
1047 result = lowerPredicatedReductionWithStartValue<
1048 LLVM::VPReduceAddOp, ReductionNeutralZero, LLVM::VPReduceFAddOp,
1049 ReductionNeutralZero>(rewriter, loc, llvmType, operand, acc,
1050 maskOp.getMask());
1051 break;
1052 case vector::CombiningKind::MUL:
1053 result = lowerPredicatedReductionWithStartValue<
1054 LLVM::VPReduceMulOp, ReductionNeutralIntOne, LLVM::VPReduceFMulOp,
1055 ReductionNeutralFPOne>(rewriter, loc, llvmType, operand, acc,
1056 maskOp.getMask());
1057 break;
1058 case vector::CombiningKind::MINUI:
1059 result = lowerPredicatedReductionWithStartValue<LLVM::VPReduceUMinOp,
1060 ReductionNeutralUIntMax>(
1061 rewriter, loc, llvmType, operand, acc, maskOp.getMask());
1062 break;
1063 case vector::CombiningKind::MINSI:
1064 result = lowerPredicatedReductionWithStartValue<LLVM::VPReduceSMinOp,
1065 ReductionNeutralSIntMax>(
1066 rewriter, loc, llvmType, operand, acc, maskOp.getMask());
1067 break;
1068 case vector::CombiningKind::MAXUI:
1069 result = lowerPredicatedReductionWithStartValue<LLVM::VPReduceUMaxOp,
1070 ReductionNeutralUIntMin>(
1071 rewriter, loc, llvmType, operand, acc, maskOp.getMask());
1072 break;
1073 case vector::CombiningKind::MAXSI:
1074 result = lowerPredicatedReductionWithStartValue<LLVM::VPReduceSMaxOp,
1075 ReductionNeutralSIntMin>(
1076 rewriter, loc, llvmType, operand, acc, maskOp.getMask());
1077 break;
1078 case vector::CombiningKind::AND:
1079 result = lowerPredicatedReductionWithStartValue<LLVM::VPReduceAndOp,
1080 ReductionNeutralAllOnes>(
1081 rewriter, loc, llvmType, operand, acc, maskOp.getMask());
1082 break;
1083 case vector::CombiningKind::OR:
1084 result = lowerPredicatedReductionWithStartValue<LLVM::VPReduceOrOp,
1085 ReductionNeutralZero>(
1086 rewriter, loc, llvmType, operand, acc, maskOp.getMask());
1087 break;
1088 case vector::CombiningKind::XOR:
1089 result = lowerPredicatedReductionWithStartValue<LLVM::VPReduceXorOp,
1090 ReductionNeutralZero>(
1091 rewriter, loc, llvmType, operand, acc, maskOp.getMask());
1092 break;
1093 case vector::CombiningKind::MINNUMF:
1094 result =
1095 lowerPredicatedReductionWithStartValue<LLVM::VPReduceFMinOp,
1096 ReductionNeutralFPNegQNaN>(
1097 rewriter, loc, llvmType, operand, acc, maskOp.getMask());
1098 break;
1099 case vector::CombiningKind::MAXNUMF:
1100 result = lowerPredicatedReductionWithStartValue<LLVM::VPReduceFMaxOp,
1101 ReductionNeutralFPQNaN>(
1102 rewriter, loc, llvmType, operand, acc, maskOp.getMask());
1103 break;
1104 case CombiningKind::MAXIMUMF:
1105 // `ninf` promises no infinity reaches the reduction, so the neutral start
1106 // value must stay finite.
1107 result =
1108 noInfs
1109 ? lowerPredicatedReductionWithStartValue<
1110 LLVM::VPReduceFMaximumOp, ReductionNeutralFPLowestFinite>(
1111 rewriter, loc, llvmType, operand, acc, maskOp.getMask())
1112 : lowerPredicatedReductionWithStartValue<
1113 LLVM::VPReduceFMaximumOp, ReductionNeutralFPNegInf>(
1114 rewriter, loc, llvmType, operand, acc, maskOp.getMask());
1115 break;
1116 case CombiningKind::MINIMUMF:
1117 result =
1118 noInfs
1119 ? lowerPredicatedReductionWithStartValue<
1120 LLVM::VPReduceFMinimumOp, ReductionNeutralFPLargestFinite>(
1121 rewriter, loc, llvmType, operand, acc, maskOp.getMask())
1122 : lowerPredicatedReductionWithStartValue<
1123 LLVM::VPReduceFMinimumOp, ReductionNeutralFPPosInf>(
1124 rewriter, loc, llvmType, operand, acc, maskOp.getMask());
1125 break;
1126 case CombiningKind::MAXIMUMNUMF:
1127 result = lowerMaskedReductionWithRegular<LLVM::vector_reduce_fmaximumnum,
1128 MaskNeutralFMaximumNum>(
1129 rewriter, loc, llvmType, operand, acc, maskOp.getMask(), fmf);
1130 break;
1131 case CombiningKind::MINIMUMNUMF:
1132 result = lowerMaskedReductionWithRegular<LLVM::vector_reduce_fminimumnum,
1133 MaskNeutralFMinimumNum>(
1134 rewriter, loc, llvmType, operand, acc, maskOp.getMask(), fmf);
1135 break;
1136 }
1137
1138 // Replace `vector.mask` operation altogether.
1139 rewriter.replaceOp(maskOp, result);
1140 return success();
1141 }
1142};
1143
1144class VectorShuffleOpConversion
1145 : public ConvertOpToLLVMPattern<vector::ShuffleOp> {
1146public:
1147 using ConvertOpToLLVMPattern<vector::ShuffleOp>::ConvertOpToLLVMPattern;
1148
1149 LogicalResult
1150 matchAndRewrite(vector::ShuffleOp shuffleOp, OpAdaptor adaptor,
1151 ConversionPatternRewriter &rewriter) const override {
1152 auto loc = shuffleOp->getLoc();
1153 auto v1Type = shuffleOp.getV1VectorType();
1154 auto v2Type = shuffleOp.getV2VectorType();
1155 auto vectorType = shuffleOp.getResultVectorType();
1156 Type llvmType = typeConverter->convertType(vectorType);
1157 ArrayRef<int64_t> mask = shuffleOp.getMask();
1158
1159 // Bail if result type cannot be lowered.
1160 if (!llvmType)
1161 return failure();
1162
1163 // Get rank and dimension sizes.
1164 int64_t rank = vectorType.getRank();
1165#ifndef NDEBUG
1166 bool wellFormed0DCase =
1167 v1Type.getRank() == 0 && v2Type.getRank() == 0 && rank == 1;
1168 bool wellFormedNDCase =
1169 v1Type.getRank() == rank && v2Type.getRank() == rank;
1170 assert((wellFormed0DCase || wellFormedNDCase) && "op is not well-formed");
1171#endif
1172
1173 // For rank 0 and 1, where both operands have *exactly* the same vector
1174 // type, there is direct shuffle support in LLVM. Use it!
1175 if (rank <= 1 && v1Type == v2Type) {
1176 Value llvmShuffleOp = LLVM::ShuffleVectorOp::create(
1177 rewriter, loc, adaptor.getV1(), adaptor.getV2(),
1178 llvm::to_vector_of<int32_t>(mask));
1179 rewriter.replaceOp(shuffleOp, llvmShuffleOp);
1180 return success();
1181 }
1182
1183 // For all other cases, insert the individual values individually.
1184 int64_t v1Dim = v1Type.getDimSize(0);
1185 Type eltType;
1186 if (auto arrayType = dyn_cast<LLVM::LLVMArrayType>(llvmType))
1187 eltType = arrayType.getElementType();
1188 else
1189 eltType = cast<VectorType>(llvmType).getElementType();
1190 Value insert = LLVM::PoisonOp::create(rewriter, loc, llvmType);
1191 int64_t insPos = 0;
1192 for (int64_t extPos : mask) {
1193 Value value = adaptor.getV1();
1194 if (extPos >= v1Dim) {
1195 extPos -= v1Dim;
1196 value = adaptor.getV2();
1197 }
1198 Value extract = extractOne(rewriter, *getTypeConverter(), loc, value,
1199 eltType, rank, extPos);
1200 insert = insertOne(rewriter, *getTypeConverter(), loc, insert, extract,
1201 llvmType, rank, insPos++);
1202 }
1203 rewriter.replaceOp(shuffleOp, insert);
1204 return success();
1205 }
1206};
1207
1208class VectorExtractOpConversion
1209 : public ConvertOpToLLVMPattern<vector::ExtractOp> {
1210public:
1211 using ConvertOpToLLVMPattern<vector::ExtractOp>::ConvertOpToLLVMPattern;
1212
1213 LogicalResult
1214 matchAndRewrite(vector::ExtractOp extractOp, OpAdaptor adaptor,
1215 ConversionPatternRewriter &rewriter) const override {
1216 auto loc = extractOp->getLoc();
1217 auto resultType = extractOp.getResult().getType();
1218 auto llvmResultType = typeConverter->convertType(resultType);
1219 // Bail if result type cannot be lowered.
1220 if (!llvmResultType)
1221 return failure();
1222
1223 SmallVector<OpFoldResult> positionVec = getMixedValues(
1224 adaptor.getStaticPosition(), adaptor.getDynamicPosition(), rewriter);
1225
1226 // The Vector -> LLVM lowering models N-D vectors as nested aggregates of
1227 // 1-d vectors. This nesting is modeled using arrays. We do this conversion
1228 // from a N-d vector extract to a nested aggregate vector extract in two
1229 // steps:
1230 // - Extract a member from the nested aggregate. The result can be
1231 // a lower rank nested aggregate or a vector (1-D). This is done using
1232 // `llvm.extractvalue`.
1233 // - Extract a scalar out of the vector if needed. This is done using
1234 // `llvm.extractelement`.
1235
1236 // Determine if we need to extract a member out of the aggregate. We
1237 // always need to extract a member if the input rank >= 2.
1238 bool extractsAggregate = extractOp.getSourceVectorType().getRank() >= 2;
1239 // Determine if we need to extract a scalar as the result. We extract
1240 // a scalar if the extract is full rank, i.e., the number of indices is
1241 // equal to source vector rank.
1242 bool extractsScalar = static_cast<int64_t>(positionVec.size()) ==
1243 extractOp.getSourceVectorType().getRank();
1244
1245 // Since the LLVM type converter converts 0-d vectors to 1-d vectors, we
1246 // need to add a position for this change.
1247 if (extractOp.getSourceVectorType().getRank() == 0) {
1248 Type idxType = typeConverter->convertType(rewriter.getIndexType());
1249 positionVec.push_back(rewriter.getZeroAttr(idxType));
1250 }
1251
1252 Value extracted = adaptor.getSource();
1253 if (extractsAggregate) {
1254 ArrayRef<OpFoldResult> position(positionVec);
1255 if (extractsScalar) {
1256 // If we are extracting a scalar from the extracted member, we drop
1257 // the last index, which will be used to extract the scalar out of the
1258 // vector.
1259 position = position.drop_back();
1260 }
1261 // llvm.extractvalue does not support dynamic dimensions.
1262 if (!llvm::all_of(position, llvm::IsaPred<Attribute>)) {
1263 return failure();
1264 }
1265 extracted = LLVM::ExtractValueOp::create(rewriter, loc, extracted,
1266 getAsIntegers(position));
1267 }
1268
1269 if (extractsScalar) {
1270 extracted = LLVM::ExtractElementOp::create(
1271 rewriter, loc, extracted,
1272 getAsLLVMValue(rewriter, loc, positionVec.back()));
1273 }
1274
1275 rewriter.replaceOp(extractOp, extracted);
1276 return success();
1277 }
1278};
1279
1280/// Conversion pattern that turns a vector.fma on a 1-D vector
1281/// into an llvm.intr.fmuladd. This is a trivial 1-1 conversion.
1282/// This does not match vectors of n >= 2 rank.
1283///
1284/// Example:
1285/// ```
1286/// vector.fma %a, %a, %a : vector<8xf32>
1287/// ```
1288/// is converted to:
1289/// ```
1290/// llvm.intr.fmuladd %va, %va, %va:
1291/// (!llvm."<8 x f32>">, !llvm<"<8 x f32>">, !llvm<"<8 x f32>">)
1292/// -> !llvm."<8 x f32>">
1293/// ```
1294class VectorFMAOp1DConversion : public ConvertOpToLLVMPattern<vector::FMAOp> {
1295public:
1296 using ConvertOpToLLVMPattern<vector::FMAOp>::ConvertOpToLLVMPattern;
1297
1298 LogicalResult
1299 matchAndRewrite(vector::FMAOp fmaOp, OpAdaptor adaptor,
1300 ConversionPatternRewriter &rewriter) const override {
1301 VectorType vType = fmaOp.getVectorType();
1302 if (vType.getRank() > 1)
1303 return failure();
1304
1305 rewriter.replaceOpWithNewOp<LLVM::FMulAddOp>(
1306 fmaOp, adaptor.getLhs(), adaptor.getRhs(), adaptor.getAcc());
1307 return success();
1308 }
1309};
1310
1311class VectorInsertOpConversion
1312 : public ConvertOpToLLVMPattern<vector::InsertOp> {
1313public:
1314 using ConvertOpToLLVMPattern<vector::InsertOp>::ConvertOpToLLVMPattern;
1315
1316 LogicalResult
1317 matchAndRewrite(vector::InsertOp insertOp, OpAdaptor adaptor,
1318 ConversionPatternRewriter &rewriter) const override {
1319 auto loc = insertOp->getLoc();
1320 auto destVectorType = insertOp.getDestVectorType();
1321 auto llvmResultType = typeConverter->convertType(destVectorType);
1322 // Bail if result type cannot be lowered.
1323 if (!llvmResultType)
1324 return failure();
1325
1326 SmallVector<OpFoldResult> positionVec = getMixedValues(
1327 adaptor.getStaticPosition(), adaptor.getDynamicPosition(), rewriter);
1328
1329 // The logic in this pattern mirrors VectorExtractOpConversion. Refer to
1330 // its explanatory comment about how N-D vectors are converted as nested
1331 // aggregates (llvm.array's) of 1D vectors.
1332 //
1333 // The innermost dimension of the destination vector, when converted to a
1334 // nested aggregate form, will always be a 1D vector.
1335 //
1336 // * If the insertion is happening into the innermost dimension of the
1337 // destination vector:
1338 // - If the destination is a nested aggregate, extract a 1D vector out of
1339 // the aggregate. This can be done using llvm.extractvalue. The
1340 // destination is now guaranteed to be a 1D vector, to which we are
1341 // inserting.
1342 // - Do the insertion into the 1D destination vector, and make the result
1343 // the new source nested aggregate. This can be done using
1344 // llvm.insertelement.
1345 // * Insert the source nested aggregate into the destination nested
1346 // aggregate.
1347
1348 // Determine if we need to extract/insert a 1D vector out of the aggregate.
1349 bool isNestedAggregate = isa<LLVM::LLVMArrayType>(llvmResultType);
1350 // Determine if we need to insert a scalar into the 1D vector.
1351 bool insertIntoInnermostDim =
1352 static_cast<int64_t>(positionVec.size()) == destVectorType.getRank();
1353
1354 ArrayRef<OpFoldResult> positionOf1DVectorWithinAggregate(
1355 positionVec.begin(),
1356 insertIntoInnermostDim ? positionVec.size() - 1 : positionVec.size());
1357 OpFoldResult positionOfScalarWithin1DVector;
1358 if (destVectorType.getRank() == 0) {
1359 // Since the LLVM type converter converts 0D vectors to 1D vectors, we
1360 // need to create a 0 here as the position into the 1D vector.
1361 Type idxType = typeConverter->convertType(rewriter.getIndexType());
1362 positionOfScalarWithin1DVector = rewriter.getZeroAttr(idxType);
1363 } else if (insertIntoInnermostDim) {
1364 positionOfScalarWithin1DVector = positionVec.back();
1365 }
1366
1367 // We are going to mutate this 1D vector until it is either the final
1368 // result (in the non-aggregate case) or the value that needs to be
1369 // inserted into the aggregate result.
1370 Value sourceAggregate = adaptor.getValueToStore();
1371 if (insertIntoInnermostDim) {
1372 // Scalar-into-1D-vector case, so we know we will have to create a
1373 // InsertElementOp. The question is into what destination.
1374 if (isNestedAggregate) {
1375 // Aggregate case: the destination for the InsertElementOp needs to be
1376 // extracted from the aggregate.
1377 if (!llvm::all_of(positionOf1DVectorWithinAggregate,
1378 llvm::IsaPred<Attribute>)) {
1379 // llvm.extractvalue does not support dynamic dimensions.
1380 return failure();
1381 }
1382 sourceAggregate = LLVM::ExtractValueOp::create(
1383 rewriter, loc, adaptor.getDest(),
1384 getAsIntegers(positionOf1DVectorWithinAggregate));
1385 } else {
1386 // No-aggregate case. The destination for the InsertElementOp is just
1387 // the insertOp's destination.
1388 sourceAggregate = adaptor.getDest();
1389 }
1390 // Insert the scalar into the 1D vector.
1391 sourceAggregate = LLVM::InsertElementOp::create(
1392 rewriter, loc, sourceAggregate.getType(), sourceAggregate,
1393 adaptor.getValueToStore(),
1394 getAsLLVMValue(rewriter, loc, positionOfScalarWithin1DVector));
1395 }
1396
1397 Value result = sourceAggregate;
1398 if (isNestedAggregate) {
1399 if (!llvm::all_of(positionOf1DVectorWithinAggregate,
1400 llvm::IsaPred<Attribute>)) {
1401 // llvm.insertvalue does not support dynamic dimensions.
1402 return failure();
1403 }
1404 result = LLVM::InsertValueOp::create(
1405 rewriter, loc, adaptor.getDest(), sourceAggregate,
1406 getAsIntegers(positionOf1DVectorWithinAggregate));
1407 }
1408
1409 rewriter.replaceOp(insertOp, result);
1410 return success();
1411 }
1412};
1413
1414/// Lower vector.scalable.insert ops to LLVM vector.insert
1415struct VectorScalableInsertOpLowering
1416 : public ConvertOpToLLVMPattern<vector::ScalableInsertOp> {
1417 using ConvertOpToLLVMPattern<
1418 vector::ScalableInsertOp>::ConvertOpToLLVMPattern;
1419
1420 LogicalResult
1421 matchAndRewrite(vector::ScalableInsertOp insOp, OpAdaptor adaptor,
1422 ConversionPatternRewriter &rewriter) const override {
1423 rewriter.replaceOpWithNewOp<LLVM::vector_insert>(
1424 insOp, adaptor.getDest(), adaptor.getValueToStore(), adaptor.getPos());
1425 return success();
1426 }
1427};
1428
1429/// Lower vector.scalable.extract ops to LLVM vector.extract
1430struct VectorScalableExtractOpLowering
1431 : public ConvertOpToLLVMPattern<vector::ScalableExtractOp> {
1432 using ConvertOpToLLVMPattern<
1433 vector::ScalableExtractOp>::ConvertOpToLLVMPattern;
1434
1435 LogicalResult
1436 matchAndRewrite(vector::ScalableExtractOp extOp, OpAdaptor adaptor,
1437 ConversionPatternRewriter &rewriter) const override {
1438 rewriter.replaceOpWithNewOp<LLVM::vector_extract>(
1439 extOp, typeConverter->convertType(extOp.getResultVectorType()),
1440 adaptor.getSource(), adaptor.getPos());
1441 return success();
1442 }
1443};
1444
1445/// Rank reducing rewrite for n-D FMA into (n-1)-D FMA where n > 1.
1446///
1447/// Example:
1448/// ```
1449/// %d = vector.fma %a, %b, %c : vector<2x4xf32>
1450/// ```
1451/// is rewritten into:
1452/// ```
1453/// %r = vector.broadcast %f0 : f32 to vector<2x4xf32>
1454/// %va = vector.extractvalue %a[0] : vector<2x4xf32>
1455/// %vb = vector.extractvalue %b[0] : vector<2x4xf32>
1456/// %vc = vector.extractvalue %c[0] : vector<2x4xf32>
1457/// %vd = vector.fma %va, %vb, %vc : vector<4xf32>
1458/// %r2 = vector.insertvalue %vd, %r[0] : vector<4xf32> into vector<2x4xf32>
1459/// %va2 = vector.extractvalue %a2[1] : vector<2x4xf32>
1460/// %vb2 = vector.extractvalue %b2[1] : vector<2x4xf32>
1461/// %vc2 = vector.extractvalue %c2[1] : vector<2x4xf32>
1462/// %vd2 = vector.fma %va2, %vb2, %vc2 : vector<4xf32>
1463/// %r3 = vector.insertvalue %vd2, %r2[1] : vector<4xf32> into vector<2x4xf32>
1464/// // %r3 holds the final value.
1465/// ```
1466class VectorFMAOpNDRewritePattern : public OpRewritePattern<FMAOp> {
1467public:
1468 using Base::Base;
1469
1470 void initialize() {
1471 // This pattern recursively unpacks one dimension at a time. The recursion
1472 // bounded as the rank is strictly decreasing.
1473 setHasBoundedRewriteRecursion();
1474 }
1475
1476 LogicalResult matchAndRewrite(FMAOp op,
1477 PatternRewriter &rewriter) const override {
1478 auto vType = op.getVectorType();
1479 if (vType.getRank() < 2)
1480 return failure();
1481
1482 auto loc = op.getLoc();
1483 auto elemType = vType.getElementType();
1484 Value zero = arith::ConstantOp::create(rewriter, loc, elemType,
1485 rewriter.getZeroAttr(elemType));
1486 Value desc = vector::BroadcastOp::create(rewriter, loc, vType, zero);
1487 for (int64_t i = 0, e = vType.getShape().front(); i != e; ++i) {
1488 Value extrLHS = ExtractOp::create(rewriter, loc, op.getLhs(), i);
1489 Value extrRHS = ExtractOp::create(rewriter, loc, op.getRhs(), i);
1490 Value extrACC = ExtractOp::create(rewriter, loc, op.getAcc(), i);
1491 Value fma = FMAOp::create(rewriter, loc, extrLHS, extrRHS, extrACC);
1492 desc = InsertOp::create(rewriter, loc, fma, desc, i);
1493 }
1494 rewriter.replaceOp(op, desc);
1495 return success();
1496 }
1497};
1498
1499/// Returns the strides if the memory underlying `memRefType` has a contiguous
1500/// static layout.
1501static std::optional<SmallVector<int64_t, 4>>
1502computeContiguousStrides(MemRefType memRefType) {
1503 int64_t offset;
1505 if (failed(memRefType.getStridesAndOffset(strides, offset)))
1506 return std::nullopt;
1507 if (!strides.empty() && strides.back() != 1)
1508 return std::nullopt;
1509 // If no layout or identity layout, this is contiguous by definition.
1510 if (memRefType.getLayout().isIdentity())
1511 return strides;
1512
1513 // Otherwise, we must determine contiguity form shapes. This can only ever
1514 // work in static cases because MemRefType is underspecified to represent
1515 // contiguous dynamic shapes in other ways than with just empty/identity
1516 // layout.
1517 auto sizes = memRefType.getShape();
1518 for (int index = 0, e = strides.size() - 1; index < e; ++index) {
1519 if (ShapedType::isDynamic(sizes[index + 1]) ||
1520 ShapedType::isDynamic(strides[index]) ||
1521 ShapedType::isDynamic(strides[index + 1]))
1522 return std::nullopt;
1523 if (strides[index] != strides[index + 1] * sizes[index + 1])
1524 return std::nullopt;
1525 }
1526 return strides;
1527}
1528
1529class VectorTypeCastOpConversion
1530 : public ConvertOpToLLVMPattern<vector::TypeCastOp> {
1531public:
1532 using ConvertOpToLLVMPattern<vector::TypeCastOp>::ConvertOpToLLVMPattern;
1533
1534 LogicalResult
1535 matchAndRewrite(vector::TypeCastOp castOp, OpAdaptor adaptor,
1536 ConversionPatternRewriter &rewriter) const override {
1537 auto loc = castOp->getLoc();
1538 MemRefType sourceMemRefType =
1539 cast<MemRefType>(castOp.getOperand().getType());
1540 MemRefType targetMemRefType = castOp.getType();
1541
1542 // Only static shape casts supported atm.
1543 if (!sourceMemRefType.hasStaticShape() ||
1544 !targetMemRefType.hasStaticShape())
1545 return failure();
1546
1547 auto llvmSourceDescriptorTy =
1548 dyn_cast<LLVM::LLVMStructType>(adaptor.getOperands()[0].getType());
1549 if (!llvmSourceDescriptorTy)
1550 return failure();
1551 MemRefDescriptor sourceMemRef(adaptor.getOperands()[0]);
1552
1553 auto llvmTargetDescriptorTy = dyn_cast_or_null<LLVM::LLVMStructType>(
1554 typeConverter->convertType(targetMemRefType));
1555 if (!llvmTargetDescriptorTy)
1556 return failure();
1557
1558 // Only contiguous source buffers supported atm.
1559 auto sourceStrides = computeContiguousStrides(sourceMemRefType);
1560 if (!sourceStrides)
1561 return failure();
1562 auto targetStrides = computeContiguousStrides(targetMemRefType);
1563 if (!targetStrides)
1564 return failure();
1565 // Only support static strides for now, regardless of contiguity.
1566 if (llvm::any_of(*targetStrides, ShapedType::isDynamic))
1567 return failure();
1568
1569 // The offset, size and stride fields of a memref descriptor use the
1570 // converted index type.
1571 Type indexTy = getTypeConverter()->getIndexType();
1572
1573 // Create descriptor.
1574 auto desc = MemRefDescriptor::poison(rewriter, loc, llvmTargetDescriptorTy);
1575 // Set allocated ptr.
1576 Value allocated = sourceMemRef.allocatedPtr(rewriter, loc);
1577 desc.setAllocatedPtr(rewriter, loc, allocated);
1578
1579 // Set aligned ptr.
1580 Value ptr = sourceMemRef.alignedPtr(rewriter, loc);
1581 desc.setAlignedPtr(rewriter, loc, ptr);
1582 // Fill offset 0.
1583 desc.setOffset(rewriter, loc,
1584 LLVM::createIndexAttrConstant(rewriter, loc, indexTy, 0));
1585
1586 // Fill size and stride descriptors in memref.
1587 for (const auto &indexedSize :
1588 llvm::enumerate(targetMemRefType.getShape())) {
1589 int64_t index = indexedSize.index();
1590 desc.setSize(rewriter, loc, index,
1591 LLVM::createIndexAttrConstant(rewriter, loc, indexTy,
1592 indexedSize.value()));
1593 desc.setStride(rewriter, loc, index,
1594 LLVM::createIndexAttrConstant(rewriter, loc, indexTy,
1595 (*targetStrides)[index]));
1596 }
1597
1598 rewriter.replaceOp(castOp, {desc});
1599 return success();
1600 }
1601};
1602
1603/// Conversion pattern for a `vector.create_mask` (1-D scalable vectors only).
1604/// Non-scalable versions of this operation are handled in Vector Transforms.
1605class VectorCreateMaskOpConversion
1606 : public OpConversionPattern<vector::CreateMaskOp> {
1607public:
1608 explicit VectorCreateMaskOpConversion(MLIRContext *context,
1609 bool enableIndexOpt)
1610 : OpConversionPattern<vector::CreateMaskOp>(context),
1611 force32BitVectorIndices(enableIndexOpt) {}
1612
1613 LogicalResult
1614 matchAndRewrite(vector::CreateMaskOp op, OpAdaptor adaptor,
1615 ConversionPatternRewriter &rewriter) const override {
1616 auto dstType = op.getType();
1617 if (dstType.getRank() != 1 || !cast<VectorType>(dstType).isScalable())
1618 return failure();
1619 IntegerType idxType =
1620 force32BitVectorIndices ? rewriter.getI32Type() : rewriter.getI64Type();
1621 auto loc = op->getLoc();
1622 Value indices = LLVM::StepVectorOp::create(
1623 rewriter, loc,
1624 LLVM::getVectorType(idxType, dstType.getShape()[0],
1625 /*isScalable=*/true));
1626 Value maskBound = adaptor.getOperands()[0];
1627 // When using 32-bit indices, cap the bound at INT32_MAX in index type
1628 // before casting. For scalable vectors the runtime size (vscale * dim) is
1629 // unknown at compile time, so we can't clamp to `dim` as in the fixed-size
1630 // path. Clamping to INT32_MAX is safe because any realistic scalable vector
1631 // size fits well below this limit, so a bound >= vscale*dim still produces
1632 // an all-true mask after the comparison.
1633 if (force32BitVectorIndices) {
1634 Value maxBound =
1635 arith::ConstantIndexOp::create(rewriter, loc, (1LL << 31) - 1);
1636 maskBound = arith::MinSIOp::create(rewriter, loc, maskBound, maxBound);
1637 }
1638 auto bound =
1639 getValueOrCreateCastToIndexLike(rewriter, loc, idxType, maskBound);
1640 Value bounds = BroadcastOp::create(rewriter, loc, indices.getType(), bound);
1641 Value comp = arith::CmpIOp::create(rewriter, loc, arith::CmpIPredicate::slt,
1642 indices, bounds);
1643 rewriter.replaceOp(op, comp);
1644 return success();
1645 }
1646
1647private:
1648 const bool force32BitVectorIndices;
1649};
1650
1651class VectorPrintOpConversion : public ConvertOpToLLVMPattern<vector::PrintOp> {
1652 SymbolTableCollection *symbolTables = nullptr;
1653
1654public:
1655 explicit VectorPrintOpConversion(
1656 const LLVMTypeConverter &typeConverter,
1657 SymbolTableCollection *symbolTables = nullptr)
1658 : ConvertOpToLLVMPattern<vector::PrintOp>(typeConverter),
1659 symbolTables(symbolTables) {}
1660
1661 // Lowering implementation that relies on a small runtime support library,
1662 // which only needs to provide a few printing methods (single value for all
1663 // data types, opening/closing bracket, comma, newline). The lowering splits
1664 // the vector into elementary printing operations. The advantage of this
1665 // approach is that the library can remain unaware of all low-level
1666 // implementation details of vectors while still supporting output of any
1667 // shaped and dimensioned vector.
1668 //
1669 // Note: This lowering only handles scalars, n-D vectors are broken into
1670 // printing scalars in loops in VectorToSCF.
1671 //
1672 // TODO: rely solely on libc in future? something else?
1673 //
1674 LogicalResult
1675 matchAndRewrite(vector::PrintOp printOp, OpAdaptor adaptor,
1676 ConversionPatternRewriter &rewriter) const override {
1677 auto parent = printOp->getParentOfType<ModuleOp>();
1678 if (!parent)
1679 return failure();
1680
1681 auto loc = printOp->getLoc();
1682
1683 if (auto value = adaptor.getSource()) {
1684 Type printType = printOp.getPrintType();
1685 if (isa<VectorType>(printType)) {
1686 // Vectors should be broken into elementary print ops in VectorToSCF.
1687 return failure();
1688 }
1689 if (failed(emitScalarPrint(rewriter, parent, loc, printType, value)))
1690 return failure();
1691 }
1692
1693 auto punct = printOp.getPunctuation();
1694 if (auto stringLiteral = printOp.getStringLiteral()) {
1695 auto createResult =
1696 LLVM::createPrintStrCall(rewriter, loc, parent, "vector_print_str",
1697 *stringLiteral, *getTypeConverter(),
1698 /*addNewline=*/false);
1699 if (createResult.failed())
1700 return failure();
1701
1702 } else if (punct != PrintPunctuation::NoPunctuation) {
1703 FailureOr<LLVM::LLVMFuncOp> op = [&]() {
1704 switch (punct) {
1705 case PrintPunctuation::Close:
1706 return LLVM::lookupOrCreatePrintCloseFn(rewriter, parent,
1707 symbolTables);
1708 case PrintPunctuation::Open:
1709 return LLVM::lookupOrCreatePrintOpenFn(rewriter, parent,
1710 symbolTables);
1711 case PrintPunctuation::Comma:
1712 return LLVM::lookupOrCreatePrintCommaFn(rewriter, parent,
1713 symbolTables);
1714 case PrintPunctuation::NewLine:
1715 return LLVM::lookupOrCreatePrintNewlineFn(rewriter, parent,
1716 symbolTables);
1717 default:
1718 llvm_unreachable("unexpected punctuation");
1719 }
1720 }();
1721 if (failed(op))
1722 return failure();
1723 emitCall(rewriter, printOp->getLoc(), op.value());
1724 }
1725
1726 rewriter.eraseOp(printOp);
1727 return success();
1728 }
1729
1730private:
1731 enum class PrintConversion {
1732 // clang-format off
1733 None,
1734 ZeroExt64,
1735 SignExt64,
1736 Bitcast16
1737 // clang-format on
1738 };
1739
1740 LogicalResult emitScalarPrint(ConversionPatternRewriter &rewriter,
1741 ModuleOp parent, Location loc, Type printType,
1742 Value value) const {
1743 if (typeConverter->convertType(printType) == nullptr)
1744 return failure();
1745
1746 // Make sure element type has runtime support.
1747 PrintConversion conversion = PrintConversion::None;
1748 FailureOr<Operation *> printer;
1749 if (printType.isF32()) {
1750 printer = LLVM::lookupOrCreatePrintF32Fn(rewriter, parent, symbolTables);
1751 } else if (printType.isF64()) {
1752 printer = LLVM::lookupOrCreatePrintF64Fn(rewriter, parent, symbolTables);
1753 } else if (printType.isF16()) {
1754 conversion = PrintConversion::Bitcast16; // bits!
1755 printer = LLVM::lookupOrCreatePrintF16Fn(rewriter, parent, symbolTables);
1756 } else if (printType.isBF16()) {
1757 conversion = PrintConversion::Bitcast16; // bits!
1758 printer = LLVM::lookupOrCreatePrintBF16Fn(rewriter, parent, symbolTables);
1759 } else if (printType.isIndex()) {
1760 printer = LLVM::lookupOrCreatePrintU64Fn(rewriter, parent, symbolTables);
1761 } else if (auto intTy = dyn_cast<IntegerType>(printType)) {
1762 // Integers need a zero or sign extension on the operand
1763 // (depending on the source type) as well as a signed or
1764 // unsigned print method. Up to 64-bit is supported.
1765 unsigned width = intTy.getWidth();
1766 if (intTy.isUnsigned()) {
1767 if (width <= 64) {
1768 if (width < 64)
1769 conversion = PrintConversion::ZeroExt64;
1770 printer =
1771 LLVM::lookupOrCreatePrintU64Fn(rewriter, parent, symbolTables);
1772 } else {
1773 return failure();
1774 }
1775 } else {
1776 assert(intTy.isSignless() || intTy.isSigned());
1777 if (width <= 64) {
1778 // Note that we *always* zero extend booleans (1-bit integers),
1779 // so that true/false is printed as 1/0 rather than -1/0.
1780 if (width == 1)
1781 conversion = PrintConversion::ZeroExt64;
1782 else if (width < 64)
1783 conversion = PrintConversion::SignExt64;
1784 printer =
1785 LLVM::lookupOrCreatePrintI64Fn(rewriter, parent, symbolTables);
1786 } else {
1787 return failure();
1788 }
1789 }
1790 } else if (auto floatTy = dyn_cast<FloatType>(printType)) {
1791 // Print other floating-point types using the APFloat runtime library.
1792 int32_t sem =
1793 llvm::APFloatBase::SemanticsToEnum(floatTy.getFloatSemantics());
1794 Value semValue = LLVM::ConstantOp::create(
1795 rewriter, loc, rewriter.getI32Type(),
1796 rewriter.getIntegerAttr(rewriter.getI32Type(), sem));
1797 Value floatBits =
1798 LLVM::ZExtOp::create(rewriter, loc, rewriter.getI64Type(), value);
1799 printer =
1800 LLVM::lookupOrCreateApFloatPrintFn(rewriter, parent, symbolTables);
1801 emitCall(rewriter, loc, printer.value(),
1802 ValueRange({semValue, floatBits}));
1803 return success();
1804 } else {
1805 return failure();
1806 }
1807 if (failed(printer))
1808 return failure();
1809
1810 switch (conversion) {
1811 case PrintConversion::ZeroExt64:
1812 value = arith::ExtUIOp::create(
1813 rewriter, loc, IntegerType::get(rewriter.getContext(), 64), value);
1814 break;
1815 case PrintConversion::SignExt64:
1816 value = arith::ExtSIOp::create(
1817 rewriter, loc, IntegerType::get(rewriter.getContext(), 64), value);
1818 break;
1819 case PrintConversion::Bitcast16:
1820 value = LLVM::BitcastOp::create(
1821 rewriter, loc, IntegerType::get(rewriter.getContext(), 16), value);
1822 break;
1823 case PrintConversion::None:
1824 break;
1825 }
1826 emitCall(rewriter, loc, printer.value(), value);
1827 return success();
1828 }
1829
1830 // Helper to emit a call.
1831 static void emitCall(ConversionPatternRewriter &rewriter, Location loc,
1832 Operation *ref, ValueRange params = ValueRange()) {
1833 LLVM::CallOp::create(rewriter, loc, TypeRange(), SymbolRefAttr::get(ref),
1834 params);
1835 }
1836};
1837
1838/// A broadcast of a scalar is lowered to an insertelement + a shufflevector
1839/// operation. Only broadcasts to 0-d and 1-d vectors are lowered by this
1840/// pattern, the higher rank cases are handled by another pattern.
1841struct VectorBroadcastScalarToLowRankLowering
1842 : public ConvertOpToLLVMPattern<vector::BroadcastOp> {
1843 using ConvertOpToLLVMPattern<vector::BroadcastOp>::ConvertOpToLLVMPattern;
1844
1845 LogicalResult
1846 matchAndRewrite(vector::BroadcastOp broadcast, OpAdaptor adaptor,
1847 ConversionPatternRewriter &rewriter) const override {
1848 if (isa<VectorType>(broadcast.getSourceType()))
1849 return rewriter.notifyMatchFailure(
1850 broadcast, "broadcast from vector type not handled");
1851
1852 VectorType resultType = broadcast.getType();
1853 if (resultType.getRank() > 1)
1854 return rewriter.notifyMatchFailure(broadcast,
1855 "broadcast to 2+-d handled elsewhere");
1856
1857 // First insert it into a poison vector so we can shuffle it.
1858 auto vectorType = typeConverter->convertType(broadcast.getType());
1859 Value poison =
1860 LLVM::PoisonOp::create(rewriter, broadcast.getLoc(), vectorType);
1861 auto zero = LLVM::ConstantOp::create(
1862 rewriter, broadcast.getLoc(),
1863 typeConverter->convertType(rewriter.getIntegerType(32)),
1864 rewriter.getZeroAttr(rewriter.getIntegerType(32)));
1865
1866 // For 0-d vector, we simply do `insertelement`.
1867 if (resultType.getRank() == 0) {
1868 rewriter.replaceOpWithNewOp<LLVM::InsertElementOp>(
1869 broadcast, vectorType, poison, adaptor.getSource(), zero);
1870 return success();
1871 }
1872
1873 auto v =
1874 LLVM::InsertElementOp::create(rewriter, broadcast.getLoc(), vectorType,
1875 poison, adaptor.getSource(), zero);
1876
1877 // For 1-d vector, we additionally do a `shufflevector`.
1878 int64_t width = cast<VectorType>(broadcast.getType()).getDimSize(0);
1879 SmallVector<int32_t> zeroValues(width, 0);
1880
1881 // Shuffle the value across the desired number of elements.
1882 auto shuffle = rewriter.createOrFold<LLVM::ShuffleVectorOp>(
1883 broadcast.getLoc(), v, poison, zeroValues);
1884 rewriter.replaceOp(broadcast, shuffle);
1885 return success();
1886 }
1887};
1888
1889/// The broadcast of a scalar is lowered to an insertelement + a shufflevector
1890/// operation. Only broadcasts to 2+-d vector result types are lowered by this
1891/// pattern, the 1-d case is handled by another pattern. Broadcasts from vectors
1892/// are not converted to LLVM, only broadcasts from scalars are.
1893struct VectorBroadcastScalarToNdLowering
1894 : public ConvertOpToLLVMPattern<BroadcastOp> {
1895 using ConvertOpToLLVMPattern<BroadcastOp>::ConvertOpToLLVMPattern;
1896
1897 LogicalResult
1898 matchAndRewrite(BroadcastOp broadcast, OpAdaptor adaptor,
1899 ConversionPatternRewriter &rewriter) const override {
1900 if (isa<VectorType>(broadcast.getSourceType()))
1901 return rewriter.notifyMatchFailure(
1902 broadcast, "broadcast from vector type not handled");
1903
1904 VectorType resultType = broadcast.getType();
1905 if (resultType.getRank() <= 1)
1906 return rewriter.notifyMatchFailure(
1907 broadcast, "broadcast to 1-d or 0-d handled elsewhere");
1908
1909 // First insert it into a poison vector so we can shuffle it.
1910 auto loc = broadcast.getLoc();
1911 auto vectorTypeInfo =
1912 LLVM::detail::extractNDVectorTypeInfo(resultType, *getTypeConverter());
1913 auto llvmNDVectorTy = vectorTypeInfo.llvmNDVectorTy;
1914 auto llvm1DVectorTy = vectorTypeInfo.llvm1DVectorTy;
1915 if (!llvmNDVectorTy || !llvm1DVectorTy)
1916 return failure();
1917
1918 // Construct returned value.
1919 Value desc = LLVM::PoisonOp::create(rewriter, loc, llvmNDVectorTy);
1920
1921 // Construct a 1-D vector with the broadcasted value that we insert in all
1922 // the places within the returned descriptor.
1923 Value vdesc = LLVM::PoisonOp::create(rewriter, loc, llvm1DVectorTy);
1924 auto zero = LLVM::ConstantOp::create(
1925 rewriter, loc, typeConverter->convertType(rewriter.getIntegerType(32)),
1926 rewriter.getZeroAttr(rewriter.getIntegerType(32)));
1927 Value v = LLVM::InsertElementOp::create(rewriter, loc, llvm1DVectorTy,
1928 vdesc, adaptor.getSource(), zero);
1929
1930 // Shuffle the value across the desired number of elements.
1931 int64_t width = resultType.getDimSize(resultType.getRank() - 1);
1932 SmallVector<int32_t> zeroValues(width, 0);
1933 v = LLVM::ShuffleVectorOp::create(rewriter, loc, v, v, zeroValues);
1934
1935 // Iterate of linear index, convert to coords space and insert broadcasted
1936 // 1-D vector in each position.
1937 nDVectorIterate(vectorTypeInfo, rewriter, [&](ArrayRef<int64_t> position) {
1938 desc = LLVM::InsertValueOp::create(rewriter, loc, desc, v, position);
1939 });
1940 rewriter.replaceOp(broadcast, desc);
1941 return success();
1942 }
1943};
1944
1945/// Conversion pattern for a `vector.interleave`.
1946/// This supports fixed-sized vectors and scalable vectors.
1947struct VectorInterleaveOpLowering
1948 : public ConvertOpToLLVMPattern<vector::InterleaveOp> {
1950
1951 LogicalResult
1952 matchAndRewrite(vector::InterleaveOp interleaveOp, OpAdaptor adaptor,
1953 ConversionPatternRewriter &rewriter) const override {
1954 VectorType resultType = interleaveOp.getResultVectorType();
1955 // n-D interleaves should have been lowered already.
1956 if (resultType.getRank() != 1)
1957 return rewriter.notifyMatchFailure(interleaveOp,
1958 "InterleaveOp not rank 1");
1959 // If the result is rank 1, then this directly maps to LLVM.
1960 if (resultType.isScalable()) {
1961 rewriter.replaceOpWithNewOp<LLVM::vector_interleave2>(
1962 interleaveOp, typeConverter->convertType(resultType),
1963 adaptor.getLhs(), adaptor.getRhs());
1964 return success();
1965 }
1966 // Lower fixed-size interleaves to a shufflevector. While the
1967 // vector.interleave2 intrinsic supports fixed and scalable vectors, the
1968 // langref still recommends fixed-vectors use shufflevector, see:
1969 // https://llvm.org/docs/LangRef.html#id876.
1970 int64_t resultVectorSize = resultType.getNumElements();
1971 SmallVector<int32_t> interleaveShuffleMask;
1972 interleaveShuffleMask.reserve(resultVectorSize);
1973 for (int i = 0, end = resultVectorSize / 2; i < end; ++i) {
1974 interleaveShuffleMask.push_back(i);
1975 interleaveShuffleMask.push_back((resultVectorSize / 2) + i);
1976 }
1977 rewriter.replaceOpWithNewOp<LLVM::ShuffleVectorOp>(
1978 interleaveOp, adaptor.getLhs(), adaptor.getRhs(),
1979 interleaveShuffleMask);
1980 return success();
1981 }
1982};
1983
1984/// Conversion pattern for a `vector.deinterleave`.
1985/// This supports fixed-sized vectors and scalable vectors.
1986struct VectorDeinterleaveOpLowering
1987 : public ConvertOpToLLVMPattern<vector::DeinterleaveOp> {
1989
1990 LogicalResult
1991 matchAndRewrite(vector::DeinterleaveOp deinterleaveOp, OpAdaptor adaptor,
1992 ConversionPatternRewriter &rewriter) const override {
1993 VectorType resultType = deinterleaveOp.getResultVectorType();
1994 VectorType sourceType = deinterleaveOp.getSourceVectorType();
1995 auto loc = deinterleaveOp.getLoc();
1996
1997 // Note: n-D deinterleave operations should be lowered to the 1-D before
1998 // converting to LLVM.
1999 if (resultType.getRank() != 1)
2000 return rewriter.notifyMatchFailure(deinterleaveOp,
2001 "DeinterleaveOp not rank 1");
2002
2003 if (resultType.isScalable()) {
2004 const auto *llvmTypeConverter = this->getTypeConverter();
2005 auto deinterleaveResults = deinterleaveOp.getResultTypes();
2006 auto packedOpResults =
2007 llvmTypeConverter->packOperationResults(deinterleaveResults);
2008 auto intrinsic = LLVM::vector_deinterleave2::create(
2009 rewriter, loc, packedOpResults, adaptor.getSource());
2010
2011 auto evenResult = LLVM::ExtractValueOp::create(
2012 rewriter, loc, intrinsic->getResult(0), 0);
2013 auto oddResult = LLVM::ExtractValueOp::create(rewriter, loc,
2014 intrinsic->getResult(0), 1);
2015
2016 rewriter.replaceOp(deinterleaveOp, ValueRange{evenResult, oddResult});
2017 return success();
2018 }
2019 // Lower fixed-size deinterleave to two shufflevectors. While the
2020 // vector.deinterleave2 intrinsic supports fixed and scalable vectors, the
2021 // langref still recommends fixed-vectors use shufflevector, see:
2022 // https://llvm.org/docs/LangRef.html#id889.
2023 int64_t resultVectorSize = resultType.getNumElements();
2024 SmallVector<int32_t> evenShuffleMask;
2025 SmallVector<int32_t> oddShuffleMask;
2026
2027 evenShuffleMask.reserve(resultVectorSize);
2028 oddShuffleMask.reserve(resultVectorSize);
2029
2030 for (int i = 0; i < sourceType.getNumElements(); ++i) {
2031 if (i % 2 == 0)
2032 evenShuffleMask.push_back(i);
2033 else
2034 oddShuffleMask.push_back(i);
2035 }
2036
2037 auto poison = LLVM::PoisonOp::create(rewriter, loc, sourceType);
2038 auto evenShuffle = LLVM::ShuffleVectorOp::create(
2039 rewriter, loc, adaptor.getSource(), poison, evenShuffleMask);
2040 auto oddShuffle = LLVM::ShuffleVectorOp::create(
2041 rewriter, loc, adaptor.getSource(), poison, oddShuffleMask);
2042
2043 rewriter.replaceOp(deinterleaveOp, ValueRange{evenShuffle, oddShuffle});
2044 return success();
2045 }
2046};
2047
2048/// Conversion pattern for a `vector.from_elements`.
2049struct VectorFromElementsLowering
2050 : public ConvertOpToLLVMPattern<vector::FromElementsOp> {
2052
2053 LogicalResult
2054 matchAndRewrite(vector::FromElementsOp fromElementsOp, OpAdaptor adaptor,
2055 ConversionPatternRewriter &rewriter) const override {
2056 Location loc = fromElementsOp.getLoc();
2057 VectorType vectorType = fromElementsOp.getType();
2058 // Only support 1-D vectors. Multi-dimensional vectors should have been
2059 // transformed to 1-D vectors by the vector-to-vector transformations before
2060 // this.
2061 if (vectorType.getRank() > 1)
2062 return rewriter.notifyMatchFailure(fromElementsOp,
2063 "rank > 1 vectors are not supported");
2064 Type llvmType = typeConverter->convertType(vectorType);
2065 Type llvmIndexType = typeConverter->convertType(rewriter.getIndexType());
2066 Value result = LLVM::PoisonOp::create(rewriter, loc, llvmType);
2067 for (auto [idx, val] : llvm::enumerate(adaptor.getElements())) {
2068 auto constIdx =
2069 LLVM::ConstantOp::create(rewriter, loc, llvmIndexType, idx);
2070 result = LLVM::InsertElementOp::create(rewriter, loc, llvmType, result,
2071 val, constIdx);
2072 }
2073 rewriter.replaceOp(fromElementsOp, result);
2074 return success();
2075 }
2076};
2077
2078/// Conversion pattern for a `vector.to_elements`.
2079struct VectorToElementsLowering
2080 : public ConvertOpToLLVMPattern<vector::ToElementsOp> {
2082
2083 LogicalResult
2084 matchAndRewrite(vector::ToElementsOp toElementsOp, OpAdaptor adaptor,
2085 ConversionPatternRewriter &rewriter) const override {
2086 Location loc = toElementsOp.getLoc();
2087 auto idxType = typeConverter->convertType(rewriter.getIndexType());
2088 Value source = adaptor.getSource();
2089
2090 SmallVector<Value> results(toElementsOp->getNumResults());
2091 for (auto [idx, element] : llvm::enumerate(toElementsOp.getElements())) {
2092 // Create an extractelement operation only for results that are not dead.
2093 if (element.use_empty())
2094 continue;
2095
2096 auto constIdx = LLVM::ConstantOp::create(
2097 rewriter, loc, idxType, rewriter.getIntegerAttr(idxType, idx));
2098 auto llvmType = typeConverter->convertType(element.getType());
2099
2100 Value result = LLVM::ExtractElementOp::create(rewriter, loc, llvmType,
2101 source, constIdx);
2102 results[idx] = result;
2103 }
2104
2105 rewriter.replaceOp(toElementsOp, results);
2106 return success();
2107 }
2108};
2109
2110/// Conversion pattern for vector.step.
2111struct VectorStepOpLowering : public ConvertOpToLLVMPattern<vector::StepOp> {
2113
2114 LogicalResult
2115 matchAndRewrite(vector::StepOp stepOp, OpAdaptor adaptor,
2116 ConversionPatternRewriter &rewriter) const override {
2117 Type llvmType = typeConverter->convertType(stepOp.getType());
2118 rewriter.replaceOpWithNewOp<LLVM::StepVectorOp>(stepOp, llvmType);
2119 return success();
2120 }
2121};
2122
2123/// Progressive lowering of a `vector.contract %a, %b, %c` with row-major matmul
2124/// semantics to:
2125/// ```
2126/// %flattened_a = vector.shape_cast %a
2127/// %flattened_b = vector.shape_cast %b
2128/// %flattened_d = vector.matrix_multiply %flattened_a, %flattened_b
2129/// %d = vector.shape_cast %%flattened_d
2130/// %e = add %c, %d
2131/// ```
2132/// `vector.matrix_multiply` later lowers to `llvm.matrix.multiply`.
2133class ContractionOpToMatmulOpLowering
2134 : public vector::MaskableOpRewritePattern<vector::ContractionOp> {
2135public:
2136 using MaskableOpRewritePattern::MaskableOpRewritePattern;
2137
2138 ContractionOpToMatmulOpLowering(MLIRContext *context,
2139 PatternBenefit benefit = 100)
2140 : MaskableOpRewritePattern<vector::ContractionOp>(context, benefit) {}
2141
2142 FailureOr<Value>
2143 matchAndRewriteMaskableOp(vector::ContractionOp op, MaskingOpInterface maskOp,
2144 PatternRewriter &rewriter) const override;
2145};
2146
2147/// Lower a qualifying `vector.contract %a, %b, %c` (with row-major matmul
2148/// semantics directly into `llvm.intr.matrix.multiply`:
2149/// BEFORE:
2150/// ```mlir
2151/// %res = vector.contract #matmat_trait %lhs, %rhs, %acc
2152/// : vector<2x4xf32>, vector<4x3xf32> into vector<2x3xf32>
2153/// ```
2154///
2155/// AFTER:
2156/// ```mlir
2157/// %lhs = vector.shape_cast %arg0 : vector<2x4xf32> to vector<8xf32>
2158/// %rhs = vector.shape_cast %arg1 : vector<4x3xf32> to vector<12xf32>
2159/// %matmul = llvm.intr.matrix.multiply %lhs, %rhs
2160/// %res = arith.addf %acc, %matmul : vector<2x3xf32>
2161/// ```
2162//
2163/// Scalable vectors are not supported.
2164FailureOr<Value> ContractionOpToMatmulOpLowering::matchAndRewriteMaskableOp(
2165 vector::ContractionOp op, MaskingOpInterface maskOp,
2166 PatternRewriter &rew) const {
2167 // TODO: Support vector.mask.
2168 if (maskOp)
2169 return failure();
2170
2171 auto iteratorTypes = op.getIteratorTypes().getValue();
2172 if (!isParallelIterator(iteratorTypes[0]) ||
2173 !isParallelIterator(iteratorTypes[1]) ||
2174 !isReductionIterator(iteratorTypes[2]))
2175 return failure();
2176
2177 Type opResType = op.getType();
2178 VectorType vecType = dyn_cast<VectorType>(opResType);
2179 if (vecType && vecType.isScalable()) {
2180 // Note - this is sufficient to reject all cases with scalable vectors.
2181 return failure();
2182 }
2183
2184 Type elementType = op.getLhsType().getElementType();
2185 if (!elementType.isIntOrFloat())
2186 return failure();
2187
2188 Type dstElementType = vecType ? vecType.getElementType() : opResType;
2189 if (elementType != dstElementType)
2190 return failure();
2191
2192 // Perform lhs + rhs transpositions to conform to matmul row-major semantics.
2193 // Bail out if the contraction cannot be put in this form.
2194 MLIRContext *ctx = op.getContext();
2195 Location loc = op.getLoc();
2196 AffineExpr m, n, k;
2197 bindDims(rew.getContext(), m, n, k);
2198 // LHS must be A(m, k) or A(k, m).
2199 Value lhs = op.getLhs();
2200 auto lhsMap = op.getIndexingMapsArray()[0];
2201 if (lhsMap == AffineMap::get(3, 0, {k, m}, ctx))
2202 lhs = vector::TransposeOp::create(rew, loc, lhs, ArrayRef<int64_t>{1, 0});
2203 else if (lhsMap != AffineMap::get(3, 0, {m, k}, ctx))
2204 return failure();
2205
2206 // RHS must be B(k, n) or B(n, k).
2207 Value rhs = op.getRhs();
2208 auto rhsMap = op.getIndexingMapsArray()[1];
2209 if (rhsMap == AffineMap::get(3, 0, {n, k}, ctx))
2210 rhs = vector::TransposeOp::create(rew, loc, rhs, ArrayRef<int64_t>{1, 0});
2211 else if (rhsMap != AffineMap::get(3, 0, {k, n}, ctx))
2212 return failure();
2213
2214 // At this point lhs and rhs are in row-major.
2215 VectorType lhsType = cast<VectorType>(lhs.getType());
2216 VectorType rhsType = cast<VectorType>(rhs.getType());
2217 int64_t lhsRows = lhsType.getDimSize(0);
2218 int64_t lhsColumns = lhsType.getDimSize(1);
2219 int64_t rhsColumns = rhsType.getDimSize(1);
2220
2221 Type flattenedLHSType =
2222 VectorType::get(lhsType.getNumElements(), lhsType.getElementType());
2223 lhs = vector::ShapeCastOp::create(rew, loc, flattenedLHSType, lhs);
2224
2225 Type flattenedRHSType =
2226 VectorType::get(rhsType.getNumElements(), rhsType.getElementType());
2227 rhs = vector::ShapeCastOp::create(rew, loc, flattenedRHSType, rhs);
2228
2229 Value mul = LLVM::MatrixMultiplyOp::create(
2230 rew, loc,
2231 VectorType::get(lhsRows * rhsColumns,
2232 cast<VectorType>(lhs.getType()).getElementType()),
2233 lhs, rhs, lhsRows, lhsColumns, rhsColumns);
2234
2235 mul = vector::ShapeCastOp::create(
2236 rew, loc,
2237 VectorType::get({lhsRows, rhsColumns},
2238 getElementTypeOrSelf(op.getAcc().getType())),
2239 mul);
2240
2241 // ACC must be C(m, n) or C(n, m).
2242 auto accMap = op.getIndexingMapsArray()[2];
2243 if (accMap == AffineMap::get(3, 0, {n, m}, ctx))
2244 mul = vector::TransposeOp::create(rew, loc, mul, ArrayRef<int64_t>{1, 0});
2245 else if (accMap != AffineMap::get(3, 0, {m, n}, ctx))
2246 llvm_unreachable("invalid contraction semantics");
2247
2248 Value res = isa<IntegerType>(elementType)
2249 ? static_cast<Value>(
2250 arith::AddIOp::create(rew, loc, op.getAcc(), mul))
2251 : static_cast<Value>(
2252 arith::AddFOp::create(rew, loc, op.getAcc(), mul));
2253
2254 return res;
2255}
2256
2257/// Lowers vector.transpose directly to llvm.intr.matrix.transpose
2258///
2259/// BEFORE:
2260/// ```mlir
2261/// %tr = vector.transpose %vec, [1, 0] : vector<2x4xf32> to vector<4x2xf32>
2262/// ```
2263/// AFTER:
2264/// ```mlir
2265/// %vec_cs = vector.shape_cast %vec : vector<2x4xf32> to vector<8xf32>
2266/// %tr = llvm.intr.matrix.transpose %vec_sc
2267/// {columns = 2 : i32, rows = 4 : i32} : vector<8xf32> into vector<8xf32>
2268/// %res = vector.shape_cast %tr : vector<8xf32> to vector<4x2xf32>
2269/// ```
2270class TransposeOpToMatrixTransposeOpLowering
2271 : public OpRewritePattern<vector::TransposeOp> {
2272public:
2273 using Base::Base;
2274
2275 LogicalResult matchAndRewrite(vector::TransposeOp op,
2276 PatternRewriter &rewriter) const override {
2277 auto loc = op.getLoc();
2278
2279 Value input = op.getVector();
2280 VectorType inputType = op.getSourceVectorType();
2281 VectorType resType = op.getResultVectorType();
2282
2283 if (inputType.isScalable())
2284 return rewriter.notifyMatchFailure(
2285 op, "This lowering does not support scalable vectors");
2286
2287 // Set up convenience transposition table.
2288 ArrayRef<int64_t> transp = op.getPermutation();
2289
2290 if (resType.getRank() != 2 || transp[0] != 1 || transp[1] != 0) {
2291 return failure();
2292 }
2293
2294 Type flattenedType =
2295 VectorType::get(resType.getNumElements(), resType.getElementType());
2296 auto matrix =
2297 vector::ShapeCastOp::create(rewriter, loc, flattenedType, input);
2298 auto rows = rewriter.getI32IntegerAttr(resType.getShape()[0]);
2299 auto columns = rewriter.getI32IntegerAttr(resType.getShape()[1]);
2300 Value trans = LLVM::MatrixTransposeOp::create(rewriter, loc, flattenedType,
2301 matrix, rows, columns);
2302 rewriter.replaceOpWithNewOp<vector::ShapeCastOp>(op, resType, trans);
2303 return success();
2304 }
2305};
2306
2307} // namespace
2308
2310 RewritePatternSet &patterns) {
2311 patterns.add<VectorFMAOpNDRewritePattern>(patterns.getContext());
2312}
2313
2315 RewritePatternSet &patterns, PatternBenefit benefit) {
2316 patterns.add<ContractionOpToMatmulOpLowering>(patterns.getContext(), benefit);
2317}
2318
2320 RewritePatternSet &patterns, PatternBenefit benefit) {
2321 patterns.add<TransposeOpToMatrixTransposeOpLowering>(patterns.getContext(),
2322 benefit);
2323}
2324
2325/// Populate the given list with patterns that convert from Vector to LLVM.
2327 const LLVMTypeConverter &converter, RewritePatternSet &patterns,
2328 bool reassociateFPReductions, bool force32BitVectorIndices,
2329 bool useVectorAlignment, bool enableGEPInboundsNuw) {
2330 // This function populates only ConversionPatterns, not RewritePatterns.
2331 MLIRContext *ctx = converter.getDialect()->getContext();
2332 patterns.add<VectorReductionOpConversion>(converter, reassociateFPReductions);
2333 patterns.add<VectorCreateMaskOpConversion>(ctx, force32BitVectorIndices);
2334 patterns.add<VectorLoadStoreConversion<vector::LoadOp>,
2335 VectorLoadStoreConversion<vector::MaskedLoadOp>,
2336 VectorLoadStoreConversion<vector::StoreOp>,
2337 VectorLoadStoreConversion<vector::MaskedStoreOp>>(
2338 converter, useVectorAlignment, enableGEPInboundsNuw);
2339 patterns.add<VectorGatherOpConversion, VectorScatterOpConversion>(
2340 converter, useVectorAlignment);
2341 patterns.add<VectorBitCastOpConversion, VectorShuffleOpConversion,
2342 VectorExtractOpConversion, VectorFMAOp1DConversion,
2343 VectorInsertOpConversion, VectorPrintOpConversion,
2344 VectorTypeCastOpConversion, VectorScaleOpConversion,
2345 VectorExpandLoadOpConversion, VectorCompressStoreOpConversion,
2346 VectorBroadcastScalarToLowRankLowering,
2347 VectorBroadcastScalarToNdLowering,
2348 VectorScalableInsertOpLowering, VectorScalableExtractOpLowering,
2349 MaskedReductionOpConversion, VectorInterleaveOpLowering,
2350 VectorDeinterleaveOpLowering, VectorFromElementsLowering,
2351 VectorToElementsLowering, VectorStepOpLowering>(converter);
2352}
2353
2354namespace {
2355struct VectorToLLVMDialectInterface : public ConvertToLLVMPatternInterface {
2356 VectorToLLVMDialectInterface(Dialect *dialect)
2357 : ConvertToLLVMPatternInterface(dialect) {}
2358
2359 using ConvertToLLVMPatternInterface::ConvertToLLVMPatternInterface;
2360 void loadDependentDialects(MLIRContext *context) const final {
2361 context->loadDialect<LLVM::LLVMDialect>();
2362 }
2363
2364 /// Hook for derived dialect interface to provide conversion patterns
2365 /// and mark dialect legal for the conversion target.
2366 void populateConvertToLLVMConversionPatterns(
2367 ConversionTarget &target, LLVMTypeConverter &typeConverter,
2368 RewritePatternSet &patterns) const final {
2369 populateVectorToLLVMConversionPatterns(typeConverter, patterns);
2370 }
2371};
2372} // namespace
2373
2375 DialectRegistry &registry) {
2376 registry.addExtension(+[](MLIRContext *ctx, vector::VectorDialect *dialect) {
2377 dialect->addInterfaces<VectorToLLVMDialectInterface>();
2378 });
2379}
return success()
static Value getIndexedPtrs(ConversionPatternRewriter &rewriter, Location loc, const LLVMTypeConverter &typeConverter, MemRefType memRefType, Value llvmMemref, Value base, Value index, VectorType vectorType)
LogicalResult getVectorToLLVMAlignment(const LLVMTypeConverter &typeConverter, VectorType vectorType, MemRefType memrefType, unsigned &align, bool useVectorAlignment)
LogicalResult getVectorAlignment(const LLVMTypeConverter &typeConverter, VectorType vectorType, unsigned &align)
LogicalResult getMemRefAlignment(const LLVMTypeConverter &typeConverter, MemRefType memrefType, unsigned &align)
static Value extractOne(ConversionPatternRewriter &rewriter, const LLVMTypeConverter &typeConverter, Location loc, Value val, Type llvmType, int64_t rank, int64_t pos)
static Value insertOne(ConversionPatternRewriter &rewriter, const LLVMTypeConverter &typeConverter, Location loc, Value val1, Value val2, Type llvmType, int64_t rank, int64_t pos)
static Value getAsLLVMValue(OpBuilder &builder, Location loc, OpFoldResult foldResult)
Convert foldResult into a Value.
static LogicalResult isMemRefTypeSupported(MemRefType memRefType, const LLVMTypeConverter &converter)
LogicalResult initialize(unsigned origNumLoops, ArrayRef< ReassociationIndices > foldedIterationDims)
lhs
static Value broadcast(Location loc, Value toBroadcast, unsigned numElements, const TypeConverter &typeConverter, ConversionPatternRewriter &rewriter)
Broadcasts the value to vector with numElements number of elements.
static void printOp(llvm::raw_ostream &os, Operation *op, OpPrintingFlags &flags)
Definition Unit.cpp:18
#define mul(a, b)
static AffineMap get(MLIRContext *context)
Returns a zero result affine map with no dimensions or symbols: () -> ().
IntegerAttr getI32IntegerAttr(int32_t value)
Definition Builders.cpp:208
TypedAttr getZeroAttr(Type type)
Definition Builders.cpp:333
MLIRContext * getContext() const
Definition Builders.h:56
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
static DenseElementsAttr get(ShapedType type, ArrayRef< Attribute > values)
Constructs a dense elements attribute from an array of element values.
The DialectRegistry maps a dialect namespace to a constructor for the matching dialect.
bool addExtension(TypeID extensionID, std::unique_ptr< DialectExtensionBase > extension)
Add the given extension to the registry.
Dialects are groups of MLIR operations, types and attributes, as well as behavior associated with the...
Definition Dialect.h:38
Conversion from types to the LLVM IR dialect.
const llvm::DataLayout & getDataLayout() const
Returns the data layout to use during and after conversion.
FailureOr< unsigned > getMemRefAddressSpace(BaseMemRefType type) const
Return the LLVM address space corresponding to the memory space of the memref type type or failure if...
LLVM::LLVMDialect * getDialect() const
Returns the LLVM dialect.
Utility class to translate MLIR LLVM dialect types to LLVM IR.
Definition TypeToLLVM.h:39
unsigned getPreferredAlignment(Type type, const llvm::DataLayout &layout)
Returns the preferred alignment for the type given the data layout.
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
Helper class to produce LLVM dialect operations extracting or inserting elements of a MemRef descript...
LLVM::LLVMPointerType getElementPtrType()
Returns the (LLVM) pointer type this descriptor contains.
Generic implementation of one-to-one conversion from "SourceOp" to "TargetOp" where the latter belong...
Definition Pattern.h:336
This class helps build Operations.
Definition Builders.h:210
This class represents a single result from folding an operation.
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,...
OpTy replaceOpWithNewOp(Operation *op, Args &&...args)
Replace the results of the given (original) op with a new op that is created without verification (re...
Instances of the Type class are uniqued, have an immutable identifier and an optional mutable compone...
Definition Types.h:74
bool isIntOrIndex() const
Return true if this is an integer (of any signedness) or an index type.
Definition Types.cpp:114
bool isIntOrFloat() const
Return true if this is an integer (of any signedness) or a float type.
Definition Types.cpp:118
unsigned getIntOrFloatBitWidth() const
Return the bit width of an integer or a float type, assert failure on other types.
Definition Types.cpp:124
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
Location getLoc() const
Return the location of this value.
Definition Value.cpp:24
static ConstantIndexOp create(OpBuilder &builder, Location location, int64_t value)
Definition ArithOps.cpp:398
void printType(Type type, AsmPrinter &printer)
Prints an LLVM Dialect type.
void nDVectorIterate(const NDVectorTypeInfo &info, OpBuilder &builder, function_ref< void(ArrayRef< int64_t >)> fun)
NDVectorTypeInfo extractNDVectorTypeInfo(VectorType vectorType, const LLVMTypeConverter &converter)
FailureOr< LLVM::LLVMFuncOp > lookupOrCreatePrintBF16Fn(OpBuilder &b, Operation *moduleOp, SymbolTableCollection *symbolTables=nullptr)
FailureOr< LLVM::LLVMFuncOp > lookupOrCreatePrintOpenFn(OpBuilder &b, Operation *moduleOp, SymbolTableCollection *symbolTables=nullptr)
Value getStridedElementPtr(OpBuilder &builder, Location loc, const LLVMTypeConverter &converter, MemRefType type, Value memRefDesc, ValueRange indices, LLVM::GEPNoWrapFlags noWrapFlags=LLVM::GEPNoWrapFlags::none)
Performs the index computation to get to the element at indices of the memory pointed to by memRefDes...
Definition Pattern.cpp:620
Type getVectorType(Type elementType, unsigned numElements, bool isScalable=false)
Creates an LLVM dialect-compatible vector type with the given element type and length.
FailureOr< LLVM::LLVMFuncOp > lookupOrCreatePrintCommaFn(OpBuilder &b, Operation *moduleOp, SymbolTableCollection *symbolTables=nullptr)
FailureOr< LLVM::LLVMFuncOp > lookupOrCreatePrintI64Fn(OpBuilder &b, Operation *moduleOp, SymbolTableCollection *symbolTables=nullptr)
Helper functions to look up or create the declaration for commonly used external C function calls.
Value createIndexAttrConstant(OpBuilder &builder, Location loc, Type resultType, int64_t value)
Creates an llvm.mlir.constant producing value as resultType, which is expected to be the converted in...
Definition Pattern.cpp:58
FailureOr< LLVM::LLVMFuncOp > lookupOrCreatePrintNewlineFn(OpBuilder &b, Operation *moduleOp, SymbolTableCollection *symbolTables=nullptr)
FailureOr< LLVM::LLVMFuncOp > lookupOrCreatePrintCloseFn(OpBuilder &b, Operation *moduleOp, SymbolTableCollection *symbolTables=nullptr)
FailureOr< LLVM::LLVMFuncOp > lookupOrCreatePrintU64Fn(OpBuilder &b, Operation *moduleOp, SymbolTableCollection *symbolTables=nullptr)
FailureOr< LLVM::LLVMFuncOp > lookupOrCreatePrintF32Fn(OpBuilder &b, Operation *moduleOp, SymbolTableCollection *symbolTables=nullptr)
FailureOr< LLVM::LLVMFuncOp > lookupOrCreateApFloatPrintFn(OpBuilder &b, Operation *moduleOp, SymbolTableCollection *symbolTables=nullptr)
LogicalResult createPrintStrCall(OpBuilder &builder, Location loc, ModuleOp moduleOp, StringRef symbolName, StringRef string, const LLVMTypeConverter &typeConverter, bool addNewline=true, std::optional< StringRef > runtimeFunctionName={}, SymbolTableCollection *symbolTables=nullptr)
Generate IR that prints the given string to stdout.
FailureOr< LLVM::LLVMFuncOp > lookupOrCreatePrintF16Fn(OpBuilder &b, Operation *moduleOp, SymbolTableCollection *symbolTables=nullptr)
FailureOr< LLVM::LLVMFuncOp > lookupOrCreatePrintF64Fn(OpBuilder &b, Operation *moduleOp, SymbolTableCollection *symbolTables=nullptr)
LLVM::FastmathFlags convertArithFastMathFlagsToLLVM(arith::FastMathFlags arithFMF)
Maps arithmetic fastmath enum values to LLVM enum values.
bool hasNegativeStaticStride(MemRefType memRefTy)
Returns true if any stride of memRefTy is statically known to be negative.
detail::InFlightRemark failed(Location loc, RemarkOpts opts)
Report an optimization remark that failed.
Definition Remarks.h:734
bool isReductionIterator(Attribute attr)
Returns true if attr has "reduction" iterator type semantics.
Definition VectorOps.h:156
void populateVectorContractToMatrixMultiply(RewritePatternSet &patterns, PatternBenefit benefit=100)
Populate the pattern set with the following patterns:
void populateVectorRankReducingFMAPattern(RewritePatternSet &patterns)
Populates a pattern that rank-reduces n-D FMAs into (n-1)-D FMAs where n > 1.
bool isParallelIterator(Attribute attr)
Returns true if attr has "parallel" iterator type semantics.
Definition VectorOps.h:151
void registerConvertVectorToLLVMInterface(DialectRegistry &registry)
SmallVector< int64_t > getAsIntegers(ArrayRef< Value > values)
Returns the integer numbers in values.
void populateVectorTransposeToFlatTranspose(RewritePatternSet &patterns, PatternBenefit benefit=100)
Populate the pattern set with the following patterns:
Value createReductionNeutralValue(OpBuilder &builder, Location loc, Type type, vector::CombiningKind kind)
Creates a constant filled with the neutral (identity) value for the given reduction kind.
Include the generated interface declarations.
SmallVector< OpFoldResult > getMixedValues(ArrayRef< int64_t > staticValues, ValueRange dynamicValues, MLIRContext *context)
Return a vector of OpFoldResults with the same size a staticValues, but all elements for which Shaped...
void populateVectorToLLVMConversionPatterns(const LLVMTypeConverter &converter, RewritePatternSet &patterns, bool reassociateFPReductions=false, bool force32BitVectorIndices=false, bool useVectorAlignment=false, bool enableGEPInboundsNuw=false)
Collect a set of patterns to convert from the Vector dialect to LLVM.
void bindDims(MLIRContext *ctx, AffineExprTy &...exprs)
Bind a list of AffineExpr references to DimExpr at positions: [0 .
Definition AffineExpr.h:311
Value getValueOrCreateCastToIndexLike(OpBuilder &b, Location loc, Type targetType, Value value)
Create a cast from an index-like value (index or integer) to another index-like value.
Definition Utils.cpp:122
Type getElementTypeOrSelf(Type type)
Return the element type or return the type itself.
OpRewritePattern is a wrapper around RewritePattern that allows for matching and rewriting against an...
A pattern for ops that implement MaskableOpInterface and that might be masked (i.e.