changh95's picture
tt-model push diffusion-planner-p150 (container)
be62f78 verified
Raw History Blame Contribute Delete
15.6 kB
# SPDX-License-Identifier: Apache-2.0
"""C01: open the Blackhole p150 the way every published number is measured.
Default: dispatch on the idle Ethernet cores (``ttnn.DispatchCoreConfig(ttnn.DispatchCoreType.ETH)``), which frees the
Tensix dispatch column and gives a 12x10 = 120-core compute grid on a p150b (11x10 with stock WORKER dispatch). ETH
dispatch needs ``patches/tt-metal-eth-dispatch.patch`` on tt-metal 44d6650; the patch adds the marker
``single_chip_arch_1cq_no_dispatch_s`` to ``tt_metal/impl/dispatch/topology.cpp``, which ``dispatch="auto"`` looks
for. If the ETH open fails, a ``RuntimeWarning`` is issued and WORKER dispatch is used (``allow_fallback=False``
raises instead, for the container smoke). Never hard-code the grid: use :func:`compute_grid` /
``device.compute_with_storage_grid_size()``.
``ttnn`` is imported inside the functions, so importing this module has no side effects.
Example::
from .ttaw.device import device_session, describe_device
with device_session(dispatch="eth", num_command_queues=2, trace_region_size=64 << 20) as dev:
print(describe_device(dev)) # {'dispatch': 'eth', 'grid': '12x10', 'cores': 120, 'num_command_queues': 2, ...}
"""
from __future__ import annotations
import contextlib
import importlib.util
import os
import warnings
from dataclasses import asdict, dataclass
from pathlib import Path
from typing import Any, Dict, Iterator, Mapping, Optional, Tuple
__all__ = [
"DISPATCH_MODES",
"PATCH_MARKER",
"DEFAULT_L1_SMALL_SIZE",
"DEFAULT_TRACE_REGION_SIZE",
"P150_ETH_GRID",
"DeviceConfig",
"tt_metal_home",
"eth_dispatch_patch_present",
"reshape_patch_present",
"RESHAPE_PATCH_MARKER",
"resolve_dispatch",
"open_device",
"describe_device",
"open_info",
"close_device",
"device_session",
"compute_grid",
"core_grid",
"full_core_range_set",
]
DISPATCH_MODES = ("eth", "worker", "auto")
PATCH_MARKER = "single_chip_arch_1cq_no_dispatch_s"
PATCH_RELPATH = Path("tt_metal", "impl", "dispatch", "topology.cpp")
DEFAULT_L1_SMALL_SIZE = 32768 # conv / pool config tensors (CNN demos use 24576-32768)
DEFAULT_TRACE_REGION_SIZE = 64 << 20 # DRAM bytes for the trace command buffers of all captured traces
P150_ETH_GRID = (12, 10) # p150b compute grid with ETH dispatch (research/DISPATCH.md)
# id(device) -> (device, what open_device actually did: dispatch used, CQs, sizes); read by describe_device. The
# device itself is kept so that its id cannot be reused by another object while the entry exists (a device opened
# outside ttaw must never inherit a dead device's record); close_device removes the entry.
_OPEN_INFO: Dict[int, Tuple[Any, Dict[str, Any]]] = {}
# id(device) -> device for every device closed by close_device: a second close is a no-op. The reference keeps the
# closed object alive, so its id cannot be reused by a later device (ttnn device objects take no attributes).
_CLOSED: Dict[int, Any] = {}
def _normalize_dispatch(dispatch: Optional[str]) -> str:
mode = (dispatch or "eth").strip().lower()
if mode not in DISPATCH_MODES:
raise ValueError(f"dispatch={dispatch!r}: expected one of {DISPATCH_MODES}")
return mode
@dataclass(frozen=True)
class DeviceConfig:
"""Everything :func:`open_device` needs, so a model can pin it, log it and A/B it from the environment."""
device_id: int = 0
dispatch: str = "eth"
num_command_queues: int = 1
l1_small_size: int = DEFAULT_L1_SMALL_SIZE
trace_region_size: int = DEFAULT_TRACE_REGION_SIZE
worker_l1_size: Optional[int] = None
allow_fallback: bool = True
def __post_init__(self) -> None:
object.__setattr__(self, "dispatch", _normalize_dispatch(self.dispatch))
if self.num_command_queues not in (1, 2):
raise ValueError(f"num_command_queues={self.num_command_queues}: expected 1 or 2")
if self.l1_small_size < 0 or self.trace_region_size < 0:
raise ValueError("l1_small_size and trace_region_size must be >= 0")
if self.worker_l1_size is not None and self.worker_l1_size <= 0:
raise ValueError("worker_l1_size must be a positive byte count (or None for tt-metal's default)")
@classmethod
def from_env(cls, prefix: str, env: Optional[Mapping[str, str]] = None, **defaults: Any) -> "DeviceConfig":
"""Model defaults overridden by ``<PREFIX>_DISPATCH`` (eth|worker|auto), ``<PREFIX>_NUM_CQS``,
``<PREFIX>_L1_SMALL``, ``<PREFIX>_TRACE_REGION``, ``<PREFIX>_WORKER_L1_SIZE`` and ``TT_DEVICE_ID``
(BUNDLE_CONVENTIONS.md section 7.6). Empty or ``0`` values keep the default. Read once, at build."""
env = os.environ if env is None else env
values = dict(defaults)
def _int(name: str) -> Optional[int]:
raw = (env.get(name) or "").strip()
if not raw:
return None
try:
value = int(float(raw))
except ValueError:
raise ValueError(f"{name}={raw!r} is not an integer") from None
return value or None
dispatch = (env.get(f"{prefix}_DISPATCH") or "").strip()
if dispatch:
values["dispatch"] = dispatch
for key, name in (("num_command_queues", f"{prefix}_NUM_CQS"), ("l1_small_size", f"{prefix}_L1_SMALL"),
("trace_region_size", f"{prefix}_TRACE_REGION"),
("worker_l1_size", f"{prefix}_WORKER_L1_SIZE")):
value = _int(name)
if value is not None:
values[key] = value
device_id = (env.get("TT_DEVICE_ID") or "").strip()
if device_id:
values["device_id"] = int(device_id)
return cls(**values)
def open(self):
"""``open_device(**self)``."""
return open_device(**asdict(self))
def tt_metal_home() -> Optional[Path]:
"""The tt-metal tree: ``$TT_METAL_HOME``, else ``$TT_METAL_RUNTIME_ROOT``, else the tree holding the ``ttnn``
package (located with ``importlib.util.find_spec``, so ttnn is not imported)."""
for var in ("TT_METAL_HOME", "TT_METAL_RUNTIME_ROOT"):
value = os.environ.get(var)
if value and Path(value).is_dir():
return Path(value)
try:
spec = importlib.util.find_spec("ttnn")
except (ImportError, ValueError):
spec = None
if spec is None or not spec.origin:
return None
for parent in Path(spec.origin).resolve().parents:
if (parent / "tt_metal").is_dir():
return parent
return None
def eth_dispatch_patch_present(home: Optional[os.PathLike] = None) -> bool:
"""True when the tt-metal tree carries the ETH-dispatch patch (its marker string is in ``topology.cpp``)."""
root = tt_metal_home() if home is None else Path(home)
if root is None:
return False
try:
return PATCH_MARKER in (Path(root) / PATCH_RELPATH).read_text(errors="ignore")
except OSError:
return False
# patches/tt-metal-reshape-rm-sys1419.patch (ttaw 0.23.2): single-kernel ROW_MAJOR reshape for small Blackhole DRAM
# destination pages (SYS-1419: the dual-kernel path hangs the chip under ETH dispatch; logs/meteor/hangfix/)
RESHAPE_PATCH_MARKER = "dual_kernel_min_dram_dest_page_bytes"
RESHAPE_PATCH_RELPATH = "ttnn/cpp/ttnn/operations/data_movement/reshape_view/device/reshape_rm_program_factory.cpp"
RESHAPE_PATCH_ENV = "TT_METAL_RESHAPE_RM_DUAL_DRAM_MIN_PAGE_BYTES"
def reshape_patch_present(home: Optional[os.PathLike] = None, *, min_page_bytes: int = 4096) -> bool:
"""True when the tt-metal tree carries the reshape patch (marker in ``reshape_rm_program_factory.cpp``) and its
runtime override (``TT_METAL_RESHAPE_RM_DUAL_DRAM_MIN_PAGE_BYTES``) does not lower the guard below
``min_page_bytes``. Like :func:`eth_dispatch_patch_present` it reads the source tree: the library must have been
rebuilt from it (the container image always is)."""
raw = (os.environ.get(RESHAPE_PATCH_ENV) or "").strip()
if raw:
try:
if int(raw) < int(min_page_bytes):
return False
except ValueError:
return False
root = tt_metal_home() if home is None else Path(home)
if root is None:
return False
try:
return RESHAPE_PATCH_MARKER in (Path(root) / RESHAPE_PATCH_RELPATH).read_text(errors="ignore")
except OSError:
return False
def resolve_dispatch(dispatch: str = "auto", *, home: Optional[os.PathLike] = None) -> str:
"""``"eth"`` / ``"worker"`` pass through; ``"auto"`` becomes ``"eth"`` when the patch marker is found, else
``"worker"`` with a ``RuntimeWarning`` (the grid is then 11x10 and the published numbers do not apply)."""
mode = _normalize_dispatch(dispatch)
if mode != "auto":
return mode
if eth_dispatch_patch_present(home):
return "eth"
warnings.warn(
"the tt-metal ETH-dispatch patch was not found (marker "
f"{PATCH_MARKER!r} missing from {PATCH_RELPATH} under {tt_metal_home()}); using WORKER dispatch: the compute "
"grid is 11x10 on a p150 and the published numbers do not apply (apply patches/tt-metal-eth-dispatch.patch)",
RuntimeWarning, stacklevel=3)
return "worker"
def open_device(device_id: int = 0, *, dispatch: str = "eth", num_command_queues: int = 1,
l1_small_size: int = DEFAULT_L1_SMALL_SIZE, trace_region_size: int = DEFAULT_TRACE_REGION_SIZE,
worker_l1_size: Optional[int] = None, allow_fallback: bool = True):
"""Open one chip and return the ttnn (1x1 mesh) device.
Args:
device_id: chip id as ttnn sees it.
dispatch: ``"eth"`` (default; 12x10 on a p150b), ``"worker"`` (A/B only; 11x10) or ``"auto"`` (ETH when the
patch marker is present, else WORKER with a ``RuntimeWarning``).
num_command_queues: 1, or 2 for trace + 2CQ (input upload on CQ1).
l1_small_size: L1_SMALL bytes per core (conv / pool config tensors).
trace_region_size: DRAM bytes reserved for trace buffers (0 = dynamic top-down buffers).
worker_l1_size: allocatable L1 per core in bytes; smaller values grow the kernel-config ring buffer
(TT_PLATFORM.md section 1). ``None`` keeps tt-metal's default.
allow_fallback: when the ETH open fails, warn and open with WORKER dispatch (``True``) or re-raise.
"""
import ttnn
cfg = DeviceConfig(device_id=device_id, dispatch=dispatch, num_command_queues=num_command_queues,
l1_small_size=l1_small_size, trace_region_size=trace_region_size,
worker_l1_size=worker_l1_size, allow_fallback=allow_fallback)
mode = resolve_dispatch(cfg.dispatch)
params: Dict[str, Any] = dict(l1_small_size=cfg.l1_small_size, trace_region_size=cfg.trace_region_size,
num_command_queues=cfg.num_command_queues)
if cfg.worker_l1_size is not None:
params["worker_l1_size"] = cfg.worker_l1_size
device = None
fallback: Optional[str] = None
if mode == "eth":
try:
device = ttnn.open_device(device_id=cfg.device_id,
dispatch_core_config=ttnn.DispatchCoreConfig(ttnn.DispatchCoreType.ETH), **params)
except Exception as exc: # noqa: BLE001 -- tt-metal without the ETH-dispatch patch, or no idle ETH cores
if not cfg.allow_fallback:
raise
fallback = f"{type(exc).__name__}: {exc}"
warnings.warn(f"ETH dispatch is not available ({fallback}); falling back to WORKER dispatch: the compute "
"grid is 11x10 on a p150 and the published numbers do not apply "
"(apply patches/tt-metal-eth-dispatch.patch)", RuntimeWarning, stacklevel=2)
mode = "worker"
if device is None:
device = ttnn.open_device(device_id=cfg.device_id,
dispatch_core_config=ttnn.DispatchCoreConfig(ttnn.DispatchCoreType.WORKER), **params)
_CLOSED.pop(id(device), None)
_OPEN_INFO[id(device)] = (device, {
"dispatch": mode, "dispatch_requested": cfg.dispatch, "fallback": fallback, "device_id": cfg.device_id,
"num_command_queues": cfg.num_command_queues, "l1_small_size": cfg.l1_small_size,
"trace_region_size": cfg.trace_region_size, "worker_l1_size": cfg.worker_l1_size,
})
gx, gy = compute_grid(device)
if mode == "eth" and (gx, gy) != P150_ETH_GRID:
warnings.warn(f"compute grid is {gx}x{gy}, expected {P150_ETH_GRID[0]}x{P150_ETH_GRID[1]} with ETH dispatch "
"on a p150b", RuntimeWarning, stacklevel=2)
return device
def compute_grid(device) -> Tuple[int, int]:
"""``(x, y)`` of ``device.compute_with_storage_grid_size()``: (12, 10) with ETH dispatch on a p150b."""
g = device.compute_with_storage_grid_size()
return int(g.x), int(g.y)
def core_grid(device):
"""``ttnn.CoreGrid`` covering the whole compute grid (for ``core_grid=`` / program configs)."""
import ttnn
x, y = compute_grid(device)
return ttnn.CoreGrid(y=y, x=x)
def full_core_range_set(device):
"""``ttnn.CoreRangeSet`` of one rectangle covering the whole compute grid (for ``generic_op`` descriptors)."""
import ttnn
x, y = compute_grid(device)
return ttnn.CoreRangeSet([ttnn.CoreRange(ttnn.CoreCoord(0, 0), ttnn.CoreCoord(x - 1, y - 1))])
def open_info(device) -> Dict[str, Any]:
"""What :func:`open_device` recorded for ``device`` (dispatch used, CQs, sizes), ``{}`` for a device opened
elsewhere."""
entry = _OPEN_INFO.get(id(device))
return dict(entry[1]) if entry is not None and entry[0] is device else {}
def _arch_name(device) -> str:
try:
arch = device.arch()
except Exception: # noqa: BLE001 -- a fake or partially closed device
return "unknown"
name = getattr(arch, "name", None) or str(arch).rsplit(".", 1)[-1]
return name.lower()
def describe_device(device, device_id: Optional[int] = None) -> Dict[str, Any]:
"""What ``/info`` and ``model.info`` report: dispatch actually used, CQs, grid, cores, arch and open parameters.
A device opened outside :func:`open_device` reports ``dispatch="unknown"``."""
gx, gy = compute_grid(device)
info: Dict[str, Any] = {"dispatch": "unknown", "num_command_queues": None}
info.update(open_info(device))
info.update({"grid": f"{gx}x{gy}", "grid_x": gx, "grid_y": gy, "cores": gx * gy, "arch": _arch_name(device),
"eth_patch": eth_dispatch_patch_present()})
if device_id is not None:
info["device_id"] = device_id
return info
def close_device(device) -> None:
"""Close a device (synchronizes first, as ``ttnn.close_device`` does). Idempotent: closing a device this
function already closed does nothing."""
import ttnn
if _CLOSED.get(id(device)) is device:
return
if open_info(device):
del _OPEN_INFO[id(device)]
try:
ttnn.close_device(device)
finally:
_CLOSED[id(device)] = device
@contextlib.contextmanager
def device_session(device_id: int = 0, **kwargs: Any) -> Iterator[Any]:
"""``with device_session(dispatch="eth", num_command_queues=2) as dev:`` -- opens with :func:`open_device`
(same keyword arguments) and always closes, also when the body raises."""
device = open_device(device_id, **kwargs)
try:
yield device
finally:
close_device(device)