| 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_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") |
| |
| am = target.argmax(dim=-1) |
| draft_tokens = am[:, :k].clone() |
| draft_tokens[0, 2] = (draft_tokens[0, 2] + 1) % V |
| 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]) |
|
|