Skip to content

Commit 97d8d30

Browse files
authored
[Backend]add_triton_version_event (flagos-ai#3336)
* add_triton_version_event * fix codestyle * fix bug * fix codestyle * fix code mode
1 parent c6ad911 commit 97d8d30

7 files changed

Lines changed: 115 additions & 31 deletions

File tree

src/flag_gems/__init__.py

Lines changed: 9 additions & 4 deletions
Original file line numberDiff line numberDiff line change
@@ -14,15 +14,20 @@
1414
from flag_gems.ops import * # noqa: F403
1515
from flag_gems.patches import * # noqa: F403
1616
from flag_gems.runtime import flagtune
17-
from flag_gems.runtime.register import Register
17+
from flag_gems.runtime.backend import SpecOpRegistrar
18+
from flag_gems.runtime.op_registrar import GeneralOpRegistrar
1819

1920
__version__ = "5.0.2"
2021
device = runtime.device.name
2122
vendor_name = runtime.device.vendor_name
23+
backend_info = runtime.device
2224
aten_lib = torch.library.Library("aten", "IMPL")
23-
registrar = Register
25+
26+
# Register all ops in the current backend with SpecOpRegistrar to support architecture-specialized implementations
27+
SpecOpRegistrar(globals()).apply()
28+
29+
registrar = GeneralOpRegistrar
2430
current_work_registrar = None
25-
runtime.replace_customized_ops(globals())
2631
AUTOGRAD_DISPATCH_KEY = torch._C.DispatchKey.Autograd.name
2732

2833

@@ -683,7 +688,7 @@ def __init__(self, exclude=None, include=None, record=False, once=False, path=No
683688
self.lib = torch.library.Library("aten", "IMPL")
684689
self.exclude = exclude if isinstance(exclude, (list, tuple, set, str)) else []
685690
self.include = include if isinstance(include, (list, tuple, set, str)) else []
686-
self.registrar = Register
691+
self.registrar = GeneralOpRegistrar
687692
self.record = record
688693
self.once = once
689694
self.path = path

src/flag_gems/runtime/__init__.py

Lines changed: 3 additions & 21 deletions
Original file line numberDiff line numberDiff line change
@@ -1,9 +1,9 @@
11
from . import backend, common, error
22
from .backend.device import DeviceDetector
3-
from .configloader import ConfigLoader
3+
from .configs_loader import TunedConfigLoader
44
from .flagtune import flagtune, flagtune_enabled
55

6-
config_loader = ConfigLoader()
6+
config_loader = TunedConfigLoader()
77
device = DeviceDetector()
88

99
"""
@@ -26,24 +26,6 @@ def get_heuristic_config(op_name):
2626
return config_loader.get_heuristics_config(op_name)
2727

2828

29-
def replace_customized_ops(_globals):
30-
event = backend.BackendArchEvent()
31-
arch_specific_ops = event.get_arch_ops() if event.has_arch else None
32-
extended_ops = backend.get_customized_ops(device.vendor_name)
33-
if device.vendor != common.vendors.NVIDIA:
34-
try:
35-
for fn_name, fn in extended_ops:
36-
_globals[fn_name] = fn
37-
except RuntimeError as e:
38-
error.customized_op_replace_error(e)
39-
if arch_specific_ops:
40-
try:
41-
for fn_name, fn in arch_specific_ops:
42-
_globals[fn_name] = fn
43-
except RuntimeError as e:
44-
error.customized_op_replace_error(e)
45-
46-
4729
def get_expand_config(op_name, yaml_path=None):
4830
return config_loader.get_expand_config(op_name=op_name, yaml_path=yaml_path)
4931

@@ -57,7 +39,7 @@ def ops_get_configs(op_name, pre_hook=None, yaml_path=None):
5739

5840

5941
__all__ = [
60-
"ConfigLoader",
42+
"TunedConfigLoader",
6143
"DeviceDetector",
6244
"backend",
6345
"common",

src/flag_gems/runtime/backend/__init__.py

Lines changed: 87 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -8,6 +8,7 @@
88

99
from ..common import vendors
1010
from . import backend_utils
11+
from .backend_utils import BackendEventBase
1112

1213

1314
class 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+
150231
def _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+
298384
def get_unused_ops(vendor_name=None):
299385
global vendor_module # noqa: F824
300386
get_vendor_module(vendor_name)

src/flag_gems/runtime/backend/backend_utils.py

Lines changed: 11 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -39,6 +39,17 @@ def get_tune_config(vendor_name=None, file_mode="r", file_path=None):
3939
return config
4040

4141

42+
class BackendEventBase:
43+
def __init__(self):
44+
...
45+
46+
def get_ops(self):
47+
...
48+
49+
def is_available(self):
50+
...
51+
52+
4253
@functools.lru_cache(maxsize=None)
4354
def _load_expand_config(file_path, file_mode="r"):
4455
with open(file_path, file_mode) as file:
Lines changed: 2 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -9,12 +9,12 @@
99
from .backend.device import DeviceDetector
1010

1111

12-
class ConfigLoader(object):
12+
class TunedConfigLoader(object):
1313
_instance = None
1414

1515
def __new__(cls, *args, **kargs):
1616
if cls._instance is None:
17-
cls._instance = super(ConfigLoader, cls).__new__(cls)
17+
cls._instance = super(TunedConfigLoader, cls).__new__(cls)
1818
return cls._instance
1919

2020
def __init__(self):
Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -4,7 +4,7 @@
44
from .backend.device import DeviceDetector
55

66

7-
class Register:
7+
class GeneralOpRegistrar:
88
def __init__(
99
self,
1010
config,

src/flag_gems/runtime/precision_register.py

Lines changed: 2 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -16,7 +16,7 @@
1616
precision_config,
1717
write_precision_result,
1818
)
19-
from .register import Register
19+
from .op_registrar import GeneralOpRegistrar
2020

2121
# Maximum tensor element count allowed for precision check
2222
# (skip if exceeded to avoid large tensor copy overhead)
@@ -177,7 +177,7 @@ def wrapper(*args, **kwargs):
177177
return wrapper
178178

179179

180-
class PrecisionCheckRegister(Register):
180+
class PrecisionCheckRegister(GeneralOpRegistrar):
181181
"""Register subclass that wraps every operator with precision checking.
182182
183183
This class is only instantiated when the user has explicitly called

0 commit comments

Comments
 (0)