Skip to content

Commit a90f20c

Browse files
committed
[SPEC] Fix spec testing.py
1 parent a777fcc commit a90f20c

1 file changed

Lines changed: 14 additions & 13 deletions

File tree

python/triton/testing.py

Lines changed: 14 additions & 13 deletions
Original file line numberDiff line numberDiff line change
@@ -8,6 +8,7 @@
88
from typing import Any, Dict, List
99
from . import language as tl
1010
from . import runtime
11+
from .backends import backends as _available_backends
1112

1213

1314
def nvsmi(attrs):
@@ -124,24 +125,24 @@ def do_bench_cudagraph(fn, rep=20, grad_to_none=None, quantiles=None, return_mod
124125
return _summarize_statistics(ret, quantiles, return_mode)
125126

126127

128+
# flagtree: supports specifying device_type to select runtime_driver
129+
@functools.lru_cache(maxsize=None)
130+
def _get_backend_driver(backend_name: str):
131+
if backend_name not in _available_backends:
132+
available = ", ".join(sorted(_available_backends.keys()))
133+
raise RuntimeError(f"Unsupported device_type/backend '{backend_name}'. "
134+
f"Available Triton backends: [{available}]")
135+
driver_cls = _available_backends[backend_name].driver
136+
if not driver_cls.is_active():
137+
raise RuntimeError(f"Backend '{backend_name}' is not active.")
138+
return driver_cls()
139+
140+
127141
# flagtree: supports specifying device_type to select runtime_driver
128142
def _get_runtime_driver_active(device_type: str | None):
129143
_DEVICE_TYPE_TO_BACKEND = {
130144
"cuda": "nvidia", "nvidia": "nvidia", "hip": "amd", "amd": "amd", "musa": "mthreads", "mthreads": "mthreads"
131145
}
132-
133-
@functools.lru_cache(maxsize=None)
134-
def _get_backend_driver(backend_name: str):
135-
from .backends import backends as _discovered_backends
136-
if backend_name not in _discovered_backends:
137-
available = ", ".join(sorted(_discovered_backends.keys()))
138-
raise RuntimeError(f"Unsupported device_type/backend '{backend_name}'. "
139-
f"Available Triton backends: [{available}]")
140-
driver_cls = _discovered_backends[backend_name].driver
141-
if not driver_cls.is_active():
142-
raise RuntimeError(f"Backend '{backend_name}' is not active.")
143-
return driver_cls()
144-
145146
if device_type is None:
146147
return runtime.driver.active
147148
backend_name = _DEVICE_TYPE_TO_BACKEND.get(device_type.lower(), device_type.lower())

0 commit comments

Comments
 (0)