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