adminaikka's picture
Fix span merging: group by entity type before merging (pipelines interleave)
06f10e3 verified
Raw
History Blame Contribute Delete
5.81 kB
"""
aikka OpenMed NER — Gradio Space (ZeroGPU).
Two OpenMed token-classification pipelines (drugs + diseases) exposed through a
single stable API route `/ner`. Loaded once at module import; inference runs
under @spaces.GPU so ZeroGPU allocates a GPU on demand.
No secrets live here — the OpenMed models are public. The Space itself is
private, so callers authenticate with a Hugging Face token at the HTTP layer.
"""
import json
import time
import traceback
import gradio as gr
import spaces
from transformers import pipeline
DRUG_MODEL = "OpenMed/OpenMed-NER-PharmaDetect-ModernClinical-149M"
DISEASE_MODEL = "OpenMed/OpenMed-NER-DiseaseDetect-ModernClinical-149M"
MAX_TEXTS = 50
MAX_CHARS = 2000
# Loaded once at module load (CPU). The GPU is only claimed inside ner().
_drug_pipe = pipeline(
"token-classification", model=DRUG_MODEL, aggregation_strategy="simple"
)
_disease_pipe = pipeline(
"token-classification", model=DISEASE_MODEL, aggregation_strategy="simple"
)
def _merge_spans(entities, text):
"""
OpenMed ModernBERT checkpoints emit sub-word fragments even with
aggregation_strategy='simple' (e.g. 'Tremfya' -> T/rem/f/ya). Merge
consecutive same-type entities whose character spans touch or are
separated only by whitespace/hyphen into one clean span, re-deriving the
surface from the original text and keeping the max fragment score.
"""
if not entities:
return []
# Merge PER TYPE. The drug and disease pipelines run independently and can
# both fire on the same token; sorting the combined list by offset then
# interleaves their fragments (T/drug, T/disease, rem/drug, rem/disease...)
# so a single pass would never find two same-type neighbours and nothing
# would ever merge. Grouping by type first is what makes merging reliable.
merged = []
by_type = {}
for e in entities:
by_type.setdefault(e["type"], []).append(e)
for etype, group in by_type.items():
ents = sorted(group, key=lambda x: (x["start"], x["end"]))
cur = dict(ents[0])
for e in ents[1:]:
gap = text[cur["end"]:e["start"]] if e["start"] >= cur["end"] else ""
contiguous = e["start"] <= cur["end"] + 1 and gap.strip(" -") == ""
if contiguous:
cur["end"] = max(cur["end"], e["end"])
cur["score"] = max(cur["score"], e["score"])
else:
merged.append(cur)
cur = dict(e)
merged.append(cur)
for m in merged:
m["surface"] = text[m["start"]:m["end"]].strip()
m["score"] = round(float(m["score"]), 4)
merged.sort(key=lambda x: (x["start"], x["type"]))
return [m for m in merged if m["surface"]]
def _dedupe(entities, text):
"""Merge sub-word spans, then keep the best-scoring (surface, type) pair."""
merged = _merge_spans(entities, text)
best = {}
for e in merged:
key = (e["surface"].lower(), e["type"])
if key not in best or e["score"] > best[key]["score"]:
best[key] = e
return sorted(best.values(), key=lambda x: (x["start"], -x["score"]))
def _run_pipe(pipe, text, etype):
out = []
for r in pipe(text):
out.append(
{
"surface": text[r["start"]:r["end"]] if r.get("end") is not None else r.get("word", ""),
"type": etype,
"score": round(float(r.get("score", 0.0)), 4),
"start": int(r.get("start", -1)),
"end": int(r.get("end", -1)),
}
)
return out
@spaces.GPU
def ner(texts_json: str, targets: str = "both") -> str:
"""
texts_json : JSON string, list of texts (<=50, each truncated to 2000 chars)
targets : "drug" | "disease" | "both" (default "both")
returns : JSON string
{results:[{text_index, entities:[{surface,type,score,start,end}]}],
models:{drug, disease}, latency_ms}
"""
t0 = time.time()
targets = (targets or "both").strip().lower()
if targets not in ("drug", "disease", "both"):
targets = "both"
try:
texts = json.loads(texts_json) if isinstance(texts_json, str) else texts_json
if isinstance(texts, str):
texts = [texts]
if not isinstance(texts, list):
raise ValueError("texts_json must be a JSON array of strings")
except Exception as e:
return json.dumps({"error": f"invalid texts_json: {e}", "results": []})
texts = [str(t)[:MAX_CHARS] for t in texts[:MAX_TEXTS]]
results = []
for i, text in enumerate(texts):
try:
ents = []
if targets in ("drug", "both"):
ents += _run_pipe(_drug_pipe, text, "drug")
if targets in ("disease", "both"):
ents += _run_pipe(_disease_pipe, text, "disease")
results.append({"text_index": i, "entities": _dedupe(ents, text)})
except Exception as e:
# One bad text must never fail the whole batch.
results.append(
{"text_index": i, "entities": [], "error": str(e)[:200]}
)
return json.dumps(
{
"results": results,
"models": {"drug": DRUG_MODEL, "disease": DISEASE_MODEL},
"latency_ms": int((time.time() - t0) * 1000),
}
)
demo = gr.Interface(
fn=ner,
inputs=[
gr.Textbox(label="texts_json", value='["Tremfya (guselkumab) for plaque psoriasis."]'),
gr.Textbox(label="targets (drug|disease|both)", value="both"),
],
outputs=gr.Textbox(label="result_json"),
title="aikka OpenMed NER",
description="Clinical drug + disease NER (OpenMed, Apache-2.0). API route: /ner",
api_name="ner",
)
if __name__ == "__main__":
demo.launch()