HFSwissParlTopicTagger / predictor.py
kapllan's picture
First commit.
80f21d1
Raw
History Blame Contribute Delete
5.8 kB
import json as js
import logging
import os
import re
from pathlib import Path
import sys
import types
from typing import Dict, List, Tuple
import joblib
import omikuji
from huggingface_hub import snapshot_download
import requests
from langdetect import detect
from langdetect.lang_detect_exception import LangDetectException
# Constants
HF_REPOS = [
'kapllan/omikuji-bonsai-parliament-de-spacy',
'kapllan/omikuji-bonsai-parliament-fr-spacy',
'kapllan/omikuji-bonsai-parliament-it-spacy'
]
BASE_PATH = Path(__file__).resolve().parent
MODEL_CACHE_ROOT = Path(
os.getenv(
"SWISSPARL_MODEL_DIR",
Path.home() / ".cache" / "swissparl-topic-tagger" / "models",
)
)
logger = logging.getLogger(__name__)
# State
id2label = None
topics_hierarchy = None
MODEL_CACHE: Dict[str, Tuple[object, object]] = {}
def download_file(url, save_path):
if not os.path.exists(save_path):
response = requests.get(url, stream=True)
if response.status_code == 200:
with open(save_path, 'wb') as f:
for chunk in response.iter_content(chunk_size=1024):
f.write(chunk)
def get_repo_dir(repo_id: str) -> Path:
repo_dir = MODEL_CACHE_ROOT / repo_id
repo_dir.parent.mkdir(parents=True, exist_ok=True)
return repo_dir
def get_language_repo_dir(language: str) -> Path:
repo_name = f'omikuji-bonsai-parliament-{language}-spacy'
legacy_dir = Path.cwd() / 'kapllan' / repo_name
if legacy_dir.exists():
return legacy_dir
bundled_dir = BASE_PATH / 'kapllan' / repo_name
if bundled_dir.exists():
return bundled_dir
return MODEL_CACHE_ROOT / 'kapllan' / repo_name
def register_annif_spacy_compat() -> None:
try:
__import__("annif.analyzer.spacy")
return
except ModuleNotFoundError:
pass
annif_module = sys.modules.setdefault("annif", types.ModuleType("annif"))
analyzer_module = sys.modules.setdefault(
"annif.analyzer", types.ModuleType("annif.analyzer")
)
spacy_module = types.ModuleType("annif.analyzer.spacy")
class SpacyAnalyzer:
def __setstate__(self, state):
self.__dict__.update(state)
def tokenize_words(self, text):
doc = self.nlp(text)
tokens = [token.text for token in doc if not token.is_space]
if getattr(self, "lowercase", False):
return [token.lower() for token in tokens]
return tokens
spacy_module.SpacyAnalyzer = SpacyAnalyzer
annif_module.analyzer = analyzer_module
analyzer_module.spacy = spacy_module
sys.modules["annif.analyzer.spacy"] = spacy_module
def find_resource_path(filename: str) -> Path | None:
candidates = [
Path(os.getenv("SWISSPARL_RESOURCE_DIR", "")) / filename,
BASE_PATH / filename,
BASE_PATH.parent / filename,
Path.cwd() / filename,
]
seen: set[Path] = set()
for candidate in candidates:
if not str(candidate):
continue
resolved = candidate.resolve(strict=False)
if resolved in seen:
continue
seen.add(resolved)
if resolved.exists():
return resolved
return None
def load_resources():
global id2label, topics_hierarchy
# Download Omikuji models
for repo_id in HF_REPOS:
repo_dir = get_repo_dir(repo_id)
if not repo_dir.exists():
snapshot_download(repo_id=repo_id, local_dir=str(repo_dir))
id2label_path = find_resource_path("id2label.json")
hierarchy_path = find_resource_path("topics_hierarchy.json")
if id2label_path is None or hierarchy_path is None:
search_roots = [BASE_PATH, BASE_PATH.parent, Path.cwd()]
resource_dir = os.getenv("SWISSPARL_RESOURCE_DIR")
if resource_dir:
search_roots.insert(0, Path(resource_dir))
logger.warning(
"id2label.json or topics_hierarchy.json not found. Looked in %s",
", ".join(str(path.resolve(strict=False)) for path in search_roots),
)
return
with id2label_path.open("r") as f:
id2label = js.load(f)
with hierarchy_path.open("r") as f:
topics_hierarchy = js.load(f)
def predict_lang(text: str) -> str:
try:
return detect(text)
except LangDetectException:
return "unknown"
def find_model(language: str):
if language in ['de', 'fr', 'it']:
cached = MODEL_CACHE.get(language)
if cached is not None:
return cached
register_annif_spacy_compat()
repo_dir = get_language_repo_dir(language)
path_to_vectorizer = repo_dir / 'vectorizer'
path_to_model = repo_dir / 'omikuji-model'
if path_to_vectorizer.exists() and path_to_model.exists():
vectorizer = joblib.load(path_to_vectorizer)
model = omikuji.Model.load(str(path_to_model))
MODEL_CACHE[language] = (vectorizer, model)
return vectorizer, model
return None, None
def predict_topic(text: str, top_k: int = 1000) -> Tuple[List[Tuple[str, float]], str]:
if id2label is None:
load_resources()
results = []
language = predict_lang(text)
vectorizer, model = find_model(language)
if vectorizer is not None and model is not None:
vector = vectorizer.transform([text])
for row in vector:
if row.nnz == 0: continue
feature_values = [(col, row[0, col]) for col in row.nonzero()[1]]
for subj_id, score in model.predict(feature_values, top_k=top_k):
label = id2label.get(str(subj_id), str(subj_id)) if id2label else str(subj_id)
results.append((label, round(score, 2)))
return results, language