ayushKishor commited on
Commit
2fe3d02
·
1 Parent(s): 15f6b5f

Support NVIDIA and Mistral provider setup

Browse files
Files changed (2) hide show
  1. mp1/pluto/dispatcher.py +10 -3
  2. mp1/pluto/modes.py +52 -3
mp1/pluto/dispatcher.py CHANGED
@@ -28,6 +28,10 @@ NVIDIA_RERANK_URLS = (
28
  "https://ai.api.nvidia.com/v1/retrieval/nvidia/reranking",
29
  "https://integrate.api.nvidia.com/v1/ranking",
30
  )
 
 
 
 
31
 
32
  # Groq has strict RPM limits; semaphore prevents parallel-call overflow.
33
  _groq_semaphore = threading.BoundedSemaphore(2)
@@ -234,6 +238,9 @@ def _call_nvidia(cfg: ModeConfig, prompt: str) -> str:
234
  env_var, api_key = _resolve_nvidia_api_key(cfg.model_id)
235
  if not api_key:
236
  raise ValueError(f"{env_var} or NVIDIA_API_KEY not set")
 
 
 
237
 
238
  prefix = str(prompt)[:120]
239
  use_reasoning = any(
@@ -250,7 +257,7 @@ def _call_nvidia(cfg: ModeConfig, prompt: str) -> str:
250
  if "nano" in cfg.model_id and use_reasoning:
251
  payload["nvidia"] = {"reasoning": True}
252
 
253
- for attempt in range(6):
254
  try:
255
  response = requests.post(
256
  NVIDIA_CHAT_URL,
@@ -259,7 +266,7 @@ def _call_nvidia(cfg: ModeConfig, prompt: str) -> str:
259
  "Content-Type": "application/json",
260
  },
261
  json=payload,
262
- timeout=(20, 120),
263
  )
264
  if response.status_code == 200:
265
  data = response.json()
@@ -289,7 +296,7 @@ def _call_nvidia(cfg: ModeConfig, prompt: str) -> str:
289
  "504",
290
  ]
291
  )
292
- if (is_rate or is_transient) and attempt < 5:
293
  delay = min(60.0, 5.0 * (1.5**attempt))
294
  reason = "rate-limit" if is_rate else "transient error"
295
  print(f" [RETRY] NVIDIA {reason} — waiting {delay:.1f}s (attempt {attempt + 1})")
 
28
  "https://ai.api.nvidia.com/v1/retrieval/nvidia/reranking",
29
  "https://integrate.api.nvidia.com/v1/ranking",
30
  )
31
+ NVIDIA_MAX_RETRIES_WITH_FALLBACK = 2
32
+ NVIDIA_MAX_RETRIES_DEFAULT = 6
33
+ NVIDIA_TIMEOUT_WITH_FALLBACK = (10, 45)
34
+ NVIDIA_TIMEOUT_DEFAULT = (20, 120)
35
 
36
  # Groq has strict RPM limits; semaphore prevents parallel-call overflow.
37
  _groq_semaphore = threading.BoundedSemaphore(2)
 
238
  env_var, api_key = _resolve_nvidia_api_key(cfg.model_id)
239
  if not api_key:
240
  raise ValueError(f"{env_var} or NVIDIA_API_KEY not set")
241
+ has_fallback = bool(os.getenv("GROQ_API_KEY", "") or os.getenv("MISTRAL_API_KEY", ""))
242
+ max_retries = NVIDIA_MAX_RETRIES_WITH_FALLBACK if has_fallback else NVIDIA_MAX_RETRIES_DEFAULT
243
+ request_timeout = NVIDIA_TIMEOUT_WITH_FALLBACK if has_fallback else NVIDIA_TIMEOUT_DEFAULT
244
 
245
  prefix = str(prompt)[:120]
246
  use_reasoning = any(
 
257
  if "nano" in cfg.model_id and use_reasoning:
258
  payload["nvidia"] = {"reasoning": True}
259
 
260
+ for attempt in range(max_retries):
261
  try:
262
  response = requests.post(
263
  NVIDIA_CHAT_URL,
 
266
  "Content-Type": "application/json",
267
  },
268
  json=payload,
269
+ timeout=request_timeout,
270
  )
271
  if response.status_code == 200:
272
  data = response.json()
 
296
  "504",
297
  ]
298
  )
299
+ if (is_rate or is_transient) and attempt < max_retries - 1:
300
  delay = min(60.0, 5.0 * (1.5**attempt))
301
  reason = "rate-limit" if is_rate else "transient error"
302
  print(f" [RETRY] NVIDIA {reason} — waiting {delay:.1f}s (attempt {attempt + 1})")
mp1/pluto/modes.py CHANGED
@@ -32,7 +32,7 @@ class ModeConfig:
32
  temperature: float
33
  max_tokens: int
34
  compute_profile: str
35
- provider: str # "groq" | "mistral"
36
 
37
  def to_log_dict(self) -> dict:
38
  return {
@@ -58,7 +58,7 @@ def _build_registry() -> dict[str, ModeConfig]:
58
  Embedding + reranking are handled separately in embedder.py and dispatcher.py
59
  (they use /v1/embeddings and scoring endpoints, not chat completions).
60
 
61
- Fallback: if NVIDIA_API_KEY absent, fall back to Groq (same model sizes).
62
  """
63
  # Check for any NVIDIA key
64
  nvidia_keys = [
@@ -68,6 +68,7 @@ def _build_registry() -> dict[str, ModeConfig]:
68
  ]
69
  nvidia_ready = any(os.getenv(k) for k in nvidia_keys)
70
  groq_key = os.getenv("GROQ_API_KEY", "").strip()
 
71
 
72
  if nvidia_ready:
73
  return {
@@ -157,9 +158,57 @@ def _build_registry() -> dict[str, ModeConfig]:
157
  provider="groq",
158
  ),
159
  }
 
 
160
  return _build_unconfigured_registry()
161
 
162
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
163
  def _build_unconfigured_registry() -> dict[str, ModeConfig]:
164
  """Return placeholder modes so imports work without provider credentials."""
165
  return {
@@ -210,7 +259,7 @@ MODE_REGISTRY: dict[str, ModeConfig] = _build_registry()
210
 
211
 
212
  def _missing_provider_error() -> EnvironmentError:
213
- return EnvironmentError("Neither NVIDIA_API_KEY nor GROQ_API_KEY is set.")
214
 
215
 
216
  def _is_unconfigured() -> bool:
 
32
  temperature: float
33
  max_tokens: int
34
  compute_profile: str
35
+ provider: str # "nvidia" | "groq" | "mistral"
36
 
37
  def to_log_dict(self) -> dict:
38
  return {
 
58
  Embedding + reranking are handled separately in embedder.py and dispatcher.py
59
  (they use /v1/embeddings and scoring endpoints, not chat completions).
60
 
61
+ Fallback: if NVIDIA_API_KEY absent, fall back to Groq or Mistral.
62
  """
63
  # Check for any NVIDIA key
64
  nvidia_keys = [
 
68
  ]
69
  nvidia_ready = any(os.getenv(k) for k in nvidia_keys)
70
  groq_key = os.getenv("GROQ_API_KEY", "").strip()
71
+ mistral_key = os.getenv("MISTRAL_API_KEY", "").strip()
72
 
73
  if nvidia_ready:
74
  return {
 
158
  provider="groq",
159
  ),
160
  }
161
+ if mistral_key:
162
+ return _build_mistral_registry()
163
  return _build_unconfigured_registry()
164
 
165
 
166
+ def _build_mistral_registry() -> dict[str, ModeConfig]:
167
+ """Use Mistral for every mode when it is the only configured chat provider."""
168
+ return {
169
+ "MODE_QUICK": ModeConfig(
170
+ mode_name="MODE_QUICK",
171
+ model_id="mistral-small-latest",
172
+ temperature=0.1,
173
+ max_tokens=1024,
174
+ compute_profile="fallback",
175
+ provider="mistral",
176
+ ),
177
+ "MODE_REASONING": ModeConfig(
178
+ mode_name="MODE_REASONING",
179
+ model_id="mistral-small-latest",
180
+ temperature=0.3,
181
+ max_tokens=4096,
182
+ compute_profile="fallback",
183
+ provider="mistral",
184
+ ),
185
+ "MODE_VISION": ModeConfig(
186
+ mode_name="MODE_VISION",
187
+ model_id="mistral-small-latest",
188
+ temperature=0.1,
189
+ max_tokens=4096,
190
+ compute_profile="fallback",
191
+ provider="mistral",
192
+ ),
193
+ "MODE_ULTRA": ModeConfig(
194
+ mode_name="MODE_ULTRA",
195
+ model_id="mistral-small-latest",
196
+ temperature=0.2,
197
+ max_tokens=4096,
198
+ compute_profile="fallback",
199
+ provider="mistral",
200
+ ),
201
+ "MODE_GEMINI": ModeConfig(
202
+ mode_name="MODE_GEMINI",
203
+ model_id="mistral-small-latest",
204
+ temperature=0.0,
205
+ max_tokens=4096,
206
+ compute_profile="fallback",
207
+ provider="mistral",
208
+ ),
209
+ }
210
+
211
+
212
  def _build_unconfigured_registry() -> dict[str, ModeConfig]:
213
  """Return placeholder modes so imports work without provider credentials."""
214
  return {
 
259
 
260
 
261
  def _missing_provider_error() -> EnvironmentError:
262
+ return EnvironmentError("None of NVIDIA_API_KEY, GROQ_API_KEY, or MISTRAL_API_KEY is set.")
263
 
264
 
265
  def _is_unconfigured() -> bool: