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))))