diff --git a/.github/scripts/test-python.sh b/.github/scripts/test-python.sh index 49523aaf11..ea8679a2a0 100755 --- a/.github/scripts/test-python.sh +++ b/.github/scripts/test-python.sh @@ -8,6 +8,17 @@ log() { echo -e "$(date '+%Y-%m-%d %H:%M:%S') (${fname}:${BASH_LINENO[0]}:${FUNCNAME[1]}) $*" } +log "test Google MedASR" +curl -SL -O https://github.com/k2-fsa/sherpa-onnx/releases/download/asr-models/sherpa-onnx-medasr-ctc-en-int8-2025-12-25.tar.bz2 +tar xvf sherpa-onnx-medasr-ctc-en-int8-2025-12-25.tar.bz2 +rm sherpa-onnx-medasr-ctc-en-int8-2025-12-25.tar.bz2 +ls -lh sherpa-onnx-medasr-ctc-en-int8-2025-12-25 + +ls -lh sherpa-onnx-medasr-ctc-en-int8-2025-12-25/test_wavs + +python3 ./python-api-examples/offline-medasr-ctc-decode-files.py +rm -rf sherpa-onnx-medasr-ctc-en-int8-2025-12-25 + log "test omnilingual ASR" curl -SL -O https://github.com/k2-fsa/sherpa-onnx/releases/download/asr-models/sherpa-onnx-omnilingual-asr-1600-languages-300M-ctc-int8-2025-11-12.tar.bz2 tar xvf sherpa-onnx-omnilingual-asr-1600-languages-300M-ctc-int8-2025-11-12.tar.bz2 diff --git a/.github/workflows/export-medasr-ctc-to-onnx.yaml b/.github/workflows/export-medasr-ctc-to-onnx.yaml index 80aa921fd7..9bfafe251c 100644 --- a/.github/workflows/export-medasr-ctc-to-onnx.yaml +++ b/.github/workflows/export-medasr-ctc-to-onnx.yaml @@ -3,7 +3,7 @@ name: export-medasr-ctc-to-onnx on: push: branches: - - export-medasr-onnx + - cpp-medasr-2 workflow_dispatch: concurrency: @@ -42,7 +42,9 @@ jobs: run: | cd scripts/medasr - curl -SL -O https://huggingface.co/csukuangfj/sherpa-onnx-medasr-ctc-en-int8-2025-12-25/resolve/main/test_wavs/0.wav + for i in $(seq 0 5); do + curl -SL -O https://huggingface.co/csukuangfj/sherpa-onnx-medasr-ctc-en-int8-2025-12-25/resolve/main/test_wavs/$i.wav + done curl -SL -O https://huggingface.co/csukuangfj/sherpa-onnx-medasr-ctc-en-int8-2025-12-25/resolve/main/test_wavs/transcript.txt @@ -53,7 +55,9 @@ jobs: run: | cd scripts/medasr - python3 test_onnx.py --model ./model.onnx --tokens ./tokens.txt --wav ./0.wav + for i in $(seq 0 5); do + python3 test_onnx.py --model ./model.onnx --tokens ./tokens.txt --wav ./$i.wav + done cat transcript.txt @@ -62,7 +66,9 @@ jobs: run: | cd scripts/medasr - python3 test_onnx.py --model ./model.int8.onnx --tokens ./tokens.txt --wav ./0.wav + for i in $(seq 0 5); do + python3 test_onnx.py --model ./model.int8.onnx --tokens ./tokens.txt --wav ./$i.wav + done cat transcript.txt diff --git a/python-api-examples/offline-medasr-ctc-decode-files.py b/python-api-examples/offline-medasr-ctc-decode-files.py new file mode 100755 index 0000000000..d231eb852d --- /dev/null +++ b/python-api-examples/offline-medasr-ctc-decode-files.py @@ -0,0 +1,142 @@ +#!/usr/bin/env python3 + +""" +This file shows how to use a non-streaming Google MedASR CTC model from +https://huggingface.co/google/medasr +to decode files. + +Please download model files from +https://github.com/k2-fsa/sherpa-onnx/releases/tag/asr-models + +For instance, + +wget https://github.com/k2-fsa/sherpa-onnx/releases/download/asr-models/sherpa-onnx-medasr-ctc-en-int8-2025-12-25.tar.bz2 +tar xvf sherpa-onnx-medasr-ctc-en-int8-2025-12-25.tar.bz2 +rm sherpa-onnx-medasr-ctc-en-int8-2025-12-25.tar.bz2 +""" + +import time +from pathlib import Path + +import librosa +import numpy as np +import sherpa_onnx + + +def create_recognizer(): + model = "./sherpa-onnx-medasr-ctc-en-int8-2025-12-25/model.int8.onnx" + tokens = "./sherpa-onnx-medasr-ctc-en-int8-2025-12-25/tokens.txt" + test_wav_0 = "./sherpa-onnx-medasr-ctc-en-int8-2025-12-25/test_wavs/0.wav" + test_wav_1 = "./sherpa-onnx-medasr-ctc-en-int8-2025-12-25/test_wavs/1.wav" + test_wav_2 = "./sherpa-onnx-medasr-ctc-en-int8-2025-12-25/test_wavs/2.wav" + test_wav_3 = "./sherpa-onnx-medasr-ctc-en-int8-2025-12-25/test_wavs/3.wav" + test_wav_4 = "./sherpa-onnx-medasr-ctc-en-int8-2025-12-25/test_wavs/4.wav" + test_wav_5 = "./sherpa-onnx-medasr-ctc-en-int8-2025-12-25/test_wavs/5.wav" + + for f in [ + model, + tokens, + test_wav_0, + test_wav_1, + test_wav_2, + test_wav_3, + test_wav_4, + test_wav_5, + ]: + if not Path(f).is_file(): + print(f"{f} does not exist") + + raise ValueError( + """Please download model files from + https://github.com/k2-fsa/sherpa-onnx/releases/tag/asr-models + """ + ) + return ( + sherpa_onnx.OfflineRecognizer.from_medasr_ctc( + model=model, + tokens=tokens, + num_threads=2, + ), + test_wav_0, + test_wav_1, + test_wav_2, + test_wav_3, + test_wav_4, + test_wav_5, + ) + + +def load_audio(filename): + audio, sample_rate = librosa.load(filename, sr=16000) + assert sample_rate == 16000, sample_rate + + return np.ascontiguousarray(audio) + + +def decode_single_file(recognizer, filename): + samples = load_audio(filename) + + start_time = time.time() + + stream = recognizer.create_stream() + stream.accept_waveform(sample_rate=16000, waveform=samples) + recognizer.decode_stream(stream) + + end_time = time.time() + elapsed_seconds = end_time - start_time + audio_duration = len(samples) / 16000 + real_time_factor = elapsed_seconds / audio_duration + + print("---") + print(filename) + print(stream.result) + print(f"Elapsed seconds: {elapsed_seconds:.3f}") + print(f"Audio duration in seconds: {audio_duration:.3f}") + print(f"RTF: {elapsed_seconds:.3f}/{audio_duration:.3f} = {real_time_factor:.3f}") + print() + + +def decode_multiple_files(recognizer, filenames): + streams = [] + + start_time = time.time() + + audio_duration = 0 + + for filename in filenames: + samples = load_audio(filename) + audio_duration += len(samples) / 16000 + + stream = recognizer.create_stream() + stream.accept_waveform(sample_rate=16000, waveform=samples) + streams.append(stream) + + recognizer.decode_streams(streams) + + end_time = time.time() + elapsed_seconds = end_time - start_time + real_time_factor = elapsed_seconds / audio_duration + + for name, stream in zip(filenames, streams): + print("---") + print(name) + print(stream.result) + print() + + print(f"Elapsed seconds: {elapsed_seconds:.3f}") + print(f"Audio duration in seconds: {audio_duration:.3f}") + print(f"RTF: {elapsed_seconds:.3f}/{audio_duration:.3f} = {real_time_factor:.3f}") + print() + print() + + +def main(): + recognizer, *filenames = create_recognizer() + + decode_single_file(recognizer, filenames[0]) + decode_single_file(recognizer, filenames[1]) + decode_multiple_files(recognizer, filenames[2:]) + + +if __name__ == "__main__": + main() diff --git a/scripts/medasr/export_onnx.py b/scripts/medasr/export_onnx.py index a0fce456b9..2f7bd7e7a5 100755 --- a/scripts/medasr/export_onnx.py +++ b/scripts/medasr/export_onnx.py @@ -107,6 +107,7 @@ def main(): "model_author": "google", "maintainer": "k2-fsa", "vocab_size": processor.tokenizer.vocab_size, + "subsampling_factor": 4, "url": "https://github.com/Google-Health/medasr", "license": "https://developers.google.com/health-ai-developer-foundations/terms", } diff --git a/sherpa-onnx/csrc/CMakeLists.txt b/sherpa-onnx/csrc/CMakeLists.txt index 5c3338fe48..d21995b5b6 100755 --- a/sherpa-onnx/csrc/CMakeLists.txt +++ b/sherpa-onnx/csrc/CMakeLists.txt @@ -40,6 +40,8 @@ set(sources offline-fire-red-asr-model.cc offline-lm-config.cc offline-lm.cc + offline-medasr-ctc-model-config.cc + offline-medasr-ctc-model.cc offline-model-config.cc offline-moonshine-greedy-search-decoder.cc offline-moonshine-model-config.cc diff --git a/sherpa-onnx/csrc/offline-ctc-model.cc b/sherpa-onnx/csrc/offline-ctc-model.cc index 4becf1afb2..4cc5d407f6 100644 --- a/sherpa-onnx/csrc/offline-ctc-model.cc +++ b/sherpa-onnx/csrc/offline-ctc-model.cc @@ -21,6 +21,7 @@ #include "sherpa-onnx/csrc/file-utils.h" #include "sherpa-onnx/csrc/macros.h" #include "sherpa-onnx/csrc/offline-dolphin-model.h" +#include "sherpa-onnx/csrc/offline-medasr-ctc-model.h" #include "sherpa-onnx/csrc/offline-nemo-enc-dec-ctc-model.h" #include "sherpa-onnx/csrc/offline-omnilingual-asr-ctc-model.h" #include "sherpa-onnx/csrc/offline-tdnn-ctc-model.h" @@ -126,6 +127,8 @@ std::unique_ptr OfflineCtcModel::Create( return std::make_unique(config); } else if (!config.omnilingual.model.empty()) { return std::make_unique(config); + } else if (!config.medasr.model.empty()) { + return std::make_unique(config); } // TODO(fangjun): Refactor it. We don't need to use model_type here @@ -192,6 +195,8 @@ std::unique_ptr OfflineCtcModel::Create( return std::make_unique(mgr, config); } else if (!config.omnilingual.model.empty()) { return std::make_unique(mgr, config); + } else if (!config.medasr.model.empty()) { + return std::make_unique(mgr, config); } // TODO(fangjun): Refactor it. We don't need to use model_type here diff --git a/sherpa-onnx/csrc/offline-medasr-ctc-model-config.cc b/sherpa-onnx/csrc/offline-medasr-ctc-model-config.cc new file mode 100644 index 0000000000..1fdac49db9 --- /dev/null +++ b/sherpa-onnx/csrc/offline-medasr-ctc-model-config.cc @@ -0,0 +1,40 @@ +// sherpa-onnx/csrc/offline-medasr-ctc-model-config.cc +// +// Copyright (c) 2025 Xiaomi Corporation + +#include "sherpa-onnx/csrc/offline-medasr-ctc-model-config.h" + +#include +#include + +#include "sherpa-onnx/csrc/file-utils.h" +#include "sherpa-onnx/csrc/macros.h" + +namespace sherpa_onnx { + +void OfflineMedAsrCtcModelConfig::Register(ParseOptions *po) { + po->Register( + "medasr", &model, + "Path to model.onnx from MedASR. Please see " + "https://github.com/k2-fsa/sherpa-onnx/pull/2934 for available models"); +} + +bool OfflineMedAsrCtcModelConfig::Validate() const { + if (!FileExists(model)) { + SHERPA_ONNX_LOGE("MedASR model: '%s' does not exist", model.c_str()); + return false; + } + + return true; +} + +std::string OfflineMedAsrCtcModelConfig::ToString() const { + std::ostringstream os; + + os << "OfflineMedAsrCtcModelConfig("; + os << "model=\"" << model << "\")"; + + return os.str(); +} + +} // namespace sherpa_onnx diff --git a/sherpa-onnx/csrc/offline-medasr-ctc-model-config.h b/sherpa-onnx/csrc/offline-medasr-ctc-model-config.h new file mode 100644 index 0000000000..02629e1b5d --- /dev/null +++ b/sherpa-onnx/csrc/offline-medasr-ctc-model-config.h @@ -0,0 +1,29 @@ +// sherpa-onnx/csrc/offline-medasr-ctc-model-config.h +// +// Copyright (c) 2025 Xiaomi Corporation + +#ifndef SHERPA_ONNX_CSRC_OFFLINE_MEDASR_CTC_MODEL_CONFIG_H_ +#define SHERPA_ONNX_CSRC_OFFLINE_MEDASR_CTC_MODEL_CONFIG_H_ + +#include + +#include "sherpa-onnx/csrc/parse-options.h" + +namespace sherpa_onnx { + +struct OfflineMedAsrCtcModelConfig { + std::string model; + + OfflineMedAsrCtcModelConfig() = default; + explicit OfflineMedAsrCtcModelConfig(const std::string &model) + : model(model) {} + + void Register(ParseOptions *po); + bool Validate() const; + + std::string ToString() const; +}; + +} // namespace sherpa_onnx + +#endif // SHERPA_ONNX_CSRC_OFFLINE_MEDASR_CTC_MODEL_CONFIG_H_ diff --git a/sherpa-onnx/csrc/offline-medasr-ctc-model.cc b/sherpa-onnx/csrc/offline-medasr-ctc-model.cc new file mode 100644 index 0000000000..de1255cf4d --- /dev/null +++ b/sherpa-onnx/csrc/offline-medasr-ctc-model.cc @@ -0,0 +1,198 @@ +// sherpa-onnx/csrc/offline-medasr-ctc-model.cc +// +// Copyright (c) 2025 Xiaomi Corporation + +#include "sherpa-onnx/csrc/offline-medasr-ctc-model.h" + +#include +#include +#include +#include +#include + +#if __ANDROID_API__ >= 9 +#include "android/asset_manager.h" +#include "android/asset_manager_jni.h" +#endif + +#if __OHOS__ +#include "rawfile/raw_file_manager.h" +#endif + +#include "sherpa-onnx/csrc/file-utils.h" +#include "sherpa-onnx/csrc/macros.h" +#include "sherpa-onnx/csrc/onnx-utils.h" +#include "sherpa-onnx/csrc/session.h" +#include "sherpa-onnx/csrc/text-utils.h" + +namespace sherpa_onnx { + +namespace { + +std::vector GetMask(Ort::Value length) { + auto shape = length.GetTensorTypeAndShapeInfo().GetShape(); + if (shape.size() != 1) { + SHERPA_ONNX_LOGE("Invalid length dim %zu", shape.size()); + SHERPA_ONNX_EXIT(-1); + } + + auto batch_size = shape[0]; + + const int64_t *p = length.GetTensorData(); + + int64_t max_len = *std::max_element(p, p + batch_size); + + std::vector ans(batch_size * max_len, 0); + + int64_t *p_mask = ans.data(); + + for (int32_t i = 0; i < batch_size; ++i) { + auto len = p[i]; + std::fill(p_mask, p_mask + len, 1); + + p_mask += max_len; + } + + return ans; +} + +} // namespace + +class OfflineMedAsrCtcModel::Impl { + public: + explicit Impl(const OfflineModelConfig &config) + : config_(config), + env_(ORT_LOGGING_LEVEL_ERROR), + sess_opts_(GetSessionOptions(config)), + allocator_{} { + auto buf = ReadFile(config_.medasr.model); + Init(buf.data(), buf.size()); + } + + template + Impl(Manager *mgr, const OfflineModelConfig &config) + : config_(config), + env_(ORT_LOGGING_LEVEL_ERROR), + sess_opts_(GetSessionOptions(config)), + allocator_{} { + auto buf = ReadFile(mgr, config_.medasr.model); + Init(buf.data(), buf.size()); + } + + std::vector Forward(Ort::Value features, + Ort::Value features_length) { + std::vector mask = GetMask(std::move(features_length)); + + auto memory_info = + Ort::MemoryInfo::CreateCpu(OrtDeviceAllocator, OrtMemTypeDefault); + + std::vector shape = + features.GetTensorTypeAndShapeInfo().GetShape(); + shape.resize(2); + + Ort::Value mask_tensor = Ort::Value::CreateTensor( + memory_info, mask.data(), mask.size(), shape.data(), shape.size()); + + std::array inputs = {std::move(features), + std::move(mask_tensor)}; + + return sess_->Run({}, input_names_ptr_.data(), inputs.data(), inputs.size(), + output_names_ptr_.data(), output_names_ptr_.size()); + } + + int32_t VocabSize() const { return vocab_size_; } + + int32_t SubsamplingFactor() const { return subsampling_factor_; } + + OrtAllocator *Allocator() { return allocator_; } + + private: + void Init(void *model_data, size_t model_data_length) { + sess_ = std::make_unique(env_, model_data, model_data_length, + sess_opts_); + + GetInputNames(sess_.get(), &input_names_, &input_names_ptr_); + + GetOutputNames(sess_.get(), &output_names_, &output_names_ptr_); + + // get meta data + Ort::ModelMetadata meta_data = sess_->GetModelMetadata(); + if (config_.debug) { + std::ostringstream os; + PrintModelMetadata(os, meta_data); +#if __OHOS__ + SHERPA_ONNX_LOGE("%{public}s\n", os.str().c_str()); +#else + SHERPA_ONNX_LOGE("%s\n", os.str().c_str()); +#endif + } + + Ort::AllocatorWithDefaultOptions allocator; // used in the macro below + + std::string model_type; + SHERPA_ONNX_READ_META_DATA_STR(model_type, "model_type"); + if (model_type != "medasr_ctc") { + SHERPA_ONNX_LOGE("Expect model type medasr_ctc. Given: '%s'", + model_type.c_str()); + SHERPA_ONNX_EXIT(-1); + } + + SHERPA_ONNX_READ_META_DATA(vocab_size_, "vocab_size"); + SHERPA_ONNX_READ_META_DATA_WITH_DEFAULT(subsampling_factor_, + "subsampling_factor", 4); + } + + private: + OfflineModelConfig config_; + Ort::Env env_; + Ort::SessionOptions sess_opts_; + Ort::AllocatorWithDefaultOptions allocator_; + + std::unique_ptr sess_; + + std::vector input_names_; + std::vector input_names_ptr_; + + std::vector output_names_; + std::vector output_names_ptr_; + + int32_t vocab_size_ = 0; + int32_t subsampling_factor_ = 0; +}; + +OfflineMedAsrCtcModel::OfflineMedAsrCtcModel(const OfflineModelConfig &config) + : impl_(std::make_unique(config)) {} + +template +OfflineMedAsrCtcModel::OfflineMedAsrCtcModel(Manager *mgr, + const OfflineModelConfig &config) + : impl_(std::make_unique(mgr, config)) {} + +OfflineMedAsrCtcModel::~OfflineMedAsrCtcModel() = default; + +std::vector OfflineMedAsrCtcModel::Forward( + Ort::Value features, Ort::Value features_length) { + return impl_->Forward(std::move(features), std::move(features_length)); +} + +int32_t OfflineMedAsrCtcModel::VocabSize() const { return impl_->VocabSize(); } + +int32_t OfflineMedAsrCtcModel::SubsamplingFactor() const { + return impl_->SubsamplingFactor(); +} + +OrtAllocator *OfflineMedAsrCtcModel::Allocator() const { + return impl_->Allocator(); +} + +#if __ANDROID_API__ >= 9 +template OfflineMedAsrCtcModel::OfflineMedAsrCtcModel( + AAssetManager *mgr, const OfflineModelConfig &config); +#endif + +#if __OHOS__ +template OfflineMedAsrCtcModel::OfflineMedAsrCtcModel( + NativeResourceManager *mgr, const OfflineModelConfig &config); +#endif + +} // namespace sherpa_onnx diff --git a/sherpa-onnx/csrc/offline-medasr-ctc-model.h b/sherpa-onnx/csrc/offline-medasr-ctc-model.h new file mode 100644 index 0000000000..76085effeb --- /dev/null +++ b/sherpa-onnx/csrc/offline-medasr-ctc-model.h @@ -0,0 +1,65 @@ +// sherpa-onnx/csrc/offline-medasr-ctc-model.h +// +// Copyright (c) 2025 Xiaomi Corporation +#ifndef SHERPA_ONNX_CSRC_OFFLINE_MEDASR_CTC_MODEL_H_ +#define SHERPA_ONNX_CSRC_OFFLINE_MEDASR_CTC_MODEL_H_ +#include +#include +#include +#include + +#include "onnxruntime_cxx_api.h" // NOLINT +#include "sherpa-onnx/csrc/offline-ctc-model.h" +#include "sherpa-onnx/csrc/offline-model-config.h" + +namespace sherpa_onnx { + +/** This class implements the CTC model from MedASR. + * + * See + * https://github.com/k2-fsa/sherpa-onnx/blob/master/scripts/medasr/export_onnx.py + * https://github.com/k2-fsa/sherpa-onnx/blob/master/scripts/medasr/test_onnx.py + * https://github.com/k2-fsa/sherpa-onnx/blob/master/scripts/medasr/run.sh + * + */ +class OfflineMedAsrCtcModel : public OfflineCtcModel { + public: + explicit OfflineMedAsrCtcModel(const OfflineModelConfig &config); + + template + OfflineMedAsrCtcModel(Manager *mgr, const OfflineModelConfig &config); + + ~OfflineMedAsrCtcModel() override; + + /** Run the forward method of the model. + * + * @param features A tensor of shape (N, T, C). + * @param features_length A 1-D tensor of shape (N,) containing number of + * valid frames in `features` before padding. + * Its dtype is int64_t. + * + * @return Return a vector containing: + * - log_probs: A 3-D tensor of shape (N, T', vocab_size). + * - log_probs_length A 1-D tensor of shape (N,). Its dtype is int64_t + */ + std::vector Forward(Ort::Value features, + Ort::Value features_length) override; + + /** Return the vocabulary size of the model + */ + int32_t VocabSize() const override; + + int32_t SubsamplingFactor() const override; + + /** Return an allocator for allocating memory + */ + OrtAllocator *Allocator() const override; + + private: + class Impl; + std::unique_ptr impl_; +}; + +} // namespace sherpa_onnx + +#endif // SHERPA_ONNX_CSRC_OFFLINE_MEDASR_CTC_MODEL_H_ diff --git a/sherpa-onnx/csrc/offline-model-config.cc b/sherpa-onnx/csrc/offline-model-config.cc index 49189f9fd1..cae0ea78fc 100644 --- a/sherpa-onnx/csrc/offline-model-config.cc +++ b/sherpa-onnx/csrc/offline-model-config.cc @@ -25,6 +25,7 @@ void OfflineModelConfig::Register(ParseOptions *po) { dolphin.Register(po); canary.Register(po); omnilingual.Register(po); + medasr.Register(po); po->Register("telespeech-ctc", &telespeech_ctc, "Path to model.onnx for telespeech ctc"); @@ -156,6 +157,10 @@ bool OfflineModelConfig::Validate() const { return omnilingual.Validate(); } + if (!medasr.model.empty()) { + return medasr.Validate(); + } + if (!telespeech_ctc.empty() && !FileExists(telespeech_ctc)) { SHERPA_ONNX_LOGE("telespeech_ctc: '%s' does not exist", telespeech_ctc.c_str()); @@ -186,6 +191,7 @@ std::string OfflineModelConfig::ToString() const { os << "dolphin=" << dolphin.ToString() << ", "; os << "canary=" << canary.ToString() << ", "; os << "omnilingual=" << omnilingual.ToString() << ", "; + os << "medasr=" << medasr.ToString() << ", "; os << "telespeech_ctc=\"" << telespeech_ctc << "\", "; os << "tokens=\"" << tokens << "\", "; os << "num_threads=" << num_threads << ", "; diff --git a/sherpa-onnx/csrc/offline-model-config.h b/sherpa-onnx/csrc/offline-model-config.h index 6ef84edc8d..dec5c95ea2 100644 --- a/sherpa-onnx/csrc/offline-model-config.h +++ b/sherpa-onnx/csrc/offline-model-config.h @@ -9,6 +9,7 @@ #include "sherpa-onnx/csrc/offline-canary-model-config.h" #include "sherpa-onnx/csrc/offline-dolphin-model-config.h" #include "sherpa-onnx/csrc/offline-fire-red-asr-model-config.h" +#include "sherpa-onnx/csrc/offline-medasr-ctc-model-config.h" #include "sherpa-onnx/csrc/offline-moonshine-model-config.h" #include "sherpa-onnx/csrc/offline-nemo-enc-dec-ctc-model-config.h" #include "sherpa-onnx/csrc/offline-omnilingual-asr-ctc-model-config.h" @@ -36,6 +37,7 @@ struct OfflineModelConfig { OfflineDolphinModelConfig dolphin; OfflineCanaryModelConfig canary; OfflineOmnilingualAsrCtcModelConfig omnilingual; + OfflineMedAsrCtcModelConfig medasr; std::string telespeech_ctc; std::string tokens; @@ -71,6 +73,7 @@ struct OfflineModelConfig { const OfflineDolphinModelConfig &dolphin, const OfflineCanaryModelConfig &canary, const OfflineOmnilingualAsrCtcModelConfig &omnilingual, + const OfflineMedAsrCtcModelConfig &medasr, const std::string &telespeech_ctc, const std::string &tokens, int32_t num_threads, bool debug, const std::string &provider, const std::string &model_type, @@ -89,6 +92,7 @@ struct OfflineModelConfig { dolphin(dolphin), canary(canary), omnilingual(omnilingual), + medasr(medasr), telespeech_ctc(telespeech_ctc), tokens(tokens), num_threads(num_threads), diff --git a/sherpa-onnx/csrc/offline-recognizer-ctc-impl.h b/sherpa-onnx/csrc/offline-recognizer-ctc-impl.h index 11f7b81d88..93505043f4 100644 --- a/sherpa-onnx/csrc/offline-recognizer-ctc-impl.h +++ b/sherpa-onnx/csrc/offline-recognizer-ctc-impl.h @@ -38,6 +38,11 @@ OfflineRecognitionResult Convert(const OfflineCtcDecoderResult &src, // tdnn models from yesno have a SIL token, we should remove it. continue; } + + if (sym_table.Contains("") && src.tokens[i] == sym_table[""]) { + // Skip for Google MedASR + continue; + } auto sym = sym_table[src.tokens[i]]; text.append(sym); @@ -58,6 +63,10 @@ OfflineRecognitionResult Convert(const OfflineCtcDecoderResult &src, text = sym_table.DecodeByteBpe(text); } + if (!text.empty() && text.front() == ' ') { + text.erase(0, 1); + } + r.text = std::move(text); float frame_shift_s = frame_shift_ms / 1000. * subsampling_factor; @@ -144,6 +153,17 @@ class OfflineRecognizerCtcImpl : public OfflineRecognizerImpl { config_.feat_config.normalize_samples = false; } + if (!config_.model_config.medasr.model.empty()) { + config_.feat_config.low_freq = 125; + config_.feat_config.high_freq = 7500; + config_.feat_config.remove_dc_offset = false; + config_.feat_config.dither = 0; + config_.feat_config.preemph_coeff = 0; + config_.feat_config.window_type = "hanning"; + config_.feat_config.feature_dim = 128; + config_.feat_config.snip_edges = true; + } + config_.feat_config.nemo_normalize_type = model_->FeatureNormalizationMethod(); diff --git a/sherpa-onnx/csrc/offline-recognizer-impl.cc b/sherpa-onnx/csrc/offline-recognizer-impl.cc index 7aabf8d83c..d894f955ca 100644 --- a/sherpa-onnx/csrc/offline-recognizer-impl.cc +++ b/sherpa-onnx/csrc/offline-recognizer-impl.cc @@ -217,6 +217,7 @@ std::unique_ptr OfflineRecognizerImpl::Create( !config.model_config.tdnn.model.empty() || !config.model_config.wenet_ctc.model.empty() || !config.model_config.omnilingual.model.empty() || + !config.model_config.medasr.model.empty() || !config.model_config.dolphin.model.empty()) { return std::make_unique(config); } @@ -547,6 +548,7 @@ std::unique_ptr OfflineRecognizerImpl::Create( !config.model_config.tdnn.model.empty() || !config.model_config.wenet_ctc.model.empty() || !config.model_config.omnilingual.model.empty() || + !config.model_config.medasr.model.empty() || !config.model_config.dolphin.model.empty()) { return std::make_unique(mgr, config); } diff --git a/sherpa-onnx/python/csrc/CMakeLists.txt b/sherpa-onnx/python/csrc/CMakeLists.txt index 0b55ae91f9..69b47c500e 100644 --- a/sherpa-onnx/python/csrc/CMakeLists.txt +++ b/sherpa-onnx/python/csrc/CMakeLists.txt @@ -14,6 +14,7 @@ set(srcs offline-dolphin-model-config.cc offline-fire-red-asr-model-config.cc offline-lm-config.cc + offline-medasr-ctc-model-config.cc offline-model-config.cc offline-moonshine-model-config.cc offline-nemo-enc-dec-ctc-model-config.cc diff --git a/sherpa-onnx/python/csrc/offline-medasr-ctc-model-config.cc b/sherpa-onnx/python/csrc/offline-medasr-ctc-model-config.cc new file mode 100644 index 0000000000..8d13a2efb3 --- /dev/null +++ b/sherpa-onnx/python/csrc/offline-medasr-ctc-model-config.cc @@ -0,0 +1,22 @@ +// sherpa-onnx/python/csrc/offline-medasr-ctc-model-config.cc +// +// Copyright (c) 2025 Xiaomi Corporation + +#include "sherpa-onnx/csrc/offline-medasr-ctc-model-config.h" + +#include +#include + +#include "sherpa-onnx/python/csrc/offline-medasr-ctc-model-config.h" + +namespace sherpa_onnx { + +void PybindOfflineMedAsrCtcModelConfig(py::module *m) { + using PyClass = OfflineMedAsrCtcModelConfig; + py::class_(*m, "OfflineMedAsrCtcModelConfig") + .def(py::init(), py::arg("model")) + .def_readwrite("model", &PyClass::model) + .def("__str__", &PyClass::ToString); +} + +} // namespace sherpa_onnx diff --git a/sherpa-onnx/python/csrc/offline-medasr-ctc-model-config.h b/sherpa-onnx/python/csrc/offline-medasr-ctc-model-config.h new file mode 100644 index 0000000000..f09a78b3c1 --- /dev/null +++ b/sherpa-onnx/python/csrc/offline-medasr-ctc-model-config.h @@ -0,0 +1,16 @@ +// sherpa-onnx/python/csrc/offline-medasr-ctc-model-config.h +// +// Copyright (c) 2023 Xiaomi Corporation + +#ifndef SHERPA_ONNX_PYTHON_CSRC_OFFLINE_MEDASR_CTC_MODEL_CONFIG_H_ +#define SHERPA_ONNX_PYTHON_CSRC_OFFLINE_MEDASR_CTC_MODEL_CONFIG_H_ + +#include "sherpa-onnx/python/csrc/sherpa-onnx.h" + +namespace sherpa_onnx { + +void PybindOfflineMedAsrCtcModelConfig(py::module *m); + +} + +#endif // SHERPA_ONNX_PYTHON_CSRC_OFFLINE_MEDASR_CTC_MODEL_CONFIG_H_ diff --git a/sherpa-onnx/python/csrc/offline-model-config.cc b/sherpa-onnx/python/csrc/offline-model-config.cc index 6c1286b8fd..59fa03ab3d 100644 --- a/sherpa-onnx/python/csrc/offline-model-config.cc +++ b/sherpa-onnx/python/csrc/offline-model-config.cc @@ -11,6 +11,7 @@ #include "sherpa-onnx/python/csrc/offline-canary-model-config.h" #include "sherpa-onnx/python/csrc/offline-dolphin-model-config.h" #include "sherpa-onnx/python/csrc/offline-fire-red-asr-model-config.h" +#include "sherpa-onnx/python/csrc/offline-medasr-ctc-model-config.h" #include "sherpa-onnx/python/csrc/offline-moonshine-model-config.h" #include "sherpa-onnx/python/csrc/offline-nemo-enc-dec-ctc-model-config.h" #include "sherpa-onnx/python/csrc/offline-omnilingual-asr-ctc-model-config.h" @@ -38,6 +39,7 @@ void PybindOfflineModelConfig(py::module *m) { PybindOfflineDolphinModelConfig(m); PybindOfflineCanaryModelConfig(m); PybindOfflineOmnilingualAsrCtcModelConfig(m); + PybindOfflineMedAsrCtcModelConfig(m); using PyClass = OfflineModelConfig; py::class_(*m, "OfflineModelConfig") @@ -54,9 +56,10 @@ void PybindOfflineModelConfig(py::module *m) { const OfflineDolphinModelConfig &, const OfflineCanaryModelConfig &, const OfflineOmnilingualAsrCtcModelConfig &, - const std::string &, const std::string &, int32_t, bool, + const OfflineMedAsrCtcModelConfig &, const std::string &, + const std::string &, int32_t, bool, const std::string &, const std::string &, const std::string &, - const std::string &, const std::string &>(), + const std::string &>(), py::arg("transducer") = OfflineTransducerModelConfig(), py::arg("paraformer") = OfflineParaformerModelConfig(), py::arg("nemo_ctc") = OfflineNemoEncDecCtcModelConfig(), @@ -70,6 +73,7 @@ void PybindOfflineModelConfig(py::module *m) { py::arg("dolphin") = OfflineDolphinModelConfig(), py::arg("canary") = OfflineCanaryModelConfig(), py::arg("omnilingual") = OfflineOmnilingualAsrCtcModelConfig(), + py::arg("medasr") = OfflineMedAsrCtcModelConfig(), py::arg("telespeech_ctc") = "", py::arg("tokens") = "", py::arg("num_threads") = 1, py::arg("debug") = false, py::arg("provider") = "cpu", py::arg("model_type") = "", @@ -87,6 +91,7 @@ void PybindOfflineModelConfig(py::module *m) { .def_readwrite("dolphin", &PyClass::dolphin) .def_readwrite("canary", &PyClass::canary) .def_readwrite("omnilingual", &PyClass::omnilingual) + .def_readwrite("medasr", &PyClass::medasr) .def_readwrite("telespeech_ctc", &PyClass::telespeech_ctc) .def_readwrite("tokens", &PyClass::tokens) .def_readwrite("num_threads", &PyClass::num_threads) diff --git a/sherpa-onnx/python/csrc/offline-wenet-ctc-model-config.cc b/sherpa-onnx/python/csrc/offline-wenet-ctc-model-config.cc index d8bce13a2f..8e8e73c6ae 100644 --- a/sherpa-onnx/python/csrc/offline-wenet-ctc-model-config.cc +++ b/sherpa-onnx/python/csrc/offline-wenet-ctc-model-config.cc @@ -1,4 +1,4 @@ -// sherpa-onnx/python/csrc/offline-wenet-model-config.cc +// sherpa-onnx/python/csrc/offline-wenet-ctc-model-config.cc // // Copyright (c) 2023 Xiaomi Corporation diff --git a/sherpa-onnx/python/csrc/offline-wenet-ctc-model-config.h b/sherpa-onnx/python/csrc/offline-wenet-ctc-model-config.h index ea92c46f18..b9df4f0725 100644 --- a/sherpa-onnx/python/csrc/offline-wenet-ctc-model-config.h +++ b/sherpa-onnx/python/csrc/offline-wenet-ctc-model-config.h @@ -1,4 +1,4 @@ -// sherpa-onnx/python/csrc/offline-wenet-model-config.h +// sherpa-onnx/python/csrc/offline-wenet-ctc-model-config.h // // Copyright (c) 2023 Xiaomi Corporation diff --git a/sherpa-onnx/python/sherpa_onnx/offline_recognizer.py b/sherpa-onnx/python/sherpa_onnx/offline_recognizer.py index 1194c664ba..7149f87a44 100644 --- a/sherpa-onnx/python/sherpa_onnx/offline_recognizer.py +++ b/sherpa-onnx/python/sherpa_onnx/offline_recognizer.py @@ -8,6 +8,7 @@ HomophoneReplacerConfig, OfflineCanaryModelConfig, OfflineOmnilingualAsrCtcModelConfig, + OfflineMedAsrCtcModelConfig, OfflineCtcFstDecoderConfig, OfflineDolphinModelConfig, OfflineFireRedAsrModelConfig, @@ -536,6 +537,56 @@ def from_dolphin_ctc( self.config = recognizer_config return self + @classmethod + def from_medasr_ctc( + cls, + model: str, + tokens: str, + num_threads: int = 1, + decoding_method: str = "greedy_search", + debug: bool = False, + provider: str = "cpu", + ): + """ + Please refer to + ``_ + to download pre-trained models. + + Args: + model: + Path to ``model.onnx``. + tokens: + Path to ``tokens.txt``. Each line in ``tokens.txt`` contains two + columns:: + + symbol integer_id + + num_threads: + Number of threads for neural network computation. + decoding_method: + The only supported decoding method is greedy_search. + debug: + True to show debug messages. + provider: + onnxruntime execution providers. Valid values are: cpu, cuda, coreml. + """ + self = cls.__new__(cls) + model_config = OfflineModelConfig( + medasr=OfflineMedAsrCtcModelConfig(model=model), + tokens=tokens, + num_threads=num_threads, + debug=debug, + provider=provider, + ) + + recognizer_config = OfflineRecognizerConfig( + model_config=model_config, + decoding_method=decoding_method, + ) + self.recognizer = _Recognizer(recognizer_config) + self.config = recognizer_config + return self + @classmethod def from_omnilingual_asr_ctc( cls,