#!/usr/bin/env python3 """Multi-source PubMedBERT-NLI verification endpoint.""" from typing import Any import torch from fastapi import HTTPException from logger import setup_logger from models.loader import load_nli_model, load_nli_tokenizer logger = setup_logger("endpoints.verify") _CANONICAL = { "entailment": "entailment", "entails": "entailment", "neutral": "neutral", "contradiction": "contradiction", "contradicts": "contradiction", } def _label_map(model) -> dict[int, str]: raw = getattr(model.config, "id2label", {}) or {} mapped: dict[int, str] = {} for key, value in raw.items(): label = _CANONICAL.get(str(value).strip().lower()) if label: mapped[int(key)] = label if set(mapped.values()) != {"entailment", "neutral", "contradiction"}: raise RuntimeError(f"NLI id2label incompatible: {raw}") return mapped def _nli_score(claim: str, evidence: str) -> dict[str, float | str]: if not evidence.strip(): return { "entailment": 0.0, "neutral": 1.0, "contradiction": 0.0, "verdict": "neutral", } tokenizer = load_nli_tokenizer() model = load_nli_model() inputs = tokenizer( evidence, claim, return_tensors="pt", truncation=True, max_length=512, padding=True, ) device = next(model.parameters()).device inputs = {key: value.to(device) for key, value in inputs.items()} with torch.no_grad(): probs = torch.softmax(model(**inputs).logits, dim=-1)[0].detach().cpu() id2label = _label_map(model) scores = {"entailment": 0.0, "neutral": 0.0, "contradiction": 0.0} for index, probability in enumerate(probs.tolist()): scores[id2label[index]] = round(float(probability), 6) verdict = max(scores, key=scores.get) return {**scores, "verdict": verdict} def _aggregate_sources( claim: str, sources: list[dict[str, Any]], entailment_threshold: float, contradiction_threshold: float, ) -> dict[str, Any]: per_source = [] for source in sources: text = str(source.get("text", "")).strip() if not text: continue scores = _nli_score(claim, text) per_source.append({ "pmid": str(source.get("pmid", "")), **scores, }) if not per_source: return { "verdict": "NOT_SUPPORTED", "reason": "aucune source exploitable", "entailment": 0.0, "neutral": 1.0, "contradiction": 0.0, "sourceVerdicts": [], } best_entailment = max(per_source, key=lambda item: item["entailment"]) best_contradiction = max(per_source, key=lambda item: item["contradiction"]) ent = float(best_entailment["entailment"]) con = float(best_contradiction["contradiction"]) neu = min(float(item["neutral"]) for item in per_source) if ent >= entailment_threshold and con >= contradiction_threshold: verdict = "PARTIAL" reason = "sources discordantes: soutien et contradiction détectés" elif ent >= entailment_threshold and ent >= con: verdict = "SUPPORTED" reason = f"soutenu par PMID {best_entailment.get('pmid') or '?'}" elif con >= contradiction_threshold and con > ent: verdict = "CONTRADICTED" reason = f"contredit par PMID {best_contradiction.get('pmid') or '?'}" elif ent >= max(0.45, entailment_threshold * 0.70): verdict = "PARTIAL" reason = "soutien partiel sous le seuil principal" else: verdict = "NOT_SUPPORTED" reason = "aucune source ne soutient suffisamment le claim" return { "verdict": verdict, "reason": reason, "entailment": ent, "neutral": neu, "contradiction": con, "bestSupportPmid": best_entailment.get("pmid"), "bestContradictionPmid": best_contradiction.get("pmid"), "sourceVerdicts": per_source, } async def post_verify(request: dict) -> dict: if "items" in request: items = request.get("items", []) entailment_threshold = float(request.get("entailment_threshold", 0.65)) contradiction_threshold = float(request.get("contradiction_threshold", 0.50)) verdicts = [] for item in items: item_id = item.get("id", 0) claim = str(item.get("claim", "")).strip() if not claim: verdicts.append({ "id": item_id, "verdict": "NOT_SUPPORTED", "reason": "claim vide", }) continue try: result = _aggregate_sources( claim, list(item.get("sources", [])), entailment_threshold, contradiction_threshold, ) verdicts.append({"id": item_id, **result}) except Exception as exc: logger.error("NLI error item %s: %s", item_id, exc) verdicts.append({ "id": item_id, "verdict": "NOT_SUPPORTED", "reason": str(exc), }) return {"verdicts": verdicts} claim = str(request.get("claim", "")).strip() abstract = str(request.get("abstract", "")).strip() if not claim: raise HTTPException(status_code=422, detail="claim required") if not abstract: raise HTTPException(status_code=422, detail="abstract required") try: scores = _nli_score(claim, abstract) return { "claim": claim, "verdict": scores["verdict"], "scores": { "entailment": scores["entailment"], "neutral": scores["neutral"], "contradiction": scores["contradiction"], }, "supported": scores["verdict"] == "entailment", } except Exception as exc: logger.error("Verify error: %s", exc) raise HTTPException(status_code=500, detail=str(exc))