diff --git a/.github/scripts/test-rust.sh b/.github/scripts/test-rust.sh new file mode 100755 index 0000000000..22712e27b9 --- /dev/null +++ b/.github/scripts/test-rust.sh @@ -0,0 +1,9 @@ +#!/usr/bin/env bash + +set -ex + +cd rust-api-examples + +./run-version.sh + +./run-streaming-zipformer.sh diff --git a/.github/workflows/test-rust-package.yaml b/.github/workflows/test-rust-package.yaml index fc7a819495..6f403964c6 100644 --- a/.github/workflows/test-rust-package.yaml +++ b/.github/workflows/test-rust-package.yaml @@ -70,12 +70,7 @@ jobs: echo "RUSTFLAGS: $RUSTFLAGS" - # Test using published crates.io dependencies - - name: Test Published Crates + - name: Run test shell: bash run: | - cd rust-api-examples - git checkout Cargo.toml - cargo clean - # cargo test --locked --all-features - cargo run --example version + ./.github/scripts/test-rust.sh diff --git a/.github/workflows/test-rust.yaml b/.github/workflows/test-rust.yaml index ffb255dc31..b60f67a531 100644 --- a/.github/workflows/test-rust.yaml +++ b/.github/workflows/test-rust.yaml @@ -88,8 +88,12 @@ jobs: sed -i.bak 's|^sherpa-onnx *=.*|sherpa-onnx = { path = "../sherpa-onnx/rust/sherpa-onnx" }|' Cargo.toml - git diff . cargo clean cargo run --example version + + - name: Run test + shell: bash + run: | + ./.github/scripts/test-rust.sh diff --git a/c-api-examples/streaming-zipformer-c-api.c b/c-api-examples/streaming-zipformer-c-api.c index 6011186ea1..b54f8ba81f 100644 --- a/c-api-examples/streaming-zipformer-c-api.c +++ b/c-api-examples/streaming-zipformer-c-api.c @@ -62,6 +62,7 @@ int32_t main() { memset(&recognizer_config, 0, sizeof(recognizer_config)); recognizer_config.decoding_method = "greedy_search"; recognizer_config.model_config = online_model_config; + recognizer_config.enable_endpoint = 1; const SherpaOnnxOnlineRecognizer *recognizer = SherpaOnnxCreateOnlineRecognizer(&recognizer_config); diff --git a/rust-api-examples/Cargo.lock b/rust-api-examples/Cargo.lock index f1eb03ace8..cf16558f68 100644 --- a/rust-api-examples/Cargo.lock +++ b/rust-api-examples/Cargo.lock @@ -2,24 +2,122 @@ # It is not intended for manual editing. version = 4 +[[package]] +name = "itoa" +version = "1.0.17" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "92ecc6618181def0457392ccd0ee51198e065e016d1d527a7ac1b6dc7c1f09d2" + +[[package]] +name = "memchr" +version = "2.8.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "f8ca58f447f06ed17d5fc4043ce1b10dd205e060fb3ce5b979b8ed8e59ff3f79" + +[[package]] +name = "proc-macro2" +version = "1.0.106" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "8fd00f0bb2e90d81d1044c2b32617f68fcb9fa3bb7640c23e9c748e53fb30934" +dependencies = [ + "unicode-ident", +] + +[[package]] +name = "quote" +version = "1.0.44" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "21b2ebcf727b7760c461f091f9f0f539b77b8e87f2fd88131e7f1b433b3cece4" +dependencies = [ + "proc-macro2", +] + [[package]] name = "rust-api-examples" -version = "0.1.0" +version = "0.1.1" dependencies = [ "sherpa-onnx", ] +[[package]] +name = "serde" +version = "1.0.228" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "9a8e94ea7f378bd32cbbd37198a4a91436180c5bb472411e48b5ec2e2124ae9e" +dependencies = [ + "serde_core", + "serde_derive", +] + +[[package]] +name = "serde_core" +version = "1.0.228" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "41d385c7d4ca58e59fc732af25c3983b67ac852c1a25000afe1175de458b67ad" +dependencies = [ + "serde_derive", +] + +[[package]] +name = "serde_derive" +version = "1.0.228" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "d540f220d3187173da220f885ab66608367b6574e925011a9353e4badda91d79" +dependencies = [ + "proc-macro2", + "quote", + "syn", +] + +[[package]] +name = "serde_json" +version = "1.0.149" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "83fc039473c5595ace860d8c4fafa220ff474b3fc6bfdb4293327f1a37e94d86" +dependencies = [ + "itoa", + "memchr", + "serde", + "serde_core", + "zmij", +] + [[package]] name = "sherpa-onnx" -version = "0.1.0" +version = "0.1.1" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "d294fdded2188c7dd3255f59bbc3182ae0b1fb8a2536de214d0db47da74d8524" +checksum = "9678e5b5315d1bad0b78c7be2e27820ec9fce51d814b844aa8b5cd64fc36e417" dependencies = [ + "serde", + "serde_json", "sherpa-onnx-sys", ] [[package]] name = "sherpa-onnx-sys" -version = "0.1.0" +version = "0.1.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "50208e12ba107c6d3063d28e69d1855fd81d032cf197fc04c4e343f7f7f1ac7f" + +[[package]] +name = "syn" +version = "2.0.117" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "e665b8803e7b1d2a727f4023456bbbbe74da67099c585258af0ad9c5013b9b99" +dependencies = [ + "proc-macro2", + "quote", + "unicode-ident", +] + +[[package]] +name = "unicode-ident" +version = "1.0.24" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "e6e4313cd5fcd3dad5cafa179702e2b244f760991f45397d14d4ebf38247da75" + +[[package]] +name = "zmij" +version = "1.0.21" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "5cef134a691ff7d30551b81ca0171f1c9349ae150f4fef1acc81f9f95105080e" +checksum = "b8848ee67ecc8aedbaf3e4122217aff892639231befc6a1b58d29fff4c2cabaa" diff --git a/rust-api-examples/Cargo.toml b/rust-api-examples/Cargo.toml index 472b2ed880..587b76168e 100644 --- a/rust-api-examples/Cargo.toml +++ b/rust-api-examples/Cargo.toml @@ -1,9 +1,9 @@ [package] name = "rust-api-examples" -version = "0.1.0" +version = "0.1.1" edition = "2021" [dependencies] -sherpa-onnx = "0.1.0" +sherpa-onnx = "0.1.1" # sherpa-onnx = { path = "../sherpa-onnx/rust/sherpa-onnx" } diff --git a/rust-api-examples/README.md b/rust-api-examples/README.md index 7e304b88c4..be6c5fd919 100644 --- a/rust-api-examples/README.md +++ b/rust-api-examples/README.md @@ -31,6 +31,8 @@ export RUSTFLAGS="-C link-arg=-Wl,-rpath,$SHERPA_ONNX_LIB_DIR" ## Run it +### Example 1: Show sherpa-onnx version + ```bash cargo run --example version ``` @@ -41,6 +43,17 @@ otool -l target/debug/examples/version | grep -A2 LC_RPATH ``` to check the RPATH. +### Example 2: ASR with streaming zipformer + +```bash +wget https://github.com/k2-fsa/sherpa-onnx/releases/download/asr-models/sherpa-onnx-streaming-zipformer-en-2023-06-21.tar.bz2 + +tar xvf sherpa-onnx-streaming-zipformer-en-2023-06-21.tar.bz2 +rm sherpa-onnx-streaming-zipformer-en-2023-06-21.tar.bz2 + +cargo run --example streaming_zipformer +``` + # Alternative rust bindings for sherpa-onnx Please see also https://github.com/thewh1teagle/sherpa-rs diff --git a/rust-api-examples/examples/streaming_zipformer.rs b/rust-api-examples/examples/streaming_zipformer.rs new file mode 100644 index 0000000000..54609606b2 --- /dev/null +++ b/rust-api-examples/examples/streaming_zipformer.rs @@ -0,0 +1,93 @@ +// Copyright (c) 2026 Xiaomi Corporation +// +// This file demonstrates how to use streaming Zipformer with sherpa-onnx's +// Rust API for speech recognition. +// +// See ../README.md for how to run it +// +// Note that even if we use a wave file as an example, this model supports +// real-time streaming speech recognition. You can read audio samples +// from a microphone. + +use sherpa_onnx::{OnlineRecognizer, OnlineRecognizerConfig, Wave}; + +fn main() { + let wav_path = "sherpa-onnx-streaming-zipformer-en-2023-06-21/test_wavs/1.wav"; + let encoder_path = + "sherpa-onnx-streaming-zipformer-en-2023-06-21/encoder-epoch-99-avg-1.int8.onnx"; + let decoder_path = "sherpa-onnx-streaming-zipformer-en-2023-06-21/decoder-epoch-99-avg-1.onnx"; + let joiner_path = + "sherpa-onnx-streaming-zipformer-en-2023-06-21/joiner-epoch-99-avg-1.int8.onnx"; + let tokens_path = "sherpa-onnx-streaming-zipformer-en-2023-06-21/tokens.txt"; + let provider = "cpu"; + + let wave = Wave::read(wav_path).expect("Failed to read WAV file"); + + let mut recognizer_config = OnlineRecognizerConfig::default(); + recognizer_config.model_config.transducer.encoder = Some(encoder_path.to_string()); + recognizer_config.model_config.transducer.decoder = Some(decoder_path.to_string()); + recognizer_config.model_config.transducer.joiner = Some(joiner_path.to_string()); + recognizer_config.model_config.tokens = Some(tokens_path.to_string()); + recognizer_config.model_config.provider = Some(provider.to_string()); + recognizer_config.enable_endpoint = true; + + // set to true to see verbose logs + recognizer_config.model_config.debug = true; + + recognizer_config.decoding_method = Some("greedy_search".to_string()); + + let recognizer = + OnlineRecognizer::create(&recognizer_config).expect("Failed to create OnlineRecognizer"); + + let stream = recognizer.create_stream(); + + let mut segment_id = 0; + + // use any positive value as you like + const CHUNK_SIZE: usize = 3200; + + println!( + "Sample rate: {}, num samples: {}, duration: {:.2}s", + wave.sample_rate(), + wave.num_samples(), + wave.num_samples() as f32 / wave.sample_rate() as f32 + ); + + for chunk in wave.samples().chunks(CHUNK_SIZE) { + stream.accept_waveform(wave.sample_rate(), chunk); + + while recognizer.is_ready(&stream) { + recognizer.decode(&stream); + + if let Some(result) = recognizer.get_result(&stream) { + if !result.text.is_empty() { + println!("Segment {}: {}", segment_id, result.text); + } + } + + if recognizer.is_endpoint(&stream) { + recognizer.reset(&stream); + segment_id += 1; + } + } + } + + // Tail padding (~0.3s) + let tail_padding_len = (wave.sample_rate() as f32 * 0.3).round() as usize; + let tail_padding = vec![0.0f32; tail_padding_len]; + + stream.accept_waveform(wave.sample_rate(), &tail_padding); + + stream.input_finished(); + + while recognizer.is_ready(&stream) { + recognizer.decode(&stream); + if let Some(result) = recognizer.get_result(&stream) { + if !result.text.is_empty() { + println!("Segment {}: {}", segment_id, result.text); + } + } + } + + println!("Transcription finished."); +} diff --git a/rust-api-examples/run-streaming-zipformer.sh b/rust-api-examples/run-streaming-zipformer.sh new file mode 100755 index 0000000000..8866340c38 --- /dev/null +++ b/rust-api-examples/run-streaming-zipformer.sh @@ -0,0 +1,12 @@ +#!/usr/bin/env bash +set -ex + +if [ ! -f ./sherpa-onnx-streaming-zipformer-en-2023-06-21/encoder-epoch-99-avg-1.int8.onnx ]; then + curl -SsL -O https://github.com/k2-fsa/sherpa-onnx/releases/download/asr-models/sherpa-onnx-streaming-zipformer-en-2023-06-21.tar.bz2 + + tar xvf sherpa-onnx-streaming-zipformer-en-2023-06-21.tar.bz2 + rm sherpa-onnx-streaming-zipformer-en-2023-06-21.tar.bz2 + ls -lh sherpa-onnx-streaming-zipformer-en-2023-06-21 +fi + +cargo run --example streaming_zipformer diff --git a/rust-api-examples/run-version.sh b/rust-api-examples/run-version.sh new file mode 100755 index 0000000000..1f9fc86250 --- /dev/null +++ b/rust-api-examples/run-version.sh @@ -0,0 +1,3 @@ +#!/usr/bin/env bash +set -ex +cargo run --example version diff --git a/sherpa-onnx/rust/sherpa-onnx-sys/Cargo.toml b/sherpa-onnx/rust/sherpa-onnx-sys/Cargo.toml index c7a9476184..dd761cb870 100644 --- a/sherpa-onnx/rust/sherpa-onnx-sys/Cargo.toml +++ b/sherpa-onnx/rust/sherpa-onnx-sys/Cargo.toml @@ -1,6 +1,6 @@ [package] name = "sherpa-onnx-sys" -version = "0.1.0" +version = "0.1.1" edition = "2021" description = "Raw FFI bindings to the sherpa-onnx C API" license = "Apache-2.0" diff --git a/sherpa-onnx/rust/sherpa-onnx-sys/src/lib.rs b/sherpa-onnx/rust/sherpa-onnx-sys/src/lib.rs index 12ab0da685..f294eb0622 100644 --- a/sherpa-onnx/rust/sherpa-onnx-sys/src/lib.rs +++ b/sherpa-onnx/rust/sherpa-onnx-sys/src/lib.rs @@ -10,3 +10,9 @@ extern "C" { pub fn SherpaOnnxGetGitDate() -> *const c_char; pub fn SherpaOnnxFileExists(filename: *const c_char) -> c_int; } + +pub mod online_asr; +pub mod wave; + +pub use online_asr::*; +pub use wave::*; diff --git a/sherpa-onnx/rust/sherpa-onnx-sys/src/online_asr.rs b/sherpa-onnx/rust/sherpa-onnx-sys/src/online_asr.rs new file mode 100644 index 0000000000..6c62ea1538 --- /dev/null +++ b/sherpa-onnx/rust/sherpa-onnx-sys/src/online_asr.rs @@ -0,0 +1,189 @@ +#![allow(non_camel_case_types)] +#![allow(non_snake_case)] +#![allow(non_upper_case_globals)] + +use std::os::raw::{c_char, c_float}; + +#[repr(C)] +#[derive(Debug, Copy, Clone)] +pub struct OnlineTransducerModelConfig { + pub encoder: *const c_char, + pub decoder: *const c_char, + pub joiner: *const c_char, +} + +#[repr(C)] +#[derive(Debug, Copy, Clone)] +pub struct OnlineParaformerModelConfig { + pub encoder: *const c_char, + pub decoder: *const c_char, +} + +#[repr(C)] +#[derive(Debug, Copy, Clone)] +pub struct OnlineZipformer2CtcModelConfig { + pub model: *const c_char, +} + +#[repr(C)] +#[derive(Debug, Copy, Clone)] +pub struct OnlineNemoCtcModelConfig { + pub model: *const c_char, +} + +#[repr(C)] +#[derive(Debug, Copy, Clone)] +pub struct OnlineToneCtcModelConfig { + pub model: *const c_char, +} + +#[repr(C)] +#[derive(Debug, Copy, Clone)] +pub struct OnlineModelConfig { + pub transducer: OnlineTransducerModelConfig, + pub paraformer: OnlineParaformerModelConfig, + pub zipformer2_ctc: OnlineZipformer2CtcModelConfig, + + pub tokens: *const c_char, + pub num_threads: i32, + pub provider: *const c_char, + pub debug: i32, + + pub model_type: *const c_char, + + // cjkchar | bpe | cjkchar+bpe + pub modeling_unit: *const c_char, + + pub bpe_vocab: *const c_char, + + pub tokens_buf: *const u8, + pub tokens_buf_size: i32, + + pub nemo_ctc: OnlineNemoCtcModelConfig, + pub t_one_ctc: OnlineToneCtcModelConfig, +} + +#[repr(C)] +#[derive(Debug, Copy, Clone)] +pub struct FeatureConfig { + pub sample_rate: i32, + pub feature_dim: i32, +} + +#[repr(C)] +#[derive(Debug, Copy, Clone)] +pub struct OnlineCtcFstDecoderConfig { + pub graph: *const c_char, + pub max_active: i32, +} + +#[repr(C)] +#[derive(Debug, Copy, Clone)] +pub struct HomophoneReplacerConfig { + pub dict_dir: *const c_char, + pub lexicon: *const c_char, + pub rule_fsts: *const c_char, +} + +#[repr(C)] +#[derive(Debug, Copy, Clone)] +pub struct OnlineRecognizerConfig { + pub feat_config: FeatureConfig, + pub model_config: OnlineModelConfig, + + // greedy_search | modified_beam_search + pub decoding_method: *const c_char, + + pub max_active_paths: i32, + + pub enable_endpoint: i32, + + pub rule1_min_trailing_silence: c_float, + pub rule2_min_trailing_silence: c_float, + pub rule3_min_utterance_length: c_float, + + pub hotwords_file: *const c_char, + pub hotwords_score: c_float, + + pub ctc_fst_decoder_config: OnlineCtcFstDecoderConfig, + + pub rule_fsts: *const c_char, + pub rule_fars: *const c_char, + + pub blank_penalty: c_float, + + pub hotwords_buf: *const u8, + pub hotwords_buf_size: i32, + + pub hr: HomophoneReplacerConfig, +} + +#[repr(C)] +pub struct OnlineRecognizer { + _private: [u8; 0], +} + +#[repr(C)] +pub struct OnlineStream { + _private: [u8; 0], +} + +extern "C" { + pub fn SherpaOnnxCreateOnlineRecognizer( + config: *const OnlineRecognizerConfig, + ) -> *const OnlineRecognizer; + + pub fn SherpaOnnxDestroyOnlineRecognizer(recognizer: *const OnlineRecognizer); + + pub fn SherpaOnnxCreateOnlineStream(recognizer: *const OnlineRecognizer) + -> *const OnlineStream; + + pub fn SherpaOnnxCreateOnlineStreamWithHotwords( + recognizer: *const OnlineRecognizer, + hotwords: *const c_char, + ) -> *const OnlineStream; + + pub fn SherpaOnnxDestroyOnlineStream(stream: *const OnlineStream); + + pub fn SherpaOnnxOnlineStreamAcceptWaveform( + stream: *const OnlineStream, + sample_rate: i32, + samples: *const f32, + n: i32, + ); + + pub fn SherpaOnnxIsOnlineStreamReady( + recognizer: *const OnlineRecognizer, + stream: *const OnlineStream, + ) -> i32; + + pub fn SherpaOnnxDecodeOnlineStream( + recognizer: *const OnlineRecognizer, + stream: *const OnlineStream, + ); + + pub fn SherpaOnnxDecodeMultipleOnlineStreams( + recognizer: *const OnlineRecognizer, + streams: *const *const OnlineStream, + n: i32, + ); + + pub fn SherpaOnnxGetOnlineStreamResultAsJson( + recognizer: *const OnlineRecognizer, + stream: *const OnlineStream, + ) -> *const c_char; + + pub fn SherpaOnnxDestroyOnlineStreamResultJson(s: *const c_char); + + pub fn SherpaOnnxOnlineStreamReset( + recognizer: *const OnlineRecognizer, + stream: *const OnlineStream, + ); + + pub fn SherpaOnnxOnlineStreamInputFinished(stream: *const OnlineStream); + + pub fn SherpaOnnxOnlineStreamIsEndpoint( + recognizer: *const OnlineRecognizer, + stream: *const OnlineStream, + ) -> i32; +} diff --git a/sherpa-onnx/rust/sherpa-onnx-sys/src/wave.rs b/sherpa-onnx/rust/sherpa-onnx-sys/src/wave.rs new file mode 100644 index 0000000000..7355d5381d --- /dev/null +++ b/sherpa-onnx/rust/sherpa-onnx-sys/src/wave.rs @@ -0,0 +1,22 @@ +#![allow(non_camel_case_types)] +#![allow(non_snake_case)] +#![allow(non_upper_case_globals)] + +use std::os::raw::{c_char, c_int}; + +#[repr(C)] +#[derive(Debug, Copy, Clone)] +pub struct SherpaOnnxWave { + /// Samples normalized to [-1, 1] + pub samples: *const f32, + pub sample_rate: c_int, + pub num_samples: c_int, +} + +extern "C" { + /// Read a WAV file. Returns NULL on error. + pub fn SherpaOnnxReadWave(filename: *const c_char) -> *const SherpaOnnxWave; + + /// Free memory allocated by SherpaOnnxReadWave + pub fn SherpaOnnxFreeWave(wave: *const SherpaOnnxWave); +} diff --git a/sherpa-onnx/rust/sherpa-onnx/Cargo.toml b/sherpa-onnx/rust/sherpa-onnx/Cargo.toml index c008e3d95e..a7ed81bbe1 100644 --- a/sherpa-onnx/rust/sherpa-onnx/Cargo.toml +++ b/sherpa-onnx/rust/sherpa-onnx/Cargo.toml @@ -1,6 +1,6 @@ [package] name = "sherpa-onnx" -version = "0.1.0" +version = "0.1.1" edition = "2021" description = "Safe Rust wrapper for sherpa-onnx speech recognition toolkit" license = "Apache-2.0" @@ -20,4 +20,6 @@ include = [ ] [dependencies] -sherpa-onnx-sys = { path = "../sherpa-onnx-sys", version = "0.1.0" } +sherpa-onnx-sys = { path = "../sherpa-onnx-sys", version = "0.1.1" } +serde = { version = "1.0", features = ["derive"] } +serde_json = "1.0" diff --git a/sherpa-onnx/rust/sherpa-onnx/src/lib.rs b/sherpa-onnx/rust/sherpa-onnx/src/lib.rs index 4b074c395c..957f08d7de 100644 --- a/sherpa-onnx/rust/sherpa-onnx/src/lib.rs +++ b/sherpa-onnx/rust/sherpa-onnx/src/lib.rs @@ -1,3 +1,7 @@ +mod online_asr; mod utils; +mod wave; +pub use online_asr::*; pub use utils::*; +pub use wave::*; diff --git a/sherpa-onnx/rust/sherpa-onnx/src/online_asr.rs b/sherpa-onnx/rust/sherpa-onnx/src/online_asr.rs new file mode 100644 index 0000000000..3fcf659db4 --- /dev/null +++ b/sherpa-onnx/rust/sherpa-onnx/src/online_asr.rs @@ -0,0 +1,470 @@ +use serde::Deserialize; +use std::ffi::{CStr, CString}; +use std::os::raw::c_char; +use std::ptr; + +use sherpa_onnx_sys as sys; + +#[derive(Clone, Debug)] +pub struct OnlineTransducerModelConfig { + pub encoder: Option, + pub decoder: Option, + pub joiner: Option, +} + +impl Default for OnlineTransducerModelConfig { + fn default() -> Self { + Self { + encoder: None, + decoder: None, + joiner: None, + } + } +} + +impl OnlineTransducerModelConfig { + fn to_sys(&self, cstrings: &mut Vec) -> sys::OnlineTransducerModelConfig { + sys::OnlineTransducerModelConfig { + encoder: to_c_ptr(&self.encoder, cstrings), + decoder: to_c_ptr(&self.decoder, cstrings), + joiner: to_c_ptr(&self.joiner, cstrings), + } + } +} + +#[derive(Clone, Debug)] +pub struct OnlineParaformerModelConfig { + pub encoder: Option, + pub decoder: Option, +} + +impl Default for OnlineParaformerModelConfig { + fn default() -> Self { + Self { + encoder: None, + decoder: None, + } + } +} + +impl OnlineParaformerModelConfig { + fn to_sys(&self, cstrings: &mut Vec) -> sys::OnlineParaformerModelConfig { + sys::OnlineParaformerModelConfig { + encoder: to_c_ptr(&self.encoder, cstrings), + decoder: to_c_ptr(&self.decoder, cstrings), + } + } +} + +#[derive(Clone, Debug)] +pub struct OnlineZipformer2CtcModelConfig { + pub model: Option, +} + +impl Default for OnlineZipformer2CtcModelConfig { + fn default() -> Self { + Self { model: None } + } +} + +impl OnlineZipformer2CtcModelConfig { + fn to_sys(&self, cstrings: &mut Vec) -> sys::OnlineZipformer2CtcModelConfig { + sys::OnlineZipformer2CtcModelConfig { + model: to_c_ptr(&self.model, cstrings), + } + } +} + +#[derive(Clone, Debug)] +pub struct OnlineNemoCtcModelConfig { + pub model: Option, +} + +impl Default for OnlineNemoCtcModelConfig { + fn default() -> Self { + Self { model: None } + } +} + +impl OnlineNemoCtcModelConfig { + fn to_sys(&self, cstrings: &mut Vec) -> sys::OnlineNemoCtcModelConfig { + sys::OnlineNemoCtcModelConfig { + model: to_c_ptr(&self.model, cstrings), + } + } +} + +#[derive(Clone, Debug)] +pub struct OnlineToneCtcModelConfig { + pub model: Option, +} + +impl Default for OnlineToneCtcModelConfig { + fn default() -> Self { + Self { model: None } + } +} + +impl OnlineToneCtcModelConfig { + fn to_sys(&self, cstrings: &mut Vec) -> sys::OnlineToneCtcModelConfig { + sys::OnlineToneCtcModelConfig { + model: to_c_ptr(&self.model, cstrings), + } + } +} + +#[derive(Clone, Debug)] +pub struct OnlineModelConfig { + pub transducer: OnlineTransducerModelConfig, + pub paraformer: OnlineParaformerModelConfig, + pub zipformer2_ctc: OnlineZipformer2CtcModelConfig, + pub nemo_ctc: OnlineNemoCtcModelConfig, + pub t_one_ctc: OnlineToneCtcModelConfig, + + pub tokens: Option, + pub num_threads: i32, + pub provider: Option, + pub debug: bool, + + pub model_type: Option, + pub modeling_unit: Option, // cjkchar | bpe | cjkchar+bpe + pub bpe_vocab: Option, + + /// Optional in-memory tokens + pub tokens_buf: Option>, +} + +impl Default for OnlineModelConfig { + fn default() -> Self { + Self { + transducer: Default::default(), + paraformer: Default::default(), + zipformer2_ctc: Default::default(), + nemo_ctc: Default::default(), + t_one_ctc: Default::default(), + + tokens: None, + num_threads: 1, + provider: Some("cpu".to_string()), + debug: false, + + model_type: None, + modeling_unit: None, + bpe_vocab: None, + tokens_buf: None, + } + } +} + +impl OnlineModelConfig { + fn to_sys(&self, cstrings: &mut Vec) -> sys::OnlineModelConfig { + sys::OnlineModelConfig { + transducer: self + .transducer + .to_sys(cstrings), + paraformer: self + .paraformer + .to_sys(cstrings), + zipformer2_ctc: self + .zipformer2_ctc + .to_sys(cstrings), + nemo_ctc: self + .nemo_ctc + .to_sys(cstrings), + t_one_ctc: self + .t_one_ctc + .to_sys(cstrings), + + tokens: to_c_ptr(&self.tokens, cstrings), + num_threads: self.num_threads, + provider: to_c_ptr(&self.provider, cstrings), + debug: self.debug as i32, + + model_type: to_c_ptr(&self.model_type, cstrings), + modeling_unit: to_c_ptr(&self.modeling_unit, cstrings), + bpe_vocab: to_c_ptr(&self.bpe_vocab, cstrings), + + tokens_buf: self + .tokens_buf + .as_ref() + .map_or(ptr::null(), |buf| buf.as_ptr() as *const _), + tokens_buf_size: self + .tokens_buf + .as_ref() + .map_or(0, |buf| buf.len() as i32), + } + } +} + +#[derive(Clone, Debug)] +pub struct OnlineCtcFstDecoderConfig { + pub graph: Option, + pub max_active: i32, +} + +impl Default for OnlineCtcFstDecoderConfig { + fn default() -> Self { + Self { + graph: None, + max_active: 0, + } + } +} + +impl OnlineCtcFstDecoderConfig { + /// Convert to sys struct using `to_c_ptr()` + pub(crate) fn to_sys(&self, cstrings: &mut Vec) -> sys::OnlineCtcFstDecoderConfig { + sys::OnlineCtcFstDecoderConfig { + graph: to_c_ptr(&self.graph, cstrings), + max_active: self.max_active, + } + } +} + +#[derive(Clone, Debug)] +pub struct HomophoneReplacerConfig { + pub lexicon: Option, + pub rule_fsts: Option, +} + +impl Default for HomophoneReplacerConfig { + fn default() -> Self { + Self { + lexicon: None, + rule_fsts: None, + } + } +} + +impl HomophoneReplacerConfig { + pub(crate) fn to_sys(&self, cstrings: &mut Vec) -> sys::HomophoneReplacerConfig { + sys::HomophoneReplacerConfig { + dict_dir: ptr::null(), // not used any more internally + lexicon: to_c_ptr(&self.lexicon, cstrings), + rule_fsts: to_c_ptr(&self.rule_fsts, cstrings), + } + } +} + +#[derive(Clone, Debug)] +pub struct OnlineRecognizerConfig { + pub feat_config: sys::FeatureConfig, + pub model_config: OnlineModelConfig, + + /// Decoding method: greedy_search | modified_beam_search + pub decoding_method: Option, + + /// Used only when decoding_method is modified_beam_search + pub max_active_paths: i32, + + /// Endpoint detection + pub enable_endpoint: bool, + + pub rule1_min_trailing_silence: f32, + pub rule2_min_trailing_silence: f32, + pub rule3_min_utterance_length: f32, + + pub hotwords_file: Option, + pub hotwords_score: f32, + + pub ctc_fst_decoder_config: OnlineCtcFstDecoderConfig, + + pub rule_fsts: Option, + pub rule_fars: Option, + + pub blank_penalty: f32, + + pub hotwords_buf: Option>, + + pub hr: HomophoneReplacerConfig, +} + +impl Default for OnlineRecognizerConfig { + fn default() -> Self { + Self { + feat_config: sys::FeatureConfig { + sample_rate: 16000, + feature_dim: 80, + }, + model_config: Default::default(), + decoding_method: None, + max_active_paths: 0, + enable_endpoint: false, + rule1_min_trailing_silence: 0.0, + rule2_min_trailing_silence: 0.0, + rule3_min_utterance_length: 0.0, + hotwords_file: None, + hotwords_score: 0.0, + ctc_fst_decoder_config: Default::default(), + rule_fsts: None, + rule_fars: None, + blank_penalty: 0.0, + hotwords_buf: None, + hr: Default::default(), + } + } +} + +impl OnlineRecognizerConfig { + /// Convert to sys struct for FFI call + pub(crate) fn to_sys(&self, cstrings: &mut Vec) -> sys::OnlineRecognizerConfig { + sys::OnlineRecognizerConfig { + feat_config: self.feat_config, + model_config: self + .model_config + .to_sys(cstrings), + decoding_method: to_c_ptr(&self.decoding_method, cstrings), + max_active_paths: self.max_active_paths, + enable_endpoint: self.enable_endpoint as i32, + rule1_min_trailing_silence: self.rule1_min_trailing_silence, + rule2_min_trailing_silence: self.rule2_min_trailing_silence, + rule3_min_utterance_length: self.rule3_min_utterance_length, + hotwords_file: to_c_ptr(&self.hotwords_file, cstrings), + hotwords_score: self.hotwords_score, + ctc_fst_decoder_config: self + .ctc_fst_decoder_config + .to_sys(cstrings), + rule_fsts: to_c_ptr(&self.rule_fsts, cstrings), + rule_fars: to_c_ptr(&self.rule_fars, cstrings), + blank_penalty: self.blank_penalty, + hotwords_buf: self + .hotwords_buf + .as_ref() + .map_or(ptr::null(), |buf| buf.as_ptr() as *const _), + hotwords_buf_size: self + .hotwords_buf + .as_ref() + .map_or(0, |buf| buf.len() as i32), + hr: self + .hr + .to_sys(cstrings), + } + } +} + +pub struct OnlineRecognizer { + ptr: *const sys::OnlineRecognizer, +} + +impl OnlineRecognizer { + pub fn create(config: &OnlineRecognizerConfig) -> Option { + let mut cstrings = Vec::new(); + + let sys_config = config.to_sys(&mut cstrings); + + let ptr = unsafe { sys::SherpaOnnxCreateOnlineRecognizer(&sys_config) }; + + if ptr.is_null() { + None + } else { + Some(Self { ptr }) + } + } + + pub fn create_stream(&self) -> OnlineStream { + let ptr = unsafe { sys::SherpaOnnxCreateOnlineStream(self.ptr) }; + OnlineStream { ptr } + } + + pub fn create_stream_with_hotwords(&self, hotwords: &str) -> OnlineStream { + let c = CString::new(hotwords).unwrap(); + let ptr = unsafe { sys::SherpaOnnxCreateOnlineStreamWithHotwords(self.ptr, c.as_ptr()) }; + OnlineStream { ptr } + } + + pub fn decode(&self, stream: &OnlineStream) { + unsafe { sys::SherpaOnnxDecodeOnlineStream(self.ptr, stream.ptr) } + } + + pub fn decode_multiple_streams(&self, streams: &[&OnlineStream]) { + let ptrs: Vec<*const sys::OnlineStream> = streams + .iter() + .map(|s| s.ptr) + .collect(); + unsafe { + sys::SherpaOnnxDecodeMultipleOnlineStreams(self.ptr, ptrs.as_ptr(), ptrs.len() as i32) + } + } + + pub fn reset(&self, stream: &OnlineStream) { + unsafe { sys::SherpaOnnxOnlineStreamReset(self.ptr, stream.ptr) } + } + + pub fn is_endpoint(&self, stream: &OnlineStream) -> bool { + unsafe { sys::SherpaOnnxOnlineStreamIsEndpoint(self.ptr, stream.ptr) != 0 } + } + + pub fn is_ready(&self, stream: &OnlineStream) -> bool { + unsafe { sys::SherpaOnnxIsOnlineStreamReady(self.ptr, stream.ptr) != 0 } + } + + pub fn get_result(&self, stream: &OnlineStream) -> Option { + unsafe { + let cstr = sys::SherpaOnnxGetOnlineStreamResultAsJson(self.ptr, stream.ptr); + if cstr.is_null() { + return None; + } + let s = CStr::from_ptr(cstr) + .to_string_lossy() + .into_owned(); + sys::SherpaOnnxDestroyOnlineStreamResultJson(cstr); + serde_json::from_str(&s).ok() + } + } +} + +#[derive(Clone, Debug, Deserialize)] +pub struct RecognizerResult { + pub text: String, + pub tokens: Vec, + pub timestamps: Option>, + pub segment: Option, + pub start_time: Option, + pub is_final: bool, +} + +impl Drop for OnlineRecognizer { + fn drop(&mut self) { + unsafe { + sys::SherpaOnnxDestroyOnlineRecognizer(self.ptr); + } + } +} + +pub struct OnlineStream { + ptr: *const sys::OnlineStream, +} + +impl OnlineStream { + pub fn accept_waveform(&self, sample_rate: i32, samples: &[f32]) { + unsafe { + sys::SherpaOnnxOnlineStreamAcceptWaveform( + self.ptr, + sample_rate, + samples.as_ptr(), + samples.len() as i32, + ) + } + } + + pub fn input_finished(&self) { + unsafe { sys::SherpaOnnxOnlineStreamInputFinished(self.ptr) } + } +} + +impl Drop for OnlineStream { + fn drop(&mut self) { + unsafe { sys::SherpaOnnxDestroyOnlineStream(self.ptr) } + } +} + +fn to_c_ptr(opt: &Option, storage: &mut Vec) -> *const c_char { + if let Some(s) = opt { + let c = CString::new(s.as_str()).unwrap(); + let ptr = c.as_ptr(); + storage.push(c); + ptr + } else { + ptr::null() + } +} diff --git a/sherpa-onnx/rust/sherpa-onnx/src/utils.rs b/sherpa-onnx/rust/sherpa-onnx/src/utils.rs index 75e95fcf4b..f6f28ca312 100644 --- a/sherpa-onnx/rust/sherpa-onnx/src/utils.rs +++ b/sherpa-onnx/rust/sherpa-onnx/src/utils.rs @@ -8,20 +8,12 @@ use std::os::raw::c_char; /// If the C string is not valid UTF-8, a lossy UTF-8 conversion is used /// and the resulting string is leaked to obtain a `'static` lifetime. fn c_str_to_static_str(ptr: *const c_char) -> &'static str { - if ptr.is_null() { - return ""; - } + assert!(!ptr.is_null(), "C string pointer is null"); + unsafe { - match CStr::from_ptr(ptr).to_str() { - Ok(s) => s, - Err(_) => { - // Fallback to a lossy conversion and leak the String to get a `'static` str. - let owned = CStr::from_ptr(ptr) - .to_string_lossy() - .into_owned(); - Box::leak(owned.into_boxed_str()) - } - } + CStr::from_ptr(ptr) + .to_str() + .unwrap() } } diff --git a/sherpa-onnx/rust/sherpa-onnx/src/wave.rs b/sherpa-onnx/rust/sherpa-onnx/src/wave.rs new file mode 100644 index 0000000000..66937c6132 --- /dev/null +++ b/sherpa-onnx/rust/sherpa-onnx/src/wave.rs @@ -0,0 +1,59 @@ +use std::ffi::CString; +use std::slice; + +use sherpa_onnx_sys as sys; + +#[derive(Debug)] +pub struct Wave { + inner: *const sys::SherpaOnnxWave, +} + +impl Wave { + /// Read a WAV file using SherpaOnnx C API. + pub fn read(filename: &str) -> Option { + let c_filename = CString::new(filename).unwrap(); + let wave_ptr = unsafe { sys::SherpaOnnxReadWave(c_filename.as_ptr()) }; + if wave_ptr.is_null() { + None + } else { + Some(Self { inner: wave_ptr }) + } + } + + /// Get sample rate + pub fn sample_rate(&self) -> i32 { + unsafe { (*self.inner).sample_rate } + } + + /// Get number of samples + pub fn num_samples(&self) -> i32 { + unsafe { (*self.inner).num_samples } + } + + /// Get a slice of normalized samples + pub fn samples(&self) -> &[f32] { + unsafe { + let ptr = (*self.inner).samples; + let len = (*self.inner).num_samples as usize; + + if ptr.is_null() || len == 0 { + &[] + } else { + slice::from_raw_parts(ptr, len) + } + } + } +} + +impl Drop for Wave { + fn drop(&mut self) { + unsafe { + if !self + .inner + .is_null() + { + sys::SherpaOnnxFreeWave(self.inner); + } + } + } +}