changh95's picture
tt-model push diffusion-planner-p150 (container)
4d9b003 verified
Raw History Blame Contribute Delete
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
@pytest.fixture(scope="module")
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()
@pytest.fixture(scope="module")
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"])