Download code/tt_diffusion_planner/ttaw/device.py from changh95/diffusion-planner-p150: direct link, hf CLI and curl.
- Browser
- Download file 15.6 kB
-
https://huggingface.co/changh95/diffusion-planner-p150/resolve/main/code/tt_diffusion_planner/ttaw/device.py
- Command line
-
hf download hf://changh95/diffusion-planner-p150/code/tt_diffusion_planner/ttaw/device.py
-
curl -L -o device.py https://huggingface.co/changh95/diffusion-planner-p150/resolve/main/code/tt_diffusion_planner/ttaw/device.py
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 | |
| 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)") | |
| 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 | |
| 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) | |