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)