Spaces:
Sleeping
Sleeping
File size: 6,470 Bytes
5b33b3c 2898fad 5b33b3c 042d836 2898fad 5e563eb 2898fad 5b33b3c 042d836 2898fad 042d836 2898fad 042d836 2898fad 5b33b3c 042d836 2898fad 5e563eb 2898fad 5e563eb 2898fad 5e563eb 2898fad 5e563eb 2898fad 5b33b3c | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 | 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)
|