beyond-tokens-decoding / tests /test_parallel.py
wang2226's picture
Beyond Tokens decoding playground: contrastive, guided and parallel decoding
371d90c verified
Raw History Blame Contribute Delete
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)]
@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)