Spaces:
Running on Zero
Running on Zero
Download tests/test_parallel.py from wang2226/beyond-tokens-decoding: direct link, hf CLI and curl.
- Browser
- Download file 5.84 kB
-
https://huggingface.co/spaces/wang2226/beyond-tokens-decoding/resolve/main/tests/test_parallel.py
- Command line
-
hf download hf://spaces/wang2226/beyond-tokens-decoding/tests/test_parallel.py
-
curl -L -o test_parallel.py https://huggingface.co/spaces/wang2226/beyond-tokens-decoding/resolve/main/tests/test_parallel.py
5.84 kB
| """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)] | |
| 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] | |
| 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 | |
| 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 | |
| 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) | |