|
8 | 8 | from typing import Any, Dict, List |
9 | 9 | from . import language as tl |
10 | 10 | from . import runtime |
| 11 | +from .backends import backends as _available_backends |
11 | 12 |
|
12 | 13 |
|
13 | 14 | def nvsmi(attrs): |
@@ -124,24 +125,24 @@ def do_bench_cudagraph(fn, rep=20, grad_to_none=None, quantiles=None, return_mod |
124 | 125 | return _summarize_statistics(ret, quantiles, return_mode) |
125 | 126 |
|
126 | 127 |
|
| 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 | + |
127 | 141 | # flagtree: supports specifying device_type to select runtime_driver |
128 | 142 | def _get_runtime_driver_active(device_type: str | None): |
129 | 143 | _DEVICE_TYPE_TO_BACKEND = { |
130 | 144 | "cuda": "nvidia", "nvidia": "nvidia", "hip": "amd", "amd": "amd", "musa": "mthreads", "mthreads": "mthreads" |
131 | 145 | } |
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 | | - |
145 | 146 | if device_type is None: |
146 | 147 | return runtime.driver.active |
147 | 148 | backend_name = _DEVICE_TYPE_TO_BACKEND.get(device_type.lower(), device_type.lower()) |
|
0 commit comments