superpoint-p150 / code /models /tt /device_open.py
changh95's picture
p150 ETH-dispatch compliance (2026-10-05): default ETH dispatch, 1 CQ, 12x10 in Python API and server; numbers re-measured
026da6c verified
Raw History Blame Contribute Delete
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