AlokBharadwaj commited on
Commit
85acbc9
·
1 Parent(s): dd0651d

fix inference client bug

Browse files
Files changed (1) hide show
  1. src/entailment.py +31 -14
src/entailment.py CHANGED
@@ -1,7 +1,9 @@
 
1
  import json
2
  import time
3
  from concurrent.futures import ThreadPoolExecutor
4
 
 
5
  import numpy as np
6
  import torch
7
 
@@ -67,37 +69,52 @@ def _parse_classification(data):
67
  return _label_to_id(best.get("label", ""))
68
 
69
 
70
- def _api_classify_pair(client, model, premise, hypothesis):
71
- """One serverless text-classification call for a (premise, hypothesis) pair."""
72
- payload = {"inputs": {"text": premise, "text_pair": hypothesis}}
73
- for attempt in range(4):
 
 
 
 
 
 
 
 
 
74
  try:
75
- raw = client.post(json=payload, model=model, task="text-classification")
76
- return _parse_classification(raw)
77
- except Exception as e:
78
- msg = str(e).lower()
79
- # 503 = model loading (cold start); 429 = rate limit → wait & retry
80
- if any(x in msg for x in ("503", "loading", "429", "rate limit", "too many requests")):
81
  time.sleep(min(2 ** attempt * 2, 20))
82
  continue
83
- print(f"[entailment-api] classify error: {e}", flush=True)
84
- return 1
 
 
 
 
 
 
 
 
 
85
  return 1
86
 
87
 
88
- def api_check_implication(pairs, client):
89
  """
90
  Entailment via HF serverless Inference API (parallel requests).
91
  Returns list[int] in {0,1,2}, aligned with `pairs`.
92
  """
93
  if not pairs:
94
  return []
 
95
  total = len(pairs)
96
  print(f"[entailment-api] Classifying {total} pairs via Inference API "
97
  f"(model={ENTAILMENT_API_MODEL}, workers={ENTAILMENT_API_WORKERS})...", flush=True)
98
 
99
  def work(pair):
100
- return _api_classify_pair(client, ENTAILMENT_API_MODEL, pair[0], pair[1])
101
 
102
  with ThreadPoolExecutor(max_workers=ENTAILMENT_API_WORKERS) as ex:
103
  results = list(ex.map(work, pairs))
 
1
+ import os
2
  import json
3
  import time
4
  from concurrent.futures import ThreadPoolExecutor
5
 
6
+ import requests
7
  import numpy as np
8
  import torch
9
 
 
69
  return _label_to_id(best.get("label", ""))
70
 
71
 
72
+ _API_URL = "https://api-inference.huggingface.co/models/" + ENTAILMENT_API_MODEL
73
+ _warned_once = {"done": False}
74
+
75
+
76
+ def _api_classify_pair(token, premise, hypothesis):
77
+ """One serverless text-classification call for a (premise, hypothesis) pair.
78
+ Uses the raw Inference API with the sentence-pair payload."""
79
+ headers = {"Authorization": f"Bearer {token}"} if token else {}
80
+ payload = {
81
+ "inputs": {"text": premise, "text_pair": hypothesis},
82
+ "options": {"wait_for_model": True},
83
+ }
84
+ for attempt in range(5):
85
  try:
86
+ r = requests.post(_API_URL, headers=headers, json=payload, timeout=60)
87
+ if r.status_code in (429, 503):
 
 
 
 
88
  time.sleep(min(2 ** attempt * 2, 20))
89
  continue
90
+ if r.status_code >= 400:
91
+ if not _warned_once["done"]:
92
+ print(f"[entailment-api] HTTP {r.status_code}: {r.text[:300]}", flush=True)
93
+ _warned_once["done"] = True
94
+ return 1
95
+ return _parse_classification(r.json())
96
+ except Exception as e:
97
+ if not _warned_once["done"]:
98
+ print(f"[entailment-api] request error: {e}", flush=True)
99
+ _warned_once["done"] = True
100
+ time.sleep(min(2 ** attempt, 8))
101
  return 1
102
 
103
 
104
+ def api_check_implication(pairs, client=None):
105
  """
106
  Entailment via HF serverless Inference API (parallel requests).
107
  Returns list[int] in {0,1,2}, aligned with `pairs`.
108
  """
109
  if not pairs:
110
  return []
111
+ token = getattr(client, "token", None) or os.environ.get("HF_TOKEN")
112
  total = len(pairs)
113
  print(f"[entailment-api] Classifying {total} pairs via Inference API "
114
  f"(model={ENTAILMENT_API_MODEL}, workers={ENTAILMENT_API_WORKERS})...", flush=True)
115
 
116
  def work(pair):
117
+ return _api_classify_pair(token, pair[0], pair[1])
118
 
119
  with ThreadPoolExecutor(max_workers=ENTAILMENT_API_WORKERS) as ex:
120
  results = list(ex.map(work, pairs))