File size: 2,176 Bytes
c699c4c
 
 
 
 
 
 
026da6c
 
 
c699c4c
 
 
026da6c
c699c4c
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
026da6c
c699c4c
026da6c
 
c699c4c
 
 
 
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
# SPDX-FileCopyrightText: © 2026 Tenstorrent USA, Inc.
# SPDX-License-Identifier: Apache-2.0
"""Minimal pytest fixtures for the SuperPoint tests (the HF snapshot shipped without this file).

``device`` / ``device_params`` (indirect) / ``--device-id``. Dispatch + grid selection for the
Galaxy BH bring-up via env ``SP_DISPATCH``:

  auto (default)    ETH dispatch + 1 CQ + 12x10 grid when the tt-metal ETH-dispatch patch is present
                    (the p150 target configuration), else worker with a warning
                    (models/tt/device_open.py, the same rule as the Python API and the server)
  eth               ETH dispatch, compute grid fixed at 12x10 (p150 + ETH dispatch equivalent;
                    project policy: 12x10 only). ETH dispatch supports only 1 command queue on
                    Galaxy, so tests asking for 2 CQs fail loudly instead of silently changing.
  worker            explicit opt-in: stock Tensix-column dispatch (Galaxy BH: 12x10; p150: 11x10)
"""
import os

import pytest


def pytest_addoption(parser):
    parser.addoption("--device-id", action="store", default=None, type=int)


def resolve_device_id(cli, env=None) -> int:
    """--device-id > $TT_DEVICE_ID > $DEVICE_ID > 0 (inside a chipenv.sh shell the pinned chip is
    always device 0)."""
    env = os.environ if env is None else env
    if cli is not None:
        return int(cli)
    for key in ("TT_DEVICE_ID", "DEVICE_ID"):
        raw = str(env.get(key, "")).strip()
        if raw:
            return int(raw)
    return 0


@pytest.fixture
def device_params(request):
    return getattr(request, "param", {}) or {}


@pytest.fixture
def device(request, device_params):
    import ttnn

    params = dict(device_params)
    dev_id = resolve_device_id(request.config.getoption("--device-id"))
    from models.tt.device_open import open_ttnn_device, resolve_dispatch

    mode = resolve_dispatch(os.environ.get("SP_DISPATCH", "auto"))
    dev = open_ttnn_device(dev_id, dispatch=mode, **params)
    g = dev.compute_with_storage_grid_size()
    print(f"conftest: dispatch={mode} compute_grid={g.x}x{g.y} params={params}")
    yield dev
    ttnn.close_device(dev)