brintos

brintos / llvm-project-archived public Read only

0
0
Text · 9.5 KiB · 7faa222 Raw
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 &registry) {259  registry.addExtensions<VectorTransformDialectExtension>();260}261