Srishti280992 commited on
Commit
a57fb0d
·
verified ·
1 Parent(s): c1b0f6d

Update app.py

Browse files
Files changed (1) hide show
  1. app.py +38 -6
app.py CHANGED
@@ -144,12 +144,18 @@ class OptionalHFInferenceClient:
144
  detailed errors instead of hiding provider/model failures.
145
  """
146
 
147
- def __init__(self, model_id: str, token: str | None = None) -> None:
 
 
 
 
 
148
  """Initialize the optional Hugging Face inference adapter."""
149
 
150
  from huggingface_hub import InferenceClient
151
 
152
  self.model_id = model_id.strip().strip('"').strip("'").strip()
 
153
 
154
  api_key = (
155
  token
@@ -165,10 +171,17 @@ class OptionalHFInferenceClient:
165
  "and add it as a Space secret named HF_TOKEN."
166
  )
167
 
168
- self.client = InferenceClient(
169
- model=self.model_id,
170
- api_key=api_key,
171
- )
 
 
 
 
 
 
 
172
 
173
  def generate(
174
  self,
@@ -376,6 +389,12 @@ def build_optional_model_client() -> tuple[Any | None, AppModelStatus]:
376
  or os.getenv("HUGGINGFACEHUB_API_TOKEN")
377
  )
378
 
 
 
 
 
 
 
379
  if not raw_model_id:
380
  return None, AppModelStatus(
381
  enabled=False,
@@ -399,7 +418,11 @@ def build_optional_model_client() -> tuple[Any | None, AppModelStatus]:
399
  )
400
 
401
  try:
402
- client = OptionalHFInferenceClient(model_id=model_id, token=token)
 
 
 
 
403
  return client, AppModelStatus(
404
  enabled=True,
405
  model_id=model_id,
@@ -430,6 +453,13 @@ def model_health_check_ui() -> str:
430
  "Check that WORLDSMITHAI_MODEL_ID is set in Space settings."
431
  )
432
 
 
 
 
 
 
 
 
433
  try:
434
  response = MODEL_CLIENT.generate(
435
  prompt='Return exactly this JSON object and nothing else: {"ok": true}',
@@ -450,6 +480,7 @@ def model_health_check_ui() -> str:
450
  return (
451
  "Model client is configured and returned text.\n\n"
452
  f"Model id: {MODEL_STATUS.model_id}\n\n"
 
453
  f"Response:\n{response}"
454
  )
455
  except Exception as exc:
@@ -457,6 +488,7 @@ def model_health_check_ui() -> str:
457
  return (
458
  "Model client is configured but inference failed.\n\n"
459
  f"Model id: {MODEL_STATUS.model_id}\n\n"
 
460
  f"Error type: {exc.__class__.__name__}\n"
461
  f"Error:\n{exc}"
462
  )
 
144
  detailed errors instead of hiding provider/model failures.
145
  """
146
 
147
+ def __init__(
148
+ self,
149
+ model_id: str,
150
+ token: str | None = None,
151
+ provider: str | None = None,
152
+ ) -> None:
153
  """Initialize the optional Hugging Face inference adapter."""
154
 
155
  from huggingface_hub import InferenceClient
156
 
157
  self.model_id = model_id.strip().strip('"').strip("'").strip()
158
+ self.provider = None if provider is None or not str(provider).strip() else str(provider).strip()
159
 
160
  api_key = (
161
  token
 
171
  "and add it as a Space secret named HF_TOKEN."
172
  )
173
 
174
+ if self.provider:
175
+ self.client = InferenceClient(
176
+ model=self.model_id,
177
+ provider=self.provider,
178
+ api_key=api_key,
179
+ )
180
+ else:
181
+ self.client = InferenceClient(
182
+ model=self.model_id,
183
+ api_key=api_key,
184
+ )
185
 
186
  def generate(
187
  self,
 
389
  or os.getenv("HUGGINGFACEHUB_API_TOKEN")
390
  )
391
 
392
+ provider = (
393
+ os.getenv("WORLDSMITHAI_PROVIDER")
394
+ or os.getenv("HF_INFERENCE_PROVIDER")
395
+ or None
396
+ )
397
+
398
  if not raw_model_id:
399
  return None, AppModelStatus(
400
  enabled=False,
 
418
  )
419
 
420
  try:
421
+ client = OptionalHFInferenceClient(
422
+ model_id=model_id,
423
+ token=token,
424
+ provider=provider,
425
+ )
426
  return client, AppModelStatus(
427
  enabled=True,
428
  model_id=model_id,
 
453
  "Check that WORLDSMITHAI_MODEL_ID is set in Space settings."
454
  )
455
 
456
+ provider_label = (
457
+ getattr(MODEL_CLIENT, "provider", None)
458
+ or os.getenv("WORLDSMITHAI_PROVIDER")
459
+ or os.getenv("HF_INFERENCE_PROVIDER")
460
+ or "auto"
461
+ )
462
+
463
  try:
464
  response = MODEL_CLIENT.generate(
465
  prompt='Return exactly this JSON object and nothing else: {"ok": true}',
 
480
  return (
481
  "Model client is configured and returned text.\n\n"
482
  f"Model id: {MODEL_STATUS.model_id}\n\n"
483
+ f"Provider: {provider_label}\n\n"
484
  f"Response:\n{response}"
485
  )
486
  except Exception as exc:
 
488
  return (
489
  "Model client is configured but inference failed.\n\n"
490
  f"Model id: {MODEL_STATUS.model_id}\n\n"
491
+ f"Provider: {provider_label}\n\n"
492
  f"Error type: {exc.__class__.__name__}\n"
493
  f"Error:\n{exc}"
494
  )