|
| 1 | +#pragma once |
| 2 | +#include "amd/include/Dialect/TritonAMDGPU/IR/Dialect.h" |
| 3 | +#include "amd/include/TritonAMDGPUTransforms/Passes.h" |
| 4 | +#include "nvidia/include/Dialect/NVGPU/IR/Dialect.h" |
| 5 | +#include "nvidia/include/Dialect/NVWS/IR/Dialect.h" |
| 6 | +#include "proton/Dialect/include/Conversion/ProtonGPUToLLVM/Passes.h" |
| 7 | +#include "proton/Dialect/include/Conversion/ProtonGPUToLLVM/ProtonAMDGPUToLLVM/Passes.h" |
| 8 | +#include "proton/Dialect/include/Conversion/ProtonGPUToLLVM/ProtonNvidiaGPUToLLVM/Passes.h" |
| 9 | +#include "proton/Dialect/include/Conversion/ProtonToProtonGPU/Passes.h" |
| 10 | +#include "proton/Dialect/include/Dialect/Proton/IR/Dialect.h" |
| 11 | +#include "proton/Dialect/include/Dialect/ProtonGPU/IR/Dialect.h" |
| 12 | +#include "proton/Dialect/include/Dialect/ProtonGPU/Transforms/Passes.h" |
| 13 | +#include "triton/Dialect/Gluon/Transforms/Passes.h" |
| 14 | +#include "triton/Dialect/Triton/IR/Dialect.h" |
| 15 | +#include "triton/Dialect/TritonGPU/IR/Dialect.h" |
| 16 | +#include "triton/Dialect/TritonInstrument/IR/Dialect.h" |
| 17 | +#include "triton/Dialect/TritonNvidiaGPU/IR/Dialect.h" |
| 18 | +#ifdef __MCTLE__ |
| 19 | +#include "mctle/dialect/include/IR/Dialect.h" |
| 20 | +#include "mctle/dialect/include/Transforms/Passes.h" |
| 21 | +#endif |
| 22 | + |
| 23 | +// Below headers will allow registration to ROCm passes |
| 24 | +#include "TritonAMDGPUToLLVM/Passes.h" |
| 25 | +#include "TritonAMDGPUTransforms/Passes.h" |
| 26 | +#include "TritonAMDGPUTransforms/TritonGPUConversion.h" |
| 27 | + |
| 28 | +#include "triton/Dialect/Triton/Transforms/Passes.h" |
| 29 | +#include "triton/Dialect/TritonGPU/Transforms/Passes.h" |
| 30 | +#include "triton/Dialect/TritonInstrument/Transforms/Passes.h" |
| 31 | +#include "triton/Dialect/TritonNvidiaGPU/Transforms/Passes.h" |
| 32 | + |
| 33 | +#include "nvidia/hopper/include/Transforms/Passes.h" |
| 34 | +#include "nvidia/include/Dialect/NVWS/Transforms/Passes.h" |
| 35 | +#include "nvidia/include/NVGPUToLLVM/Passes.h" |
| 36 | +#include "nvidia/include/TritonNVIDIAGPUToLLVM/Passes.h" |
| 37 | +#include "triton/Conversion/TritonGPUToLLVM/Passes.h" |
| 38 | +#include "triton/Conversion/TritonToTritonGPU/Passes.h" |
| 39 | +#include "triton/Target/LLVMIR/Passes.h" |
| 40 | + |
| 41 | +#include "mlir/Dialect/LLVMIR/NVVMDialect.h" |
| 42 | +#include "mlir/Dialect/LLVMIR/ROCDLDialect.h" |
| 43 | +#include "mlir/Dialect/LLVMIR/Transforms/InlinerInterfaceImpl.h" |
| 44 | +#include "mlir/InitAllPasses.h" |
| 45 | + |
| 46 | +#include "mlir/Conversion/ArithToLLVM/ArithToLLVM.h" |
| 47 | +#include "mlir/Conversion/ControlFlowToLLVM/ControlFlowToLLVM.h" |
| 48 | +#include "mlir/Conversion/MathToLLVM/MathToLLVM.h" |
| 49 | +#include "mlir/Conversion/NVVMToLLVM/NVVMToLLVM.h" |
| 50 | +#include "mlir/Conversion/UBToLLVM/UBToLLVM.h" |
| 51 | + |
| 52 | +namespace mlir { |
| 53 | +namespace test { |
| 54 | +void registerTestAliasPass(); |
| 55 | +void registerTestAlignmentPass(); |
| 56 | +void registerAMDTestAlignmentPass(); |
| 57 | +void registerTestAllocationPass(); |
| 58 | +void registerTestMembarPass(); |
| 59 | +void registerTestAMDGPUMembarPass(); |
| 60 | +void registerTestTritonAMDGPURangeAnalysis(); |
| 61 | +void registerTestLoopPeelingPass(); |
| 62 | +namespace proton { |
| 63 | +void registerTestScopeIdAllocationPass(); |
| 64 | +} // namespace proton |
| 65 | +} // namespace test |
| 66 | +} // namespace mlir |
| 67 | + |
| 68 | +inline void registerTritonDialects(mlir::DialectRegistry ®istry) { |
| 69 | + mlir::registerAllPasses(); |
| 70 | + mlir::triton::registerTritonPasses(); |
| 71 | + mlir::triton::gpu::registerTritonGPUPasses(); |
| 72 | + mlir::triton::nvidia_gpu::registerTritonNvidiaGPUPasses(); |
| 73 | + mlir::triton::instrument::registerTritonInstrumentPasses(); |
| 74 | + mlir::triton::gluon::registerGluonPasses(); |
| 75 | +#ifdef __MCTLE__ |
| 76 | + mlir::triton::mctle::registerPasses(); |
| 77 | +#endif |
| 78 | + mlir::test::registerTestAliasPass(); |
| 79 | + mlir::test::registerTestAlignmentPass(); |
| 80 | + mlir::test::registerAMDTestAlignmentPass(); |
| 81 | + mlir::test::registerTestAllocationPass(); |
| 82 | + mlir::test::registerTestMembarPass(); |
| 83 | + mlir::test::registerTestLoopPeelingPass(); |
| 84 | + mlir::test::registerTestAMDGPUMembarPass(); |
| 85 | + mlir::test::registerTestTritonAMDGPURangeAnalysis(); |
| 86 | + mlir::triton::registerConvertTritonToTritonGPUPass(); |
| 87 | + mlir::triton::registerRelayoutTritonGPUPass(); |
| 88 | + mlir::triton::gpu::registerAllocateSharedMemoryPass(); |
| 89 | + mlir::triton::gpu::registerTritonGPUAllocateWarpGroups(); |
| 90 | + mlir::triton::gpu::registerTritonGPUGlobalScratchAllocationPass(); |
| 91 | + mlir::triton::registerConvertWarpSpecializeToLLVM(); |
| 92 | + mlir::triton::registerConvertTritonGPUToLLVMPass(); |
| 93 | + mlir::triton::registerConvertNVGPUToLLVMPass(); |
| 94 | + mlir::triton::registerAllocateSharedMemoryNvPass(); |
| 95 | + mlir::registerLLVMDIScope(); |
| 96 | + mlir::LLVM::registerInlinerInterface(registry); |
| 97 | + mlir::NVVM::registerInlinerInterface(registry); |
| 98 | + mlir::registerLLVMDILocalVariable(); |
| 99 | + |
| 100 | + // TritonAMDGPUToLLVM passes |
| 101 | + mlir::triton::registerAllocateAMDGPUSharedMemory(); |
| 102 | + mlir::triton::registerConvertTritonAMDGPUToLLVM(); |
| 103 | + mlir::triton::registerConvertBuiltinFuncToLLVM(); |
| 104 | + mlir::triton::registerOptimizeAMDLDSUsage(); |
| 105 | + |
| 106 | + mlir::ub::registerConvertUBToLLVMInterface(registry); |
| 107 | + mlir::registerConvertNVVMToLLVMInterface(registry); |
| 108 | + mlir::registerConvertMathToLLVMInterface(registry); |
| 109 | + mlir::cf::registerConvertControlFlowToLLVMInterface(registry); |
| 110 | + mlir::arith::registerConvertArithToLLVMInterface(registry); |
| 111 | + |
| 112 | + // TritonAMDGPUTransforms passes |
| 113 | + mlir::registerTritonAMDGPUAccelerateMatmul(); |
| 114 | + mlir::registerTritonAMDGPUOptimizeEpilogue(); |
| 115 | + mlir::registerTritonAMDGPUHoistLayoutConversions(); |
| 116 | + mlir::registerTritonAMDGPUReorderInstructions(); |
| 117 | + mlir::registerTritonAMDGPUBlockPingpong(); |
| 118 | + mlir::registerTritonAMDGPUPipeline(); |
| 119 | + mlir::registerTritonAMDGPUScheduleLoops(); |
| 120 | + mlir::registerTritonAMDGPUCanonicalizePointers(); |
| 121 | + mlir::registerTritonAMDGPUConvertToBufferOps(); |
| 122 | + mlir::registerTritonAMDGPUInThreadTranspose(); |
| 123 | + mlir::registerTritonAMDGPUCoalesceAsyncCopy(); |
| 124 | + mlir::registerTritonAMDGPUUpdateAsyncWaitCount(); |
| 125 | + mlir::triton::registerTritonAMDGPUInsertInstructionSchedHints(); |
| 126 | + mlir::triton::registerTritonAMDGPULowerInstructionSchedHints(); |
| 127 | + mlir::registerTritonAMDFoldTrueCmpI(); |
| 128 | + mlir::triton::amdgpu::registerTritonAMDGPUOptimizeDotOperands(); |
| 129 | + |
| 130 | + // NVWS passes |
| 131 | + mlir::triton::registerNVWSTransformsPasses(); |
| 132 | + |
| 133 | + // NVGPU transform passes |
| 134 | + mlir::registerNVHopperTransformsPasses(); |
| 135 | + |
| 136 | + // Proton passes |
| 137 | + mlir::test::proton::registerTestScopeIdAllocationPass(); |
| 138 | + mlir::triton::proton::registerConvertProtonToProtonGPU(); |
| 139 | + mlir::triton::proton::gpu::registerConvertProtonNvidiaGPUToLLVM(); |
| 140 | + mlir::triton::proton::gpu::registerConvertProtonAMDGPUToLLVM(); |
| 141 | + mlir::triton::proton::gpu::registerAllocateProtonSharedMemoryPass(); |
| 142 | + mlir::triton::proton::gpu::registerAllocateProtonGlobalScratchBufferPass(); |
| 143 | + mlir::triton::proton::gpu::registerScheduleBufferStorePass(); |
| 144 | + mlir::triton::proton::gpu::registerAddSchedBarriersPass(); |
| 145 | + |
| 146 | + registry.insert< |
| 147 | + mlir::triton::TritonDialect, mlir::cf::ControlFlowDialect, |
| 148 | + mlir::triton::nvidia_gpu::TritonNvidiaGPUDialect, |
| 149 | + mlir::triton::gpu::TritonGPUDialect, |
| 150 | + mlir::triton::instrument::TritonInstrumentDialect, |
| 151 | + mlir::math::MathDialect, mlir::arith::ArithDialect, mlir::scf::SCFDialect, |
| 152 | + mlir::gpu::GPUDialect, mlir::LLVM::LLVMDialect, mlir::NVVM::NVVMDialect, |
| 153 | + mlir::triton::nvgpu::NVGPUDialect, mlir::triton::nvws::NVWSDialect, |
| 154 | + mlir::triton::amdgpu::TritonAMDGPUDialect, |
| 155 | + mlir::triton::proton::ProtonDialect, |
| 156 | + mlir::triton::proton::gpu::ProtonGPUDialect, mlir::ROCDL::ROCDLDialect, |
| 157 | +#ifdef __MCTLE__ |
| 158 | + mlir::triton::mctle::McTleDialect, |
| 159 | +#endif |
| 160 | + mlir::triton::gluon::GluonDialect>(); |
| 161 | +} |
0 commit comments