261 lines · cpp
1//===- VectorTransformOps.cpp - Implementation of Vector transform ops ----===//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-exception6//7//===----------------------------------------------------------------------===//8 9#include "mlir/Dialect/Vector/TransformOps/VectorTransformOps.h"10 11#include "mlir/Conversion/LLVMCommon/TypeConverter.h"12#include "mlir/Conversion/VectorToLLVM/ConvertVectorToLLVM.h"13#include "mlir/Conversion/VectorToSCF/VectorToSCF.h"14#include "mlir/Dialect/LLVMIR/LLVMDialect.h"15#include "mlir/Dialect/Transform/IR/TransformDialect.h"16#include "mlir/Dialect/Transform/Interfaces/TransformInterfaces.h"17#include "mlir/Dialect/Vector/IR/VectorOps.h"18#include "mlir/Dialect/Vector/Transforms/LoweringPatterns.h"19#include "mlir/Dialect/Vector/Transforms/VectorRewritePatterns.h"20#include "mlir/Dialect/Vector/Transforms/VectorTransforms.h"21#include "mlir/Dialect/X86Vector/Transforms.h"22 23using namespace mlir;24using namespace mlir::vector;25using namespace mlir::transform;26 27//===----------------------------------------------------------------------===//28// Apply...ConversionPatternsOp29//===----------------------------------------------------------------------===//30 31void transform::ApplyVectorToLLVMConversionPatternsOp::populatePatterns(32 TypeConverter &typeConverter, RewritePatternSet &patterns) {33 populateVectorToLLVMConversionPatterns(34 static_cast<LLVMTypeConverter &>(typeConverter), patterns,35 getReassociateFpReductions(), getForce_32bitVectorIndices(),36 getUseVectorAlignment());37}38 39LogicalResult40transform::ApplyVectorToLLVMConversionPatternsOp::verifyTypeConverter(41 transform::TypeConverterBuilderOpInterface builder) {42 if (builder.getTypeConverterType() != "LLVMTypeConverter")43 return emitOpError("expected LLVMTypeConverter");44 return success();45}46 47//===----------------------------------------------------------------------===//48// Apply...PatternsOp49//===----------------------------------------------------------------------===//50 51void transform::ApplyCastAwayVectorLeadingOneDimPatternsOp::populatePatterns(52 RewritePatternSet &patterns) {53 vector::populateCastAwayVectorLeadingOneDimPatterns(patterns);54}55 56void transform::ApplyFoldArithExtensionPatternsOp::populatePatterns(57 RewritePatternSet &patterns) {58 vector::populateFoldArithExtensionPatterns(patterns);59}60 61void transform::ApplyFoldElementwiseToVectorPatternsOp::populatePatterns(62 RewritePatternSet &patterns) {63 vector::populateElementwiseToVectorOpsPatterns(patterns);64}65 66void transform::ApplyVectorReductionToContractPatternsOp::populatePatterns(67 RewritePatternSet &patterns) {68 vector::populateVectorReductionToContractPatterns(patterns);69}70 71void transform::ApplyLowerCreateMaskPatternsOp::populatePatterns(72 RewritePatternSet &patterns) {73 vector::populateVectorMaskOpLoweringPatterns(patterns);74}75 76void transform::ApplyRankReducingSubviewPatternsOp::populatePatterns(77 RewritePatternSet &patterns) {78 vector::populateVectorTransferDropUnitDimsPatterns(patterns);79}80 81void transform::ApplyTransferPermutationPatternsOp::populatePatterns(82 RewritePatternSet &patterns) {83 vector::populateVectorTransferPermutationMapLoweringPatterns(patterns);84}85 86void transform::ApplyDropUnitDimWithShapeCastPatternsOp::populatePatterns(87 RewritePatternSet &patterns) {88 vector::populateDropUnitDimWithShapeCastPatterns(patterns);89}90 91void transform::ApplyDropInnerMostUnitDimsFromXferOpsPatternsOp::92 populatePatterns(RewritePatternSet &patterns) {93 vector::populateDropInnerMostUnitDimsXferOpPatterns(patterns);94}95 96void transform::ApplyLowerBitCastPatternsOp::populatePatterns(97 RewritePatternSet &patterns) {98 vector::populateVectorBitCastLoweringPatterns(patterns);99}100 101void transform::ApplyLowerBroadcastPatternsOp::populatePatterns(102 RewritePatternSet &patterns) {103 populateVectorBroadcastLoweringPatterns(patterns);104}105 106void transform::ApplyLowerContractionPatternsOp::populatePatterns(107 RewritePatternSet &patterns) {108 populateVectorContractLoweringPatterns(patterns, getLoweringStrategy(),109 /*benefit=*/1,110 /*disableOuterProductLowering=*/true);111}112 113void transform::ApplyLowerMasksPatternsOp::populatePatterns(114 RewritePatternSet &patterns) {115 populateVectorMaskOpLoweringPatterns(patterns);116}117 118void transform::ApplyLowerMaskedTransfersPatternsOp::populatePatterns(119 RewritePatternSet &patterns) {120 populateVectorMaskLoweringPatternsForSideEffectingOps(patterns);121}122 123void transform::ApplyMaterializeMasksPatternsOp::populatePatterns(124 RewritePatternSet &patterns) {125 populateVectorMaskMaterializationPatterns(patterns,126 /*force32BitVectorIndices=*/false);127}128 129void transform::ApplyLowerMultiReductionPatternsOp::populatePatterns(130 RewritePatternSet &patterns) {131 vector::VectorTransformsOptions vectorTransformOptions;132 vectorTransformOptions.setVectorMultiReductionLowering(getLoweringStrategy());133 vector::populateVectorMultiReductionLoweringPatterns(134 patterns, vectorTransformOptions.vectorMultiReductionLowering);135}136 137void transform::ApplyLowerOuterProductPatternsOp::populatePatterns(138 RewritePatternSet &patterns) {139 populateVectorOuterProductLoweringPatterns(patterns);140}141 142void transform::ApplyLowerGatherPatternsOp::populatePatterns(143 RewritePatternSet &patterns) {144 vector::populateVectorGatherLoweringPatterns(patterns);145}146 147void transform::ApplyUnrollFromElementsPatternsOp::populatePatterns(148 RewritePatternSet &patterns) {149 vector::populateVectorFromElementsUnrollPatterns(patterns);150}151 152void transform::ApplyUnrollToElementsPatternsOp::populatePatterns(153 RewritePatternSet &patterns) {154 vector::populateVectorToElementsUnrollPatterns(patterns);155}156 157void transform::ApplyLowerScanPatternsOp::populatePatterns(158 RewritePatternSet &patterns) {159 vector::populateVectorScanLoweringPatterns(patterns);160}161 162void transform::ApplyLowerShapeCastPatternsOp::populatePatterns(163 RewritePatternSet &patterns) {164 vector::populateVectorShapeCastLoweringPatterns(patterns);165}166 167void transform::ApplyLowerTransferPatternsOp::populatePatterns(168 RewritePatternSet &patterns) {169 vector::populateVectorTransferLoweringPatterns(patterns,170 getMaxTransferRank());171}172 173void transform::ApplyLowerTransposePatternsOp::populatePatterns(174 RewritePatternSet &patterns) {175 vector::populateVectorTransposeLoweringPatterns(patterns,176 getLoweringStrategy());177 if (getAvx2LoweringStrategy()) {178 auto avx2LoweringOptions =179 x86vector::avx2::LoweringOptions().setTransposeOptions(180 x86vector::avx2::TransposeLoweringOptions()181 .lower4x8xf32(true)182 .lower8x8xf32(true));183 x86vector::avx2::populateSpecializedTransposeLoweringPatterns(184 patterns, avx2LoweringOptions, /*benefit=*/10);185 }186}187 188void transform::ApplyLowerInterleavePatternsOp::populatePatterns(189 RewritePatternSet &patterns) {190 vector::populateVectorInterleaveLoweringPatterns(patterns);191}192 193void transform::ApplyInterleaveToShufflePatternsOp::populatePatterns(194 RewritePatternSet &patterns) {195 vector::populateVectorInterleaveToShufflePatterns(patterns);196}197 198void transform::ApplyRewriteNarrowTypePatternsOp::populatePatterns(199 RewritePatternSet &patterns) {200 populateVectorNarrowTypeRewritePatterns(patterns);201 populateVectorTransposeNarrowTypeRewritePatterns(patterns);202}203 204void transform::ApplySplitTransferFullPartialPatternsOp::populatePatterns(205 RewritePatternSet &patterns) {206 vector::VectorTransformsOptions vectorTransformOptions;207 vectorTransformOptions.setVectorTransferSplit(getSplitTransferStrategy());208 populateVectorTransferFullPartialPatterns(patterns, vectorTransformOptions);209}210 211void transform::ApplyTransferToScfPatternsOp::populatePatterns(212 RewritePatternSet &patterns) {213 VectorTransferToSCFOptions vectorTransferToSCFOptions =214 VectorTransferToSCFOptions()215 .enableFullUnroll(getFullUnroll())216 .setTargetRank(getMaxTransferRank());217 populateVectorToSCFConversionPatterns(patterns, vectorTransferToSCFOptions);218}219 220void transform::ApplySinkVectorPatternsOp::populatePatterns(221 RewritePatternSet &patterns) {222 vector::populateSinkVectorOpsPatterns(patterns);223}224 225void transform::ApplySinkVectorMemPatternsOp::populatePatterns(226 RewritePatternSet &patterns) {227 vector::populateSinkVectorMemOpsPatterns(patterns);228}229 230//===----------------------------------------------------------------------===//231// Transform op registration232//===----------------------------------------------------------------------===//233 234namespace {235/// Registers new ops and declares PDL as dependent dialect since the additional236/// ops are using PDL types for operands and results.237class VectorTransformDialectExtension238 : public transform::TransformDialectExtension<239 VectorTransformDialectExtension> {240public:241 MLIR_DEFINE_EXPLICIT_INTERNAL_INLINE_TYPE_ID(VectorTransformDialectExtension)242 243 VectorTransformDialectExtension() {244 declareGeneratedDialect<vector::VectorDialect>();245 declareGeneratedDialect<LLVM::LLVMDialect>();246 registerTransformOps<247#define GET_OP_LIST248#include "mlir/Dialect/Vector/TransformOps/VectorTransformOps.cpp.inc"249 >();250 }251};252} // namespace253 254#define GET_OP_CLASSES255#include "mlir/Dialect/Vector/TransformOps/VectorTransformOps.cpp.inc"256 257void mlir::vector::registerTransformDialectExtension(258 DialectRegistry ®istry) {259 registry.addExtensions<VectorTransformDialectExtension>();260}261