File size: 4,531 Bytes
371d90c
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
"""The single GPU entry point: run a method and its baseline, return a JSON-safe trace.

On ZeroGPU, ``@spaces.GPU`` attaches a GPU only for the duration of this call. The
arguments and the returned trace are pickled, so the trace holds plain Python types.
"""

from __future__ import annotations

import math
import time

import spaces
import torch
import transformers

import models
from decoding import contrastive, guided, parallel
from decoding.common import sync
from params import SPECS, clamp_params

# Rough seconds per decoding forward pass on a ZeroGPU slice (main, helper), bf16, eager.
PER_FWD = {"qwen3": (0.032, 0.026), "llama32": (0.026, 0.016), "gemma3": (0.026, 0.018), "smollm2": (0.022, 0.028)}
DEFAULT_FWD = (0.035, 0.03)


def estimate_duration(paradigm: str, family: str, P: dict) -> int:
    """Seconds to request from ZeroGPU; generous enough not to be cut off, small enough for quotas."""
    P = P or {}
    m, h = PER_FWD.get(family, DEFAULT_FWD)
    spec = SPECS.get(paradigm, {})
    cap = spec.get("max_new_tokens", ("int", 1, 128, 64))
    try:
        n = min(max(int(P.get("max_new_tokens") or cap[3]), 1), cap[2])
    except (TypeError, ValueError):
        n = cap[3]
    method = P.get("method") or spec.get("method", ("", (), ""))[2]
    if method == "cd":
        per_tok = 2 * m + 2 * h
    elif method in ("cad", "rose"):
        per_tok = 4 * m
    elif method == "dola":
        per_tok = 2.6 * m
    elif method == "classifier":
        L, k = int(P.get("lookahead") or 0), int(P.get("k") or 10)
        per_tok = 2 * m + L * m * (1 + k / 40) + 0.01
    elif method == "regex":
        per_tok = 2 * m + 0.01
    elif method == "speculative":
        per_tok = 2 * m + 1.2 * h
    else:  # pld, jacobi
        per_tok = 2.3 * m
    return int(min(90, max(15, 10 + 1.3 * n * per_tok)))


def warmup(fam) -> None:
    """A tiny forward per model so kernel setup doesn't land inside the timed runs."""
    ids = torch.tensor([[fam.tok.eos_token_id or 0] * 2])
    for model in (fam.main, fam.helper):
        if model is not None:
            model(input_ids=ids.to(model.device), use_cache=False)
            sync(model.device)


def run_method(paradigm: str, fam, P: dict) -> dict:
    method = P["method"]
    if paradigm == "contrastive":
        if method == "cd":
            return contrastive.run_cd(fam, P)
        if method in ("cad", "rose"):
            return contrastive.run_two_prompt(fam, P)
        return contrastive.run_dola(fam, P)
    if paradigm == "guided":
        if method == "classifier":
            return guided.run_classifier(fam, P, models.SCORERS)
        return guided.run_regex(fam, P)
    return {"speculative": parallel.run_speculative, "pld": parallel.run_pld, "jacobi": parallel.run_jacobi}[method](fam, P)


def sanitize(obj):
    """Tuples to lists, non-finite floats to None, so the trace is strict JSON."""
    if isinstance(obj, dict):
        return {str(k): sanitize(v) for k, v in obj.items()}
    if isinstance(obj, (list, tuple)):
        return [sanitize(v) for v in obj]
    if isinstance(obj, float):
        return obj if math.isfinite(obj) else None
    if isinstance(obj, (str, int, bool)) or obj is None:
        return obj
    if isinstance(obj, torch.Tensor):
        return sanitize(obj.tolist())
    return str(obj)


def device_name(device: torch.device) -> str:
    if device.type == "cuda":
        return torch.cuda.get_device_name(device)
    return device.type.upper()


@spaces.GPU(duration=estimate_duration)
def gpu_run(paradigm: str, family: str, P: dict) -> dict:
    t0 = time.perf_counter()
    fam = models.REGISTRY.get(family)
    if fam is None:
        raise ValueError(f"Model family {family!r} is not available.")
    P = clamp_params(paradigm, P)
    with torch.inference_mode():
        warmup(fam)
        result = run_method(paradigm, fam, P)
    trace = {
        "schema": "dd.trace/v1",
        "paradigm": paradigm,
        "method": P["method"],
        "family": family,
        "family_label": fam.label,
        "models": fam.info(),
        "params": P,
        **result,
        "env": {
            "device": device_name(fam.main.device),
            "dtype": str(next(fam.main.parameters()).dtype).replace("torch.", ""),
            "torch": torch.__version__,
            "transformers": transformers.__version__,
            "gpu_seconds": round(time.perf_counter() - t0, 2),
            "requested_seconds": estimate_duration(paradigm, family, P),
        },
    }
    return sanitize(trace)