temsa's picture
Publish rc7 with spec-driven scanner release
32bcb86 verified
Raw
History Blame Contribute Delete
10 kB
#!/usr/bin/env python3
import json
import os
import tempfile
from pathlib import Path
from typing import Any
os.environ.setdefault("TRANSFORMERS_NO_TF", "1")
os.environ.setdefault("TRANSFORMERS_NO_FLAX", "1")
os.environ.setdefault("TRANSFORMERS_NO_TORCHVISION", "1")
os.environ["USE_TF"] = "0"
os.environ["USE_FLAX"] = "0"
os.environ["USE_TORCH"] = "1"
import numpy as np
import re
from huggingface_hub import HfApi, hf_hub_download
from transformers import AutoConfig, AutoTokenizer
TOKEN_RE = re.compile(r"[A-Za-z0-9]+|[^\w\s]", re.UNICODE)
DEFAULT_ONNX_FILES = [
"onnx/model_quantized.onnx",
"model_quantized.onnx",
"onnx/model.onnx",
"model.onnx",
]
EIRCODE_RE = re.compile(r"^(?:[ACDEFHKNPRTVWXY]\d{2}|D6W)\s?[0-9ACDEFHKNPRTVWXY]{4}$", re.IGNORECASE)
TOKENIZER_FILES = [
"tokenizer_config.json",
"tokenizer.json",
"special_tokens_map.json",
"vocab.txt",
"vocab.json",
"merges.txt",
"added_tokens.json",
"sentencepiece.bpe.model",
"spiece.model",
]
def tokenize_with_spans(text: str):
return [(m.group(0), m.start(), m.end()) for m in TOKEN_RE.finditer(text)]
def normalize_label(label: str) -> str:
label = (label or "").strip()
if label.startswith("B-") or label.startswith("I-"):
label = label[2:]
return label.upper()
def looks_like_eircode(value: str) -> bool:
return EIRCODE_RE.match(value.strip()) is not None
def _sanitize_tokenizer_dir(tokenizer_path: Path) -> str:
tokenizer_cfg_path = tokenizer_path / "tokenizer_config.json"
if not tokenizer_cfg_path.exists():
return str(tokenizer_path)
data = json.loads(tokenizer_cfg_path.read_text(encoding="utf-8"))
if "fix_mistral_regex" not in data:
return str(tokenizer_path)
tmpdir = Path(tempfile.mkdtemp(prefix="openmed_onnx_tokenizer_"))
keep = set(TOKENIZER_FILES)
for child in tokenizer_path.iterdir():
if child.is_file() and child.name in keep:
target = tmpdir / child.name
target.write_bytes(child.read_bytes())
data.pop("fix_mistral_regex", None)
(tmpdir / "tokenizer_config.json").write_text(
json.dumps(data, ensure_ascii=False, indent=2) + "\n",
encoding="utf-8",
)
return str(tmpdir)
def _materialize_remote_tokenizer(repo_id: str) -> str:
api = HfApi()
files = set(api.list_repo_files(repo_id=repo_id, repo_type="model"))
tmpdir = Path(tempfile.mkdtemp(prefix="openmed_remote_tokenizer_"))
copied = False
for name in TOKENIZER_FILES:
if name not in files:
continue
src = hf_hub_download(repo_id=repo_id, filename=name, repo_type="model")
(tmpdir / Path(name).name).write_bytes(Path(src).read_bytes())
copied = True
if not copied:
return repo_id
return _sanitize_tokenizer_dir(tmpdir)
def safe_auto_tokenizer(tokenizer_ref: str):
tokenizer_path = Path(tokenizer_ref)
if tokenizer_path.exists():
tokenizer_ref = _sanitize_tokenizer_dir(tokenizer_path)
else:
tokenizer_ref = _materialize_remote_tokenizer(tokenizer_ref)
try:
return AutoTokenizer.from_pretrained(tokenizer_ref, use_fast=True, fix_mistral_regex=True)
except Exception:
pass
try:
return AutoTokenizer.from_pretrained(tokenizer_ref, use_fast=True, fix_mistral_regex=False)
except TypeError:
pass
try:
return AutoTokenizer.from_pretrained(tokenizer_ref, use_fast=True)
except Exception:
return AutoTokenizer.from_pretrained(tokenizer_ref, use_fast=False)
def _load_tokenizer(tokenizer_ref: str):
return safe_auto_tokenizer(tokenizer_ref)
def _resolve_local_onnx(model_path: Path, preferred: str | None = None) -> Path:
candidates = ([preferred] if preferred else []) + DEFAULT_ONNX_FILES
for candidate in candidates:
if not candidate:
continue
path = model_path / candidate
if path.exists():
return path
raise FileNotFoundError(f"No ONNX file found under {model_path}")
def _resolve_remote_onnx(model_ref: str, preferred: str | None = None) -> Path:
api = HfApi()
files = set(api.list_repo_files(repo_id=model_ref, repo_type="model"))
candidates = ([preferred] if preferred else []) + DEFAULT_ONNX_FILES
for candidate in candidates:
if candidate and candidate in files:
return Path(hf_hub_download(repo_id=model_ref, filename=candidate, repo_type="model"))
raise FileNotFoundError(f"No ONNX file published for {model_ref}")
def load_onnx_token_classifier(
model_ref: str,
onnx_file: str | None = None,
providers: list[str] | None = None,
):
import onnxruntime as ort
model_path = Path(model_ref)
if model_path.exists():
onnx_path = _resolve_local_onnx(model_path, preferred=onnx_file)
config = AutoConfig.from_pretrained(model_ref)
tokenizer = safe_auto_tokenizer(model_ref)
else:
onnx_path = _resolve_remote_onnx(model_ref, preferred=onnx_file)
config = AutoConfig.from_pretrained(model_ref)
tokenizer = safe_auto_tokenizer(model_ref)
session = ort.InferenceSession(str(onnx_path), providers=providers or ["CPUExecutionProvider"])
return session, tokenizer, config, onnx_path
def _run_onnx(session, encoded: dict[str, Any]) -> np.ndarray:
feed = {}
input_names = {item.name for item in session.get_inputs()}
for key, value in encoded.items():
if key == "offset_mapping":
continue
if key in input_names:
feed[key] = value
outputs = session.run(None, feed)
return outputs[0]
def _softmax(logits: np.ndarray, axis: int = -1) -> np.ndarray:
shifted = logits - np.max(logits, axis=axis, keepdims=True)
exp = np.exp(shifted)
return exp / np.clip(np.sum(exp, axis=axis, keepdims=True), 1e-12, None)
def _split_tag(label: str) -> tuple[str, str]:
if label.startswith("B-") or label.startswith("I-"):
return label[:1], label[2:]
return "B", label
def simple_aggregate_spans_onnx(
text: str,
session,
tokenizer,
config,
min_score: float = 0.5,
) -> list[dict[str, Any]]:
encoded = tokenizer(text, return_offsets_mapping=True, return_tensors="np", truncation=True)
logits = _run_onnx(session, encoded)[0]
probs = _softmax(logits, axis=-1)
pred_ids = probs.argmax(axis=-1)
id2label = {int(k): v for k, v in config.id2label.items()}
offsets = encoded["offset_mapping"][0].tolist()
attention_mask = encoded.get("attention_mask")
if attention_mask is None:
attention = [1] * len(offsets)
else:
attention = attention_mask[0].tolist()
spans: list[dict[str, Any]] = []
active: dict[str, Any] | None = None
for idx, ((start, end), keep) in enumerate(zip(offsets, attention)):
if not keep or start == end:
if active is not None:
spans.append(active)
active = None
continue
label = id2label[int(pred_ids[idx])]
if label == "O":
if active is not None:
spans.append(active)
active = None
continue
score = float(probs[idx, int(pred_ids[idx])])
if score < min_score:
if active is not None:
spans.append(active)
active = None
continue
prefix, entity = _split_tag(label)
if (
active is None
or prefix == "B"
or entity != active["entity_group"]
or int(start) > int(active["end"]) + 1
):
if active is not None:
spans.append(active)
active = {
"entity_group": entity,
"start": int(start),
"end": int(end),
"score": score,
}
else:
active["end"] = int(end)
active["score"] = max(float(active["score"]), score)
if active is not None:
spans.append(active)
for span in spans:
span["word"] = text[span["start"] : span["end"]]
return spans
def ppsn_label_ids_from_config(config) -> list[int]:
ids = []
for raw_id, raw_label in config.id2label.items():
label_id = int(raw_id)
label = str(raw_label or "").strip()
if label.endswith("PPSN"):
ids.append(label_id)
return sorted(ids)
def word_aligned_ppsn_spans_onnx(
text: str,
session,
tokenizer,
config,
threshold: float = 0.4,
) -> list[dict[str, Any]]:
pieces = tokenize_with_spans(text)
if not pieces:
return []
words = [word for word, _, _ in pieces]
encoded = tokenizer(words, is_split_into_words=True, return_tensors="np", truncation=True)
word_ids = encoded.word_ids(batch_index=0)
logits = _run_onnx(session, encoded)[0]
probs = _softmax(logits, axis=-1)
label_ids = ppsn_label_ids_from_config(config)
word_scores: list[float] = []
for word_index in range(len(pieces)):
score = 0.0
for token_index, wid in enumerate(word_ids):
if wid != word_index:
continue
for label_id in label_ids:
score = max(score, float(probs[token_index, label_id]))
word_scores.append(score)
spans: list[dict[str, Any]] = []
active = None
for (_, start, end), score in zip(pieces, word_scores):
if score >= threshold:
if active is None:
active = {"start": start, "end": end, "score": score}
else:
active["end"] = end
active["score"] = max(active["score"], score)
elif active is not None:
spans.append(active)
active = None
if active is not None:
spans.append(active)
for span in spans:
span["text"] = text[span["start"] : span["end"]]
span["label"] = "PPSN"
span["source"] = "onnx"
return spans