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)
|