probabot / src /entailment.py
AlokBharadwaj's picture
Toggle switch between GPU and API inference
54b4fa6
Raw History Blame
11.4 kB
import os
import json
import time
from concurrent.futures import ThreadPoolExecutor
import requests
import numpy as np
import torch
try:
import spaces
except Exception:
class _SpacesShim:
@staticmethod
def GPU(fn=None, **kwargs):
if fn is not None:
return fn
def decorator(f):
return f
return decorator
spaces = _SpacesShim()
from src.config import (
ENTAILMENT_MODEL_NAME,
ENTAILMENT_BATCH_SIZE,
ENTAILMENT_USE_API,
ENTAILMENT_API_MODEL,
ENTAILMENT_API_WORKERS,
)
_deberta_tokenizer = None
_deberta_model = None
# ---------------------------------------------------------------------------
# Label mapping (model-agnostic): map a label STRING to our convention
# 0 = contradiction, 1 = neutral, 2 = entailment
# This works regardless of a model's internal label ordering.
# ---------------------------------------------------------------------------
def _label_to_id(label):
l = str(label).lower()
if "contradict" in l:
return 0
if "neutral" in l:
return 1
if "entail" in l:
return 2
# Fallback for generic "LABEL_0/1/2" (standard MNLI: 0=contra,1=neutral,2=entail)
if l.endswith("0"):
return 0
if l.endswith("1"):
return 1
if l.endswith("2"):
return 2
return 1 # safest default: neutral
def _parse_classification(data):
"""Extract the argmax label id from a text-classification API response.
Accepts [{'label','score'}, ...] or nested [[...]]."""
if isinstance(data, (bytes, bytearray, str)):
data = json.loads(data)
if isinstance(data, list) and data and isinstance(data[0], list):
data = data[0]
if not isinstance(data, list) or not data:
return 1
best = max(data, key=lambda d: d.get("score", 0.0))
return _label_to_id(best.get("label", ""))
_API_URL = "https://api-inference.huggingface.co/models/" + ENTAILMENT_API_MODEL
_warned_once = {"done": False}
def _api_classify_pair(token, premise, hypothesis):
"""One serverless text-classification call for a (premise, hypothesis) pair.
Uses the raw Inference API with the sentence-pair payload."""
headers = {"Authorization": f"Bearer {token}"} if token else {}
payload = {
"inputs": {"text": premise, "text_pair": hypothesis},
"options": {"wait_for_model": True},
}
for attempt in range(5):
try:
r = requests.post(_API_URL, headers=headers, json=payload, timeout=60)
if r.status_code in (429, 503):
time.sleep(min(2 ** attempt * 2, 20))
continue
if r.status_code >= 400:
if not _warned_once["done"]:
print(f"[entailment-api] HTTP {r.status_code}: {r.text[:300]}", flush=True)
_warned_once["done"] = True
return 1
return _parse_classification(r.json())
except Exception as e:
if not _warned_once["done"]:
print(f"[entailment-api] request error: {e}", flush=True)
_warned_once["done"] = True
time.sleep(min(2 ** attempt, 8))
return 1
def api_check_implication(pairs, client=None):
"""
Entailment via HF serverless Inference API (parallel requests).
Returns list[int] in {0,1,2}, aligned with `pairs`.
"""
if not pairs:
return []
token = getattr(client, "token", None) or os.environ.get("HF_TOKEN")
total = len(pairs)
print(f"[entailment-api] Classifying {total} pairs via Inference API "
f"(model={ENTAILMENT_API_MODEL}, workers={ENTAILMENT_API_WORKERS})...", flush=True)
def work(pair):
return _api_classify_pair(token, pair[0], pair[1])
with ThreadPoolExecutor(max_workers=ENTAILMENT_API_WORKERS) as ex:
results = list(ex.map(work, pairs))
print(f"[entailment-api] Done ({total} pairs).", flush=True)
return results
@spaces.GPU
def batch_check_implication(pairs):
"""
pairs: list of (premise, hypothesis) str tuples
Returns list of int (0=contradiction, 1=neutral, 2=entailment).
DeBERTa is loaded lazily inside this GPU-bound function.
"""
global _deberta_tokenizer, _deberta_model
if not pairs:
return []
if _deberta_model is None:
from transformers import AutoModelForSequenceClassification, AutoTokenizer
device = "cuda" if torch.cuda.is_available() else "cpu"
print(f"[entailment] Loading DeBERTa onto {device} (first GPU call)...", flush=True)
_deberta_tokenizer = AutoTokenizer.from_pretrained(ENTAILMENT_MODEL_NAME)
_deberta_model = AutoModelForSequenceClassification.from_pretrained(
ENTAILMENT_MODEL_NAME
).to(device)
_deberta_model.eval()
print("[entailment] DeBERTa loaded.", flush=True)
device = next(_deberta_model.parameters()).device
results = []
total = len(pairs)
n_batches = (total + ENTAILMENT_BATCH_SIZE - 1) // ENTAILMENT_BATCH_SIZE
print(f"[entailment] Running entailment on {total} pairs in {n_batches} batches "
f"of {ENTAILMENT_BATCH_SIZE} (device={device})...", flush=True)
for bi, i in enumerate(range(0, total, ENTAILMENT_BATCH_SIZE), 1):
batch = pairs[i : i + ENTAILMENT_BATCH_SIZE]
premises = [p for p, _ in batch]
hypotheses = [h for _, h in batch]
inputs = _deberta_tokenizer(
premises, hypotheses, padding=True, truncation=True, return_tensors="pt"
).to(device)
with torch.no_grad():
logits = _deberta_model(**inputs).logits
preds = torch.argmax(logits, dim=1).tolist()
results.extend(preds)
print(f"[entailment] batch {bi}/{n_batches} done "
f"({min(i + ENTAILMENT_BATCH_SIZE, total)}/{total} pairs)", flush=True)
print(f"[entailment] Entailment complete ({total} pairs).", flush=True)
return results
def get_semantic_ids(strings_list, strict_entailment=False):
"""
3-phase:
1. Gather all ordered pairs (both directions).
2. One batched GPU call to batch_check_implication.
3. Pure-Python greedy clustering via lookup dict.
"""
n = len(strings_list)
if n == 0:
return []
if n == 1:
return [0]
# Phase 1
pairs = []
pair_indices = []
for i in range(n):
for j in range(n):
if i != j:
pairs.append((strings_list[i], strings_list[j]))
pair_indices.append((i, j))
# Phase 2
raw = batch_check_implication(pairs)
# Phase 3
impl = {(i, j): r for (i, j), r in zip(pair_indices, raw)}
def are_equivalent(i, j):
i1 = impl.get((i, j), 1)
i2 = impl.get((j, i), 1)
if strict_entailment:
return i1 == 2 and i2 == 2
return 0 not in [i1, i2] and [i1, i2] != [1, 1]
semantic_ids = [-1] * n
next_id = 0
for i in range(n):
if semantic_ids[i] == -1:
semantic_ids[i] = next_id
for j in range(i + 1, n):
if semantic_ids[j] == -1 and are_equivalent(i, j):
semantic_ids[j] = next_id
next_id += 1
return semantic_ids
def _cluster_from_impl(n, impl, strict_entailment):
"""Pure-Python greedy clustering given an implication lookup dict."""
def are_equivalent(i, j):
i1 = impl.get((i, j), 1)
i2 = impl.get((j, i), 1)
if strict_entailment:
return i1 == 2 and i2 == 2
return 0 not in [i1, i2] and [i1, i2] != [1, 1]
semantic_ids = [-1] * n
next_id = 0
for i in range(n):
if semantic_ids[i] == -1:
semantic_ids[i] = next_id
for j in range(i + 1, n):
if semantic_ids[j] == -1 and are_equivalent(i, j):
semantic_ids[j] = next_id
next_id += 1
return semantic_ids
def get_semantic_ids_batched(groups, client=None, strict_entailment=False, use_api=None):
"""
Cluster MANY groups of strings using a SINGLE pass of entailment checks.
groups: list of list[str]
client: HF InferenceClient (required when using the API path)
use_api: True β†’ Inference API, False β†’ local GPU/CPU, None β†’ config default
Returns: list of list[int] β€” semantic ids per group, aligned with `groups`.
Pattern 3: gather pairs (both directions) from ALL groups β†’ one entailment
pass (serverless API OR one batched GPU call) β†’ pure-Python greedy
clustering per group via lookup.
"""
if use_api is None:
use_api = ENTAILMENT_USE_API
all_pairs = []
group_meta = [] # (offset, n_uniq, local_indices, orig_to_uniq)
for strings in groups:
# Deduplicate identical answers (normalized) β€” these are trivially
# equivalent, so we only run entailment on the UNIQUE strings and
# expand cluster ids back afterwards. Big reduction in entailment calls.
uniq = []
norm_to_uniq = {}
orig_to_uniq = []
for s in strings:
key = s.strip().lower()
if key not in norm_to_uniq:
norm_to_uniq[key] = len(uniq)
uniq.append(s)
orig_to_uniq.append(norm_to_uniq[key])
m = len(uniq)
offset = len(all_pairs)
local_indices = []
for i in range(m):
for j in range(m):
if i != j:
all_pairs.append((uniq[i], uniq[j]))
local_indices.append((i, j))
group_meta.append((offset, m, local_indices, orig_to_uniq))
print(f"[entailment] {len(all_pairs)} pairs after dedup "
f"(across {len(groups)} groups)", flush=True)
# One entailment pass for the entire pipeline
if use_api and client is not None:
raw = api_check_implication(all_pairs, client)
else:
# Retry transient GPU faults (e.g. ECC errors on a bad ZeroGPU node) β€”
# each call re-acquires a fresh worker/GPU.
raw = None
last_err = None
for attempt in range(3):
try:
raw = batch_check_implication(all_pairs)
break
except Exception as e:
last_err = e
print(f"[entailment] GPU attempt {attempt+1}/3 failed: {e}", flush=True)
time.sleep(3)
if raw is None:
raise RuntimeError(f"Entailment GPU call failed after retries: {last_err}")
results = []
for offset, m, local_indices, orig_to_uniq in group_meta:
n_orig = len(orig_to_uniq)
if m == 0:
results.append([])
continue
if m == 1:
# all originals identical β†’ one cluster
results.append([0] * n_orig)
continue
impl = {
(i, j): raw[offset + k] for k, (i, j) in enumerate(local_indices)
}
uniq_ids = _cluster_from_impl(m, impl, strict_entailment)
# expand unique-cluster ids back to every original string
results.append([uniq_ids[orig_to_uniq[t]] for t in range(n_orig)])
return results
def cluster_assignment_entropy(semantic_ids):
if not semantic_ids:
return 0.0
counts = np.bincount(semantic_ids)
probs = counts / counts.sum()
# guard log(0)
return float(-np.sum(probs * np.log(np.where(probs > 0, probs, 1.0))))