Skip to content

Commit c9f1046

Browse files
authored
[Metax][TLE] Add Metax TLE support (#742)
1 parent 4ea86a6 commit c9f1046

50 files changed

Lines changed: 4271 additions & 9 deletions

Some content is hidden

Large Commits have some content hidden by default. Use the searchbox below for content that may be hidden.

.github/workflows/metax3.6-build-and-test.yml

Lines changed: 13 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -95,3 +95,16 @@ jobs:
9595
not test_mma_remark and \
9696
not test_remark_vectorization and \
9797
not test_passing_tuple_to_make_tensor_descriptor"
98+
99+
- name: FlagTree TLE Tile Unit Test on Metax
100+
if: steps.check_backend.outputs.should_skip != 'true'
101+
shell: bash
102+
run: |
103+
set -x
104+
source ~/env.sh
105+
cd python/test/tle/unit
106+
python3 -m pytest -s \
107+
test_extract_tile_dynamic_index.py \
108+
test_extract_tile_static_index.py \
109+
test_insert_tile_dynamic_index.py \
110+
test_insert_tile_static_index.py

CMakeLists.txt

Lines changed: 12 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -61,7 +61,11 @@ elseif(FLAGTREE_BACKEND STREQUAL "hcu")
6161
add_definitions(-D__HCU__)
6262
elseif(FLAGTREE_BACKEND STREQUAL "metax")
6363
add_definitions(-DUSE_MACA)
64-
# TODO(metax): support TLE
64+
# metax use mctle to replace tle
65+
option(BUILD_MCTLE "use maca triton language extensions" ON)
66+
list(APPEND TRITON_PLUGIN_NAMES "mctle")
67+
add_definitions(-D__MCTLE__)
68+
6569
set(FLAGTREE_TLE OFF)
6670
remove_definitions(-D__TLE__)
6771
list(REMOVE_ITEM LLVM_TABLEGEN_FLAGS -D__TLE__)
@@ -324,6 +328,11 @@ endif()
324328
include_directories(${PROJECT_SOURCE_DIR}/third_party)
325329
include_directories(${PROJECT_BINARY_DIR}/third_party) # Tablegen'd files
326330
331+
if(BUILD_MCTLE)
332+
include_directories(${PROJECT_SOURCE_DIR}/third_party/metax/plugin)
333+
include_directories(${PROJECT_BINARY_DIR}/third_party/metax/plugin)
334+
endif()
335+
327336
# link_directories(${LLVM_LIBRARY_DIR})
328337
if (FLAGTREE_BACKEND MATCHES "^(xpu|cambricon|aipu|tsingmicro|enflame|rpu|thrive|tileir)$")
329338
include_directories(${PROJECT_SOURCE_DIR}/include)
@@ -612,6 +621,8 @@ if(TRITON_BUILD_PYTHON_MODULE)
612621
Python3::Module
613622
pybind11::headers
614623
)
624+
elseif(FLAGTREE_BACKEND STREQUAL "metax" AND BUILD_MCTLE)
625+
list(APPEND TRITON_LIBRARIES MLIRTargetLLVMIRImport)
615626
endif()
616627
617628
if(CMAKE_SYSTEM_PROCESSOR MATCHES "aarch64" OR # Linux arm64

python/setup_tools/setup_helper.py

Lines changed: 2 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -607,9 +607,9 @@ def uninstall_triton():
607607
cache.store(
608608
file="metaxTritonPlugin.so",
609609
condition=("metax" == flagtree_backend) and (not configs.flagtree_plugin),
610-
url="https://baai-cp-web.ks3-cn-beijing.ksyuncs.com/trans/metaxTritonPlugin-cpython3.12-x86_64_v0.6.0.tar.gz",
610+
url="https://baai-cp-web.ks3-cn-beijing.ksyuncs.com/trans/metaxTritonPlugin-cpython3.12-x86_64_v0.6.1.tar.gz",
611611
copy_dst_path=f"third_party/{flagtree_backend}",
612-
md5_digest="415a08bd",
612+
md5_digest="afb7ab8f",
613613
)
614614

615615
# thrive

third_party/metax/CMakeLists.txt

Lines changed: 7 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -10,7 +10,14 @@ else()
1010
endif()
1111

1212
if(TRITON_BUILD_PYTHON_MODULE)
13+
if(BUILD_MCTLE)
14+
include_directories(${CMAKE_CURRENT_SOURCE_DIR}/plugin)
15+
include_directories(${CMAKE_CURRENT_BINARY_DIR}/plugin)
16+
add_subdirectory(plugin/mctle)
17+
endif()
18+
1319
if(FLAGTREE_PLUGIN)
20+
include_directories(plugin)
1421
add_subdirectory(plugin)
1522
else()
1623
find_library(BackendTritonPluginLib

third_party/metax/backend/compiler.py

Lines changed: 14 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -5,6 +5,11 @@
55
enable_dist = True
66
except ImportError:
77
enable_dist = False
8+
try:
9+
from triton._C.libtriton import mctle
10+
enable_mctle = True
11+
except ImportError:
12+
enable_mctle = False
813
from triton import knobs
914

1015
from dataclasses import dataclass, field
@@ -218,6 +223,8 @@ def get_module_map(self) -> Dict[str, ModuleType]:
218223

219224
def load_dialects(self, ctx):
220225
metax.load_dialects(ctx)
226+
if enable_mctle:
227+
mctle.load_dialects(ctx)
221228
if enable_dist:
222229
distributed.ir.load_dialects(ctx)
223230

@@ -264,6 +271,13 @@ def make_ttgir(mod, metadata, opt, capability):
264271
passes.ttgpuir.add_f32_dot_tc(pm, emuTF32)
265272
passes.ttgpuir.add_remove_layout_conversions(pm)
266273
passes.ttgpuir.add_optimize_thread_locality(pm)
274+
if enable_mctle:
275+
#mctle.passes.add_reject_dot_op(pm)
276+
mctle.passes.add_early_assign_memory_space(pm)
277+
mctle.passes.add_select_encodings(pm)
278+
mctle.passes.add_insert_local_pointer_barriers(pm)
279+
mctle.passes.add_optimize_local_pointer_loads(pm)
280+
mctle.passes.add_optimize_local_pointer_stores(pm)
267281
if opt.pipeline == "cpasync" or opt.pipeline == "cpasync-mixed":
268282
disable_prefetch = True
269283
metax.passes.ttgpuir.add_tritonmetaxgpu_change_layout_for_int8_pass(pm, opt.num_stages, opt.pipeline)
Lines changed: 96 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,96 @@
1+
get_property(triton_libs GLOBAL PROPERTY TRITON_LIBS)
2+
3+
add_llvm_executable(triton-opt triton-opt.cpp PARTIAL_SOURCES_INTENDED)
4+
5+
# TODO: what's this?
6+
llvm_update_compile_flags(triton-opt)
7+
target_link_libraries(triton-opt PRIVATE
8+
${triton_libs}
9+
# tests
10+
TritonTestAnalysis
11+
TritonTestDialect
12+
TritonAMDGPUTestAnalysis
13+
TritonTestProton
14+
# MLIR core
15+
MLIROptLib
16+
MLIRPass
17+
MLIRRegisterAllDialects
18+
MLIRRegisterAllPasses
19+
MLIRTransforms
20+
)
21+
22+
mlir_check_all_link_libraries(triton-opt)
23+
24+
add_llvm_executable(triton-reduce triton-reduce.cpp PARTIAL_SOURCES_INTENDED)
25+
mlir_check_all_link_libraries(triton-reduce)
26+
27+
llvm_update_compile_flags(triton-reduce)
28+
target_link_libraries(triton-reduce PRIVATE
29+
${triton_libs}
30+
# tests
31+
TritonTestAnalysis
32+
TritonTestDialect
33+
TritonAMDGPUTestAnalysis
34+
TritonTestProton
35+
# MLIR core
36+
MLIRReduceLib
37+
MLIRPass
38+
MLIRRegisterAllDialects
39+
MLIRRegisterAllPasses
40+
MLIRTransforms
41+
)
42+
43+
mlir_check_all_link_libraries(triton-reduce)
44+
45+
add_llvm_executable(triton-lsp triton-lsp.cpp PARTIAL_SOURCES_INTENDED)
46+
47+
llvm_update_compile_flags(triton-lsp)
48+
target_link_libraries(triton-lsp PRIVATE
49+
${triton_libs}
50+
# tests
51+
TritonTestAnalysis
52+
TritonTestDialect
53+
TritonAMDGPUTestAnalysis
54+
TritonTestProton
55+
# MLIR core
56+
MLIRLspServerLib
57+
MLIRPass
58+
MLIRRegisterAllDialects
59+
MLIRRegisterAllPasses
60+
MLIRTransforms
61+
)
62+
63+
mlir_check_all_link_libraries(triton-lsp)
64+
65+
66+
add_llvm_executable(triton-llvm-opt
67+
triton-llvm-opt.cpp
68+
69+
PARTIAL_SOURCES_INTENDED
70+
DEPENDS
71+
intrinsics_gen
72+
SUPPORT_PLUGINS
73+
)
74+
target_link_libraries(triton-llvm-opt PRIVATE
75+
TritonLLVMIR
76+
77+
LLVMAnalysis
78+
LLVMCore
79+
LLVMSupport
80+
LLVMOption
81+
LLVMCodeGen
82+
)
83+
export_executable_symbols_for_plugins(triton-llvm-opt)
84+
85+
86+
add_llvm_executable(triton-tensor-layout triton-tensor-layout.cpp PARTIAL_SOURCES_INTENDED)
87+
target_link_libraries(triton-tensor-layout PRIVATE
88+
${triton_libs}
89+
TritonTestAnalysis
90+
TritonTestDialect
91+
TritonTestProton
92+
TritonAMDGPUTestAnalysis
93+
MLIRRegisterAllDialects
94+
MLIRRegisterAllPasses
95+
MLIRTransforms
96+
)
Lines changed: 161 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,161 @@
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 &registry) {
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

Comments
 (0)