changh95's picture
tt-model push diffusion-planner-p150 (container)
be62f78 verified
Raw History Blame Contribute Delete
22 kB
# SPDX-License-Identifier: Apache-2.0
"""Host tests of the bundle wiring (no device, no weights): the vendored ttaw copy is intact and every module
imports without ttnn / torch; the identity in api.py / device.py and the numerics knobs of tt/config.py match
tt-model.yaml; the device configuration resolves like the server's; the stdlib client and smoke test work against
this bundle's app served over real HTTP (ETH dispatch and 12x10 asserted, serve-profile pins, agreement with a stored
reference output); the pip project is the repo-root pyproject.toml; code/scripts/container_smoke.sh (run against fake
tt-model / docker tools) keeps its evidence and always stops the container.
TT_VISIBLE_DEVICES=none python -m pytest -q code/tt_diffusion_planner/tests/test_bundle_host.py
"""
from __future__ import annotations
import contextlib
import fnmatch
import hashlib
import json
import os
import shutil
import socket
import subprocess
import sys
import threading
from pathlib import Path
from typing import Iterator
import numpy as np
import pytest
import tt_diffusion_planner
from tt_diffusion_planner import device as tdev
from tt_diffusion_planner.api import DiffusionPlanner
from tt_diffusion_planner.tests.stubs import ETH_P150, StubModel, sample_inputs
PKG_DIR = Path(tt_diffusion_planner.__file__).resolve().parent
CODE_DIR, BUNDLE_DIR, TTAW_DIR = PKG_DIR.parent, PKG_DIR.parents[1], PKG_DIR / "ttaw"
PATCH_MARKER = "single_chip_arch_1cq_no_dispatch_s" # added to topology.cpp by patches/tt-metal-eth-dispatch.patch
VENDOR_IGNORE = ("__pycache__", "*.pyc", "*.pyo", ".pytest_cache", "*.bin", "*.log", ".DS_Store", "VENDORED.json")
BLOCKER = """
import importlib, importlib.abc, sys
BLOCKED = set(sys.argv[2].split(","))
class Block(importlib.abc.MetaPathFinder):
def find_spec(self, name, path, target=None):
if name.split(".")[0] in BLOCKED:
raise ImportError("blocked import of " + name)
sys.meta_path.insert(0, Block())
sys.path.insert(0, sys.argv[1])
for target in sys.argv[3:]:
module, _, attr = target.partition(":")
mod = importlib.import_module(module)
if attr:
getattr(mod, attr)
leaked = sorted(k for k in sys.modules if k.split(".")[0] in BLOCKED)
print("ok" if not leaked else "leaked: " + " ".join(leaked))
"""
# ------------------------------------------------------------------------------------------ vendored ttaw
def test_vendored_ttaw_matches_its_manifest():
"""The copy under code/tt_diffusion_planner/ttaw is exactly what common/tools/vendor.py wrote (never edit it
here)."""
manifest = json.loads((TTAW_DIR / "VENDORED.json").read_text())
recorded = manifest["files"]
actual = {}
for p in sorted(TTAW_DIR.rglob("*")):
rel = p.relative_to(TTAW_DIR)
if p.is_file() and not any(fnmatch.fnmatch(part, pat) for part in rel.parts for pat in VENDOR_IGNORE):
actual[rel.as_posix()] = hashlib.sha256(p.read_bytes()).hexdigest()
problems = ([f"missing: {k}" for k in sorted(set(recorded) - set(actual))]
+ [f"extra: {k}" for k in sorted(set(actual) - set(recorded))]
+ [f"modified: {k}" for k in sorted(set(recorded) & set(actual)) if recorded[k] != actual[k]])
assert problems == [], "re-vendor with common/tools/vendor.py instead of editing the copy"
from tt_diffusion_planner import ttaw
assert manifest["package"] == "ttaw" and manifest["version"] == ttaw.__version__
_P = "tt_diffusion_planner"
@pytest.mark.parametrize("blocked,targets", [
("ttnn,torch", [f"{_P}:DiffusionPlanner", f"{_P}:Output", f"{_P}:open_device", f"{_P}:load_inputs",
f"{_P}:INPUT_SCHEMA", f"{_P}.host", f"{_P}.host.solver", f"{_P}.reference.config", f"{_P}.device",
f"{_P}.io", f"{_P}.server", f"{_P}.server.smoke_test"]),
# the CPU reference also runs where ttnn is absent (research venv)
("ttnn", [f"{_P}.reference", f"{_P}.reference.pipeline", f"{_P}.reference.goldens", f"{_P}.reference.ort"]),
])
def test_imports_have_no_side_effects(blocked, targets):
if blocked == "ttnn,torch":
pytest.importorskip("fastapi")
targets = targets + ["tt_diffusion_planner.server.app:app"]
out = subprocess.run([sys.executable, "-c", BLOCKER, str(CODE_DIR), blocked] + targets, capture_output=True,
text=True)
assert out.returncode == 0 and out.stdout.strip() == "ok", out.stderr[-2000:]
# ----------------------------------------------------------------------------------- identity and config
@pytest.fixture(scope="module")
def manifest():
yaml = pytest.importorskip("yaml")
path = BUNDLE_DIR / "tt-model.yaml"
if not path.is_file():
pytest.skip("no tt-model.yaml next to code/ (installed package or container image)")
return yaml.safe_load(path.read_text())
def test_identity_matches_manifest(manifest):
cls = DiffusionPlanner
assert manifest["name"] == cls.MODEL_NAME and manifest["repo"].endswith("/" + cls.MODEL_NAME)
w = manifest["weights"]
assert (w["repo"], w["revision"]) == (cls.DEFAULT_REPO, cls.DEFAULT_REVISION)
assert list(w.get("allow_patterns") or []) == list(cls.ALLOW_PATTERNS or [])
assert manifest["runtime"]["app"] == "tt_diffusion_planner.server.app:app"
env = manifest["serve"]["env"]
assert env["TT_WEIGHTS_REVISION"] == cls.DEFAULT_REVISION
assert int(env["DIFFUSION_PLANNER_NUM_CQS"]) == tdev.DEVICE_DEFAULTS["num_command_queues"]
for profile in [{"name": "serve", "env": {}}] + list(manifest.get("serve_profiles") or []):
merged = {**env, **(profile.get("env") or {})}
assert merged.get("DIFFUSION_PLANNER_DISPATCH") == "eth", f"{profile['name']}: the p150 target is ETH dispatch"
assert merged.get("DIFFUSION_PLANNER_VARIANT") in cls.VARIANTS, profile["name"]
verify = "\n".join(manifest["verify"])
assert PATCH_MARKER in verify and cls.DEFAULT_REVISION in verify
def test_serve_env_pins_the_numerics(manifest):
"""The image serves the gated graph: ``serve.env`` pins every numerics knob of ``tt/config.py`` at its published
default (``KNOBS.serve_env()``, the ttaw.knobs convention) and an empty ``DIFFUSION_PLANNER_PRECISION`` (the
``DEFAULT_PRECISION`` policy, no override rule); no serve profile overrides them (VERIFY_PORT L2)."""
from tt_diffusion_planner.tt.config import KNOBS
env = manifest["serve"]["env"]
pins = KNOBS.serve_env()
assert set(pins) == {f"DIFFUSION_PLANNER_{k}" for k in
("LN_FP32", "HIDDEN_FP32", "SPLIT_MATMUL", "ATTN_FP32_ACC", "ATTN_MATMUL",
"ENC_CH2D", "ATTN_FAST",
"DEC_MMCFG", "LN_KERNEL", "LN_RESID", "SPLIT_KCAT", "LN_SFPU_BCAST", "LN_SPLIT", "KCAT_ACT", "ATTN_SMASK",
"ATTN_SMSM", "ENC_KCAT", "LIN_ACT", "ATTN_FUSED", "KCAT_EMIT", "LN_TR", "KCAT_L1", "KCAT_ACT_ONCE", "ATTN_L1", "DEC_L1", "ENC_L1", "FUS_L1",
"COMPACT", "AGENT_BUCKETS", "HOST_FAST", "INPUT_TRIM")}
assert {k: env.get(k) for k in pins} == pins
assert env.get("DIFFUSION_PLANNER_PRECISION") == ""
for profile in manifest.get("serve_profiles") or []:
assert not set(profile.get("env") or {}) & (set(pins) | {"DIFFUSION_PLANNER_PRECISION"}), profile["name"]
def test_device_config_resolution(monkeypatch):
for name in ("DISPATCH", "NUM_CQS", "L1_SMALL", "TRACE_REGION", "WORKER_L1_SIZE"):
monkeypatch.delenv(f"{tdev.ENV_PREFIX}_{name}", raising=False)
monkeypatch.delenv("TT_DEVICE_ID", raising=False)
cfg = tdev.device_config()
assert (cfg.dispatch, cfg.device_id, cfg.allow_fallback) == ("eth", 0, True)
assert {k: getattr(cfg, k) for k in tdev.DEVICE_DEFAULTS} == tdev.DEVICE_DEFAULTS
assert cfg == DiffusionPlanner.device_config()
monkeypatch.setenv("DIFFUSION_PLANNER_DISPATCH", "worker")
monkeypatch.setenv("DIFFUSION_PLANNER_TRACE_REGION", str(1 << 20))
monkeypatch.setenv("TT_DEVICE_ID", "1")
cfg = tdev.device_config(num_command_queues=2, worker_l1_size=None)
assert (cfg.dispatch, cfg.trace_region_size, cfg.device_id, cfg.num_command_queues) == ("worker", 1 << 20, 1, 2)
assert cfg == DiffusionPlanner.device_config(num_command_queues=2)
assert tdev.device_config(dispatch="eth").dispatch == "eth" # an explicit argument beats the environment
# ------------------------------------------------------------------------------- client and smoke test
def test_client_script_runs_standalone(tmp_path):
"""server/client.py runs with any python3 in isolated mode: no package install, no numpy needed."""
np.savez(tmp_path / "scene.npz", **sample_inputs())
subprocess.run([sys.executable, "-I", str(PKG_DIR / "server" / "client.py"),
"--inputs", str(tmp_path / "scene.npz"),
"--param", "stopping_threshold=0.5", "--out", str(tmp_path / "req.json")], check=True,
capture_output=True, text=True)
req = json.loads((tmp_path / "req.json").read_text())
assert req["inputs"]["format"] == "npz" and req["params"] == {"stopping_threshold": 0.5}
class WorkerStub(StubModel):
"""A server whose ETH open fell back to WORKER dispatch."""
device_info = dict(ETH_P150, dispatch="worker", fallback="RuntimeError: no ETH", grid="11x10", grid_x=11,
cores=110)
def _free_port() -> int:
with socket.socket() as s:
s.bind(("127.0.0.1", 0))
return s.getsockname()[1]
@contextlib.contextmanager
def _serve(stub_cls) -> Iterator[str]:
"""This bundle's real app under uvicorn on a free port with ``stub_cls`` as the model; yields the base URL."""
uvicorn = pytest.importorskip("uvicorn")
pytest.importorskip("fastapi")
from tt_diffusion_planner.server import app as server
from tt_diffusion_planner.server import smoke_test
state = server.app.state.ttaw
saved, state.model_factory = state.model_factory, stub_cls
srv = uvicorn.Server(uvicorn.Config(server.app, host="127.0.0.1", port=_free_port(), log_level="warning",
lifespan="on"))
thread = threading.Thread(target=srv.run, daemon=True)
thread.start()
base = f"http://127.0.0.1:{srv.config.port}"
try:
assert smoke_test.ttaw_client.wait_ready(base, wait_s=30, poll_s=0.1)["status"] == "ok"
yield base
finally:
srv.should_exit = True
thread.join(timeout=30)
state.model_factory = saved
def _staged_manifest(tmp_path: Path) -> Path:
"""A staged tt_kernel_manifest.json (wire format) with a default profile and one whose pins differ."""
env = {"TT_WEIGHTS_REVISION": DiffusionPlanner.DEFAULT_REVISION, "DIFFUSION_PLANNER_DISPATCH": "eth",
"DIFFUSION_PLANNER_NUM_CQS": str(tdev.DEVICE_DEFAULTS["num_command_queues"]),
"DIFFUSION_PLANNER_VARIANT": DiffusionPlanner.DEFAULT_VARIANT}
wire = {"name": DiffusionPlanner.MODEL_NAME,
"weights": {"repo_id": DiffusionPlanner.DEFAULT_REPO, "revision": DiffusionPlanner.DEFAULT_REVISION},
"container": {"serve": {"env": env}, "default_profile": None,
"serve_profiles": [{"name": "default", "env": {}},
{"name": "other", "env": {"DIFFUSION_PLANNER_VARIANT": "not-a-variant"}}]}}
path = tmp_path / "tt_kernel_manifest.json"
path.write_text(json.dumps(wire))
return path
def test_smoke_test_against_live_server(tmp_path, monkeypatch, capsys):
from tt_diffusion_planner.server import smoke_test
monkeypatch.setenv("TT_MESH_SHAPE", "1x1")
monkeypatch.setenv("TT_MODEL_WEIGHTS_REVISION", DiffusionPlanner.DEFAULT_REVISION)
for name in ("HF_MODEL", "DIFFUSION_PLANNER_VARIANT", "DIFFUSION_PLANNER_DISPATCH", "DIFFUSION_PLANNER_NUM_CQS"):
monkeypatch.delenv(name, raising=False)
sample = tmp_path / "sample.npz"
np.savez(sample, **sample_inputs())
reference = StubModel()(inputs=str(sample)).to_dict() # the stub plays the CPU reference of the sample
(tmp_path / "sample.reference.json").write_text(json.dumps(reference))
shifted = json.loads(json.dumps(reference))
for row in shifted["trajectory"]:
row[0] += 2.0
(tmp_path / "shifted.json").write_text(json.dumps(shifted))
args = ["--input", str(sample), "--out", str(tmp_path / "out.json")]
pinned = args + ["--manifest", str(_staged_manifest(tmp_path))]
with _serve(StubModel) as url:
assert smoke_test.main(["--url", url] + pinned) == 0
out = capsys.readouterr().out
assert out.startswith("PASS diffusion-planner-p150: profile=default") and "dispatch=eth grid=12x10" in out
assert "reference=sample.reference.json (trajectory ade=0.0000 fde=0.0000" in out
assert json.loads((tmp_path / "out.json").read_text())["num_poses"] == 80
assert smoke_test.main(["--url", url, "--profile", "other"] + pinned) == 1
assert "profile 'other' pins DIFFUSION_PLANNER_VARIANT=not-a-variant" in capsys.readouterr().out
assert smoke_test.main(["--url", url, "--reference", str(tmp_path / "shifted.json")] + args) == 1
assert "trajectory ADE 2.0000 > 0.3" in capsys.readouterr().out
assert smoke_test.main(["--url", url, "--reference", str(tmp_path / "missing.json")] + args) == 1
assert "reference " in capsys.readouterr().out
assert smoke_test.main(["--url", url, "--input", str(tmp_path / "missing.npz")]) == 1
assert "/predict: FileNotFoundError" in capsys.readouterr().out # a FAIL line, not a traceback
with _serve(WorkerStub) as url:
assert smoke_test.main(["--url", url] + args) == 1
out = capsys.readouterr().out
assert out.startswith("FAIL diffusion-planner-p150:")
assert "dispatch is 'worker', expected 'eth'" in out and "grid is '11x10', expected '12x10'" in out
# ------------------------------------------------------------------------- pip project and container smoke
def test_pip_project_at_repo_root():
"""The pip project is the repo-root pyproject.toml over code/ (never code/pyproject.toml: tt-model copies code/
over the tt-metal tree before building ttnn, PACKAGING_PILOT.md problem 1)."""
tomllib = pytest.importorskip("tomllib")
path = BUNDLE_DIR / "pyproject.toml"
if not (BUNDLE_DIR / "tt-model.yaml").is_file():
pytest.skip("not a bundle source checkout (installed package or container image)")
assert not (CODE_DIR / "pyproject.toml").exists(), "the pip project lives at the repo root, never in code/"
project = tomllib.loads(path.read_text())
find = project["tool"]["setuptools"]["packages"]["find"]
assert find["where"] == ["code"] and any(fnmatch.fnmatch("tt_diffusion_planner", pat) for pat in find["include"])
data = project["tool"]["setuptools"]["package-data"]
assert {"API.md", "VENDORED.json"} <= set(data["tt_diffusion_planner.ttaw"])
assert {"kernels/*.cpp", "kernels/*.hpp", "kernels/*.h"} <= set(data["tt_diffusion_planner.ttaw.ops"])
assert (BUNDLE_DIR / project["project"]["readme"]).is_file()
assert project["tool"]["setuptools"]["dynamic"]["version"]["attr"] == "tt_diffusion_planner.__version__"
yaml = pytest.importorskip("yaml")
shipped = [p for e in yaml.safe_load((BUNDLE_DIR / "tt-model.yaml").read_text())["source"]["extra_code"]
for p in e["paths"]]
assert "pyproject.toml" not in shipped and "setup.py" not in shipped
setuptools = pytest.importorskip("setuptools")
found = set(setuptools.find_packages(where=str(CODE_DIR), include=find["include"], exclude=find.get("exclude", [])))
assert {"tt_diffusion_planner", "tt_diffusion_planner.ttaw", "tt_diffusion_planner.ttaw.server",
"tt_diffusion_planner.ttaw.ops", "tt_diffusion_planner.host", "tt_diffusion_planner.reference"} <= found
assert not any(p.startswith("tt_diffusion_planner.tests") for p in found)
FAKE_TT_MODEL = r'''
import json, os, signal, subprocess, sys, time, urllib.request
from pathlib import Path
state = Path(os.environ["FAKE_STATE"])
args = sys.argv[1:]
with open(state / "calls.txt", "a") as f:
f.write(" ".join(args) + "\n")
opt = lambda name: args[args.index(name) + 1] if name in args else None
if args[0] == "serve":
port, manifest = int(opt("--port")), json.load(open(args[-1]))
(state / "container").write_text("tt-model-" + manifest["name"] + "-" + (opt("--profile") or "default"))
(state / "container.log").write_text("boot: Opening device 0 (dispatch eth, 1 CQ)\n")
if os.environ.get("FAKE_SERVE_FAIL"):
sys.exit("fake serve: the boot failed")
srv = subprocess.Popen([sys.executable, str(state / "server.py"), str(port)], start_new_session=True)
(state / "server.pid").write_text(str(srv.pid))
for _ in range(200):
try:
urllib.request.urlopen("http://127.0.0.1:%d/health" % port, timeout=1)
break
except OSError:
time.sleep(0.05)
elif args[0] == "stop":
if (state / "server.pid").exists():
os.kill(int((state / "server.pid").read_text()), signal.SIGTERM)
with open(state / "container.log", "a") as f:
f.write("shutdown: traces released, device closed\n")
(state / "stopped").write_text("1")
elif args[0] == "logs":
print("fallback: tt-model logs")
else:
sys.exit(2)
'''
FAKE_DOCKER = r'''
import os, sys, time
from pathlib import Path
state = Path(os.environ["FAKE_STATE"])
args = sys.argv[1:]
with open(state / "docker_calls.txt", "a") as f:
f.write(" ".join(args) + "\n")
known = (state / "container").read_text() if (state / "container").exists() else None
if args[0] == "inspect":
sys.exit(0 if args[-1] == known else 1)
if args[0] == "logs" and args[-1] == known:
deadline = time.time() + 20
while "--follow" in args and not (state / "stopped").exists() and time.time() < deadline:
time.sleep(0.05)
sys.stdout.write((state / "container.log").read_text())
sys.exit(0)
sys.exit(1)
'''
FAKE_SERVER = r'''
import json, sys
from http.server import BaseHTTPRequestHandler, HTTPServer
INFO = {"model": "fake", "device": {"dispatch": "eth", "grid": "12x10", "cores": 120}}
class Handler(BaseHTTPRequestHandler):
def _send(self, code, body):
data = json.dumps(body).encode()
self.send_response(code)
self.send_header("Content-Type", "application/json")
self.send_header("Content-Length", str(len(data)))
self.end_headers()
self.wfile.write(data)
def do_GET(self):
self._send(200, {"status": "ok"} if self.path.startswith("/health") else INFO if self.path == "/info" else {})
def do_POST(self):
self._send(500, {"detail": "fake server"})
def log_message(self, *args):
pass
HTTPServer(("127.0.0.1", int(sys.argv[1])), Handler).serve_forever()
'''
def _fake_tools(tmp_path: Path) -> dict:
"""PATH with fake `tt-model` / `docker` / `python3` (the test interpreter) and a stub server: no docker, no chip."""
state, bin_dir = tmp_path / "state", tmp_path / "bin"
state.mkdir()
bin_dir.mkdir()
(state / "server.py").write_text(FAKE_SERVER)
for name, body in (("tt-model", FAKE_TT_MODEL), ("docker", FAKE_DOCKER)):
(bin_dir / name).write_text(f"#!{sys.executable}\n{body}")
(bin_dir / name).chmod(0o755)
(bin_dir / "python3").write_text(f'#!/bin/sh\nexec "{sys.executable}" "$@"\n')
(bin_dir / "python3").chmod(0o755)
env = {k: v for k, v in os.environ.items() if k not in ("SMOKE_OUT", "HF_TOKEN", "HUGGING_FACE_HUB_TOKEN")}
env.update(PATH=f"{bin_dir}:{env.get('PATH', '/usr/bin:/bin')}", FAKE_STATE=str(state))
for tool in ("tt-model", "docker", "python3"): # never the real tools: a real serve would claim the chip
assert shutil.which(tool, path=env["PATH"]) == str(bin_dir / tool)
return env
@pytest.mark.parametrize("serve_fails", [False, True])
def test_container_smoke_keeps_evidence_and_stops(tmp_path, serve_fails):
"""code/scripts/container_smoke.sh passes --profile through, saves /info once READY and the container log through
the shutdown (whatever the outcome), and always stops the container."""
staged = tmp_path / "staged"
staged.mkdir()
manifest = _staged_manifest(staged)
env, logs, port = _fake_tools(tmp_path), tmp_path / "evidence", _free_port()
profile = [] if serve_fails else ["--profile", "other"]
if serve_fails:
env["FAKE_SERVE_FAIL"] = "1"
r = subprocess.run(["bash", str(CODE_DIR / "scripts" / "container_smoke.sh"), str(staged), str(port), *profile,
f"--log-dir={logs}"], env=env, capture_output=True, text=True, timeout=180)
assert r.returncode == 1, r.stdout + r.stderr # the stub serves no valid /predict: the smoke FAILS
calls = (tmp_path / "state" / "calls.txt").read_text().splitlines()
assert calls[0] == " ".join(["serve", "--port", str(port), *profile, str(manifest)])
assert calls[-1] == " ".join(["stop", *profile, str(manifest)]) and len(calls) == 2
stem = "diffusion-planner-p150" + ("" if serve_fails else "-other") + "-"
files = {p.name.split(".", 1)[1]: p for p in logs.iterdir()}
assert all(p.name.startswith(stem) for p in files.values()), sorted(files)
log = files["container.log"].read_text()
assert "boot: Opening device 0" in log and "shutdown: traces released" in log # followed through the stop
result = json.loads(files["result.json"].read_text())
assert (result["rc"], result["result"], result["profile"]) == (1, "FAIL", "default" if serve_fails else "other")
if serve_fails:
assert "info.json" not in files
else:
assert json.loads(files["info.json"].read_text())["device"] == {"dispatch": "eth", "grid": "12x10",
"cores": 120}
assert "FAIL diffusion-planner-p150:" in r.stdout