Download code/tt_diffusion_planner/tests/test_tt_graph_host.py from changh95/diffusion-planner-p150: direct link, hf CLI and curl.
- Browser
- Download file 8.96 kB
-
https://huggingface.co/changh95/diffusion-planner-p150/resolve/main/code/tt_diffusion_planner/tests/test_tt_graph_host.py
- Command line
-
hf download hf://changh95/diffusion-planner-p150/code/tt_diffusion_planner/tests/test_tt_graph_host.py
-
curl -L -o test_tt_graph_host.py https://huggingface.co/changh95/diffusion-planner-p150/resolve/main/code/tt_diffusion_planner/tests/test_tt_graph_host.py
8.96 kB
| # SPDX-License-Identifier: Apache-2.0 | |
| """The whole ttnn graph of ``tt.model.TtDiffusionPlanner`` on the FAKE ttnn (numpy numerics: bf16 tensors hold | |
| bf16-rounded values, fp32 tensors float32), against the fp32 CPU reference's goldens: catches wiring mistakes | |
| (shapes, transposes, token order, masks, solver indices, packing) without a device. The device numbers are gated | |
| in ``test_pcc_device.py`` / ``test_e2e_device.py``. | |
| The fake op set is ``common/tests/host/fake_ttnn.py`` + ``fake_ttnn_cnn.py`` + ``fake_ttnn_attention.py`` (loaded | |
| by path like ``fake_ttnn_plugin``) plus LayerNorm / mean / subtract / tanh-GELU defined here; skipped when the | |
| workspace's common tree is absent (an installed package). | |
| TT_VISIBLE_DEVICES=none python -m pytest -q -p fake_ttnn_plugin code/tt_diffusion_planner/tests/test_tt_graph_host.py | |
| """ | |
| from __future__ import annotations | |
| import importlib.util | |
| import math | |
| import sys | |
| from pathlib import Path | |
| import numpy as np | |
| import pytest | |
| from tt_diffusion_planner.host import pipeline as hp | |
| from tt_diffusion_planner.reference import config as C | |
| from tt_diffusion_planner.ttaw.metrics import pcc | |
| PKG = Path(__file__).resolve().parents[1] | |
| FAKES = PKG.parents[3] / "common" / "tests" / "host" | |
| GOLDENS = PKG.parents[3] / "research" / "diffusion-planner" / "goldens" | |
| SCENE = "kashiwanoha_dense" | |
| def _load(name: str): | |
| path = FAKES / f"{name}.py" | |
| if not path.is_file(): | |
| pytest.skip(f"{path} not found (workspace fake ttnn op sets)") | |
| if name in sys.modules: | |
| return sys.modules[name] | |
| spec = importlib.util.spec_from_file_location(name, path) | |
| mod = importlib.util.module_from_spec(spec) | |
| sys.modules[name] = mod | |
| spec.loader.exec_module(mod) | |
| return mod | |
| def _install_planner_ops(fake): | |
| """LayerNorm, mean, subtract, rsqrt, scalar add and a linear with "gelu_tanh" on top of the CNN / attention | |
| fakes.""" | |
| names = ("linear", "layer_norm", "mean", "subtract", "rsqrt", "add", "gelu", "GeluVariant") | |
| saved = {n: getattr(fake, n) for n in names if hasattr(fake, n)} | |
| base_linear, base_add = fake.linear, fake.add | |
| variants = type("GeluVariant", (), {"Accurate": "accurate", "FastLut": "fast_lut", "Tanh": "tanh"}) | |
| def gelu_tanh(x): | |
| return 0.5 * x * (1.0 + np.tanh(math.sqrt(2.0 / math.pi) * (x + 0.044715 * x ** 3))) | |
| def linear(a, b, *, bias=None, activation=None, dtype=None, compute_kernel_config=None, **kw): | |
| if activation != "gelu_tanh": | |
| return base_linear(a, b, bias=bias, activation=activation, dtype=dtype, | |
| compute_kernel_config=compute_kernel_config, **kw) | |
| inputs = [a, b] + ([bias] if bias is not None else []) | |
| def fn(x, y, *bb): | |
| z = np.matmul(np.asarray(x, np.float64), np.asarray(y, np.float64)) | |
| return gelu_tanh(z + (bb[0].reshape(-1) if bb else 0.0)) | |
| return fake._op("linear_gelu_tanh", inputs, fn, dtype or a.dtype, fake.TILE_LAYOUT) | |
| def layer_norm(x, *, epsilon=1e-12, weight=None, bias=None, compute_kernel_config=None, **kw): | |
| inputs = [x] + [t for t in (weight, bias) if t is not None] | |
| def fn(a, *gb): | |
| a = np.asarray(a, np.float64) | |
| y = (a - a.mean(-1, keepdims=True)) / np.sqrt(a.var(-1, keepdims=True) + epsilon) | |
| if weight is not None: | |
| y = y * np.asarray(gb[0], np.float64).reshape(-1) | |
| if bias is not None: | |
| y = y + np.asarray(gb[-1], np.float64).reshape(-1) | |
| return y | |
| return fake._op("layer_norm", inputs, fn, x.dtype, x.layout, (float(epsilon),)) | |
| def mean(x, dim, keepdim=False, **kw): | |
| return fake._op("mean", [x], lambda a: np.asarray(a, np.float64).mean(axis=dim, keepdims=keepdim), x.dtype, | |
| x.layout, (dim, keepdim)) | |
| def subtract(a, b, **kw): | |
| return fake._op("subtract", [a, b], lambda x, y: np.asarray(x, np.float64) - y, a.dtype, a.layout) | |
| def rsqrt(x, *, fast_and_approximate_mode=True, **kw): | |
| return fake._op("rsqrt", [x], lambda a: 1.0 / np.sqrt(np.asarray(a, np.float64)), x.dtype, x.layout) | |
| def add(a, b, **kw): | |
| if isinstance(b, (int, float)): | |
| value = float(b) | |
| return fake._op("add_scalar", [a], lambda x: np.asarray(x, np.float64) + value, a.dtype, a.layout, | |
| (value,)) | |
| return base_add(a, b, **kw) | |
| def gelu(x, *, variant=None, fast_and_approximate_mode=False, **kw): | |
| if variant == variants.Tanh: | |
| return fake._op("gelu_tanh", [x], lambda a: gelu_tanh(np.asarray(a, np.float64)), x.dtype, x.layout) | |
| erf = np.vectorize(math.erf) | |
| return fake._op("gelu", [x], lambda a: 0.5 * a * (1.0 + erf(np.asarray(a, np.float64) / math.sqrt(2.0))), | |
| x.dtype, x.layout) | |
| for name, fn in (("linear", linear), ("layer_norm", layer_norm), ("mean", mean), ("subtract", subtract), | |
| ("rsqrt", rsqrt), ("add", add), ("gelu", gelu), ("GeluVariant", variants)): | |
| setattr(fake, name, fn) | |
| def restore(): | |
| for name in names: | |
| if name in saved: | |
| setattr(fake, name, saved[name]) | |
| elif hasattr(fake, name): | |
| delattr(fake, name) | |
| return restore | |
| def fake_planner(): | |
| import ttnn # the fake (fake_ttnn_plugin) | |
| if not hasattr(ttnn, "_op"): | |
| pytest.skip("needs the fake ttnn (-p fake_ttnn_plugin)") | |
| from tt_diffusion_planner.reference.weights import find_weights_dir, load_weights | |
| wd = find_weights_dir() | |
| if wd is None: | |
| pytest.skip("weights not found") | |
| ttnn.reset() | |
| undo = [_load("fake_ttnn_cnn").install(ttnn), _load("fake_ttnn_attention").install(ttnn)] | |
| undo.append(_install_planner_ops(ttnn)) | |
| from tt_diffusion_planner.ttaw.device import close_device, open_device | |
| from tt_diffusion_planner.tt.model import TtDiffusionPlanner | |
| weights = load_weights(wd) | |
| dev = open_device(0, dispatch="eth", num_command_queues=1) | |
| tt = TtDiffusionPlanner(dev, weights, debug=True) | |
| yield tt, weights | |
| tt.release() | |
| close_device(dev) | |
| for fn in reversed(undo): | |
| fn() | |
| def scene(fake_planner): | |
| tt, weights = fake_planner | |
| path = GOLDENS / f"{SCENE}.npz" | |
| if not path.is_file(): | |
| pytest.skip(f"research goldens {path} not found") | |
| with np.load(path, allow_pickle=False) as z: | |
| gold = {k: z[k] for k in z.files if k != "__meta__"} | |
| raw = {k: gold[f"in.{k}"] for k in C.INPUT_NAMES} | |
| return hp.prepare(raw, weights.normalization.observation), gold | |
| def test_encoder_taps_match_the_reference(fake_planner, scene): | |
| tt, _ = fake_planner | |
| prep, gold = scene | |
| taps = tt.encoder_taps(prep, eager=True) | |
| for name, _ in C.TOKEN_LAYOUT: | |
| rows = np.flatnonzero(gold[f"host.valid.{name}"]) | |
| if rows.size: | |
| assert pcc(taps[f"enc.{name}"][rows], gold[f"enc.{name}"][rows]) > 0.999, name | |
| for c in ("ego", "neighbor", "lane", "route", "line_string"): | |
| rows = gold[f"enc.{c}.pre.rows"] | |
| assert pcc(taps[f"enc.{c}.pre"][rows], gold[f"enc.{c}.pre"]) > 0.9999, c | |
| assert pcc(taps[f"enc.{c}.mixer"][rows], gold[f"enc.{c}.mixer"]) > 0.999, c | |
| tok = np.flatnonzero(gold["host.token_valid"]) | |
| assert pcc(taps["enc.encoding"][tok], gold["enc.encoding"][tok]) > 0.999 | |
| def test_decode_once_matches_the_reference(fake_planner, scene): | |
| tt, _ = fake_planner | |
| prep, gold = scene | |
| rows = gold["dec.rows"] | |
| for k in (0, 10): | |
| got = tt.decode_once(prep, gold["dec.x_in"][k], float(gold["dec.t"][k]), encoding=gold["enc.encoding"], | |
| eager=True) | |
| assert not got[:, 0].any() # the masked t = 0 columns | |
| assert pcc(got[rows][:, 1:], gold["dec.out"][k][:, 1:]) > 0.999, k | |
| def test_plan_and_trace(fake_planner, scene): | |
| """The plan eagerly vs the reference's final x0 / logits, then captured: replay == eager.""" | |
| tt, weights = fake_planner | |
| prep, gold = scene | |
| eager = tt.forward(prep, eager=True) | |
| rows = gold["dec.rows"] | |
| assert pcc(eager["final_x0"][rows], gold["final_x0"][rows]) > 0.999 | |
| np.testing.assert_array_equal(eager["final_x0"][:, 0], prep.decoder.current_states) # prefix constraint | |
| assert np.abs(eager["logit"] - gold["turn.logit"]).max() < 0.05 | |
| assert len(eager["denoising_steps"]) == C.DPM_SOLVER_STEPS + 1 | |
| np.testing.assert_array_equal(eager["denoising_steps"][-1][0], eager["final_x0"][0]) | |
| tt.capture() | |
| replay = tt.forward(prep) | |
| for k in ("final_x0", "logit"): | |
| np.testing.assert_array_equal(replay[k], eager[k]) | |
| out = hp.make_output(replay["final_x0"], replay["logit"], prep, weights.normalization, | |
| {k: v[3] for k, v in hp.RUNTIME_PARAMS.items()}) | |
| assert int(out.turn_indicator["command"]) == int(gold["out.turn_command"]) | |