From a0f300068ed28d49ce9684fa3100428c1df7ca5d Mon Sep 17 00:00:00 2001 From: hueso Date: Fri, 7 Aug 2026 04:23:44 -0300 Subject: [PATCH] STT plugin based on onnx-asr --- stt_voice_messages/__init__.py | 1 + stt_voice_messages/gtk/config_dialog.py | 100 +++++++++++ stt_voice_messages/gtk/sttbox.py | 155 +++++++++++++++++ stt_voice_messages/model.py | 164 ++++++++++++++++++ ...jim.Plugin.stt_voice_messages.metainfo.xml | 11 ++ stt_voice_messages/plugin-manifest.json | 20 +++ stt_voice_messages/stt_voice_messages.png | Bin 0 -> 371 bytes stt_voice_messages/stt_voice_messages.py | 70 ++++++++ 8 files changed, 521 insertions(+) create mode 100644 stt_voice_messages/__init__.py create mode 100644 stt_voice_messages/gtk/config_dialog.py create mode 100644 stt_voice_messages/gtk/sttbox.py create mode 100644 stt_voice_messages/model.py create mode 100644 stt_voice_messages/org.gajim.Gajim.Plugin.stt_voice_messages.metainfo.xml create mode 100644 stt_voice_messages/plugin-manifest.json create mode 100644 stt_voice_messages/stt_voice_messages.png create mode 100644 stt_voice_messages/stt_voice_messages.py diff --git a/stt_voice_messages/__init__.py b/stt_voice_messages/__init__.py new file mode 100644 index 0000000..81ba31c --- /dev/null +++ b/stt_voice_messages/__init__.py @@ -0,0 +1 @@ +from .stt_voice_messages import STTVoiceMessagesPlugin # pyright: ignore # noqa: F401 diff --git a/stt_voice_messages/gtk/config_dialog.py b/stt_voice_messages/gtk/config_dialog.py new file mode 100644 index 0000000..dad57e4 --- /dev/null +++ b/stt_voice_messages/gtk/config_dialog.py @@ -0,0 +1,100 @@ +# This file is part of Gajim. +# +# Gajim is free software: you can redistribute it and/or modify +# it under the terms of the GNU General Public License as published by +# the Free Software Foundation, either version 3 of the License, or +# (at your option) any later version. +# +# Gajim is distributed in the hope that it will be useful, +# but WITHOUT ANY WARRANTY; without even the implied warranty of +# MERCHANTABILITY or FITNESS FOR A PARTICULAR PURPOSE. See the +# GNU General Public License for more details. +# +# You should have received a copy of the GNU General Public License +# along with Gajim. If not, see . + +from __future__ import annotations + +from typing import Any +from typing import TYPE_CHECKING + +from gi.repository import Gtk + +from gajim.gtk.const import Setting +from gajim.gtk.const import SettingKind +from gajim.gtk.const import SettingType +from gajim.gtk.settings import SettingsDialog +from gajim.plugins.plugins_i18n import _ + +if TYPE_CHECKING: + from ..stt_voice_messages import STTVoiceMessagesPlugin + + +class STTVoiceMessagesConfigDialog(SettingsDialog): + def __init__(self, plugin: STTVoiceMessagesPlugin, parent: Gtk.Window) -> None: + + self.plugin = plugin + model = plugin.model + + if not model.available(): + status = _('onnx-asr is not installed. Run: pip install "onnx-asr[hub]"') + elif model.is_loaded: + status = _("Model is loaded.") + elif model.will_download: + status = _("Model files will be downloaded on first use.") + else: + status = _("Model files are downloaded and ready.") + + settings = [ + Setting( + SettingKind.SWITCH, + _("Transcribe automatically"), + SettingType.VALUE, + bool(plugin.config["auto_transcribe"]), + callback=self._on_setting, + data="auto_transcribe", + desc=_("Transcribe voice messages as soon as they are displayed"), + ), + Setting( + SettingKind.ENTRY, + _("Model"), + SettingType.VALUE, + str(plugin.config["model_id"]), + callback=self._on_setting, + data="model_id", + desc=_("onnx-asr model name or Hugging Face repository"), + ), + Setting( + SettingKind.DROPDOWN, + _("Quantization"), + SettingType.VALUE, + str(plugin.config["quantization"]), + callback=self._on_setting, + data="quantization", + props={ + "data": { + "int8": _("int8 (fast, ~600 MB download)"), + "fp32": _("Full precision (~2.4 GB download)"), + } + }, + desc=_("Applied the next time the model is loaded"), + ), + Setting(SettingKind.GENERIC, _("Status"), SettingType.VALUE, desc=status), + ] + + SettingsDialog.__init__( + self, + parent, + _("STT Voice Messages Configuration"), + Gtk.DialogFlags.MODAL, + settings, + "", + ) + + def _on_setting(self, value: Any, data: Any) -> None: + self.plugin.config[data] = value + if data in ("model_id", "quantization"): + self.plugin.model.set_config( + str(self.plugin.config["model_id"]), + str(self.plugin.config["quantization"]), + ) diff --git a/stt_voice_messages/gtk/sttbox.py b/stt_voice_messages/gtk/sttbox.py new file mode 100644 index 0000000..dc7e167 --- /dev/null +++ b/stt_voice_messages/gtk/sttbox.py @@ -0,0 +1,155 @@ +# This file is part of Gajim. +# +# Gajim is free software: you can redistribute it and/or modify +# it under the terms of the GNU General Public License as published by +# the Free Software Foundation, either version 3 of the License, or +# (at your option) any later version. +# +# Gajim is distributed in the hope that it will be useful, +# but WITHOUT ANY WARRANTY; without even the implied warranty of +# MERCHANTABILITY or FITNESS FOR A PARTICULAR PURPOSE. See the +# GNU General Public License for more details. +# +# You should have received a copy of the GNU General Public License +# along with Gajim. If not, see . + +from __future__ import annotations + +from typing import Any + +import logging +from pathlib import Path + +from gi.repository import Adw +from gi.repository import Gio +from gi.repository import GLib +from gi.repository import GObject +from gi.repository import Gtk + +from gajim.plugins.plugins_i18n import _ + +from ..model import OnnxAsrModel + +log = logging.getLogger("gajim.p.stt_voice_messages_sttbox") + + +class STTBox(Gtk.Box): + def __init__(self, model: OnnxAsrModel, audio_file: Path) -> None: + Gtk.Box.__init__(self, orientation=Gtk.Orientation.VERTICAL, spacing=6) + self.set_margin_top(6) + self.set_halign(Gtk.Align.CENTER) + self.set_visible(False) + + self._model = model + self._audio_file = audio_file + self._task: Gio.Task | None = None + + self._transcribe_button = Gtk.Button.new_from_icon_name( + "lucide-square-pen-symbolic" + ) + self._transcribe_button.set_tooltip_text(_("Transcribe voice message")) + self._transcribe_button.set_valign(Gtk.Align.CENTER) + self._transcribe_button.connect("clicked", self._on_transcribe_clicked) + + self._spinner = Adw.Spinner(valign=Gtk.Align.CENTER, visible=False) + + self._transcription_label = Gtk.Label() + self._transcription_label.set_max_width_chars(40) + self._transcription_label.set_wrap(True) + self._transcription_label.set_xalign(0) + self._transcription_label.set_selectable(True) + + self.append(self._spinner) + self.append(self._transcription_label) + + @property + def button(self) -> Gtk.Button: + return self._transcribe_button + + def transcribe(self) -> None: + self._on_transcribe_clicked(self._transcribe_button) + + def _on_transcribe_clicked(self, _button: Gtk.Button) -> None: + if not self._model.available(): + self._show_status( + _('onnx-asr is not installed. Run: pip install "onnx-asr[hub]"'), + busy=False, + ) + return + + if self._model.is_loaded: + text = _("Transcribing…") + elif self._model.will_download: + text = _("Downloading model…") + else: + text = _("Loading model…") + + self._show_status(text, busy=True) + self._run_bg(self._model.load, self._on_load_done) + + def _on_load_done(self) -> None: + try: + self._finish() + except Exception as e: + self._show_status(_("Error: {}").format(e), busy=False) + return + + self._show_status(_("Transcribing…"), busy=True) + self._run_bg( + lambda: self._model.recognize(self._audio_file, self._on_partial), + self._show_result, + ) + + def _on_partial(self, text: str) -> None: + GLib.idle_add(self._set_partial, text) + + def _set_partial(self, text: str) -> None: + self._transcription_label.set_text(text) + + def _show_result(self) -> None: + try: + text = self._finish() + except Exception as e: + self._show_status(_("Error: {}").format(e), busy=False) + return + + text = text.strip() + if not text: + text = _("No speech detected.") + self._show_status(text, busy=False) + + def _run_bg(self, fn, on_done) -> None: + self._bg_fn = fn + + def _cb(_source: GObject.Object, _result: Gio.AsyncResult, _data: None) -> None: + on_done() + + self._task = Gio.Task.new(self, None, _cb, None) + self._task.run_in_thread(self._thread_cb) + + @staticmethod + def _thread_cb( + task: Gio.Task, + source: "STTBox", + _data: None, + _cancel: Gio.Cancellable, + ) -> None: + try: + task.return_value(source._bg_fn()) + except Exception as e: + task.return_value(e) + + def _finish(self) -> Any: + task = self._task + self._task = None + value = task.propagate_value().value + if isinstance(value, Exception): + raise value + return value + + def _show_status(self, text: str, busy: bool) -> None: + self._transcription_label.set_text(text) + self._transcription_label.set_visible(True) + self.set_visible(True) + self._spinner.set_visible(busy) + self._transcribe_button.set_sensitive(not busy) diff --git a/stt_voice_messages/model.py b/stt_voice_messages/model.py new file mode 100644 index 0000000..e976705 --- /dev/null +++ b/stt_voice_messages/model.py @@ -0,0 +1,164 @@ +# This file is part of Gajim. +# +# Gajim is free software: you can redistribute it and/or modify +# it under the terms of the GNU General Public License as published by +# the Free Software Foundation, either version 3 of the License, or +# (at your option) any later version. +# +# Gajim is distributed in the hope that it will be useful, +# but WITHOUT ANY WARRANTY; without even the implied warranty of +# MERCHANTABILITY or FITNESS FOR A PARTICULAR PURPOSE. See the +# GNU General Public License for more details. +# +# You should have received a copy of the GNU General Public License +# along with Gajim. If not, see . + +from __future__ import annotations + +from typing import Any + +import gc +import logging +import threading +from collections.abc import Callable +from importlib.util import find_spec +from pathlib import Path + +import numpy as np +from gi.repository import GLib +from gi.repository import Gst + +log = logging.getLogger("gajim.p.sttvm_model") + +_IDLE_UNLOAD_SECONDS = 300 + + +def load_audio(path: Path, sample_rate: int = 16000) -> np.ndarray: + Gst.init(None) + pipeline = Gst.parse_launch( + "filesrc name=src ! decodebin ! audioconvert ! audioresample ! " + f"audio/x-raw,format=F32LE,rate={sample_rate},channels=1 ! " + "appsink name=sink sync=false" + ) + + pipeline.get_by_name("src").set_property("location", str(path)) + sink = pipeline.get_by_name("sink") + + chunks: list[np.ndarray] = [] + try: + pipeline.set_state(Gst.State.PLAYING) + while (sample := sink.emit("try-pull-sample", 10 * Gst.SECOND)) is not None: + buf = sample.get_buffer() + _, info = buf.map(Gst.MapFlags.READ) + chunks.append(np.frombuffer(bytes(info.data), dtype=np.float32)) + buf.unmap(info) + pipeline.set_state(Gst.State.NULL) + # wait for async cleanup so decodebin/appsink release memory + pipeline.get_state(Gst.CLOCK_TIME_NONE) + finally: + del pipeline + gc.collect() + + if not chunks: + raise RuntimeError(f"Could not decode audio: {path}") + return np.concatenate(chunks) + + +class OnnxAsrModel: + def __init__(self, model_id: str, quantization: str) -> None: + self._model_id = model_id + self._quantization = quantization + self._model: Any = None + self._loaded = False + self._busy = False + self._recognize_lock = threading.Lock() + self._unload_source: int | None = None + + @staticmethod + def available() -> bool: + return find_spec("onnx_asr") is not None + + @property + def is_loaded(self) -> bool: + return self._loaded + + @property + def model_id(self) -> str: + return self._model_id + + @property + def will_download(self) -> bool: + if self._loaded or not self.available(): + return False + from huggingface_hub import try_to_load_from_cache + from onnx_asr.resolver import model_repos + + repo = model_repos.get(self._model_id, self._model_id) + if "/" not in repo: + # Local path, nothing to download + return False + return not isinstance(try_to_load_from_cache(repo, "config.json"), str) + + def set_config(self, model_id: str, quantization: str) -> None: + if model_id == self._model_id and quantization == self._quantization: + return + self.unload_now() + self._model_id = model_id + self._quantization = quantization + + def load(self) -> None: + if self._loaded: + self._schedule_unload() + return + + import onnx_asr + + log.debug("Loading model %s", self._model_id) + providers = ["CUDAExecutionProvider", "CPUExecutionProvider"] + quantization = None if self._quantization == "fp32" else self._quantization + model = onnx_asr.load_model(self._model_id, quantization=quantization, providers=providers) + self._model = model.with_vad(onnx_asr.load_vad("silero")) + self._loaded = True + self._schedule_unload() + + def recognize( + self, + audio_file: Path, + on_partial: Callable[[str], None] | None = None, + ) -> str: + # global lock, one transcription at a time + with self._recognize_lock: + self.load() + self._busy = True + try: + audio = load_audio(audio_file) + parts: list[str] = [] + text = "" + for segment in self._model.recognize(audio, sample_rate=16000): + parts.append(segment.text.strip()) + text = " ".join(part for part in parts if part) + if on_partial is not None: + on_partial(text) + finally: + self._busy = False + self._schedule_unload() + return text + + def unload_now(self) -> None: + self._model = None + self._loaded = False + gc.collect() + + def _schedule_unload(self) -> None: + if self._unload_source is not None: + GLib.source_remove(self._unload_source) + self._unload_source = GLib.timeout_add_seconds( + _IDLE_UNLOAD_SECONDS, self._unload_idle + ) + + def _unload_idle(self) -> bool: + if self._busy: + return GLib.SOURCE_CONTINUE + self._unload_source = None + self.unload_now() + return GLib.SOURCE_REMOVE diff --git a/stt_voice_messages/org.gajim.Gajim.Plugin.stt_voice_messages.metainfo.xml b/stt_voice_messages/org.gajim.Gajim.Plugin.stt_voice_messages.metainfo.xml new file mode 100644 index 0000000..dd3ed6e --- /dev/null +++ b/stt_voice_messages/org.gajim.Gajim.Plugin.stt_voice_messages.metainfo.xml @@ -0,0 +1,11 @@ + + + org.gajim.Gajim.Plugin.stt_voice_messages + org.gajim.Gajim + STT Voice Messages Plugin + Transcribes voice messages to text + https://gajim.org/ + CC-BY-SA-3.0 + GPL-3.0-only + gajim-devel_AT_gajim.org + diff --git a/stt_voice_messages/plugin-manifest.json b/stt_voice_messages/plugin-manifest.json new file mode 100644 index 0000000..aa8b87b --- /dev/null +++ b/stt_voice_messages/plugin-manifest.json @@ -0,0 +1,20 @@ +{ + "authors": [ + "hueso" + ], + "description": "Transcribes voice messages to text.", + "homepage": "https://dev.gajim.org/gajim/gajim-plugins/wikis/STTVoiceMessagesPlugin", + "config_dialog": true, + "name": "STT Voice Messages", + "platforms": [ + "others", + "linux", + "darwin", + "win32" + ], + "requirements": [ + "gajim>=2.5.0" + ], + "short_name": "stt_voice_messages", + "version": "0.1.0" +} diff --git a/stt_voice_messages/stt_voice_messages.png b/stt_voice_messages/stt_voice_messages.png new file mode 100644 index 0000000000000000000000000000000000000000..f7cea2fa10d4b7b7b8324757f9483aaf0212d3b3 GIT binary patch literal 371 zcmV-(0gV2MP)ftCGW@rMlzaWr`*?P1kkk`<_0} z^Lz}z0hVDHu&(Q#K$7da20;+c0n;==*LBb|4c(KH$8n^dK*{%0Kvh+c=Q(`N2SNx~ zmc0fbrx!)Bzat<5;5hCYplw@v+NWh%2C^)l1MnsJX`24)f+&ihC<;yfc8TcdUvM0U zwpEfO;CUV}99-8$mj?+4+qeUMS(d;s%-2CW9c9n+Y$7d55-RbQge&$EU;q`ytQB#? Ri8lZM002ovPDHLkV1lR+lr8`O literal 0 HcmV?d00001 diff --git a/stt_voice_messages/stt_voice_messages.py b/stt_voice_messages/stt_voice_messages.py new file mode 100644 index 0000000..b80fed1 --- /dev/null +++ b/stt_voice_messages/stt_voice_messages.py @@ -0,0 +1,70 @@ +# This file is part of Gajim. +# +# Gajim is free software: you can redistribute it and/or modify +# it under the terms of the GNU General Public License as published by +# the Free Software Foundation, either version 3 of the License, or +# (at your option) any later version. +# +# Gajim is distributed in the hope that it will be useful, +# but WITHOUT ANY WARRANTY; without even the implied warranty of +# MERCHANTABILITY or FITNESS FOR A PARTICULAR PURPOSE. See the +# GNU General Public License for more details. +# +# You should have received a copy of the GNU General Public License +# along with Gajim. If not, see . + +from __future__ import annotations + +import logging +from functools import partial +from pathlib import Path + +from gi.repository import Gtk + +from gajim.plugins import GajimPlugin +from gajim.plugins.plugins_i18n import _ + +from .gtk.config_dialog import STTVoiceMessagesConfigDialog +from .gtk.sttbox import STTBox +from .model import OnnxAsrModel + +log = logging.getLogger("gajim.p.stt_voice_messages") + + +class STTVoiceMessagesPlugin(GajimPlugin): + def init(self) -> None: + self.description = _("Transcribes voice messages to text.") + + self.config_default_values = { + "auto_transcribe": (False, ""), + "model_id": ("nemo-parakeet-tdt-0.6b-v3", ""), + "quantization": ("int8", ""), + } + + self._model = OnnxAsrModel( + str(self.config["model_id"]), str(self.config["quantization"]) + ) + self.config_dialog = partial(STTVoiceMessagesConfigDialog, self) + + self.gui_extension_points = { + "preview_audio": (self._on_preview_audio, None), + } + + @property + def model(self) -> OnnxAsrModel: + return self._model + + def deactivate(self) -> None: + self._model.unload_now() + + def _on_preview_audio( + self, drawing_box: Gtk.Box, control_box: Gtk.Box, audio_file: Path + ) -> None: + content_box = drawing_box.get_parent().get_parent() + + stt_box = STTBox(self._model, audio_file) + control_box.append(stt_box.button) + content_box.append(stt_box) + + if self.config["auto_transcribe"]: + stt_box.transcribe()