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 ®istry) 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