react-native-nitro-onnx 0.1.0

This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
Files changed (91) hide show
  1. package/NOTICE +241 -0
  2. package/NitroOnnxSpeech.podspec +51 -0
  3. package/Package.swift +52 -0
  4. package/README.md +398 -0
  5. package/android/CMakeLists.txt +117 -0
  6. package/android/build.gradle +65 -0
  7. package/android/src/main/AndroidManifest.xml +6 -0
  8. package/android/src/main/assets/silero_vad.onnx +0 -0
  9. package/android/src/main/cpp/cpp-adapter.cpp +24 -0
  10. package/android/src/main/java/com/margelo/nitro/onnx/speech/OnnxSpeechPackage.kt +67 -0
  11. package/assets/silero_vad.onnx +0 -0
  12. package/cpp/AndroidPthreadCompat.cpp +16 -0
  13. package/cpp/AsrEngine.cpp +277 -0
  14. package/cpp/AsrEngine.hpp +110 -0
  15. package/cpp/AudioFileReader.cpp +177 -0
  16. package/cpp/AudioFileReader.hpp +26 -0
  17. package/cpp/AudioUtils.cpp +28 -0
  18. package/cpp/AudioUtils.hpp +38 -0
  19. package/cpp/ModelSingleton.hpp +51 -0
  20. package/cpp/NitroOnnxSpeech.cpp +68 -0
  21. package/cpp/NitroOnnxSpeech.hpp +42 -0
  22. package/cpp/OfflineAsr.cpp +81 -0
  23. package/cpp/OfflineAsr.hpp +33 -0
  24. package/cpp/ResourceDir.cpp +29 -0
  25. package/cpp/ResourceDir.hpp +18 -0
  26. package/cpp/SpeakerEngine.cpp +170 -0
  27. package/cpp/SpeakerEngine.hpp +70 -0
  28. package/cpp/SpeakerManager.cpp +93 -0
  29. package/cpp/SpeakerManager.hpp +43 -0
  30. package/cpp/StreamingAsr.cpp +123 -0
  31. package/cpp/StreamingAsr.hpp +55 -0
  32. package/cpp/ThreadPool.cpp +41 -0
  33. package/cpp/ThreadPool.hpp +62 -0
  34. package/cpp/Tts.cpp +132 -0
  35. package/cpp/Tts.hpp +38 -0
  36. package/cpp/TtsEngine.cpp +222 -0
  37. package/cpp/TtsEngine.hpp +77 -0
  38. package/cpp/Vad.cpp +126 -0
  39. package/cpp/Vad.hpp +62 -0
  40. package/cpp/VadEngine.cpp +217 -0
  41. package/cpp/VadEngine.hpp +126 -0
  42. package/ios/OnnxSpeechInitializer.mm +37 -0
  43. package/ios/PrivacyInfo.xcprivacy +14 -0
  44. package/lib/index.d.ts +11 -0
  45. package/lib/index.d.ts.map +1 -0
  46. package/lib/index.js +19 -0
  47. package/lib/index.js.map +1 -0
  48. package/lib/specs/OnnxSpeech.nitro.d.ts +264 -0
  49. package/lib/specs/OnnxSpeech.nitro.d.ts.map +1 -0
  50. package/lib/specs/OnnxSpeech.nitro.js +6 -0
  51. package/lib/specs/OnnxSpeech.nitro.js.map +1 -0
  52. package/nitro.json +19 -0
  53. package/nitrogen/generated/.gitattributes +1 -0
  54. package/nitrogen/generated/android/NitroOnnxSpeech+autolinking.cmake +86 -0
  55. package/nitrogen/generated/android/NitroOnnxSpeech+autolinking.gradle +27 -0
  56. package/nitrogen/generated/android/NitroOnnxSpeechOnLoad.cpp +49 -0
  57. package/nitrogen/generated/android/NitroOnnxSpeechOnLoad.hpp +34 -0
  58. package/nitrogen/generated/android/kotlin/com/margelo/nitro/onnx/speech/NitroOnnxSpeechOnLoad.kt +35 -0
  59. package/nitrogen/generated/ios/NitroOnnxSpeech+autolinking.rb +62 -0
  60. package/nitrogen/generated/ios/NitroOnnxSpeech-Swift-Cxx-Bridge.cpp +17 -0
  61. package/nitrogen/generated/ios/NitroOnnxSpeech-Swift-Cxx-Bridge.hpp +27 -0
  62. package/nitrogen/generated/ios/NitroOnnxSpeech-Swift-Cxx-Umbrella.hpp +38 -0
  63. package/nitrogen/generated/ios/NitroOnnxSpeechAutolinking.mm +35 -0
  64. package/nitrogen/generated/ios/NitroOnnxSpeechAutolinking.swift +16 -0
  65. package/nitrogen/generated/shared/c++/AsrModelConfig.hpp +142 -0
  66. package/nitrogen/generated/shared/c++/AsrModelType.hpp +112 -0
  67. package/nitrogen/generated/shared/c++/AsrResult.hpp +105 -0
  68. package/nitrogen/generated/shared/c++/HybridOfflineAsrSpec.cpp +25 -0
  69. package/nitrogen/generated/shared/c++/HybridOfflineAsrSpec.hpp +73 -0
  70. package/nitrogen/generated/shared/c++/HybridOnnxSpeechSpec.cpp +27 -0
  71. package/nitrogen/generated/shared/c++/HybridOnnxSpeechSpec.hpp +82 -0
  72. package/nitrogen/generated/shared/c++/HybridSpeakerManagerSpec.cpp +28 -0
  73. package/nitrogen/generated/shared/c++/HybridSpeakerManagerSpec.hpp +77 -0
  74. package/nitrogen/generated/shared/c++/HybridStreamingAsrSpec.cpp +32 -0
  75. package/nitrogen/generated/shared/c++/HybridStreamingAsrSpec.hpp +81 -0
  76. package/nitrogen/generated/shared/c++/HybridTtsSpec.cpp +26 -0
  77. package/nitrogen/generated/shared/c++/HybridTtsSpec.hpp +74 -0
  78. package/nitrogen/generated/shared/c++/HybridVadSpec.cpp +31 -0
  79. package/nitrogen/generated/shared/c++/HybridVadSpec.hpp +81 -0
  80. package/nitrogen/generated/shared/c++/RegisteredSpeaker.hpp +91 -0
  81. package/nitrogen/generated/shared/c++/SpeakerEmbeddingConfig.hpp +91 -0
  82. package/nitrogen/generated/shared/c++/TtsModelConfig.hpp +178 -0
  83. package/nitrogen/generated/shared/c++/TtsModelType.hpp +88 -0
  84. package/nitrogen/generated/shared/c++/TtsResult.hpp +91 -0
  85. package/nitrogen/generated/shared/c++/VadConfig.hpp +100 -0
  86. package/nitrogen/generated/shared/c++/VadSegment.hpp +91 -0
  87. package/package.json +62 -0
  88. package/scripts/prepare-sherpa-onnx.js +236 -0
  89. package/scripts/test-cpp.js +24 -0
  90. package/src/index.ts +45 -0
  91. package/src/specs/OnnxSpeech.nitro.ts +326 -0
@@ -0,0 +1,16 @@
1
+ // ------------------------------------------------------------------------------
2
+ // AndroidPthreadCompat.cpp
3
+ // Android Bionic omits pthread_cancel(). sherpa-onnx's prebuilt .so references
4
+ // it, so we provide a no-op stub that returns ESRCH (no such thread).
5
+ // ------------------------------------------------------------------------------
6
+
7
+ #ifdef __ANDROID__
8
+
9
+ #include <cerrno>
10
+ #include <pthread.h>
11
+
12
+ extern "C" int pthread_cancel(pthread_t) {
13
+ return ESRCH;
14
+ }
15
+
16
+ #endif
@@ -0,0 +1,277 @@
1
+ // ------------------------------------------------------------------------------
2
+ // AsrEngine.cpp
3
+ // ------------------------------------------------------------------------------
4
+ #include "AsrEngine.hpp"
5
+
6
+ #include "AudioFileReader.hpp"
7
+ #include "sherpa-onnx/c-api/c-api.h"
8
+
9
+ #include <cstring>
10
+ #include <fstream>
11
+ #include <stdexcept>
12
+
13
+ namespace margelo::nitro::onnx::speech {
14
+
15
+ namespace {
16
+
17
+ std::string joinPath(const std::string& dir, const std::string& file) {
18
+ if (file.empty()) {
19
+ return dir;
20
+ }
21
+ if (file[0] == '/') {
22
+ return file;
23
+ }
24
+ if (dir.empty()) {
25
+ return file;
26
+ }
27
+ if (dir.back() == '/') {
28
+ return dir + file;
29
+ }
30
+ return dir + "/" + file;
31
+ }
32
+
33
+ ModelSingleton<const SherpaOnnxOfflineRecognizer> gOfflineRecognizerCache;
34
+ ModelSingleton<const SherpaOnnxOnlineRecognizer> gOnlineRecognizerCache;
35
+
36
+ } // namespace
37
+
38
+ // ------------------------------------------------------------------------------
39
+ // Offline ASR
40
+ // ------------------------------------------------------------------------------
41
+
42
+ OfflineAsrEngine::OfflineAsrEngine(std::shared_ptr<ThreadPool> threadPool)
43
+ : threadPool_(std::move(threadPool)) {}
44
+
45
+ OfflineAsrEngine::~OfflineAsrEngine() {
46
+ unload();
47
+ }
48
+
49
+ void OfflineAsrEngine::load(const AsrEngineConfig& config) {
50
+ unload();
51
+ config_ = config;
52
+
53
+ const std::string key = config_.modelDir + "|" + std::to_string(static_cast<int>(config_.type));
54
+ auto cached = gOfflineRecognizerCache.getOrCreate(key, [this](const std::string&) {
55
+ SherpaOnnxOfflineRecognizerConfig c;
56
+ std::memset(&c, 0, sizeof(c));
57
+
58
+ // All joined paths must outlive the call to CreateOfflineRecognizer.
59
+ std::string tokens = joinPath(config_.modelDir, config_.tokensPath);
60
+ std::string whisperEncoder = joinPath(config_.modelDir, config_.whisperEncoder);
61
+ std::string whisperDecoder = joinPath(config_.modelDir, config_.whisperDecoder);
62
+ std::string encoder = joinPath(config_.modelDir, config_.encoder);
63
+ std::string decoder = joinPath(config_.modelDir, config_.decoder);
64
+ std::string joiner = joinPath(config_.modelDir, config_.joiner);
65
+ std::string model = joinPath(config_.modelDir, config_.model);
66
+
67
+ switch (config_.type) {
68
+ case AsrModelType::WHISPER:
69
+ c.model_config.whisper.encoder = whisperEncoder.c_str();
70
+ c.model_config.whisper.decoder = whisperDecoder.c_str();
71
+ c.model_config.whisper.language = config_.language.c_str();
72
+ c.model_config.whisper.tail_paddings = 2;
73
+ break;
74
+ case AsrModelType::TRANSDUCER:
75
+ case AsrModelType::ZIPFORMER:
76
+ case AsrModelType::CONFORMER:
77
+ c.model_config.transducer.encoder = encoder.c_str();
78
+ c.model_config.transducer.decoder = decoder.c_str();
79
+ c.model_config.transducer.joiner = joiner.c_str();
80
+ break;
81
+ case AsrModelType::PARAFORMER:
82
+ case AsrModelType::WENET:
83
+ case AsrModelType::TELESPEECH:
84
+ case AsrModelType::SENSE_VOICE:
85
+ c.model_config.paraformer.model = model.c_str();
86
+ break;
87
+ case AsrModelType::MOONSHINE:
88
+ case AsrModelType::DOLPHIN:
89
+ case AsrModelType::NEMO:
90
+ c.model_config.nemo_ctc.model = model.c_str();
91
+ if (config_.type == AsrModelType::NEMO) {
92
+ c.model_config.model_type = "nemo";
93
+ }
94
+ break;
95
+ }
96
+
97
+ c.model_config.tokens = tokens.c_str();
98
+ c.model_config.num_threads = config_.numThreads;
99
+ c.model_config.debug = 0;
100
+ c.decoding_method = config_.decodingMethod.c_str();
101
+ c.max_active_paths = config_.maxActivePaths;
102
+
103
+ const SherpaOnnxOfflineRecognizer* rec = SherpaOnnxCreateOfflineRecognizer(&c);
104
+ if (rec == nullptr) {
105
+ throw std::runtime_error("Failed to create offline ASR recognizer");
106
+ }
107
+ return std::shared_ptr<const SherpaOnnxOfflineRecognizer>(
108
+ rec, SherpaOnnxDestroyOfflineRecognizer);
109
+ });
110
+
111
+ recognizer_ = cached;
112
+ }
113
+
114
+ bool OfflineAsrEngine::isLoaded() const {
115
+ return recognizer_ != nullptr;
116
+ }
117
+
118
+ AsrEngineResult OfflineAsrEngine::recognize(const std::vector<float>& samples) {
119
+ if (recognizer_ == nullptr) {
120
+ throw std::runtime_error("Offline ASR not loaded");
121
+ }
122
+
123
+ const SherpaOnnxOfflineStream* stream = SherpaOnnxCreateOfflineStream(recognizer_.get());
124
+ SherpaOnnxAcceptWaveformOffline(stream, 16000, samples.data(), static_cast<int32_t>(samples.size()));
125
+ SherpaOnnxDecodeOfflineStream(recognizer_.get(), stream);
126
+
127
+ const char* json = SherpaOnnxGetOfflineStreamResultAsJson(stream);
128
+ AsrEngineResult result;
129
+ result.json = json ? json : "";
130
+
131
+ // Parse the JSON to extract text. In a full implementation, use a JSON
132
+ // library to also populate timestamps and score.
133
+ const char* textKey = "\"text\":\"";
134
+ const char* textStart = std::strstr(result.json.c_str(), textKey);
135
+ if (textStart != nullptr) {
136
+ textStart += std::strlen(textKey);
137
+ const char* textEnd = std::strstr(textStart, "\"");
138
+ if (textEnd != nullptr) {
139
+ result.text = std::string(textStart, textEnd);
140
+ }
141
+ }
142
+ result.endMs = samplesToMs(static_cast<int32_t>(samples.size()));
143
+
144
+ SherpaOnnxDestroyOfflineStream(stream);
145
+ return result;
146
+ }
147
+
148
+ AsrEngineResult OfflineAsrEngine::recognizeFile(const std::string& path) {
149
+ std::vector<float> samples;
150
+ if (path.size() >= 4 && path.substr(path.size() - 4) == ".wav") {
151
+ samples = readWavFile(path);
152
+ } else {
153
+ samples = readRawPcmFile(path);
154
+ }
155
+ return recognize(samples);
156
+ }
157
+
158
+ void OfflineAsrEngine::unload() {
159
+ recognizer_.reset();
160
+ }
161
+
162
+ // ------------------------------------------------------------------------------
163
+ // Streaming ASR
164
+ // ------------------------------------------------------------------------------
165
+
166
+ StreamingAsrEngine::StreamingAsrEngine(std::shared_ptr<ThreadPool> threadPool)
167
+ : threadPool_(std::move(threadPool)) {}
168
+
169
+ StreamingAsrEngine::~StreamingAsrEngine() {
170
+ unload();
171
+ }
172
+
173
+ void StreamingAsrEngine::load(
174
+ const AsrEngineConfig& config,
175
+ std::shared_ptr<StreamingAsrListener> listener) {
176
+ unload();
177
+ config_ = config;
178
+ listener_ = std::move(listener);
179
+
180
+ const std::string key = config_.modelDir + "|streaming|" + std::to_string(static_cast<int>(config_.type));
181
+ auto cached = gOnlineRecognizerCache.getOrCreate(key, [this](const std::string&) {
182
+ SherpaOnnxOnlineRecognizerConfig c;
183
+ std::memset(&c, 0, sizeof(c));
184
+
185
+ std::string tokens = joinPath(config_.modelDir, config_.tokensPath);
186
+ std::string encoder = joinPath(config_.modelDir, config_.encoder);
187
+ std::string decoder = joinPath(config_.modelDir, config_.decoder);
188
+ std::string joiner = joinPath(config_.modelDir, config_.joiner);
189
+
190
+ switch (config_.type) {
191
+ case AsrModelType::TRANSDUCER:
192
+ case AsrModelType::ZIPFORMER:
193
+ case AsrModelType::CONFORMER:
194
+ c.model_config.transducer.encoder = encoder.c_str();
195
+ c.model_config.transducer.decoder = decoder.c_str();
196
+ c.model_config.transducer.joiner = joiner.c_str();
197
+ break;
198
+ default:
199
+ throw std::runtime_error("Streaming ASR does not support this model type in the scaffold");
200
+ }
201
+
202
+ c.model_config.tokens = tokens.c_str();
203
+ c.model_config.num_threads = config_.numThreads;
204
+ c.decoding_method = config_.decodingMethod.c_str();
205
+ c.max_active_paths = config_.maxActivePaths;
206
+ c.enable_endpoint = 1;
207
+
208
+ const SherpaOnnxOnlineRecognizer* rec = SherpaOnnxCreateOnlineRecognizer(&c);
209
+ if (rec == nullptr) {
210
+ throw std::runtime_error("Failed to create streaming ASR recognizer");
211
+ }
212
+ return std::shared_ptr<const SherpaOnnxOnlineRecognizer>(rec, SherpaOnnxDestroyOnlineRecognizer);
213
+ });
214
+
215
+ recognizer_ = cached;
216
+ stream_ = std::shared_ptr<const SherpaOnnxOnlineStream>(
217
+ SherpaOnnxCreateOnlineStream(recognizer_.get()),
218
+ SherpaOnnxDestroyOnlineStream);
219
+ }
220
+
221
+ bool StreamingAsrEngine::isLoaded() const {
222
+ return recognizer_ != nullptr && stream_ != nullptr;
223
+ }
224
+
225
+ void StreamingAsrEngine::acceptWaveform(const std::vector<float>& samples) {
226
+ if (recognizer_ == nullptr || stream_ == nullptr) {
227
+ throw std::runtime_error("Streaming ASR not loaded");
228
+ }
229
+ SherpaOnnxOnlineStreamAcceptWaveform(stream_.get(), 16000, samples.data(), static_cast<int32_t>(samples.size()));
230
+
231
+ if (listener_ && SherpaOnnxIsOnlineStreamReady(recognizer_.get(), stream_.get())) {
232
+ SherpaOnnxDecodeOnlineStream(recognizer_.get(), stream_.get());
233
+ const char* json = SherpaOnnxGetOnlineStreamResultAsJson(recognizer_.get(), stream_.get());
234
+ if (json != nullptr) {
235
+ AsrEngineResult result;
236
+ result.json = json;
237
+ // TODO: parse text from JSON.
238
+ listener_->onPartialResult(result);
239
+ SherpaOnnxDestroyOnlineStreamResultJson(json);
240
+ }
241
+ }
242
+ }
243
+
244
+ AsrEngineResult StreamingAsrEngine::finalize() {
245
+ if (recognizer_ == nullptr || stream_ == nullptr) {
246
+ throw std::runtime_error("Streaming ASR not loaded");
247
+ }
248
+ SherpaOnnxOnlineStreamInputFinished(stream_.get());
249
+ SherpaOnnxDecodeOnlineStream(recognizer_.get(), stream_.get());
250
+ const char* json = SherpaOnnxGetOnlineStreamResultAsJson(recognizer_.get(), stream_.get());
251
+ AsrEngineResult result;
252
+ if (json != nullptr) {
253
+ result.json = json;
254
+ // TODO: parse text from JSON.
255
+ SherpaOnnxDestroyOnlineStreamResultJson(json);
256
+ }
257
+ if (listener_) {
258
+ listener_->onFinalResult(result);
259
+ }
260
+ return result;
261
+ }
262
+
263
+ void StreamingAsrEngine::reset() {
264
+ if (recognizer_ != nullptr) {
265
+ stream_ = std::shared_ptr<const SherpaOnnxOnlineStream>(
266
+ SherpaOnnxCreateOnlineStream(recognizer_.get()),
267
+ SherpaOnnxDestroyOnlineStream);
268
+ }
269
+ }
270
+
271
+ void StreamingAsrEngine::unload() {
272
+ stream_.reset();
273
+ recognizer_.reset();
274
+ listener_.reset();
275
+ }
276
+
277
+ } // namespace margelo::nitro::onnx::speech
@@ -0,0 +1,110 @@
1
+ // ------------------------------------------------------------------------------
2
+ // AsrEngine.hpp
3
+ // Offline and streaming ASR wrappers backed by sherpa-onnx.
4
+ // Supports Whisper, Transducer, Paraformer, Zipformer, Conformer, Wenet,
5
+ // Telespeech, Moonshine, Dolphin, NeMo and SenseVoice through a unified
6
+ // AsrModelConfig structure.
7
+ // ------------------------------------------------------------------------------
8
+ #pragma once
9
+
10
+ #include "AudioUtils.hpp"
11
+ #include "ModelSingleton.hpp"
12
+ #include "ThreadPool.hpp"
13
+
14
+ #include "AsrModelType.hpp"
15
+
16
+ #include <memory>
17
+ #include <string>
18
+ #include <vector>
19
+
20
+ struct SherpaOnnxOfflineRecognizer;
21
+ struct SherpaOnnxOnlineRecognizer;
22
+ struct SherpaOnnxOfflineStream;
23
+ struct SherpaOnnxOnlineStream;
24
+
25
+ namespace margelo::nitro::onnx::speech {
26
+
27
+ /** Unified native configuration for any supported ASR model. */
28
+ struct AsrEngineConfig {
29
+ AsrModelType type = AsrModelType::WHISPER;
30
+ std::string modelDir;
31
+ std::string tokensPath;
32
+ std::string whisperEncoder;
33
+ std::string whisperDecoder;
34
+ std::string encoder;
35
+ std::string decoder;
36
+ std::string joiner;
37
+ std::string model;
38
+ std::string config;
39
+ int32_t numThreads = 4;
40
+ std::string decodingMethod = "greedy_search";
41
+ int32_t maxActivePaths = 4;
42
+ std::string language = "en";
43
+ bool useItn = true;
44
+ };
45
+
46
+ /** Native recognition result before conversion to the generated AsrResult. */
47
+ struct AsrEngineResult {
48
+ std::string text;
49
+ float score = 0.0f;
50
+ float startMs = 0.0f;
51
+ float endMs = 0.0f;
52
+ std::vector<float> timestamps;
53
+ std::string json;
54
+ };
55
+
56
+ /** Streaming ASR event listener. */
57
+ class StreamingAsrListener {
58
+ public:
59
+ virtual ~StreamingAsrListener() = default;
60
+ virtual void onPartialResult(const AsrEngineResult& result) = 0;
61
+ virtual void onFinalResult(const AsrEngineResult& result) = 0;
62
+ virtual void onError(const std::string& error) = 0;
63
+ };
64
+
65
+ /** Offline ASR engine. */
66
+ class OfflineAsrEngine final {
67
+ public:
68
+ explicit OfflineAsrEngine(std::shared_ptr<ThreadPool> threadPool);
69
+ ~OfflineAsrEngine();
70
+
71
+ OfflineAsrEngine(const OfflineAsrEngine&) = delete;
72
+ OfflineAsrEngine& operator=(const OfflineAsrEngine&) = delete;
73
+
74
+ void load(const AsrEngineConfig& config);
75
+ bool isLoaded() const;
76
+ AsrEngineResult recognize(const std::vector<float>& samples);
77
+ AsrEngineResult recognizeFile(const std::string& path);
78
+ void unload();
79
+
80
+ private:
81
+ std::shared_ptr<ThreadPool> threadPool_;
82
+ AsrEngineConfig config_;
83
+ std::shared_ptr<const SherpaOnnxOfflineRecognizer> recognizer_;
84
+ };
85
+
86
+ /** Streaming ASR engine. */
87
+ class StreamingAsrEngine final {
88
+ public:
89
+ explicit StreamingAsrEngine(std::shared_ptr<ThreadPool> threadPool);
90
+ ~StreamingAsrEngine();
91
+
92
+ StreamingAsrEngine(const StreamingAsrEngine&) = delete;
93
+ StreamingAsrEngine& operator=(const StreamingAsrEngine&) = delete;
94
+
95
+ void load(const AsrEngineConfig& config, std::shared_ptr<StreamingAsrListener> listener);
96
+ bool isLoaded() const;
97
+ void acceptWaveform(const std::vector<float>& samples);
98
+ AsrEngineResult finalize();
99
+ void reset();
100
+ void unload();
101
+
102
+ private:
103
+ std::shared_ptr<ThreadPool> threadPool_;
104
+ AsrEngineConfig config_;
105
+ std::shared_ptr<StreamingAsrListener> listener_;
106
+ std::shared_ptr<const SherpaOnnxOnlineRecognizer> recognizer_;
107
+ std::shared_ptr<const SherpaOnnxOnlineStream> stream_;
108
+ };
109
+
110
+ } // namespace margelo::nitro::onnx::speech
@@ -0,0 +1,177 @@
1
+ // ------------------------------------------------------------------------------
2
+ // AudioFileReader.cpp
3
+ // ------------------------------------------------------------------------------
4
+ #include "AudioFileReader.hpp"
5
+
6
+ #include <cmath>
7
+ #include <cstring>
8
+ #include <fstream>
9
+ #include <stdexcept>
10
+
11
+ namespace margelo::nitro::onnx::speech {
12
+
13
+ namespace {
14
+
15
+ constexpr int32_t kTargetSampleRate = 16000;
16
+
17
+ struct WavHeader {
18
+ uint16_t audioFormat = 0;
19
+ uint16_t numChannels = 0;
20
+ uint32_t sampleRate = 0;
21
+ uint16_t bitsPerSample = 0;
22
+ };
23
+
24
+ WavHeader parseFmtChunk(const uint8_t* data, size_t size) {
25
+ if (size < 16) {
26
+ throw std::runtime_error("WAV fmt chunk too small");
27
+ }
28
+ WavHeader h;
29
+ std::memcpy(&h.audioFormat, data, 2);
30
+ std::memcpy(&h.numChannels, data + 2, 2);
31
+ std::memcpy(&h.sampleRate, data + 4, 4);
32
+ std::memcpy(&h.bitsPerSample, data + 12, 2);
33
+
34
+ if (h.audioFormat != 1 && h.audioFormat != 3) {
35
+ throw std::runtime_error("Unsupported WAV audio format (only PCM=1 and IEEE float=3 supported)");
36
+ }
37
+ if (h.numChannels == 0 || h.sampleRate == 0 || h.bitsPerSample == 0) {
38
+ throw std::runtime_error("Invalid WAV header: zero channels, sample rate, or bits per sample");
39
+ }
40
+ return h;
41
+ }
42
+
43
+ std::vector<float> decodeSamples(const uint8_t* data, size_t byteCount, uint16_t audioFormat, uint16_t bitsPerSample) {
44
+ if (audioFormat == 1 && bitsPerSample == 16) {
45
+ const size_t sampleCount = byteCount / 2;
46
+ std::vector<float> out(sampleCount);
47
+ const auto* raw = reinterpret_cast<const int16_t*>(data);
48
+ constexpr float scale = 1.0f / 32768.0f;
49
+ for (size_t i = 0; i < sampleCount; ++i) {
50
+ out[i] = static_cast<float>(raw[i]) * scale;
51
+ }
52
+ return out;
53
+ }
54
+ if (audioFormat == 1 && bitsPerSample == 32) {
55
+ const size_t sampleCount = byteCount / 4;
56
+ std::vector<float> out(sampleCount);
57
+ const auto* raw = reinterpret_cast<const int32_t*>(data);
58
+ constexpr float scale = 1.0f / 2147483648.0f;
59
+ for (size_t i = 0; i < sampleCount; ++i) {
60
+ out[i] = static_cast<float>(raw[i]) * scale;
61
+ }
62
+ return out;
63
+ }
64
+ if (audioFormat == 3 && bitsPerSample == 32) {
65
+ const size_t sampleCount = byteCount / 4;
66
+ std::vector<float> out(sampleCount);
67
+ std::memcpy(out.data(), data, sampleCount * sizeof(float));
68
+ return out;
69
+ }
70
+ throw std::runtime_error("Unsupported WAV sample format: format=" +
71
+ std::to_string(audioFormat) + " bits=" + std::to_string(bitsPerSample));
72
+ }
73
+
74
+ std::vector<float> toMono(const std::vector<float>& samples, uint16_t channels) {
75
+ if (channels == 1) return samples;
76
+ const size_t frames = samples.size() / channels;
77
+ std::vector<float> out(frames);
78
+ const float inv = 1.0f / static_cast<float>(channels);
79
+ for (size_t i = 0; i < frames; ++i) {
80
+ float sum = 0.0f;
81
+ for (uint16_t c = 0; c < channels; ++c) {
82
+ sum += samples[i * channels + c];
83
+ }
84
+ out[i] = sum * inv;
85
+ }
86
+ return out;
87
+ }
88
+
89
+ std::vector<float> resample(const std::vector<float>& samples, uint32_t fromRate, uint32_t toRate) {
90
+ if (fromRate == toRate || samples.empty()) return samples;
91
+ const double ratio = static_cast<double>(fromRate) / static_cast<double>(toRate);
92
+ const size_t outSize = static_cast<size_t>(std::ceil(samples.size() / ratio));
93
+ std::vector<float> out(outSize);
94
+ for (size_t i = 0; i < outSize; ++i) {
95
+ const double srcPos = static_cast<double>(i) * ratio;
96
+ const size_t idx = static_cast<size_t>(srcPos);
97
+ const double frac = srcPos - static_cast<double>(idx);
98
+ if (idx + 1 < samples.size()) {
99
+ out[i] = static_cast<float>(samples[idx] * (1.0 - frac) + samples[idx + 1] * frac);
100
+ } else {
101
+ out[i] = samples[idx];
102
+ }
103
+ }
104
+ return out;
105
+ }
106
+
107
+ } // namespace
108
+
109
+ std::vector<float> readWavFile(const std::string& path) {
110
+ std::ifstream file(path, std::ios::binary);
111
+ if (!file) {
112
+ throw std::runtime_error("Cannot open WAV file: " + path);
113
+ }
114
+
115
+ file.seekg(0, std::ios::end);
116
+ const auto fileSize = file.tellg();
117
+ file.seekg(0, std::ios::beg);
118
+
119
+ if (fileSize < 44) {
120
+ throw std::runtime_error("File too small to be a valid WAV: " + path);
121
+ }
122
+
123
+ std::vector<uint8_t> buf(static_cast<size_t>(fileSize));
124
+ file.read(reinterpret_cast<char*>(buf.data()), fileSize);
125
+
126
+ if (std::memcmp(buf.data(), "RIFF", 4) != 0 || std::memcmp(buf.data() + 8, "WAVE", 4) != 0) {
127
+ throw std::runtime_error("Not a valid WAV file: " + path);
128
+ }
129
+
130
+ const uint8_t* dataChunk = nullptr;
131
+ size_t dataChunkSize = 0;
132
+ std::optional<WavHeader> header;
133
+
134
+ size_t pos = 12;
135
+ while (pos + 8 <= buf.size()) {
136
+ const char* chunkId = reinterpret_cast<const char*>(buf.data() + pos);
137
+ uint32_t chunkSize = 0;
138
+ std::memcpy(&chunkSize, buf.data() + pos + 4, 4);
139
+
140
+ if (std::memcmp(chunkId, "fmt ", 4) == 0) {
141
+ header = parseFmtChunk(buf.data() + pos + 8, chunkSize);
142
+ } else if (std::memcmp(chunkId, "data", 4) == 0) {
143
+ dataChunk = buf.data() + pos + 8;
144
+ dataChunkSize = chunkSize;
145
+ }
146
+
147
+ pos += 8 + chunkSize;
148
+ if (chunkSize % 2 != 0) ++pos;
149
+ }
150
+
151
+ if (!header.has_value() || dataChunk == nullptr) {
152
+ throw std::runtime_error("WAV file missing fmt or data chunk: " + path);
153
+ }
154
+
155
+ auto samples = decodeSamples(dataChunk, dataChunkSize, header->audioFormat, header->bitsPerSample);
156
+ samples = toMono(samples, header->numChannels);
157
+ samples = resample(samples, header->sampleRate, kTargetSampleRate);
158
+ return samples;
159
+ }
160
+
161
+ std::vector<float> readRawPcmFile(const std::string& path) {
162
+ std::ifstream file(path, std::ios::binary | std::ios::ate);
163
+ if (!file) {
164
+ throw std::runtime_error("Cannot open raw PCM file: " + path);
165
+ }
166
+ const auto byteCount = file.tellg();
167
+ if (byteCount <= 0 || static_cast<size_t>(byteCount) % sizeof(float) != 0) {
168
+ throw std::runtime_error("Raw PCM file size is not a multiple of sizeof(float): " + path);
169
+ }
170
+ file.seekg(0, std::ios::beg);
171
+ const size_t sampleCount = static_cast<size_t>(byteCount) / sizeof(float);
172
+ std::vector<float> samples(sampleCount);
173
+ file.read(reinterpret_cast<char*>(samples.data()), byteCount);
174
+ return samples;
175
+ }
176
+
177
+ } // namespace margelo::nitro::onnx::speech
@@ -0,0 +1,26 @@
1
+ // ------------------------------------------------------------------------------
2
+ // AudioFileReader.hpp
3
+ // Minimal WAV reader that outputs 16 kHz mono f32 PCM.
4
+ // Supports s16 and f32 sample formats; resamples from other rates via linear
5
+ // interpolation.
6
+ // ------------------------------------------------------------------------------
7
+ #pragma once
8
+
9
+ #include <cstdint>
10
+ #include <string>
11
+ #include <vector>
12
+
13
+ namespace margelo::nitro::onnx::speech {
14
+
15
+ /**
16
+ * Read a WAV file and return 16 kHz mono f32 PCM samples.
17
+ * Throws std::runtime_error on invalid format or I/O failure.
18
+ */
19
+ std::vector<float> readWavFile(const std::string& path);
20
+
21
+ /**
22
+ * Read a raw PCM file (16 kHz mono f32le) and return the samples directly.
23
+ */
24
+ std::vector<float> readRawPcmFile(const std::string& path);
25
+
26
+ } // namespace margelo::nitro::onnx::speech
@@ -0,0 +1,28 @@
1
+ // ------------------------------------------------------------------------------
2
+ // AudioUtils.cpp
3
+ // ------------------------------------------------------------------------------
4
+ #include "AudioUtils.hpp"
5
+
6
+ #include <cstring>
7
+ #include <stdexcept>
8
+
9
+ namespace margelo::nitro::onnx::speech {
10
+
11
+ std::vector<float> bytesToFloatVector(const uint8_t* data, size_t byteCount) {
12
+ if (byteCount % sizeof(float) != 0) {
13
+ throw std::invalid_argument("Audio byte count must be a multiple of sizeof(float)");
14
+ }
15
+ const size_t sampleCount = byteCount / sizeof(float);
16
+ std::vector<float> result(sampleCount);
17
+ std::memcpy(result.data(), data, byteCount);
18
+ return result;
19
+ }
20
+
21
+ std::vector<uint8_t> floatVectorToBytes(const std::vector<float>& samples) {
22
+ const size_t byteCount = samples.size() * sizeof(float);
23
+ std::vector<uint8_t> result(byteCount);
24
+ std::memcpy(result.data(), samples.data(), byteCount);
25
+ return result;
26
+ }
27
+
28
+ } // namespace margelo::nitro::onnx::speech
@@ -0,0 +1,38 @@
1
+ // ------------------------------------------------------------------------------
2
+ // AudioUtils.hpp
3
+ // Helpers for converting between Nitro ArrayBuffers and sherpa-onnx float
4
+ // sample vectors. All audio is expected to be 16 kHz mono f32 PCM.
5
+ // ------------------------------------------------------------------------------
6
+ #pragma once
7
+
8
+ #include <cstdint>
9
+ #include <vector>
10
+
11
+ namespace margelo::nitro::onnx::speech {
12
+
13
+ /** Number of audio samples that represent one millisecond at 16 kHz. */
14
+ constexpr float kSamplesPerMs = 16.0f;
15
+
16
+ /** Convert milliseconds to sample count at 16 kHz. */
17
+ inline int32_t msToSamples(int32_t ms) {
18
+ return static_cast<int32_t>(ms * kSamplesPerMs);
19
+ }
20
+
21
+ /** Convert sample count at 16 kHz to milliseconds. */
22
+ inline float samplesToMs(int32_t samples) {
23
+ return static_cast<float>(samples) / kSamplesPerMs;
24
+ }
25
+
26
+ /**
27
+ * Copy raw bytes into a float vector.
28
+ * The caller must ensure the byte count is a multiple of sizeof(float).
29
+ */
30
+ std::vector<float> bytesToFloatVector(const uint8_t* data, size_t byteCount);
31
+
32
+ /**
33
+ * Copy a float vector into a newly allocated byte buffer.
34
+ * Returned buffer size is samples * sizeof(float).
35
+ */
36
+ std::vector<uint8_t> floatVectorToBytes(const std::vector<float>& samples);
37
+
38
+ } // namespace margelo::nitro::onnx::speech