kenlm_layer.py 4.0 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129
  1. # -*- coding: utf-8 -*-
  2. """
  3. KenLM-Korrekturschicht.
  4. Nutzt N-Gram Sprachmodell fuer kontextuelle Kandidatenauswahl.
  5. Sehr schnell (<1ms/Query).
  6. """
  7. from pathlib import Path
  8. from trixy_core.stt.layers.base import CorrectionLayer
  9. from trixy_core.utils.debug import pdebug, pwarn
  10. try:
  11. import kenlm
  12. _HAS_KENLM = True
  13. except ImportError:
  14. _HAS_KENLM = False
  15. class KenLMLayer(CorrectionLayer):
  16. """KenLM N-Gram basierte kontextuelle Korrektur."""
  17. NAME = "kenlm"
  18. def __init__(self) -> None:
  19. self._model: "kenlm.Model | None" = None
  20. self._protected_words: set[str] = set()
  21. def is_available(self) -> bool:
  22. """Prueft ob kenlm installiert ist."""
  23. return _HAS_KENLM
  24. def initialize(self, config: dict, language: str, protected_words: list[str]) -> bool:
  25. """Laedt N-Gram Modell aus konfiguriertem Pfad."""
  26. if not _HAS_KENLM:
  27. return False
  28. model_path = config.get("kenlm_model_path", "")
  29. if not model_path or not Path(model_path).exists():
  30. pdebug(f"[KenLM] Modell nicht gefunden: {model_path}")
  31. return False
  32. try:
  33. self._model = kenlm.Model(str(model_path))
  34. self._protected_words = {w.lower() for w in protected_words}
  35. pdebug(f"[KenLM] Modell geladen: {model_path} (order={self._model.order})")
  36. return True
  37. except Exception as e:
  38. pwarn(f"[KenLM] Modell-Laden fehlgeschlagen: {e}")
  39. self._model = None
  40. return False
  41. def correct(self, text: str) -> str:
  42. """Bewertet Text und waehlt beste Kandidaten per N-Gram Score."""
  43. if not self._model or not text.strip():
  44. return text
  45. try:
  46. words = text.lower().split()
  47. if len(words) < 2:
  48. return text
  49. # Score des Original-Satzes
  50. original_score = self._model.score(text.lower(), bos=True, eos=True)
  51. # Kandidaten generieren: pro Wort Edit-Distanz-1 Varianten
  52. best_text = text.lower()
  53. best_score = original_score
  54. for i, word in enumerate(words):
  55. if word in self._protected_words:
  56. continue
  57. candidates = self._generate_candidates(word)
  58. for candidate in candidates:
  59. test_words = words.copy()
  60. test_words[i] = candidate
  61. test_text = " ".join(test_words)
  62. score = self._model.score(test_text, bos=True, eos=True)
  63. if score > best_score:
  64. best_score = score
  65. best_text = test_text
  66. if best_text != text.lower():
  67. pdebug(f"[KenLM] '{text}' → '{best_text}' (score: {original_score:.2f} → {best_score:.2f})")
  68. return best_text
  69. return text
  70. except Exception as e:
  71. pdebug(f"[KenLM] Fehler bei Korrektur: {e}")
  72. return text
  73. def _generate_candidates(self, word: str) -> list[str]:
  74. """Generiert Edit-Distanz-1 Kandidaten fuer ein Wort."""
  75. candidates = set()
  76. letters = "abcdefghijklmnopqrstuvwxyzäöüß"
  77. # Loeschungen
  78. for i in range(len(word)):
  79. candidates.add(word[:i] + word[i + 1:])
  80. # Transpositionen
  81. for i in range(len(word) - 1):
  82. candidates.add(word[:i] + word[i + 1] + word[i] + word[i + 2:])
  83. # Ersetzungen
  84. for i in range(len(word)):
  85. for c in letters:
  86. if c != word[i]:
  87. candidates.add(word[:i] + c + word[i + 1:])
  88. # Einfuegungen
  89. for i in range(len(word) + 1):
  90. for c in letters:
  91. candidates.add(word[:i] + c + word[i:])
  92. # Leere und zu kurze entfernen
  93. candidates.discard("")
  94. candidates.discard(word)
  95. return list(candidates)
  96. def shutdown(self) -> None:
  97. """Gibt Ressourcen frei."""
  98. self._model = None