# 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