"""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)