File size: 5,654 Bytes
026da6c
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
# 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