File size: 5,723 Bytes
0dff1a5
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
"""Adapter for a fine-tuned CAMeLBERT-MSA span detector (Subtask 1A) hosted on Hugging Face, plus a simulation mode.

Two ways to get detection spans into the verification pipeline:

* ``entities_to_spans``  turns the output of a Hugging Face *token-classification* call (aggregated entities, or raw
                         ``B-Ayah`` / ``I-Ayah`` / ``B-Hadith`` / ``I-Hadith`` tokens) into ``[{label, start, end, score}]``.
                         The browser page calls the hosted model and hands the entities to ``verify_with_entities``.
* ``simulate_spans``     a stand-in used when no hosted model is configured or reachable: the bundled hybrid detector
                         produces the spans, and they are labelled as simulated so nobody mistakes them for model output.

No model weights live in this repository. Train with ``research/train_detector.py`` (``--model CAMeL-Lab/bert-base-arabic-camelbert-msa``),
push the result to the Hub and set its id in the page configuration (see docs/DEPLOYMENT.md).
"""
from __future__ import annotations

import json
import re
import urllib.error
import urllib.request
from typing import Dict, Iterable, List, Optional

HF_ENDPOINT = "https://router.huggingface.co/hf-inference/models/{model}"   # configurable; not verified against the live service

LABEL_ALIASES = {"ayah": "Ayah", "quran": "Ayah", "label_1": "Ayah", "label_2": "Ayah",
                 "hadith": "Hadith", "label_3": "Hadith", "label_4": "Hadith"}
DEFAULT_MIN_SCORE = 0.5


def _label_of(raw: str) -> Optional[str]:
    name = re.sub(r"^[BI]-", "", str(raw or ""), flags=re.I).strip().lower()
    return LABEL_ALIASES.get(name)


def entities_to_spans(text: str, entities: Iterable[dict], min_score: float = DEFAULT_MIN_SCORE,
                      max_gap: int = 1) -> List[dict]:
    """Merge Hugging Face token-classification output into character spans.

    Works with ``aggregation_strategy="simple"`` output (``entity_group``) and with raw output (``entity`` = ``B-Ayah`` ...).
    Pieces of the same label separated by at most ``max_gap`` characters are joined; low-score pieces are dropped."""
    pieces = []
    for item in entities or []:
        label = _label_of(item.get("entity_group") or item.get("entity"))
        start, end = item.get("start"), item.get("end")
        if label is None or start is None or end is None or end <= start:
            continue
        pieces.append({"label": label, "start": int(start), "end": int(end), "score": float(item.get("score", 1.0)),
                       "begin": str(item.get("entity", "")).upper().startswith("B-")})
    pieces.sort(key=lambda p: (p["start"], p["end"]))
    merged: List[dict] = []
    for piece in pieces:
        last = merged[-1] if merged else None
        if last and last["label"] == piece["label"] and not piece["begin"] and piece["start"] - last["end"] <= max_gap:
            last["end"] = max(last["end"], piece["end"])
            last["scores"].append(piece["score"])
        else:
            merged.append({"label": piece["label"], "start": piece["start"], "end": piece["end"], "scores": [piece["score"]]})
    spans = []
    for item in merged:
        start, end = item["start"], item["end"]
        while start < end and text[start].isspace():
            start += 1
        while end > start and text[end - 1].isspace():
            end -= 1
        score = sum(item["scores"]) / len(item["scores"])
        if end > start and score >= min_score:
            spans.append({"label": item["label"], "start": start, "end": end, "score": round(score, 4)})
    return spans


def simulate_spans(pipeline, text: str) -> List[dict]:
    """Spans from the bundled hybrid detector, flagged as simulated (``score`` is ``None``)."""
    return [{"label": s.label, "start": s.start, "end": s.end, "score": None} for s in pipeline.detect(text)]


def analyze_with_spans(pipeline, text: str, spans: List[dict], engine: str = "camelbert") -> dict:
    """Run verification + correction on externally detected spans and record which engine produced them."""
    result = pipeline.analyze_spans(text, [{"label": s["label"], "start": s["start"], "end": s["end"]} for s in spans])
    scores: Dict[tuple, Optional[float]] = {(s["start"], s["end"]): s.get("score") for s in spans}
    for report in result["spans"]:
        score = scores.get((report["start"], report["end"]))
        report["detection"] = {"backend": engine, "confidence": None if score is None else round(score, 4)}
    result["detector"] = engine
    return result


def query_hosted_model(text: str, model: str, token: str = "", endpoint: str = HF_ENDPOINT, timeout: float = 30.0) -> List[dict]:
    """Call a hosted token-classification model (server-side use: the local Gradio app and tests with a mock server).

    Returns ``entities_to_spans`` output. Raises ``RuntimeError`` with a short message on any failure so the caller can fall
    back to the simulation."""
    headers = {"Content-Type": "application/json"}
    if token:
        headers["Authorization"] = f"Bearer {token}"
    body = json.dumps({"inputs": text, "parameters": {"aggregation_strategy": "simple"}}).encode("utf-8")
    request = urllib.request.Request(endpoint.format(model=model), data=body, headers=headers, method="POST")
    try:
        with urllib.request.urlopen(request, timeout=timeout) as response:
            payload = json.loads(response.read().decode("utf-8"))
    except (urllib.error.URLError, TimeoutError, json.JSONDecodeError) as exc:
        raise RuntimeError("hosted model unavailable") from exc
    if not isinstance(payload, list):
        raise RuntimeError("unexpected response from the hosted model")
    return entities_to_spans(text, payload)