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