algorithm-finder / src /zerogpu_backend.py
Michael Arana
Add ZeroGPU backend so any HF model runs on Space GPU (default XHToken/Spark-X2.5-4B)
042d836
Raw History Blame Contribute Delete
5.66 kB
"""ZeroGPU local inference backend: run ANY HF causal-LM on the Space GPU.
- Model files download once to /data (CPU, no GPU quota) via ensure_cached.
- Inference runs inside _gpu_infer (@spaces.GPU) so a GPU is allocated only
for generation. Weights load straight onto CUDA, then are released.
- Locally without `spaces`/CUDA the decorator is a no-op (CPU fallback).
Any model ID works, incl. ones with no Inference Provider such as
XHToken/Spark-X2.5-4B (custom modeling code via trust_remote_code=True).
"""
import gc
import os
import threading
PERSISTENT_CACHE = "/data/hf-hub" if os.path.isdir("/data") else None
DEFAULT_ZEROGPU_MODEL = "XHToken/Spark-X2.5-4B"
MAX_NEW_TOKENS_HARD_CAP = 8192
MAX_GPU_SECONDS = 300
try:
import spaces as _spaces
_gpu = _spaces.GPU
except Exception:
_spaces = None
def _gpu(*dargs, **dkwargs):
def wrap(fn):
return fn
if len(dargs) == 1 and callable(dargs[0]) and not dkwargs:
return dargs[0]
return wrap
_lock = threading.Lock()
def _cache_dir():
if PERSISTENT_CACHE:
os.makedirs(PERSISTENT_CACHE, exist_ok=True)
return PERSISTENT_CACHE
return None
def _gpu_seconds(model_id, prompt, max_tokens, temperature, token) -> int:
try:
return int(min(MAX_GPU_SECONDS, max(90, int(max_tokens) // 25)))
except Exception:
return 180
def ensure_cached(model_id: str, token: str = None) -> str:
from huggingface_hub import snapshot_download
try:
return snapshot_download(
repo_id=model_id,
token=token,
cache_dir=_cache_dir(),
allow_patterns=["*.json", "*.safetensors", "*.bin", "*.model", "*.txt", "*.py"],
)
except Exception as e:
msg = str(e)
if "401" in msg or "403" in msg or "gated" in msg.lower() or "access" in msg.lower():
raise RuntimeError(
f"Cannot download model '{model_id}': access denied. "
"Set an HF_TOKEN secret with access to this model."
) from e
raise RuntimeError(f"Cannot download model '{model_id}': {e}") from e
@_gpu(duration=_gpu_seconds)
def _gpu_infer(model_id, prompt, max_tokens, temperature, token=None):
import torch
from transformers import AutoTokenizer, AutoModelForCausalLM
max_tokens = int(max(64, min(MAX_NEW_TOKENS_HARD_CAP, max_tokens or 2048)))
use_cuda = torch.cuda.is_available()
device = "cuda" if use_cuda else "cpu"
auth = {"token": token} if token else {}
cache = _cache_dir()
cache_kw = {"cache_dir": cache} if cache else {}
model = None
try:
tokenizer = AutoTokenizer.from_pretrained(
model_id, trust_remote_code=True, local_files_only=False, **auth, **cache_kw
)
if tokenizer.pad_token is None:
tokenizer.pad_token = tokenizer.eos_token
try:
formatted = tokenizer.apply_chat_template(
[{"role": "user", "content": prompt}],
add_generation_prompt=True,
tokenize=False,
)
except Exception:
formatted = prompt
inputs = tokenizer(formatted, return_tensors="pt")
input_ids = inputs["input_ids"].to(device)
mask = inputs.get("attention_mask")
if mask is not None:
mask = mask.to(device)
model = AutoModelForCausalLM.from_pretrained(
model_id,
trust_remote_code=True,
torch_dtype=torch.bfloat16 if use_cuda else torch.float32,
low_cpu_mem_usage=True,
device_map=device,
**auth,
**cache_kw,
)
model.eval()
gen_kw = {
"max_new_tokens": max_tokens,
"do_sample": bool(temperature and temperature > 0),
"pad_token_id": tokenizer.eos_token_id,
"eos_token_id": tokenizer.eos_token_id,
}
if temperature and temperature > 0:
gen_kw["temperature"] = float(temperature)
gen_kw["top_p"] = 0.95
with torch.no_grad():
out = model.generate(input_ids, attention_mask=mask, **gen_kw)
text = tokenizer.decode(out[0][input_ids.shape[-1]:], skip_special_tokens=True)
return (text or "").strip()
finally:
try:
del model
except Exception:
pass
gc.collect()
try:
if torch.cuda.is_available():
torch.cuda.empty_cache()
except Exception:
pass
def generate_text(model_id, prompt, max_tokens=2048, temperature=0.2, token=None):
"""Cache on CPU, then run one metered GPU call."""
with _lock:
ensure_cached(model_id, token)
try:
text = _gpu_infer(model_id, prompt, max_tokens, temperature, token)
except Exception as e:
msg = str(e).lower()
if "out of memory" in msg or "cuda oom" in msg:
raise RuntimeError(
f"Model '{model_id}' ran out of GPU memory. "
"Use a smaller model (e.g. XHToken/Spark-X2.5-4B)."
) from e
if "quota" in msg or "no gpu" in msg or "allocat" in msg:
raise RuntimeError(
f"ZeroGPU allocation failed: {e}. Quota may be exhausted; try later."
) from e
raise RuntimeError(f"ZeroGPU inference failed for model '{model_id}': {e}") from e
if not text:
raise RuntimeError(
f"Model '{model_id}' returned empty text. Try again or reduce max tokens."
)
return text