| 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 |
|
|
| |
| 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__) |
|
|
| |
| 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 |
| |
| |
| 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 |
|
|