diff --git a/CMakeLists.txt b/CMakeLists.txt index 39c16a8..9024c4f 100644 --- a/CMakeLists.txt +++ b/CMakeLists.txt @@ -6,7 +6,11 @@ add_subdirectory(whisper.cpp) pybind11_add_module(_pywhispercpp src/main.cpp + src/parakeet_bindings.cpp ) -target_link_libraries (_pywhispercpp PRIVATE whisper) - +target_link_libraries(_pywhispercpp PRIVATE whisper parakeet) +target_include_directories(_pywhispercpp PRIVATE + ${CMAKE_CURRENT_SOURCE_DIR}/whisper.cpp/include + ${CMAKE_CURRENT_SOURCE_DIR}/whisper.cpp/ggml/include +) diff --git a/README.md b/README.md index 018f18e..db8f1fc 100644 --- a/README.md +++ b/README.md @@ -126,6 +126,19 @@ segments = model.transcribe('file.mp3', new_segment_callback=print) * You can pass any `whisper.cpp` [parameter](https://absadiki.github.io/pywhispercpp/#pywhispercpp.constants.PARAMS_SCHEMA) as a keyword argument to the `Model` class or to the `transcribe` function. * Check the [Model](https://absadiki.github.io/pywhispercpp/#pywhispercpp.model.Model) class documentation for more details. +### Parakeet (NVIDIA) + +Parakeet is a separate whisper.cpp API from Whisper. Point `ParakeetModel` at a local GGUF (for example from [ggml-org/parakeet-GGUF](https://huggingface.co/ggml-org/parakeet-GGUF)); there is no automatic download helper yet. + +```python +from pywhispercpp import ParakeetModel + +model = ParakeetModel("/path/to/ggml-parakeet-tdt-0.6b-v3-q5_0.bin", n_threads=4) +segments = model.transcribe("file.wav") +for segment in segments: + print(segment.text) +``` + # Examples ## CLI diff --git a/pywhispercpp/__init__.py b/pywhispercpp/__init__.py index 8b13789..22f6f85 100644 --- a/pywhispercpp/__init__.py +++ b/pywhispercpp/__init__.py @@ -1 +1,4 @@ +from pywhispercpp.model import Model, Segment +from pywhispercpp.parakeet_model import ParakeetModel +__all__ = ["Model", "Segment", "ParakeetModel"] diff --git a/pywhispercpp/parakeet_model.py b/pywhispercpp/parakeet_model.py new file mode 100644 index 0000000..8b33275 --- /dev/null +++ b/pywhispercpp/parakeet_model.py @@ -0,0 +1,124 @@ +#!/usr/bin/env python +# -*- coding: utf-8 -*- + +""" +Thin Python helper for whisper.cpp Parakeet (`parakeet.h`). + +Models are local GGUF files (e.g. https://huggingface.co/ggml-org/parakeet-GGUF). + +Example:: + + from pywhispercpp import ParakeetModel + + model = ParakeetModel("/path/to/ggml-parakeet-tdt-0.6b-v3-q5_0.bin", n_threads=4) + for seg in model.transcribe("audio.wav"): + print(seg.text) +""" + +from __future__ import annotations + +import os +from pathlib import Path +from typing import List, Optional, Union + +import numpy as np + +import _pywhispercpp as pw +from pywhispercpp.model import Model, Segment + + +class ParakeetModel: + """ + Load a local Parakeet GGUF and run ``parakeet_full``. + + Reuses :class:`~pywhispercpp.model.Segment` and ``Model._load_audio`` so we + do not duplicate wav/ffmpeg decoding. + """ + + def __init__( + self, + model_path: str, + n_threads: Optional[int] = None, + use_gpu: bool = True, + gpu_device: int = 0, + **params, + ): + """ + :param model_path: Path to a Parakeet GGUF model file. + :param n_threads: Inference threads (default min(4, CPU count)). + :param use_gpu: Forwarded to ``parakeet_context_params.use_gpu``. + :param gpu_device: Forwarded to ``parakeet_context_params.gpu_device``. + :param params: Extra fields on ``parakeet_full_params`` (e.g. ``offset_ms``). + """ + self._ctx = None + path = Path(model_path).expanduser() + if not path.is_file(): + raise FileNotFoundError( + f"Parakeet model not found: {model_path}. " + "Download a GGUF from https://huggingface.co/ggml-org/parakeet-GGUF" + ) + self.model_path = str(path.resolve()) + + cparams = pw.parakeet_context_default_params() + cparams.use_gpu = use_gpu + cparams.gpu_device = gpu_device + + self._params = pw.parakeet_full_default_params( + pw.parakeet_sampling_strategy.PARAKEET_SAMPLING_GREEDY + ) + self._params.n_threads = int( + n_threads if n_threads is not None else min(4, os.cpu_count() or 4) + ) + for key, value in params.items(): + setattr(self._params, key, value) + + self._ctx = pw.parakeet_init_from_file_with_params(self.model_path, cparams) + if self._ctx is None: + raise RuntimeError(f"Failed to load Parakeet model: {self.model_path}") + + def transcribe( + self, + media: Union[str, np.ndarray], + **params, + ) -> List[Segment]: + """ + Transcribe a media path or float32 mono PCM at 16 kHz. + + :param media: File path (wav preferred; other formats need ffmpeg) or + 1-D float32 ``numpy.ndarray``. + :param params: Overrides for ``parakeet_full_params`` for this call. + :return: List of :class:`~pywhispercpp.model.Segment`. + """ + if isinstance(media, np.ndarray): + audio = np.ascontiguousarray(media, dtype=np.float32).ravel() + else: + if not Path(media).exists(): + raise FileNotFoundError(media) + audio = Model._load_audio(str(media)) + + for key, value in params.items(): + setattr(self._params, key, value) + + ret = pw.parakeet_full(self._ctx, self._params, audio, int(audio.size)) + if ret != 0: + raise RuntimeError(f"parakeet_full failed with code {ret}") + + n = pw.parakeet_full_n_segments(self._ctx) + segments: List[Segment] = [] + for i in range(n): + t0 = pw.parakeet_full_get_segment_t0(self._ctx, i) + t1 = pw.parakeet_full_get_segment_t1(self._ctx, i) + text = pw.parakeet_full_get_segment_text(self._ctx, i).decode( + "utf-8", errors="replace" + ).strip() + segments.append(Segment(t0, t1, text)) + return segments + + def __del__(self): + ctx = getattr(self, "_ctx", None) + if ctx is not None: + try: + pw.parakeet_free(ctx) + except Exception: + pass + self._ctx = None diff --git a/src/main.cpp b/src/main.cpp index 6bc3c00..feb94b9 100644 --- a/src/main.cpp +++ b/src/main.cpp @@ -16,6 +16,7 @@ #include #include "whisper.h" +#include "parakeet_bindings.h" #define STRINGIFY(x) #x @@ -1384,8 +1385,7 @@ PYBIND11_MODULE(_pywhispercpp, m) { m.def("whisper_vad_free_segments", &whisper_vad_free_segments_wrapper); m.def("whisper_vad_free", &whisper_vad_free_wrapper); - - + init_parakeet_bindings(m); #ifdef VERSION_INFO m.attr("__version__") = MACRO_STRINGIFY(VERSION_INFO); diff --git a/src/parakeet_bindings.cpp b/src/parakeet_bindings.cpp new file mode 100644 index 0000000..79e954b --- /dev/null +++ b/src/parakeet_bindings.cpp @@ -0,0 +1,126 @@ +/** + * Python bindings for whisper.cpp Parakeet API (parakeet.h). + */ + +#include +#include + +#include "parakeet.h" + +#include + +namespace py = pybind11; + +// Opaque context wrapper (parakeet_context is incomplete in the public header). +struct parakeet_context_wrapper { + parakeet_context * ptr; +}; + +static py::object parakeet_init_from_file_with_params_wrapper( + const char * path_model, + struct parakeet_context_params cparams) { + // Release the GIL only for the native load. py::cast / py::none need the GIL. + struct parakeet_context * ctx; + { + py::gil_scoped_release release; + ctx = parakeet_init_from_file_with_params(path_model, cparams); + } + if (!ctx) { + return py::none(); + } + struct parakeet_context_wrapper w; + w.ptr = ctx; + return py::cast(w); +} + +static void parakeet_free_wrapper(py::object ctx_obj) { + if (ctx_obj.is_none()) { + return; + } + auto * ctx_w = ctx_obj.cast(); + if (ctx_w && ctx_w->ptr) { + parakeet_free(ctx_w->ptr); + ctx_w->ptr = nullptr; + } +} + +static int parakeet_full_wrapper( + struct parakeet_context_wrapper * ctx_w, + struct parakeet_full_params params, + py::array_t samples, + int n_samples) { + if (!ctx_w || !ctx_w->ptr) { + return -1; + } + py::buffer_info buf = samples.request(); + float * samples_ptr = static_cast(buf.ptr); + + py::gil_scoped_release release; + return parakeet_full(ctx_w->ptr, params, samples_ptr, n_samples); +} + +static int parakeet_full_n_segments_wrapper(struct parakeet_context_wrapper * ctx_w) { + if (!ctx_w || !ctx_w->ptr) { + return 0; + } + py::gil_scoped_release release; + return parakeet_full_n_segments(ctx_w->ptr); +} + +static int64_t parakeet_full_get_segment_t0_wrapper( + struct parakeet_context_wrapper * ctx_w, int i_segment) { + return parakeet_full_get_segment_t0(ctx_w->ptr, i_segment); +} + +static int64_t parakeet_full_get_segment_t1_wrapper( + struct parakeet_context_wrapper * ctx_w, int i_segment) { + return parakeet_full_get_segment_t1(ctx_w->ptr, i_segment); +} + +static py::bytes parakeet_full_get_segment_text_wrapper( + struct parakeet_context_wrapper * ctx_w, int i_segment) { + const char * c_array = parakeet_full_get_segment_text(ctx_w->ptr, i_segment); + if (!c_array) { + return py::bytes("", 0); + } + return py::bytes(c_array, strlen(c_array)); +} + +void init_parakeet_bindings(py::module_ & m) { + m.attr("PARAKEET_SAMPLE_RATE") = PARAKEET_SAMPLE_RATE; + + py::enum_(m, "parakeet_sampling_strategy") + .value("PARAKEET_SAMPLING_GREEDY", parakeet_sampling_strategy::PARAKEET_SAMPLING_GREEDY) + .export_values(); + + py::class_(m, "parakeet_context"); + + py::class_(m, "parakeet_context_params") + .def(py::init<>()) + .def_readwrite("use_gpu", ¶keet_context_params::use_gpu) + .def_readwrite("gpu_device", ¶keet_context_params::gpu_device); + + py::class_(m, "parakeet_full_params") + .def(py::init<>()) + .def_readwrite("strategy", ¶keet_full_params::strategy) + .def_readwrite("n_threads", ¶keet_full_params::n_threads) + .def_readwrite("offset_ms", ¶keet_full_params::offset_ms) + .def_readwrite("duration_ms", ¶keet_full_params::duration_ms) + .def_readwrite("no_context", ¶keet_full_params::no_context) + .def_readwrite("audio_ctx", ¶keet_full_params::audio_ctx); + + m.def("parakeet_version", ¶keet_version); + m.def("parakeet_context_default_params", ¶keet_context_default_params); + m.def("parakeet_full_default_params", ¶keet_full_default_params, py::arg("strategy")); + // No call_guard gil_scoped_release: this binding returns py::object via py::cast. + m.def("parakeet_init_from_file_with_params", + ¶keet_init_from_file_with_params_wrapper, + py::arg("path_model"), py::arg("params")); + m.def("parakeet_free", ¶keet_free_wrapper); + m.def("parakeet_full", ¶keet_full_wrapper, + py::arg("ctx"), py::arg("params"), py::arg("samples"), py::arg("n_samples")); + m.def("parakeet_full_n_segments", ¶keet_full_n_segments_wrapper); + m.def("parakeet_full_get_segment_t0", ¶keet_full_get_segment_t0_wrapper); + m.def("parakeet_full_get_segment_t1", ¶keet_full_get_segment_t1_wrapper); + m.def("parakeet_full_get_segment_text", ¶keet_full_get_segment_text_wrapper); +} diff --git a/src/parakeet_bindings.h b/src/parakeet_bindings.h new file mode 100644 index 0000000..bd9f4a8 --- /dev/null +++ b/src/parakeet_bindings.h @@ -0,0 +1,5 @@ +#pragma once + +#include + +void init_parakeet_bindings(pybind11::module_ & m); diff --git a/tests/test_parakeet.py b/tests/test_parakeet.py new file mode 100644 index 0000000..900b38a --- /dev/null +++ b/tests/test_parakeet.py @@ -0,0 +1,63 @@ +#!/usr/bin/env python +# -*- coding: utf-8 -*- + +"""Parakeet binding smoke tests (no large model download).""" + +import unittest +from unittest import TestCase + +import _pywhispercpp as pw + + +class TestParakeetBindings(TestCase): + def test_import_and_constants(self): + self.assertEqual(pw.PARAKEET_SAMPLE_RATE, 16000) + + def test_parakeet_version(self): + version = pw.parakeet_version() + self.assertIsInstance(version, str) + self.assertGreater(len(version), 0) + + def test_context_default_params(self): + params = pw.parakeet_context_default_params() + params.use_gpu = False + params.gpu_device = 0 + self.assertFalse(params.use_gpu) + self.assertEqual(params.gpu_device, 0) + + def test_full_default_params(self): + params = pw.parakeet_full_default_params( + pw.parakeet_sampling_strategy.PARAKEET_SAMPLING_GREEDY + ) + self.assertGreater(params.n_threads, 0) + params.n_threads = 2 + self.assertEqual(params.n_threads, 2) + + def test_init_missing_file_returns_none(self): + params = pw.parakeet_context_default_params() + params.use_gpu = False + ctx = pw.parakeet_init_from_file_with_params( + "/nonexistent/path/parakeet-model.bin", + params, + ) + self.assertIsNone(ctx) + + def test_free_none_safe(self): + pw.parakeet_free(None) + + def test_python_package_export(self): + from pywhispercpp import ParakeetModel, Segment, Model + + self.assertTrue(callable(ParakeetModel)) + self.assertTrue(callable(Model)) + self.assertTrue(callable(Segment)) + + def test_parakeet_model_missing_path(self): + from pywhispercpp import ParakeetModel + + with self.assertRaises(FileNotFoundError): + ParakeetModel("/no/such/parakeet.gguf") + + +if __name__ == "__main__": + unittest.main() diff --git a/whisper.cpp b/whisper.cpp index 9386f23..080bbbe 160000 --- a/whisper.cpp +++ b/whisper.cpp @@ -1 +1 @@ -Subproject commit 9386f239401074690479731c1e41683fbbeac557 +Subproject commit 080bbbe85230f624f0b52127f1ae1218247989f9