Spaces:
Paused
Paused
Download src/entailment.py from cryotud-AJ/probabot: direct link, hf CLI and curl.
- Browser
- Download file 11.4 kB
-
https://huggingface.co/spaces/cryotud-AJ/probabot/resolve/6d6c490ad7b5f432349fcf096d2bac0bf28b456a/src/entailment.py
- Command line
-
hf download hf://spaces/cryotud-AJ/probabot@6d6c490ad7b5f432349fcf096d2bac0bf28b456a/src/entailment.py
-
curl -L -o entailment.py https://huggingface.co/spaces/cryotud-AJ/probabot/resolve/6d6c490ad7b5f432349fcf096d2bac0bf28b456a/src/entailment.py
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: | |
| 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 | |
| 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)))) | |