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