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)