File size: 15,584 Bytes
4d9b003
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
be62f78
 
4d9b003
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
be62f78
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
4d9b003
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
203
204
205
206
207
208
209
210
211
212
213
214
215
216
217
218
219
220
221
222
223
224
225
226
227
228
229
230
231
232
233
234
235
236
237
238
239
240
241
242
243
244
245
246
247
248
249
250
251
252
253
254
255
256
257
258
259
260
261
262
263
264
265
266
267
268
269
270
271
272
273
274
275
276
277
278
279
280
281
282
283
284
285
286
287
288
289
290
291
292
293
294
295
296
297
298
299
300
301
302
303
304
305
306
307
308
309
310
311
312
313
314
315
316
317
318
319
320
321
322
323
324
325
326
327
328
329
330
331
332
333
334
335
336
337
338
339
340
# 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)