"""Parallel decoding must be lossless (greedy) and distribution-preserving (sampling).""" from __future__ import annotations import random import pytest import torch from conftest import ARCHS from decoding.common import make_generator, sample from decoding.parallel import ( expected_tokens, ngram_lookup, predicted_speedup, run_jacobi, run_pld, run_speculative, verify, ) def rand_prompt(n: int, vocab: int, seed: int = 0) -> list[int]: rng = random.Random(seed) return [rng.randrange(3, vocab) for _ in range(n)] @pytest.mark.parametrize("arch", ARCHS) @pytest.mark.parametrize("gamma", [1, 3, 5]) def test_speculative_greedy_equals_autoregressive(tiny_family, arch, gamma): fam = tiny_family(arch, scale=0.08) P = {"max_new_tokens": 40, "gamma": gamma, "mode": "greedy", "temperature": 1.0, "seed": 0} res = run_speculative(fam, {**P, "prompt_ids": rand_prompt(12, 1000, seed=gamma)}) base, meth = res["runs"]["baseline"], res["runs"]["method"] assert meth["ids"] == base["ids"] assert res["summary"]["identical"] is True s = res["summary"] assert s["drafted"] > 0 and s["target_forwards"] <= len(meth["ids"]) assert len(meth["kinds"]) == len(meth["ids"]) == len(meth["tok_iter"]) def test_speculative_sampling_runs_and_counts(tiny_family): fam = tiny_family("qwen3", scale=0.05) P = {"max_new_tokens": 30, "gamma": 4, "mode": "sampling", "temperature": 0.8, "seed": 7} res = run_speculative(fam, {**P, "prompt_ids": rand_prompt(10, 1000)}) s = res["summary"] assert len(res["runs"]["method"]["ids"]) <= 30 assert 0.0 <= s["alpha_hat"] <= 1.0 assert s["accepted"] <= s["drafted"] assert "identical" not in s # identity only holds for greedy def test_speculative_sampling_preserves_target_distribution(): """Leviathan et al.: the first emitted token is distributed exactly as p.""" p = torch.tensor([0.05, 0.40, 0.10, 0.25, 0.15, 0.05]) q = torch.tensor([0.30, 0.10, 0.20, 0.10, 0.10, 0.20]) ps = torch.stack([p, p]) gen = make_generator(0) counts = torch.zeros(6) trials = 40_000 for _ in range(trials): d = sample(q, gen) n, nxt, _, _ = verify(ps, [q], [d], greedy=False, gen=gen) counts[d if n >= 1 else nxt] += 1 assert torch.allclose(counts / trials, p, atol=0.01) def test_identical_draft_is_always_accepted(): p = torch.tensor([0.1, 0.6, 0.3]) gen = make_generator(1) for _ in range(200): d = sample(p, gen) n, _, kind, probs = verify(torch.stack([p, p]), [p], [d], greedy=False, gen=gen) assert n == 1 and kind == "bonus" and probs == [1.0] @pytest.mark.parametrize("arch", ARCHS) def test_prompt_lookup_equals_autoregressive(tiny_family, arch): fam = tiny_family(arch) motif = rand_prompt(6, 1000, seed=5) prompt = motif * 3 + rand_prompt(4, 1000, seed=6) + motif[:3] P = {"max_new_tokens": 40, "num_pred": 6, "ngram_max": 3, "ngram_min": 1} res = run_pld(fam, {**P, "prompt_ids": prompt}) assert res["runs"]["method"]["ids"] == res["runs"]["baseline"]["ids"] assert res["summary"]["draft_hits"] > 0 @pytest.mark.parametrize("arch", ARCHS) @pytest.mark.parametrize("block", [1, 4, 8]) def test_jacobi_equals_autoregressive(tiny_family, arch, block): fam = tiny_family(arch) P = {"max_new_tokens": 30, "block": block, "init": "repeat", "seed": 0} res = run_jacobi(fam, {**P, "prompt_ids": rand_prompt(10, 1000, seed=block)}) meth = res["runs"]["method"] assert meth["ids"] == res["runs"]["baseline"]["ids"] assert all(it["jacobi"]["fixed"] >= 1 for it in meth["iters"]) assert res["summary"]["iterations"] <= len(meth["ids"]) def test_jacobi_random_init_equals_autoregressive(tiny_family): fam = tiny_family("gemma3") P = {"max_new_tokens": 25, "block": 6, "init": "random", "seed": 3} res = run_jacobi(fam, {**P, "prompt_ids": rand_prompt(9, 1000)}) assert res["runs"]["method"]["ids"] == res["runs"]["baseline"]["ids"] class Biased(torch.nn.Module): """Wraps a model so that one token always wins: drafts and Jacobi guesses then converge.""" def __init__(self, model, token: int): super().__init__() self.inner, self.token = model, token self.config, self.generation_config = model.config, model.generation_config @property def device(self): return self.inner.device def get_output_embeddings(self): return self.inner.get_output_embeddings() def forward(self, **kw): out = self.inner(**kw) out.logits[..., self.token] += 1e4 return out def test_multi_token_commits(tiny_family): fam = tiny_family("llama") fam.main = Biased(fam.main, 7) prompt = rand_prompt(8, 1000) + [7] * 8 jac = run_jacobi(fam, {"max_new_tokens": 20, "block": 5, "init": "repeat", "seed": 0, "prompt_ids": prompt}) assert jac["runs"]["method"]["ids"] == [7] * 20 == jac["runs"]["baseline"]["ids"] assert jac["runs"]["method"]["iters"][0]["jacobi"]["fixed"] == 6 # the whole block plus one bonus token pld = run_pld(fam, {"max_new_tokens": 20, "num_pred": 5, "ngram_max": 3, "ngram_min": 1, "prompt_ids": prompt}) assert pld["runs"]["method"]["ids"] == [7] * 20 assert pld["summary"]["tau"] > 3 def test_ngram_lookup_prefers_longest_recent_match(): S = [1, 2, 3, 9, 1, 2, 3, 8, 7, 1, 2, 3] cont, (start, n) = ngram_lookup(S, 3, 1, 2) assert n == 3 and cont == [8, 7] and start == 7 assert ngram_lookup([1, 2, 3], 3, 1, 2) == ([], None) def test_predicted_speedup_limits(): assert predicted_speedup(1.0, 4, 0.5) == pytest.approx(5 / 3) assert predicted_speedup(0.0, 4, 0.5) == pytest.approx(1 / 3) assert predicted_speedup(0.7, 0, 0.5) == 1.0 assert expected_tokens(0.5, 3) == pytest.approx(1.875)