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
6.47 kB
import os
from huggingface_hub import InferenceClient
from huggingface_hub.utils import HfHubHTTPError
from .prompts import CODER_PROMPT, RESEARCHER_PROMPT, VALIDATOR_PROMPT, REALWORLD_PROMPT, SUMMARY_PROMPT
DEFAULT_MODEL = "XHToken/Spark-X2.5-4B"
SUPPORTED_EXAMPLES = "Qwen/Qwen2.5-7B-Instruct, Qwen/Qwen2.5-14B-Instruct, meta-llama/Meta-Llama-3-8B-Instruct"
class LLMClient:
def __init__(self, model_id: str = None, token: str = None, provider: str = None,
backend: str = None):
self.model_id = model_id or os.getenv("HF_MODEL_ID", DEFAULT_MODEL)
self.token = token or os.getenv("HF_TOKEN") or os.getenv("HUGGINGFACE_HUB_TOKEN")
self.provider = provider or os.getenv("HF_PROVIDER")
self.backend = (backend or os.getenv("LLM_BACKEND") or "auto").lower()
# Newer huggingface_hub prefers api_key; older versions accept token.
self.client = self._build_client(self.token, self.provider)
def using_local_gpu(self) -> bool:
if self.backend in ("local", "zerogpu", "gpu"):
return True
if self.backend in ("api", "serverless", "providers"):
return False
return self.provider is None or self.provider.lower() in ("", "local", "zerogpu")
@staticmethod
def _build_client(token, provider):
base = {"timeout": 120}
if provider:
base["provider"] = provider
if not token:
return InferenceClient(**base)
try:
return InferenceClient(api_key=token, **base)
except TypeError:
return InferenceClient(token=token, **base)
def generate(self, prompt: str, max_tokens: int = 2048, temperature: float = 0.2) -> str:
if self.using_local_gpu():
try:
from .zerogpu_backend import generate_text
return generate_text(self.model_id, prompt, max_tokens, temperature, self.token)
except ImportError as e:
raise RuntimeError(
"ZeroGPU backend needs torch + transformers. "
f"Install requirements.txt ({e})."
) from e
except RuntimeError:
raise
except Exception as e:
raise self._friendly_error(e) from e
try:
return self._generate_with_client(self.client, prompt, max_tokens, temperature)
except Exception as e:
raise self._friendly_error(e) from e
def _generate_with_client(self, client, prompt: str, max_tokens: int, temperature: float) -> str:
try:
completion = client.chat_completion(
messages=[{"role": "user", "content": prompt}],
model=self.model_id,
max_tokens=max_tokens,
temperature=temperature,
)
text = completion.choices[0].message.content
if text and text.strip():
return text.strip()
except HfHubHTTPError:
raise
except Exception:
pass
response = client.text_generation(
prompt=prompt,
model=self.model_id,
max_new_tokens=max_tokens,
temperature=temperature,
do_sample=temperature > 0,
)
if isinstance(response, str):
return response.strip()
return str(response).strip()
@staticmethod
def _is_unsupported_model_error(e: Exception) -> bool:
if e is None:
return False
msg = str(e).lower()
return (
"not supported by any provider" in msg
or "model_not_supported" in msg
or "no provider" in msg
or ("provider" in msg and "not supported" in msg)
or "availableinferenceproviders" in msg.replace(" ", "")
)
def _friendly_error(self, e: Exception) -> RuntimeError:
msg = str(e)
status = getattr(getattr(e, "response", None), "status_code", None)
if self._is_unsupported_model_error(e) or status == 400 or "bad request" in msg.lower():
return RuntimeError(
f"Model '{self.model_id}' has no Inference Provider on this Space "
"(its page shows empty availableInferenceProviders). It cannot run serverless. "
f"Use {SUPPORTED_EXAMPLES}, and only add an HF_PROVIDER override "
"if you verified that provider serves the chosen model."
)
if status in (401, 403) or ("401" in msg or "403" in msg
or "unauthorized" in msg.lower() or "forbidden" in msg.lower()
or "gated" in msg.lower()):
return RuntimeError(
f"LLM auth error for model '{self.model_id}'. "
"Set a valid HF_TOKEN secret (with access to gated models like Llama/Gemma) "
"or use a public model such as Qwen/Qwen2.5-7B-Instruct."
)
if status == 404 or "404" in msg or "not found" in msg.lower():
return RuntimeError(
f"LLM model '{self.model_id}' not found via Inference Providers. "
"Pick a supported model ID (e.g. Qwen/Qwen2.5-7B-Instruct)."
)
return RuntimeError(f"LLM request failed for model '{self.model_id}': {e}")
def get_coder_prompt(self, problem: str, objective: str) -> str:
return CODER_PROMPT.format(problem=problem, objective=objective)
def get_researcher_prompt(self, problem: str, baseline: str, objective: str, n: int) -> str:
return RESEARCHER_PROMPT.format(problem=problem, baseline=baseline, objective=objective, n=n)
def get_validator_prompt(self, problem: str, objective: str, user_metric: str, metrics_table: str, n: int) -> str:
return VALIDATOR_PROMPT.format(problem=problem, objective=objective, user_metric=user_metric, metrics_table=metrics_table, n=n)
def get_realworld_prompt(self, problem: str, winner_code: str, baseline_code: str, objective: str, n: int) -> str:
return REALWORLD_PROMPT.format(problem=problem, winner_code=winner_code, baseline_code=baseline_code, objective=objective, n=n)
def get_summary_prompt(self, problem: str, objective: str, winner_index: int, metrics_table: str, realworld_results: str) -> str:
return SUMMARY_PROMPT.format(problem=problem, objective=objective, winner_index=winner_index, metrics_table=metrics_table, realworld_results=realworld_results)