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