import pytest import torch import kernels sdo = kernels.get_kernel("phanerozoic/spec-decode-ops", version=1, trust_remote_code=True) def hf_filter_reference(logits, temperature, top_k, top_p, min_p): """Reference filter chain with transformers semantics.""" l = logits.float() / temperature if top_k > 0: kth = torch.topk(l, top_k, dim=-1).values[..., -1, None] l = l.masked_fill(l < kth, float("-inf")) if top_p < 1.0: sl, si = torch.sort(l, descending=True, dim=-1) p = torch.softmax(sl, dim=-1) cum = p.cumsum(dim=-1) keep_sorted = (cum - p) < top_p # keep tokens whose exclusive cumsum < top_p keep_sorted[..., 0] = True keep = torch.zeros_like(keep_sorted).scatter(-1, si, keep_sorted) l = l.masked_fill(~keep, float("-inf")) if min_p > 0.0: pmax = l.max(dim=-1, keepdim=True).values l = l.masked_fill(l < pmax + torch.log(torch.tensor(min_p)), float("-inf")) return l @pytest.mark.kernels_ci @pytest.mark.parametrize("top_k,top_p,min_p", [(50, 1.0, 0.0), (0, 0.9, 0.0), (0, 1.0, 0.05), (40, 0.92, 0.02)]) def test_kept_set_matches_reference(top_k, top_p, min_p): torch.manual_seed(0) logits = torch.randn(8, 4096, device="cuda") * 3.0 ours = sdo.filter_logits(logits, temperature=0.7, top_k=top_k, top_p=top_p, min_p=min_p) ref = hf_filter_reference(logits, 0.7, top_k, top_p, min_p) assert torch.equal(torch.isinf(ours), torch.isinf(ref)) @pytest.mark.kernels_ci def test_greedy_matches_argmax(): torch.manual_seed(1) logits = torch.randn(16, 32000, device="cuda", dtype=torch.bfloat16) tok = sdo.sample(logits, temperature=0.0) assert torch.equal(tok, logits.float().argmax(dim=-1)) @pytest.mark.kernels_ci def test_verify_greedy_contract(): torch.manual_seed(2) B, k, V = 4, 5, 1024 target = torch.randn(B, k + 1, V, device="cuda") # draft tokens equal to target argmax for the first j positions am = target.argmax(dim=-1) # [B, k+1] draft_tokens = am[:, :k].clone() draft_tokens[0, 2] = (draft_tokens[0, 2] + 1) % V # force mismatch at position 2 draft_logits = torch.randn(B, k, V, device="cuda") alen, out = sdo.verify(target, draft_logits, draft_tokens, temperature=0.0) assert alen[0].item() == 2 assert (alen[1:] == k).all() assert torch.equal(out[1, : k + 1], am[1])