brintos

brintos / llvm-project-archived public Read only

0
0
Text · 6.2 KiB · f958edf Raw
152 lines · cpp
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-exception6//7//===----------------------------------------------------------------------===//8 9#include "mlir/Conversion/VectorToLLVM/ConvertVectorToLLVMPass.h"10 11#include "mlir/Conversion/LLVMCommon/ConversionTarget.h"12#include "mlir/Conversion/LLVMCommon/TypeConverter.h"13#include "mlir/Dialect/AMX/AMXDialect.h"14#include "mlir/Dialect/AMX/Transforms.h"15#include "mlir/Dialect/Arith/IR/Arith.h"16#include "mlir/Dialect/ArmNeon/ArmNeonDialect.h"17#include "mlir/Dialect/ArmNeon/Transforms.h"18#include "mlir/Dialect/ArmSVE/IR/ArmSVEDialect.h"19#include "mlir/Dialect/ArmSVE/Transforms/Transforms.h"20#include "mlir/Dialect/LLVMIR/LLVMDialect.h"21#include "mlir/Dialect/MemRef/IR/MemRef.h"22#include "mlir/Dialect/Tensor/IR/Tensor.h"23#include "mlir/Dialect/Vector/Transforms/LoweringPatterns.h"24#include "mlir/Dialect/Vector/Transforms/VectorRewritePatterns.h"25#include "mlir/Dialect/X86Vector/Transforms.h"26#include "mlir/Dialect/X86Vector/X86VectorDialect.h"27#include "mlir/Pass/Pass.h"28#include "mlir/Transforms/GreedyPatternRewriteDriver.h"29 30namespace mlir {31#define GEN_PASS_DEF_CONVERTVECTORTOLLVMPASS32#include "mlir/Conversion/Passes.h.inc"33} // namespace mlir34 35using namespace mlir;36using namespace mlir::vector;37 38namespace {39struct ConvertVectorToLLVMPass40    : public impl::ConvertVectorToLLVMPassBase<ConvertVectorToLLVMPass> {41 42  using Base::Base;43 44  // Override explicitly to allow conditional dialect dependence.45  void getDependentDialects(DialectRegistry &registry) const override {46    registry.insert<LLVM::LLVMDialect>();47    registry.insert<arith::ArithDialect>();48    registry.insert<memref::MemRefDialect>();49    registry.insert<tensor::TensorDialect>();50    if (armNeon)51      registry.insert<arm_neon::ArmNeonDialect>();52    if (armSVE)53      registry.insert<arm_sve::ArmSVEDialect>();54    if (amx)55      registry.insert<amx::AMXDialect>();56    if (x86Vector)57      registry.insert<x86vector::X86VectorDialect>();58  }59  void runOnOperation() override;60};61} // namespace62 63void ConvertVectorToLLVMPass::runOnOperation() {64  // Perform progressive lowering of operations on slices and all contraction65  // operations. Also materializes masks, lowers vector.step, rank-reduces FMA,66  // applies folding and DCE.67  {68    RewritePatternSet patterns(&getContext());69    populateVectorToVectorCanonicalizationPatterns(patterns);70    populateVectorBitCastLoweringPatterns(patterns);71    populateVectorBroadcastLoweringPatterns(patterns);72    populateVectorContractLoweringPatterns(patterns, vectorContractLowering);73    if (vectorContractLowering == vector::VectorContractLowering::LLVMIntr) {74      // This pattern creates a dependency on the LLVM dialect, hence we don't75      // include it in `populateVectorContractLoweringPatterns` that is part of76      // the Vector dialect (and should not depend on LLVM).77      populateVectorContractToMatrixMultiply(patterns);78    }79    populateVectorMaskOpLoweringPatterns(patterns);80    populateVectorShapeCastLoweringPatterns(patterns);81    populateVectorInterleaveLoweringPatterns(patterns);82    populateVectorTransposeLoweringPatterns(patterns, vectorTransposeLowering);83    if (vectorTransposeLowering == vector::VectorTransposeLowering::LLVMIntr) {84      // This pattern creates a dependency on the LLVM dialect, hence we don't85      // include it in `populateVectorTransposeLoweringPatterns` that is part of86      // the Vector dialect (and should not depend on LLVM).87      populateVectorTransposeToFlatTranspose(patterns);88    }89    // Vector transfer ops with rank > 1 should be lowered with VectorToSCF.90    populateVectorTransferLoweringPatterns(patterns, /*maxTransferRank=*/1);91    populateVectorMaskMaterializationPatterns(patterns,92                                              force32BitVectorIndices);93    populateVectorInsertExtractStridedSliceTransforms(patterns);94    populateVectorStepLoweringPatterns(patterns);95    populateVectorRankReducingFMAPattern(patterns);96    populateVectorGatherLoweringPatterns(patterns);97    populateVectorFromElementsUnrollPatterns(patterns);98    populateVectorToElementsUnrollPatterns(patterns);99    if (armI8MM) {100      if (armNeon)101        arm_neon::populateLowerContractionToNeonI8MMPatterns(patterns);102      if (armSVE)103        populateLowerContractionToSVEI8MMPatterns(patterns);104    }105    if (armBF16) {106      if (armNeon)107        arm_neon::populateLowerContractionToNeonBFMMLAPatterns(patterns);108      if (armSVE)109        populateLowerContractionToSVEBFMMLAPatterns(patterns);110    }111    (void)applyPatternsGreedily(getOperation(), std::move(patterns));112  }113 114  // Convert to the LLVM IR dialect.115  LowerToLLVMOptions options(&getContext());116  LLVMTypeConverter converter(&getContext(), options);117  RewritePatternSet patterns(&getContext());118  populateVectorTransferLoweringPatterns(patterns);119  populateVectorToLLVMConversionPatterns(120      converter, patterns, reassociateFPReductions, force32BitVectorIndices,121      useVectorAlignment);122 123  // Architecture specific augmentations.124  LLVMConversionTarget target(getContext());125  target.addLegalDialect<arith::ArithDialect>();126  target.addLegalDialect<memref::MemRefDialect>();127  target.addLegalOp<UnrealizedConversionCastOp>();128 129  if (armNeon) {130    // TODO: we may or may not want to include in-dialect lowering to131    // LLVM-compatible operations here. So far, all operations in the dialect132    // can be translated to LLVM IR so there is no conversion necessary.133    target.addLegalDialect<arm_neon::ArmNeonDialect>();134  }135  if (armSVE) {136    configureArmSVELegalizeForExportTarget(target);137    populateArmSVELegalizeForLLVMExportPatterns(converter, patterns);138  }139  if (amx) {140    configureAMXLegalizeForExportTarget(target);141    populateAMXLegalizeForLLVMExportPatterns(converter, patterns);142  }143  if (x86Vector) {144    configureX86VectorLegalizeForExportTarget(target);145    populateX86VectorLegalizeForLLVMExportPatterns(converter, patterns);146  }147 148  if (failed(149          applyPartialConversion(getOperation(), target, std::move(patterns))))150    signalPassFailure();151}152