88
99from ..common import vendors
1010from . import backend_utils
11+ from .backend_utils import BackendEventBase
1112
1213
1314class BackendState :
@@ -42,7 +43,63 @@ def __init__(self):
4243_state = BackendState ()
4344
4445
45- class BackendArchEvent :
46+ class TritonVersionEvent (BackendEventBase ):
47+ _instance = None
48+ has_version_spec = False
49+
50+ def __new__ (cls , * args , ** kwargs ):
51+ if cls ._instance is None :
52+ cls ._instance = super ().__new__ (cls )
53+ return cls ._instance
54+
55+ def __init__ (self , version = None ):
56+ self .has_version_spec = False
57+ self .version = version if version is not None else self .get_version ()
58+ self .dir = self .get_version_spec_dir ()
59+ if self .dir and Path (self .dir ).exists ():
60+ self .module = self .get_version_spec_module ()
61+ self .has_version_spec = True
62+
63+ def is_available (self ):
64+ return self .has_version_spec
65+
66+ def get_version_spec_dir (self , path = None ):
67+ dir_name = f"triton_{ self .version } "
68+ backend_path = Path (path or _state .vendor_module .__path__ [0 ])
69+ backend_path = backend_path .parent if backend_path .is_file () else backend_path
70+ excluded = ("ops" , "fused" )
71+ return {
72+ p .name : str (p )
73+ for p in backend_path .iterdir ()
74+ if p .is_dir () and p .name not in excluded and not p .name .startswith ("_" )
75+ }.get (dir_name , None )
76+
77+ def get_functions_from_module (self , module ):
78+ return inspect .getmembers (module , inspect .isfunction ) if module else []
79+
80+ def get_version_spec_module (self ):
81+ module_name = f"triton_{ self .version } "
82+ path_dir = os .path .dirname (self .dir )
83+ sys .path .insert (0 , str (path_dir ))
84+ version_module = importlib .import_module (module_name )
85+ sys .path .remove (str (path_dir ))
86+ return version_module
87+
88+ def get_ops (self ):
89+ return self .get_version_ops ()
90+
91+ def get_version_ops (self ):
92+ pass
93+
94+ def get_version (self ):
95+ try :
96+ import triton
97+ except ImportError :
98+ return None
99+ return triton .__version__
100+
101+
102+ class BackendArchEvent (BackendEventBase ):
46103 has_arch : bool = False
47104 _instance = None
48105 _initialized : bool = False
@@ -67,6 +124,9 @@ def __init__(self, backend=None):
67124 self .autotune_configs = self .get_autotune_configs ()
68125 self .heuristics_configs = self .get_heuristics_configs ()
69126
127+ def is_available (self ):
128+ return self .has_arch
129+
70130 def get_functions_from_module (self , module ):
71131 return inspect .getmembers (module , inspect .isfunction ) if module else []
72132
@@ -126,6 +186,10 @@ def get_arch_module(self):
126186 sys .path .remove (str (path_dir ))
127187 return current_arch_module
128188
189+ def get_ops (self ):
190+ """Provide a unified interface for the upper layer"""
191+ return self .get_arch_ops ()
192+
129193 def get_arch_ops (self ):
130194 arch_specialized_ops = []
131195 sys .path .append (self .current_arch_path )
@@ -147,6 +211,23 @@ def get_arch_ops(self):
147211 return arch_specialized_ops
148212
149213
214+ class SpecOpRegistrar :
215+ def __init__ (self , _globals ):
216+ self ._globals = _globals
217+
218+ def apply (self ):
219+ spec_events = self ._get_specific_events ()
220+ for event in spec_events :
221+ if not event .is_available ():
222+ continue
223+ operators = event .get_ops ()
224+ for fn_name , fn in operators :
225+ self ._globals [fn_name ] = fn
226+
227+ def _get_specific_events (self ):
228+ return (BackendArchEvent (), TritonVersionEvent ())
229+
230+
150231def _import_module_safe (module_name , vendor_name , module_type ):
151232 """Helper to import a module with proper error handling."""
152233 try :
@@ -295,6 +376,11 @@ def get_customized_ops(vendor_name=None):
295376 return _state .customized_ops
296377
297378
379+ def get_ops (vendor_name = None ):
380+ """Provide a unified interface for the upper layer"""
381+ return get_customized_ops (vendor_name )
382+
383+
298384def get_unused_ops (vendor_name = None ):
299385 global vendor_module # noqa: F824
300386 get_vendor_module (vendor_name )
0 commit comments