Spaces:
Sleeping
Sleeping
Michael Arana
Add ZeroGPU backend so any HF model runs on Space GPU (default XHToken/Spark-X2.5-4B)
042d836 Download src/llm.py from gamer729/algorithm-finder: direct link, hf CLI and curl.
- Browser
- Download file 6.47 kB
-
https://huggingface.co/spaces/gamer729/algorithm-finder/resolve/main/src/llm.py
- Command line
-
hf download hf://spaces/gamer729/algorithm-finder/src/llm.py
-
curl -L -o llm.py https://huggingface.co/spaces/gamer729/algorithm-finder/resolve/main/src/llm.py
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") | |
| 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() | |
| 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) | |