Download code/models/tt/device_open.py from changh95/superpoint-p150: direct link, hf CLI and curl.
- Browser
- Download file 5.65 kB
-
https://huggingface.co/changh95/superpoint-p150/resolve/main/code/models/tt/device_open.py
- Command line
-
hf download hf://changh95/superpoint-p150/code/models/tt/device_open.py
-
curl -L -o device_open.py https://huggingface.co/changh95/superpoint-p150/resolve/main/code/models/tt/device_open.py
5.65 kB
| # SPDX-License-Identifier: Apache-2.0 | |
| """Open one Blackhole chip the way the published SuperPoint numbers were measured. | |
| Used by the Python API (``tt_superpoint.device`` re-exports this module) and by the HTTP | |
| server (``models/server/app.py``), so both open the chip the same way. | |
| Default (``dispatch="auto"``): ETH dispatch, compute grid capped at 12x10, 1 command queue. | |
| This is the p150 target configuration: on a p150 the stock (Tensix/worker) dispatch takes one | |
| Tensix column and leaves only 11x10 compute cores; ETH dispatch frees that column (12x10). | |
| ETH dispatch needs the tt-metal ETH-dispatch patch and supports only ``num_command_queues=1``. | |
| When the patch is not found, ``"auto"`` falls back to the default (Tensix) dispatch and | |
| warns (11x10 on a p150). ``SP_DISPATCH=worker`` forces the Tensix dispatch (explicit opt-in; | |
| on this Galaxy it happens to give 12x10, on a p150 it gives 11x10). | |
| """ | |
| from __future__ import annotations | |
| import os | |
| import warnings | |
| from pathlib import Path | |
| from typing import Optional, Tuple | |
| #: Device-open constants of the server and of the published numbers. | |
| L1_SMALL_SIZE = 32 * 1024 | |
| TRACE_REGION_SIZE = 32 * 1024 * 1024 | |
| #: The compute grid every kernel of this port is tuned for (p150, ETH dispatch). | |
| GRID: Tuple[int, int] = (12, 10) | |
| _GRID_ENV = "TT_METAL_CORE_GRID_OVERRIDE_TODEPRECATE" | |
| #: A symbol that the ETH-dispatch patch adds to tt-metal (tt_metal/impl/dispatch/topology.cpp). | |
| _PATCH_FILE = Path("tt_metal", "impl", "dispatch", "topology.cpp") | |
| _PATCH_MARKER = "single_chip_arch_1cq_no_dispatch_s" | |
| def tt_metal_home() -> Optional[Path]: | |
| """The tt-metal source tree: ``$TT_METAL_HOME``, else the tree that holds the imported ttnn.""" | |
| env = os.environ.get("TT_METAL_HOME") | |
| if env: | |
| return Path(env) | |
| try: | |
| import ttnn | |
| p = Path(ttnn.__file__).resolve() | |
| for parent in p.parents: | |
| if (parent / "tt_metal").is_dir(): | |
| return parent | |
| except Exception: # noqa: BLE001 - detection only | |
| pass | |
| return None | |
| def eth_dispatch_patch_present(home: Optional[Path] = None) -> bool: | |
| """True when the tt-metal tree contains the ETH-dispatch patch.""" | |
| home = tt_metal_home() if home is None else Path(home) | |
| if home is None: | |
| return False | |
| try: | |
| return _PATCH_MARKER in (home / _PATCH_FILE).read_text(errors="ignore") | |
| except OSError: | |
| return False | |
| def resolve_dispatch(dispatch: str = "auto") -> str: | |
| """``"auto"`` -> ``"eth"`` when the patched tt-metal is present, else ``"worker"`` (warns). | |
| ``$SP_DISPATCH`` (eth / worker) overrides ``"auto"``.""" | |
| d = (dispatch or "auto").lower() | |
| if d == "auto": | |
| d = os.environ.get("SP_DISPATCH", "auto").lower() or "auto" | |
| if d == "auto": | |
| if eth_dispatch_patch_present(): | |
| return "eth" | |
| warnings.warn( | |
| "superpoint: the tt-metal ETH-dispatch patch was not found (TT_METAL_HOME=" | |
| f"{os.environ.get('TT_METAL_HOME')!r}); opening the device with the default (Tensix) " | |
| "dispatch. Outputs are the same, but on a p150 the compute grid is then 11x10, not " | |
| "12x10; the published timings were measured with ETH dispatch, 1 command queue and " | |
| "a 12x10 grid. Set SP_DISPATCH=worker to silence this warning.", | |
| RuntimeWarning, | |
| stacklevel=3, | |
| ) | |
| return "worker" | |
| if d not in ("eth", "worker"): | |
| raise ValueError(f"dispatch must be 'auto', 'eth' or 'worker', got {dispatch!r}") | |
| return d | |
| def open_ttnn_device(device_id: int = 0, *, dispatch: str = "eth", **kwargs): | |
| """``ttnn.open_device`` with ``dispatch`` ``"eth"`` (ETH dispatch, compute grid capped at 12x10) | |
| or ``"worker"`` (default Tensix dispatch). ``kwargs`` go to ``ttnn.open_device``.""" | |
| import ttnn | |
| if dispatch == "eth": | |
| # tt-metal reads the grid cap when the first device of the process opens | |
| os.environ.setdefault(_GRID_ENV, f"{GRID[0] - 1},{GRID[1] - 1}") | |
| kwargs.setdefault("dispatch_core_config", ttnn.DispatchCoreConfig(ttnn.DispatchCoreType.ETH)) | |
| if kwargs.get("num_command_queues", 1) != 1: | |
| raise ValueError("ETH dispatch on a Galaxy Blackhole chip supports only num_command_queues=1") | |
| elif dispatch != "worker": | |
| raise ValueError(f"dispatch must be 'eth' or 'worker', got {dispatch!r}") | |
| dev = ttnn.open_device(device_id=int(device_id), **kwargs) | |
| g = dev.compute_with_storage_grid_size() | |
| if dispatch == "eth" and (g.x, g.y) != GRID: | |
| warnings.warn( | |
| f"superpoint: compute grid is {g.x}x{g.y}, not {GRID[0]}x{GRID[1]}. Set " | |
| f"{_GRID_ENV}={GRID[0] - 1},{GRID[1] - 1} before the first device opens in this process.", | |
| RuntimeWarning, | |
| stacklevel=2, | |
| ) | |
| return dev | |
| def open_device( | |
| device_id: int = 0, | |
| *, | |
| dispatch: str = "auto", | |
| l1_small_size: int = L1_SMALL_SIZE, | |
| trace_region_size: int = TRACE_REGION_SIZE, | |
| ): | |
| """Open ``device_id`` with the published configuration. Returns ``(device, dispatch)``, | |
| where ``dispatch`` is the mode used (``"eth"`` or ``"worker"``). | |
| ``dispatch``: ``"auto"`` (ETH when the patched tt-metal is present, else Tensix with a | |
| warning), ``"eth"`` or ``"worker"``. With ETH dispatch the compute grid is capped at 12x10 | |
| through ``TT_METAL_CORE_GRID_OVERRIDE_TODEPRECATE`` (set here if not already set). | |
| """ | |
| mode = resolve_dispatch(dispatch) | |
| dev = open_ttnn_device(device_id, dispatch=mode, l1_small_size=int(l1_small_size), | |
| trace_region_size=int(trace_region_size)) | |
| return dev, mode | |