File size: 5,798 Bytes
80f21d1 | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 142 143 144 145 146 147 148 149 150 151 152 153 154 155 156 157 158 159 160 161 162 163 164 165 166 167 168 169 170 171 172 173 174 175 176 177 178 179 180 181 182 | 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
|