From 1b085c18c1bff7a516e580179a22920f946e1885 Mon Sep 17 00:00:00 2001 From: Fangjun Kuang Date: Thu, 29 Jan 2026 10:26:26 +0800 Subject: [PATCH] Refactor JNI --- java-api-examples/PocketTts.java | 31 ++-- .../com/k2fsa/sherpa/onnx/AudioTagging.java | 3 + .../com/k2fsa/sherpa/onnx/KeywordSpotter.java | 3 + .../k2fsa/sherpa/onnx/OfflinePunctuation.java | 3 + .../k2fsa/sherpa/onnx/OfflineRecognizer.java | 3 + .../onnx/OfflineSpeakerDiarization.java | 5 +- .../sherpa/onnx/OfflineSpeechDenoiser.java | 3 + .../com/k2fsa/sherpa/onnx/OfflineTts.java | 62 +++++++- .../k2fsa/sherpa/onnx/OfflineTtsCallback.java | 4 + .../k2fsa/sherpa/onnx/OnlineRecognizer.java | 3 + .../onnx/SpeakerEmbeddingExtractor.java | 3 + .../onnx/SpokenLanguageIdentification.java | 3 + .../main/java/com/k2fsa/sherpa/onnx/Vad.java | 3 + sherpa-onnx/jni/offline-tts.cc | 132 +++++------------- 14 files changed, 134 insertions(+), 127 deletions(-) diff --git a/java-api-examples/PocketTts.java b/java-api-examples/PocketTts.java index 5cdaad7ba8..34497a773f 100644 --- a/java-api-examples/PocketTts.java +++ b/java-api-examples/PocketTts.java @@ -5,8 +5,6 @@ import com.k2fsa.sherpa.onnx.*; import java.util.HashMap; import java.util.Map; -import java.util.function.Consumer; -import java.util.function.Function; public class PocketTts { public static void main(String[] args) { @@ -92,22 +90,10 @@ public Integer invoke(float[] samples) { tts.generateWithConfigAndCallback( text, genConfig, - (OfflineTtsCallback) - samples -> { - System.out.println("Lambda Integer callback: " + samples.length); - return 1; // continue - }); - } - - if (false) { - audio = - tts.generateWithConfigAndCallback( - text, - genConfig, - (Consumer) - samples -> { - System.out.println("Consumer: " + samples.length); - }); + samples -> { + System.out.println("Lambda Integer callback: " + samples.length); + return 1; // continue + }); } if (false) { @@ -115,11 +101,10 @@ public Integer invoke(float[] samples) { tts.generateWithConfigAndCallback( text, genConfig, - (Function) - samples -> { - System.out.println("Function: " + samples.length); - return 1; - }); + samples -> { + System.out.println("Consumer: " + samples.length); + // implicitly, it returns 1 internally + }); } if (audio == null) { diff --git a/sherpa-onnx/java-api/src/main/java/com/k2fsa/sherpa/onnx/AudioTagging.java b/sherpa-onnx/java-api/src/main/java/com/k2fsa/sherpa/onnx/AudioTagging.java index b4ccb09d6d..1e9dc64e08 100644 --- a/sherpa-onnx/java-api/src/main/java/com/k2fsa/sherpa/onnx/AudioTagging.java +++ b/sherpa-onnx/java-api/src/main/java/com/k2fsa/sherpa/onnx/AudioTagging.java @@ -8,6 +8,9 @@ public class AudioTagging { public AudioTagging(AudioTaggingConfig config) { LibraryLoader.maybeLoad(); ptr = newFromFile(config); + if (ptr == 0) { + throw new IllegalArgumentException("Invalid AudioTaggingConfig: failed to create native AudioTagging"); + } } public OfflineStream createStream() { diff --git a/sherpa-onnx/java-api/src/main/java/com/k2fsa/sherpa/onnx/KeywordSpotter.java b/sherpa-onnx/java-api/src/main/java/com/k2fsa/sherpa/onnx/KeywordSpotter.java index a42108c237..f929110876 100644 --- a/sherpa-onnx/java-api/src/main/java/com/k2fsa/sherpa/onnx/KeywordSpotter.java +++ b/sherpa-onnx/java-api/src/main/java/com/k2fsa/sherpa/onnx/KeywordSpotter.java @@ -8,6 +8,9 @@ public class KeywordSpotter { public KeywordSpotter(KeywordSpotterConfig config) { LibraryLoader.maybeLoad(); ptr = newFromFile(config); + if (ptr == 0) { + throw new IllegalArgumentException("Invalid KeywordSpotterConfig: failed to create native KeywordSpotter"); + } } public OnlineStream createStream(String keywords) { diff --git a/sherpa-onnx/java-api/src/main/java/com/k2fsa/sherpa/onnx/OfflinePunctuation.java b/sherpa-onnx/java-api/src/main/java/com/k2fsa/sherpa/onnx/OfflinePunctuation.java index 5597d86e49..0ccec271d6 100644 --- a/sherpa-onnx/java-api/src/main/java/com/k2fsa/sherpa/onnx/OfflinePunctuation.java +++ b/sherpa-onnx/java-api/src/main/java/com/k2fsa/sherpa/onnx/OfflinePunctuation.java @@ -8,6 +8,9 @@ public class OfflinePunctuation { public OfflinePunctuation(OfflinePunctuationConfig config) { LibraryLoader.maybeLoad(); ptr = newFromFile(config); + if (ptr == 0) { + throw new IllegalArgumentException("Invalid OfflinePunctuationConfig: failed to create native OfflinePunctuation"); + } } public String addPunctuation(String text) { diff --git a/sherpa-onnx/java-api/src/main/java/com/k2fsa/sherpa/onnx/OfflineRecognizer.java b/sherpa-onnx/java-api/src/main/java/com/k2fsa/sherpa/onnx/OfflineRecognizer.java index ac847b23b8..47e449faf4 100644 --- a/sherpa-onnx/java-api/src/main/java/com/k2fsa/sherpa/onnx/OfflineRecognizer.java +++ b/sherpa-onnx/java-api/src/main/java/com/k2fsa/sherpa/onnx/OfflineRecognizer.java @@ -9,6 +9,9 @@ public class OfflineRecognizer { public OfflineRecognizer(OfflineRecognizerConfig config) { LibraryLoader.maybeLoad(); ptr = newFromFile(config); + if (ptr == 0) { + throw new IllegalArgumentException("Invalid OfflineRecognizerConfig: failed to create native OfflineRecognizer"); + } this.config = config; } diff --git a/sherpa-onnx/java-api/src/main/java/com/k2fsa/sherpa/onnx/OfflineSpeakerDiarization.java b/sherpa-onnx/java-api/src/main/java/com/k2fsa/sherpa/onnx/OfflineSpeakerDiarization.java index 41dca34ad5..ed3cdd5546 100644 --- a/sherpa-onnx/java-api/src/main/java/com/k2fsa/sherpa/onnx/OfflineSpeakerDiarization.java +++ b/sherpa-onnx/java-api/src/main/java/com/k2fsa/sherpa/onnx/OfflineSpeakerDiarization.java @@ -8,6 +8,9 @@ public class OfflineSpeakerDiarization { public OfflineSpeakerDiarization(OfflineSpeakerDiarizationConfig config) { LibraryLoader.maybeLoad(); ptr = newFromFile(config); + if (ptr == 0) { + throw new IllegalArgumentException("Invalid OfflineSpeakerDiarizationConfig: failed to create native OfflineSpeakerDiarization"); + } } public int getSampleRate() { @@ -55,4 +58,4 @@ public void release() { private native OfflineSpeakerDiarizationSegment[] process(long ptr, float[] samples); private native OfflineSpeakerDiarizationSegment[] processWithCallback(long ptr, float[] samples, OfflineSpeakerDiarizationCallback callback, long arg); -} \ No newline at end of file +} diff --git a/sherpa-onnx/java-api/src/main/java/com/k2fsa/sherpa/onnx/OfflineSpeechDenoiser.java b/sherpa-onnx/java-api/src/main/java/com/k2fsa/sherpa/onnx/OfflineSpeechDenoiser.java index 63e0e7db65..860906ee41 100644 --- a/sherpa-onnx/java-api/src/main/java/com/k2fsa/sherpa/onnx/OfflineSpeechDenoiser.java +++ b/sherpa-onnx/java-api/src/main/java/com/k2fsa/sherpa/onnx/OfflineSpeechDenoiser.java @@ -8,6 +8,9 @@ public class OfflineSpeechDenoiser { public OfflineSpeechDenoiser(OfflineSpeechDenoiserConfig config) { LibraryLoader.maybeLoad(); ptr = newFromFile(config); + if (ptr == 0) { + throw new IllegalArgumentException("Invalid OfflineSpeechDenoiserConfig: failed to create native OfflineSpeechDenoiser"); + } } public int getSampleRate() { diff --git a/sherpa-onnx/java-api/src/main/java/com/k2fsa/sherpa/onnx/OfflineTts.java b/sherpa-onnx/java-api/src/main/java/com/k2fsa/sherpa/onnx/OfflineTts.java index 3afe8fc713..a9f1b217fd 100644 --- a/sherpa-onnx/java-api/src/main/java/com/k2fsa/sherpa/onnx/OfflineTts.java +++ b/sherpa-onnx/java-api/src/main/java/com/k2fsa/sherpa/onnx/OfflineTts.java @@ -2,12 +2,17 @@ package com.k2fsa.sherpa.onnx; +import java.util.function.Consumer; + public class OfflineTts { private long ptr = 0; public OfflineTts(OfflineTtsConfig config) { LibraryLoader.maybeLoad(); ptr = newFromFile(config); + if (ptr == 0) { + throw new IllegalArgumentException("Invalid OfflineTtsConfig: failed to create native OfflineTts"); + } } /** Returns the sample rate of the TTS engine. */ @@ -15,6 +20,10 @@ public int getSampleRate() { return getSampleRate(ptr); } + public int getNumSpeakers() { + return getNumSpeakers(ptr); + } + /** Generates audio for the given text using the default speaker (sid=0) and speed=1.0. */ public GeneratedAudio generate(String text) { return generate(text, 0, 1.0f); @@ -30,18 +39,47 @@ public GeneratedAudio generate(String text, int sid, float speed) { return generateImpl(ptr, text, sid, speed); } - public GeneratedAudio generateWithCallback(String text, Object callback) { + public GeneratedAudio generateWithCallback(String text, OfflineTtsCallback callback) { return generateWithCallback(text, 0, 1.0f, callback); } - public GeneratedAudio generateWithCallback(String text, int sid, Object callback) { + public GeneratedAudio generateWithCallback( + String text, + Consumer consumer + ) { + return generateWithCallback(text, 0, 1.0f, consumer); + } + + public GeneratedAudio generateWithCallback(String text, int sid, OfflineTtsCallback callback) { return generateWithCallback(text, sid, 1.0f, callback); } - public GeneratedAudio generateWithCallback(String text, int sid, float speed, Object callback) { + public GeneratedAudio generateWithCallback( + String text, + int sid, + Consumer consumer + ) { + + return generateWithCallback(text, sid, 1.0f, consumer); + } + + public GeneratedAudio generateWithCallback(String text, int sid, float speed, OfflineTtsCallback callback) { return generateWithCallbackImpl(ptr, text, sid, speed, callback); } + public GeneratedAudio generateWithCallback( + String text, + int sid, + float speed, + Consumer consumer + ) { + OfflineTtsCallback cb = samples -> { + consumer.accept(samples); + return 1; + }; + return generateWithCallback(text, sid, speed, cb); + } + /** * Generate audio using a GenerationConfig and a callback. * @@ -53,12 +91,24 @@ public GeneratedAudio generateWithCallback(String text, int sid, float speed, Ob public GeneratedAudio generateWithConfigAndCallback( String text, GenerationConfig config, - Object callback + OfflineTtsCallback callback ) { return generateWithConfigImpl(ptr, text, config, callback); } + public GeneratedAudio generateWithConfigAndCallback( + String text, + GenerationConfig config, + Consumer consumer + ) { + OfflineTtsCallback cb = samples -> { + consumer.accept(samples); + return 1; + }; + return generateWithConfigAndCallback(text, config, cb); + } + @Override protected void finalize() throws Throwable { release(); @@ -80,13 +130,13 @@ public void release() { private native GeneratedAudio generateImpl(long ptr, String text, int sid, float speed); - private native GeneratedAudio generateWithCallbackImpl(long ptr, String text, int sid, float speed, Object callback); + private native GeneratedAudio generateWithCallbackImpl(long ptr, String text, int sid, float speed, OfflineTtsCallback callback); private native GeneratedAudio generateWithConfigImpl( long ptr, String text, GenerationConfig config, - Object callback + OfflineTtsCallback callback ); private native long newFromFile(OfflineTtsConfig config); diff --git a/sherpa-onnx/java-api/src/main/java/com/k2fsa/sherpa/onnx/OfflineTtsCallback.java b/sherpa-onnx/java-api/src/main/java/com/k2fsa/sherpa/onnx/OfflineTtsCallback.java index 2fc1d45dde..81669f3db6 100644 --- a/sherpa-onnx/java-api/src/main/java/com/k2fsa/sherpa/onnx/OfflineTtsCallback.java +++ b/sherpa-onnx/java-api/src/main/java/com/k2fsa/sherpa/onnx/OfflineTtsCallback.java @@ -4,5 +4,9 @@ @FunctionalInterface public interface OfflineTtsCallback { + /** + * @param samples audio chunk + * @return 1 to continue, 0 to stop + */ Integer invoke(float[] samples); } diff --git a/sherpa-onnx/java-api/src/main/java/com/k2fsa/sherpa/onnx/OnlineRecognizer.java b/sherpa-onnx/java-api/src/main/java/com/k2fsa/sherpa/onnx/OnlineRecognizer.java index 1de28f7c0a..138cd1abcf 100644 --- a/sherpa-onnx/java-api/src/main/java/com/k2fsa/sherpa/onnx/OnlineRecognizer.java +++ b/sherpa-onnx/java-api/src/main/java/com/k2fsa/sherpa/onnx/OnlineRecognizer.java @@ -9,6 +9,9 @@ public class OnlineRecognizer { public OnlineRecognizer(OnlineRecognizerConfig config) { LibraryLoader.maybeLoad(); ptr = newFromFile(config); + if (ptr == 0) { + throw new IllegalArgumentException("Invalid OnlineRecognizerConfig: failed to create native OnlineRecognizer"); + } } public void decode(OnlineStream s) { diff --git a/sherpa-onnx/java-api/src/main/java/com/k2fsa/sherpa/onnx/SpeakerEmbeddingExtractor.java b/sherpa-onnx/java-api/src/main/java/com/k2fsa/sherpa/onnx/SpeakerEmbeddingExtractor.java index d2ff7ff9f5..2aace7aec8 100644 --- a/sherpa-onnx/java-api/src/main/java/com/k2fsa/sherpa/onnx/SpeakerEmbeddingExtractor.java +++ b/sherpa-onnx/java-api/src/main/java/com/k2fsa/sherpa/onnx/SpeakerEmbeddingExtractor.java @@ -8,6 +8,9 @@ public class SpeakerEmbeddingExtractor { public SpeakerEmbeddingExtractor(SpeakerEmbeddingExtractorConfig config) { LibraryLoader.maybeLoad(); ptr = newFromFile(config); + if (ptr == 0) { + throw new IllegalArgumentException("Invalid SpeakerEmbeddingExtractorConfig: failed to create native SpeakerEmbeddingExtractor"); + } } @Override diff --git a/sherpa-onnx/java-api/src/main/java/com/k2fsa/sherpa/onnx/SpokenLanguageIdentification.java b/sherpa-onnx/java-api/src/main/java/com/k2fsa/sherpa/onnx/SpokenLanguageIdentification.java index 814d9ff4fb..12619afae3 100644 --- a/sherpa-onnx/java-api/src/main/java/com/k2fsa/sherpa/onnx/SpokenLanguageIdentification.java +++ b/sherpa-onnx/java-api/src/main/java/com/k2fsa/sherpa/onnx/SpokenLanguageIdentification.java @@ -13,6 +13,9 @@ public class SpokenLanguageIdentification { public SpokenLanguageIdentification(SpokenLanguageIdentificationConfig config) { LibraryLoader.maybeLoad(); ptr = newFromFile(config); + if (ptr == 0) { + throw new IllegalArgumentException("Invalid SpokenLanguageIdentificationConfig: failed to create native SpokenLanguageIdentification"); + } String[] languages = Locale.getISOLanguages(); localeMap = new HashMap(languages.length); diff --git a/sherpa-onnx/java-api/src/main/java/com/k2fsa/sherpa/onnx/Vad.java b/sherpa-onnx/java-api/src/main/java/com/k2fsa/sherpa/onnx/Vad.java index 9a9e2fb663..0a15a725d0 100644 --- a/sherpa-onnx/java-api/src/main/java/com/k2fsa/sherpa/onnx/Vad.java +++ b/sherpa-onnx/java-api/src/main/java/com/k2fsa/sherpa/onnx/Vad.java @@ -8,6 +8,9 @@ public class Vad { public Vad(VadModelConfig config) { LibraryLoader.maybeLoad(); ptr = newFromFile(config); + if (ptr == 0) { + throw new IllegalArgumentException("Invalid VadModelConfig: failed to create native Vad"); + } } @Override diff --git a/sherpa-onnx/jni/offline-tts.cc b/sherpa-onnx/jni/offline-tts.cc index ece5e317fc..b2fe1f29a9 100644 --- a/sherpa-onnx/jni/offline-tts.cc +++ b/sherpa-onnx/jni/offline-tts.cc @@ -316,68 +316,10 @@ static jobject CreateAudioObject(JNIEnv *env, const std::vector &samples, return gen_audio_obj; } -// ----------------- Consumer ----------------- -static int32_t CallConsumerCallback(JNIEnv *env, jobject callback, - jfloatArray samples_arr) { - jclass consumer_cls = env->FindClass("java/util/function/Consumer"); - if (env->ExceptionCheck() || !env->IsInstanceOf(callback, consumer_cls)) { - env->DeleteLocalRef(consumer_cls); - return -1; // not a Consumer - } - - jmethodID accept_mid = - env->GetMethodID(consumer_cls, "accept", "(Ljava/lang/Object;)V"); - if (env->ExceptionCheck()) { - env->DeleteLocalRef(consumer_cls); - return -1; - } - - env->CallVoidMethod(callback, accept_mid, samples_arr); - if (env->ExceptionCheck()) { - env->DeleteLocalRef(consumer_cls); - return 1; // exception occurred, continue - } - - env->DeleteLocalRef(consumer_cls); - return 1; // continue -} - -// ----------------- Function ----------------- -static int32_t CallFunctionCallback(JNIEnv *env, jobject callback, - jfloatArray samples_arr) { - jclass function_cls = env->FindClass("java/util/function/Function"); - if (env->ExceptionCheck() || !env->IsInstanceOf(callback, function_cls)) { - env->DeleteLocalRef(function_cls); - return -1; // not a Function - } - - jmethodID apply_mid = env->GetMethodID( - function_cls, "apply", "(Ljava/lang/Object;)Ljava/lang/Object;"); - if (env->ExceptionCheck()) { - env->DeleteLocalRef(function_cls); - return -1; - } - - jobject result = env->CallObjectMethod(callback, apply_mid, samples_arr); - if (env->ExceptionCheck() || !result) { - env->DeleteLocalRef(function_cls); - return 1; // exception or null → continue - } - - jclass integer_cls = env->FindClass("java/lang/Integer"); - jmethodID int_val_mid = env->GetMethodID(integer_cls, "intValue", "()I"); - jint ret = env->CallIntMethod(result, int_val_mid); - - env->DeleteLocalRef(integer_cls); - env->DeleteLocalRef(result); - env->DeleteLocalRef(function_cls); - - return ret; -} +static int32_t CallCallback(JNIEnv *env, jobject callback, + jfloatArray samples_arr) { + if (!callback) return 1; -// ----------------- OfflineTtsCallback.invoke ----------------- -static int32_t CallInvokeCallback(JNIEnv *env, jobject callback, - jfloatArray samples_arr) { jclass cls = env->GetObjectClass(callback); if (env->ExceptionCheck()) { env->DeleteLocalRef(cls); @@ -408,24 +350,6 @@ static int32_t CallInvokeCallback(JNIEnv *env, jobject callback, return ret; } -static int32_t CallCallback(JNIEnv *env, jobject callback, - jfloatArray samples_arr) { - if (!callback) return 1; - - int32_t ret; - - // Try Consumer - ret = CallConsumerCallback(env, callback, samples_arr); - if (ret != -1) return ret; - - // Try Function - ret = CallFunctionCallback(env, callback, samples_arr); - if (ret != -1) return ret; - - // Fallback to invoke() - return CallInvokeCallback(env, callback, samples_arr); -} - SHERPA_ONNX_EXTERN_C JNIEXPORT jlong JNICALL Java_com_k2fsa_sherpa_onnx_OfflineTts_newFromAsset( JNIEnv *env, jobject /*obj*/, jobject asset_manager, jobject _config) { @@ -526,16 +450,23 @@ Java_com_k2fsa_sherpa_onnx_OfflineTts_generateWithCallbackImpl( auto tts = reinterpret_cast(ptr); - std::function callback_wrapper = - [env, callback](const float *samples, int32_t n, float) -> int32_t { - jfloatArray samples_arr = env->NewFloatArray(n); - env->SetFloatArrayRegion(samples_arr, 0, n, samples); - int32_t ret = CallCallback(env, callback, samples_arr); - env->DeleteLocalRef(samples_arr); - return ret; - }; + sherpa_onnx::GeneratedAudio audio; + + if (callback) { + std::function callback_wrapper = + [env, callback](const float *samples, int32_t n, float) -> int32_t { + jfloatArray samples_arr = env->NewFloatArray(n); + env->SetFloatArrayRegion(samples_arr, 0, n, samples); + int32_t ret = CallCallback(env, callback, samples_arr); + env->DeleteLocalRef(samples_arr); + return ret; + }; + + audio = tts->Generate(p_text, sid, speed, callback_wrapper); + } else { + audio = tts->Generate(p_text, sid, speed, nullptr); + } - auto audio = tts->Generate(p_text, sid, speed, callback_wrapper); env->ReleaseStringUTFChars(text, p_text); return CreateAudioObject(env, audio.samples, audio.sample_rate); @@ -550,16 +481,23 @@ Java_com_k2fsa_sherpa_onnx_OfflineTts_generateWithConfigImpl( auto gen_config = sherpa_onnx::GetGenerationConfig(env, _gen_config); auto tts = reinterpret_cast(ptr); - std::function callback_wrapper = - [env, callback](const float *samples, int32_t n, float) -> int32_t { - jfloatArray samples_arr = env->NewFloatArray(n); - env->SetFloatArrayRegion(samples_arr, 0, n, samples); - int32_t ret = CallCallback(env, callback, samples_arr); - env->DeleteLocalRef(samples_arr); - return ret; - }; + sherpa_onnx::GeneratedAudio audio; + + if (callback) { + std::function callback_wrapper = + [env, callback](const float *samples, int32_t n, float) -> int32_t { + jfloatArray samples_arr = env->NewFloatArray(n); + env->SetFloatArrayRegion(samples_arr, 0, n, samples); + int32_t ret = CallCallback(env, callback, samples_arr); + env->DeleteLocalRef(samples_arr); + return ret; + }; + + audio = tts->Generate(p_text, gen_config, callback_wrapper); + } else { + audio = tts->Generate(p_text, gen_config, nullptr); + } - auto audio = tts->Generate(p_text, gen_config, callback_wrapper); env->ReleaseStringUTFChars(text, p_text); return CreateAudioObject(env, audio.samples, audio.sample_rate);