intent_classifier.py 16 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242243244245246247248249250251252253254255256257258259260261262263264265266267268269270271272273274275276277278279280281282283284285286287288289290291292293294295296297298299300301302303304305306307308309310311312313314315316317318319320321322323324325326327328329330331332333334335336337338339340341342343344345346347348349350351352353354355356357358359360361362363364365366367368369370371372373374375376377378379380381382383384385386387388389390391392393394395396397398399400401402403404405406407408409410411412413414415416417418419420421422423424425426427428429430431432433434435436437438439440441442443444445446447448449450451452
  1. # -*- coding: utf-8 -*-
  2. """
  3. Intent-Classifier Runtime.
  4. Laedt das trainierte Intent-Modell (ONNX oder PyTorch)
  5. und klassifiziert Eingabesaetze in Intents + Slots.
  6. Inferenz auf Pi 4: <100ms
  7. Modellgroesse: ~30-50MB
  8. Beispiel:
  9. ```python
  10. classifier = IntentClassifier()
  11. await classifier.load("models/intent")
  12. result = await classifier.classify("Mach das Licht im Wohnzimmer an")
  13. # IntentPrediction(
  14. # intent="lightswitch",
  15. # confidence=0.97,
  16. # slots={"room": "Wohnzimmer", "status": "an"},
  17. # )
  18. ```
  19. """
  20. from __future__ import annotations
  21. import json
  22. import time
  23. from dataclasses import dataclass, field
  24. from pathlib import Path
  25. from typing import Any
  26. from trixy_core.utils.debug import pinfo, pdebug, perror
  27. @dataclass
  28. class IntentPrediction:
  29. """Ergebnis einer Intent-Klassifikation."""
  30. intent: str
  31. confidence: float
  32. slots: dict[str, str | list[str]] = field(default_factory=dict)
  33. tones: list[str] = field(default_factory=list) # Emotionaler Ton des Satzes
  34. inference_ms: float = 0.0
  35. # Top-N Alternativen
  36. alternatives: list[tuple[str, float]] = field(default_factory=list)
  37. @property
  38. def is_confident(self) -> bool:
  39. """True wenn Confidence ueber Schwelle (0.5)."""
  40. return self.confidence >= 0.5
  41. @property
  42. def is_friendly(self) -> bool:
  43. """True wenn Ton freundlich ist."""
  44. return any(t in self.tones for t in ("friendly", "polite", "nice", "grateful"))
  45. @property
  46. def is_rude(self) -> bool:
  47. """True wenn Ton unhoeflich/aggressiv ist."""
  48. return any(t in self.tones for t in ("rude", "aggressive"))
  49. @property
  50. def is_urgent(self) -> bool:
  51. """True wenn dringend."""
  52. return "urgent" in self.tones
  53. def to_dict(self) -> dict[str, Any]:
  54. from trixy_core.nlp.tone import TAG_TO_CATEGORY
  55. return {
  56. "intent": self.intent,
  57. "confidence": self.confidence,
  58. "slots": self.slots,
  59. "tones": [
  60. {"tag": t, "category": TAG_TO_CATEGORY.get(t, "unknown")}
  61. for t in self.tones
  62. ],
  63. "inference_ms": self.inference_ms,
  64. }
  65. class IntentClassifier:
  66. """
  67. Intent-Classifier fuer Trixy.
  68. Nutzt Sentence-Embeddings + trainiertes Klassifikationsmodell
  69. fuer schnelle Intent-Erkennung auf CPU.
  70. """
  71. def __init__(self) -> None:
  72. self._encoder = None # SentenceTransformer oder ONNX-Encoder
  73. self._tokenizer = None # HuggingFace Tokenizer (fuer ONNX-Encoder)
  74. self._classifier = None # ONNX InferenceSession oder PyTorch
  75. self._slot_tagger = None # ONNX InferenceSession fuer BIO-Tagger
  76. self._bio_tagset = None # BIOTagSet fuer Tag-Decoding
  77. self._use_onnx = False
  78. self._use_onnx_encoder = False
  79. self._has_bio_tagger = False
  80. self._intent_names: list[str] = []
  81. self._slot_names: list[str] = []
  82. self._embedding_dim: int = 0
  83. self._loaded = False
  84. self._model_dir: Path | None = None
  85. @property
  86. def is_loaded(self) -> bool:
  87. """True wenn Modell geladen."""
  88. return self._loaded
  89. @property
  90. def intent_count(self) -> int:
  91. """Anzahl bekannter Intents."""
  92. return len(self._intent_names)
  93. @property
  94. def intent_names(self) -> list[str]:
  95. """Liste aller bekannten Intent-Namen."""
  96. return list(self._intent_names)
  97. async def load(self, model_dir: str | Path) -> bool:
  98. """
  99. Laedt das trainierte Modell.
  100. Args:
  101. model_dir: Verzeichnis mit metafile.json, classifier.onnx/pth
  102. Returns:
  103. True bei Erfolg
  104. """
  105. model_dir = Path(model_dir)
  106. self._model_dir = model_dir
  107. # Metadaten laden
  108. metafile = model_dir / "metafile.json"
  109. if not metafile.is_file():
  110. pdebug(f"IntentClassifier: Kein Modell in {model_dir}")
  111. return False
  112. try:
  113. with open(metafile) as f:
  114. meta = json.load(f)
  115. self._embedding_dim = meta["embedding_dim"]
  116. base_model = meta["base_model"]
  117. # Labels laden
  118. with open(model_dir / "intent_labels.json") as f:
  119. self._intent_names = json.load(f)
  120. self._slot_names = meta.get("slot_names", [])
  121. except (json.JSONDecodeError, KeyError, OSError) as e:
  122. perror(f"IntentClassifier: Metadaten-Fehler: {e}")
  123. return False
  124. # Encoder laden (ONNX bevorzugt, Fallback SentenceTransformer)
  125. encoder_dir = model_dir / "encoder_onnx"
  126. has_onnx_encoder = encoder_dir.is_dir() and any(
  127. f.suffix == ".onnx" for f in encoder_dir.iterdir() if f.is_file()
  128. )
  129. if has_onnx_encoder:
  130. # ONNX-Encoder (kein PyTorch noetig, ~112MB)
  131. try:
  132. from optimum.onnxruntime import ORTModelForFeatureExtraction
  133. from transformers import AutoTokenizer
  134. # Quantisiertes Modell bevorzugen
  135. onnx_files = [f.name for f in encoder_dir.iterdir() if f.suffix == ".onnx"]
  136. file_name = "model_quantized.onnx" if "model_quantized.onnx" in onnx_files else None
  137. self._encoder = ORTModelForFeatureExtraction.from_pretrained(
  138. str(encoder_dir), file_name=file_name,
  139. )
  140. self._tokenizer = AutoTokenizer.from_pretrained(str(encoder_dir))
  141. self._use_onnx_encoder = True
  142. pinfo(f"IntentClassifier: ONNX-Encoder geladen ({encoder_dir})")
  143. except ImportError:
  144. pdebug("optimum nicht verfuegbar, versuche SentenceTransformer")
  145. except Exception as e:
  146. pdebug(f"ONNX-Encoder Fehler: {e}, versuche SentenceTransformer")
  147. if not self._use_onnx_encoder:
  148. # Fallback: SentenceTransformer (braucht PyTorch, ~400MB)
  149. try:
  150. from sentence_transformers import SentenceTransformer
  151. self._encoder = SentenceTransformer(base_model)
  152. pdebug(f"IntentClassifier: SentenceTransformer geladen ({base_model})")
  153. except ImportError:
  154. perror("Weder optimum noch sentence-transformers installiert")
  155. return False
  156. # Classifier laden (ONNX bevorzugt, Fallback PyTorch)
  157. onnx_path = model_dir / "intent_classifier.onnx"
  158. pth_path = model_dir / "classifier.pth"
  159. if onnx_path.is_file():
  160. try:
  161. import onnxruntime as ort
  162. self._classifier = ort.InferenceSession(
  163. str(onnx_path),
  164. providers=["CPUExecutionProvider"],
  165. )
  166. self._use_onnx = True
  167. pdebug(f"IntentClassifier: ONNX Modell geladen")
  168. except ImportError:
  169. pdebug("onnxruntime nicht verfuegbar, versuche PyTorch")
  170. if not self._use_onnx and pth_path.is_file():
  171. try:
  172. import torch
  173. num_intents = len(self._intent_names)
  174. self._classifier = torch.nn.Sequential(
  175. torch.nn.Linear(self._embedding_dim, 256),
  176. torch.nn.ReLU(),
  177. torch.nn.Dropout(0.3),
  178. torch.nn.Linear(256, num_intents),
  179. )
  180. self._classifier.load_state_dict(
  181. torch.load(pth_path, map_location="cpu", weights_only=True)
  182. )
  183. self._classifier.eval()
  184. pdebug(f"IntentClassifier: PyTorch Modell geladen")
  185. except Exception as e:
  186. perror(f"IntentClassifier: PyTorch laden fehlgeschlagen: {e}")
  187. return False
  188. if self._classifier is None:
  189. perror(f"IntentClassifier: Kein Modell gefunden in {model_dir}")
  190. return False
  191. # BIO Slot-Tagger laden (optional)
  192. bio_onnx = model_dir / "slot_tagger.onnx"
  193. bio_labels = model_dir / "bio_labels.json"
  194. if bio_onnx.is_file() and bio_labels.is_file():
  195. try:
  196. import onnxruntime as ort
  197. self._slot_tagger = ort.InferenceSession(
  198. str(bio_onnx), providers=["CPUExecutionProvider"],
  199. )
  200. with open(bio_labels) as f:
  201. from trixy_core.trainer.core.intent.bio_tagger import BIOTagSet
  202. self._bio_tagset = BIOTagSet.from_dict(json.load(f))
  203. self._has_bio_tagger = True
  204. pinfo(f"IntentClassifier: BIO Slot-Tagger geladen ({self._bio_tagset.num_tags} Tags)")
  205. except Exception as e:
  206. pdebug(f"BIO Slot-Tagger nicht geladen: {e}")
  207. self._loaded = True
  208. pinfo(
  209. f"IntentClassifier geladen: {len(self._intent_names)} Intents"
  210. f"{', ONNX-Encoder' if self._use_onnx_encoder else ', SentenceTransformer'}"
  211. f"{', BIO-Slots' if self._has_bio_tagger else ''}"
  212. )
  213. return True
  214. async def classify(
  215. self,
  216. text: str,
  217. min_confidence: float = 0.3,
  218. top_n: int = 3,
  219. ) -> IntentPrediction:
  220. """
  221. Klassifiziert einen Text in Intent + Slots.
  222. Args:
  223. text: Eingabetext
  224. min_confidence: Minimale Confidence fuer gueltigen Intent
  225. top_n: Anzahl Alternativen zurueckgeben
  226. Returns:
  227. IntentPrediction
  228. """
  229. if not self._loaded:
  230. return IntentPrediction(intent="unknown", confidence=0.0)
  231. start = time.monotonic()
  232. # 1. Embedding berechnen
  233. import numpy as np
  234. if self._use_onnx_encoder:
  235. # ONNX-Encoder: Tokenize → Encoder → Mean-Pooling
  236. inputs = self._tokenizer(
  237. text, return_tensors="np",
  238. padding=True, truncation=True, max_length=128,
  239. )
  240. outputs = self._encoder(**inputs)
  241. # Mean-Pooling ueber Token-Embeddings (ohne Padding)
  242. token_embeddings = outputs.last_hidden_state[0] # (seq_len, dim)
  243. attention_mask = inputs["attention_mask"][0] # (seq_len,)
  244. mask = attention_mask.astype(np.float32)
  245. masked = token_embeddings * mask[:, np.newaxis]
  246. embedding = masked.sum(axis=0) / mask.sum()
  247. embedding = embedding.reshape(1, -1).astype(np.float32)
  248. else:
  249. # SentenceTransformer: Direkte Embedding-Berechnung
  250. embedding = self._encoder.encode([text], show_progress_bar=False)
  251. embedding = np.array(embedding, dtype=np.float32)
  252. # 2. Classifier ausfuehren
  253. if self._use_onnx:
  254. outputs = self._classifier.run(None, {"embedding": embedding})
  255. logits = outputs[0][0]
  256. else:
  257. import torch
  258. with torch.no_grad():
  259. logits = self._classifier(torch.FloatTensor(embedding))
  260. logits = logits[0].numpy()
  261. # 3. Softmax + Top-N
  262. exp_logits = np.exp(logits - np.max(logits))
  263. probabilities = exp_logits / exp_logits.sum()
  264. top_indices = np.argsort(probabilities)[::-1][:top_n]
  265. best_idx = top_indices[0]
  266. best_confidence = float(probabilities[best_idx])
  267. best_intent = self._intent_names[best_idx]
  268. # Negativ-Klasse pruefen
  269. if best_intent == "__negative__" or best_confidence < min_confidence:
  270. best_intent = "unknown"
  271. best_confidence = 0.0
  272. # Alternativen
  273. alternatives = [
  274. (self._intent_names[idx], float(probabilities[idx]))
  275. for idx in top_indices[1:]
  276. if self._intent_names[idx] != "__negative__"
  277. ]
  278. # 4. Slot-Extraktion
  279. slots: dict[str, Any] = {}
  280. if best_intent != "unknown":
  281. # BIO-Tagger hat Vorrang (neuronale Slot-Extraktion)
  282. if self._has_bio_tagger and self._use_onnx_encoder:
  283. slots = self._extract_slots_bio(token_embeddings, inputs)
  284. # Fallback: Regelbasiert
  285. if not slots:
  286. slots = self._extract_slots(text)
  287. # 5. Tone-Analyse (Schlagwort + Satzstruktur)
  288. from trixy_core.nlp.tone import analyze_tone
  289. tone_results = analyze_tone(text)
  290. tones = [r.tag for r in tone_results]
  291. elapsed = (time.monotonic() - start) * 1000
  292. return IntentPrediction(
  293. intent=best_intent,
  294. confidence=best_confidence,
  295. slots=slots,
  296. tones=tones,
  297. inference_ms=elapsed,
  298. alternatives=alternatives,
  299. )
  300. def _extract_slots_bio(
  301. self, token_embeddings: "np.ndarray", inputs: dict,
  302. ) -> dict[str, Any]:
  303. """
  304. Extrahiert Slots via BIO-Tagger auf Token-Embeddings.
  305. Args:
  306. token_embeddings: (seq_len, 384) vom Encoder
  307. inputs: Tokenizer-Output mit input_ids
  308. Returns:
  309. Dict mit extrahierten Slots
  310. """
  311. import numpy as np
  312. if not self._slot_tagger or not self._bio_tagset:
  313. return {}
  314. try:
  315. # BIO-Logits berechnen: (1, seq_len, num_tags)
  316. token_emb_batch = token_embeddings.reshape(1, *token_embeddings.shape).astype(np.float32)
  317. bio_outputs = self._slot_tagger.run(
  318. None, {"token_embeddings": token_emb_batch},
  319. )
  320. bio_logits = bio_outputs[0][0] # (seq_len, num_tags)
  321. # Argmax → Tag-Indizes
  322. tag_indices = bio_logits.argmax(axis=1).tolist() # (seq_len,)
  323. # Token-IDs fuer Decoding
  324. token_ids = inputs["input_ids"][0].tolist() if hasattr(inputs["input_ids"], "tolist") else list(inputs["input_ids"][0])
  325. # BIO-Tags decodieren → Slots
  326. from trixy_core.trainer.core.intent.bio_tagger import decode_bio_tags
  327. slots = decode_bio_tags(
  328. tag_indices, token_ids, self._bio_tagset, self._tokenizer,
  329. )
  330. if slots:
  331. pdebug(f"IntentClassifier BIO-Slots: {slots}")
  332. return slots
  333. except Exception as e:
  334. pdebug(f"BIO Slot-Extraktion Fehler: {e}")
  335. return {}
  336. def _extract_slots(self, text: str) -> dict[str, str | list[str]]:
  337. """
  338. Extrahiert Slots aus dem Text via Pattern-Matching.
  339. Nutzt die Slot-Listen um bekannte Werte zu finden.
  340. """
  341. slots: dict[str, str | list[str]] = {}
  342. text_lower = text.lower()
  343. # Slot-Listen laden (gecacht aus Metafile)
  344. if not self._model_dir:
  345. return slots
  346. slot_file = self._model_dir / ".." / ".." / "config" / "slot_lists.json"
  347. if not slot_file.is_file():
  348. # Fallback: Default Slot-Registry
  349. from trixy_core.trainer.core.intent.slot_lists import SlotRegistry
  350. registry = SlotRegistry()
  351. for slot_name in self._slot_names:
  352. values = registry.get_values(slot_name)
  353. found = self._find_values_in_text(text_lower, values)
  354. if found:
  355. slots[slot_name] = found if len(found) > 1 else found[0]
  356. return slots
  357. try:
  358. with open(slot_file) as f:
  359. slot_data = json.load(f)
  360. for slot_name in self._slot_names:
  361. entry = slot_data.get(slot_name, {})
  362. values = entry.get("values", []) if isinstance(entry, dict) else entry
  363. found = self._find_values_in_text(text_lower, values)
  364. if found:
  365. slots[slot_name] = found if len(found) > 1 else found[0]
  366. except (json.JSONDecodeError, OSError):
  367. pass
  368. return slots
  369. @staticmethod
  370. def _find_values_in_text(text: str, values: list[str]) -> list[str]:
  371. """Findet Slot-Werte im Text (case-insensitive)."""
  372. found = []
  373. # Laengste Werte zuerst (damit "Wohnzimmer" vor "Wohn" matched)
  374. for value in sorted(values, key=len, reverse=True):
  375. if value.lower() in text:
  376. found.append(value)
  377. # Wert aus Text entfernen um Doppel-Matches zu vermeiden
  378. text = text.replace(value.lower(), "", 1)
  379. return found