Spaces:
Running on Zero
Running on Zero
Download gpu_entry.py from wang2226/beyond-tokens-decoding: direct link, hf CLI and curl.
- Browser
- Download file 4.53 kB
-
https://huggingface.co/spaces/wang2226/beyond-tokens-decoding/resolve/main/gpu_entry.py
- Command line
-
hf download hf://spaces/wang2226/beyond-tokens-decoding/gpu_entry.py
-
curl -L -o gpu_entry.py https://huggingface.co/spaces/wang2226/beyond-tokens-decoding/resolve/main/gpu_entry.py
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() | |
| 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) | |