diff --git a/.github/workflows/flutter-android.yaml b/.github/workflows/flutter-android.yaml index 3fa987cfda..1b7d04185a 100644 --- a/.github/workflows/flutter-android.yaml +++ b/.github/workflows/flutter-android.yaml @@ -16,6 +16,7 @@ concurrency: jobs: asr: + if: false name: asr ${{ matrix.index }}/${{ matrix.total }} runs-on: ${{ matrix.os }} strategy: @@ -272,6 +273,7 @@ jobs: git push https://csukuangfj2:$HF_TOKEN@huggingface.co/csukuangfj2/sherpa-onnx-flutter main tts: + if: false name: tts ${{ matrix.index }}/${{ matrix.total }} runs-on: ${{ matrix.os }} strategy: diff --git a/.github/workflows/flutter-linux.yaml b/.github/workflows/flutter-linux.yaml index 571ff4b446..aeb7734ce5 100644 --- a/.github/workflows/flutter-linux.yaml +++ b/.github/workflows/flutter-linux.yaml @@ -22,6 +22,7 @@ env: jobs: asr: + if: false name: asr ${{ matrix.arch }} ${{ matrix.index }}/${{ matrix.total }} runs-on: ${{ matrix.os }} strategy: @@ -151,6 +152,7 @@ jobs: git push https://csukuangfj2:$HF_TOKEN@huggingface.co/csukuangfj2/sherpa-onnx-flutter main tts: + if: false name: tts ${{ matrix.arch }} ${{ matrix.index }}/${{ matrix.total }} runs-on: ${{ matrix.os }} strategy: diff --git a/.github/workflows/flutter-macos.yaml b/.github/workflows/flutter-macos.yaml index 965422078b..424240807f 100644 --- a/.github/workflows/flutter-macos.yaml +++ b/.github/workflows/flutter-macos.yaml @@ -16,6 +16,7 @@ concurrency: jobs: asr: + if: false name: asr ${{ matrix.arch }}/${{ matrix.total }} runs-on: ${{ matrix.os }} strategy: @@ -129,6 +130,7 @@ jobs: git push https://csukuangfj2:$HF_TOKEN@huggingface.co/csukuangfj2/sherpa-onnx-flutter main tts: + if: false name: tts ${{ matrix.arch }} ${{ matrix.index }}/${{ matrix.total }} runs-on: ${{ matrix.os }} strategy: diff --git a/.github/workflows/flutter-windows-x64.yaml b/.github/workflows/flutter-windows-x64.yaml index f2e82c453d..712e73a951 100644 --- a/.github/workflows/flutter-windows-x64.yaml +++ b/.github/workflows/flutter-windows-x64.yaml @@ -16,6 +16,7 @@ concurrency: jobs: asr: + if: false name: asr ${{ matrix.index }}/${{ matrix.total }} runs-on: ${{ matrix.os }} strategy: @@ -121,6 +122,7 @@ jobs: git push https://csukuangfj2:$HF_TOKEN@huggingface.co/csukuangfj2/sherpa-onnx-flutter main tts: + if: false name: tts ${{ matrix.index }}/${{ matrix.total }} runs-on: ${{ matrix.os }} strategy: diff --git a/.github/workflows/release-dart-package.yaml b/.github/workflows/release-dart-package.yaml index 7acde91a28..5469c4391a 100644 --- a/.github/workflows/release-dart-package.yaml +++ b/.github/workflows/release-dart-package.yaml @@ -234,6 +234,18 @@ jobs: - uses: dart-lang/setup-dart@v1 + - name: Zip package for inspection + shell: bash + run: | + cd /tmp + zip -r /tmp/sherpa_onnx_linux.zip to_be_published + ls -lh /tmp/sherpa_onnx_linux.zip + + - uses: actions/upload-artifact@v4 + with: + name: package-sherpa_onnx_linux + path: /tmp/sherpa_onnx_linux.zip + - name: Release shell: bash run: | @@ -285,6 +297,11 @@ jobs: git status + - name: Remove .gitignore so xcframework files are included in package + shell: bash + run: | + rm -fv flutter/sherpa_onnx_macos/macos/.gitignore + - name: Download pre-built xcframework shell: bash run: | @@ -317,6 +334,18 @@ jobs: - uses: dart-lang/setup-dart@v1 + - name: Zip package for inspection + shell: bash + run: | + cd /tmp + zip -r /tmp/sherpa_onnx_macos.zip to_be_published + ls -lh /tmp/sherpa_onnx_macos.zip + + - uses: actions/upload-artifact@v4 + with: + name: package-sherpa_onnx_macos + path: /tmp/sherpa_onnx_macos.zip + - name: Release shell: bash run: | @@ -408,6 +437,18 @@ jobs: - uses: dart-lang/setup-dart@v1 + - name: Zip package for inspection + shell: bash + run: | + cd /tmp + zip -r /tmp/sherpa_onnx_windows.zip to_be_published + ls -lh /tmp/sherpa_onnx_windows.zip + + - uses: actions/upload-artifact@v4 + with: + name: package-sherpa_onnx_windows + path: /tmp/sherpa_onnx_windows.zip + - name: Release shell: bash run: | @@ -504,6 +545,18 @@ jobs: - uses: dart-lang/setup-dart@v1 + - name: Zip package for inspection + shell: bash + run: | + cd /tmp + zip -r /tmp/sherpa_onnx_android_arm64.zip to_be_published + ls -lh /tmp/sherpa_onnx_android_arm64.zip + + - uses: actions/upload-artifact@v4 + with: + name: package-sherpa_onnx_android_arm64 + path: /tmp/sherpa_onnx_android_arm64.zip + - name: Release shell: bash run: | @@ -602,6 +655,18 @@ jobs: - uses: dart-lang/setup-dart@v1 + - name: Zip package for inspection + shell: bash + run: | + cd /tmp + zip -r /tmp/sherpa_onnx_android_armeabi.zip to_be_published + ls -lh /tmp/sherpa_onnx_android_armeabi.zip + + - uses: actions/upload-artifact@v4 + with: + name: package-sherpa_onnx_android_armeabi + path: /tmp/sherpa_onnx_android_armeabi.zip + - name: Release shell: bash run: | @@ -700,6 +765,18 @@ jobs: - uses: dart-lang/setup-dart@v1 + - name: Zip package for inspection + shell: bash + run: | + cd /tmp + zip -r /tmp/sherpa_onnx_android_x86.zip to_be_published + ls -lh /tmp/sherpa_onnx_android_x86.zip + + - uses: actions/upload-artifact@v4 + with: + name: package-sherpa_onnx_android_x86 + path: /tmp/sherpa_onnx_android_x86.zip + - name: Release shell: bash run: | @@ -798,6 +875,18 @@ jobs: - uses: dart-lang/setup-dart@v1 + - name: Zip package for inspection + shell: bash + run: | + cd /tmp + zip -r /tmp/sherpa_onnx_android_x86_64.zip to_be_published + ls -lh /tmp/sherpa_onnx_android_x86_64.zip + + - uses: actions/upload-artifact@v4 + with: + name: package-sherpa_onnx_android_x86_64 + path: /tmp/sherpa_onnx_android_x86_64.zip + - name: Release shell: bash run: | @@ -851,6 +940,11 @@ jobs: git status + - name: Remove .gitignore so xcframework files are included in package + shell: bash + run: | + rm -fv flutter/sherpa_onnx_ios/ios/.gitignore + - name: Download pre-built xcframework shell: bash run: | @@ -880,6 +974,133 @@ jobs: - uses: dart-lang/setup-dart@v1 + - name: Zip package for inspection + shell: bash + run: | + cd /tmp + zip -r /tmp/sherpa_onnx_ios.zip to_be_published + ls -lh /tmp/sherpa_onnx_ios.zip + + - uses: actions/upload-artifact@v4 + with: + name: package-sherpa_onnx_ios + path: /tmp/sherpa_onnx_ios.zip + + - name: Release + shell: bash + run: | + cd /tmp/to_be_published + du -h -d1 . + flutter pub get + flutter pub publish --dry-run + flutter pub publish --force + + sherpa_onnx_web: + # if: false + permissions: + id-token: write # Required for authentication using OIDC + name: sherpa_onnx_web + runs-on: ubuntu-latest + + steps: + - uses: actions/checkout@v4 + + - name: Update version + shell: bash + run: | + ./new-release.sh + git diff . + + - name: ccache + uses: hendrikmuhs/ccache-action@v1.2 + with: + key: release-wasm-web + + - name: Install Emscripten + uses: mymindstorm/setup-emsdk@v14 + with: + version: 4.0.23 + + - name: Build WASM module + shell: bash + run: | + export CMAKE_CXX_COMPILER_LAUNCHER=ccache + export CMAKE_C_COMPILER_LAUNCHER=ccache + export PATH="/usr/lib/ccache:/usr/local/opt/ccache/libexec:$PATH" + ./build-wasm-simd-web.sh + + - name: Copy WASM assets + shell: bash + run: | + cp build-wasm-simd-web/install/bin/wasm/web/sherpa-onnx-wasm-web.js \ + flutter/sherpa_onnx_web/assets/ + cp build-wasm-simd-web/install/bin/wasm/web/sherpa-onnx-wasm-web.wasm \ + flutter/sherpa_onnx_web/assets/ + + # Copy JS wrappers from their source directories (not symlinks). + cp wasm/tts/sherpa-onnx-tts.js flutter/sherpa_onnx_web/assets/ + cp wasm/asr/sherpa-onnx-asr.js flutter/sherpa_onnx_web/assets/ + cp wasm/vad/sherpa-onnx-vad.js flutter/sherpa_onnx_web/assets/ + cp wasm/kws/sherpa-onnx-kws.js flutter/sherpa_onnx_web/assets/ + cp wasm/nodejs/sherpa-onnx-punctuation.js flutter/sherpa_onnx_web/assets/ + cp wasm/speaker-diarization/sherpa-onnx-speaker-diarization.js flutter/sherpa_onnx_web/assets/ + cp wasm/speech-enhancement/sherpa-onnx-speech-enhancement.js flutter/sherpa_onnx_web/assets/ + + ls -lh flutter/sherpa_onnx_web/assets/ + + - name: Remove .gitignore so WASM files are included in package + shell: bash + run: | + rm -f flutter/sherpa_onnx_web/.gitignore + + - name: Fix version + shell: bash + run: | + SHERPA_ONNX_VERSION=$(grep "SHERPA_ONNX_VERSION" ./CMakeLists.txt | cut -d " " -f 2 | cut -d '"' -f 2) + + src_dir=$PWD/flutter/sherpa_onnx_web + pushd $src_dir + v="version: $SHERPA_ONNX_VERSION" + echo "v: $v" + sed -i.bak s"/^version: .*/$v/" ./pubspec.yaml + rm *.bak + git status + git diff + + - name: Copy extra files + shell: bash + run: | + dst=flutter/sherpa_onnx_web + + cp -v LICENSE $dst/ + cp -v CHANGELOG.md $dst/ + + git status + + mv -v flutter/sherpa_onnx_web /tmp/to_be_published + + ls -lh /tmp/to_be_published + + - name: Setup Flutter SDK + uses: flutter-actions/setup-flutter@v3 + with: + channel: stable + version: latest + + - uses: dart-lang/setup-dart@v1 + + - name: Zip package for inspection + shell: bash + run: | + cd /tmp + zip -r /tmp/sherpa_onnx_web.zip to_be_published + ls -lh /tmp/sherpa_onnx_web.zip + + - uses: actions/upload-artifact@v4 + with: + name: package-sherpa_onnx_web + path: /tmp/sherpa_onnx_web.zip + - name: Release shell: bash run: | @@ -890,7 +1111,7 @@ jobs: flutter pub publish --force sherpa_onnx: - needs: [sherpa_onnx_linux, sherpa_onnx_macos, sherpa_onnx_windows, sherpa_onnx_android_arm64, sherpa_onnx_android_armeabi, sherpa_onnx_android_x86, sherpa_onnx_android_x86_64, sherpa_onnx_ios] + needs: [sherpa_onnx_linux, sherpa_onnx_macos, sherpa_onnx_windows, sherpa_onnx_android_arm64, sherpa_onnx_android_armeabi, sherpa_onnx_android_x86, sherpa_onnx_android_x86_64, sherpa_onnx_ios, sherpa_onnx_web] # if: false permissions: id-token: write # Required for authentication using OIDC @@ -944,6 +1165,18 @@ jobs: ls -lh /tmp/to_be_published + - name: Zip package for inspection + shell: bash + run: | + cd /tmp + zip -r /tmp/sherpa_onnx.zip to_be_published + ls -lh /tmp/sherpa_onnx.zip + + - uses: actions/upload-artifact@v4 + with: + name: package-sherpa_onnx + path: /tmp/sherpa_onnx.zip + - name: Release shell: bash run: | diff --git a/.github/workflows/test-flutter-package.yaml b/.github/workflows/test-flutter-package.yaml index 77d2fdea34..cdd7eeb7f3 100644 --- a/.github/workflows/test-flutter-package.yaml +++ b/.github/workflows/test-flutter-package.yaml @@ -350,3 +350,72 @@ jobs: with: name: hello-world-pkg-android path: flutter-examples/hello_world/build/app/outputs/flutter-apk/app-release.apk + + web: + if: github.repository_owner == 'csukuangfj' || github.repository_owner == 'k2-fsa' + name: Web + runs-on: ubuntu-latest + steps: + - uses: actions/checkout@v4 + with: + fetch-depth: 0 + + - name: Update version + shell: bash + run: | + ./new-release.sh + git diff . + + - name: Install Emscripten + uses: mymindstorm/setup-emsdk@v14 + with: + version: 4.0.23 + + - name: Setup Flutter SDK + uses: flutter-actions/setup-flutter@v4 + with: + channel: stable + version: latest + + - name: Display flutter info + shell: bash + run: | + which flutter + which dart + flutter --version + dart --version + flutter doctor + + - name: Build WASM module + shell: bash + run: | + ./build-wasm-simd-web.sh + + - name: Copy WASM assets to web plugin + shell: bash + run: | + ./build-flutter-web-wasm.sh + + - name: Build Flutter web app + shell: bash + run: | + cd flutter-examples/hello_world + + flutter pub get + flutter build web + + ls -lh build/web/ + + - name: Zip web app + shell: bash + run: | + cd flutter-examples/hello_world/build/ + mv web hello-world-pkg-web + zip -r hello-world-pkg-web.zip hello-world-pkg-web/ + ls -lh hello-world-pkg-web.zip + + - name: Upload artifact + uses: actions/upload-artifact@v4 + with: + name: hello-world-pkg-web + path: flutter-examples/hello_world/build/hello-world-pkg-web.zip diff --git a/.github/workflows/test-flutter.yaml b/.github/workflows/test-flutter.yaml index 3c433e0f16..a7568cb450 100644 --- a/.github/workflows/test-flutter.yaml +++ b/.github/workflows/test-flutter.yaml @@ -590,3 +590,82 @@ jobs: with: name: hello-world-android path: flutter-examples/hello_world/build/app/outputs/flutter-apk/app-release.apk + + web: + name: Web + runs-on: ubuntu-latest + steps: + - uses: actions/checkout@v4 + with: + fetch-depth: 0 + + - name: Update version + shell: bash + run: | + ./new-release.sh + git diff . + + - name: Install Emscripten + uses: mymindstorm/setup-emsdk@v14 + with: + version: 4.0.23 + + - name: Setup Flutter SDK + uses: flutter-actions/setup-flutter@v4 + with: + channel: stable + version: latest + + - name: Display flutter info + shell: bash + run: | + which flutter + which dart + flutter --version + dart --version + flutter doctor + + - name: Build WASM module + shell: bash + run: | + ./build-wasm-simd-web.sh + + - name: Copy WASM assets to web plugin + shell: bash + run: | + ./build-flutter-web-wasm.sh + + - name: Use local sherpa_onnx_web + shell: bash + run: | + cd flutter/sherpa_onnx + sed -i.bak 's| sherpa_onnx_web: ^1.13.4| # sherpa_onnx_web: ^1.13.4\n sherpa_onnx_web:\n path: ../sherpa_onnx_web|' pubspec.yaml + rm -f pubspec.yaml.bak + cat pubspec.yaml + + - name: Build Flutter web app + shell: bash + run: | + cd flutter-examples/hello_world + + sed -i.bak 's|sherpa_onnx: ^1.13.4|sherpa_onnx:\n path: ../../flutter/sherpa_onnx|' pubspec.yaml + rm -f pubspec.yaml.bak + + flutter pub get + flutter build web + + ls -lh build/web/ + + - name: Zip web app + shell: bash + run: | + cd flutter-examples/hello_world/build/ + mv web hello-world-web + zip -r hello-world-web.zip hello-world-web/ + ls -lh hello-world-web.zip + + - name: Upload artifact + uses: actions/upload-artifact@v4 + with: + name: hello-world-web + path: flutter-examples/hello_world/build/hello-world-web.zip diff --git a/.gitignore b/.gitignore index 831ae4e12c..30b60b80af 100755 --- a/.gitignore +++ b/.gitignore @@ -193,3 +193,6 @@ sherpa-onnx-spleeter-2stems-fp16 sherpa-onnx-cohere-transcribe-14-lang-int8-2026-04-01 espeak-ng-data sherpa-onnx-nemo* +kokoro-int8-en-v0_19 +vits-inflect-*-v2 +assets-2 diff --git a/CMakeLists.txt b/CMakeLists.txt index 5bd2241659..d8b8261296 100644 --- a/CMakeLists.txt +++ b/CMakeLists.txt @@ -65,6 +65,7 @@ option(SHERPA_ONNX_ENABLE_WASM_KWS "Whether to enable WASM for KWS" OFF) option(SHERPA_ONNX_ENABLE_WASM_VAD "Whether to enable WASM for VAD" OFF) option(SHERPA_ONNX_ENABLE_WASM_VAD_ASR "Whether to enable WASM for VAD+ASR" OFF) option(SHERPA_ONNX_ENABLE_WASM_NODEJS "Whether to enable WASM for NodeJS" OFF) +option(SHERPA_ONNX_ENABLE_WASM_WEB "Whether to enable WASM for Web browser" OFF) option(SHERPA_ONNX_ENABLE_WASM_SPEECH_ENHANCEMENT "Whether to enable WASM for speech enhancement" OFF) option(SHERPA_ONNX_ENABLE_BINARY "Whether to build binaries" ${SUGGEST_BUILD_BINARIES}) option(SHERPA_ONNX_ENABLE_TTS "Whether to build TTS related code" ON) @@ -225,6 +226,7 @@ message(STATUS "SHERPA_ONNX_ENABLE_WASM_KWS ${SHERPA_ONNX_ENABLE_WASM_KWS}") message(STATUS "SHERPA_ONNX_ENABLE_WASM_VAD ${SHERPA_ONNX_ENABLE_WASM_VAD}") message(STATUS "SHERPA_ONNX_ENABLE_WASM_VAD_ASR ${SHERPA_ONNX_ENABLE_WASM_VAD_ASR}") message(STATUS "SHERPA_ONNX_ENABLE_WASM_NODEJS ${SHERPA_ONNX_ENABLE_WASM_NODEJS}") +message(STATUS "SHERPA_ONNX_ENABLE_WASM_WEB ${SHERPA_ONNX_ENABLE_WASM_WEB}") message(STATUS "SHERPA_ONNX_ENABLE_WASM_SPEECH_ENHANCEMENT ${SHERPA_ONNX_ENABLE_WASM_SPEECH_ENHANCEMENT}") message(STATUS "SHERPA_ONNX_ENABLE_BINARY ${SHERPA_ONNX_ENABLE_BINARY}") message(STATUS "SHERPA_ONNX_ENABLE_TTS ${SHERPA_ONNX_ENABLE_TTS}") @@ -330,6 +332,12 @@ if(SHERPA_ONNX_ENABLE_WASM_NODEJS) add_definitions(-DSHERPA_ONNX_ENABLE_WASM_KWS=1) endif() +if(SHERPA_ONNX_ENABLE_WASM_WEB) + if(NOT SHERPA_ONNX_ENABLE_WASM) + message(FATAL_ERROR "Please set SHERPA_ONNX_ENABLE_WASM to ON if you enable WASM for Web") + endif() +endif() + if(SHERPA_ONNX_ENABLE_WASM) add_definitions(-DSHERPA_ONNX_ENABLE_WASM=1) endif() diff --git a/build-flutter-web-wasm.sh b/build-flutter-web-wasm.sh new file mode 100755 index 0000000000..c76b04f668 --- /dev/null +++ b/build-flutter-web-wasm.sh @@ -0,0 +1,37 @@ +#!/usr/bin/env bash +# Copyright (c) 2026 Xiaomi Corporation +# +# Build the WASM module for Flutter web and copy assets to the web plugin. +# JS wrapper files are symlinked (not copied) to avoid code duplication. +# For publishing, see release-dart-package.yaml which replaces symlinks +# with real files. + +set -ex + +SCRIPT_DIR=$(cd "$(dirname "$0")" && pwd) +cd "$SCRIPT_DIR" + +# Build the WASM module +./build-wasm-simd-web.sh + +# Copy WASM module to the Flutter web plugin's assets/ +cp build-wasm-simd-web/install/bin/wasm/web/sherpa-onnx-wasm-web.js \ + flutter/sherpa_onnx_web/assets/ + +cp build-wasm-simd-web/install/bin/wasm/web/sherpa-onnx-wasm-web.wasm \ + flutter/sherpa_onnx_web/assets/ + +# Create symlinks for JS wrappers (avoid code duplication) +cd flutter/sherpa_onnx_web/assets +ln -sf ../../../wasm/asr/sherpa-onnx-asr.js . +ln -sf ../../../wasm/tts/sherpa-onnx-tts.js . +ln -sf ../../../wasm/vad/sherpa-onnx-vad.js . +ln -sf ../../../wasm/kws/sherpa-onnx-kws.js . +ln -sf ../../../wasm/nodejs/sherpa-onnx-punctuation.js . +ln -sf ../../../wasm/speaker-diarization/sherpa-onnx-speaker-diarization.js . +ln -sf ../../../wasm/speech-enhancement/sherpa-onnx-speech-enhancement.js . +cd "$SCRIPT_DIR" + +echo "" +echo "Done! WASM assets copied and JS symlinks created." +echo "You can now run: cd flutter-examples/hello_world && flutter run -d chrome" diff --git a/build-wasm-simd-web.sh b/build-wasm-simd-web.sh new file mode 100755 index 0000000000..a1f806d6fa --- /dev/null +++ b/build-wasm-simd-web.sh @@ -0,0 +1,63 @@ +#!/usr/bin/env bash +# Copyright (c) 2026 Xiaomi Corporation +# +# This script is to build sherpa-onnx for WebAssembly (Web browser) + +set -ex + +if [ x"$EMSCRIPTEN" == x"" ]; then + if ! command -v emcc &> /dev/null; then + echo "Please install emscripten first" + echo "" + echo "You can use the following commands to install it:" + echo "" + echo "git clone https://github.com/emscripten-core/emsdk.git" + echo "cd emsdk" + echo "git pull" + echo "./emsdk install 4.0.23" + echo "./emsdk activate 4.0.23" + echo "source ./emsdk_env.sh" + exit 1 + else + EMSCRIPTEN=$(dirname $(realpath $(which emcc))) + emcc --version + fi +fi + +export EMSCRIPTEN=$EMSCRIPTEN +echo "EMSCRIPTEN: $EMSCRIPTEN" +if [ ! -f $EMSCRIPTEN/cmake/Modules/Platform/Emscripten.cmake ]; then + echo "Cannot find $EMSCRIPTEN/cmake/Modules/Platform/Emscripten.cmake" + echo "Please make sure you have installed emsdk correctly" + echo "Hint: emsdk 4.0.23 is known to work. Other versions may not work" + exit 1 +fi + +mkdir -p build-wasm-simd-web +pushd build-wasm-simd-web + +export SHERPA_ONNX_IS_USING_BUILD_WASM_SH=ON + +cmake \ + -DCMAKE_INSTALL_PREFIX=./install \ + -DCMAKE_BUILD_TYPE=Release \ + -DCMAKE_TOOLCHAIN_FILE=$EMSCRIPTEN/cmake/Modules/Platform/Emscripten.cmake \ + \ + -DSHERPA_ONNX_ENABLE_PYTHON=OFF \ + -DSHERPA_ONNX_ENABLE_TESTS=OFF \ + -DSHERPA_ONNX_ENABLE_CHECK=OFF \ + -DBUILD_SHARED_LIBS=OFF \ + -DSHERPA_ONNX_ENABLE_PORTAUDIO=OFF \ + -DSHERPA_ONNX_ENABLE_JNI=OFF \ + -DSHERPA_ONNX_ENABLE_C_API=ON \ + -DSHERPA_ONNX_ENABLE_WEBSOCKET=OFF \ + -DSHERPA_ONNX_ENABLE_GPU=OFF \ + -DSHERPA_ONNX_ENABLE_WASM=ON \ + -DSHERPA_ONNX_ENABLE_WASM_WEB=ON \ + -DSHERPA_ONNX_ENABLE_BINARY=OFF \ + -DSHERPA_ONNX_LINK_LIBSTDCPP_STATICALLY=OFF \ + .. +make -j3 +make install + +ls -lh install/bin/wasm/web diff --git a/flutter-examples/hello_world/README.md b/flutter-examples/hello_world/README.md index f5b46808c1..9d7e9fb227 100644 --- a/flutter-examples/hello_world/README.md +++ b/flutter-examples/hello_world/README.md @@ -114,6 +114,37 @@ flutter build linux flutter build windows ``` +## Running on Web (Chrome) + +```bash +cd flutter-examples/hello_world +flutter pub get +flutter run -d chrome +``` + +### Build for deployment + +```bash +flutter build web +``` + +The output is in `build/web/`. Serve it with any HTTP server: + +```bash +cd build/web +python3 -m http.server 8080 +``` + +Then start your browser and visit . + +### Notes + +- The web build uses WebAssembly (via Emscripten) to run the sherpa-onnx C API + in the browser. No server-side processing is needed. +- The WASM binary is ~15-20MB as it includes all features (ASR, TTS, VAD, etc.) +- First load may take a few seconds while the WASM module is downloaded and + compiled + ## Running on iOS simulator ### Step 1: List available simulators diff --git a/flutter-examples/hello_world/ios/Flutter/Debug.xcconfig b/flutter-examples/hello_world/ios/Flutter/Debug.xcconfig index 592ceee85b..ec97fc6f30 100644 --- a/flutter-examples/hello_world/ios/Flutter/Debug.xcconfig +++ b/flutter-examples/hello_world/ios/Flutter/Debug.xcconfig @@ -1 +1,2 @@ +#include? "Pods/Target Support Files/Pods-Runner/Pods-Runner.debug.xcconfig" #include "Generated.xcconfig" diff --git a/flutter-examples/hello_world/ios/Flutter/Release.xcconfig b/flutter-examples/hello_world/ios/Flutter/Release.xcconfig index 592ceee85b..c4855bfe20 100644 --- a/flutter-examples/hello_world/ios/Flutter/Release.xcconfig +++ b/flutter-examples/hello_world/ios/Flutter/Release.xcconfig @@ -1 +1,2 @@ +#include? "Pods/Target Support Files/Pods-Runner/Pods-Runner.release.xcconfig" #include "Generated.xcconfig" diff --git a/flutter-examples/hello_world/lib/main.dart b/flutter-examples/hello_world/lib/main.dart index 786a52f5a1..db9768c21d 100644 --- a/flutter-examples/hello_world/lib/main.dart +++ b/flutter-examples/hello_world/lib/main.dart @@ -2,9 +2,14 @@ import 'package:flutter/material.dart'; import 'package:sherpa_onnx/sherpa_onnx.dart'; -void main() { +void main() async { WidgetsFlutterBinding.ensureInitialized(); - initBindings(); + + // Works on all platforms: loads native lib on desktop/mobile, WASM on web. + // IMPORTANT: You must call initBindingsAsync() in every isolate that uses + // sherpa-onnx APIs — including the main isolate and any worker isolates. + await initBindingsAsync(); + runApp(const MyApp()); } diff --git a/flutter-examples/hello_world/macos/Flutter/Flutter-Debug.xcconfig b/flutter-examples/hello_world/macos/Flutter/Flutter-Debug.xcconfig index c2efd0b608..4b81f9b2d2 100644 --- a/flutter-examples/hello_world/macos/Flutter/Flutter-Debug.xcconfig +++ b/flutter-examples/hello_world/macos/Flutter/Flutter-Debug.xcconfig @@ -1 +1,2 @@ +#include? "Pods/Target Support Files/Pods-Runner/Pods-Runner.debug.xcconfig" #include "ephemeral/Flutter-Generated.xcconfig" diff --git a/flutter-examples/hello_world/macos/Flutter/Flutter-Release.xcconfig b/flutter-examples/hello_world/macos/Flutter/Flutter-Release.xcconfig index c2efd0b608..5caa9d1579 100644 --- a/flutter-examples/hello_world/macos/Flutter/Flutter-Release.xcconfig +++ b/flutter-examples/hello_world/macos/Flutter/Flutter-Release.xcconfig @@ -1 +1,2 @@ +#include? "Pods/Target Support Files/Pods-Runner/Pods-Runner.release.xcconfig" #include "ephemeral/Flutter-Generated.xcconfig" diff --git a/flutter-examples/hello_world/macos/Runner.xcodeproj/project.pbxproj b/flutter-examples/hello_world/macos/Runner.xcodeproj/project.pbxproj index 22a58778c7..2601f2f750 100644 --- a/flutter-examples/hello_world/macos/Runner.xcodeproj/project.pbxproj +++ b/flutter-examples/hello_world/macos/Runner.xcodeproj/project.pbxproj @@ -27,7 +27,9 @@ 33CC10F32044A3C60003C045 /* Assets.xcassets in Resources */ = {isa = PBXBuildFile; fileRef = 33CC10F22044A3C60003C045 /* Assets.xcassets */; }; 33CC10F62044A3C60003C045 /* MainMenu.xib in Resources */ = {isa = PBXBuildFile; fileRef = 33CC10F42044A3C60003C045 /* MainMenu.xib */; }; 33CC11132044BFA00003C045 /* MainFlutterWindow.swift in Sources */ = {isa = PBXBuildFile; fileRef = 33CC11122044BFA00003C045 /* MainFlutterWindow.swift */; }; + 46D48026BE7F7A7044F470E3 /* Pods_Runner.framework in Frameworks */ = {isa = PBXBuildFile; fileRef = A42DD1D4F7710EBD4FBA6234 /* Pods_Runner.framework */; }; 78A318202AECB46A00862997 /* FlutterGeneratedPluginSwiftPackage in Frameworks */ = {isa = PBXBuildFile; productRef = 78A3181F2AECB46A00862997 /* FlutterGeneratedPluginSwiftPackage */; }; + E47DAED2ECD26A0CD51F7FE0 /* Pods_RunnerTests.framework in Frameworks */ = {isa = PBXBuildFile; fileRef = 756EFF32029A77AA56984168 /* Pods_RunnerTests.framework */; }; /* End PBXBuildFile section */ /* Begin PBXContainerItemProxy section */ @@ -65,7 +67,7 @@ 331C80D7294CF71000263BE5 /* RunnerTests.swift */ = {isa = PBXFileReference; lastKnownFileType = sourcecode.swift; path = RunnerTests.swift; sourceTree = ""; }; 333000ED22D3DE5D00554162 /* Warnings.xcconfig */ = {isa = PBXFileReference; lastKnownFileType = text.xcconfig; path = Warnings.xcconfig; sourceTree = ""; }; 335BBD1A22A9A15E00E9071D /* GeneratedPluginRegistrant.swift */ = {isa = PBXFileReference; fileEncoding = 4; lastKnownFileType = sourcecode.swift; path = GeneratedPluginRegistrant.swift; sourceTree = ""; }; - 33CC10ED2044A3C60003C045 /* hello_world.app */ = {isa = PBXFileReference; explicitFileType = wrapper.application; includeInIndex = 0; path = "hello_world.app"; sourceTree = BUILT_PRODUCTS_DIR; }; + 33CC10ED2044A3C60003C045 /* hello_world.app */ = {isa = PBXFileReference; explicitFileType = wrapper.application; includeInIndex = 0; path = hello_world.app; sourceTree = BUILT_PRODUCTS_DIR; }; 33CC10F02044A3C60003C045 /* AppDelegate.swift */ = {isa = PBXFileReference; lastKnownFileType = sourcecode.swift; path = AppDelegate.swift; sourceTree = ""; }; 33CC10F22044A3C60003C045 /* Assets.xcassets */ = {isa = PBXFileReference; lastKnownFileType = folder.assetcatalog; name = Assets.xcassets; path = Runner/Assets.xcassets; sourceTree = ""; }; 33CC10F52044A3C60003C045 /* Base */ = {isa = PBXFileReference; lastKnownFileType = file.xib; name = Base; path = Base.lproj/MainMenu.xib; sourceTree = ""; }; @@ -77,9 +79,17 @@ 33E51913231747F40026EE4D /* DebugProfile.entitlements */ = {isa = PBXFileReference; lastKnownFileType = text.plist.entitlements; path = DebugProfile.entitlements; sourceTree = ""; }; 33E51914231749380026EE4D /* Release.entitlements */ = {isa = PBXFileReference; fileEncoding = 4; lastKnownFileType = text.plist.entitlements; path = Release.entitlements; sourceTree = ""; }; 33E5194F232828860026EE4D /* AppInfo.xcconfig */ = {isa = PBXFileReference; lastKnownFileType = text.xcconfig; path = AppInfo.xcconfig; sourceTree = ""; }; + 35D5819F4EDF9ADC485A7FE4 /* Pods-Runner.debug.xcconfig */ = {isa = PBXFileReference; includeInIndex = 1; lastKnownFileType = text.xcconfig; name = "Pods-Runner.debug.xcconfig"; path = "Target Support Files/Pods-Runner/Pods-Runner.debug.xcconfig"; sourceTree = ""; }; + 45A500409F656DD81A0C9518 /* Pods-RunnerTests.profile.xcconfig */ = {isa = PBXFileReference; includeInIndex = 1; lastKnownFileType = text.xcconfig; name = "Pods-RunnerTests.profile.xcconfig"; path = "Target Support Files/Pods-RunnerTests/Pods-RunnerTests.profile.xcconfig"; sourceTree = ""; }; + 5A14EC9C87C06A220919A1F8 /* Pods-RunnerTests.release.xcconfig */ = {isa = PBXFileReference; includeInIndex = 1; lastKnownFileType = text.xcconfig; name = "Pods-RunnerTests.release.xcconfig"; path = "Target Support Files/Pods-RunnerTests/Pods-RunnerTests.release.xcconfig"; sourceTree = ""; }; + 756EFF32029A77AA56984168 /* Pods_RunnerTests.framework */ = {isa = PBXFileReference; explicitFileType = wrapper.framework; includeInIndex = 0; path = Pods_RunnerTests.framework; sourceTree = BUILT_PRODUCTS_DIR; }; 78E0A7A72DC9AD7400C4905E /* FlutterGeneratedPluginSwiftPackage */ = {isa = PBXFileReference; lastKnownFileType = wrapper; name = FlutterGeneratedPluginSwiftPackage; path = ephemeral/Packages/FlutterGeneratedPluginSwiftPackage; sourceTree = ""; }; 7AFA3C8E1D35360C0083082E /* Release.xcconfig */ = {isa = PBXFileReference; lastKnownFileType = text.xcconfig; path = Release.xcconfig; sourceTree = ""; }; 9740EEB21CF90195004384FC /* Debug.xcconfig */ = {isa = PBXFileReference; fileEncoding = 4; lastKnownFileType = text.xcconfig; path = Debug.xcconfig; sourceTree = ""; }; + A42DD1D4F7710EBD4FBA6234 /* Pods_Runner.framework */ = {isa = PBXFileReference; explicitFileType = wrapper.framework; includeInIndex = 0; path = Pods_Runner.framework; sourceTree = BUILT_PRODUCTS_DIR; }; + CE67B695759FBB16AC6572CC /* Pods-Runner.profile.xcconfig */ = {isa = PBXFileReference; includeInIndex = 1; lastKnownFileType = text.xcconfig; name = "Pods-Runner.profile.xcconfig"; path = "Target Support Files/Pods-Runner/Pods-Runner.profile.xcconfig"; sourceTree = ""; }; + CEE73865BA6A731E8763DE4E /* Pods-Runner.release.xcconfig */ = {isa = PBXFileReference; includeInIndex = 1; lastKnownFileType = text.xcconfig; name = "Pods-Runner.release.xcconfig"; path = "Target Support Files/Pods-Runner/Pods-Runner.release.xcconfig"; sourceTree = ""; }; + F42D5F22F76E25BEF1D5E36D /* Pods-RunnerTests.debug.xcconfig */ = {isa = PBXFileReference; includeInIndex = 1; lastKnownFileType = text.xcconfig; name = "Pods-RunnerTests.debug.xcconfig"; path = "Target Support Files/Pods-RunnerTests/Pods-RunnerTests.debug.xcconfig"; sourceTree = ""; }; /* End PBXFileReference section */ /* Begin PBXFrameworksBuildPhase section */ @@ -87,6 +97,7 @@ isa = PBXFrameworksBuildPhase; buildActionMask = 2147483647; files = ( + E47DAED2ECD26A0CD51F7FE0 /* Pods_RunnerTests.framework in Frameworks */, ); runOnlyForDeploymentPostprocessing = 0; }; @@ -95,6 +106,7 @@ buildActionMask = 2147483647; files = ( 78A318202AECB46A00862997 /* FlutterGeneratedPluginSwiftPackage in Frameworks */, + 46D48026BE7F7A7044F470E3 /* Pods_Runner.framework in Frameworks */, ); runOnlyForDeploymentPostprocessing = 0; }; @@ -128,6 +140,7 @@ 331C80D6294CF71000263BE5 /* RunnerTests */, 33CC10EE2044A3C60003C045 /* Products */, D73912EC22F37F3D000D13A0 /* Frameworks */, + 543A97B5A08257836DFCE019 /* Pods */, ); sourceTree = ""; }; @@ -176,9 +189,25 @@ path = Runner; sourceTree = ""; }; + 543A97B5A08257836DFCE019 /* Pods */ = { + isa = PBXGroup; + children = ( + 35D5819F4EDF9ADC485A7FE4 /* Pods-Runner.debug.xcconfig */, + CEE73865BA6A731E8763DE4E /* Pods-Runner.release.xcconfig */, + CE67B695759FBB16AC6572CC /* Pods-Runner.profile.xcconfig */, + F42D5F22F76E25BEF1D5E36D /* Pods-RunnerTests.debug.xcconfig */, + 5A14EC9C87C06A220919A1F8 /* Pods-RunnerTests.release.xcconfig */, + 45A500409F656DD81A0C9518 /* Pods-RunnerTests.profile.xcconfig */, + ); + name = Pods; + path = Pods; + sourceTree = ""; + }; D73912EC22F37F3D000D13A0 /* Frameworks */ = { isa = PBXGroup; children = ( + A42DD1D4F7710EBD4FBA6234 /* Pods_Runner.framework */, + 756EFF32029A77AA56984168 /* Pods_RunnerTests.framework */, ); name = Frameworks; sourceTree = ""; @@ -190,6 +219,7 @@ isa = PBXNativeTarget; buildConfigurationList = 331C80DE294CF71000263BE5 /* Build configuration list for PBXNativeTarget "RunnerTests" */; buildPhases = ( + D0FC17B8CD58C0383A413FCB /* [CP] Check Pods Manifest.lock */, 331C80D1294CF70F00263BE5 /* Sources */, 331C80D2294CF70F00263BE5 /* Frameworks */, 331C80D3294CF70F00263BE5 /* Resources */, @@ -208,6 +238,7 @@ isa = PBXNativeTarget; buildConfigurationList = 33CC10FB2044A3C60003C045 /* Build configuration list for PBXNativeTarget "Runner" */; buildPhases = ( + 01A634E6E6E687EBB3089286 /* [CP] Check Pods Manifest.lock */, 33CC10E92044A3C60003C045 /* Sources */, 33CC10EA2044A3C60003C045 /* Frameworks */, 33CC10EB2044A3C60003C045 /* Resources */, @@ -268,7 +299,7 @@ ); mainGroup = 33CC10E42044A3C60003C045; packageReferences = ( - 781AD8BC2B33823900A9FFBB /* XCLocalSwiftPackageReference "Flutter/ephemeral/Packages/FlutterGeneratedPluginSwiftPackage" */, + 781AD8BC2B33823900A9FFBB /* XCLocalSwiftPackageReference "FlutterGeneratedPluginSwiftPackage" */, ); productRefGroup = 33CC10EE2044A3C60003C045 /* Products */; projectDirPath = ""; @@ -301,6 +332,28 @@ /* End PBXResourcesBuildPhase section */ /* Begin PBXShellScriptBuildPhase section */ + 01A634E6E6E687EBB3089286 /* [CP] Check Pods Manifest.lock */ = { + isa = PBXShellScriptBuildPhase; + buildActionMask = 2147483647; + files = ( + ); + inputFileListPaths = ( + ); + inputPaths = ( + "${PODS_PODFILE_DIR_PATH}/Podfile.lock", + "${PODS_ROOT}/Manifest.lock", + ); + name = "[CP] Check Pods Manifest.lock"; + outputFileListPaths = ( + ); + outputPaths = ( + "$(DERIVED_FILE_DIR)/Pods-Runner-checkManifestLockResult.txt", + ); + runOnlyForDeploymentPostprocessing = 0; + shellPath = /bin/sh; + shellScript = "diff \"${PODS_PODFILE_DIR_PATH}/Podfile.lock\" \"${PODS_ROOT}/Manifest.lock\" > /dev/null\nif [ $? != 0 ] ; then\n # print error to STDERR\n echo \"error: The sandbox is not in sync with the Podfile.lock. Run 'pod install' or update your CocoaPods installation.\" >&2\n exit 1\nfi\n# This output is used by Xcode 'outputs' to avoid re-running this script phase.\necho \"SUCCESS\" > \"${SCRIPT_OUTPUT_FILE_0}\"\n"; + showEnvVarsInLog = 0; + }; 3399D490228B24CF009A79C7 /* ShellScript */ = { isa = PBXShellScriptBuildPhase; alwaysOutOfDate = 1; @@ -339,6 +392,28 @@ shellPath = /bin/sh; shellScript = "\"$FLUTTER_ROOT\"/packages/flutter_tools/bin/macos_assemble.sh && touch Flutter/ephemeral/tripwire"; }; + D0FC17B8CD58C0383A413FCB /* [CP] Check Pods Manifest.lock */ = { + isa = PBXShellScriptBuildPhase; + buildActionMask = 2147483647; + files = ( + ); + inputFileListPaths = ( + ); + inputPaths = ( + "${PODS_PODFILE_DIR_PATH}/Podfile.lock", + "${PODS_ROOT}/Manifest.lock", + ); + name = "[CP] Check Pods Manifest.lock"; + outputFileListPaths = ( + ); + outputPaths = ( + "$(DERIVED_FILE_DIR)/Pods-RunnerTests-checkManifestLockResult.txt", + ); + runOnlyForDeploymentPostprocessing = 0; + shellPath = /bin/sh; + shellScript = "diff \"${PODS_PODFILE_DIR_PATH}/Podfile.lock\" \"${PODS_ROOT}/Manifest.lock\" > /dev/null\nif [ $? != 0 ] ; then\n # print error to STDERR\n echo \"error: The sandbox is not in sync with the Podfile.lock. Run 'pod install' or update your CocoaPods installation.\" >&2\n exit 1\nfi\n# This output is used by Xcode 'outputs' to avoid re-running this script phase.\necho \"SUCCESS\" > \"${SCRIPT_OUTPUT_FILE_0}\"\n"; + showEnvVarsInLog = 0; + }; /* End PBXShellScriptBuildPhase section */ /* Begin PBXSourcesBuildPhase section */ @@ -390,6 +465,7 @@ /* Begin XCBuildConfiguration section */ 331C80DB294CF71000263BE5 /* Debug */ = { isa = XCBuildConfiguration; + baseConfigurationReference = F42D5F22F76E25BEF1D5E36D /* Pods-RunnerTests.debug.xcconfig */; buildSettings = { BUNDLE_LOADER = "$(TEST_HOST)"; CURRENT_PROJECT_VERSION = 1; @@ -404,6 +480,7 @@ }; 331C80DC294CF71000263BE5 /* Release */ = { isa = XCBuildConfiguration; + baseConfigurationReference = 5A14EC9C87C06A220919A1F8 /* Pods-RunnerTests.release.xcconfig */; buildSettings = { BUNDLE_LOADER = "$(TEST_HOST)"; CURRENT_PROJECT_VERSION = 1; @@ -418,6 +495,7 @@ }; 331C80DD294CF71000263BE5 /* Profile */ = { isa = XCBuildConfiguration; + baseConfigurationReference = 45A500409F656DD81A0C9518 /* Pods-RunnerTests.profile.xcconfig */; buildSettings = { BUNDLE_LOADER = "$(TEST_HOST)"; CURRENT_PROJECT_VERSION = 1; @@ -712,7 +790,7 @@ /* End XCConfigurationList section */ /* Begin XCLocalSwiftPackageReference section */ - 781AD8BC2B33823900A9FFBB /* XCLocalSwiftPackageReference "Flutter/ephemeral/Packages/FlutterGeneratedPluginSwiftPackage" */ = { + 781AD8BC2B33823900A9FFBB /* XCLocalSwiftPackageReference "FlutterGeneratedPluginSwiftPackage" */ = { isa = XCLocalSwiftPackageReference; relativePath = Flutter/ephemeral/Packages/FlutterGeneratedPluginSwiftPackage; }; diff --git a/flutter-examples/hello_world/macos/Runner.xcworkspace/contents.xcworkspacedata b/flutter-examples/hello_world/macos/Runner.xcworkspace/contents.xcworkspacedata index 1d526a16ed..21a3cc14c7 100644 --- a/flutter-examples/hello_world/macos/Runner.xcworkspace/contents.xcworkspacedata +++ b/flutter-examples/hello_world/macos/Runner.xcworkspace/contents.xcworkspacedata @@ -4,4 +4,7 @@ + + diff --git a/flutter-examples/tts/.gitignore b/flutter-examples/tts/.gitignore index 29a3a5017f..79c113f9b5 100644 --- a/flutter-examples/tts/.gitignore +++ b/flutter-examples/tts/.gitignore @@ -5,9 +5,11 @@ *.swp .DS_Store .atom/ +.build/ .buildlog/ .history .svn/ +.swiftpm/ migrate_working_dir/ # IntelliJ related diff --git a/flutter-examples/tts/README.md b/flutter-examples/tts/README.md index 3eb3e7dd46..3a75bd4aa1 100644 --- a/flutter-examples/tts/README.md +++ b/flutter-examples/tts/README.md @@ -9,6 +9,7 @@ It works on the following platforms: - Linux - macOS (both arm64 and x86_64 are supported) - Windows + - Web Pre-built APPs for this folder can be found at @@ -43,9 +44,9 @@ Then please do the following: ```bash cd flutter-examples/tts/assets -wget https://github.com/k2-fsa/sherpa-onnx/releases/download/tts-models/vits-piper-en_US-libritts_r-medium.tar.bz2 -tar xf vits-piper-en_US-libritts_r-medium.tar.bz2 -rm vits-piper-en_US-libritts_r-medium.tar.bz2 +wget https://github.com/k2-fsa/sherpa-onnx/releases/download/tts-models/vits-piper-en_US-amy-low.tar.bz2 +tar xf vits-piper-en_US-amy-low.tar.bz2 +rm vits-piper-en_US-amy-low.tar.bz2 cd .. ./generate-asset-list.py @@ -56,16 +57,11 @@ cd .. - 2. Change the code to use the downloaded model. - We have given several examples for different models in [./lib/model.dart](./lib/model.dart). - For our selected model, we need to change [./lib/model.dart](./lib/model.dart) so that it looks like below: + We have given several examples for different models in [./lib/model_config.dart](./lib/model_config.dart). + For our selected model, we need to change [./lib/model_config.dart](./lib/model_config.dart) so that it looks like below: ``` -// Example 6 -// https://github.com/k2-fsa/sherpa-onnx/releases/tag/tts-models -// https://github.com/k2-fsa/sherpa-onnx/releases/download/tts-models/vits-piper-en_US-libritts_r-medium.tar.bz2 -modelDir = 'vits-piper-en_US-libritts_r-medium'; -modelName = 'en_US-libritts_r-medium.onnx'; -dataDir = 'vits-piper-en_US-libritts_r-medium/espeak-ng-data'; +const int selectedModelIndex = 0; ``` - 3. That's it. @@ -114,7 +110,23 @@ flutter build windows flutter build apk --split-per-abi ``` - - 5. For iOS + - 5. For web + +``` +flutter run -d chrome +``` + +Also, you can use + +``` +flutter build web +cd build/web +python3 -m http.server 6006 +``` +and then start your browser and access +or you can serve it with any http server. + + - 6. For iOS First, connect your iPhone to your computer and use `flutter devices` to show available devices. You will see something like below: diff --git a/flutter-examples/tts/generate-asset-list.py b/flutter-examples/tts/generate-asset-list.py index f04be92af7..9eadab8665 100755 --- a/flutter-examples/tts/generate-asset-list.py +++ b/flutter-examples/tts/generate-asset-list.py @@ -14,20 +14,27 @@ def main(): target = "./assets/" space = " " - subfolders = [] + entries = [] patterns_to_skip = ["1.5x", "2.x", "3.x", "4.x"] + has_root_files = False for root, dirs, files in os.walk(target): + # If there are files directly in assets/, add the directory itself. + if root == target: + has_root_files = any(not f.startswith('.') for f in files) + if has_root_files: + entries.append("{space}- assets/".format(space=space)) for d in dirs: path = os.path.join(root, d).replace("\\", "/") if os.listdir(path): path = path.lstrip('./') if any(path.endswith(pattern) for pattern in patterns_to_skip): continue - subfolders.append("{space}- {path}/".format(space=space, path=path)) + entries.append("{space}- {path}/".format(space=space, path=path)) - assert subfolders, "The subfolders list is empty." - - subfolders = sorted(subfolders) + if not entries: + print("Warning: no assets found in ./assets/. " + "Add model files and run this script again.") + entries = sorted(entries) loc_of_flutter = -1 loc_of_flutter_asset = -1 @@ -75,24 +82,24 @@ def main(): f.write(line) if index + 1 == loc_of_flutter_asset: f.write(" assets:\n") - for folder in subfolders: - f.write("{folder}\n".format(folder=folder)) + for entry in entries: + f.write("{entry}\n".format(entry=entry)) else: if index + 1 < loc_of_end_flutter or index + 1 > loc_of_end_flutter: f.write(line) if index + 1 == loc_of_end_flutter: f.write(" assets:\n") - for indexOfFolder, folder in enumerate(subfolders): - f.write("{folder}\n".format(folder=folder)) - if indexOfFolder == len(subfolders) - 1: + for indexOfEntry, entry in enumerate(entries): + f.write("{entry}\n".format(entry=entry)) + if indexOfEntry == len(entries) - 1: f.write("\n") break if loc_of_end_flutter == len(lines) + 1: f.write("\n") f.write(" assets:\n") - for folder in subfolders: - f.write("{folder}\n".format(folder=folder)) + for entry in entries: + f.write("{entry}\n".format(entry=entry)) if __name__ == "__main__": main() diff --git a/flutter-examples/tts/ios/Flutter/Debug.xcconfig b/flutter-examples/tts/ios/Flutter/Debug.xcconfig index 592ceee85b..ec97fc6f30 100644 --- a/flutter-examples/tts/ios/Flutter/Debug.xcconfig +++ b/flutter-examples/tts/ios/Flutter/Debug.xcconfig @@ -1 +1,2 @@ +#include? "Pods/Target Support Files/Pods-Runner/Pods-Runner.debug.xcconfig" #include "Generated.xcconfig" diff --git a/flutter-examples/tts/ios/Flutter/Release.xcconfig b/flutter-examples/tts/ios/Flutter/Release.xcconfig index 592ceee85b..c4855bfe20 100644 --- a/flutter-examples/tts/ios/Flutter/Release.xcconfig +++ b/flutter-examples/tts/ios/Flutter/Release.xcconfig @@ -1 +1,2 @@ +#include? "Pods/Target Support Files/Pods-Runner/Pods-Runner.release.xcconfig" #include "Generated.xcconfig" diff --git a/flutter-examples/tts/lib/audio_list.dart b/flutter-examples/tts/lib/audio_list.dart new file mode 100644 index 0000000000..400ea9a4d9 --- /dev/null +++ b/flutter-examples/tts/lib/audio_list.dart @@ -0,0 +1,109 @@ +// Copyright (c) 2026 Xiaomi Corporation +import 'package:flutter/foundation.dart' show kIsWeb; +import 'package:flutter/material.dart'; +import 'package:audioplayers/audioplayers.dart'; + +import './generated_audio.dart'; +import './web_audio.dart' if (dart.library.io) './web_audio_stub.dart' + as web_audio; + +/// Displays a list of generated audio items with play, stop, download, and save. +class AudioList extends StatelessWidget { + final List items; + final AudioPlayer? player; + final void Function(GeneratedAudioItem item, int index) onSaveAs; + + const AudioList({ + super.key, + required this.items, + required this.player, + required this.onSaveAs, + }); + + @override + Widget build(BuildContext context) { + if (items.isEmpty) return const SizedBox.shrink(); + + return ListView.builder( + itemCount: items.length, + itemBuilder: (context, index) { + final item = items[index]; + final idx = items.length - index; + return Card( + margin: const EdgeInsets.symmetric(vertical: 4), + child: Padding( + padding: const EdgeInsets.symmetric(horizontal: 8, vertical: 4), + child: Row( + children: [ + // Play button. + IconButton( + icon: const Icon(Icons.play_arrow), + tooltip: 'Play', + onPressed: () => _play(item), + ), + // Stop button. + IconButton( + icon: const Icon(Icons.stop), + tooltip: 'Stop', + onPressed: () => _stop(), + ), + // Label and duration. + Expanded( + child: Column( + crossAxisAlignment: CrossAxisAlignment.start, + mainAxisSize: MainAxisSize.min, + children: [ + Text( + '$idx. ${item.label}', + maxLines: 1, + overflow: TextOverflow.ellipsis, + ), + Text( + '${item.duration.toStringAsPrecision(3)}s | ' + 'RTF ${item.duration > 0 ? (item.elapsed / item.duration).toStringAsPrecision(3) : '-'}', + style: Theme.of(context).textTheme.bodySmall, + ), + ], + ), + ), + // Download (web only). + if (kIsWeb) + IconButton( + icon: const Icon(Icons.download), + tooltip: 'Download', + onPressed: () { + web_audio.downloadWavBytes( + item.wavBytes!, '$idx-${item.label}.wav'); + }, + ), + // Save as. + IconButton( + icon: const Icon(Icons.save_as), + tooltip: 'Save as', + onPressed: () => onSaveAs(item, idx), + ), + ], + ), + ), + ); + }, + ); + } + + Future _play(GeneratedAudioItem item) async { + if (kIsWeb) { + web_audio.playWavBytes(item.wavBytes!); + } else { + await player?.stop(); + await player?.play(DeviceFileSource(item.filePath!)); + } + } + + Future _stop() async { + if (kIsWeb) { + web_audio.stopPlayback(); + } else { + await player?.stop(); + } + } +} diff --git a/flutter-examples/tts/lib/generated_audio.dart b/flutter-examples/tts/lib/generated_audio.dart new file mode 100644 index 0000000000..450287c133 --- /dev/null +++ b/flutter-examples/tts/lib/generated_audio.dart @@ -0,0 +1,65 @@ +// Copyright (c) 2026 Xiaomi Corporation +import 'dart:typed_data'; + +/// Represents a generated audio item with metadata. +class GeneratedAudioItem { + /// Short label derived from the input text. + final String label; + + /// WAV file bytes (non-null on web, null on native where file is on disk). + final Uint8List? wavBytes; + + /// File path on disk (non-null on native, null on web). + final String? filePath; + + /// Duration of the generated audio in seconds. + final double duration; + + /// Time taken to generate the audio in seconds. + final double elapsed; + + /// Sample rate of the audio. + final int sampleRate; + + /// Generation ID to distinguish from previous generations. + final int generationId; + + GeneratedAudioItem({ + required this.label, + this.wavBytes, + this.filePath, + required this.duration, + required this.elapsed, + required this.sampleRate, + this.generationId = 0, + }); + + /// Create a label from input text (first 30 characters). + static String makeLabel(String text) { + final trimmed = text.trim(); + if (trimmed.length <= 30) return trimmed; + return '${trimmed.substring(0, 27)}...'; + } +} + +/// A chunk of audio samples received during streaming generation. +class AudioChunk { + /// PCM audio samples (Float32, mono). + final Float32List samples; + + /// Progress of generation (0.0 to 1.0). + final double progress; + + /// Sample rate of the audio. + final int sampleRate; + + /// Generation ID to distinguish from previous generations. + final int generationId; + + AudioChunk({ + required this.samples, + required this.progress, + required this.sampleRate, + this.generationId = 0, + }); +} diff --git a/flutter-examples/tts/lib/info.dart b/flutter-examples/tts/lib/info.dart deleted file mode 100644 index e1b3b19e0c..0000000000 --- a/flutter-examples/tts/lib/info.dart +++ /dev/null @@ -1,40 +0,0 @@ -// Copyright (c) 2024 Xiaomi Corporation -import 'package:flutter/material.dart'; -import 'package:url_launcher/url_launcher.dart'; - -class InfoScreen extends StatelessWidget { - @override - Widget build(BuildContext context) { - const double height = 20; - return Container( - child: Padding( - padding: const EdgeInsets.all(8.0), - child: Column( - crossAxisAlignment: CrossAxisAlignment.start, - children: [ - Text('Everything is open-sourced.'), - SizedBox(height: height), - InkWell( - child: Text('Code: https://github.com/k2-fsa/sherpa-onnx'), - onTap: () => launch('https://k2-fsa.github.io/sherpa/onnx/'), - ), - SizedBox(height: height), - InkWell( - child: Text('Doc: https://k2-fsa.github.io/sherpa/onnx/'), - onTap: () => launch('https://k2-fsa.github.io/sherpa/onnx/'), - ), - SizedBox(height: height), - Text('QQ 群: 744602236'), - SizedBox(height: height), - InkWell( - child: Text( - '微信群: https://k2-fsa.github.io/sherpa/social-groups.html'), - onTap: () => - launch('https://k2-fsa.github.io/sherpa/social-groups.html'), - ), - ], - ), - ), - ); - } -} diff --git a/flutter-examples/tts/lib/isolate_tts.dart b/flutter-examples/tts/lib/isolate_tts.dart deleted file mode 100644 index e087a3a9dd..0000000000 --- a/flutter-examples/tts/lib/isolate_tts.dart +++ /dev/null @@ -1,251 +0,0 @@ -import 'dart:io'; -import 'dart:isolate'; - -import 'package:flutter/material.dart'; -import 'package:flutter/services.dart'; -import 'package:media_kit/media_kit.dart'; -import 'package:path/path.dart' as p; -import 'package:path_provider/path_provider.dart'; -import 'package:sherpa_onnx/sherpa_onnx.dart' as sherpa_onnx; - -import 'utils.dart'; - -class _IsolateTask { - final SendPort sendPort; - - RootIsolateToken? rootIsolateToken; - - _IsolateTask(this.sendPort, this.rootIsolateToken); -} - -class _PortModel { - final String method; - - final SendPort? sendPort; - dynamic data; - - _PortModel({ - required this.method, - this.sendPort, - this.data, - }); -} - -class _TtsManager { - /// 主进程通信端口 - final ReceivePort receivePort; - - final Isolate isolate; - - final SendPort isolatePort; - - _TtsManager({ - required this.receivePort, - required this.isolate, - required this.isolatePort, - }); -} - -class IsolateTts { - static late final _TtsManager _ttsManager; - - /// 获取线程里的通信端口 - static SendPort get _sendPort => _ttsManager.isolatePort; - - static late sherpa_onnx.OfflineTts _tts; - - static late Player _player; - - static Future init() async { - ReceivePort port = ReceivePort(); - RootIsolateToken? rootIsolateToken = RootIsolateToken.instance; - - Isolate isolate = await Isolate.spawn( - _isolateEntry, - _IsolateTask(port.sendPort, rootIsolateToken), - errorsAreFatal: false, - ); - port.listen((msg) async { - if (msg is SendPort) { - print(11); - _ttsManager = - _TtsManager(receivePort: port, isolate: isolate, isolatePort: msg); - return; - } - }); - } - - static Future _isolateEntry(_IsolateTask task) async { - if (task.rootIsolateToken != null) { - BackgroundIsolateBinaryMessenger.ensureInitialized( - task.rootIsolateToken!); - } - MediaKit.ensureInitialized(); - _player = Player(); - sherpa_onnx.initBindings(); - final receivePort = ReceivePort(); - task.sendPort.send(receivePort.sendPort); - - String modelDir = ''; - String modelName = ''; - String ruleFsts = ''; - String ruleFars = ''; - String lexicon = ''; - String dataDir = ''; - - // Example 7 - // https://github.com/k2-fsa/sherpa-onnx/releases/tag/tts-models - // https://github.com/k2-fsa/sherpa-onnx/releases/download/tts-models/vits-melo-tts-zh_en.tar.bz2 - modelDir = 'vits-melo-tts-zh_en'; - modelName = 'model.onnx'; - lexicon = 'lexicon.txt'; - - if (modelName == '') { - throw Exception( - 'You are supposed to select a model by changing the code before you run the app'); - } - - final Directory directory = await getApplicationSupportDirectory(); - modelName = p.join(directory.path, modelDir, modelName); - - if (ruleFsts != '') { - final all = ruleFsts.split(','); - var tmp = []; - for (final f in all) { - tmp.add(p.join(directory.path, f)); - } - ruleFsts = tmp.join(','); - } - - if (ruleFars != '') { - final all = ruleFars.split(','); - var tmp = []; - for (final f in all) { - tmp.add(p.join(directory.path, f)); - } - ruleFars = tmp.join(','); - } - - if (lexicon != '') { - lexicon = p.join(directory.path, modelDir, lexicon); - } - - if (dataDir != '') { - dataDir = p.join(directory.path, dataDir); - } - - final tokens = p.join(directory.path, modelDir, 'tokens.txt'); - - final vits = sherpa_onnx.OfflineTtsVitsModelConfig( - model: modelName, - lexicon: lexicon, - tokens: tokens, - dataDir: dataDir, - ); - - final modelConfig = sherpa_onnx.OfflineTtsModelConfig( - vits: vits, - numThreads: 2, - debug: true, - provider: 'cpu', - ); - - final config = sherpa_onnx.OfflineTtsConfig( - model: modelConfig, - ruleFsts: ruleFsts, - ruleFars: ruleFars, - maxNumSenetences: 1, - ); - // print(config); - receivePort.listen((msg) async { - print(msg); - if (msg is _PortModel) { - switch (msg.method) { - case 'generate': - { - _PortModel _v = msg; - final stopwatch = Stopwatch(); - stopwatch.start(); - final genConfig = sherpa_onnx.OfflineTtsGenerationConfig( - sid: _v.data['sid'], - speed: _v.data['speed'], - silenceScale: 0.2, - ); - final audio = - _tts.generateWithConfig(text: _v.data['text'], config: genConfig); - final suffix = - '-sid-${_v.data['sid']}-speed-${_v.data['speed'].toStringAsPrecision(2)}'; - final filename = await generateWaveFilename(suffix); - - final ok = sherpa_onnx.writeWave( - filename: filename, - samples: audio.samples, - sampleRate: audio.sampleRate, - ); - - if (ok) { - stopwatch.stop(); - double elapsed = stopwatch.elapsed.inMilliseconds.toDouble(); - - double waveDuration = audio.samples.length.toDouble() / - audio.sampleRate.toDouble(); - - print('Saved to\n$filename\n' - 'Elapsed: ${(elapsed / 1000).toStringAsPrecision(4)} s\n' - 'Wave duration: ${waveDuration.toStringAsPrecision(4)} s\n' - 'RTF: ${(elapsed / 1000).toStringAsPrecision(4)}/${waveDuration.toStringAsPrecision(4)} ' - '= ${(elapsed / 1000 / waveDuration).toStringAsPrecision(3)} '); - - await _player.open(Media('file:///$filename')); - await _player.play(); - } - } - break; - } - } - }); - _tts = sherpa_onnx.OfflineTts(config); - } - - static Future generate( - {required String text, int sid = 0, double speed = 1.0}) async { - ReceivePort receivePort = ReceivePort(); - _sendPort.send(_PortModel( - method: 'generate', - data: {'text': text, 'sid': sid, 'speed': speed}, - sendPort: receivePort.sendPort, - )); - await receivePort.first; - receivePort.close(); - } -} - -/// 这里是页面 -class IsolateTtsView extends StatefulWidget { - const IsolateTtsView({super.key}); - - @override - State createState() => _IsolateTtsViewState(); -} - -class _IsolateTtsViewState extends State { - @override - void initState() { - super.initState(); - IsolateTts.init(); - } - - @override - Widget build(BuildContext context) { - return Scaffold( - body: Center( - child: ElevatedButton( - onPressed: () { - IsolateTts.generate(text: '这是已退出的 isolate TTS'); - }, - child: Text('Isolate TTS'), - ), - ), - ); - } -} diff --git a/flutter-examples/tts/lib/main.dart b/flutter-examples/tts/lib/main.dart index 78042254ab..699c7f5f22 100644 --- a/flutter-examples/tts/lib/main.dart +++ b/flutter-examples/tts/lib/main.dart @@ -1,9 +1,9 @@ // Copyright (c) 2024 Xiaomi Corporation import 'package:flutter/material.dart'; +import 'package:url_launcher/url_launcher.dart'; -import './info.dart'; -import './tts.dart'; -import 'isolate_tts.dart'; +import './tts_screen.dart'; +import './model_config.dart' show selectedModelDir, selectedModelUrl; void main() { runApp(const MyApp()); @@ -15,61 +15,223 @@ class MyApp extends StatelessWidget { @override Widget build(BuildContext context) { return MaterialApp( - title: 'Next-gen Kaldi flutter demo', + title: 'sherpa-onnx TTS Demo', theme: ThemeData( colorScheme: ColorScheme.fromSeed(seedColor: Colors.deepPurple), useMaterial3: true, ), - home: const MyHomePage(title: 'Next-gen Kaldi with Flutter'), + home: const HomePage(), ); } } -class MyHomePage extends StatefulWidget { - const MyHomePage({super.key, required this.title}); - - final String title; +class HomePage extends StatefulWidget { + const HomePage({super.key}); @override - State createState() => _MyHomePageState(); + State createState() => _HomePageState(); } -class _MyHomePageState extends State { +class _HomePageState extends State { int _currentIndex = 0; - final List _tabs = [ - TtsScreen(), - InfoScreen(), - IsolateTtsView(), - ]; + @override Widget build(BuildContext context) { return Scaffold( - appBar: AppBar( - title: Text(widget.title), + body: IndexedStack( + index: _currentIndex, + children: const [ + TtsScreen(), + InfoScreen(), + ], ), - body: _tabs[_currentIndex], bottomNavigationBar: BottomNavigationBar( currentIndex: _currentIndex, - onTap: (int index) { - setState(() { - _currentIndex = index; - }); - }, - items: [ - BottomNavigationBarItem( - icon: Icon(Icons.home), - label: 'Home', + onTap: (i) => setState(() => _currentIndex = i), + items: const [ + BottomNavigationBarItem(icon: Icon(Icons.home), label: 'TTS'), + BottomNavigationBarItem(icon: Icon(Icons.info), label: 'Info'), + ], + ), + ); + } +} + +class InfoScreen extends StatelessWidget { + const InfoScreen({super.key}); + + @override + Widget build(BuildContext context) { + final theme = Theme.of(context); + final linkStyle = TextStyle( + color: theme.colorScheme.primary, + fontSize: 13, + ); + + return Scaffold( + appBar: AppBar(title: const Text('Info')), + body: ListView( + padding: const EdgeInsets.all(16), + children: [ + // ── Current model card ── + Card( + child: Padding( + padding: const EdgeInsets.all(16), + child: Column( + crossAxisAlignment: CrossAxisAlignment.start, + children: [ + Row( + children: [ + Icon(Icons.record_voice_over, + color: theme.colorScheme.primary), + const SizedBox(width: 8), + Text('Current Model', + style: theme.textTheme.titleMedium + ?.copyWith(fontWeight: FontWeight.bold)), + ], + ), + const Divider(), + Text(selectedModelDir, + style: theme.textTheme.bodyLarge + ?.copyWith(fontWeight: FontWeight.w600)), + const SizedBox(height: 8), + _LinkRow( + icon: Icons.download, + label: 'Download', + url: selectedModelUrl, + style: linkStyle, + ), + ], + ), + ), ), - BottomNavigationBarItem( - icon: Icon(Icons.info), - label: 'Info', + + const SizedBox(height: 12), + + // ── Resources card ── + Card( + child: Padding( + padding: const EdgeInsets.all(16), + child: Column( + crossAxisAlignment: CrossAxisAlignment.start, + children: [ + Row( + children: [ + Icon(Icons.link, color: theme.colorScheme.primary), + const SizedBox(width: 8), + Text('Resources', + style: theme.textTheme.titleMedium + ?.copyWith(fontWeight: FontWeight.bold)), + ], + ), + const Divider(), + _LinkRow( + icon: Icons.code, + label: 'GitHub', + url: 'https://github.com/k2-fsa/sherpa-onnx', + style: linkStyle, + ), + _LinkRow( + icon: Icons.menu_book, + label: 'Documentation', + url: 'https://k2-fsa.github.io/sherpa/onnx/', + style: linkStyle, + ), + _LinkRow( + icon: Icons.surround_sound, + label: 'All TTS Models', + url: 'https://k2-fsa.github.io/sherpa/onnx/tts/all/', + style: linkStyle, + ), + _LinkRow( + icon: Icons.cloud_download, + label: 'TTS Model Releases', + url: 'https://github.com/k2-fsa/sherpa-onnx/releases/tag/tts-models', + style: linkStyle, + ), + _LinkRow( + icon: Icons.play_circle_outline, + label: 'Online Demo', + url: 'https://huggingface.co/spaces/k2-fsa/text-to-speech', + style: linkStyle, + ), + ], + ), + ), ), - BottomNavigationBarItem( - icon: Icon(Icons.multiline_chart), - label: 'isolate', + + const SizedBox(height: 16), + + Center( + child: GestureDetector( + onTap: () => launchUrl(Uri.parse('https://github.com/k2-fsa/sherpa-onnx')), + child: Text.rich( + TextSpan( + text: 'Powered by ', + style: theme.textTheme.bodySmall + ?.copyWith(color: theme.colorScheme.outline), + children: [ + TextSpan( + text: 'sherpa-onnx', + style: theme.textTheme.bodySmall?.copyWith( + color: theme.colorScheme.primary, + decoration: TextDecoration.underline, + ), + ), + ], + ), + ), + ), ), ], ), ); } } + +/// A tappable row with icon, label, and URL. +class _LinkRow extends StatelessWidget { + final IconData icon; + final String label; + final String url; + final TextStyle style; + + const _LinkRow({ + required this.icon, + required this.label, + required this.url, + required this.style, + }); + + @override + Widget build(BuildContext context) { + return Padding( + padding: const EdgeInsets.symmetric(vertical: 4), + child: InkWell( + onTap: () => launchUrl(Uri.parse(url)), + borderRadius: BorderRadius.circular(8), + child: Padding( + padding: const EdgeInsets.symmetric(vertical: 6, horizontal: 4), + child: Row( + children: [ + Icon(icon, size: 18, color: style.color), + const SizedBox(width: 10), + Expanded( + child: Column( + crossAxisAlignment: CrossAxisAlignment.start, + children: [ + Text(label, style: style.copyWith(fontSize: 14)), + Text(url, + style: style.copyWith(fontSize: 11, color: Colors.grey), + overflow: TextOverflow.ellipsis), + ], + ), + ), + Icon(Icons.open_in_new, size: 14, color: style.color), + ], + ), + ), + ), + ); + } +} diff --git a/flutter-examples/tts/lib/model.dart b/flutter-examples/tts/lib/model.dart index 80be43376b..c9524a0ac4 100644 --- a/flutter-examples/tts/lib/model.dart +++ b/flutter-examples/tts/lib/model.dart @@ -1,214 +1,136 @@ // Copyright (c) 2024 Xiaomi Corporation - import "dart:io"; import 'package:flutter/services.dart'; import 'package:path_provider/path_provider.dart'; import 'package:path/path.dart' as p; import 'package:sherpa_onnx/sherpa_onnx.dart' as sherpa_onnx; +import './model_config.dart'; + +/// Resolve a relative path to absolute. +String _abs(String base, String relative) => + relative.isEmpty ? '' : p.join(base, relative); + +/// Resolve comma-separated relative paths to absolute. +String _absMulti(String base, String csv) => + csv.isEmpty ? '' : csv.split(',').map((f) => _abs(base, f.trim())).join(','); + +/// Prepare model config: copy assets to disk and resolve all paths. +Future prepareModelConfig() async { + await _copyAllAssetFiles(); + + final cfg = selectedTtsConfig; + final d = (await getApplicationSupportDirectory()).path; + final m = cfg.model; + + return sherpa_onnx.OfflineTtsConfig( + model: sherpa_onnx.OfflineTtsModelConfig( + vits: sherpa_onnx.OfflineTtsVitsModelConfig( + model: _abs(d, m.vits.model), + lexicon: _abs(d, m.vits.lexicon), + tokens: _abs(d, m.vits.tokens), + dataDir: _abs(d, m.vits.dataDir), + noiseScale: m.vits.noiseScale, + noiseScaleW: m.vits.noiseScaleW, + lengthScale: m.vits.lengthScale, + ), + kokoro: sherpa_onnx.OfflineTtsKokoroModelConfig( + model: _abs(d, m.kokoro.model), + voices: _abs(d, m.kokoro.voices), + tokens: _abs(d, m.kokoro.tokens), + dataDir: _abs(d, m.kokoro.dataDir), + lexicon: _absMulti(d, m.kokoro.lexicon), + lang: m.kokoro.lang, + lengthScale: m.kokoro.lengthScale, + ), + kitten: sherpa_onnx.OfflineTtsKittenModelConfig( + model: _abs(d, m.kitten.model), + voices: _abs(d, m.kitten.voices), + tokens: _abs(d, m.kitten.tokens), + dataDir: _abs(d, m.kitten.dataDir), + lengthScale: m.kitten.lengthScale, + ), + matcha: sherpa_onnx.OfflineTtsMatchaModelConfig( + acousticModel: _abs(d, m.matcha.acousticModel), + vocoder: _abs(d, m.matcha.vocoder), + tokens: _abs(d, m.matcha.tokens), + dataDir: _abs(d, m.matcha.dataDir), + lexicon: _abs(d, m.matcha.lexicon), + noiseScale: m.matcha.noiseScale, + lengthScale: m.matcha.lengthScale, + ), + pocket: sherpa_onnx.OfflineTtsPocketModelConfig( + lmFlow: _abs(d, m.pocket.lmFlow), + lmMain: _abs(d, m.pocket.lmMain), + encoder: _abs(d, m.pocket.encoder), + decoder: _abs(d, m.pocket.decoder), + textConditioner: _abs(d, m.pocket.textConditioner), + vocabJson: _abs(d, m.pocket.vocabJson), + tokenScoresJson: _abs(d, m.pocket.tokenScoresJson), + voiceEmbeddingCacheCapacity: m.pocket.voiceEmbeddingCacheCapacity, + ), + supertonic: sherpa_onnx.OfflineTtsSupertonicModelConfig( + durationPredictor: _abs(d, m.supertonic.durationPredictor), + textEncoder: _abs(d, m.supertonic.textEncoder), + vectorEstimator: _abs(d, m.supertonic.vectorEstimator), + vocoder: _abs(d, m.supertonic.vocoder), + ttsJson: _abs(d, m.supertonic.ttsJson), + unicodeIndexer: _abs(d, m.supertonic.unicodeIndexer), + voiceStyle: _abs(d, m.supertonic.voiceStyle), + ), + zipvoice: sherpa_onnx.OfflineTtsZipVoiceModelConfig( + tokens: _abs(d, m.zipvoice.tokens), + encoder: _abs(d, m.zipvoice.encoder), + decoder: _abs(d, m.zipvoice.decoder), + vocoder: _abs(d, m.zipvoice.vocoder), + dataDir: _abs(d, m.zipvoice.dataDir), + lexicon: _abs(d, m.zipvoice.lexicon), + featScale: m.zipvoice.featScale, + tShift: m.zipvoice.tShift, + targetRms: m.zipvoice.targetRms, + guidanceScale: m.zipvoice.guidanceScale, + ), + numThreads: m.numThreads, + debug: m.debug, + provider: m.provider, + ), + ruleFsts: _absMulti(d, cfg.ruleFsts), + ruleFars: _absMulti(d, cfg.ruleFars), + maxNumSenetences: cfg.maxNumSenetences, + ); +} -import './utils.dart'; - -Future createOfflineTts() async { - // sherpa_onnx requires that model files are in the local disk, so we - // need to copy all asset files to disk. - await copyAllAssetFiles(); - - sherpa_onnx.initBindings(); - - // Such a design is to make it easier to build flutter APPs with - // github actions for a variety of tts models - // - // See https://github.com/k2-fsa/sherpa-onnx/blob/master/scripts/flutter/generate-tts.py - // for details - - String modelDir = ''; - String modelName = ''; - String voices = ''; // for Kokoro or Kitten - bool isKitten = false; - String ruleFsts = ''; - String ruleFars = ''; - String lexicon = ''; - String dataDir = ''; - - // You can select an example below and change it accordingly to match your - // selected tts model - - // ============================================================ - // Your change starts here - // ============================================================ - - // Example 1: - // modelDir = 'vits-vctk'; - // modelName = 'vits-vctk.onnx'; - // lexicon = 'lexicon.txt'; - - // Example 2: - // https://github.com/k2-fsa/sherpa-onnx/releases/tag/tts-models - // https://github.com/k2-fsa/sherpa-onnx/releases/download/tts-models/vits-piper-en_US-amy-low.tar.bz2 - // modelDir = 'vits-piper-en_US-amy-low'; - // modelName = 'en_US-amy-low.onnx'; - // dataDir = 'vits-piper-en_US-amy-low/espeak-ng-data'; - - // Example 3: - // https://github.com/k2-fsa/sherpa-onnx/releases/download/tts-models/vits-icefall-zh-aishell3.tar.bz2 - // modelDir = 'vits-icefall-zh-aishell3'; - // modelName = 'model.onnx'; - // ruleFsts = 'vits-icefall-zh-aishell3/phone.fst,vits-icefall-zh-aishell3/date.fst,vits-icefall-zh-aishell3/number.fst,vits-icefall-zh-aishell3/new_heteronym.fst'; - // ruleFars = 'vits-icefall-zh-aishell3/rule.far'; - // lexicon = 'lexicon.txt'; - - // Example 4: - // https://k2-fsa.github.io/sherpa/onnx/tts/pretrained_models/vits.html#csukuangfj-vits-zh-hf-fanchen-c-chinese-187-speakers - // modelDir = 'vits-zh-hf-fanchen-C'; - // modelName = 'vits-zh-hf-fanchen-C.onnx'; - // lexicon = 'lexicon.txt'; - - // Example 5: - // https://github.com/k2-fsa/sherpa-onnx/releases/download/tts-models/vits-coqui-de-css10.tar.bz2 - // modelDir = 'vits-coqui-de-css10'; - // modelName = 'model.onnx'; - - // Example 6 - // https://github.com/k2-fsa/sherpa-onnx/releases/tag/tts-models - // https://github.com/k2-fsa/sherpa-onnx/releases/download/tts-models/vits-piper-en_US-libritts_r-medium.tar.bz2 - // modelDir = 'vits-piper-en_US-libritts_r-medium'; - // modelName = 'en_US-libritts_r-medium.onnx'; - // dataDir = 'vits-piper-en_US-libritts_r-medium/espeak-ng-data'; - - // Example 7 - // https://github.com/k2-fsa/sherpa-onnx/releases/tag/tts-models - // https://github.com/k2-fsa/sherpa-onnx/releases/download/tts-models/vits-melo-tts-zh_en.tar.bz2 - // modelDir = 'vits-melo-tts-zh_en'; - // modelName = 'model.onnx'; - // lexicon = 'lexicon.txt'; - - // Example 8 - // https://k2-fsa.github.io/sherpa/onnx/tts/pretrained_models/kokoro.html#kokoro-en-v0-19-english-11-speakers - // modelDir = 'kokoro-en-v0_19'; - // modelName = 'model.onnx'; - // voices = 'voices.bin'; - // dataDir = 'kokoro-en-v0_19/espeak-ng-data'; - - // Example 9 - // https://k2-fsa.github.io/sherpa/onnx/tts/pretrained_models/kokoro.html - // modelDir = 'kokoro-multi-lang-v1_0'; - // modelName = 'model.onnx'; - // voices = 'voices.bin'; - // dataDir = 'kokoro-multi-lang-v1_0/espeak-ng-data'; - // lexicon = 'kokoro-multi-lang-v1_0/lexicon-us-en.txt,kokoro-multi-lang-v1_0/lexicon-zh.txt'; - - // Example 10 - // https://github.com/k2-fsa/sherpa-onnx/releases/tag/tts-models - // modelDir = 'kitten-nano-en-v0_8-fp32'; - // modelName = 'model.fp32.onnx'; - // voices = 'voices.bin'; - // dataDir = 'kitten-nano-en-v0_8-fp32/espeak-ng-data'; - // isKitten = true; - - // ============================================================ - // Please don't change the remaining part of this function - // ============================================================ - if (modelName == '') { - throw Exception( - 'You are supposed to select a model by changing the code before you run the app'); - } - - final Directory directory = await getApplicationSupportDirectory(); - modelName = p.join(directory.path, modelDir, modelName); - - if (ruleFsts != '') { - final all = ruleFsts.split(','); - var tmp = []; - for (final f in all) { - tmp.add(p.join(directory.path, f)); - } - ruleFsts = tmp.join(','); - } - - if (ruleFars != '') { - final all = ruleFars.split(','); - var tmp = []; - for (final f in all) { - tmp.add(p.join(directory.path, f)); - } - ruleFars = tmp.join(','); - } - - if (lexicon.contains(',')) { - final all = lexicon.split(','); - var tmp = []; - for (final f in all) { - tmp.add(p.join(directory.path, f)); - } - lexicon = tmp.join(','); - } else if (lexicon != '') { - lexicon = p.join(directory.path, modelDir, lexicon); - } +/// Create an OfflineTts from a resolved config. +sherpa_onnx.OfflineTts createTtsFromConfig(sherpa_onnx.OfflineTtsConfig cfg) { + return sherpa_onnx.OfflineTts(cfg); +} - if (dataDir != '') { - dataDir = p.join(directory.path, dataDir); - } +// ── Asset copy helpers ─────────────────────────────────────────────────── - final tokens = p.join(directory.path, modelDir, 'tokens.txt'); - if (voices != '') { - voices = p.join(directory.path, modelDir, voices); +Future _copyAllAssetFiles() async { + final AssetManifest assetManifest = + await AssetManifest.loadFromAssetBundle(rootBundle); + final List assets = assetManifest.listAssets(); + for (final src in assets) { + final dst = _stripLeadingDirectory(src); + await _copyAssetFile(src, dst); } +} - late final sherpa_onnx.OfflineTtsVitsModelConfig vits; - late final sherpa_onnx.OfflineTtsKokoroModelConfig kokoro; - late final sherpa_onnx.OfflineTtsKittenModelConfig kitten; - - if (isKitten) { - vits = sherpa_onnx.OfflineTtsVitsModelConfig(); - kokoro = sherpa_onnx.OfflineTtsKokoroModelConfig(); - kitten = sherpa_onnx.OfflineTtsKittenModelConfig( - model: modelName, - voices: voices, - tokens: tokens, - dataDir: dataDir, - ); - } else if (voices != '') { - vits = sherpa_onnx.OfflineTtsVitsModelConfig(); - kitten = sherpa_onnx.OfflineTtsKittenModelConfig(); - kokoro = sherpa_onnx.OfflineTtsKokoroModelConfig( - model: modelName, - voices: voices, - tokens: tokens, - dataDir: dataDir, - lexicon: lexicon, - ); - } else { - vits = sherpa_onnx.OfflineTtsVitsModelConfig( - model: modelName, - lexicon: lexicon, - tokens: tokens, - dataDir: dataDir, - ); +String _stripLeadingDirectory(String src, {int n = 1}) { + return p.joinAll(p.split(src).sublist(n)); +} - kokoro = sherpa_onnx.OfflineTtsKokoroModelConfig(); - kitten = sherpa_onnx.OfflineTtsKittenModelConfig(); +Future _copyAssetFile(String src, [String? dst]) async { + final Directory directory = await getApplicationSupportDirectory(); + if (dst == null) dst = p.basename(src); + final target = p.join(directory.path, dst); + bool exists = await File(target).exists(); + final data = await rootBundle.load(src); + if (!exists || File(target).lengthSync() != data.lengthInBytes) { + final List bytes = + data.buffer.asUint8List(data.offsetInBytes, data.lengthInBytes); + await (await File(target).create(recursive: true)).writeAsBytes(bytes); } - - final modelConfig = sherpa_onnx.OfflineTtsModelConfig( - vits: vits, - kokoro: kokoro, - kitten: kitten, - numThreads: 2, - debug: true, - provider: 'cpu', - ); - - final config = sherpa_onnx.OfflineTtsConfig( - model: modelConfig, - ruleFsts: ruleFsts, - ruleFars: ruleFars, - maxNumSenetences: 1, - ); - // print(config); - - final tts = sherpa_onnx.OfflineTts(config); - print('tts created successfully'); - - return tts; + return target; } diff --git a/flutter-examples/tts/lib/model_config.dart b/flutter-examples/tts/lib/model_config.dart new file mode 100644 index 0000000000..913e77c1bf --- /dev/null +++ b/flutter-examples/tts/lib/model_config.dart @@ -0,0 +1,213 @@ +// Copyright (c) 2026 Xiaomi Corporation +// Shared model selection — edit this file to change the TTS model. +// Both native (model.dart) and web (model_web.dart) use this. +// +import 'package:sherpa_onnx/sherpa_onnx.dart'; + +// Change the index below to select a different model. +/// Select which TTS model to use (0-10). +const int selectedModelIndex = 0; + +/// Model directory name, extracted from the first non-empty model path. +final String selectedModelDir = () { + final m = selectedTtsConfig.model; + final path = m.vits.model.isNotEmpty + ? m.vits.model + : m.matcha.acousticModel.isNotEmpty + ? m.matcha.acousticModel + : m.kokoro.model.isNotEmpty + ? m.kokoro.model + : m.kitten.model.isNotEmpty + ? m.kitten.model + : m.pocket.lmFlow.isNotEmpty + ? m.pocket.lmFlow + : m.supertonic.durationPredictor.isNotEmpty + ? m.supertonic.durationPredictor + : m.zipvoice.encoder.isNotEmpty + ? m.zipvoice.encoder + : ''; + return path.contains('/') ? path.split('/').first : path; +}(); + +/// Download URL for the selected model. +final String selectedModelUrl = + 'https://github.com/k2-fsa/sherpa-onnx/releases/download/tts-models/$selectedModelDir.tar.bz2'; + +/// Available TTS models. +final OfflineTtsConfig selectedTtsConfig = switch (selectedModelIndex) { + // ── VITS Piper (English) ────────────────────────────────────────────── + 0 => OfflineTtsConfig( + model: OfflineTtsModelConfig( + vits: OfflineTtsVitsModelConfig( + model: 'vits-piper-en_US-amy-low/en_US-amy-low.onnx', + tokens: 'vits-piper-en_US-amy-low/tokens.txt', + dataDir: 'vits-piper-en_US-amy-low/espeak-ng-data', + ), + numThreads: 2, + debug: true, + ), + maxNumSenetences: 1, + ), + + // ── VITS Piper (Chinese) ───────────────────────────────────────────── + 1 => OfflineTtsConfig( + model: OfflineTtsModelConfig( + vits: OfflineTtsVitsModelConfig( + model: 'vits-piper-zh_CN-xiao_ya-medium/zh_CN-xiao_ya-medium.onnx', + tokens: 'vits-piper-zh_CN-xiao_ya-medium/tokens.txt', + lexicon: 'vits-piper-zh_CN-xiao_ya-medium/lexicon.txt', + ), + numThreads: 2, + debug: true, + ), + ruleFsts: 'vits-piper-zh_CN-xiao_ya-medium/phone.fst,vits-piper-zh_CN-xiao_ya-medium/date.fst,vits-piper-zh_CN-xiao_ya-medium/number.fst', + maxNumSenetences: 1, + ), + + // ── VITS Piper (English, libritts) ──────────────────────────────────── + 2 => OfflineTtsConfig( + model: OfflineTtsModelConfig( + vits: OfflineTtsVitsModelConfig( + model: 'vits-piper-en_US-libritts_r-medium/en_US-libritts_r-medium.onnx', + tokens: 'vits-piper-en_US-libritts_r-medium/tokens.txt', + dataDir: 'vits-piper-en_US-libritts_r-medium/espeak-ng-data', + ), + numThreads: 2, + debug: true, + ), + maxNumSenetences: 1, + ), + + // ── VITS (English, inflect-nano-v2) ────────────────────────────────── + // https://k2-fsa.github.io/sherpa/onnx/tts/all/English/vits-inflect-en-nano-v2.html + 3 => OfflineTtsConfig( + model: OfflineTtsModelConfig( + vits: OfflineTtsVitsModelConfig( + model: 'vits-inflect-en-nano-v2/model.onnx', + tokens: 'vits-inflect-en-nano-v2/tokens.txt', + dataDir: 'vits-inflect-en-nano-v2/espeak-ng-data', + ), + numThreads: 2, + debug: true, + ), + maxNumSenetences: 1, + ), + + // ── Kokoro (English) ────────────────────────────────────────────────── + // warning: It is super slow with single threaded wasm + 4 => OfflineTtsConfig( + model: OfflineTtsModelConfig( + kokoro: OfflineTtsKokoroModelConfig( + model: 'kokoro-int8-en-v0_19/model.int8.onnx', + voices: 'kokoro-int8-en-v0_19/voices.bin', + tokens: 'kokoro-int8-en-v0_19/tokens.txt', + dataDir: 'kokoro-int8-en-v0_19/espeak-ng-data', + ), + numThreads: 2, + debug: true, + ), + maxNumSenetences: 1, + ), + + // ── Kokoro (Chinese + English) ──────────────────────────────────────── + // warning: It is super slow with single threaded wasm + 5 => OfflineTtsConfig( + model: OfflineTtsModelConfig( + kokoro: OfflineTtsKokoroModelConfig( + model: 'kokoro-multi-lang-v1_0/model.onnx', + voices: 'kokoro-multi-lang-v1_0/voices.bin', + tokens: 'kokoro-multi-lang-v1_0/tokens.txt', + dataDir: 'kokoro-multi-lang-v1_0/espeak-ng-data', + lexicon: 'kokoro-multi-lang-v1_0/lexicon-us-en.txt,kokoro-multi-lang-v1_0/lexicon-zh.txt', + ), + numThreads: 2, + debug: true, + ), + maxNumSenetences: 1, + ), + + // ── MatchaTTS (English) ─────────────────────────────────────────────── + 6 => OfflineTtsConfig( + model: OfflineTtsModelConfig( + matcha: OfflineTtsMatchaModelConfig( + acousticModel: 'matcha-icefall-en_US-ljspeech/model-steps-3.onnx', + vocoder: 'vocos-22khz-univ.onnx', + tokens: 'matcha-icefall-en_US-ljspeech/tokens.txt', + dataDir: 'matcha-icefall-en_US-ljspeech/espeak-ng-data', + ), + numThreads: 2, + debug: true, + ), + maxNumSenetences: 1, + ), + + // ── MatchaTTS (Chinese + English) ──────────────────────────────────── + // https://k2-fsa.github.io/sherpa/onnx/tts/all/Chinese-English/matcha-icefall-zh-en.html + 7 => OfflineTtsConfig( + model: OfflineTtsModelConfig( + matcha: OfflineTtsMatchaModelConfig( + acousticModel: 'matcha-icefall-zh-en/model-steps-3.onnx', + vocoder: 'vocos-16khz-univ.onnx', + lexicon: 'matcha-icefall-zh-en/lexicon.txt', + tokens: 'matcha-icefall-zh-en/tokens.txt', + dataDir: 'matcha-icefall-zh-en/espeak-ng-data', + ), + numThreads: 2, + debug: true, + ), + ruleFsts: 'matcha-icefall-zh-en/phone-zh.fst,matcha-icefall-zh-en/date-zh.fst,matcha-icefall-zh-en/number-zh.fst', + maxNumSenetences: 1, + ), + + // ── KittenTTS (English) ─────────────────────────────────────────────── + 8 => OfflineTtsConfig( + model: OfflineTtsModelConfig( + kitten: OfflineTtsKittenModelConfig( + model: 'kitten-nano-en-v0_1-fp16/model.fp16.onnx', + voices: 'kitten-nano-en-v0_1-fp16/voices.bin', + tokens: 'kitten-nano-en-v0_1-fp16/tokens.txt', + dataDir: 'kitten-nano-en-v0_1-fp16/espeak-ng-data', + ), + numThreads: 2, + debug: true, + ), + maxNumSenetences: 1, + ), + + // ── Pocket TTS (English) ────────────────────────────────────────────── + 9 => OfflineTtsConfig( + model: OfflineTtsModelConfig( + pocket: OfflineTtsPocketModelConfig( + lmFlow: 'sherpa-onnx-pocket-tts-int8-2026-01-26/lm_flow.int8.onnx', + lmMain: 'sherpa-onnx-pocket-tts-int8-2026-01-26/lm_main.int8.onnx', + encoder: 'sherpa-onnx-pocket-tts-int8-2026-01-26/encoder.onnx', + decoder: 'sherpa-onnx-pocket-tts-int8-2026-01-26/decoder.int8.onnx', + textConditioner: 'sherpa-onnx-pocket-tts-int8-2026-01-26/text_conditioner.onnx', + vocabJson: 'sherpa-onnx-pocket-tts-int8-2026-01-26/vocab.json', + tokenScoresJson: 'sherpa-onnx-pocket-tts-int8-2026-01-26/token_scores.json', + voiceEmbeddingCacheCapacity: 50, + ), + numThreads: 2, + debug: true, + ), + ), + + // ── Supertonic TTS (English) ────────────────────────────────────────── + 10 => OfflineTtsConfig( + model: OfflineTtsModelConfig( + supertonic: OfflineTtsSupertonicModelConfig( + durationPredictor: 'sherpa-onnx-supertonic-3-tts-int8-2026-05-11/duration_predictor.int8.onnx', + textEncoder: 'sherpa-onnx-supertonic-3-tts-int8-2026-05-11/text_encoder.int8.onnx', + vectorEstimator: 'sherpa-onnx-supertonic-3-tts-int8-2026-05-11/vector_estimator.int8.onnx', + vocoder: 'sherpa-onnx-supertonic-3-tts-int8-2026-05-11/vocoder.int8.onnx', + ttsJson: 'sherpa-onnx-supertonic-3-tts-int8-2026-05-11/tts.json', + unicodeIndexer: 'sherpa-onnx-supertonic-3-tts-int8-2026-05-11/unicode_indexer.bin', + voiceStyle: 'sherpa-onnx-supertonic-3-tts-int8-2026-05-11/voice.bin', + ), + numThreads: 2, + debug: true, + ), + ), + + _ => throw ArgumentError('Invalid selectedModelIndex: $selectedModelIndex. Must be 0-10.'), +}; diff --git a/flutter-examples/tts/lib/model_web.dart b/flutter-examples/tts/lib/model_web.dart new file mode 100644 index 0000000000..729200240a --- /dev/null +++ b/flutter-examples/tts/lib/model_web.dart @@ -0,0 +1,104 @@ +// Copyright (c) 2026 Xiaomi Corporation +// Web-specific TTS model loading. +// +// This file provides three utilities for the Flutter web TTS demo: +// +// 1. [prepareModelConfig] — returns the selected OfflineTtsConfig (paths are +// relative to the WASM virtual filesystem, not absolute). +// +// 2. [loadModelFileBytes] — loads all model-related files from Flutter assets +// and returns a map of { relativePath: bytes } to be written into the WASM FS. +// +// 3. [configToJs] — converts an OfflineTtsConfig to a JSObject (via JSON +// round-trip) for passing to the Web Worker. +// +// Message flow (Dart → Worker): +// +// worker_web.dart sends an 'init' message with: +// { +// type: 'init', +// jsGlueSource: String, // sherpa-onnx-wasm-web.js source +// ttsJsSource: String, // sherpa-onnx-tts.js source +// wasmBinary: ArrayBuffer, // compiled WASM module +// modelFiles: Object, // { "path/in/fs": ArrayBuffer, ... } +// config: Object, // OfflineTtsConfig as JSON (via configToJs) +// } +// +// The config JSON uses the keys from OfflineTtsConfig.toJson(): +// { +// "model": { +// "vits": { "model": "...", "lexicon": "...", ... }, +// "matcha": { ... }, "kokoro": { ... }, ... +// "numThreads": 2, "debug": true, "provider": "cpu" +// }, +// "ruleFsts": "...", "ruleFars": "...", +// "maxNumSentences": 1, "silenceScale": 0.2 +// } +// +// See tts-worker.js for the Worker → Dart message formats. + +import 'dart:convert'; +import 'dart:js_interop'; +import 'dart:js_interop_unsafe'; +import 'dart:typed_data'; + +import 'package:flutter/services.dart'; +import 'package:sherpa_onnx/sherpa_onnx.dart' as sherpa_onnx; +import './model_config.dart'; + +/// Prepare model config for web (paths relative to WASM FS). +Future prepareModelConfig() async { + return selectedTtsConfig; +} + +/// Load model file bytes from Flutter assets. +Future> loadModelFileBytes() async { + final assetManifest = await AssetManifest.loadFromAssetBundle(rootBundle); + final allAssets = assetManifest.listAssets(); + final cfg = selectedTtsConfig; + + // Collect all model directory names from the config. + final modelDirs = {}; + for (final path in [ + cfg.model.vits.model, cfg.model.vits.tokens, cfg.model.vits.dataDir, + cfg.model.kokoro.model, cfg.model.kokoro.voices, cfg.model.kokoro.tokens, + cfg.model.kitten.model, cfg.model.kitten.voices, cfg.model.kitten.tokens, + cfg.model.matcha.acousticModel, cfg.model.matcha.vocoder, + cfg.model.pocket.lmFlow, cfg.model.pocket.lmMain, + cfg.model.supertonic.durationPredictor, + cfg.model.zipvoice.encoder, cfg.model.zipvoice.tokens, + ]) { + if (path.isNotEmpty) { + modelDirs.add(path.split('/').first); + } + } + + final modelAssets = allAssets.where((a) { + for (final dir in modelDirs) { + if (a.contains(dir)) return true; + } + return false; + }).toList(); + + final result = {}; + for (final asset in modelAssets) { + final bytes = await _loadAsset(asset); + final relativePath = asset.replaceFirst('assets/', ''); + result[relativePath] = bytes; + } + return result; +} + +/// Convert OfflineTtsConfig to JS config for the worker. +/// Uses toJson() from the config classes and converts to JSObject via JSON. +JSObject configToJs(sherpa_onnx.OfflineTtsConfig cfg) { + final jsonStr = jsonEncode(cfg.toJson()); + final jsonObj = globalContext.getProperty('JSON'.toJS) as JSObject; + final jsonParse = jsonObj.getProperty('parse'.toJS) as JSFunction; + return jsonParse.callAsFunction(jsonObj, jsonStr.toJS) as JSObject; +} + +Future _loadAsset(String assetPath) async { + final data = await rootBundle.load(assetPath); + return data.buffer.asUint8List(data.offsetInBytes, data.lengthInBytes); +} diff --git a/flutter-examples/tts/lib/play_bytes.dart b/flutter-examples/tts/lib/play_bytes.dart new file mode 100644 index 0000000000..aa895c258c --- /dev/null +++ b/flutter-examples/tts/lib/play_bytes.dart @@ -0,0 +1,15 @@ +// Native: write WAV bytes to temp file and play via AudioPlayer. +import 'dart:io'; +import 'dart:typed_data'; + +import 'package:audioplayers/audioplayers.dart'; +import 'package:path_provider/path_provider.dart'; + +Future playWavBytes(AudioPlayer player, Uint8List wavBytes) async { + final dir = await getTemporaryDirectory(); + final file = + File('${dir.path}/tts_chunk_${DateTime.now().microsecondsSinceEpoch}.wav'); + await file.writeAsBytes(wavBytes); + // play() automatically stops any previous playback. + await player.play(DeviceFileSource(file.path)); +} diff --git a/flutter-examples/tts/lib/play_bytes_stub.dart b/flutter-examples/tts/lib/play_bytes_stub.dart new file mode 100644 index 0000000000..de75528196 --- /dev/null +++ b/flutter-examples/tts/lib/play_bytes_stub.dart @@ -0,0 +1,5 @@ +// Web stub. +import 'dart:typed_data'; +import 'package:audioplayers/audioplayers.dart'; + +Future playWavBytes(AudioPlayer player, Uint8List wavBytes) async {} diff --git a/flutter-examples/tts/lib/play_ref.dart b/flutter-examples/tts/lib/play_ref.dart new file mode 100644 index 0000000000..d2ca5d7da1 --- /dev/null +++ b/flutter-examples/tts/lib/play_ref.dart @@ -0,0 +1,13 @@ +// Native: write WAV bytes to a temp file and play via AudioPlayer. +import 'dart:io'; +import 'dart:typed_data'; + +import 'package:audioplayers/audioplayers.dart'; + +Future playRefWavBytes(AudioPlayer player, Uint8List wavBytes) async { + await player.stop(); + final dir = await Directory.systemTemp.createTemp('sherpa_ref'); + final file = File('${dir.path}/ref.wav'); + await file.writeAsBytes(wavBytes); + await player.play(DeviceFileSource(file.path)); +} diff --git a/flutter-examples/tts/lib/play_ref_stub.dart b/flutter-examples/tts/lib/play_ref_stub.dart new file mode 100644 index 0000000000..ba489c1bf2 --- /dev/null +++ b/flutter-examples/tts/lib/play_ref_stub.dart @@ -0,0 +1,6 @@ +// Web stub — reference audio is played via web_audio on web, not through this. +import 'dart:typed_data'; + +import 'package:audioplayers/audioplayers.dart'; + +Future playRefWavBytes(AudioPlayer player, Uint8List wavBytes) async {} diff --git a/flutter-examples/tts/lib/save_file.dart b/flutter-examples/tts/lib/save_file.dart new file mode 100644 index 0000000000..81b2de3b38 --- /dev/null +++ b/flutter-examples/tts/lib/save_file.dart @@ -0,0 +1,30 @@ +// Native file save helper with file picker dialog. +import 'dart:io'; + +import 'package:file_picker/file_picker.dart'; + +/// Show a native save dialog and write the file. +/// Returns the chosen path, or null if cancelled. +Future saveFileAs(String sourcePath, String suggestedName) async { + try { + final bytes = File(sourcePath).readAsBytesSync(); + + final result = await FilePicker.platform.saveFile( + dialogTitle: 'Save audio file', + fileName: suggestedName, + type: FileType.custom, + allowedExtensions: ['wav'], + ); + + print('FilePicker result: $result'); + + if (result == null) return null; + + await File(result).writeAsBytes(bytes); + return result; + } catch (e, st) { + print('Error in saveFileAs: $e'); + print(st); + rethrow; + } +} diff --git a/flutter-examples/tts/lib/save_file_stub.dart b/flutter-examples/tts/lib/save_file_stub.dart new file mode 100644 index 0000000000..c6d1d4ffe4 --- /dev/null +++ b/flutter-examples/tts/lib/save_file_stub.dart @@ -0,0 +1,4 @@ +// Web stub for save_file.dart. +Future saveFileAs(String sourcePath, String destPath) async { + return destPath; +} diff --git a/flutter-examples/tts/lib/tts.dart b/flutter-examples/tts/lib/tts.dart deleted file mode 100644 index 2bce4087b0..0000000000 --- a/flutter-examples/tts/lib/tts.dart +++ /dev/null @@ -1,249 +0,0 @@ -// Copyright (c) 2024 Xiaomi Corporation -import 'dart:async'; - -import 'package:flutter/foundation.dart'; -import 'package:flutter/services.dart'; - -import 'package:flutter/material.dart'; - -import 'package:audioplayers/audioplayers.dart'; -import 'package:sherpa_onnx/sherpa_onnx.dart' as sherpa_onnx; - -import './model.dart'; -import './utils.dart'; - -class TtsScreen extends StatefulWidget { - const TtsScreen({super.key}); - - @override - State createState() => _TtsScreenState(); -} - -class _TtsScreenState extends State { - late final TextEditingController _controller_text_input; - late final TextEditingController _controller_sid; - late final TextEditingController _controller_hint; - late final AudioPlayer _player; - String _title = 'Text to speech'; - String _lastFilename = ''; - bool _isInitialized = false; - int _maxSpeakerID = 0; - double _speed = 1.0; - - sherpa_onnx.OfflineTts? _tts; - - @override - void initState() { - _controller_text_input = TextEditingController(); - _controller_hint = TextEditingController(); - _controller_sid = TextEditingController(text: '0'); - - super.initState(); - } - - Future _init() async { - if (!_isInitialized) { - sherpa_onnx.initBindings(); - - _tts?.free(); - _tts = await createOfflineTts(); - - _player = AudioPlayer(); - - _isInitialized = true; - } - } - - @override - Widget build(BuildContext context) { - return MaterialApp( - home: Scaffold( - appBar: AppBar( - title: Text(_title), - ), - body: Padding( - padding: EdgeInsets.all(10), - child: Column( - // mainAxisAlignment: MainAxisAlignment.center, - children: [ - TextField( - decoration: InputDecoration( - labelText: "Speaker ID (0-$_maxSpeakerID)", - hintText: 'Please input your speaker ID', - ), - keyboardType: TextInputType.number, - maxLines: 1, - controller: _controller_sid, - onTapOutside: (PointerDownEvent event) { - FocusManager.instance.primaryFocus?.unfocus(); - }, - inputFormatters: [FilteringTextInputFormatter.digitsOnly]), - Slider( - // decoration: InputDecoration( - // labelText: "speech speed", - // ), - label: "Speech speed ${_speed.toStringAsPrecision(2)}", - min: 0.5, - max: 3.0, - divisions: 25, - value: _speed, - onChanged: (value) { - setState(() { - _speed = value; - }); - }, - ), - const SizedBox(height: 5), - TextField( - decoration: InputDecoration( - border: OutlineInputBorder(), - hintText: 'Please enter your text here', - ), - maxLines: 5, - controller: _controller_text_input, - onTapOutside: (PointerDownEvent event) { - FocusManager.instance.primaryFocus?.unfocus(); - }, - ), - const SizedBox(height: 5), - Row(mainAxisAlignment: MainAxisAlignment.center, children: [ - OutlinedButton( - child: Text("Generate"), - onPressed: () async { - await _init(); - await _player?.stop(); - - setState(() { - _maxSpeakerID = _tts?.numSpeakers ?? 0; - if (_maxSpeakerID > 0) { - _maxSpeakerID -= 1; - } - }); - - if (_tts == null) { - _controller_hint.value = TextEditingValue( - text: 'Failed to initialize tts', - ); - return; - } - - _controller_hint.value = TextEditingValue( - text: '', - ); - - final text = _controller_text_input.text.trim(); - if (text == '') { - _controller_hint.value = TextEditingValue( - text: 'Please first input your text to generate', - ); - return; - } - - final sid = int.tryParse(_controller_sid.text.trim()) ?? 0; - - final stopwatch = Stopwatch(); - stopwatch.start(); - final genConfig = sherpa_onnx.OfflineTtsGenerationConfig( - sid: sid, - speed: _speed, - silenceScale: 0.2, - ); - final audio = - _tts!.generateWithConfig(text: text, config: genConfig); - final suffix = '-sid-$sid-speed-${_speed.toStringAsPrecision(2)}'; - final filename = await generateWaveFilename(suffix); - - final ok = sherpa_onnx.writeWave( - filename: filename, - samples: audio.samples, - sampleRate: audio.sampleRate, - ); - - if (ok) { - stopwatch.stop(); - double elapsed = stopwatch.elapsed.inMilliseconds.toDouble(); - - double waveDuration = audio.samples.length.toDouble() / audio.sampleRate.toDouble(); - - _controller_hint.value = TextEditingValue( - text: 'Saved to\n$filename\n' - 'Elapsed: ${(elapsed / 1000).toStringAsPrecision(4)} s\n' - 'Wave duration: ${waveDuration.toStringAsPrecision(4)} s\n' - 'RTF: ${(elapsed / 1000).toStringAsPrecision(4)}/${waveDuration.toStringAsPrecision(4)} ' - '= ${(elapsed / 1000 / waveDuration).toStringAsPrecision(3)} ', - ); - _lastFilename = filename; - - await _player?.play(DeviceFileSource(_lastFilename)); - } else { - _controller_hint.value = TextEditingValue( - text: 'Failed to save generated audio', - ); - } - }, - ), - const SizedBox(width: 5), - OutlinedButton( - child: Text("Clear"), - onPressed: () { - _controller_text_input.value = TextEditingValue( - text: '', - ); - - _controller_hint.value = TextEditingValue( - text: '', - ); - }, - ), - const SizedBox(width: 5), - OutlinedButton( - child: Text("Play"), - onPressed: () async { - if (_lastFilename == '') { - _controller_hint.value = TextEditingValue( - text: 'No generated wave file found', - ); - return; - } - await _player?.stop(); - await _player?.play(DeviceFileSource(_lastFilename)); - _controller_hint.value = TextEditingValue( - text: 'Playing\n$_lastFilename', - ); - }, - ), - const SizedBox(width: 5), - OutlinedButton( - child: Text("Stop"), - onPressed: () async { - await _player?.stop(); - _controller_hint.value = TextEditingValue( - text: '', - ); - }, - ), - ]), - const SizedBox(height: 5), - TextField( - decoration: InputDecoration( - border: OutlineInputBorder(), - hintText: 'Logs will be shown here.\n' - 'The first run is slower due to model initialization.', - ), - maxLines: 6, - controller: _controller_hint, - readOnly: true, - ), - ], - ), - ), - ), - ); - } - - @override - void dispose() { - _tts?.free(); - super.dispose(); - } -} diff --git a/flutter-examples/tts/lib/tts_controls.dart b/flutter-examples/tts/lib/tts_controls.dart new file mode 100644 index 0000000000..750f8828fa --- /dev/null +++ b/flutter-examples/tts/lib/tts_controls.dart @@ -0,0 +1,156 @@ +// Copyright (c) 2026 Xiaomi Corporation +import 'package:flutter/material.dart'; +import 'package:flutter/services.dart'; + +/// Input controls for TTS: speaker ID, speed slider, text input, action buttons. +class TtsControls extends StatelessWidget { + final int maxSpeakerID; + final double speed; + final ValueChanged onSpeedChanged; + final TextEditingController textController; + final TextEditingController sidController; + final VoidCallback onGenerate; + final VoidCallback onClear; + final VoidCallback? onStop; + final bool isGenerating; + + // Reference audio support (shown only for Pocket TTS). + final bool showReferenceAudio; + final String? referenceAudioLabel; + final VoidCallback? onPickReferenceAudio; + final VoidCallback? onPlayReferenceAudio; + final VoidCallback? onStopReferenceAudio; + final bool isRefPlaying; + final TextEditingController? numStepsController; + + const TtsControls({ + super.key, + required this.maxSpeakerID, + required this.speed, + required this.onSpeedChanged, + required this.textController, + required this.sidController, + required this.onGenerate, + required this.onClear, + this.onStop, + this.isGenerating = false, + this.showReferenceAudio = false, + this.referenceAudioLabel, + this.onPickReferenceAudio, + this.onPlayReferenceAudio, + this.onStopReferenceAudio, + this.isRefPlaying = false, + this.numStepsController, + }); + + @override + Widget build(BuildContext context) { + return Column( + mainAxisSize: MainAxisSize.min, + children: [ + if (showReferenceAudio) ...[ + Row( + children: [ + Expanded( + child: Text( + referenceAudioLabel != null + ? 'Reference: $referenceAudioLabel' + : 'No reference audio selected', + style: TextStyle( + fontSize: 13, + color: referenceAudioLabel != null ? null : Colors.grey, + ), + overflow: TextOverflow.ellipsis, + ), + ), + const SizedBox(width: 8), + OutlinedButton( + onPressed: onPickReferenceAudio, + child: const Text('Pick WAV'), + ), + if (referenceAudioLabel != null) ...[ + const SizedBox(width: 5), + OutlinedButton( + onPressed: isRefPlaying ? onStopReferenceAudio : onPlayReferenceAudio, + child: Text(isRefPlaying ? 'Stop' : 'Play'), + ), + ], + ], + ), + const SizedBox(height: 4), + TextField( + decoration: const InputDecoration( + labelText: 'Num steps', + hintText: '5', + isDense: true, + contentPadding: + EdgeInsets.symmetric(horizontal: 12, vertical: 8), + ), + keyboardType: TextInputType.number, + maxLines: 1, + controller: numStepsController, + onTapOutside: (_) => FocusManager.instance.primaryFocus?.unfocus(), + inputFormatters: [FilteringTextInputFormatter.digitsOnly], + ), + const SizedBox(height: 4), + ], + TextField( + decoration: InputDecoration( + labelText: 'Speaker ID (0-$maxSpeakerID)', + hintText: 'Speaker ID', + isDense: true, + contentPadding: + const EdgeInsets.symmetric(horizontal: 12, vertical: 8), + ), + keyboardType: TextInputType.number, + maxLines: 1, + controller: sidController, + onTapOutside: (_) => FocusManager.instance.primaryFocus?.unfocus(), + inputFormatters: [FilteringTextInputFormatter.digitsOnly], + ), + Slider( + label: 'Speed ${speed.toStringAsPrecision(2)}', + min: 0.5, + max: 3.0, + divisions: 25, + value: speed, + onChanged: onSpeedChanged, + ), + TextField( + decoration: const InputDecoration( + border: OutlineInputBorder(), + hintText: 'Enter text to synthesize', + contentPadding: + EdgeInsets.symmetric(horizontal: 12, vertical: 8), + ), + maxLines: 8, + minLines: 4, + controller: textController, + onTapOutside: (_) => FocusManager.instance.primaryFocus?.unfocus(), + ), + const SizedBox(height: 4), + Row( + mainAxisAlignment: MainAxisAlignment.center, + children: [ + OutlinedButton( + onPressed: isGenerating ? null : onGenerate, + child: Text(isGenerating ? 'Generating...' : 'Generate'), + ), + const SizedBox(width: 5), + OutlinedButton( + onPressed: onClear, + child: const Text('Clear'), + ), + if (onStop != null) ...[ + const SizedBox(width: 5), + OutlinedButton( + onPressed: onStop, + child: const Text('Stop'), + ), + ], + ], + ), + ], + ); + } +} diff --git a/flutter-examples/tts/lib/tts_manager.dart b/flutter-examples/tts/lib/tts_manager.dart new file mode 100644 index 0000000000..45abed08f6 --- /dev/null +++ b/flutter-examples/tts/lib/tts_manager.dart @@ -0,0 +1,487 @@ +// Copyright (c) 2026 Xiaomi Corporation +// +// TTS manager — handles TTS lifecycle for both native and web. +// +// On native: communicates with a background isolate via SendPort/ReceivePort. +// On web: delegates to TtsWorker which communicates with a Web Worker. +// +// See worker_web.dart and tts-worker.js for the web message protocol. +// +// ── Native Isolate Message Protocol ──────────────────────────────────────── +// +// The main isolate and background isolate communicate over a pair of +// SendPort/ReceivePort. The very first message from the background isolate +// is its SendPort (bidirectional setup). After that, messages flow as typed +// Dart objects. +// +// Main → Background: +// +// OfflineTtsConfig — initial config; triggers TTS creation. +// The background isolate calls sherpa_onnx.OfflineTts(config) +// and replies with _Ready or _WorkerError. +// +// _GenerateRequest — synthesize speech. +// .text String — text to synthesize +// .sid int — speaker ID +// .speed double — speech rate +// .generationId int — id for matching chunks/result +// .referenceAudio Float32List? — optional PCM samples for voice cloning +// .referenceSampleRate int — sample rate of reference audio +// .numSteps int — diffusion steps (Pocket TTS) +// +// _DisposeRequest — free the OfflineTts and close the isolate. +// +// Background → Main: +// +// SendPort — first message; the main isolate uses this to send back. +// SendPort — second+ messages; per-generation cancel port. +// Sent each time _handleGenerate starts. +// Main isolate sends `true` on this port to cancel. +// +// _Ready — TTS created successfully. +// .numSpeakers int +// +// _AudioChunk — streaming audio chunk (sent during generation). +// .samples Float32List — PCM samples +// .progress double — 0.0–1.0 +// .sampleRate int +// .generationId int +// +// _GenerateDone — generation complete. +// .samples Float32List — full PCM audio +// .sampleRate int +// .duration double — audio length in seconds +// .elapsed double — wall-clock time in seconds +// .generationId int +// +// _WorkerLog — debug/info message. +// .message String +// +// _WorkerError — error message. +// .message String +// +// ─────────────────────────────────────────────────────────────────────────── + +import 'dart:async'; +import 'dart:isolate'; +import 'dart:typed_data'; + +import 'package:flutter/foundation.dart' show kDebugMode, kIsWeb; +import 'package:sherpa_onnx/sherpa_onnx.dart' as sherpa_onnx; + +import './generated_audio.dart'; +import './model.dart' if (dart.library.js_interop) './model_web.dart' as m; +import './utils.dart' if (dart.library.js_interop) './utils_web.dart' as u; +import './web_audio.dart' if (dart.library.io) './web_audio_stub.dart' + as web_audio; +import './worker_web.dart' if (dart.library.io) './worker_stub.dart' + as worker_lib; + +/// State of the TTS engine. +enum TtsState { uninitialized, initializing, initialized } + +// ── Messages from main isolate → background isolate (native only) ──────── + +sealed class _ToWorker {} + +class _GenerateRequest extends _ToWorker { + final String text; + final int sid; + final double speed; + final int generationId; + final Float32List? referenceAudio; + final int referenceSampleRate; + final int numSteps; + _GenerateRequest(this.text, this.sid, this.speed, this.generationId, + {this.referenceAudio, this.referenceSampleRate = 0, this.numSteps = 5}); +} + +class _DisposeRequest extends _ToWorker {} + +// ── Messages from background isolate → main isolate (native only) ──────── + +sealed class _FromWorker {} + +class _Ready extends _FromWorker { + final int numSpeakers; + _Ready(this.numSpeakers); +} + +class _GenerateDone extends _FromWorker { + final Float32List samples; + final int sampleRate; + final double duration; + final double elapsed; + final int generationId; + _GenerateDone(this.samples, this.sampleRate, this.duration, this.elapsed, + this.generationId); +} + +class _AudioChunk extends _FromWorker { + final Float32List samples; + final double progress; + final int sampleRate; + final int generationId; + _AudioChunk(this.samples, this.progress, this.sampleRate, this.generationId); +} + +class _WorkerError extends _FromWorker { + final String message; + _WorkerError(this.message); +} + +class _WorkerLog extends _FromWorker { + final String message; + _WorkerLog(this.message); +} + +// ── Pending generate tracking ──────────────────────────────────────────── + +class _PendingGenerate { + final String text; + final int sid; + final double speed; + final int generationId; + _PendingGenerate({ + required this.text, + required this.sid, + required this.speed, + this.generationId = 0, + }); +} + +// ── TtsManager ─────────────────────────────────────────────────────────── + +/// Manages TTS lifecycle with isolate-based execution on native +/// and Web Worker-based execution on web. +class TtsManager { + final _logController = StreamController.broadcast(); + final _audioController = StreamController.broadcast(); + final _initController = StreamController.broadcast(); + final _chunkController = StreamController.broadcast(); + + Stream get logStream => _logController.stream; + Stream get audioStream => _audioController.stream; + Stream get initStream => _initController.stream; + Stream get chunkStream => _chunkController.stream; + + // Native: isolate-based. + Isolate? _isolate; + SendPort? _sendPort; + final Map _pending = {}; + + // Web: Web Worker-based. + worker_lib.TtsWorker? _worker; + + TtsState _state = TtsState.uninitialized; + int _numSpeakers = 0; + int _nextId = 0; + + int get numSpeakers => _numSpeakers; + TtsState get state => _state; + bool get isInitialized => _state == TtsState.initialized; + + /// Initialize the TTS engine. + Future init() async { + if (_state != TtsState.uninitialized) return; + _state = TtsState.initializing; + + if (kIsWeb) { + await _initWeb(); + } else { + await _initNative(); + } + } + + /// Generate audio from text. + int generate({ + required String text, + int sid = 0, + double speed = 1.0, + int generationId = 0, + Float32List? referenceAudio, + int referenceSampleRate = 0, + int numSteps = 5, + }) { + if (_state != TtsState.initialized) { + _logController.add('Error: TTS not initialized'); + return -1; + } + + if (kDebugMode) { + print('[tts_manager] generate: text="$text", sid=$sid, speed=$speed'); + } + + final id = _nextId++; + + if (kIsWeb) { + _worker?.generate( + text: text, sid: sid, speed: speed, generationId: generationId, + referenceAudio: referenceAudio, referenceSampleRate: referenceSampleRate, + numSteps: numSteps, + ); + } else { + _pending[id] = _PendingGenerate( + text: text, sid: sid, speed: speed, generationId: generationId, + ); + _sendPort!.send(_GenerateRequest(text, sid, speed, generationId, + referenceAudio: referenceAudio, referenceSampleRate: referenceSampleRate, + numSteps: numSteps, + )); + } + + return id; + } + + /// Cancel the current generation. + /// On web: terminates the worker (TTS is recreated on next Generate). + /// On native: sends cancel signal to the isolate (TTS stays alive). + void cancel() { + if (kIsWeb) { + // The WASM call blocks the worker, so cancel messages are queued + // and never processed. Terminate the worker instead. + _worker?.dispose(); + _worker = null; + _state = TtsState.uninitialized; + } else { + // Send cancel signal to the isolate. + // The callback checks this and returns 0 to stop generation. + _cancelPort?.send(true); + } + _pending.clear(); + } + + /// Port for sending cancel signals to the background isolate. + SendPort? _cancelPort; + + /// Dispose the TTS engine and release resources. + void dispose() { + _state = TtsState.uninitialized; + if (kIsWeb) { + _worker?.dispose(); + _worker = null; + } else { + _sendPort?.send(_DisposeRequest()); + _isolate?.kill(priority: Isolate.immediate); + _isolate = null; + _sendPort = null; + } + _logController.close(); + _audioController.close(); + _initController.close(); + _chunkController.close(); + } + + // ── Web (Web Worker) ────────────────────────────────────────────────── + + Future _initWeb() async { + try { + final readyCompleter = Completer(); + + _worker = worker_lib.TtsWorker( + onReady: (numSpeakers) { + _numSpeakers = numSpeakers; + _state = TtsState.initialized; + _logController.add('TTS ready (speakers: $_numSpeakers)'); + _initController.add(null); + if (!readyCompleter.isCompleted) readyCompleter.complete(); + }, + onChunk: (chunk) { + _chunkController.add(chunk); + }, + onDone: (item) { + _audioController.add(item); + }, + onError: (msg) { + _state = TtsState.uninitialized; + _logController.add('Error: $msg'); + if (!readyCompleter.isCompleted) readyCompleter.completeError(msg); + }, + ); + + await _worker!.init(); + await readyCompleter.future; + } catch (e) { + _state = TtsState.uninitialized; + _logController.add('Error: $e'); + } + } + + // ── Native (isolate-based) ────────────────────────────────────────────── + + Future _initNative() async { + try { + _logController.add('Preparing model...'); + final config = await m.prepareModelConfig(); + + if (kDebugMode) { + print('[tts_manager] config: $config'); + } + + // IMPORTANT: sherpa-onnx must be initialized in every isolate that calls + // its Dart API. The main isolate needs initBindings() for writeWave(), + // AudioPlayer, and other operations that happen after generation. + sherpa_onnx.initBindings(); + + _logController.add('Starting TTS isolate...'); + final receivePort = ReceivePort(); + _isolate = await Isolate.spawn(_workerEntry, receivePort.sendPort); + + final readyCompleter = Completer(); + + bool gotSendPort = false; + receivePort.listen((message) { + if (message is SendPort && !gotSendPort) { + // First SendPort: main communication port. + gotSendPort = true; + _sendPort = message; + message.send(config); + } else if (message is SendPort && gotSendPort) { + // Per-generation cancel port (sent each time _handleGenerate runs). + _cancelPort = message; + } else if (message is _Ready) { + _numSpeakers = message.numSpeakers; + _state = TtsState.initialized; + _logController.add('TTS ready (speakers: $_numSpeakers)'); + _initController.add(null); + if (!readyCompleter.isCompleted) readyCompleter.complete(); + } else if (message is _AudioChunk) { + _chunkController.add(AudioChunk( + samples: Float32List.fromList(message.samples), + progress: message.progress, + sampleRate: message.sampleRate, + generationId: message.generationId, + )); + } else if (message is _GenerateDone) { + _handleGenerateDone(message); + } else if (message is _WorkerLog) { + _logController.add(message.message); + } else if (message is _WorkerError) { + _state = TtsState.uninitialized; + _logController.add('Error: ${message.message}'); + if (!readyCompleter.isCompleted) { + readyCompleter.completeError(message.message); + } + } + }); + + return readyCompleter.future; + } catch (e) { + _state = TtsState.uninitialized; + _logController.add('Error: $e'); + rethrow; + } + } + + static void _workerEntry(SendPort mainSendPort) { + final receivePort = ReceivePort(); + mainSendPort.send(receivePort.sendPort); + + sherpa_onnx.OfflineTts? tts; + + receivePort.listen((message) { + if (message is sherpa_onnx.OfflineTtsConfig) { + try { + // IMPORTANT: sherpa-onnx must be initialized in every isolate. + sherpa_onnx.initBindings(); + tts = sherpa_onnx.OfflineTts(message); + mainSendPort.send(_Ready(tts!.numSpeakers)); + } catch (e) { + mainSendPort.send(_WorkerError('$e')); + } + } else if (message is _GenerateRequest && tts != null) { + _handleGenerate(mainSendPort, tts!, message); + } else if (message is _DisposeRequest) { + tts?.free(); + tts = null; + receivePort.close(); + } + }); + } + + static void _handleGenerate(SendPort mainSendPort, sherpa_onnx.OfflineTts tts, + _GenerateRequest req) { + // Create a fresh cancel port for each generation. + final cancelPort = ReceivePort(); + mainSendPort.send(cancelPort.sendPort); + + try { + final stopwatch = Stopwatch()..start(); + bool cancelled = false; + + cancelPort.listen((_) { + cancelled = true; + }); + + final genConfig = sherpa_onnx.OfflineTtsGenerationConfig( + sid: req.sid, + speed: req.speed, + silenceScale: 0.2, + referenceAudio: req.referenceAudio, + referenceSampleRate: req.referenceSampleRate, + numSteps: req.numSteps, + ); + + final sampleRate = tts.sampleRate; + final genId = req.generationId; + final audio = tts.generateWithConfig( + text: req.text, + config: genConfig, + onProgress: (samples, progress) { + if (cancelled) return 0; // stop generation + mainSendPort.send(_AudioChunk( + Float32List.fromList(samples), + progress, + sampleRate, + genId, + )); + return 1; // continue generation + }, + ); + + stopwatch.stop(); + final elapsed = stopwatch.elapsedMilliseconds / 1000.0; + final duration = audio.samples.length / audio.sampleRate; + + mainSendPort.send(_GenerateDone( + Float32List.fromList(audio.samples), + audio.sampleRate, + duration, + elapsed, + genId, + )); + } catch (e) { + mainSendPort.send(_WorkerError('$e')); + } finally { + cancelPort.close(); + } + } + + void _handleGenerateDone(_GenerateDone msg) async { + if (_pending.isEmpty) return; + final entry = _pending.entries.first; + _pending.remove(entry.key); + final text = entry.value.text; + + final label = GeneratedAudioItem.makeLabel(text); + final suffix = + '-sid-${entry.value.sid}-speed-${entry.value.speed.toStringAsPrecision(2)}'; + final filename = await u.generateWaveFilename(suffix); + final ok = sherpa_onnx.writeWave( + filename: filename, + samples: msg.samples, + sampleRate: msg.sampleRate, + ); + + if (ok) { + _audioController.add(GeneratedAudioItem( + label: label, + filePath: filename, + duration: msg.duration, + elapsed: msg.elapsed, + sampleRate: msg.sampleRate, + generationId: msg.generationId, + )); + } + } +} diff --git a/flutter-examples/tts/lib/tts_screen.dart b/flutter-examples/tts/lib/tts_screen.dart new file mode 100644 index 0000000000..f270fe096b --- /dev/null +++ b/flutter-examples/tts/lib/tts_screen.dart @@ -0,0 +1,444 @@ +// Copyright (c) 2026 Xiaomi Corporation +import 'dart:async'; +import 'dart:collection'; +import 'dart:typed_data'; + +import 'package:flutter/foundation.dart' show kIsWeb; +import 'package:flutter/material.dart'; + +import 'package:audioplayers/audioplayers.dart'; +import 'package:sherpa_onnx/sherpa_onnx.dart' as sherpa_onnx; + +import './generated_audio.dart'; +import './tts_manager.dart'; +import './tts_controls.dart'; +import './audio_list.dart'; +import './web_audio.dart' if (dart.library.io) './web_audio_stub.dart' + as web_audio; +import './save_file.dart' if (dart.library.js_interop) './save_file_stub.dart' + as save_file; +import './play_bytes.dart' if (dart.library.js_interop) './play_bytes_stub.dart' + as play_bytes; +import './play_ref.dart' if (dart.library.js_interop) './play_ref_stub.dart' + as play_ref; +import './wav_encoder.dart'; +import './model_config.dart' show selectedModelIndex; +import 'package:file_picker/file_picker.dart'; + +class TtsScreen extends StatefulWidget { + const TtsScreen({super.key}); + + @override + State createState() => _TtsScreenState(); +} + +class _TtsScreenState extends State { + final _textController = TextEditingController(); + final _sidController = TextEditingController(text: '0'); + final _logController = TextEditingController(); + final _numStepsController = TextEditingController(text: '5'); + + late final TtsManager _manager; + AudioPlayer? _player; + + final List _audioItems = []; + int _maxSpeakerID = 0; + double _speed = 1.0; + bool _isGenerating = false; + + // Reference audio for Pocket TTS voice cloning. + static const bool _isPocketTts = selectedModelIndex == 9; + Float32List? _referenceAudio; + int _referenceSampleRate = 0; + String? _referenceAudioName; + bool _isRefPlaying = false; + // On native: incremented on each Generate/Stop to ignore stale chunks. + // On web: unused (worker is terminated on Stop, no stale chunks). + int _generationId = 0; + + double _generationProgress = 0.0; + + // Streaming playback (native). + final List _chunkBuffer = []; + int _chunkSampleRate = 0; + static const int _chunkThresholdSamples = 16000; // ~1s at 16kHz + + // Queue of encoded WAV segments waiting to be played. + final Queue _playQueue = Queue(); + bool _isPlayingSegment = false; + + @override + void initState() { + super.initState(); + _manager = TtsManager(); + + if (!kIsWeb) { + _player = AudioPlayer(); + // Listen for playback completion to play next queued segment. + _player!.onPlayerComplete.listen((_) { + _isPlayingSegment = false; + _playNextSegment(); + if (mounted && _isRefPlaying) { + setState(() => _isRefPlaying = false); + } + }); + } + + _manager.logStream.listen((msg) { + if (mounted) { + setState(() => _logController.text = msg); + } + }); + + _manager.initStream.listen((_) { + if (mounted) { + setState(() { + _maxSpeakerID = _manager.numSpeakers; + if (_maxSpeakerID > 0) _maxSpeakerID--; + }); + } + }); + + // Stream audio chunks for real-time playback. + _manager.chunkStream.listen((chunk) { + if (!mounted) return; + // On native: ignore chunks from a previous generation. + if (!kIsWeb && chunk.generationId != _generationId) return; + + setState(() { + _generationProgress = chunk.progress; + _logController.text = + 'Generating... ${(chunk.progress * 100).toStringAsFixed(0)}%'; + }); + + if (kIsWeb) { + web_audio.playAudioChunk(chunk.samples, chunk.sampleRate); + } else { + _chunkBuffer.add(chunk.samples); + _chunkSampleRate = chunk.sampleRate; + + int totalSamples = 0; + for (final c in _chunkBuffer) { + totalSamples += c.length; + } + + if (totalSamples >= _chunkThresholdSamples) { + _flushChunkBuffer(); + } + } + }); + + // When generation completes, flush remaining and add to list. + _manager.audioStream.listen((item) { + if (!mounted) return; + // On native: ignore results from a previous generation. + if (!kIsWeb && item.generationId != _generationId) return; + + _generationProgress = 0.0; + if (!kIsWeb && _chunkBuffer.isNotEmpty) { + _flushChunkBuffer(); + } + _chunkBuffer.clear(); + + final rtf = item.elapsed / item.duration; + final status = 'Duration: ${item.duration.toStringAsFixed(2)}s\n' + 'Elapsed: ${item.elapsed.toStringAsFixed(2)}s\n' + 'RTF: ${item.elapsed.toStringAsFixed(2)} / ${item.duration.toStringAsFixed(2)} = ${rtf.toStringAsFixed(3)}'; + + setState(() { + _isGenerating = false; + _audioItems.insert(0, item); + _logController.text = status; + }); + }); + } + + /// Encode buffered chunks as WAV and enqueue for sequential playback. + void _flushChunkBuffer() { + if (_chunkBuffer.isEmpty) return; + + int total = 0; + for (final c in _chunkBuffer) { + total += c.length; + } + final merged = Float32List(total); + int offset = 0; + for (final c in _chunkBuffer) { + merged.setRange(offset, offset + c.length, c); + offset += c.length; + } + _chunkBuffer.clear(); + + final wavBytes = encodeWav(merged, _chunkSampleRate); + _playQueue.add(wavBytes); + _playNextSegment(); + } + + /// Play the next segment from the queue if not already playing. + Future _playNextSegment() async { + if (_isPlayingSegment || _playQueue.isEmpty || _player == null) return; + _isPlayingSegment = true; + final wavBytes = _playQueue.removeFirst(); + try { + await play_bytes.playWavBytes(_player!, wavBytes); + } catch (_) { + _isPlayingSegment = false; + } + } + + Future _initIfNeeded() async { + if (_manager.state != TtsState.uninitialized) return; + try { + await _manager.init(); + } catch (_) {} + } + + Future _onPickReferenceAudio() async { + final result = await FilePicker.platform.pickFiles( + type: FileType.custom, + allowedExtensions: ['wav'], + withData: true, + ); + if (result == null || result.files.isEmpty) return; + + final file = result.files.first; + final bytes = file.bytes; + if (bytes == null) { + setState(() => _logController.text = 'Error: Could not read file'); + return; + } + + final wav = decodeWav(bytes); + if (wav == null) { + setState(() => _logController.text = + 'Error: Unsupported WAV format (need 16-bit PCM or 32-bit float)'); + return; + } + setState(() { + _referenceAudio = wav.samples; + _referenceSampleRate = wav.sampleRate; + _referenceAudioName = file.name; + }); + } + + Future _playReferenceAudio() async { + if (_referenceAudio == null || _referenceAudio!.isEmpty) return; + final wavBytes = encodeWav(_referenceAudio!, _referenceSampleRate); + setState(() => _isRefPlaying = true); + if (kIsWeb) { + web_audio.playWavBytes(wavBytes); + // Web Audio API doesn't have a completion callback easily; + // reset after a rough duration estimate. + final duration = _referenceAudio!.length / _referenceSampleRate; + Future.delayed(Duration(milliseconds: (duration * 1000).ceil() + 200), () { + if (mounted) setState(() => _isRefPlaying = false); + }); + } else { + if (_player != null) { + await play_ref.playRefWavBytes(_player!, wavBytes); + } + } + } + + Future _stopReferenceAudio() async { + if (kIsWeb) { + web_audio.stopPlayback(); + } else { + await _player?.stop(); + } + setState(() => _isRefPlaying = false); + } + + Future _playAudio(GeneratedAudioItem item) async { + if (kIsWeb) { + web_audio.playWavBytes(item.wavBytes!); + } else { + await _player?.stop(); + await _player?.play(DeviceFileSource(item.filePath!)); + } + } + + Future _onGenerate() async { + await _initIfNeeded(); + + if (!kIsWeb) { + await _player?.stop(); + _isPlayingSegment = false; + _chunkBuffer.clear(); + _playQueue.clear(); + } + if (_isRefPlaying) { + setState(() => _isRefPlaying = false); + } + + final text = _textController.text.trim(); + if (text.isEmpty) { + setState(() => _logController.text = 'Please enter text to synthesize'); + return; + } + + if (_isPocketTts && _referenceAudio == null) { + setState(() => _logController.text = 'Please pick a reference WAV file first'); + return; + } + + final sid = int.tryParse(_sidController.text.trim()) ?? 0; + + if (!kIsWeb) _generationId++; + if (kIsWeb) { + web_audio.resetChunkPlayback(); + } + final numSteps = int.tryParse(_numStepsController.text.trim()) ?? 5; + final id = _manager.generate( + text: text, sid: sid, speed: _speed, generationId: _generationId, + referenceAudio: _isPocketTts ? _referenceAudio : null, + referenceSampleRate: _isPocketTts ? _referenceSampleRate : 0, + numSteps: _isPocketTts ? numSteps : 5, + ); + if (id < 0) { + // Generate failed (not initialized). + return; + } + setState(() => _isGenerating = true); + } + + Future _onSaveAs(GeneratedAudioItem item, int index) async { + if (kIsWeb) { + final controller = TextEditingController(text: '$index-${item.label}.wav'); + final filename = await showDialog( + context: context, + builder: (context) => AlertDialog( + title: const Text('Save as'), + content: TextField( + controller: controller, + decoration: const InputDecoration( + labelText: 'Filename', + border: OutlineInputBorder(), + ), + autofocus: true, + ), + actions: [ + TextButton( + onPressed: () => Navigator.pop(context), + child: const Text('Cancel'), + ), + TextButton( + onPressed: () => + Navigator.pop(context, controller.text.trim()), + child: const Text('Save'), + ), + ], + ), + ); + if (filename != null && filename.isNotEmpty) { + web_audio.downloadWavBytes(item.wavBytes!, filename); + } + } else { + try { + final savedPath = + await save_file.saveFileAs(item.filePath!, '${item.label}.wav'); + if (savedPath != null && mounted) { + ScaffoldMessenger.of(context).showSnackBar( + SnackBar(content: Text('Saved to $savedPath')), + ); + } + } catch (e) { + if (mounted) { + ScaffoldMessenger.of(context).showSnackBar( + SnackBar(content: Text('Error: $e')), + ); + } + } + } + } + + @override + Widget build(BuildContext context) { + return Scaffold( + appBar: AppBar(title: const Text('Text to Speech')), + body: Padding( + padding: const EdgeInsets.all(10), + child: Column( + children: [ + TtsControls( + maxSpeakerID: _maxSpeakerID, + speed: _speed, + onSpeedChanged: (v) => setState(() => _speed = v), + textController: _textController, + sidController: _sidController, + onGenerate: _onGenerate, + onClear: () { + _textController.clear(); + _logController.clear(); + }, + showReferenceAudio: _isPocketTts, + referenceAudioLabel: _referenceAudioName, + onPickReferenceAudio: _onPickReferenceAudio, + onPlayReferenceAudio: _playReferenceAudio, + onStopReferenceAudio: _stopReferenceAudio, + isRefPlaying: _isRefPlaying, + numStepsController: _numStepsController, + onStop: () { + // Stop generation. + // On native: increment generationId to invalidate stale chunks. + // On web: worker is terminated, no stale chunks possible. + _manager.cancel(); + if (!kIsWeb) _generationId++; + // Stop playback and clear queues. + if (kIsWeb) { + web_audio.stopPlayback(); + web_audio.resetChunkPlayback(); + } else { + _player?.stop(); + _isPlayingSegment = false; + _chunkBuffer.clear(); + _playQueue.clear(); + } + setState(() { + _isGenerating = false; + _generationProgress = 0.0; + }); + }, + isGenerating: _isGenerating || + _manager.state == TtsState.initializing, + ), + const SizedBox(height: 4), + TextField( + decoration: const InputDecoration( + border: OutlineInputBorder(), + hintText: 'Status', + isDense: true, + contentPadding: + EdgeInsets.symmetric(horizontal: 12, vertical: 8), + ), + maxLines: 3, + controller: _logController, + readOnly: true, + ), + if (_audioItems.isNotEmpty) ...[ + const SizedBox(height: 4), + Expanded( + child: AudioList( + items: _audioItems, + player: _player, + onSaveAs: _onSaveAs, + ), + ), + ], + ], + ), + ), + ); + } + + @override + void dispose() { + _manager.dispose(); + _player?.dispose(); + _textController.dispose(); + _sidController.dispose(); + _logController.dispose(); + _numStepsController.dispose(); + super.dispose(); + } +} diff --git a/flutter-examples/tts/lib/utils_web.dart b/flutter-examples/tts/lib/utils_web.dart new file mode 100644 index 0000000000..11bb8e4687 --- /dev/null +++ b/flutter-examples/tts/lib/utils_web.dart @@ -0,0 +1,10 @@ +// Copyright (c) 2026 Xiaomi Corporation +// Web-specific utilities. +import './wav_encoder.dart'; +export './wav_encoder.dart' show encodeWav; + +/// Generate a filename (on web, this is just a hint for display). +Future generateWaveFilename([String suffix = '']) async { + DateTime now = DateTime.now(); + return '${now.year}-${now.month.toString().padLeft(2, '0')}-${now.day.toString().padLeft(2, '0')}-${now.hour.toString().padLeft(2, '0')}-${now.minute.toString().padLeft(2, '0')}-${now.second.toString().padLeft(2, '0')}$suffix.wav'; +} diff --git a/flutter-examples/tts/lib/wav_encoder.dart b/flutter-examples/tts/lib/wav_encoder.dart new file mode 100644 index 0000000000..f673486df5 --- /dev/null +++ b/flutter-examples/tts/lib/wav_encoder.dart @@ -0,0 +1,135 @@ +// Copyright (c) 2026 Xiaomi Corporation +// Shared WAV encoding/decoding utility — works on all platforms (including web). +import 'dart:typed_data'; + +/// Result of decoding a WAV file. +class WavData { + final Float32List samples; + final int sampleRate; + const WavData({required this.samples, required this.sampleRate}); +} + +/// Decode WAV bytes to mono float samples + sample rate. +/// Supports 16-bit PCM and 32-bit float PCM. +/// Returns null if the format is unsupported or the file is invalid. +WavData? decodeWav(Uint8List bytes) { + if (bytes.length < 44) return null; + final bd = bytes.buffer.asByteData(bytes.offsetInBytes, bytes.lengthInBytes); + + // Check RIFF/WAVE header. + if (bytes[0] != 0x52 || bytes[1] != 0x49 || bytes[2] != 0x46 || bytes[3] != 0x46) { + return null; // Not RIFF + } + if (bytes[8] != 0x57 || bytes[9] != 0x41 || bytes[10] != 0x56 || bytes[11] != 0x45) { + return null; // Not WAVE + } + + // Find "fmt " chunk. + int offset = 12; + int audioFormat = 0; + int numChannels = 0; + int sampleRate = 0; + int bitsPerSample = 0; + bool foundFmt = false; + + while (offset + 8 <= bytes.length) { + final chunkId = String.fromCharCodes(bytes.sublist(offset, offset + 4)); + final chunkSize = bd.getUint32(offset + 4, Endian.little); + if (chunkId == 'fmt ') { + audioFormat = bd.getUint16(offset + 8, Endian.little); + numChannels = bd.getUint16(offset + 10, Endian.little); + sampleRate = bd.getUint32(offset + 12, Endian.little); + bitsPerSample = bd.getUint16(offset + 22, Endian.little); + foundFmt = true; + offset += 8 + chunkSize; + break; + } + offset += 8 + chunkSize; + } + if (!foundFmt) return null; + + // Find "data" chunk. + offset = 12; + Uint8List? dataBytes; + while (offset + 8 <= bytes.length) { + final chunkId = String.fromCharCodes(bytes.sublist(offset, offset + 4)); + final chunkSize = bd.getUint32(offset + 4, Endian.little); + if (chunkId == 'data') { + dataBytes = Uint8List.view(bytes.buffer, bytes.offsetInBytes + offset + 8, chunkSize); + break; + } + offset += 8 + chunkSize; + } + if (dataBytes == null) return null; + + // Decode to mono float samples. + if (audioFormat == 1 && bitsPerSample == 16) { + // 16-bit PCM + final numSamples = dataBytes.length ~/ (2 * numChannels); + final samples = Float32List(numSamples); + final dbd = dataBytes.buffer.asByteData(dataBytes.offsetInBytes, dataBytes.lengthInBytes); + for (int i = 0; i < numSamples; i++) { + // Mix to mono if stereo: average channels. + double sum = 0; + for (int ch = 0; ch < numChannels; ch++) { + sum += dbd.getInt16((i * numChannels + ch) * 2, Endian.little) / 32768.0; + } + samples[i] = sum / numChannels; + } + return WavData(samples: samples, sampleRate: sampleRate); + } else if (audioFormat == 3 && bitsPerSample == 32) { + // 32-bit float PCM + final numSamples = dataBytes.length ~/ (4 * numChannels); + final samples = Float32List(numSamples); + final dbd = dataBytes.buffer.asByteData(dataBytes.offsetInBytes, dataBytes.lengthInBytes); + for (int i = 0; i < numSamples; i++) { + double sum = 0; + for (int ch = 0; ch < numChannels; ch++) { + sum += dbd.getFloat32((i * numChannels + ch) * 4, Endian.little); + } + samples[i] = sum / numChannels; + } + return WavData(samples: samples, sampleRate: sampleRate); + } + + return null; // Unsupported format +} + +/// Encode Float32List PCM samples to WAV bytes (16-bit mono). +Uint8List encodeWav(Float32List samples, int sampleRate) { + const numChannels = 1; + const bitsPerSample = 16; + final byteRate = sampleRate * numChannels * bitsPerSample ~/ 8; + final blockAlign = numChannels * bitsPerSample ~/ 8; + final dataSize = samples.length * 2; + final totalSize = 44 + dataSize; + + final buffer = Uint8List(totalSize); + final bd = buffer.buffer.asByteData(); + + // RIFF header + buffer.setRange(0, 4, [0x52, 0x49, 0x46, 0x46]); // "RIFF" + bd.setUint32(4, totalSize - 8, Endian.little); + buffer.setRange(8, 12, [0x57, 0x41, 0x56, 0x45]); // "WAVE" + + // fmt chunk + buffer.setRange(12, 16, [0x66, 0x6d, 0x74, 0x20]); // "fmt " + bd.setUint32(16, 16, Endian.little); + bd.setUint16(20, 1, Endian.little); // PCM + bd.setUint16(22, numChannels, Endian.little); + bd.setUint32(24, sampleRate, Endian.little); + bd.setUint32(28, byteRate, Endian.little); + bd.setUint16(32, blockAlign, Endian.little); + bd.setUint16(34, bitsPerSample, Endian.little); + + // data chunk + buffer.setRange(36, 40, [0x64, 0x61, 0x74, 0x61]); // "data" + bd.setUint32(40, dataSize, Endian.little); + + for (int i = 0; i < samples.length; i++) { + final s = (samples[i] * 32767).clamp(-32768, 32767).toInt(); + bd.setInt16(44 + i * 2, s, Endian.little); + } + + return buffer; +} diff --git a/flutter-examples/tts/lib/web_audio.dart b/flutter-examples/tts/lib/web_audio.dart new file mode 100644 index 0000000000..f76e3cb0e3 --- /dev/null +++ b/flutter-examples/tts/lib/web_audio.dart @@ -0,0 +1,131 @@ +// Copyright (c) 2026 Xiaomi Corporation +// Web audio playback using Web Audio API. +import 'dart:js_interop'; +import 'dart:js_interop_unsafe'; +import 'dart:typed_data'; + +import './wav_encoder.dart'; +export './wav_encoder.dart' show encodeWav; + +/// Download WAV bytes as a file in the browser. +void downloadWavBytes(Uint8List wavBytes, String filename) { + globalContext['_sherpaDownloadBytes'] = wavBytes.toJS; + globalContext['_sherpaDownloadFilename'] = filename.toJS; + + final eval = globalContext.getProperty('eval'.toJS) as JSFunction; + eval.callAsFunction(null, ''' + (function() { + var bytes = window._sherpaDownloadBytes; + var name = window._sherpaDownloadFilename || 'audio.wav'; + var blob = new Blob([bytes], {type: 'audio/wav'}); + var url = URL.createObjectURL(blob); + var a = document.createElement('a'); + a.href = url; + a.download = name; + document.body.appendChild(a); + a.click(); + document.body.removeChild(a); + URL.revokeObjectURL(url); + window._sherpaDownloadBytes = null; + window._sherpaDownloadFilename = null; + })() + '''.toJS); +} + +/// Play a chunk of audio samples using the Web Audio API for streaming playback. +/// Chunks are scheduled sequentially so they play without gaps. +/// Call [resetChunkPlayback] before starting a new generation. +void playAudioChunk(Float32List samples, int sampleRate) { + // Create a unique ID for this chunk to avoid race conditions. + final id = _chunkId++; + globalContext['_sherpaChunk_$id'] = samples.toJS; + + final eval = globalContext.getProperty('eval'.toJS) as JSFunction; + eval.callAsFunction(null, ''' + (function() { + var id = $id; + var sr = $sampleRate; + // Capture data immediately in closure to avoid race with next chunk. + var samples = window['_sherpaChunk_' + id]; + setTimeout(function() { + if (!window._sherpaAudioCtx) { + window._sherpaAudioCtx = new (window.AudioContext || window.webkitAudioContext)(); + window._sherpaNextTime = 0; + } + var ctx = window._sherpaAudioCtx; + if (!samples) return; + var buf = ctx.createBuffer(1, samples.length, sr); + buf.getChannelData(0).set(samples); + var source = ctx.createBufferSource(); + source.buffer = buf; + source.connect(ctx.destination); + var startTime = Math.max(ctx.currentTime, window._sherpaNextTime); + source.start(startTime); + window._sherpaNextTime = startTime + buf.duration; + delete window['_sherpaChunk_' + id]; + }, 0); + })() + '''.toJS); +} + +int _chunkId = 0; + +/// Reset the chunk playback scheduler. Call before starting a new generation. +/// Closes the old AudioContext to stop any previously scheduled chunks. +void resetChunkPlayback() { + final eval = globalContext.getProperty('eval'.toJS) as JSFunction; + eval.callAsFunction(null, ''' + (function() { + if (window._sherpaAudioCtx) { + window._sherpaAudioCtx.close(); + window._sherpaAudioCtx = null; + } + window._sherpaNextTime = 0; + })() + '''.toJS); +} + +/// Play WAV bytes using the browser's Audio API. +/// Stops any previously playing audio first. +void playWavBytes(Uint8List wavBytes) { + globalContext['_sherpaWavBytes'] = wavBytes.toJS; + + final eval = globalContext.getProperty('eval'.toJS) as JSFunction; + eval.callAsFunction(null, ''' + (function() { + // Stop previous audio. + if (window._sherpaCurrentAudio) { + window._sherpaCurrentAudio.pause(); + window._sherpaCurrentAudio.currentTime = 0; + } + var bytes = window._sherpaWavBytes; + var blob = new Blob([bytes], {type: 'audio/wav'}); + var url = URL.createObjectURL(blob); + var audio = new Audio(url); + window._sherpaCurrentAudio = audio; + audio.play(); + window._sherpaWavBytes = null; + })() + '''.toJS); +} + +/// Stop all audio playback (both Audio elements and AudioContext). +void stopPlayback() { + final eval = globalContext.getProperty('eval'.toJS) as JSFunction; + eval.callAsFunction(null, ''' + (function() { + // Stop Audio element playback. + if (window._sherpaCurrentAudio) { + window._sherpaCurrentAudio.pause(); + window._sherpaCurrentAudio.currentTime = 0; + window._sherpaCurrentAudio = null; + } + // Stop AudioContext (streaming chunks). + if (window._sherpaAudioCtx) { + window._sherpaAudioCtx.close(); + window._sherpaAudioCtx = null; + window._sherpaNextTime = 0; + } + })() + '''.toJS); +} diff --git a/flutter-examples/tts/lib/web_audio_stub.dart b/flutter-examples/tts/lib/web_audio_stub.dart new file mode 100644 index 0000000000..5353010758 --- /dev/null +++ b/flutter-examples/tts/lib/web_audio_stub.dart @@ -0,0 +1,9 @@ +// Native stub for web_audio.dart. +import 'dart:typed_data'; + +Uint8List encodeWav(Float32List samples, int sampleRate) => Uint8List(0); +void playWavBytes(Uint8List wavBytes) {} +void stopPlayback() {} +void downloadWavBytes(Uint8List wavBytes, String filename) {} +void playAudioChunk(Float32List samples, int sampleRate) {} +void resetChunkPlayback() {} diff --git a/flutter-examples/tts/lib/worker_stub.dart b/flutter-examples/tts/lib/worker_stub.dart new file mode 100644 index 0000000000..d0ad38771e --- /dev/null +++ b/flutter-examples/tts/lib/worker_stub.dart @@ -0,0 +1,30 @@ +// Native stub for TtsWorker (not used on native). +import 'dart:typed_data'; +import './generated_audio.dart'; + +typedef OnChunkCallback = void Function(AudioChunk chunk); +typedef OnDoneCallback = void Function(GeneratedAudioItem item); +typedef OnReadyCallback = void Function(int numSpeakers); +typedef OnErrorCallback = void Function(String message); + +class TtsWorker { + TtsWorker({ + required OnChunkCallback onChunk, + required OnDoneCallback onDone, + required OnReadyCallback onReady, + required OnErrorCallback onError, + }); + + Future init() async {} + void generate({ + required String text, + int sid = 0, + double speed = 1.0, + int generationId = 0, + Float32List? referenceAudio, + int referenceSampleRate = 0, + int numSteps = 5, + }) {} + void cancel() {} + void dispose() {} +} diff --git a/flutter-examples/tts/lib/worker_web.dart b/flutter-examples/tts/lib/worker_web.dart new file mode 100644 index 0000000000..a83bbffd40 --- /dev/null +++ b/flutter-examples/tts/lib/worker_web.dart @@ -0,0 +1,204 @@ +// Web Worker support for TTS generation. +import 'dart:async'; +import 'dart:convert'; +import 'dart:js_interop'; +import 'dart:js_interop_unsafe'; +import 'dart:typed_data'; +import 'package:flutter/foundation.dart' show kDebugMode; +import 'package:flutter/services.dart'; +import 'package:web/web.dart' as web; + +import './generated_audio.dart'; +import './model_web.dart' as m; +import './web_audio.dart' as web_audio; + +typedef OnChunkCallback = void Function(AudioChunk chunk); +typedef OnDoneCallback = void Function(GeneratedAudioItem item); +typedef OnReadyCallback = void Function(int numSpeakers); +typedef OnErrorCallback = void Function(String message); + +/// Manages a Web Worker for TTS generation. +class TtsWorker { + web.Worker? _worker; + final OnChunkCallback onChunk; + final OnDoneCallback onDone; + final OnReadyCallback onReady; + final OnErrorCallback onError; + + String _pendingLabel = ''; + + TtsWorker({ + required this.onChunk, + required this.onDone, + required this.onReady, + required this.onError, + }); + + /// Initialize the worker: load WASM and model files, send to worker. + Future init() async { + final modelFiles = await m.loadModelFileBytes(); + final config = await m.prepareModelConfig(); + + if (kDebugMode) { + print('[worker_web] config: ${config.toString()}'); + print('[worker_web] modelFiles: ${modelFiles.length} files'); + } + + // Create Web Worker. + _worker = web.Worker('tts-worker.js'.toJS); + + // Listen for messages from the worker. + _worker!.onmessage = (web.MessageEvent event) { + _handleMessage(event); + }.toJS; + + // Handle worker startup failures (e.g. failed to load tts-worker.js). + _worker!.onerror = (web.ErrorEvent event) { + onError('Worker error: ${event.message}'); + }.toJS; + + // Load JS glue source, TTS helpers, and WASM binary from Flutter assets. + final jsGlueSource = await _loadAssetAsString( + 'packages/sherpa_onnx_web/assets/sherpa-onnx-wasm-web.js'); + final ttsJsSource = await _loadAssetAsString( + 'packages/sherpa_onnx_web/assets/sherpa-onnx-tts.js'); + final wasmData = await _loadAssetBytes( + 'packages/sherpa_onnx_web/assets/sherpa-onnx-wasm-web.wasm'); + + // Build model files map. + final jsModelFiles = JSObject(); + for (final entry in modelFiles.entries) { + jsModelFiles[entry.key] = entry.value.buffer.toJS; + } + + // Convert OfflineTtsConfig to JS format for the worker. + final jsConfig = m.configToJs(config); + + // Send init message with JS glue, TTS helpers, WASM binary, model files, and config. + final initMsg = JSObject(); + initMsg['type'] = 'init'.toJS; + initMsg['jsGlueSource'] = jsGlueSource.toJS; + initMsg['ttsJsSource'] = ttsJsSource.toJS; + initMsg['wasmBinary'] = wasmData.buffer.toJS; + initMsg['modelFiles'] = jsModelFiles; + initMsg['config'] = jsConfig; + _worker!.postMessage(initMsg); + } + + /// Start audio generation. + void generate({ + required String text, + int sid = 0, + double speed = 1.0, + int generationId = 0, + Float32List? referenceAudio, + int referenceSampleRate = 0, + int numSteps = 5, + }) { + _pendingLabel = GeneratedAudioItem.makeLabel(text); + final msg = JSObject(); + msg['type'] = 'generate'.toJS; + msg['text'] = text.toJS; + msg['sid'] = sid.toJS; + msg['speed'] = speed.toJS; + msg['generationId'] = generationId.toJS; + msg['numSteps'] = numSteps.toJS; + + if (referenceAudio != null && referenceAudio.isNotEmpty) { + msg['referenceAudio'] = referenceAudio.buffer.toJS; + msg['referenceSampleRate'] = referenceSampleRate.toJS; + } + _worker?.postMessage(msg); + } + + /// Cancel the current generation. + void cancel() { + final msg = JSObject(); + msg['type'] = 'cancel'.toJS; + _worker?.postMessage(msg); + } + + /// Dispose the worker. + void dispose() { + _worker?.terminate(); + _worker = null; + } + + /// Load a Flutter asset as a UTF-8 string. + static Future _loadAssetAsString(String assetPath) async { + final data = await rootBundle.load(assetPath); + return utf8.decode(data.buffer.asUint8List()); + } + + /// Load a Flutter asset as bytes. + static Future _loadAssetBytes(String assetPath) async { + final data = await rootBundle.load(assetPath); + return data.buffer.asUint8List(); + } + + /// Convert a JS ArrayBuffer to a Dart ByteBuffer. + static ByteBuffer _toByteBuffer(JSAny jsValue) { + // Wrap in Uint8Array and use toDart to copy bytes. + final uint8Ctor = + globalContext.getProperty('Uint8Array'.toJS) as JSFunction; + final view = uint8Ctor.callAsConstructor(jsValue) as JSUint8Array; + return view.toDart.buffer; + } + + void _handleMessage(web.MessageEvent event) { + final data = event.data! as JSObject; + final type = (data.getProperty('type'.toJS)! as JSString).toDart; + + if (type == 'ready') { + final numSpeakers = + (data.getProperty('numSpeakers'.toJS)! as JSNumber).toDartInt; + onReady(numSpeakers); + } else if (type == 'chunk') { + final samplesBuffer = _toByteBuffer(data.getProperty('samples'.toJS)!); + final samples = samplesBuffer.asFloat32List(); + final progress = + (data.getProperty('progress'.toJS)! as JSNumber).toDartDouble; + final sampleRate = + (data.getProperty('sampleRate'.toJS)! as JSNumber).toDartInt; + final genId = + (data.getProperty('generationId'.toJS) as JSNumber?)?.toDartInt ?? 0; + onChunk(AudioChunk( + samples: Float32List.fromList(samples), + progress: progress, + sampleRate: sampleRate, + generationId: genId, + )); + } else if (type == 'done') { + final samplesBuffer = _toByteBuffer(data.getProperty('samples'.toJS)!); + final samples = samplesBuffer.asFloat32List(); + final sampleRate = + (data.getProperty('sampleRate'.toJS)! as JSNumber).toDartInt; + final duration = + (data.getProperty('duration'.toJS)! as JSNumber).toDartDouble; + final elapsed = + (data.getProperty('elapsed'.toJS)! as JSNumber).toDartDouble; + final genId = + (data.getProperty('generationId'.toJS) as JSNumber?)?.toDartInt ?? 0; + + final wavBytes = + web_audio.encodeWav(Float32List.fromList(samples), sampleRate); + onDone(GeneratedAudioItem( + label: _pendingLabel, + generationId: genId, + wavBytes: wavBytes, + duration: duration, + elapsed: elapsed, + sampleRate: sampleRate, + )); + _pendingLabel = ''; + } else if (type == 'log') { + final msg = + (data.getProperty('message'.toJS)! as JSString).toDart; + print('[tts-worker] $msg'); + } else if (type == 'error') { + final msg = + (data.getProperty('message'.toJS)! as JSString).toDart; + onError(msg); + } + } +} diff --git a/flutter-examples/tts/macos/Flutter/Flutter-Debug.xcconfig b/flutter-examples/tts/macos/Flutter/Flutter-Debug.xcconfig index c2efd0b608..4b81f9b2d2 100644 --- a/flutter-examples/tts/macos/Flutter/Flutter-Debug.xcconfig +++ b/flutter-examples/tts/macos/Flutter/Flutter-Debug.xcconfig @@ -1 +1,2 @@ +#include? "Pods/Target Support Files/Pods-Runner/Pods-Runner.debug.xcconfig" #include "ephemeral/Flutter-Generated.xcconfig" diff --git a/flutter-examples/tts/macos/Flutter/Flutter-Release.xcconfig b/flutter-examples/tts/macos/Flutter/Flutter-Release.xcconfig index c2efd0b608..5caa9d1579 100644 --- a/flutter-examples/tts/macos/Flutter/Flutter-Release.xcconfig +++ b/flutter-examples/tts/macos/Flutter/Flutter-Release.xcconfig @@ -1 +1,2 @@ +#include? "Pods/Target Support Files/Pods-Runner/Pods-Runner.release.xcconfig" #include "ephemeral/Flutter-Generated.xcconfig" diff --git a/flutter-examples/tts/macos/Runner.xcodeproj/project.pbxproj b/flutter-examples/tts/macos/Runner.xcodeproj/project.pbxproj index 25d487e9a9..12059386c3 100644 --- a/flutter-examples/tts/macos/Runner.xcodeproj/project.pbxproj +++ b/flutter-examples/tts/macos/Runner.xcodeproj/project.pbxproj @@ -27,6 +27,9 @@ 33CC10F32044A3C60003C045 /* Assets.xcassets in Resources */ = {isa = PBXBuildFile; fileRef = 33CC10F22044A3C60003C045 /* Assets.xcassets */; }; 33CC10F62044A3C60003C045 /* MainMenu.xib in Resources */ = {isa = PBXBuildFile; fileRef = 33CC10F42044A3C60003C045 /* MainMenu.xib */; }; 33CC11132044BFA00003C045 /* MainFlutterWindow.swift in Sources */ = {isa = PBXBuildFile; fileRef = 33CC11122044BFA00003C045 /* MainFlutterWindow.swift */; }; + 78A318202AECB46A00862997 /* FlutterGeneratedPluginSwiftPackage in Frameworks */ = {isa = PBXBuildFile; productRef = 78A3181F2AECB46A00862997 /* FlutterGeneratedPluginSwiftPackage */; }; + A54264B827268A256846F780 /* Pods_RunnerTests.framework in Frameworks */ = {isa = PBXBuildFile; fileRef = 8CBD2E4539AD0DA483330DEE /* Pods_RunnerTests.framework */; }; + AF711F96FA9177B104BE7FE0 /* Pods_Runner.framework in Frameworks */ = {isa = PBXBuildFile; fileRef = 11F542A77413117A32B5F590 /* Pods_Runner.framework */; }; /* End PBXBuildFile section */ /* Begin PBXContainerItemProxy section */ @@ -60,11 +63,12 @@ /* End PBXCopyFilesBuildPhase section */ /* Begin PBXFileReference section */ + 11F542A77413117A32B5F590 /* Pods_Runner.framework */ = {isa = PBXFileReference; explicitFileType = wrapper.framework; includeInIndex = 0; path = Pods_Runner.framework; sourceTree = BUILT_PRODUCTS_DIR; }; 331C80D5294CF71000263BE5 /* RunnerTests.xctest */ = {isa = PBXFileReference; explicitFileType = wrapper.cfbundle; includeInIndex = 0; path = RunnerTests.xctest; sourceTree = BUILT_PRODUCTS_DIR; }; 331C80D7294CF71000263BE5 /* RunnerTests.swift */ = {isa = PBXFileReference; lastKnownFileType = sourcecode.swift; path = RunnerTests.swift; sourceTree = ""; }; 333000ED22D3DE5D00554162 /* Warnings.xcconfig */ = {isa = PBXFileReference; lastKnownFileType = text.xcconfig; path = Warnings.xcconfig; sourceTree = ""; }; 335BBD1A22A9A15E00E9071D /* GeneratedPluginRegistrant.swift */ = {isa = PBXFileReference; fileEncoding = 4; lastKnownFileType = sourcecode.swift; path = GeneratedPluginRegistrant.swift; sourceTree = ""; }; - 33CC10ED2044A3C60003C045 /* tts.app */ = {isa = PBXFileReference; explicitFileType = wrapper.application; includeInIndex = 0; path = "tts.app"; sourceTree = BUILT_PRODUCTS_DIR; }; + 33CC10ED2044A3C60003C045 /* tts.app */ = {isa = PBXFileReference; explicitFileType = wrapper.application; includeInIndex = 0; path = tts.app; sourceTree = BUILT_PRODUCTS_DIR; }; 33CC10F02044A3C60003C045 /* AppDelegate.swift */ = {isa = PBXFileReference; lastKnownFileType = sourcecode.swift; path = AppDelegate.swift; sourceTree = ""; }; 33CC10F22044A3C60003C045 /* Assets.xcassets */ = {isa = PBXFileReference; lastKnownFileType = folder.assetcatalog; name = Assets.xcassets; path = Runner/Assets.xcassets; sourceTree = ""; }; 33CC10F52044A3C60003C045 /* Base */ = {isa = PBXFileReference; lastKnownFileType = file.xib; name = Base; path = Base.lproj/MainMenu.xib; sourceTree = ""; }; @@ -76,8 +80,16 @@ 33E51913231747F40026EE4D /* DebugProfile.entitlements */ = {isa = PBXFileReference; lastKnownFileType = text.plist.entitlements; path = DebugProfile.entitlements; sourceTree = ""; }; 33E51914231749380026EE4D /* Release.entitlements */ = {isa = PBXFileReference; fileEncoding = 4; lastKnownFileType = text.plist.entitlements; path = Release.entitlements; sourceTree = ""; }; 33E5194F232828860026EE4D /* AppInfo.xcconfig */ = {isa = PBXFileReference; lastKnownFileType = text.xcconfig; path = AppInfo.xcconfig; sourceTree = ""; }; + 62305AB298671E57A75ADEAB /* Pods-Runner.release.xcconfig */ = {isa = PBXFileReference; includeInIndex = 1; lastKnownFileType = text.xcconfig; name = "Pods-Runner.release.xcconfig"; path = "Target Support Files/Pods-Runner/Pods-Runner.release.xcconfig"; sourceTree = ""; }; + 74E5EE46C8A6AD4F1E334201 /* Pods-Runner.debug.xcconfig */ = {isa = PBXFileReference; includeInIndex = 1; lastKnownFileType = text.xcconfig; name = "Pods-Runner.debug.xcconfig"; path = "Target Support Files/Pods-Runner/Pods-Runner.debug.xcconfig"; sourceTree = ""; }; + 78E0A7A72DC9AD7400C4905E /* FlutterGeneratedPluginSwiftPackage */ = {isa = PBXFileReference; lastKnownFileType = wrapper; name = FlutterGeneratedPluginSwiftPackage; path = ephemeral/Packages/FlutterGeneratedPluginSwiftPackage; sourceTree = ""; }; 7AFA3C8E1D35360C0083082E /* Release.xcconfig */ = {isa = PBXFileReference; lastKnownFileType = text.xcconfig; path = Release.xcconfig; sourceTree = ""; }; + 89203B494313A4092CC6B704 /* Pods-Runner.profile.xcconfig */ = {isa = PBXFileReference; includeInIndex = 1; lastKnownFileType = text.xcconfig; name = "Pods-Runner.profile.xcconfig"; path = "Target Support Files/Pods-Runner/Pods-Runner.profile.xcconfig"; sourceTree = ""; }; + 8CBD2E4539AD0DA483330DEE /* Pods_RunnerTests.framework */ = {isa = PBXFileReference; explicitFileType = wrapper.framework; includeInIndex = 0; path = Pods_RunnerTests.framework; sourceTree = BUILT_PRODUCTS_DIR; }; 9740EEB21CF90195004384FC /* Debug.xcconfig */ = {isa = PBXFileReference; fileEncoding = 4; lastKnownFileType = text.xcconfig; path = Debug.xcconfig; sourceTree = ""; }; + BE84051036F8751D7BE2C845 /* Pods-RunnerTests.debug.xcconfig */ = {isa = PBXFileReference; includeInIndex = 1; lastKnownFileType = text.xcconfig; name = "Pods-RunnerTests.debug.xcconfig"; path = "Target Support Files/Pods-RunnerTests/Pods-RunnerTests.debug.xcconfig"; sourceTree = ""; }; + C7D7E877F25497D6F37C3CBB /* Pods-RunnerTests.profile.xcconfig */ = {isa = PBXFileReference; includeInIndex = 1; lastKnownFileType = text.xcconfig; name = "Pods-RunnerTests.profile.xcconfig"; path = "Target Support Files/Pods-RunnerTests/Pods-RunnerTests.profile.xcconfig"; sourceTree = ""; }; + DA4907CD669EE274589AE809 /* Pods-RunnerTests.release.xcconfig */ = {isa = PBXFileReference; includeInIndex = 1; lastKnownFileType = text.xcconfig; name = "Pods-RunnerTests.release.xcconfig"; path = "Target Support Files/Pods-RunnerTests/Pods-RunnerTests.release.xcconfig"; sourceTree = ""; }; /* End PBXFileReference section */ /* Begin PBXFrameworksBuildPhase section */ @@ -85,6 +97,7 @@ isa = PBXFrameworksBuildPhase; buildActionMask = 2147483647; files = ( + A54264B827268A256846F780 /* Pods_RunnerTests.framework in Frameworks */, ); runOnlyForDeploymentPostprocessing = 0; }; @@ -92,6 +105,8 @@ isa = PBXFrameworksBuildPhase; buildActionMask = 2147483647; files = ( + 78A318202AECB46A00862997 /* FlutterGeneratedPluginSwiftPackage in Frameworks */, + AF711F96FA9177B104BE7FE0 /* Pods_Runner.framework in Frameworks */, ); runOnlyForDeploymentPostprocessing = 0; }; @@ -125,6 +140,7 @@ 331C80D6294CF71000263BE5 /* RunnerTests */, 33CC10EE2044A3C60003C045 /* Products */, D73912EC22F37F3D000D13A0 /* Frameworks */, + 6E9A5898EEC7D374789391AF /* Pods */, ); sourceTree = ""; }; @@ -151,6 +167,7 @@ 33CEB47122A05771004F2AC0 /* Flutter */ = { isa = PBXGroup; children = ( + 78E0A7A72DC9AD7400C4905E /* FlutterGeneratedPluginSwiftPackage */, 335BBD1A22A9A15E00E9071D /* GeneratedPluginRegistrant.swift */, 33CEB47222A05771004F2AC0 /* Flutter-Debug.xcconfig */, 33CEB47422A05771004F2AC0 /* Flutter-Release.xcconfig */, @@ -172,9 +189,25 @@ path = Runner; sourceTree = ""; }; + 6E9A5898EEC7D374789391AF /* Pods */ = { + isa = PBXGroup; + children = ( + 74E5EE46C8A6AD4F1E334201 /* Pods-Runner.debug.xcconfig */, + 62305AB298671E57A75ADEAB /* Pods-Runner.release.xcconfig */, + 89203B494313A4092CC6B704 /* Pods-Runner.profile.xcconfig */, + BE84051036F8751D7BE2C845 /* Pods-RunnerTests.debug.xcconfig */, + DA4907CD669EE274589AE809 /* Pods-RunnerTests.release.xcconfig */, + C7D7E877F25497D6F37C3CBB /* Pods-RunnerTests.profile.xcconfig */, + ); + name = Pods; + path = Pods; + sourceTree = ""; + }; D73912EC22F37F3D000D13A0 /* Frameworks */ = { isa = PBXGroup; children = ( + 11F542A77413117A32B5F590 /* Pods_Runner.framework */, + 8CBD2E4539AD0DA483330DEE /* Pods_RunnerTests.framework */, ); name = Frameworks; sourceTree = ""; @@ -186,6 +219,7 @@ isa = PBXNativeTarget; buildConfigurationList = 331C80DE294CF71000263BE5 /* Build configuration list for PBXNativeTarget "RunnerTests" */; buildPhases = ( + 4E3935CD5884A7564C5D9D89 /* [CP] Check Pods Manifest.lock */, 331C80D1294CF70F00263BE5 /* Sources */, 331C80D2294CF70F00263BE5 /* Frameworks */, 331C80D3294CF70F00263BE5 /* Resources */, @@ -204,11 +238,13 @@ isa = PBXNativeTarget; buildConfigurationList = 33CC10FB2044A3C60003C045 /* Build configuration list for PBXNativeTarget "Runner" */; buildPhases = ( + 9A2B56E744F22DA6D9083811 /* [CP] Check Pods Manifest.lock */, 33CC10E92044A3C60003C045 /* Sources */, 33CC10EA2044A3C60003C045 /* Frameworks */, 33CC10EB2044A3C60003C045 /* Resources */, 33CC110E2044A8840003C045 /* Bundle Framework */, 3399D490228B24CF009A79C7 /* ShellScript */, + 5B43ECD71027391FA603C3BD /* [CP] Embed Pods Frameworks */, ); buildRules = ( ); @@ -216,6 +252,9 @@ 33CC11202044C79F0003C045 /* PBXTargetDependency */, ); name = Runner; + packageProductDependencies = ( + 78A3181F2AECB46A00862997 /* FlutterGeneratedPluginSwiftPackage */, + ); productName = Runner; productReference = 33CC10ED2044A3C60003C045 /* tts.app */; productType = "com.apple.product-type.application"; @@ -260,6 +299,9 @@ Base, ); mainGroup = 33CC10E42044A3C60003C045; + packageReferences = ( + 781AD8BC2B33823900A9FFBB /* XCLocalSwiftPackageReference "FlutterGeneratedPluginSwiftPackage" */, + ); productRefGroup = 33CC10EE2044A3C60003C045 /* Products */; projectDirPath = ""; projectRoot = ""; @@ -329,6 +371,67 @@ shellPath = /bin/sh; shellScript = "\"$FLUTTER_ROOT\"/packages/flutter_tools/bin/macos_assemble.sh && touch Flutter/ephemeral/tripwire"; }; + 4E3935CD5884A7564C5D9D89 /* [CP] Check Pods Manifest.lock */ = { + isa = PBXShellScriptBuildPhase; + buildActionMask = 2147483647; + files = ( + ); + inputFileListPaths = ( + ); + inputPaths = ( + "${PODS_PODFILE_DIR_PATH}/Podfile.lock", + "${PODS_ROOT}/Manifest.lock", + ); + name = "[CP] Check Pods Manifest.lock"; + outputFileListPaths = ( + ); + outputPaths = ( + "$(DERIVED_FILE_DIR)/Pods-RunnerTests-checkManifestLockResult.txt", + ); + runOnlyForDeploymentPostprocessing = 0; + shellPath = /bin/sh; + shellScript = "diff \"${PODS_PODFILE_DIR_PATH}/Podfile.lock\" \"${PODS_ROOT}/Manifest.lock\" > /dev/null\nif [ $? != 0 ] ; then\n # print error to STDERR\n echo \"error: The sandbox is not in sync with the Podfile.lock. Run 'pod install' or update your CocoaPods installation.\" >&2\n exit 1\nfi\n# This output is used by Xcode 'outputs' to avoid re-running this script phase.\necho \"SUCCESS\" > \"${SCRIPT_OUTPUT_FILE_0}\"\n"; + showEnvVarsInLog = 0; + }; + 5B43ECD71027391FA603C3BD /* [CP] Embed Pods Frameworks */ = { + isa = PBXShellScriptBuildPhase; + buildActionMask = 2147483647; + files = ( + ); + inputFileListPaths = ( + "${PODS_ROOT}/Target Support Files/Pods-Runner/Pods-Runner-frameworks-${CONFIGURATION}-input-files.xcfilelist", + ); + name = "[CP] Embed Pods Frameworks"; + outputFileListPaths = ( + "${PODS_ROOT}/Target Support Files/Pods-Runner/Pods-Runner-frameworks-${CONFIGURATION}-output-files.xcfilelist", + ); + runOnlyForDeploymentPostprocessing = 0; + shellPath = /bin/sh; + shellScript = "\"${PODS_ROOT}/Target Support Files/Pods-Runner/Pods-Runner-frameworks.sh\"\n"; + showEnvVarsInLog = 0; + }; + 9A2B56E744F22DA6D9083811 /* [CP] Check Pods Manifest.lock */ = { + isa = PBXShellScriptBuildPhase; + buildActionMask = 2147483647; + files = ( + ); + inputFileListPaths = ( + ); + inputPaths = ( + "${PODS_PODFILE_DIR_PATH}/Podfile.lock", + "${PODS_ROOT}/Manifest.lock", + ); + name = "[CP] Check Pods Manifest.lock"; + outputFileListPaths = ( + ); + outputPaths = ( + "$(DERIVED_FILE_DIR)/Pods-Runner-checkManifestLockResult.txt", + ); + runOnlyForDeploymentPostprocessing = 0; + shellPath = /bin/sh; + shellScript = "diff \"${PODS_PODFILE_DIR_PATH}/Podfile.lock\" \"${PODS_ROOT}/Manifest.lock\" > /dev/null\nif [ $? != 0 ] ; then\n # print error to STDERR\n echo \"error: The sandbox is not in sync with the Podfile.lock. Run 'pod install' or update your CocoaPods installation.\" >&2\n exit 1\nfi\n# This output is used by Xcode 'outputs' to avoid re-running this script phase.\necho \"SUCCESS\" > \"${SCRIPT_OUTPUT_FILE_0}\"\n"; + showEnvVarsInLog = 0; + }; /* End PBXShellScriptBuildPhase section */ /* Begin PBXSourcesBuildPhase section */ @@ -380,6 +483,7 @@ /* Begin XCBuildConfiguration section */ 331C80DB294CF71000263BE5 /* Debug */ = { isa = XCBuildConfiguration; + baseConfigurationReference = BE84051036F8751D7BE2C845 /* Pods-RunnerTests.debug.xcconfig */; buildSettings = { BUNDLE_LOADER = "$(TEST_HOST)"; CURRENT_PROJECT_VERSION = 1; @@ -394,6 +498,7 @@ }; 331C80DC294CF71000263BE5 /* Release */ = { isa = XCBuildConfiguration; + baseConfigurationReference = DA4907CD669EE274589AE809 /* Pods-RunnerTests.release.xcconfig */; buildSettings = { BUNDLE_LOADER = "$(TEST_HOST)"; CURRENT_PROJECT_VERSION = 1; @@ -408,6 +513,7 @@ }; 331C80DD294CF71000263BE5 /* Profile */ = { isa = XCBuildConfiguration; + baseConfigurationReference = C7D7E877F25497D6F37C3CBB /* Pods-RunnerTests.profile.xcconfig */; buildSettings = { BUNDLE_LOADER = "$(TEST_HOST)"; CURRENT_PROJECT_VERSION = 1; @@ -461,7 +567,7 @@ GCC_WARN_UNINITIALIZED_AUTOS = YES_AGGRESSIVE; GCC_WARN_UNUSED_FUNCTION = YES; GCC_WARN_UNUSED_VARIABLE = YES; - MACOSX_DEPLOYMENT_TARGET = 10.14; + MACOSX_DEPLOYMENT_TARGET = 10.15; MTL_ENABLE_DEBUG_INFO = NO; SDKROOT = macosx; SWIFT_COMPILATION_MODE = wholemodule; @@ -543,7 +649,7 @@ GCC_WARN_UNINITIALIZED_AUTOS = YES_AGGRESSIVE; GCC_WARN_UNUSED_FUNCTION = YES; GCC_WARN_UNUSED_VARIABLE = YES; - MACOSX_DEPLOYMENT_TARGET = 10.14; + MACOSX_DEPLOYMENT_TARGET = 10.15; MTL_ENABLE_DEBUG_INFO = YES; ONLY_ACTIVE_ARCH = YES; SDKROOT = macosx; @@ -593,7 +699,7 @@ GCC_WARN_UNINITIALIZED_AUTOS = YES_AGGRESSIVE; GCC_WARN_UNUSED_FUNCTION = YES; GCC_WARN_UNUSED_VARIABLE = YES; - MACOSX_DEPLOYMENT_TARGET = 10.14; + MACOSX_DEPLOYMENT_TARGET = 10.15; MTL_ENABLE_DEBUG_INFO = NO; SDKROOT = macosx; SWIFT_COMPILATION_MODE = wholemodule; @@ -700,6 +806,20 @@ defaultConfigurationName = Release; }; /* End XCConfigurationList section */ + +/* Begin XCLocalSwiftPackageReference section */ + 781AD8BC2B33823900A9FFBB /* XCLocalSwiftPackageReference "FlutterGeneratedPluginSwiftPackage" */ = { + isa = XCLocalSwiftPackageReference; + relativePath = Flutter/ephemeral/Packages/FlutterGeneratedPluginSwiftPackage; + }; +/* End XCLocalSwiftPackageReference section */ + +/* Begin XCSwiftPackageProductDependency section */ + 78A3181F2AECB46A00862997 /* FlutterGeneratedPluginSwiftPackage */ = { + isa = XCSwiftPackageProductDependency; + productName = FlutterGeneratedPluginSwiftPackage; + }; +/* End XCSwiftPackageProductDependency section */ }; rootObject = 33CC10E52044A3C60003C045 /* Project object */; } diff --git a/flutter-examples/tts/macos/Runner.xcodeproj/xcshareddata/xcschemes/Runner.xcscheme b/flutter-examples/tts/macos/Runner.xcodeproj/xcshareddata/xcschemes/Runner.xcscheme index 372f0df7ba..8c784f765c 100644 --- a/flutter-examples/tts/macos/Runner.xcodeproj/xcshareddata/xcschemes/Runner.xcscheme +++ b/flutter-examples/tts/macos/Runner.xcodeproj/xcshareddata/xcschemes/Runner.xcscheme @@ -5,6 +5,24 @@ + + + + + + + + + + diff --git a/flutter-examples/tts/macos/Runner.xcworkspace/contents.xcworkspacedata b/flutter-examples/tts/macos/Runner.xcworkspace/contents.xcworkspacedata index 1d526a16ed..21a3cc14c7 100644 --- a/flutter-examples/tts/macos/Runner.xcworkspace/contents.xcworkspacedata +++ b/flutter-examples/tts/macos/Runner.xcworkspace/contents.xcworkspacedata @@ -4,4 +4,7 @@ + + diff --git a/flutter-examples/tts/macos/Runner/AppDelegate.swift b/flutter-examples/tts/macos/Runner/AppDelegate.swift index d53ef64377..b3c1761412 100644 --- a/flutter-examples/tts/macos/Runner/AppDelegate.swift +++ b/flutter-examples/tts/macos/Runner/AppDelegate.swift @@ -1,9 +1,13 @@ import Cocoa import FlutterMacOS -@NSApplicationMain +@main class AppDelegate: FlutterAppDelegate { override func applicationShouldTerminateAfterLastWindowClosed(_ sender: NSApplication) -> Bool { return true } + + override func applicationSupportsSecureRestorableState(_ app: NSApplication) -> Bool { + return true + } } diff --git a/flutter-examples/tts/macos/Runner/DebugProfile.entitlements b/flutter-examples/tts/macos/Runner/DebugProfile.entitlements index dddb8a30c8..d138bd5b04 100644 --- a/flutter-examples/tts/macos/Runner/DebugProfile.entitlements +++ b/flutter-examples/tts/macos/Runner/DebugProfile.entitlements @@ -8,5 +8,7 @@ com.apple.security.network.server + com.apple.security.files.user-selected.read-write + diff --git a/flutter-examples/tts/macos/Runner/MainFlutterWindow.swift b/flutter-examples/tts/macos/Runner/MainFlutterWindow.swift index 3cc05eb234..cc5bc9c807 100644 --- a/flutter-examples/tts/macos/Runner/MainFlutterWindow.swift +++ b/flutter-examples/tts/macos/Runner/MainFlutterWindow.swift @@ -6,7 +6,8 @@ class MainFlutterWindow: NSWindow { let flutterViewController = FlutterViewController() let windowFrame = self.frame self.contentViewController = flutterViewController - self.setFrame(windowFrame, display: true) + self.setFrame(NSRect(x: windowFrame.origin.x, y: windowFrame.origin.y, + width: 900, height: 750), display: true) RegisterGeneratedPlugins(registry: flutterViewController) diff --git a/flutter-examples/tts/macos/Runner/Release.entitlements b/flutter-examples/tts/macos/Runner/Release.entitlements index 852fa1a472..19afff14a0 100644 --- a/flutter-examples/tts/macos/Runner/Release.entitlements +++ b/flutter-examples/tts/macos/Runner/Release.entitlements @@ -4,5 +4,7 @@ com.apple.security.app-sandbox + com.apple.security.files.user-selected.read-write + diff --git a/flutter-examples/tts/pubspec.yaml b/flutter-examples/tts/pubspec.yaml index 068c02a5d9..23e6cfe57b 100644 --- a/flutter-examples/tts/pubspec.yaml +++ b/flutter-examples/tts/pubspec.yaml @@ -8,7 +8,7 @@ publish_to: 'none' # Remove this line if you wish to publish to pub.dev version: 1.13.4 environment: - sdk: ">=2.17.0 <4.0.0" + sdk: ">=3.1.0 <4.0.0" flutter: ">=2.8.1" dependencies: @@ -16,6 +16,7 @@ dependencies: sdk: flutter cupertino_icons: ^1.0.6 + file_picker: ^8.0.0 path_provider: ^2.1.3 path: ^1.9.0 sherpa_onnx: ^1.13.4 @@ -24,12 +25,6 @@ dependencies: url_launcher: 6.2.6 url_launcher_linux: 3.1.0 audioplayers: ^5.0.0 - media_kit: - media_kit_libs_video: flutter: uses-material-design: true - - assets: - - assets/vits-melo-tts-zh_en/ - - assets/vits-melo-tts-zh_en/dict/ \ No newline at end of file diff --git a/flutter-examples/tts/web/index.html b/flutter-examples/tts/web/index.html new file mode 100644 index 0000000000..54c56a7ced --- /dev/null +++ b/flutter-examples/tts/web/index.html @@ -0,0 +1,46 @@ + + + + + + + + + + + + + + + + + + + + tts + + + + + + + diff --git a/flutter-examples/tts/web/tts-worker.js b/flutter-examples/tts/web/tts-worker.js new file mode 100644 index 0000000000..bc2e4eb537 --- /dev/null +++ b/flutter-examples/tts/web/tts-worker.js @@ -0,0 +1,238 @@ +// TTS Web Worker — runs WASM generation off the main thread. +// +// Config helpers (initSherpaOnnxOfflineTtsConfig, freeConfig, +// initSherpaOnnxGenerationConfig, freeSherpaOnnxGenerationConfig) +// are loaded from sherpa-onnx-tts.js via eval at init time. +// +// ── Messages: Main Thread → Worker ───────────────────────────────────────── +// +// init — Initialize the WASM module and create a TTS instance. +// { +// type: 'init', +// jsGlueSource: String, // sherpa-onnx-wasm-web.js (defines SherpaOnnx factory) +// ttsJsSource: String, // sherpa-onnx-tts.js (config/generation helpers) +// wasmBinary: ArrayBuffer, // compiled .wasm module +// modelFiles: Object, // { "relative/path": ArrayBuffer, ... } +// config: Object, // OfflineTtsConfig JSON (from toJson()) +// } +// +// generate — Synthesize speech from text. +// { +// type: 'generate', +// text: String, // text to synthesize +// sid: Number, // speaker ID (default 0) +// speed: Number, // speech rate (default 1.0) +// generationId: Number, // id for matching chunks/done (default 0) +// referenceAudio: ArrayBuffer, // optional: Float32 PCM samples for voice cloning +// referenceSampleRate: Number, // sample rate of reference audio +// numSteps: Number, // diffusion steps (default 5) +// } +// +// cancel — Abort the current generation. +// { type: 'cancel' } +// +// dispose — Destroy the TTS instance and close the worker. +// { type: 'dispose' } +// +// ── Messages: Worker → Main Thread ───────────────────────────────────────── +// +// ready — TTS initialized successfully. +// { type: 'ready', numSpeakers: Number, sampleRate: Number } +// +// chunk — Streaming audio chunk (sent during generation). +// { +// type: 'chunk', +// samples: ArrayBuffer, // Float32 PCM (transferred, not copied) +// progress: Number, // 0.0–1.0 +// sampleRate: Number, +// generationId: Number, +// } +// +// done — Generation complete. +// { +// type: 'done', +// samples: ArrayBuffer, // Float32 PCM (transferred) +// sampleRate: Number, +// duration: Number, // audio duration in seconds +// elapsed: Number, // wall-clock time in seconds +// generationId: Number, +// } +// +// log — Debug/info message from the WASM module (stdout/stderr). +// { type: 'log', message: String } +// +// error — Error message. +// { type: 'error', message: String } + +let Module = null; +let tts = null; +let _cancelled = false; + +// ── Emscripten FS helpers ──────────────────────────────────────────────── + +function getFS() { + if (Module && Module.FS) return Module.FS; + if (typeof FS !== 'undefined') return FS; + throw new Error('FS not found'); +} + +function mkdirTree(path) { + const fs = getFS(); + const parts = path.split('/'); + let current = ''; + for (const part of parts) { + if (!part) continue; + current = current + '/' + part; + try { fs.mkdir(current); } catch (_) {} + } +} + +function writeFile(path, data) { + getFS().writeFile(path, data); +} + +// ── Message handler ────────────────────────────────────────────────────── + +self.onmessage = async function(e) { + const msg = e.data; + + if (msg.type === 'init') { + try { + // 1. Load Emscripten JS glue (defines SherpaOnnx factory). + if (msg.jsGlueSource) { + self.eval(msg.jsGlueSource); + } + + // 2. Load sherpa-onnx-tts.js helpers (defines initSherpaOnnxOfflineTtsConfig, + // freeConfig, initSherpaOnnxGenerationConfig, freeSherpaOnnxGenerationConfig, etc.) + if (msg.ttsJsSource) { + self.eval(msg.ttsJsSource); + } + + // 3. Compile WASM module. + const wasmBytes = new Uint8Array(msg.wasmBinary); + Module = await SherpaOnnx({ + wasmBinary: wasmBytes, + print: (text) => self.postMessage({ type: 'log', message: text }), + printErr: (text) => self.postMessage({ type: 'log', message: '[stderr] ' + text }), + }); + + // 4. Write model files to WASM FS. + const modelFiles = msg.modelFiles; + for (const [path, bytes] of Object.entries(modelFiles)) { + const dir = path.substring(0, path.lastIndexOf('/')); + if (dir) mkdirTree(dir); + writeFile(path, new Uint8Array(bytes)); + } + + // 5. Create TTS instance using sherpa-onnx-tts.js helper. + const config = initSherpaOnnxOfflineTtsConfig(msg.config, Module); + const handle = Module._SherpaOnnxCreateOfflineTts(config.ptr); + freeConfig(config, Module); + + if (!handle) { + self.postMessage({ type: 'error', message: 'Failed to create TTS (null handle)' }); + return; + } + + const sampleRate = Module._SherpaOnnxOfflineTtsSampleRate(handle); + const numSpeakers = Module._SherpaOnnxOfflineTtsNumSpeakers(handle); + + tts = { handle, sampleRate, numSpeakers }; + self.postMessage({ type: 'ready', numSpeakers, sampleRate }); + } catch (e) { + self.postMessage({ type: 'error', message: e.message || String(e) }); + } + } + + else if (msg.type === 'generate' && tts) { + try { + const startTime = performance.now(); + + const genCfg = { + silenceScale: 0.2, + speed: msg.speed || 1.0, + sid: msg.sid || 0, + }; + + // Reference audio for voice cloning (e.g. Pocket TTS). + if (msg.referenceAudio) { + genCfg.referenceAudio = new Float32Array(msg.referenceAudio); + genCfg.referenceSampleRate = msg.referenceSampleRate || 0; + genCfg.numSteps = msg.numSteps || 5; + } + + const cfgWasm = initSherpaOnnxGenerationConfig(genCfg, Module); + + // Set up callback for streaming chunks. + _cancelled = false; + const genId = msg.generationId || 0; + const callbackPtr = Module.addFunction((samplesPtr, n, progress, arg) => { + if (_cancelled) return 0; + const samples = new Float32Array(Module.HEAPF32.buffer, samplesPtr, n).slice(); + self.postMessage({ + type: 'chunk', + samples: samples.buffer, + progress: progress, + sampleRate: tts.sampleRate, + generationId: genId, + }, [samples.buffer]); + return 1; + }, 'iiifi'); + + // Prepare text. + const textLen = Module.lengthBytesUTF8(msg.text) + 1; + const textPtr = Module._malloc(textLen); + Module.stringToUTF8(msg.text, textPtr, textLen); + + // Generate. + const audioPtr = Module._SherpaOnnxOfflineTtsGenerateWithConfig( + tts.handle, textPtr, cfgWasm.ptr, callbackPtr, 0); + + Module._free(textPtr); + freeSherpaOnnxGenerationConfig(cfgWasm, Module); + Module.removeFunction(callbackPtr); + + if (!audioPtr) { + self.postMessage({ type: 'error', message: 'Generation failed' }); + return; + } + + // Read result. + const base = audioPtr / 4; + const samplesPtr = Module.HEAPU32[base]; + const numSamples = Module.HEAP32[base + 1]; + const sampleRateOut = Module.HEAP32[base + 2]; + + const samples = new Float32Array(Module.HEAPF32.buffer, samplesPtr, numSamples).slice(); + + Module._SherpaOnnxDestroyOfflineTtsGeneratedAudio(audioPtr); + + const elapsed = (performance.now() - startTime) / 1000; + const duration = numSamples / sampleRateOut; + + self.postMessage({ + type: 'done', + samples: samples.buffer, + sampleRate: sampleRateOut, + duration: duration, + elapsed: elapsed, + generationId: genId, + }, [samples.buffer]); + } catch (e) { + self.postMessage({ type: 'error', message: e.message || String(e) }); + } + } + + else if (msg.type === 'cancel') { + _cancelled = true; + } + + else if (msg.type === 'dispose') { + if (tts) { + Module._SherpaOnnxDestroyOfflineTts(tts.handle); + tts = null; + } + self.close(); + } +}; diff --git a/flutter/sherpa_onnx/lib/sherpa_onnx.dart b/flutter/sherpa_onnx/lib/sherpa_onnx.dart index 775ead6dcb..d135226d42 100644 --- a/flutter/sherpa_onnx/lib/sherpa_onnx.dart +++ b/flutter/sherpa_onnx/lib/sherpa_onnx.dart @@ -1,6 +1,13 @@ // Copyright (c) 2024 Xiaomi Corporation -import 'dart:io'; -import 'dart:ffi'; +import 'package:flutter/foundation.dart' show kIsWeb; + +// Conditional import: native uses dart:io/dart:ffi, web uses dart:js_interop. +import 'src/init_native.dart' + if (dart.library.js_interop) 'src/web/init.dart' as init; + +// Conditional import for web WASM loader. +import 'package:sherpa_onnx_web/sherpa_onnx_web.dart' + if (dart.library.io) 'src/init_stub.dart' as web; /// Dart bindings for the public sherpa-onnx inference APIs. /// @@ -9,8 +16,9 @@ import 'dart:ffi'; /// audio tagging, spoken language identification, speech denoising, and WAV /// I/O helpers from a single entry point. /// -/// Before creating any runtime object, call [initBindings] once so the package -/// can load the underlying native `sherpa-onnx-c-api` library for the current +/// Before creating any runtime object, call [initBindings] (native) or +/// [initBindingsAsync] (all platforms including web) once so the package can +/// load the underlying native `sherpa-onnx-c-api` library for the current /// platform. /// /// For concrete end-to-end usage, see `dart-api-examples/` in the repository, @@ -24,83 +32,95 @@ import 'dart:ffi'; /// - `vad/bin/vad.dart` /// - `speaker-diarization/` -export 'src/audio_tagging.dart'; +export 'src/audio_tagging_config.dart'; +export 'src/audio_tagging.dart' + if (dart.library.js_interop) 'src/web/audio_tagging.dart'; export 'src/feature_config.dart'; export 'src/homophone_replacer_config.dart'; -export 'src/keyword_spotter.dart'; -export 'src/offline_punctuation.dart'; -export 'src/offline_recognizer.dart'; -export 'src/offline_speaker_diarization.dart'; -export 'src/offline_speech_denoiser.dart'; -export 'src/offline_stream.dart'; -export 'src/online_speech_denoiser.dart'; -export 'src/online_punctuation.dart'; -export 'src/online_recognizer.dart'; -export 'src/online_stream.dart'; -export 'src/speaker_identification.dart'; -export 'src/spoken_language_identification.dart'; -export 'src/tts.dart'; -export 'src/vad.dart'; -export 'src/version.dart'; -export 'src/wave_reader.dart'; -export 'src/wave_writer.dart'; - -import 'src/sherpa_onnx_bindings.dart'; +export 'src/keyword_spotter_config.dart'; +export 'src/keyword_spotter.dart' + if (dart.library.js_interop) 'src/web/keyword_spotter.dart'; +export 'src/offline_punctuation_config.dart'; +export 'src/offline_punctuation.dart' + if (dart.library.js_interop) 'src/web/offline_punctuation.dart'; +export 'src/offline_recognizer_config.dart'; +export 'src/offline_recognizer.dart' + if (dart.library.js_interop) 'src/web/offline_recognizer.dart'; +export 'src/offline_speaker_diarization_config.dart'; +export 'src/offline_speaker_diarization.dart' + if (dart.library.js_interop) 'src/web/offline_speaker_diarization.dart'; +export 'src/offline_speech_denoiser_config.dart'; +export 'src/offline_speech_denoiser.dart' + if (dart.library.js_interop) 'src/web/offline_speech_denoiser.dart'; +export 'src/offline_stream.dart' + if (dart.library.js_interop) 'src/web/offline_stream.dart'; +export 'src/online_speech_denoiser_config.dart'; +export 'src/online_speech_denoiser.dart' + if (dart.library.js_interop) 'src/web/online_speech_denoiser.dart'; +export 'src/online_punctuation_config.dart'; +export 'src/online_punctuation.dart' + if (dart.library.js_interop) 'src/web/online_punctuation.dart'; +export 'src/online_recognizer_config.dart'; +export 'src/online_recognizer.dart' + if (dart.library.js_interop) 'src/web/online_recognizer.dart'; +export 'src/online_stream.dart' + if (dart.library.js_interop) 'src/web/online_stream.dart'; +export 'src/speaker_identification_config.dart'; +export 'src/speaker_identification.dart' + if (dart.library.js_interop) 'src/web/speaker_identification.dart'; +export 'src/spoken_language_identification_config.dart'; +export 'src/spoken_language_identification.dart' + if (dart.library.js_interop) 'src/web/spoken_language_identification.dart'; +export 'src/tts_config.dart'; +export 'src/tts.dart' + if (dart.library.js_interop) 'src/web/tts.dart'; +export 'src/vad_config.dart'; +export 'src/vad.dart' + if (dart.library.js_interop) 'src/web/vad.dart'; +export 'src/version.dart' + if (dart.library.js_interop) 'src/web/version.dart'; +export 'src/wave_reader_config.dart'; +export 'src/wave_reader.dart' + if (dart.library.js_interop) 'src/web/wave_reader.dart'; +export 'src/wave_writer.dart' + if (dart.library.js_interop) 'src/web/wave_writer.dart'; -// it is not-empty for Dart CLI -// See ../../../dart-api-examples/vad/bin/init.dart String? _path; -// see also -// https://github.com/flutter/codelabs/blob/main/ffigen_codelab/step_05/lib/ffigen_app.dart -// https://api.flutter.dev/flutter/dart-io/Platform-class.html -final DynamicLibrary _dylib = () { - if (Platform.isMacOS) { - if (_path == null) { - // for Flutter - return DynamicLibrary.open('SherpaOnnxC.framework/SherpaOnnxC'); - } else { - // for Dart CLI without flutter - // CI places SherpaOnnxC.xcframework inside ../../../flutter/sherpa_onnx_macos/macos/sherpa_onnx_macos - return DynamicLibrary.open( - '$_path/sherpa_onnx_macos/SherpaOnnxC.xcframework/macos-arm64_x86_64/SherpaOnnxC.framework/SherpaOnnxC', - ); - } - } - - if (Platform.isIOS) { - // CI places SherpaOnnxC.xcframework inside ../../../flutter/sherpa_onnx_ios/ios/sherpa_onnx_ios - return DynamicLibrary.open('SherpaOnnxC.framework/SherpaOnnxC'); - } - - if (Platform.isAndroid || Platform.isLinux) { - if (_path == null) { - return DynamicLibrary.open('libsherpa-onnx-c-api.so'); - } else { - return DynamicLibrary.open('$_path/libsherpa-onnx-c-api.so'); - } - } - - if (Platform.isWindows) { - if (_path == null) { - return DynamicLibrary.open('sherpa-onnx-c-api.dll'); - } else { - return DynamicLibrary.open('$_path\\sherpa-onnx-c-api.dll'); - } - } - - throw UnsupportedError('Unknown platform: ${Platform.operatingSystem}'); -}(); - /// Initialize the native sherpa-onnx bindings. /// -/// Call this exactly once before using any other API from this package. +/// **Important:** This must be called in every isolate that uses sherpa-onnx. +/// Each isolate has its own FFI binding state, so calling `initBindings()` in +/// one isolate does NOT make sherpa-onnx available in other isolates. If you +/// use Dart isolates for background work (e.g., TTS generation, model loading), +/// call `initBindings()` in each isolate before calling any sherpa-onnx API. /// -/// If [p] is provided, it is treated as the directory containing the native -/// dynamic library for desktop platforms, or the framework root on Apple -/// platforms. If omitted, the package tries to load the library from the -/// default platform-specific filename. +/// On web, use [initBindingsAsync] instead. This method throws +/// [UnsupportedError] on web. void initBindings([String? p]) { + if (kIsWeb) { + throw UnsupportedError( + 'initBindings() is not supported on web. ' + 'Use initBindingsAsync() instead.', + ); + } _path ??= p; - SherpaOnnxBindings.init(_dylib); + init.initNativeBindings(_path); +} + +/// Initialize the sherpa-onnx bindings (works on all platforms including web). +/// +/// On web, this loads the WASM module and JS wrappers automatically. +/// On native platforms, this behaves the same as [initBindings]. +/// +/// **Important:** If you use Dart isolates, call `initBindings()` or +/// `initBindingsAsync()` in each isolate that calls sherpa-onnx APIs. +/// See [initBindings] for details. +Future initBindingsAsync([String? p]) async { + _path ??= p; + if (kIsWeb) { + await web.SherpaOnnxWeb.loadWasm(); + return; + } + init.initNativeBindings(_path); } diff --git a/flutter/sherpa_onnx/lib/src/audio_tagging.dart b/flutter/sherpa_onnx/lib/src/audio_tagging.dart index eb844c75a0..41cb79d409 100644 --- a/flutter/sherpa_onnx/lib/src/audio_tagging.dart +++ b/flutter/sherpa_onnx/lib/src/audio_tagging.dart @@ -3,162 +3,10 @@ import 'dart:ffi'; import 'package:ffi/ffi.dart'; import './offline_stream.dart'; +import './audio_tagging_config.dart'; import './sherpa_onnx_bindings.dart'; -/// Offline audio tagging. -/// -/// This module classifies complete audio clips and returns the most likely -/// events. See `dart-api-examples/audio-tagging/` for working examples. -/// -/// Example: -/// -/// ```dart -/// final modelConfig = AudioTaggingModelConfig( -/// zipformer: const OfflineZipformerAudioTaggingModelConfig( -/// model: './sherpa-onnx-zipformer-audio-tagging/model.int8.onnx', -/// ), -/// numThreads: 1, -/// debug: true, -/// ); -/// -/// final config = AudioTaggingConfig( -/// model: modelConfig, -/// labels: './sherpa-onnx-zipformer-audio-tagging/class_labels_indices.csv', -/// ); -/// -/// final tagger = AudioTagging(config: config); -/// final wave = readWave('./test.wav'); -/// final stream = tagger.createStream(); -/// stream.acceptWaveform(samples: wave.samples, sampleRate: wave.sampleRate); -/// final events = tagger.compute(stream: stream, topK: 5); -/// print(events); -/// stream.free(); -/// tagger.free(); -/// ``` -class OfflineZipformerAudioTaggingModelConfig { - const OfflineZipformerAudioTaggingModelConfig({this.model = ''}); - - factory OfflineZipformerAudioTaggingModelConfig.fromJson( - Map map) { - return OfflineZipformerAudioTaggingModelConfig( - model: map['model'] ?? '', - ); - } - - @override - String toString() { - return 'OfflineZipformerAudioTaggingModelConfig(model: $model)'; - } - - Map toJson() { - return { - 'model': model, - }; - } - - final String model; -} - -/// Aggregate model configuration for audio tagging. -/// -/// Configure either [zipformer] or [ced] for typical use. -class AudioTaggingModelConfig { - AudioTaggingModelConfig( - {this.zipformer = const OfflineZipformerAudioTaggingModelConfig(), - this.ced = '', - this.numThreads = 1, - this.provider = 'cpu', - this.debug = true}); - - factory AudioTaggingModelConfig.fromJson(Map map) { - return AudioTaggingModelConfig( - zipformer: - OfflineZipformerAudioTaggingModelConfig.fromJson(map['zipformer']), - ced: map['ced'] ?? '', - numThreads: map['numThreads'] ?? 1, - provider: map['provider'] ?? 'cpu', - debug: map['debug'] ?? true, - ); - } - - @override - String toString() { - return 'AudioTaggingModelConfig(zipformer: $zipformer, ced: $ced, numThreads: $numThreads, provider: $provider, debug: $debug)'; - } - - Map toJson() { - return { - 'zipformer': zipformer.toJson(), - 'ced': ced, - 'numThreads': numThreads, - 'provider': provider, - 'debug': debug, - }; - } - - final OfflineZipformerAudioTaggingModelConfig zipformer; - final String ced; - final int numThreads; - final String provider; - final bool debug; -} - -/// Top-level configuration for [AudioTagging]. -class AudioTaggingConfig { - AudioTaggingConfig({required this.model, this.labels = ''}); - - factory AudioTaggingConfig.fromJson(Map map) { - return AudioTaggingConfig( - model: AudioTaggingModelConfig.fromJson(map['model']), - labels: map['labels'] ?? '', - ); - } - - @override - String toString() { - return 'AudioTaggingConfig(model: $model, labels: $labels)'; - } - - Map toJson() { - return { - 'model': model.toJson(), - 'labels': labels, - }; - } - - final AudioTaggingModelConfig model; - final String labels; -} - -/// One predicted audio event. -class AudioEvent { - AudioEvent({required this.name, required this.index, required this.prob}); - - factory AudioEvent.fromJson(Map map) { - return AudioEvent( - name: map['name'], - index: map['index'], - prob: map['prob'], - ); - } - - @override - String toString() { - return 'AudioEvent(name: $name, index: $index, prob: $prob)'; - } - - Map toJson() { - return { - 'name': name, - 'index': index, - 'prob': prob, - }; - } - - final String name; - final int index; - final double prob; -} +export './audio_tagging_config.dart'; /// Offline audio tagger. class AudioTagging { diff --git a/flutter/sherpa_onnx/lib/src/audio_tagging_config.dart b/flutter/sherpa_onnx/lib/src/audio_tagging_config.dart new file mode 100644 index 0000000000..953eb1e6fa --- /dev/null +++ b/flutter/sherpa_onnx/lib/src/audio_tagging_config.dart @@ -0,0 +1,157 @@ +// Copyright (c) 2024 Xiaomi Corporation +// Shared config/data classes for audio tagging -- no FFI, works on all platforms. + +/// Offline audio tagging. +/// +/// This module classifies complete audio clips and returns the most likely +/// events. See `dart-api-examples/audio-tagging/` for working examples. +/// +/// Example: +/// +/// ```dart +/// final modelConfig = AudioTaggingModelConfig( +/// zipformer: const OfflineZipformerAudioTaggingModelConfig( +/// model: './sherpa-onnx-zipformer-audio-tagging/model.int8.onnx', +/// ), +/// numThreads: 1, +/// debug: true, +/// ); +/// +/// final config = AudioTaggingConfig( +/// model: modelConfig, +/// labels: './sherpa-onnx-zipformer-audio-tagging/class_labels_indices.csv', +/// ); +/// +/// final tagger = AudioTagging(config: config); +/// final wave = readWave('./test.wav'); +/// final stream = tagger.createStream(); +/// stream.acceptWaveform(samples: wave.samples, sampleRate: wave.sampleRate); +/// final events = tagger.compute(stream: stream, topK: 5); +/// print(events); +/// stream.free(); +/// tagger.free(); +/// ``` +class OfflineZipformerAudioTaggingModelConfig { + const OfflineZipformerAudioTaggingModelConfig({this.model = ''}); + + factory OfflineZipformerAudioTaggingModelConfig.fromJson( + Map map) { + return OfflineZipformerAudioTaggingModelConfig( + model: map['model'] ?? '', + ); + } + + @override + String toString() { + return 'OfflineZipformerAudioTaggingModelConfig(model: $model)'; + } + + Map toJson() { + return { + 'model': model, + }; + } + + final String model; +} + +/// Aggregate model configuration for audio tagging. +/// +/// Configure either [zipformer] or [ced] for typical use. +class AudioTaggingModelConfig { + AudioTaggingModelConfig( + {this.zipformer = const OfflineZipformerAudioTaggingModelConfig(), + this.ced = '', + this.numThreads = 1, + this.provider = 'cpu', + this.debug = true}); + + factory AudioTaggingModelConfig.fromJson(Map map) { + return AudioTaggingModelConfig( + zipformer: + OfflineZipformerAudioTaggingModelConfig.fromJson(map['zipformer']), + ced: map['ced'] ?? '', + numThreads: map['numThreads'] ?? 1, + provider: map['provider'] ?? 'cpu', + debug: map['debug'] ?? true, + ); + } + + @override + String toString() { + return 'AudioTaggingModelConfig(zipformer: $zipformer, ced: $ced, numThreads: $numThreads, provider: $provider, debug: $debug)'; + } + + Map toJson() { + return { + 'zipformer': zipformer.toJson(), + 'ced': ced, + 'numThreads': numThreads, + 'provider': provider, + 'debug': debug, + }; + } + + final OfflineZipformerAudioTaggingModelConfig zipformer; + final String ced; + final int numThreads; + final String provider; + final bool debug; +} + +/// Top-level configuration for [AudioTagging]. +class AudioTaggingConfig { + AudioTaggingConfig({required this.model, this.labels = ''}); + + factory AudioTaggingConfig.fromJson(Map map) { + return AudioTaggingConfig( + model: AudioTaggingModelConfig.fromJson(map['model']), + labels: map['labels'] ?? '', + ); + } + + @override + String toString() { + return 'AudioTaggingConfig(model: $model, labels: $labels)'; + } + + Map toJson() { + return { + 'model': model.toJson(), + 'labels': labels, + }; + } + + final AudioTaggingModelConfig model; + final String labels; +} + +/// One predicted audio event. +class AudioEvent { + AudioEvent({required this.name, required this.index, required this.prob}); + + factory AudioEvent.fromJson(Map map) { + return AudioEvent( + name: map['name'], + index: map['index'], + prob: map['prob'], + ); + } + + @override + String toString() { + return 'AudioEvent(name: $name, index: $index, prob: $prob)'; + } + + Map toJson() { + return { + 'name': name, + 'index': index, + 'prob': prob, + }; + } + + final String name; + final int index; + final double prob; +} diff --git a/flutter/sherpa_onnx/lib/src/init_native.dart b/flutter/sherpa_onnx/lib/src/init_native.dart new file mode 100644 index 0000000000..95bee333f8 --- /dev/null +++ b/flutter/sherpa_onnx/lib/src/init_native.dart @@ -0,0 +1,44 @@ +// Native platform initialization (dart:io available). +import 'dart:io'; +import 'dart:ffi'; + +import 'sherpa_onnx_bindings.dart'; + +DynamicLibrary loadDylib(String? path) { + if (Platform.isMacOS) { + if (path == null) { + return DynamicLibrary.open('SherpaOnnxC.framework/SherpaOnnxC'); + } else { + return DynamicLibrary.open( + '$path/sherpa_onnx_macos/SherpaOnnxC.xcframework/macos-arm64_x86_64/SherpaOnnxC.framework/SherpaOnnxC', + ); + } + } + + if (Platform.isIOS) { + return DynamicLibrary.open('SherpaOnnxC.framework/SherpaOnnxC'); + } + + if (Platform.isAndroid || Platform.isLinux) { + if (path == null) { + return DynamicLibrary.open('libsherpa-onnx-c-api.so'); + } else { + return DynamicLibrary.open('$path/libsherpa-onnx-c-api.so'); + } + } + + if (Platform.isWindows) { + if (path == null) { + return DynamicLibrary.open('sherpa-onnx-c-api.dll'); + } else { + return DynamicLibrary.open('$path\\sherpa-onnx-c-api.dll'); + } + } + + throw UnsupportedError('Unknown platform: ${Platform.operatingSystem}'); +} + +void initNativeBindings(String? path) { + final dylib = loadDylib(path); + SherpaOnnxBindings.init(dylib); +} diff --git a/flutter/sherpa_onnx/lib/src/init_stub.dart b/flutter/sherpa_onnx/lib/src/init_stub.dart new file mode 100644 index 0000000000..9f5134c1f2 --- /dev/null +++ b/flutter/sherpa_onnx/lib/src/init_stub.dart @@ -0,0 +1,5 @@ +// Native stub for sherpa_onnx_web. +// On native platforms, SherpaOnnxWeb is not needed. +class SherpaOnnxWeb { + static Future loadWasm() async {} +} diff --git a/flutter/sherpa_onnx/lib/src/keyword_spotter.dart b/flutter/sherpa_onnx/lib/src/keyword_spotter.dart index 18744aac41..36a116fdcc 100644 --- a/flutter/sherpa_onnx/lib/src/keyword_spotter.dart +++ b/flutter/sherpa_onnx/lib/src/keyword_spotter.dart @@ -4,114 +4,12 @@ import 'dart:ffi'; import 'package:ffi/ffi.dart'; -import './feature_config.dart'; import './online_stream.dart'; -import './online_recognizer.dart'; import './sherpa_onnx_bindings.dart'; import './utils.dart'; +import './keyword_spotter_config.dart'; -/// Streaming keyword spotting. -/// -/// See `dart-api-examples/keyword-spotter/` for end-to-end usage. -/// -/// Example: -/// -/// ```dart -/// final spotter = KeywordSpotter( -/// KeywordSpotterConfig( -/// model: onlineModelConfig, -/// keywordsFile: './keywords.txt', -/// ), -/// ); -/// -/// final stream = spotter.createStream(); -/// stream.acceptWaveform(samples: chunk, sampleRate: 16000); -/// while (spotter.isReady(stream)) { -/// spotter.decode(stream); -/// } -/// print(spotter.getResult(stream).keyword); -/// ``` -class KeywordSpotterConfig { - const KeywordSpotterConfig({ - this.feat = const FeatureConfig(), - required this.model, - this.maxActivePaths = 4, - this.numTrailingBlanks = 1, - this.keywordsScore = 1.0, - this.keywordsThreshold = 0.25, - this.keywordsFile = '', - this.keywordsBuf = '', - this.keywordsBufSize = 0, - }); - - factory KeywordSpotterConfig.fromJson(Map json) { - return KeywordSpotterConfig( - feat: json['feat'] != null - ? FeatureConfig.fromJson(json['feat'] as Map) - : const FeatureConfig(), - model: OnlineModelConfig.fromJson(json['model'] as Map), - maxActivePaths: json['maxActivePaths'] as int? ?? 4, - numTrailingBlanks: json['numTrailingBlanks'] as int? ?? 1, - keywordsScore: (json['keywordsScore'] as num?)?.toDouble() ?? 1.0, - keywordsThreshold: - (json['keywordsThreshold'] as num?)?.toDouble() ?? 0.25, - keywordsFile: json['keywordsFile'] as String? ?? '', - keywordsBuf: json['keywordsBuf'] as String? ?? '', - keywordsBufSize: json['keywordsBufSize'] as int? ?? 0, - ); - } - - @override - String toString() { - return 'KeywordSpotterConfig(feat: $feat, model: $model, maxActivePaths: $maxActivePaths, numTrailingBlanks: $numTrailingBlanks, keywordsScore: $keywordsScore, keywordsThreshold: $keywordsThreshold, keywordsFile: $keywordsFile, keywordsBuf: $keywordsBuf, keywordsBufSize: $keywordsBufSize)'; - } - - Map toJson() => { - 'feat': feat.toJson(), - 'model': model.toJson(), - 'maxActivePaths': maxActivePaths, - 'numTrailingBlanks': numTrailingBlanks, - 'keywordsScore': keywordsScore, - 'keywordsThreshold': keywordsThreshold, - 'keywordsFile': keywordsFile, - 'keywordsBuf': keywordsBuf, - 'keywordsBufSize': keywordsBufSize, - }; - - final FeatureConfig feat; - final OnlineModelConfig model; - - final int maxActivePaths; - final int numTrailingBlanks; - - final double keywordsScore; - final double keywordsThreshold; - final String keywordsFile; - final String keywordsBuf; - final int keywordsBufSize; -} - -/// Result returned by [KeywordSpotter.getResult]. -class KeywordResult { - KeywordResult({required this.keyword}); - - factory KeywordResult.fromJson(Map json) { - return KeywordResult( - keyword: json['keyword'] as String? ?? '', - ); - } - - @override - String toString() { - return 'KeywordResult(keyword: $keyword)'; - } - - Map toJson() => { - 'keyword': keyword, - }; - - final String keyword; -} +export './keyword_spotter_config.dart'; /// Streaming keyword spotter. class KeywordSpotter { diff --git a/flutter/sherpa_onnx/lib/src/keyword_spotter_config.dart b/flutter/sherpa_onnx/lib/src/keyword_spotter_config.dart new file mode 100644 index 0000000000..62e97d3ed8 --- /dev/null +++ b/flutter/sherpa_onnx/lib/src/keyword_spotter_config.dart @@ -0,0 +1,108 @@ +// Copyright (c) 2024 Xiaomi Corporation +// Shared config/data classes for keyword spotting -- no FFI, works on all platforms. + +import './feature_config.dart'; +import './online_recognizer_config.dart'; + +/// Streaming keyword spotting. +/// +/// See `dart-api-examples/keyword-spotter/` for end-to-end usage. +/// +/// Example: +/// +/// ```dart +/// final spotter = KeywordSpotter( +/// KeywordSpotterConfig( +/// model: onlineModelConfig, +/// keywordsFile: './keywords.txt', +/// ), +/// ); +/// +/// final stream = spotter.createStream(); +/// stream.acceptWaveform(samples: chunk, sampleRate: 16000); +/// while (spotter.isReady(stream)) { +/// spotter.decode(stream); +/// } +/// print(spotter.getResult(stream).keyword); +/// ``` +class KeywordSpotterConfig { + const KeywordSpotterConfig({ + this.feat = const FeatureConfig(), + required this.model, + this.maxActivePaths = 4, + this.numTrailingBlanks = 1, + this.keywordsScore = 1.0, + this.keywordsThreshold = 0.25, + this.keywordsFile = '', + this.keywordsBuf = '', + this.keywordsBufSize = 0, + }); + + factory KeywordSpotterConfig.fromJson(Map json) { + return KeywordSpotterConfig( + feat: json['feat'] != null + ? FeatureConfig.fromJson(json['feat'] as Map) + : const FeatureConfig(), + model: OnlineModelConfig.fromJson(json['model'] as Map), + maxActivePaths: json['maxActivePaths'] as int? ?? 4, + numTrailingBlanks: json['numTrailingBlanks'] as int? ?? 1, + keywordsScore: (json['keywordsScore'] as num?)?.toDouble() ?? 1.0, + keywordsThreshold: + (json['keywordsThreshold'] as num?)?.toDouble() ?? 0.25, + keywordsFile: json['keywordsFile'] as String? ?? '', + keywordsBuf: json['keywordsBuf'] as String? ?? '', + keywordsBufSize: json['keywordsBufSize'] as int? ?? 0, + ); + } + + @override + String toString() { + return 'KeywordSpotterConfig(feat: $feat, model: $model, maxActivePaths: $maxActivePaths, numTrailingBlanks: $numTrailingBlanks, keywordsScore: $keywordsScore, keywordsThreshold: $keywordsThreshold, keywordsFile: $keywordsFile, keywordsBuf: $keywordsBuf, keywordsBufSize: $keywordsBufSize)'; + } + + Map toJson() => { + 'feat': feat.toJson(), + 'model': model.toJson(), + 'maxActivePaths': maxActivePaths, + 'numTrailingBlanks': numTrailingBlanks, + 'keywordsScore': keywordsScore, + 'keywordsThreshold': keywordsThreshold, + 'keywordsFile': keywordsFile, + 'keywordsBuf': keywordsBuf, + 'keywordsBufSize': keywordsBufSize, + }; + + final FeatureConfig feat; + final OnlineModelConfig model; + + final int maxActivePaths; + final int numTrailingBlanks; + + final double keywordsScore; + final double keywordsThreshold; + final String keywordsFile; + final String keywordsBuf; + final int keywordsBufSize; +} + +/// Result returned by [KeywordSpotter.getResult]. +class KeywordResult { + KeywordResult({required this.keyword}); + + factory KeywordResult.fromJson(Map json) { + return KeywordResult( + keyword: json['keyword'] as String? ?? '', + ); + } + + @override + String toString() { + return 'KeywordResult(keyword: $keyword)'; + } + + Map toJson() => { + 'keyword': keyword, + }; + + final String keyword; +} diff --git a/flutter/sherpa_onnx/lib/src/offline_punctuation.dart b/flutter/sherpa_onnx/lib/src/offline_punctuation.dart index 915496a60c..9997d4c0f8 100644 --- a/flutter/sherpa_onnx/lib/src/offline_punctuation.dart +++ b/flutter/sherpa_onnx/lib/src/offline_punctuation.dart @@ -3,69 +3,9 @@ import 'dart:ffi'; import 'package:ffi/ffi.dart'; import './sherpa_onnx_bindings.dart'; +import './offline_punctuation_config.dart'; -/// Offline punctuation restoration. -/// -/// This is intended for complete text strings when you want one-shot -/// punctuation insertion. See `dart-api-examples/add-punctuations/`. -class OfflinePunctuationModelConfig { - OfflinePunctuationModelConfig( - {required this.ctTransformer, - this.numThreads = 1, - this.provider = 'cpu', - this.debug = true}); - - factory OfflinePunctuationModelConfig.fromJson(Map json) { - return OfflinePunctuationModelConfig( - ctTransformer: json['ctTransformer'] as String, - numThreads: json['numThreads'] as int? ?? 1, - provider: json['provider'] as String? ?? 'cpu', - debug: json['debug'] as bool? ?? true, - ); - } - - @override - String toString() { - return 'OfflinePunctuationModelConfig(ctTransformer: $ctTransformer, numThreads: $numThreads, provider: $provider, debug: $debug)'; - } - - Map toJson() => { - 'ctTransformer': ctTransformer, - 'numThreads': numThreads, - 'provider': provider, - 'debug': debug, - }; - - final String ctTransformer; - final int numThreads; - final String provider; - final bool debug; -} - -/// Top-level configuration for [OfflinePunctuation]. -class OfflinePunctuationConfig { - OfflinePunctuationConfig({ - required this.model, - }); - - factory OfflinePunctuationConfig.fromJson(Map json) { - return OfflinePunctuationConfig( - model: OfflinePunctuationModelConfig.fromJson( - json['model'] as Map), - ); - } - - @override - String toString() { - return 'OfflinePunctuationConfig(model: $model)'; - } - - Map toJson() => { - 'model': model.toJson(), - }; - - final OfflinePunctuationModelConfig model; -} +export './offline_punctuation_config.dart'; /// Offline punctuation restorer. class OfflinePunctuation { diff --git a/flutter/sherpa_onnx/lib/src/offline_punctuation_config.dart b/flutter/sherpa_onnx/lib/src/offline_punctuation_config.dart new file mode 100644 index 0000000000..fead287214 --- /dev/null +++ b/flutter/sherpa_onnx/lib/src/offline_punctuation_config.dart @@ -0,0 +1,65 @@ +// Copyright (c) 2024 Xiaomi Corporation +// Shared config/data classes for offline punctuation -- no FFI, works on all platforms. + +/// Offline punctuation restoration. +/// +/// This is intended for complete text strings when you want one-shot +/// punctuation insertion. See `dart-api-examples/add-punctuations/`. +class OfflinePunctuationModelConfig { + OfflinePunctuationModelConfig( + {required this.ctTransformer, + this.numThreads = 1, + this.provider = 'cpu', + this.debug = true}); + + factory OfflinePunctuationModelConfig.fromJson(Map json) { + return OfflinePunctuationModelConfig( + ctTransformer: json['ctTransformer'] as String, + numThreads: json['numThreads'] as int? ?? 1, + provider: json['provider'] as String? ?? 'cpu', + debug: json['debug'] as bool? ?? true, + ); + } + + @override + String toString() { + return 'OfflinePunctuationModelConfig(ctTransformer: $ctTransformer, numThreads: $numThreads, provider: $provider, debug: $debug)'; + } + + Map toJson() => { + 'ctTransformer': ctTransformer, + 'numThreads': numThreads, + 'provider': provider, + 'debug': debug, + }; + + final String ctTransformer; + final int numThreads; + final String provider; + final bool debug; +} + +/// Top-level configuration for [OfflinePunctuation]. +class OfflinePunctuationConfig { + OfflinePunctuationConfig({ + required this.model, + }); + + factory OfflinePunctuationConfig.fromJson(Map json) { + return OfflinePunctuationConfig( + model: OfflinePunctuationModelConfig.fromJson( + json['model'] as Map), + ); + } + + @override + String toString() { + return 'OfflinePunctuationConfig(model: $model)'; + } + + Map toJson() => { + 'model': model.toJson(), + }; + + final OfflinePunctuationModelConfig model; +} diff --git a/flutter/sherpa_onnx/lib/src/offline_recognizer.dart b/flutter/sherpa_onnx/lib/src/offline_recognizer.dart index 1f0cb20d5b..ed018fd61e 100644 --- a/flutter/sherpa_onnx/lib/src/offline_recognizer.dart +++ b/flutter/sherpa_onnx/lib/src/offline_recognizer.dart @@ -4,11 +4,12 @@ import 'dart:ffi'; import 'package:ffi/ffi.dart'; -import './feature_config.dart'; -import './homophone_replacer_config.dart'; import './offline_stream.dart'; import './sherpa_onnx_bindings.dart'; import './utils.dart'; +import './offline_recognizer_config.dart'; + +export './offline_recognizer_config.dart'; /// Offline speech recognition. /// @@ -45,930 +46,6 @@ import './utils.dart'; /// recognizer.free(); /// ``` -/// Model files for an offline transducer recognizer. -/// -/// This family is also used by NeMo Parakeet TDT-style examples. -class OfflineTransducerModelConfig { - const OfflineTransducerModelConfig({ - this.encoder = '', - this.decoder = '', - this.joiner = '', - }); - - factory OfflineTransducerModelConfig.fromJson(Map json) { - return OfflineTransducerModelConfig( - encoder: json['encoder'] as String? ?? '', - decoder: json['decoder'] as String? ?? '', - joiner: json['joiner'] as String? ?? '', - ); - } - - @override - String toString() { - return 'OfflineTransducerModelConfig(encoder: $encoder, decoder: $decoder, joiner: $joiner)'; - } - - Map toJson() => { - 'encoder': encoder, - 'decoder': decoder, - 'joiner': joiner, - }; - - final String encoder; - final String decoder; - final String joiner; -} - -/// Model files for an offline Paraformer recognizer. -class OfflineParaformerModelConfig { - const OfflineParaformerModelConfig({this.model = ''}); - - factory OfflineParaformerModelConfig.fromJson(Map json) { - return OfflineParaformerModelConfig(model: json['model'] as String? ?? ''); - } - - @override - String toString() { - return 'OfflineParaformerModelConfig(model: $model)'; - } - - Map toJson() => {'model': model}; - - final String model; -} - -/// Model files for an offline NeMo CTC recognizer. -class OfflineNemoEncDecCtcModelConfig { - const OfflineNemoEncDecCtcModelConfig({this.model = ''}); - - factory OfflineNemoEncDecCtcModelConfig.fromJson(Map json) { - return OfflineNemoEncDecCtcModelConfig( - model: json['model'] as String? ?? '', - ); - } - - @override - String toString() { - return 'OfflineNemoEncDecCtcModelConfig(model: $model)'; - } - - Map toJson() => {'model': model}; - - final String model; -} - -/// Model files for an offline Dolphin recognizer. -class OfflineDolphinModelConfig { - const OfflineDolphinModelConfig({this.model = ''}); - - factory OfflineDolphinModelConfig.fromJson(Map json) { - return OfflineDolphinModelConfig(model: json['model'] as String? ?? ''); - } - - @override - String toString() { - return 'OfflineDolphinModelConfig(model: $model)'; - } - - Map toJson() => {'model': model}; - - final String model; -} - -/// Model files for an offline Zipformer CTC recognizer. -class OfflineZipformerCtcModelConfig { - const OfflineZipformerCtcModelConfig({this.model = ''}); - - factory OfflineZipformerCtcModelConfig.fromJson(Map json) { - return OfflineZipformerCtcModelConfig( - model: json['model'] as String? ?? '', - ); - } - - @override - String toString() { - return 'OfflineZipformerCtcModelConfig(model: $model)'; - } - - Map toJson() => {'model': model}; - - final String model; -} - -/// Model files for an offline WeNet CTC recognizer. -class OfflineWenetCtcModelConfig { - const OfflineWenetCtcModelConfig({this.model = ''}); - - factory OfflineWenetCtcModelConfig.fromJson(Map json) { - return OfflineWenetCtcModelConfig(model: json['model'] as String? ?? ''); - } - - @override - String toString() { - return 'OfflineWenetCtcModelConfig(model: $model)'; - } - - Map toJson() => {'model': model}; - - final String model; -} - -/// Model files for the omnilingual ASR CTC recognizer. -class OfflineOmnilingualAsrCtcModelConfig { - const OfflineOmnilingualAsrCtcModelConfig({this.model = ''}); - - factory OfflineOmnilingualAsrCtcModelConfig.fromJson( - Map json, - ) { - return OfflineOmnilingualAsrCtcModelConfig( - model: json['model'] as String? ?? '', - ); - } - - @override - String toString() { - return 'OfflineOmnilingualAsrCtcModelConfig(model: $model)'; - } - - Map toJson() => {'model': model}; - - final String model; -} - -/// Model files for the MedASR CTC recognizer. -class OfflineMedAsrCtcModelConfig { - const OfflineMedAsrCtcModelConfig({this.model = ''}); - - factory OfflineMedAsrCtcModelConfig.fromJson(Map json) { - return OfflineMedAsrCtcModelConfig(model: json['model'] as String? ?? ''); - } - - @override - String toString() { - return 'OfflineMedAsrCtcModelConfig(model: $model)'; - } - - Map toJson() => {'model': model}; - - final String model; -} - -/// Model files for the Fire-Red-ASR CTC recognizer. -class OfflineFireRedAsrCtcModelConfig { - const OfflineFireRedAsrCtcModelConfig({this.model = ''}); - - factory OfflineFireRedAsrCtcModelConfig.fromJson(Map json) { - return OfflineFireRedAsrCtcModelConfig( - model: json['model'] as String? ?? '', - ); - } - - @override - String toString() { - return 'OfflineFireRedAsrCtcModelConfig(model: $model)'; - } - - Map toJson() => {'model': model}; - - final String model; -} - -/// Model files and prompt settings for FunASR-Nano. -class OfflineFunAsrNanoModelConfig { - const OfflineFunAsrNanoModelConfig({ - this.encoderAdaptor = '', - this.llm = '', - this.embedding = '', - this.tokenizer = '', - this.systemPrompt = 'You are a helpful assistant.', - this.userPrompt = '语音转写:', - this.maxNewTokens = 512, - this.temperature = 1e-6, - this.topP = 0.8, - this.seed = 42, - this.language = '', - this.itn = 1, - this.hotwords = '', - }); - - factory OfflineFunAsrNanoModelConfig.fromJson(Map json) { - return OfflineFunAsrNanoModelConfig( - encoderAdaptor: json['encoderAdaptor'] as String? ?? '', - llm: json['llm'] as String? ?? '', - embedding: json['embedding'] as String? ?? '', - tokenizer: json['tokenizer'] as String? ?? '', - systemPrompt: json['systemPrompt'] as String? ?? '', - userPrompt: json['userPrompt'] as String? ?? '', - maxNewTokens: json['maxNewTokens'] as int? ?? 512, - temperature: (json['temperature'] as num?)?.toDouble() ?? 1e-6, - topP: (json['topP'] as num?)?.toDouble() ?? 0.8, - seed: json['seed'] as int? ?? 42, - language: json['language'] as String? ?? '', - itn: json['itn'] as int? ?? 1, - hotwords: json['hotwords'] as String? ?? '', - ); - } - - @override - String toString() { - return 'OfflineFunAsrNanoModelConfig(encoderAdaptor: $encoderAdaptor, llm: $llm, embedding: $embedding, tokenizer: $tokenizer, systemPrompt: $systemPrompt, userPrompt: $userPrompt, maxNewTokens: $maxNewTokens, temperature: $temperature, topP: $topP, seed: $seed, language: $language, itn: $itn, hotwords: $hotwords)'; - } - - Map toJson() => { - 'encoderAdaptor': encoderAdaptor, - 'llm': llm, - 'embedding': embedding, - 'tokenizer': tokenizer, - 'systemPrompt': systemPrompt, - 'userPrompt': userPrompt, - 'maxNewTokens': maxNewTokens, - 'temperature': temperature, - 'topP': topP, - 'seed': seed, - 'language': language, - 'itn': itn, - 'hotwords': hotwords, - }; - - final String encoderAdaptor; - final String llm; - final String embedding; - final String tokenizer; - final String systemPrompt; - final String userPrompt; - final int maxNewTokens; - final double temperature; - final double topP; - final int seed; - final String language; - final int itn; - final String hotwords; -} - -class OfflineQwen3AsrModelConfig { - const OfflineQwen3AsrModelConfig({ - this.convFrontend = '', - this.encoder = '', - this.decoder = '', - this.tokenizer = '', - this.maxTotalLen = 512, - this.maxNewTokens = 128, - this.temperature = 1e-6, - this.topP = 0.8, - this.seed = 42, - this.hotwords = '', - }); - - factory OfflineQwen3AsrModelConfig.fromJson(Map json) { - return OfflineQwen3AsrModelConfig( - convFrontend: json['convFrontend'] as String? ?? '', - encoder: json['encoder'] as String? ?? '', - decoder: json['decoder'] as String? ?? '', - tokenizer: json['tokenizer'] as String? ?? '', - maxTotalLen: json['maxTotalLen'] as int? ?? 512, - maxNewTokens: json['maxNewTokens'] as int? ?? 128, - temperature: (json['temperature'] as num?)?.toDouble() ?? 1e-6, - topP: (json['topP'] as num?)?.toDouble() ?? 0.8, - seed: json['seed'] as int? ?? 42, - hotwords: json['hotwords'] as String? ?? '', - ); - } - - @override - String toString() { - return 'OfflineQwen3AsrModelConfig(convFrontend: $convFrontend, encoder: $encoder, decoder: $decoder, tokenizer: $tokenizer, maxTotalLen: $maxTotalLen, maxNewTokens: $maxNewTokens, temperature: $temperature, topP: $topP, seed: $seed, hotwords: $hotwords)'; - } - - Map toJson() => { - 'convFrontend': convFrontend, - 'encoder': encoder, - 'decoder': decoder, - 'tokenizer': tokenizer, - 'maxTotalLen': maxTotalLen, - 'maxNewTokens': maxNewTokens, - 'temperature': temperature, - 'topP': topP, - 'seed': seed, - 'hotwords': hotwords, - }; - - final String convFrontend; - final String encoder; - final String decoder; - final String tokenizer; - final int maxTotalLen; - final int maxNewTokens; - final double temperature; - final double topP; - final int seed; - final String hotwords; -} - -/// Model files and options for an offline Whisper recognizer. -class OfflineWhisperModelConfig { - const OfflineWhisperModelConfig({ - this.encoder = '', - this.decoder = '', - this.language = '', - this.task = '', - this.tailPaddings = -1, - this.enableTokenTimestamps = false, - this.enableSegmentTimestamps = false, - }); - - factory OfflineWhisperModelConfig.fromJson(Map json) { - return OfflineWhisperModelConfig( - encoder: json['encoder'] as String? ?? '', - decoder: json['decoder'] as String? ?? '', - language: json['language'] as String? ?? '', - task: json['task'] as String? ?? '', - tailPaddings: json['tailPaddings'] as int? ?? -1, - enableTokenTimestamps: json['enableTokenTimestamps'] as bool? ?? false, - enableSegmentTimestamps: - json['enableSegmentTimestamps'] as bool? ?? false, - ); - } - - @override - String toString() { - return 'OfflineWhisperModelConfig(encoder: $encoder, decoder: $decoder, language: $language, task: $task, tailPaddings: $tailPaddings, enableTokenTimestamps: $enableTokenTimestamps, enableSegmentTimestamps: $enableSegmentTimestamps)'; - } - - Map toJson() => { - 'encoder': encoder, - 'decoder': decoder, - 'language': language, - 'task': task, - 'tailPaddings': tailPaddings, - 'enableTokenTimestamps': enableTokenTimestamps, - 'enableSegmentTimestamps': enableSegmentTimestamps, - }; - - final String encoder; - final String decoder; - final String language; - final String task; - final int tailPaddings; - final bool enableTokenTimestamps; - final bool enableSegmentTimestamps; -} - -/// Model files and translation options for NeMo Canary. -class OfflineCanaryModelConfig { - const OfflineCanaryModelConfig({ - this.encoder = '', - this.decoder = '', - this.srcLang = 'en', - this.tgtLang = 'en', - this.usePnc = true, - }); - - factory OfflineCanaryModelConfig.fromJson(Map json) { - return OfflineCanaryModelConfig( - encoder: json['encoder'] as String? ?? '', - decoder: json['decoder'] as String? ?? '', - srcLang: json['srcLang'] as String? ?? 'en', - tgtLang: json['tgtLang'] as String? ?? 'en', - usePnc: json['usePnc'] as bool? ?? true, - ); - } - - @override - String toString() { - return 'OfflineCanaryModelConfig(encoder: $encoder, decoder: $decoder, srcLang: $srcLang, tgtLang: $tgtLang, usePnc: $usePnc)'; - } - - Map toJson() => { - 'encoder': encoder, - 'decoder': decoder, - 'srcLang': srcLang, - 'tgtLang': tgtLang, - 'usePnc': usePnc, - }; - - final String encoder; - final String decoder; - final String srcLang; - final String tgtLang; - final bool usePnc; -} - -/// Model files and text options for Cohere Transcribe. -class OfflineCohereTranscribeModelConfig { - const OfflineCohereTranscribeModelConfig({ - this.encoder = '', - this.decoder = '', - this.language = '', - this.usePunct = true, - this.useItn = true, - }); - - factory OfflineCohereTranscribeModelConfig.fromJson( - Map json, - ) { - return OfflineCohereTranscribeModelConfig( - encoder: json['encoder'] as String? ?? '', - decoder: json['decoder'] as String? ?? '', - language: json['language'] as String? ?? '', - usePunct: json['usePunct'] as bool? ?? true, - useItn: json['useItn'] as bool? ?? true, - ); - } - - @override - String toString() { - return 'OfflineCohereTranscribeModelConfig(encoder: $encoder, decoder: $decoder, language: $language, usePunct: $usePunct, useItn: $useItn)'; - } - - Map toJson() => { - 'encoder': encoder, - 'decoder': decoder, - 'language': language, - 'usePunct': usePunct, - 'useItn': useItn, - }; - - final String encoder; - final String decoder; - final String language; - final bool usePunct; - final bool useItn; -} - -/// Model files for the Fire-Red-ASR transducer recognizer. -class OfflineFireRedAsrModelConfig { - const OfflineFireRedAsrModelConfig({this.encoder = '', this.decoder = ''}); - - factory OfflineFireRedAsrModelConfig.fromJson(Map json) { - return OfflineFireRedAsrModelConfig( - encoder: json['encoder'] as String? ?? '', - decoder: json['decoder'] as String? ?? '', - ); - } - - @override - String toString() { - return 'OfflineFireRedAsrModelConfig(encoder: $encoder, decoder: $decoder)'; - } - - Map toJson() => {'encoder': encoder, 'decoder': decoder}; - - final String encoder; - final String decoder; -} - -// For Moonshine v1, you need 4 models: -// - preprocessor, encoder, uncachedDecoder, cachedDecoder -// -// For Moonshine v2, you need 2 models: -// - encoder, mergedDecoder -/// Model files for Moonshine v1 or v2. -class OfflineMoonshineModelConfig { - const OfflineMoonshineModelConfig({ - this.preprocessor = '', - this.encoder = '', - this.uncachedDecoder = '', - this.cachedDecoder = '', - this.mergedDecoder = '', - }); - - factory OfflineMoonshineModelConfig.fromJson(Map json) { - return OfflineMoonshineModelConfig( - preprocessor: json['preprocessor'] as String? ?? '', - encoder: json['encoder'] as String? ?? '', - uncachedDecoder: json['uncachedDecoder'] as String? ?? '', - cachedDecoder: json['cachedDecoder'] as String? ?? '', - mergedDecoder: json['mergedDecoder'] as String? ?? '', - ); - } - - @override - String toString() { - return 'OfflineMoonshineModelConfig(preprocessor: $preprocessor, encoder: $encoder, uncachedDecoder: $uncachedDecoder, cachedDecoder: $cachedDecoder, mergedDecoder: $mergedDecoder)'; - } - - Map toJson() => { - 'preprocessor': preprocessor, - 'encoder': encoder, - 'uncachedDecoder': uncachedDecoder, - 'cachedDecoder': cachedDecoder, - 'mergedDecoder': mergedDecoder, - }; - - final String preprocessor; - final String encoder; - final String uncachedDecoder; - final String cachedDecoder; - final String mergedDecoder; -} - -/// Model files for an offline TDNN recognizer. -class OfflineTdnnModelConfig { - const OfflineTdnnModelConfig({this.model = ''}); - - factory OfflineTdnnModelConfig.fromJson(Map json) { - return OfflineTdnnModelConfig(model: json['model'] as String? ?? ''); - } - - @override - String toString() { - return 'OfflineTdnnModelConfig(model: $model)'; - } - - Map toJson() => {'model': model}; - - final String model; -} - -/// Model files and options for SenseVoice. -/// -/// In the examples, this is typically paired with the -/// `sherpa-onnx-sense-voice-zh-en-ja-ko-yue-2024-07-17-int8` package. -class OfflineSenseVoiceModelConfig { - const OfflineSenseVoiceModelConfig({ - this.model = '', - this.language = '', - this.useInverseTextNormalization = false, - }); - - factory OfflineSenseVoiceModelConfig.fromJson(Map json) { - return OfflineSenseVoiceModelConfig( - model: json['model'] as String? ?? '', - language: json['language'] as String? ?? '', - useInverseTextNormalization: - json['useInverseTextNormalization'] as bool? ?? false, - ); - } - - @override - String toString() { - return 'OfflineSenseVoiceModelConfig(model: $model, language: $language, useInverseTextNormalization: $useInverseTextNormalization)'; - } - - Map toJson() => { - 'model': model, - 'language': language, - 'useInverseTextNormalization': useInverseTextNormalization, - }; - - final String model; - final String language; - final bool useInverseTextNormalization; -} - -/// Optional external language model settings for offline ASR. -class OfflineLMConfig { - const OfflineLMConfig({this.model = '', this.scale = 1.0}); - - factory OfflineLMConfig.fromJson(Map json) { - return OfflineLMConfig( - model: json['model'] as String? ?? '', - scale: (json['scale'] as num?)?.toDouble() ?? 1.0, - ); - } - - @override - String toString() { - return 'OfflineLMConfig(model: $model, scale: $scale)'; - } - - Map toJson() => {'model': model, 'scale': scale}; - - final String model; - final double scale; -} - -/// Aggregate model configuration for offline recognition. -/// -/// In typical use, configure exactly one model family and set the shared -/// options such as [tokens], [provider], and [numThreads]. -/// -/// For NeMo Parakeet-style transducer models, set [modelType] to -/// `nemo_transducer`, matching the repository examples. -class OfflineModelConfig { - const OfflineModelConfig({ - this.transducer = const OfflineTransducerModelConfig(), - this.paraformer = const OfflineParaformerModelConfig(), - this.nemoCtc = const OfflineNemoEncDecCtcModelConfig(), - this.whisper = const OfflineWhisperModelConfig(), - this.tdnn = const OfflineTdnnModelConfig(), - this.senseVoice = const OfflineSenseVoiceModelConfig(), - this.moonshine = const OfflineMoonshineModelConfig(), - this.fireRedAsr = const OfflineFireRedAsrModelConfig(), - this.dolphin = const OfflineDolphinModelConfig(), - this.zipformerCtc = const OfflineZipformerCtcModelConfig(), - this.canary = const OfflineCanaryModelConfig(), - this.wenetCtc = const OfflineWenetCtcModelConfig(), - this.omnilingual = const OfflineOmnilingualAsrCtcModelConfig(), - this.medasr = const OfflineMedAsrCtcModelConfig(), - this.funasrNano = const OfflineFunAsrNanoModelConfig(), - this.fireRedAsrCtc = const OfflineFireRedAsrCtcModelConfig(), - this.qwen3Asr = const OfflineQwen3AsrModelConfig(), - this.cohereTranscribe = const OfflineCohereTranscribeModelConfig(), - required this.tokens, - this.numThreads = 1, - this.debug = true, - this.provider = 'cpu', - this.modelType = '', - this.modelingUnit = '', - this.bpeVocab = '', - this.telespeechCtc = '', - }); - - factory OfflineModelConfig.fromJson(Map json) { - return OfflineModelConfig( - transducer: json['transducer'] != null - ? OfflineTransducerModelConfig.fromJson( - json['transducer'] as Map, - ) - : const OfflineTransducerModelConfig(), - paraformer: json['paraformer'] != null - ? OfflineParaformerModelConfig.fromJson( - json['paraformer'] as Map, - ) - : const OfflineParaformerModelConfig(), - nemoCtc: json['nemoCtc'] != null - ? OfflineNemoEncDecCtcModelConfig.fromJson( - json['nemoCtc'] as Map, - ) - : const OfflineNemoEncDecCtcModelConfig(), - whisper: json['whisper'] != null - ? OfflineWhisperModelConfig.fromJson( - json['whisper'] as Map, - ) - : const OfflineWhisperModelConfig(), - tdnn: json['tdnn'] != null - ? OfflineTdnnModelConfig.fromJson( - json['tdnn'] as Map, - ) - : const OfflineTdnnModelConfig(), - senseVoice: json['senseVoice'] != null - ? OfflineSenseVoiceModelConfig.fromJson( - json['senseVoice'] as Map, - ) - : const OfflineSenseVoiceModelConfig(), - moonshine: json['moonshine'] != null - ? OfflineMoonshineModelConfig.fromJson( - json['moonshine'] as Map, - ) - : const OfflineMoonshineModelConfig(), - fireRedAsr: json['fireRedAsr'] != null - ? OfflineFireRedAsrModelConfig.fromJson( - json['fireRedAsr'] as Map, - ) - : const OfflineFireRedAsrModelConfig(), - dolphin: json['dolphin'] != null - ? OfflineDolphinModelConfig.fromJson( - json['dolphin'] as Map, - ) - : const OfflineDolphinModelConfig(), - zipformerCtc: json['zipformerCtc'] != null - ? OfflineZipformerCtcModelConfig.fromJson( - json['zipformerCtc'] as Map, - ) - : const OfflineZipformerCtcModelConfig(), - canary: json['canary'] != null - ? OfflineCanaryModelConfig.fromJson( - json['canary'] as Map, - ) - : const OfflineCanaryModelConfig(), - wenetCtc: json['wenetCtc'] != null - ? OfflineWenetCtcModelConfig.fromJson( - json['wenetCtc'] as Map, - ) - : const OfflineWenetCtcModelConfig(), - omnilingual: json['omnilingual'] != null - ? OfflineOmnilingualAsrCtcModelConfig.fromJson( - json['omnilingual'] as Map, - ) - : const OfflineOmnilingualAsrCtcModelConfig(), - medasr: json['medasr'] != null - ? OfflineMedAsrCtcModelConfig.fromJson( - json['medasr'] as Map, - ) - : const OfflineMedAsrCtcModelConfig(), - funasrNano: json['funasrNano'] != null - ? OfflineFunAsrNanoModelConfig.fromJson( - json['funasrNano'] as Map, - ) - : const OfflineFunAsrNanoModelConfig(), - fireRedAsrCtc: json['fireRedAsrCtc'] != null - ? OfflineFireRedAsrCtcModelConfig.fromJson( - json['fireRedAsrCtc'] as Map, - ) - : const OfflineFireRedAsrCtcModelConfig(), - qwen3Asr: json['qwen3Asr'] != null - ? OfflineQwen3AsrModelConfig.fromJson( - json['qwen3Asr'] as Map, - ) - : const OfflineQwen3AsrModelConfig(), - cohereTranscribe: json['cohereTranscribe'] != null - ? OfflineCohereTranscribeModelConfig.fromJson( - json['cohereTranscribe'] as Map, - ) - : const OfflineCohereTranscribeModelConfig(), - tokens: json['tokens'] as String, - numThreads: json['numThreads'] as int? ?? 1, - debug: json['debug'] as bool? ?? true, - provider: json['provider'] as String? ?? 'cpu', - modelType: json['modelType'] as String? ?? '', - modelingUnit: json['modelingUnit'] as String? ?? '', - bpeVocab: json['bpeVocab'] as String? ?? '', - telespeechCtc: json['telespeechCtc'] as String? ?? '', - ); - } - - @override - String toString() { - return 'OfflineModelConfig(transducer: $transducer, paraformer: $paraformer, nemoCtc: $nemoCtc, whisper: $whisper, tdnn: $tdnn, senseVoice: $senseVoice, moonshine: $moonshine, fireRedAsr: $fireRedAsr, dolphin: $dolphin, zipformerCtc: $zipformerCtc, canary: $canary, wenetCtc: $wenetCtc, omnilingual: $omnilingual, medasr: $medasr, funasrNano: $funasrNano, fireRedAsrCtc: $fireRedAsrCtc, qwen3Asr: $qwen3Asr, cohereTranscribe: $cohereTranscribe, tokens: $tokens, numThreads: $numThreads, debug: $debug, provider: $provider, modelType: $modelType, modelingUnit: $modelingUnit, bpeVocab: $bpeVocab, telespeechCtc: $telespeechCtc)'; - } - - Map toJson() => { - 'transducer': transducer.toJson(), - 'paraformer': paraformer.toJson(), - 'nemoCtc': nemoCtc.toJson(), - 'whisper': whisper.toJson(), - 'tdnn': tdnn.toJson(), - 'senseVoice': senseVoice.toJson(), - 'moonshine': moonshine.toJson(), - 'fireRedAsr': fireRedAsr.toJson(), - 'dolphin': dolphin.toJson(), - 'zipformerCtc': zipformerCtc.toJson(), - 'canary': canary.toJson(), - 'wenetCtc': wenetCtc.toJson(), - 'omnilingual': omnilingual.toJson(), - 'medasr': medasr.toJson(), - 'funasrNano': funasrNano.toJson(), - 'fireRedAsrCtc': fireRedAsrCtc.toJson(), - 'qwen3Asr': qwen3Asr.toJson(), - 'cohereTranscribe': cohereTranscribe.toJson(), - 'tokens': tokens, - 'numThreads': numThreads, - 'debug': debug, - 'provider': provider, - 'modelType': modelType, - 'modelingUnit': modelingUnit, - 'bpeVocab': bpeVocab, - 'telespeechCtc': telespeechCtc, - }; - - final OfflineTransducerModelConfig transducer; - final OfflineParaformerModelConfig paraformer; - final OfflineNemoEncDecCtcModelConfig nemoCtc; - final OfflineWhisperModelConfig whisper; - final OfflineTdnnModelConfig tdnn; - final OfflineSenseVoiceModelConfig senseVoice; - final OfflineMoonshineModelConfig moonshine; - final OfflineFireRedAsrModelConfig fireRedAsr; - final OfflineDolphinModelConfig dolphin; - final OfflineZipformerCtcModelConfig zipformerCtc; - final OfflineCanaryModelConfig canary; - final OfflineWenetCtcModelConfig wenetCtc; - final OfflineOmnilingualAsrCtcModelConfig omnilingual; - final OfflineMedAsrCtcModelConfig medasr; - final OfflineFunAsrNanoModelConfig funasrNano; - final OfflineFireRedAsrCtcModelConfig fireRedAsrCtc; - final OfflineQwen3AsrModelConfig qwen3Asr; - final OfflineCohereTranscribeModelConfig cohereTranscribe; - - final String tokens; - final int numThreads; - final bool debug; - final String provider; - final String modelType; - final String modelingUnit; - final String bpeVocab; - final String telespeechCtc; -} - -/// Top-level configuration for [OfflineRecognizer]. -/// -/// This combines feature extraction, the selected model family, optional -/// language model settings, hotwords, grammar resources, and optional -/// homophone replacement resources. -class OfflineRecognizerConfig { - const OfflineRecognizerConfig({ - this.feat = const FeatureConfig(), - required this.model, - this.lm = const OfflineLMConfig(), - this.decodingMethod = 'greedy_search', - this.maxActivePaths = 4, - this.hotwordsFile = '', - this.hotwordsScore = 1.5, - this.ruleFsts = '', - this.ruleFars = '', - this.blankPenalty = 0.0, - this.hr = const HomophoneReplacerConfig(), - }); - - factory OfflineRecognizerConfig.fromJson(Map json) { - return OfflineRecognizerConfig( - feat: json['feat'] != null - ? FeatureConfig.fromJson(json['feat'] as Map) - : const FeatureConfig(), - model: OfflineModelConfig.fromJson(json['model'] as Map), - lm: json['lm'] != null - ? OfflineLMConfig.fromJson(json['lm'] as Map) - : const OfflineLMConfig(), - decodingMethod: json['decodingMethod'] as String? ?? 'greedy_search', - maxActivePaths: json['maxActivePaths'] as int? ?? 4, - hotwordsFile: json['hotwordsFile'] as String? ?? '', - hotwordsScore: (json['hotwordsScore'] as num?)?.toDouble() ?? 1.5, - ruleFsts: json['ruleFsts'] as String? ?? '', - ruleFars: json['ruleFars'] as String? ?? '', - blankPenalty: (json['blankPenalty'] as num?)?.toDouble() ?? 0.0, - hr: HomophoneReplacerConfig.fromJson(json['hr'] as Map), - ); - } - - @override - String toString() { - return 'OfflineRecognizerConfig(feat: $feat, model: $model, lm: $lm, decodingMethod: $decodingMethod, maxActivePaths: $maxActivePaths, hotwordsFile: $hotwordsFile, hotwordsScore: $hotwordsScore, ruleFsts: $ruleFsts, ruleFars: $ruleFars, blankPenalty: $blankPenalty, hr: $hr)'; - } - - Map toJson() => { - 'feat': feat.toJson(), - 'model': model.toJson(), - 'lm': lm.toJson(), - 'decodingMethod': decodingMethod, - 'maxActivePaths': maxActivePaths, - 'hotwordsFile': hotwordsFile, - 'hotwordsScore': hotwordsScore, - 'ruleFsts': ruleFsts, - 'ruleFars': ruleFars, - 'blankPenalty': blankPenalty, - 'hr': hr.toJson(), - }; - - final FeatureConfig feat; - final OfflineModelConfig model; - final OfflineLMConfig lm; - final String decodingMethod; - - final int maxActivePaths; - - final String hotwordsFile; - - final double hotwordsScore; - - final String ruleFsts; - final String ruleFars; - - final double blankPenalty; - final HomophoneReplacerConfig hr; -} - -/// Recognition result returned by [OfflineRecognizer.getResult]. -/// -/// Some model families populate [lang], [emotion], or [event] in addition to -/// the decoded text and token timestamps. -class OfflineRecognizerResult { - OfflineRecognizerResult({ - required this.text, - required this.tokens, - required this.timestamps, - required this.lang, - required this.emotion, - required this.event, - }); - - factory OfflineRecognizerResult.fromJson(Map json) { - return OfflineRecognizerResult( - text: json['text'] as String? ?? '', - tokens: (json['tokens'] as List?)?.map((e) => e as String).toList() ?? [], - timestamps: - (json['timestamps'] as List?) - ?.map((e) => (e as num).toDouble()) - .toList() ?? - [], - lang: json['lang'] as String? ?? '', - emotion: json['emotion'] as String? ?? '', - event: json['event'] as String? ?? '', - ); - } - - @override - String toString() { - return 'OfflineRecognizerResult(text: $text, tokens: $tokens, timestamps: $timestamps, lang: $lang, emotion: $emotion, event: $event)'; - } - - Map toJson() => { - 'text': text, - 'tokens': tokens, - 'timestamps': timestamps, - 'lang': lang, - 'emotion': emotion, - 'event': event, - }; - - final String text; - final List tokens; - final List timestamps; - final String lang; - final String emotion; - final String event; -} - /// Offline speech recognizer. /// /// Create one from an [OfflineRecognizerConfig], then create an diff --git a/flutter/sherpa_onnx/lib/src/offline_recognizer_config.dart b/flutter/sherpa_onnx/lib/src/offline_recognizer_config.dart new file mode 100644 index 0000000000..e8d37aa5fa --- /dev/null +++ b/flutter/sherpa_onnx/lib/src/offline_recognizer_config.dart @@ -0,0 +1,931 @@ +// Copyright (c) 2024 Xiaomi Corporation +// Shared config/data classes for offline recognition -- no FFI, works on all platforms. + +import './feature_config.dart'; +import './homophone_replacer_config.dart'; + +/// Model files for an offline transducer recognizer. +/// +/// This family is also used by NeMo Parakeet TDT-style examples. +class OfflineTransducerModelConfig { + const OfflineTransducerModelConfig({ + this.encoder = '', + this.decoder = '', + this.joiner = '', + }); + + factory OfflineTransducerModelConfig.fromJson(Map json) { + return OfflineTransducerModelConfig( + encoder: json['encoder'] as String? ?? '', + decoder: json['decoder'] as String? ?? '', + joiner: json['joiner'] as String? ?? '', + ); + } + + @override + String toString() { + return 'OfflineTransducerModelConfig(encoder: $encoder, decoder: $decoder, joiner: $joiner)'; + } + + Map toJson() => { + 'encoder': encoder, + 'decoder': decoder, + 'joiner': joiner, + }; + + final String encoder; + final String decoder; + final String joiner; +} + +/// Model files for an offline Paraformer recognizer. +class OfflineParaformerModelConfig { + const OfflineParaformerModelConfig({this.model = ''}); + + factory OfflineParaformerModelConfig.fromJson(Map json) { + return OfflineParaformerModelConfig(model: json['model'] as String? ?? ''); + } + + @override + String toString() { + return 'OfflineParaformerModelConfig(model: $model)'; + } + + Map toJson() => {'model': model}; + + final String model; +} + +/// Model files for an offline NeMo CTC recognizer. +class OfflineNemoEncDecCtcModelConfig { + const OfflineNemoEncDecCtcModelConfig({this.model = ''}); + + factory OfflineNemoEncDecCtcModelConfig.fromJson(Map json) { + return OfflineNemoEncDecCtcModelConfig( + model: json['model'] as String? ?? '', + ); + } + + @override + String toString() { + return 'OfflineNemoEncDecCtcModelConfig(model: $model)'; + } + + Map toJson() => {'model': model}; + + final String model; +} + +/// Model files for an offline Dolphin recognizer. +class OfflineDolphinModelConfig { + const OfflineDolphinModelConfig({this.model = ''}); + + factory OfflineDolphinModelConfig.fromJson(Map json) { + return OfflineDolphinModelConfig(model: json['model'] as String? ?? ''); + } + + @override + String toString() { + return 'OfflineDolphinModelConfig(model: $model)'; + } + + Map toJson() => {'model': model}; + + final String model; +} + +/// Model files for an offline Zipformer CTC recognizer. +class OfflineZipformerCtcModelConfig { + const OfflineZipformerCtcModelConfig({this.model = ''}); + + factory OfflineZipformerCtcModelConfig.fromJson(Map json) { + return OfflineZipformerCtcModelConfig( + model: json['model'] as String? ?? '', + ); + } + + @override + String toString() { + return 'OfflineZipformerCtcModelConfig(model: $model)'; + } + + Map toJson() => {'model': model}; + + final String model; +} + +/// Model files for an offline WeNet CTC recognizer. +class OfflineWenetCtcModelConfig { + const OfflineWenetCtcModelConfig({this.model = ''}); + + factory OfflineWenetCtcModelConfig.fromJson(Map json) { + return OfflineWenetCtcModelConfig(model: json['model'] as String? ?? ''); + } + + @override + String toString() { + return 'OfflineWenetCtcModelConfig(model: $model)'; + } + + Map toJson() => {'model': model}; + + final String model; +} + +/// Model files for the omnilingual ASR CTC recognizer. +class OfflineOmnilingualAsrCtcModelConfig { + const OfflineOmnilingualAsrCtcModelConfig({this.model = ''}); + + factory OfflineOmnilingualAsrCtcModelConfig.fromJson( + Map json, + ) { + return OfflineOmnilingualAsrCtcModelConfig( + model: json['model'] as String? ?? '', + ); + } + + @override + String toString() { + return 'OfflineOmnilingualAsrCtcModelConfig(model: $model)'; + } + + Map toJson() => {'model': model}; + + final String model; +} + +/// Model files for the MedASR CTC recognizer. +class OfflineMedAsrCtcModelConfig { + const OfflineMedAsrCtcModelConfig({this.model = ''}); + + factory OfflineMedAsrCtcModelConfig.fromJson(Map json) { + return OfflineMedAsrCtcModelConfig(model: json['model'] as String? ?? ''); + } + + @override + String toString() { + return 'OfflineMedAsrCtcModelConfig(model: $model)'; + } + + Map toJson() => {'model': model}; + + final String model; +} + +/// Model files for the Fire-Red-ASR CTC recognizer. +class OfflineFireRedAsrCtcModelConfig { + const OfflineFireRedAsrCtcModelConfig({this.model = ''}); + + factory OfflineFireRedAsrCtcModelConfig.fromJson(Map json) { + return OfflineFireRedAsrCtcModelConfig( + model: json['model'] as String? ?? '', + ); + } + + @override + String toString() { + return 'OfflineFireRedAsrCtcModelConfig(model: $model)'; + } + + Map toJson() => {'model': model}; + + final String model; +} + +/// Model files and prompt settings for FunASR-Nano. +class OfflineFunAsrNanoModelConfig { + const OfflineFunAsrNanoModelConfig({ + this.encoderAdaptor = '', + this.llm = '', + this.embedding = '', + this.tokenizer = '', + this.systemPrompt = 'You are a helpful assistant.', + this.userPrompt = '语音转写:', + this.maxNewTokens = 512, + this.temperature = 1e-6, + this.topP = 0.8, + this.seed = 42, + this.language = '', + this.itn = 1, + this.hotwords = '', + }); + + factory OfflineFunAsrNanoModelConfig.fromJson(Map json) { + return OfflineFunAsrNanoModelConfig( + encoderAdaptor: json['encoderAdaptor'] as String? ?? '', + llm: json['llm'] as String? ?? '', + embedding: json['embedding'] as String? ?? '', + tokenizer: json['tokenizer'] as String? ?? '', + systemPrompt: json['systemPrompt'] as String? ?? 'You are a helpful assistant.', + userPrompt: json['userPrompt'] as String? ?? '语音转写:', + maxNewTokens: json['maxNewTokens'] as int? ?? 512, + temperature: (json['temperature'] as num?)?.toDouble() ?? 1e-6, + topP: (json['topP'] as num?)?.toDouble() ?? 0.8, + seed: json['seed'] as int? ?? 42, + language: json['language'] as String? ?? '', + itn: json['itn'] as int? ?? 1, + hotwords: json['hotwords'] as String? ?? '', + ); + } + + @override + String toString() { + return 'OfflineFunAsrNanoModelConfig(encoderAdaptor: $encoderAdaptor, llm: $llm, embedding: $embedding, tokenizer: $tokenizer, systemPrompt: $systemPrompt, userPrompt: $userPrompt, maxNewTokens: $maxNewTokens, temperature: $temperature, topP: $topP, seed: $seed, language: $language, itn: $itn, hotwords: $hotwords)'; + } + + Map toJson() => { + 'encoderAdaptor': encoderAdaptor, + 'llm': llm, + 'embedding': embedding, + 'tokenizer': tokenizer, + 'systemPrompt': systemPrompt, + 'userPrompt': userPrompt, + 'maxNewTokens': maxNewTokens, + 'temperature': temperature, + 'topP': topP, + 'seed': seed, + 'language': language, + 'itn': itn, + 'hotwords': hotwords, + }; + + final String encoderAdaptor; + final String llm; + final String embedding; + final String tokenizer; + final String systemPrompt; + final String userPrompt; + final int maxNewTokens; + final double temperature; + final double topP; + final int seed; + final String language; + final int itn; + final String hotwords; +} + +class OfflineQwen3AsrModelConfig { + const OfflineQwen3AsrModelConfig({ + this.convFrontend = '', + this.encoder = '', + this.decoder = '', + this.tokenizer = '', + this.maxTotalLen = 512, + this.maxNewTokens = 128, + this.temperature = 1e-6, + this.topP = 0.8, + this.seed = 42, + this.hotwords = '', + }); + + factory OfflineQwen3AsrModelConfig.fromJson(Map json) { + return OfflineQwen3AsrModelConfig( + convFrontend: json['convFrontend'] as String? ?? '', + encoder: json['encoder'] as String? ?? '', + decoder: json['decoder'] as String? ?? '', + tokenizer: json['tokenizer'] as String? ?? '', + maxTotalLen: json['maxTotalLen'] as int? ?? 512, + maxNewTokens: json['maxNewTokens'] as int? ?? 128, + temperature: (json['temperature'] as num?)?.toDouble() ?? 1e-6, + topP: (json['topP'] as num?)?.toDouble() ?? 0.8, + seed: json['seed'] as int? ?? 42, + hotwords: json['hotwords'] as String? ?? '', + ); + } + + @override + String toString() { + return 'OfflineQwen3AsrModelConfig(convFrontend: $convFrontend, encoder: $encoder, decoder: $decoder, tokenizer: $tokenizer, maxTotalLen: $maxTotalLen, maxNewTokens: $maxNewTokens, temperature: $temperature, topP: $topP, seed: $seed, hotwords: $hotwords)'; + } + + Map toJson() => { + 'convFrontend': convFrontend, + 'encoder': encoder, + 'decoder': decoder, + 'tokenizer': tokenizer, + 'maxTotalLen': maxTotalLen, + 'maxNewTokens': maxNewTokens, + 'temperature': temperature, + 'topP': topP, + 'seed': seed, + 'hotwords': hotwords, + }; + + final String convFrontend; + final String encoder; + final String decoder; + final String tokenizer; + final int maxTotalLen; + final int maxNewTokens; + final double temperature; + final double topP; + final int seed; + final String hotwords; +} + +/// Model files and options for an offline Whisper recognizer. +class OfflineWhisperModelConfig { + const OfflineWhisperModelConfig({ + this.encoder = '', + this.decoder = '', + this.language = '', + this.task = '', + this.tailPaddings = -1, + this.enableTokenTimestamps = false, + this.enableSegmentTimestamps = false, + }); + + factory OfflineWhisperModelConfig.fromJson(Map json) { + return OfflineWhisperModelConfig( + encoder: json['encoder'] as String? ?? '', + decoder: json['decoder'] as String? ?? '', + language: json['language'] as String? ?? '', + task: json['task'] as String? ?? '', + tailPaddings: json['tailPaddings'] as int? ?? -1, + enableTokenTimestamps: json['enableTokenTimestamps'] as bool? ?? false, + enableSegmentTimestamps: + json['enableSegmentTimestamps'] as bool? ?? false, + ); + } + + @override + String toString() { + return 'OfflineWhisperModelConfig(encoder: $encoder, decoder: $decoder, language: $language, task: $task, tailPaddings: $tailPaddings, enableTokenTimestamps: $enableTokenTimestamps, enableSegmentTimestamps: $enableSegmentTimestamps)'; + } + + Map toJson() => { + 'encoder': encoder, + 'decoder': decoder, + 'language': language, + 'task': task, + 'tailPaddings': tailPaddings, + 'enableTokenTimestamps': enableTokenTimestamps, + 'enableSegmentTimestamps': enableSegmentTimestamps, + }; + + final String encoder; + final String decoder; + final String language; + final String task; + final int tailPaddings; + final bool enableTokenTimestamps; + final bool enableSegmentTimestamps; +} + +/// Model files and translation options for NeMo Canary. +class OfflineCanaryModelConfig { + const OfflineCanaryModelConfig({ + this.encoder = '', + this.decoder = '', + this.srcLang = 'en', + this.tgtLang = 'en', + this.usePnc = true, + }); + + factory OfflineCanaryModelConfig.fromJson(Map json) { + return OfflineCanaryModelConfig( + encoder: json['encoder'] as String? ?? '', + decoder: json['decoder'] as String? ?? '', + srcLang: json['srcLang'] as String? ?? 'en', + tgtLang: json['tgtLang'] as String? ?? 'en', + usePnc: json['usePnc'] as bool? ?? true, + ); + } + + @override + String toString() { + return 'OfflineCanaryModelConfig(encoder: $encoder, decoder: $decoder, srcLang: $srcLang, tgtLang: $tgtLang, usePnc: $usePnc)'; + } + + Map toJson() => { + 'encoder': encoder, + 'decoder': decoder, + 'srcLang': srcLang, + 'tgtLang': tgtLang, + 'usePnc': usePnc, + }; + + final String encoder; + final String decoder; + final String srcLang; + final String tgtLang; + final bool usePnc; +} + +/// Model files and text options for Cohere Transcribe. +class OfflineCohereTranscribeModelConfig { + const OfflineCohereTranscribeModelConfig({ + this.encoder = '', + this.decoder = '', + this.language = '', + this.usePunct = true, + this.useItn = true, + }); + + factory OfflineCohereTranscribeModelConfig.fromJson( + Map json, + ) { + return OfflineCohereTranscribeModelConfig( + encoder: json['encoder'] as String? ?? '', + decoder: json['decoder'] as String? ?? '', + language: json['language'] as String? ?? '', + usePunct: json['usePunct'] as bool? ?? true, + useItn: json['useItn'] as bool? ?? true, + ); + } + + @override + String toString() { + return 'OfflineCohereTranscribeModelConfig(encoder: $encoder, decoder: $decoder, language: $language, usePunct: $usePunct, useItn: $useItn)'; + } + + Map toJson() => { + 'encoder': encoder, + 'decoder': decoder, + 'language': language, + 'usePunct': usePunct, + 'useItn': useItn, + }; + + final String encoder; + final String decoder; + final String language; + final bool usePunct; + final bool useItn; +} + +/// Model files for the Fire-Red-ASR transducer recognizer. +class OfflineFireRedAsrModelConfig { + const OfflineFireRedAsrModelConfig({this.encoder = '', this.decoder = ''}); + + factory OfflineFireRedAsrModelConfig.fromJson(Map json) { + return OfflineFireRedAsrModelConfig( + encoder: json['encoder'] as String? ?? '', + decoder: json['decoder'] as String? ?? '', + ); + } + + @override + String toString() { + return 'OfflineFireRedAsrModelConfig(encoder: $encoder, decoder: $decoder)'; + } + + Map toJson() => {'encoder': encoder, 'decoder': decoder}; + + final String encoder; + final String decoder; +} + +// For Moonshine v1, you need 4 models: +// - preprocessor, encoder, uncachedDecoder, cachedDecoder +// +// For Moonshine v2, you need 2 models: +// - encoder, mergedDecoder +/// Model files for Moonshine v1 or v2. +class OfflineMoonshineModelConfig { + const OfflineMoonshineModelConfig({ + this.preprocessor = '', + this.encoder = '', + this.uncachedDecoder = '', + this.cachedDecoder = '', + this.mergedDecoder = '', + }); + + factory OfflineMoonshineModelConfig.fromJson(Map json) { + return OfflineMoonshineModelConfig( + preprocessor: json['preprocessor'] as String? ?? '', + encoder: json['encoder'] as String? ?? '', + uncachedDecoder: json['uncachedDecoder'] as String? ?? '', + cachedDecoder: json['cachedDecoder'] as String? ?? '', + mergedDecoder: json['mergedDecoder'] as String? ?? '', + ); + } + + @override + String toString() { + return 'OfflineMoonshineModelConfig(preprocessor: $preprocessor, encoder: $encoder, uncachedDecoder: $uncachedDecoder, cachedDecoder: $cachedDecoder, mergedDecoder: $mergedDecoder)'; + } + + Map toJson() => { + 'preprocessor': preprocessor, + 'encoder': encoder, + 'uncachedDecoder': uncachedDecoder, + 'cachedDecoder': cachedDecoder, + 'mergedDecoder': mergedDecoder, + }; + + final String preprocessor; + final String encoder; + final String uncachedDecoder; + final String cachedDecoder; + final String mergedDecoder; +} + +/// Model files for an offline TDNN recognizer. +class OfflineTdnnModelConfig { + const OfflineTdnnModelConfig({this.model = ''}); + + factory OfflineTdnnModelConfig.fromJson(Map json) { + return OfflineTdnnModelConfig(model: json['model'] as String? ?? ''); + } + + @override + String toString() { + return 'OfflineTdnnModelConfig(model: $model)'; + } + + Map toJson() => {'model': model}; + + final String model; +} + +/// Model files and options for SenseVoice. +/// +/// In the examples, this is typically paired with the +/// `sherpa-onnx-sense-voice-zh-en-ja-ko-yue-2024-07-17-int8` package. +class OfflineSenseVoiceModelConfig { + const OfflineSenseVoiceModelConfig({ + this.model = '', + this.language = '', + this.useInverseTextNormalization = false, + }); + + factory OfflineSenseVoiceModelConfig.fromJson(Map json) { + return OfflineSenseVoiceModelConfig( + model: json['model'] as String? ?? '', + language: json['language'] as String? ?? '', + useInverseTextNormalization: + json['useInverseTextNormalization'] as bool? ?? false, + ); + } + + @override + String toString() { + return 'OfflineSenseVoiceModelConfig(model: $model, language: $language, useInverseTextNormalization: $useInverseTextNormalization)'; + } + + Map toJson() => { + 'model': model, + 'language': language, + 'useInverseTextNormalization': useInverseTextNormalization, + }; + + final String model; + final String language; + final bool useInverseTextNormalization; +} + +/// Optional external language model settings for offline ASR. +class OfflineLMConfig { + const OfflineLMConfig({this.model = '', this.scale = 1.0}); + + factory OfflineLMConfig.fromJson(Map json) { + return OfflineLMConfig( + model: json['model'] as String? ?? '', + scale: (json['scale'] as num?)?.toDouble() ?? 1.0, + ); + } + + @override + String toString() { + return 'OfflineLMConfig(model: $model, scale: $scale)'; + } + + Map toJson() => {'model': model, 'scale': scale}; + + final String model; + final double scale; +} + +/// Aggregate model configuration for offline recognition. +/// +/// In typical use, configure exactly one model family and set the shared +/// options such as [tokens], [provider], and [numThreads]. +/// +/// For NeMo Parakeet-style transducer models, set [modelType] to +/// `nemo_transducer`, matching the repository examples. +class OfflineModelConfig { + const OfflineModelConfig({ + this.transducer = const OfflineTransducerModelConfig(), + this.paraformer = const OfflineParaformerModelConfig(), + this.nemoCtc = const OfflineNemoEncDecCtcModelConfig(), + this.whisper = const OfflineWhisperModelConfig(), + this.tdnn = const OfflineTdnnModelConfig(), + this.senseVoice = const OfflineSenseVoiceModelConfig(), + this.moonshine = const OfflineMoonshineModelConfig(), + this.fireRedAsr = const OfflineFireRedAsrModelConfig(), + this.dolphin = const OfflineDolphinModelConfig(), + this.zipformerCtc = const OfflineZipformerCtcModelConfig(), + this.canary = const OfflineCanaryModelConfig(), + this.wenetCtc = const OfflineWenetCtcModelConfig(), + this.omnilingual = const OfflineOmnilingualAsrCtcModelConfig(), + this.medasr = const OfflineMedAsrCtcModelConfig(), + this.funasrNano = const OfflineFunAsrNanoModelConfig(), + this.fireRedAsrCtc = const OfflineFireRedAsrCtcModelConfig(), + this.qwen3Asr = const OfflineQwen3AsrModelConfig(), + this.cohereTranscribe = const OfflineCohereTranscribeModelConfig(), + required this.tokens, + this.numThreads = 1, + this.debug = true, + this.provider = 'cpu', + this.modelType = '', + this.modelingUnit = '', + this.bpeVocab = '', + this.telespeechCtc = '', + }); + + factory OfflineModelConfig.fromJson(Map json) { + return OfflineModelConfig( + transducer: json['transducer'] != null + ? OfflineTransducerModelConfig.fromJson( + json['transducer'] as Map, + ) + : const OfflineTransducerModelConfig(), + paraformer: json['paraformer'] != null + ? OfflineParaformerModelConfig.fromJson( + json['paraformer'] as Map, + ) + : const OfflineParaformerModelConfig(), + nemoCtc: json['nemoCtc'] != null + ? OfflineNemoEncDecCtcModelConfig.fromJson( + json['nemoCtc'] as Map, + ) + : const OfflineNemoEncDecCtcModelConfig(), + whisper: json['whisper'] != null + ? OfflineWhisperModelConfig.fromJson( + json['whisper'] as Map, + ) + : const OfflineWhisperModelConfig(), + tdnn: json['tdnn'] != null + ? OfflineTdnnModelConfig.fromJson( + json['tdnn'] as Map, + ) + : const OfflineTdnnModelConfig(), + senseVoice: json['senseVoice'] != null + ? OfflineSenseVoiceModelConfig.fromJson( + json['senseVoice'] as Map, + ) + : const OfflineSenseVoiceModelConfig(), + moonshine: json['moonshine'] != null + ? OfflineMoonshineModelConfig.fromJson( + json['moonshine'] as Map, + ) + : const OfflineMoonshineModelConfig(), + fireRedAsr: json['fireRedAsr'] != null + ? OfflineFireRedAsrModelConfig.fromJson( + json['fireRedAsr'] as Map, + ) + : const OfflineFireRedAsrModelConfig(), + dolphin: json['dolphin'] != null + ? OfflineDolphinModelConfig.fromJson( + json['dolphin'] as Map, + ) + : const OfflineDolphinModelConfig(), + zipformerCtc: json['zipformerCtc'] != null + ? OfflineZipformerCtcModelConfig.fromJson( + json['zipformerCtc'] as Map, + ) + : const OfflineZipformerCtcModelConfig(), + canary: json['canary'] != null + ? OfflineCanaryModelConfig.fromJson( + json['canary'] as Map, + ) + : const OfflineCanaryModelConfig(), + wenetCtc: json['wenetCtc'] != null + ? OfflineWenetCtcModelConfig.fromJson( + json['wenetCtc'] as Map, + ) + : const OfflineWenetCtcModelConfig(), + omnilingual: json['omnilingual'] != null + ? OfflineOmnilingualAsrCtcModelConfig.fromJson( + json['omnilingual'] as Map, + ) + : const OfflineOmnilingualAsrCtcModelConfig(), + medasr: json['medasr'] != null + ? OfflineMedAsrCtcModelConfig.fromJson( + json['medasr'] as Map, + ) + : const OfflineMedAsrCtcModelConfig(), + funasrNano: json['funasrNano'] != null + ? OfflineFunAsrNanoModelConfig.fromJson( + json['funasrNano'] as Map, + ) + : const OfflineFunAsrNanoModelConfig(), + fireRedAsrCtc: json['fireRedAsrCtc'] != null + ? OfflineFireRedAsrCtcModelConfig.fromJson( + json['fireRedAsrCtc'] as Map, + ) + : const OfflineFireRedAsrCtcModelConfig(), + qwen3Asr: json['qwen3Asr'] != null + ? OfflineQwen3AsrModelConfig.fromJson( + json['qwen3Asr'] as Map, + ) + : const OfflineQwen3AsrModelConfig(), + cohereTranscribe: json['cohereTranscribe'] != null + ? OfflineCohereTranscribeModelConfig.fromJson( + json['cohereTranscribe'] as Map, + ) + : const OfflineCohereTranscribeModelConfig(), + tokens: json['tokens'] as String, + numThreads: json['numThreads'] as int? ?? 1, + debug: json['debug'] as bool? ?? true, + provider: json['provider'] as String? ?? 'cpu', + modelType: json['modelType'] as String? ?? '', + modelingUnit: json['modelingUnit'] as String? ?? '', + bpeVocab: json['bpeVocab'] as String? ?? '', + telespeechCtc: json['telespeechCtc'] as String? ?? '', + ); + } + + @override + String toString() { + return 'OfflineModelConfig(transducer: $transducer, paraformer: $paraformer, nemoCtc: $nemoCtc, whisper: $whisper, tdnn: $tdnn, senseVoice: $senseVoice, moonshine: $moonshine, fireRedAsr: $fireRedAsr, dolphin: $dolphin, zipformerCtc: $zipformerCtc, canary: $canary, wenetCtc: $wenetCtc, omnilingual: $omnilingual, medasr: $medasr, funasrNano: $funasrNano, fireRedAsrCtc: $fireRedAsrCtc, qwen3Asr: $qwen3Asr, cohereTranscribe: $cohereTranscribe, tokens: $tokens, numThreads: $numThreads, debug: $debug, provider: $provider, modelType: $modelType, modelingUnit: $modelingUnit, bpeVocab: $bpeVocab, telespeechCtc: $telespeechCtc)'; + } + + Map toJson() => { + 'transducer': transducer.toJson(), + 'paraformer': paraformer.toJson(), + 'nemoCtc': nemoCtc.toJson(), + 'whisper': whisper.toJson(), + 'tdnn': tdnn.toJson(), + 'senseVoice': senseVoice.toJson(), + 'moonshine': moonshine.toJson(), + 'fireRedAsr': fireRedAsr.toJson(), + 'dolphin': dolphin.toJson(), + 'zipformerCtc': zipformerCtc.toJson(), + 'canary': canary.toJson(), + 'wenetCtc': wenetCtc.toJson(), + 'omnilingual': omnilingual.toJson(), + 'medasr': medasr.toJson(), + 'funasrNano': funasrNano.toJson(), + 'fireRedAsrCtc': fireRedAsrCtc.toJson(), + 'qwen3Asr': qwen3Asr.toJson(), + 'cohereTranscribe': cohereTranscribe.toJson(), + 'tokens': tokens, + 'numThreads': numThreads, + 'debug': debug, + 'provider': provider, + 'modelType': modelType, + 'modelingUnit': modelingUnit, + 'bpeVocab': bpeVocab, + 'telespeechCtc': telespeechCtc, + }; + + final OfflineTransducerModelConfig transducer; + final OfflineParaformerModelConfig paraformer; + final OfflineNemoEncDecCtcModelConfig nemoCtc; + final OfflineWhisperModelConfig whisper; + final OfflineTdnnModelConfig tdnn; + final OfflineSenseVoiceModelConfig senseVoice; + final OfflineMoonshineModelConfig moonshine; + final OfflineFireRedAsrModelConfig fireRedAsr; + final OfflineDolphinModelConfig dolphin; + final OfflineZipformerCtcModelConfig zipformerCtc; + final OfflineCanaryModelConfig canary; + final OfflineWenetCtcModelConfig wenetCtc; + final OfflineOmnilingualAsrCtcModelConfig omnilingual; + final OfflineMedAsrCtcModelConfig medasr; + final OfflineFunAsrNanoModelConfig funasrNano; + final OfflineFireRedAsrCtcModelConfig fireRedAsrCtc; + final OfflineQwen3AsrModelConfig qwen3Asr; + final OfflineCohereTranscribeModelConfig cohereTranscribe; + + final String tokens; + final int numThreads; + final bool debug; + final String provider; + final String modelType; + final String modelingUnit; + final String bpeVocab; + final String telespeechCtc; +} + +/// Top-level configuration for [OfflineRecognizer]. +/// +/// This combines feature extraction, the selected model family, optional +/// language model settings, hotwords, grammar resources, and optional +/// homophone replacement resources. +class OfflineRecognizerConfig { + const OfflineRecognizerConfig({ + this.feat = const FeatureConfig(), + required this.model, + this.lm = const OfflineLMConfig(), + this.decodingMethod = 'greedy_search', + this.maxActivePaths = 4, + this.hotwordsFile = '', + this.hotwordsScore = 1.5, + this.ruleFsts = '', + this.ruleFars = '', + this.blankPenalty = 0.0, + this.hr = const HomophoneReplacerConfig(), + }); + + factory OfflineRecognizerConfig.fromJson(Map json) { + return OfflineRecognizerConfig( + feat: json['feat'] != null + ? FeatureConfig.fromJson(json['feat'] as Map) + : const FeatureConfig(), + model: OfflineModelConfig.fromJson(json['model'] as Map), + lm: json['lm'] != null + ? OfflineLMConfig.fromJson(json['lm'] as Map) + : const OfflineLMConfig(), + decodingMethod: json['decodingMethod'] as String? ?? 'greedy_search', + maxActivePaths: json['maxActivePaths'] as int? ?? 4, + hotwordsFile: json['hotwordsFile'] as String? ?? '', + hotwordsScore: (json['hotwordsScore'] as num?)?.toDouble() ?? 1.5, + ruleFsts: json['ruleFsts'] as String? ?? '', + ruleFars: json['ruleFars'] as String? ?? '', + blankPenalty: (json['blankPenalty'] as num?)?.toDouble() ?? 0.0, + hr: json['hr'] != null + ? HomophoneReplacerConfig.fromJson(json['hr'] as Map) + : const HomophoneReplacerConfig(), + ); + } + + @override + String toString() { + return 'OfflineRecognizerConfig(feat: $feat, model: $model, lm: $lm, decodingMethod: $decodingMethod, maxActivePaths: $maxActivePaths, hotwordsFile: $hotwordsFile, hotwordsScore: $hotwordsScore, ruleFsts: $ruleFsts, ruleFars: $ruleFars, blankPenalty: $blankPenalty, hr: $hr)'; + } + + Map toJson() => { + 'feat': feat.toJson(), + 'model': model.toJson(), + 'lm': lm.toJson(), + 'decodingMethod': decodingMethod, + 'maxActivePaths': maxActivePaths, + 'hotwordsFile': hotwordsFile, + 'hotwordsScore': hotwordsScore, + 'ruleFsts': ruleFsts, + 'ruleFars': ruleFars, + 'blankPenalty': blankPenalty, + 'hr': hr.toJson(), + }; + + final FeatureConfig feat; + final OfflineModelConfig model; + final OfflineLMConfig lm; + final String decodingMethod; + + final int maxActivePaths; + + final String hotwordsFile; + + final double hotwordsScore; + + final String ruleFsts; + final String ruleFars; + + final double blankPenalty; + final HomophoneReplacerConfig hr; +} + +/// Recognition result returned by [OfflineRecognizer.getResult]. +/// +/// Some model families populate [lang], [emotion], or [event] in addition to +/// the decoded text and token timestamps. +class OfflineRecognizerResult { + OfflineRecognizerResult({ + required this.text, + required this.tokens, + required this.timestamps, + required this.lang, + required this.emotion, + required this.event, + }); + + factory OfflineRecognizerResult.fromJson(Map json) { + return OfflineRecognizerResult( + text: json['text'] as String? ?? '', + tokens: (json['tokens'] as List?)?.map((e) => e as String).toList() ?? [], + timestamps: + (json['timestamps'] as List?) + ?.map((e) => (e as num).toDouble()) + .toList() ?? + [], + lang: json['lang'] as String? ?? '', + emotion: json['emotion'] as String? ?? '', + event: json['event'] as String? ?? '', + ); + } + + @override + String toString() { + return 'OfflineRecognizerResult(text: $text, tokens: $tokens, timestamps: $timestamps, lang: $lang, emotion: $emotion, event: $event)'; + } + + Map toJson() => { + 'text': text, + 'tokens': tokens, + 'timestamps': timestamps, + 'lang': lang, + 'emotion': emotion, + 'event': event, + }; + + final String text; + final List tokens; + final List timestamps; + final String lang; + final String emotion; + final String event; +} diff --git a/flutter/sherpa_onnx/lib/src/offline_speaker_diarization.dart b/flutter/sherpa_onnx/lib/src/offline_speaker_diarization.dart index 76859af4b9..6990a8e69f 100644 --- a/flutter/sherpa_onnx/lib/src/offline_speaker_diarization.dart +++ b/flutter/sherpa_onnx/lib/src/offline_speaker_diarization.dart @@ -6,185 +6,9 @@ import 'package:ffi/ffi.dart'; import './sherpa_onnx_bindings.dart'; import './speaker_identification.dart'; +import './offline_speaker_diarization_config.dart'; -/// Offline speaker diarization. -/// -/// This module combines segmentation, speaker embedding extraction, and -/// clustering to assign speaker labels to time spans. See -/// `dart-api-examples/speaker-diarization/` for a complete example. -class OfflineSpeakerDiarizationSegment { - const OfflineSpeakerDiarizationSegment({ - required this.start, - required this.end, - required this.speaker, - }); - - factory OfflineSpeakerDiarizationSegment.fromJson(Map json) { - return OfflineSpeakerDiarizationSegment( - start: (json['start'] as num).toDouble(), - end: (json['end'] as num).toDouble(), - speaker: json['speaker'] as int, - ); - } - - @override - String toString() { - return 'OfflineSpeakerDiarizationSegment(start: $start, end: $end, speaker: $speaker)'; - } - - Map toJson() => { - 'start': start, - 'end': end, - 'speaker': speaker, - }; - - final double start; - final double end; - final int speaker; -} - -/// Pyannote segmentation model path. -class OfflineSpeakerSegmentationPyannoteModelConfig { - const OfflineSpeakerSegmentationPyannoteModelConfig({ - this.model = '', - }); - - factory OfflineSpeakerSegmentationPyannoteModelConfig.fromJson( - Map json) { - return OfflineSpeakerSegmentationPyannoteModelConfig( - model: json['model'] as String? ?? '', - ); - } - - @override - String toString() { - return 'OfflineSpeakerSegmentationPyannoteModelConfig(model: $model)'; - } - - Map toJson() => { - 'model': model, - }; - - final String model; -} - -/// Segmentation model configuration for speaker diarization. -class OfflineSpeakerSegmentationModelConfig { - const OfflineSpeakerSegmentationModelConfig({ - this.pyannote = const OfflineSpeakerSegmentationPyannoteModelConfig(), - this.numThreads = 1, - this.debug = true, - this.provider = 'cpu', - }); - - factory OfflineSpeakerSegmentationModelConfig.fromJson( - Map json) { - return OfflineSpeakerSegmentationModelConfig( - pyannote: json['pyannote'] != null - ? OfflineSpeakerSegmentationPyannoteModelConfig.fromJson( - json['pyannote'] as Map) - : const OfflineSpeakerSegmentationPyannoteModelConfig(), - numThreads: json['numThreads'] as int? ?? 1, - debug: json['debug'] as bool? ?? true, - provider: json['provider'] as String? ?? 'cpu', - ); - } - - @override - String toString() { - return 'OfflineSpeakerSegmentationModelConfig(pyannote: $pyannote, numThreads: $numThreads, debug: $debug, provider: $provider)'; - } - - Map toJson() => { - 'pyannote': pyannote.toJson(), - 'numThreads': numThreads, - 'debug': debug, - 'provider': provider, - }; - - final OfflineSpeakerSegmentationPyannoteModelConfig pyannote; - - final int numThreads; - final bool debug; - final String provider; -} - -/// Clustering options used after segmentation and embedding extraction. -class FastClusteringConfig { - const FastClusteringConfig({ - this.numClusters = -1, - this.threshold = 0.5, - }); - - factory FastClusteringConfig.fromJson(Map json) { - return FastClusteringConfig( - numClusters: json['numClusters'] as int? ?? -1, - threshold: (json['threshold'] as num?)?.toDouble() ?? 0.5, - ); - } - - @override - String toString() { - return 'FastClusteringConfig(numClusters: $numClusters, threshold: $threshold)'; - } - - Map toJson() => { - 'numClusters': numClusters, - 'threshold': threshold, - }; - - final int numClusters; - final double threshold; -} - -/// Top-level configuration for [OfflineSpeakerDiarization]. -class OfflineSpeakerDiarizationConfig { - const OfflineSpeakerDiarizationConfig({ - this.segmentation = const OfflineSpeakerSegmentationModelConfig(), - this.embedding = const SpeakerEmbeddingExtractorConfig(model: ''), - this.clustering = const FastClusteringConfig(), - this.minDurationOn = 0.2, - this.minDurationOff = 0.5, - }); - - factory OfflineSpeakerDiarizationConfig.fromJson(Map json) { - return OfflineSpeakerDiarizationConfig( - segmentation: json['segmentation'] != null - ? OfflineSpeakerSegmentationModelConfig.fromJson( - json['segmentation'] as Map) - : const OfflineSpeakerSegmentationModelConfig(), - embedding: json['embedding'] != null - ? SpeakerEmbeddingExtractorConfig.fromJson( - json['embedding'] as Map) - : const SpeakerEmbeddingExtractorConfig(model: ''), - clustering: json['clustering'] != null - ? FastClusteringConfig.fromJson( - json['clustering'] as Map) - : const FastClusteringConfig(), - minDurationOn: (json['minDurationOn'] as num?)?.toDouble() ?? 0.2, - minDurationOff: (json['minDurationOff'] as num?)?.toDouble() ?? 0.5, - ); - } - - @override - String toString() { - return 'OfflineSpeakerDiarizationConfig(segmentation: $segmentation, embedding: $embedding, clustering: $clustering, minDurationOn: $minDurationOn, minDurationOff: $minDurationOff)'; - } - - Map toJson() => { - 'segmentation': segmentation.toJson(), - 'embedding': embedding.toJson(), - 'clustering': clustering.toJson(), - 'minDurationOn': minDurationOn, - 'minDurationOff': minDurationOff, - }; - - final OfflineSpeakerSegmentationModelConfig segmentation; - final SpeakerEmbeddingExtractorConfig embedding; - final FastClusteringConfig clustering; - final double minDurationOff; // in seconds - final double minDurationOn; // in seconds -} +export './offline_speaker_diarization_config.dart'; /// Offline speaker diarizer. class OfflineSpeakerDiarization { diff --git a/flutter/sherpa_onnx/lib/src/offline_speaker_diarization_config.dart b/flutter/sherpa_onnx/lib/src/offline_speaker_diarization_config.dart new file mode 100644 index 0000000000..474784e0de --- /dev/null +++ b/flutter/sherpa_onnx/lib/src/offline_speaker_diarization_config.dart @@ -0,0 +1,183 @@ +// Copyright (c) 2024 Xiaomi Corporation +// Shared config/data classes for offline speaker diarization -- no FFI, works on all platforms. + +import './speaker_identification_config.dart'; + +/// Offline speaker diarization. +/// +/// This module combines segmentation, speaker embedding extraction, and +/// clustering to assign speaker labels to time spans. See +/// `dart-api-examples/speaker-diarization/` for a complete example. +class OfflineSpeakerDiarizationSegment { + const OfflineSpeakerDiarizationSegment({ + required this.start, + required this.end, + required this.speaker, + }); + + factory OfflineSpeakerDiarizationSegment.fromJson(Map json) { + return OfflineSpeakerDiarizationSegment( + start: (json['start'] as num).toDouble(), + end: (json['end'] as num).toDouble(), + speaker: json['speaker'] as int, + ); + } + + @override + String toString() { + return 'OfflineSpeakerDiarizationSegment(start: $start, end: $end, speaker: $speaker)'; + } + + Map toJson() => { + 'start': start, + 'end': end, + 'speaker': speaker, + }; + + final double start; + final double end; + final int speaker; +} + +/// Pyannote segmentation model path. +class OfflineSpeakerSegmentationPyannoteModelConfig { + const OfflineSpeakerSegmentationPyannoteModelConfig({ + this.model = '', + }); + + factory OfflineSpeakerSegmentationPyannoteModelConfig.fromJson( + Map json) { + return OfflineSpeakerSegmentationPyannoteModelConfig( + model: json['model'] as String? ?? '', + ); + } + + @override + String toString() { + return 'OfflineSpeakerSegmentationPyannoteModelConfig(model: $model)'; + } + + Map toJson() => { + 'model': model, + }; + + final String model; +} + +/// Segmentation model configuration for speaker diarization. +class OfflineSpeakerSegmentationModelConfig { + const OfflineSpeakerSegmentationModelConfig({ + this.pyannote = const OfflineSpeakerSegmentationPyannoteModelConfig(), + this.numThreads = 1, + this.debug = true, + this.provider = 'cpu', + }); + + factory OfflineSpeakerSegmentationModelConfig.fromJson( + Map json) { + return OfflineSpeakerSegmentationModelConfig( + pyannote: json['pyannote'] != null + ? OfflineSpeakerSegmentationPyannoteModelConfig.fromJson( + json['pyannote'] as Map) + : const OfflineSpeakerSegmentationPyannoteModelConfig(), + numThreads: json['numThreads'] as int? ?? 1, + debug: json['debug'] as bool? ?? true, + provider: json['provider'] as String? ?? 'cpu', + ); + } + + @override + String toString() { + return 'OfflineSpeakerSegmentationModelConfig(pyannote: $pyannote, numThreads: $numThreads, debug: $debug, provider: $provider)'; + } + + Map toJson() => { + 'pyannote': pyannote.toJson(), + 'numThreads': numThreads, + 'debug': debug, + 'provider': provider, + }; + + final OfflineSpeakerSegmentationPyannoteModelConfig pyannote; + + final int numThreads; + final bool debug; + final String provider; +} + +/// Clustering options used after segmentation and embedding extraction. +class FastClusteringConfig { + const FastClusteringConfig({ + this.numClusters = -1, + this.threshold = 0.5, + }); + + factory FastClusteringConfig.fromJson(Map json) { + return FastClusteringConfig( + numClusters: json['numClusters'] as int? ?? -1, + threshold: (json['threshold'] as num?)?.toDouble() ?? 0.5, + ); + } + + @override + String toString() { + return 'FastClusteringConfig(numClusters: $numClusters, threshold: $threshold)'; + } + + Map toJson() => { + 'numClusters': numClusters, + 'threshold': threshold, + }; + + final int numClusters; + final double threshold; +} + +/// Top-level configuration for [OfflineSpeakerDiarization]. +class OfflineSpeakerDiarizationConfig { + const OfflineSpeakerDiarizationConfig({ + this.segmentation = const OfflineSpeakerSegmentationModelConfig(), + this.embedding = const SpeakerEmbeddingExtractorConfig(model: ''), + this.clustering = const FastClusteringConfig(), + this.minDurationOn = 0.2, + this.minDurationOff = 0.5, + }); + + factory OfflineSpeakerDiarizationConfig.fromJson(Map json) { + return OfflineSpeakerDiarizationConfig( + segmentation: json['segmentation'] != null + ? OfflineSpeakerSegmentationModelConfig.fromJson( + json['segmentation'] as Map) + : const OfflineSpeakerSegmentationModelConfig(), + embedding: json['embedding'] != null + ? SpeakerEmbeddingExtractorConfig.fromJson( + json['embedding'] as Map) + : const SpeakerEmbeddingExtractorConfig(model: ''), + clustering: json['clustering'] != null + ? FastClusteringConfig.fromJson( + json['clustering'] as Map) + : const FastClusteringConfig(), + minDurationOn: (json['minDurationOn'] as num?)?.toDouble() ?? 0.2, + minDurationOff: (json['minDurationOff'] as num?)?.toDouble() ?? 0.5, + ); + } + + @override + String toString() { + return 'OfflineSpeakerDiarizationConfig(segmentation: $segmentation, embedding: $embedding, clustering: $clustering, minDurationOn: $minDurationOn, minDurationOff: $minDurationOff)'; + } + + Map toJson() => { + 'segmentation': segmentation.toJson(), + 'embedding': embedding.toJson(), + 'clustering': clustering.toJson(), + 'minDurationOn': minDurationOn, + 'minDurationOff': minDurationOff, + }; + + final OfflineSpeakerSegmentationModelConfig segmentation; + final SpeakerEmbeddingExtractorConfig embedding; + final FastClusteringConfig clustering; + final double minDurationOff; // in seconds + final double minDurationOn; // in seconds +} diff --git a/flutter/sherpa_onnx/lib/src/offline_speech_denoiser.dart b/flutter/sherpa_onnx/lib/src/offline_speech_denoiser.dart index 735f35998d..11a1244b7d 100644 --- a/flutter/sherpa_onnx/lib/src/offline_speech_denoiser.dart +++ b/flutter/sherpa_onnx/lib/src/offline_speech_denoiser.dart @@ -4,146 +4,9 @@ import 'dart:typed_data'; import 'package:ffi/ffi.dart'; import './sherpa_onnx_bindings.dart'; +import './offline_speech_denoiser_config.dart'; -/// Offline speech denoising. -/// -/// Supported model families include GTCRN and DPDFNet. See the examples under -/// `dart-api-examples/speech-enhancement-gtcrn/` and -/// `dart-api-examples/speech-enhancement-dpdfnet/`. -class OfflineSpeechDenoiserGtcrnModelConfig { - const OfflineSpeechDenoiserGtcrnModelConfig({ - this.model = '', - }); - - factory OfflineSpeechDenoiserGtcrnModelConfig.fromJson( - Map json) { - return OfflineSpeechDenoiserGtcrnModelConfig( - model: json['model'] as String? ?? '', - ); - } - - @override - String toString() { - return 'OfflineSpeechDenoiserGtcrnModelConfig(model: $model)'; - } - - Map toJson() => { - 'model': model, - }; - - final String model; -} - -/// DPDFNet model path for offline speech denoising. -class OfflineSpeechDenoiserDpdfNetModelConfig { - const OfflineSpeechDenoiserDpdfNetModelConfig({ - this.model = '', - }); - - factory OfflineSpeechDenoiserDpdfNetModelConfig.fromJson( - Map json) { - return OfflineSpeechDenoiserDpdfNetModelConfig( - model: json['model'] as String? ?? '', - ); - } - - @override - String toString() { - return 'OfflineSpeechDenoiserDpdfNetModelConfig(model: $model)'; - } - - Map toJson() => { - 'model': model, - }; - - final String model; -} - -/// Aggregate model configuration for [OfflineSpeechDenoiser]. -/// -/// Configure either [gtcrn] or [dpdfnet] for typical use. -class OfflineSpeechDenoiserModelConfig { - const OfflineSpeechDenoiserModelConfig({ - this.gtcrn = const OfflineSpeechDenoiserGtcrnModelConfig(), - this.dpdfnet = const OfflineSpeechDenoiserDpdfNetModelConfig(), - this.numThreads = 1, - this.debug = true, - this.provider = 'cpu', - }); - - factory OfflineSpeechDenoiserModelConfig.fromJson(Map json) { - return OfflineSpeechDenoiserModelConfig( - gtcrn: json['gtcrn'] != null - ? OfflineSpeechDenoiserGtcrnModelConfig.fromJson( - json['gtcrn'] as Map) - : const OfflineSpeechDenoiserGtcrnModelConfig(), - dpdfnet: json['dpdfnet'] != null - ? OfflineSpeechDenoiserDpdfNetModelConfig.fromJson( - json['dpdfnet'] as Map) - : const OfflineSpeechDenoiserDpdfNetModelConfig(), - numThreads: json['numThreads'] as int? ?? 1, - debug: json['debug'] as bool? ?? true, - provider: json['provider'] as String? ?? 'cpu', - ); - } - - @override - String toString() { - return 'OfflineSpeechDenoiserModelConfig(gtcrn: $gtcrn, dpdfnet: $dpdfnet, numThreads: $numThreads, debug: $debug, provider: $provider)'; - } - - Map toJson() => { - 'gtcrn': gtcrn.toJson(), - 'dpdfnet': dpdfnet.toJson(), - 'numThreads': numThreads, - 'debug': debug, - 'provider': provider, - }; - - final OfflineSpeechDenoiserGtcrnModelConfig gtcrn; - final OfflineSpeechDenoiserDpdfNetModelConfig dpdfnet; - final int numThreads; - final bool debug; - final String provider; -} - -/// Top-level configuration for [OfflineSpeechDenoiser]. -class OfflineSpeechDenoiserConfig { - const OfflineSpeechDenoiserConfig({ - this.model = const OfflineSpeechDenoiserModelConfig(), - }); - - factory OfflineSpeechDenoiserConfig.fromJson(Map json) { - return OfflineSpeechDenoiserConfig( - model: json['model'] != null - ? OfflineSpeechDenoiserModelConfig.fromJson( - json['model'] as Map) - : const OfflineSpeechDenoiserModelConfig(), - ); - } - - @override - String toString() { - return 'OfflineSpeechDenoiserConfig(model: $model)'; - } - - Map toJson() => { - 'model': model.toJson(), - }; - - final OfflineSpeechDenoiserModelConfig model; -} - -/// Audio returned by offline or online speech denoisers. -class DenoisedAudio { - DenoisedAudio({ - required this.samples, - required this.sampleRate, - }); - - final Float32List samples; - final int sampleRate; -} +export './offline_speech_denoiser_config.dart'; /// Offline speech denoiser. class OfflineSpeechDenoiser { diff --git a/flutter/sherpa_onnx/lib/src/offline_speech_denoiser_config.dart b/flutter/sherpa_onnx/lib/src/offline_speech_denoiser_config.dart new file mode 100644 index 0000000000..d767661bc4 --- /dev/null +++ b/flutter/sherpa_onnx/lib/src/offline_speech_denoiser_config.dart @@ -0,0 +1,143 @@ +// Copyright (c) 2025 Xiaomi Corporation +// Shared config/data classes for offline speech denoising -- no FFI, works on all platforms. +import 'dart:typed_data'; + +/// Offline speech denoising. +/// +/// Supported model families include GTCRN and DPDFNet. See the examples under +/// `dart-api-examples/speech-enhancement-gtcrn/` and +/// `dart-api-examples/speech-enhancement-dpdfnet/`. +class OfflineSpeechDenoiserGtcrnModelConfig { + const OfflineSpeechDenoiserGtcrnModelConfig({ + this.model = '', + }); + + factory OfflineSpeechDenoiserGtcrnModelConfig.fromJson( + Map json) { + return OfflineSpeechDenoiserGtcrnModelConfig( + model: json['model'] as String? ?? '', + ); + } + + @override + String toString() { + return 'OfflineSpeechDenoiserGtcrnModelConfig(model: $model)'; + } + + Map toJson() => { + 'model': model, + }; + + final String model; +} + +/// DPDFNet model path for offline speech denoising. +class OfflineSpeechDenoiserDpdfNetModelConfig { + const OfflineSpeechDenoiserDpdfNetModelConfig({ + this.model = '', + }); + + factory OfflineSpeechDenoiserDpdfNetModelConfig.fromJson( + Map json) { + return OfflineSpeechDenoiserDpdfNetModelConfig( + model: json['model'] as String? ?? '', + ); + } + + @override + String toString() { + return 'OfflineSpeechDenoiserDpdfNetModelConfig(model: $model)'; + } + + Map toJson() => { + 'model': model, + }; + + final String model; +} + +/// Aggregate model configuration for [OfflineSpeechDenoiser]. +/// +/// Configure either [gtcrn] or [dpdfnet] for typical use. +class OfflineSpeechDenoiserModelConfig { + const OfflineSpeechDenoiserModelConfig({ + this.gtcrn = const OfflineSpeechDenoiserGtcrnModelConfig(), + this.dpdfnet = const OfflineSpeechDenoiserDpdfNetModelConfig(), + this.numThreads = 1, + this.debug = true, + this.provider = 'cpu', + }); + + factory OfflineSpeechDenoiserModelConfig.fromJson(Map json) { + return OfflineSpeechDenoiserModelConfig( + gtcrn: json['gtcrn'] != null + ? OfflineSpeechDenoiserGtcrnModelConfig.fromJson( + json['gtcrn'] as Map) + : const OfflineSpeechDenoiserGtcrnModelConfig(), + dpdfnet: json['dpdfnet'] != null + ? OfflineSpeechDenoiserDpdfNetModelConfig.fromJson( + json['dpdfnet'] as Map) + : const OfflineSpeechDenoiserDpdfNetModelConfig(), + numThreads: json['numThreads'] as int? ?? 1, + debug: json['debug'] as bool? ?? true, + provider: json['provider'] as String? ?? 'cpu', + ); + } + + @override + String toString() { + return 'OfflineSpeechDenoiserModelConfig(gtcrn: $gtcrn, dpdfnet: $dpdfnet, numThreads: $numThreads, debug: $debug, provider: $provider)'; + } + + Map toJson() => { + 'gtcrn': gtcrn.toJson(), + 'dpdfnet': dpdfnet.toJson(), + 'numThreads': numThreads, + 'debug': debug, + 'provider': provider, + }; + + final OfflineSpeechDenoiserGtcrnModelConfig gtcrn; + final OfflineSpeechDenoiserDpdfNetModelConfig dpdfnet; + final int numThreads; + final bool debug; + final String provider; +} + +/// Top-level configuration for [OfflineSpeechDenoiser]. +class OfflineSpeechDenoiserConfig { + const OfflineSpeechDenoiserConfig({ + this.model = const OfflineSpeechDenoiserModelConfig(), + }); + + factory OfflineSpeechDenoiserConfig.fromJson(Map json) { + return OfflineSpeechDenoiserConfig( + model: json['model'] != null + ? OfflineSpeechDenoiserModelConfig.fromJson( + json['model'] as Map) + : const OfflineSpeechDenoiserModelConfig(), + ); + } + + @override + String toString() { + return 'OfflineSpeechDenoiserConfig(model: $model)'; + } + + Map toJson() => { + 'model': model.toJson(), + }; + + final OfflineSpeechDenoiserModelConfig model; +} + +/// Audio returned by offline or online speech denoisers. +class DenoisedAudio { + DenoisedAudio({ + required this.samples, + required this.sampleRate, + }); + + final Float32List samples; + final int sampleRate; +} diff --git a/flutter/sherpa_onnx/lib/src/online_punctuation.dart b/flutter/sherpa_onnx/lib/src/online_punctuation.dart index e02aa0729d..4f07a6114e 100644 --- a/flutter/sherpa_onnx/lib/src/online_punctuation.dart +++ b/flutter/sherpa_onnx/lib/src/online_punctuation.dart @@ -2,78 +2,9 @@ import 'dart:ffi'; import 'package:ffi/ffi.dart'; import './sherpa_onnx_bindings.dart'; +import './online_punctuation_config.dart'; -/// Online punctuation restoration. -/// -/// This wrapper is intended for shorter or incremental text fragments. See -/// `dart-api-examples/add-punctuations/` for working examples. -class OnlinePunctuationModelConfig { - OnlinePunctuationModelConfig( - {required this.cnnBiLstm, - required this.bpeVocab, - this.numThreads = 1, - this.provider = 'cpu', - this.debug = true}); - - factory OnlinePunctuationModelConfig.fromJson(Map json) { - return OnlinePunctuationModelConfig( - cnnBiLstm: json['cnnBiLstm'], - bpeVocab: json['bpeVocab'], - numThreads: json['numThreads'], - provider: json['provider'], - debug: json['debug'], - ); - } - - @override - String toString() { - return 'OnlinePunctuationModelConfig(cnnBiLstm: $cnnBiLstm, ' - 'bpeVocab: $bpeVocab, numThreads: $numThreads, ' - 'provider: $provider, debug: $debug)'; - } - - Map toJson() { - return { - 'cnnBiLstm': cnnBiLstm, - 'bpeVocab': bpeVocab, - 'numThreads': numThreads, - 'provider': provider, - 'debug': debug, - }; - } - - final String cnnBiLstm; - final String bpeVocab; - final int numThreads; - final String provider; - final bool debug; -} - -/// Top-level configuration for [OnlinePunctuation]. -class OnlinePunctuationConfig { - OnlinePunctuationConfig({ - required this.model, - }); - - factory OnlinePunctuationConfig.fromJson(Map json) { - return OnlinePunctuationConfig( - model: OnlinePunctuationModelConfig.fromJson(json['model']), - ); - } - - @override - String toString() { - return 'OnlinePunctuationConfig(model: $model)'; - } - - Map toJson() { - return { - 'model': model.toJson(), - }; - } - - final OnlinePunctuationModelConfig model; -} +export './online_punctuation_config.dart'; /// Online punctuation restorer. class OnlinePunctuation { diff --git a/flutter/sherpa_onnx/lib/src/online_punctuation_config.dart b/flutter/sherpa_onnx/lib/src/online_punctuation_config.dart new file mode 100644 index 0000000000..ac47b2dad3 --- /dev/null +++ b/flutter/sherpa_onnx/lib/src/online_punctuation_config.dart @@ -0,0 +1,73 @@ +/// Shared config/data classes for online punctuation -- no FFI, works on all platforms. + +/// Online punctuation restoration. +/// +/// This wrapper is intended for shorter or incremental text fragments. See +/// `dart-api-examples/add-punctuations/` for working examples. +class OnlinePunctuationModelConfig { + OnlinePunctuationModelConfig( + {required this.cnnBiLstm, + required this.bpeVocab, + this.numThreads = 1, + this.provider = 'cpu', + this.debug = true}); + + factory OnlinePunctuationModelConfig.fromJson(Map json) { + return OnlinePunctuationModelConfig( + cnnBiLstm: json['cnnBiLstm'], + bpeVocab: json['bpeVocab'], + numThreads: json['numThreads'], + provider: json['provider'], + debug: json['debug'], + ); + } + + @override + String toString() { + return 'OnlinePunctuationModelConfig(cnnBiLstm: $cnnBiLstm, ' + 'bpeVocab: $bpeVocab, numThreads: $numThreads, ' + 'provider: $provider, debug: $debug)'; + } + + Map toJson() { + return { + 'cnnBiLstm': cnnBiLstm, + 'bpeVocab': bpeVocab, + 'numThreads': numThreads, + 'provider': provider, + 'debug': debug, + }; + } + + final String cnnBiLstm; + final String bpeVocab; + final int numThreads; + final String provider; + final bool debug; +} + +/// Top-level configuration for [OnlinePunctuation]. +class OnlinePunctuationConfig { + OnlinePunctuationConfig({ + required this.model, + }); + + factory OnlinePunctuationConfig.fromJson(Map json) { + return OnlinePunctuationConfig( + model: OnlinePunctuationModelConfig.fromJson(json['model']), + ); + } + + @override + String toString() { + return 'OnlinePunctuationConfig(model: $model)'; + } + + Map toJson() { + return { + 'model': model.toJson(), + }; + } + + final OnlinePunctuationModelConfig model; +} diff --git a/flutter/sherpa_onnx/lib/src/online_recognizer.dart b/flutter/sherpa_onnx/lib/src/online_recognizer.dart index 7fe4999ac7..5b34915761 100644 --- a/flutter/sherpa_onnx/lib/src/online_recognizer.dart +++ b/flutter/sherpa_onnx/lib/src/online_recognizer.dart @@ -4,12 +4,13 @@ import 'dart:ffi'; import 'package:ffi/ffi.dart'; -import './feature_config.dart'; -import './homophone_replacer_config.dart'; +import './online_recognizer_config.dart'; import './online_stream.dart'; import './sherpa_onnx_bindings.dart'; import './utils.dart'; +export './online_recognizer_config.dart'; + /// Streaming speech recognition. /// /// This module wraps the online ASR APIs used by the examples in @@ -38,367 +39,6 @@ import './utils.dart'; /// print(recognizer.getResult(stream).text); /// ``` -/// Model files for a streaming transducer recognizer. -class OnlineTransducerModelConfig { - const OnlineTransducerModelConfig({ - this.encoder = '', - this.decoder = '', - this.joiner = '', - }); - - factory OnlineTransducerModelConfig.fromJson(Map json) { - return OnlineTransducerModelConfig( - encoder: json['encoder'] as String? ?? '', - decoder: json['decoder'] as String? ?? '', - joiner: json['joiner'] as String? ?? '', - ); - } - - @override - String toString() { - return 'OnlineTransducerModelConfig(encoder: $encoder, decoder: $decoder, joiner: $joiner)'; - } - - Map toJson() => { - 'encoder': encoder, - 'decoder': decoder, - 'joiner': joiner, - }; - - final String encoder; - final String decoder; - final String joiner; -} - -/// Model files for a streaming Paraformer recognizer. -class OnlineParaformerModelConfig { - const OnlineParaformerModelConfig({this.encoder = '', this.decoder = ''}); - - factory OnlineParaformerModelConfig.fromJson(Map json) { - return OnlineParaformerModelConfig( - encoder: json['encoder'] as String? ?? '', - decoder: json['decoder'] as String? ?? '', - ); - } - - @override - String toString() { - return 'OnlineParaformerModelConfig(encoder: $encoder, decoder: $decoder)'; - } - - Map toJson() => { - 'encoder': encoder, - 'decoder': decoder, - }; - - final String encoder; - final String decoder; -} - -/// Model file for a streaming Zipformer2 CTC recognizer. -class OnlineZipformer2CtcModelConfig { - const OnlineZipformer2CtcModelConfig({this.model = ''}); - - factory OnlineZipformer2CtcModelConfig.fromJson(Map json) { - return OnlineZipformer2CtcModelConfig( - model: json['model'] as String? ?? '', - ); - } - - @override - String toString() { - return 'OnlineZipformer2CtcModelConfig(model: $model)'; - } - - Map toJson() => { - 'model': model, - }; - - final String model; -} - -/// Model file for a streaming NeMo CTC recognizer. -class OnlineNemoCtcModelConfig { - const OnlineNemoCtcModelConfig({this.model = ''}); - - factory OnlineNemoCtcModelConfig.fromJson(Map json) { - return OnlineNemoCtcModelConfig( - model: json['model'] as String? ?? '', - ); - } - - @override - String toString() { - return 'OnlineNemoCtcModelConfig(model: $model)'; - } - - Map toJson() => { - 'model': model, - }; - - final String model; -} - -/// Model file for a streaming tone-aware CTC recognizer. -class OnlineToneCtcModelConfig { - const OnlineToneCtcModelConfig({this.model = ''}); - - factory OnlineToneCtcModelConfig.fromJson(Map json) { - return OnlineToneCtcModelConfig( - model: json['model'] as String? ?? '', - ); - } - - @override - String toString() { - return 'OnlineToneCtcModelConfig(model: $model)'; - } - - Map toJson() => { - 'model': model, - }; - - final String model; -} - -/// Aggregate model configuration for streaming recognition. -/// -/// Configure exactly one model family for a typical deployment and supply the -/// shared tokenizer and runtime settings here. -class OnlineModelConfig { - const OnlineModelConfig({ - this.transducer = const OnlineTransducerModelConfig(), - this.paraformer = const OnlineParaformerModelConfig(), - this.zipformer2Ctc = const OnlineZipformer2CtcModelConfig(), - this.nemoCtc = const OnlineNemoCtcModelConfig(), - this.toneCtc = const OnlineToneCtcModelConfig(), - required this.tokens, - this.numThreads = 1, - this.provider = 'cpu', - this.debug = true, - this.modelType = '', - this.modelingUnit = '', - this.bpeVocab = '', - }); - - factory OnlineModelConfig.fromJson(Map json) { - return OnlineModelConfig( - transducer: OnlineTransducerModelConfig.fromJson( - json['transducer'] as Map? ?? const {}), - paraformer: OnlineParaformerModelConfig.fromJson( - json['paraformer'] as Map? ?? const {}), - zipformer2Ctc: OnlineZipformer2CtcModelConfig.fromJson( - json['zipformer2Ctc'] as Map? ?? const {}), - nemoCtc: OnlineNemoCtcModelConfig.fromJson( - json['nemoCtc'] as Map? ?? const {}), - toneCtc: OnlineToneCtcModelConfig.fromJson( - json['toneCtc'] as Map? ?? const {}), - tokens: json['tokens'] as String, - numThreads: json['numThreads'] as int? ?? 1, - provider: json['provider'] as String? ?? 'cpu', - debug: json['debug'] as bool? ?? true, - modelType: json['modelType'] as String? ?? '', - modelingUnit: json['modelingUnit'] as String? ?? '', - bpeVocab: json['bpeVocab'] as String? ?? '', - ); - } - - @override - String toString() { - return 'OnlineModelConfig(transducer: $transducer, paraformer: $paraformer, zipformer2Ctc: $zipformer2Ctc, nemoCtc: $nemoCtc, toneCtc: $toneCtc, tokens: $tokens, numThreads: $numThreads, provider: $provider, debug: $debug, modelType: $modelType, modelingUnit: $modelingUnit, bpeVocab: $bpeVocab)'; - } - - Map toJson() => { - 'transducer': transducer.toJson(), - 'paraformer': paraformer.toJson(), - 'zipformer2Ctc': zipformer2Ctc.toJson(), - 'nemoCtc': nemoCtc.toJson(), - 'toneCtc': toneCtc.toJson(), - 'tokens': tokens, - 'numThreads': numThreads, - 'provider': provider, - 'debug': debug, - 'modelType': modelType, - 'modelingUnit': modelingUnit, - 'bpeVocab': bpeVocab, - }; - - final OnlineTransducerModelConfig transducer; - final OnlineParaformerModelConfig paraformer; - final OnlineZipformer2CtcModelConfig zipformer2Ctc; - final OnlineNemoCtcModelConfig nemoCtc; - final OnlineToneCtcModelConfig toneCtc; - - final String tokens; - - final int numThreads; - - final String provider; - - final bool debug; - - final String modelType; - - final String modelingUnit; - - final String bpeVocab; -} - -/// FST decoder settings for CTC-based streaming recognition. -class OnlineCtcFstDecoderConfig { - const OnlineCtcFstDecoderConfig({this.graph = '', this.maxActive = 3000}); - - factory OnlineCtcFstDecoderConfig.fromJson(Map json) { - return OnlineCtcFstDecoderConfig( - graph: json['graph'] as String? ?? '', - maxActive: json['maxActive'] as int? ?? 3000, - ); - } - - @override - String toString() { - return 'OnlineCtcFstDecoderConfig(graph: $graph, maxActive: $maxActive)'; - } - - Map toJson() => { - 'graph': graph, - 'maxActive': maxActive, - }; - - final String graph; - final int maxActive; -} - -/// Top-level configuration for [OnlineRecognizer]. -/// -/// This combines feature extraction, the selected online model family, -/// endpointing rules, hotwords, grammar resources, and optional homophone -/// replacement resources. -class OnlineRecognizerConfig { - const OnlineRecognizerConfig({ - this.feat = const FeatureConfig(), - required this.model, - this.decodingMethod = 'greedy_search', - this.maxActivePaths = 4, - this.enableEndpoint = true, - this.rule1MinTrailingSilence = 2.4, - this.rule2MinTrailingSilence = 1.2, - this.rule3MinUtteranceLength = 20, - this.hotwordsFile = '', - this.hotwordsScore = 1.5, - this.ctcFstDecoderConfig = const OnlineCtcFstDecoderConfig(), - this.ruleFsts = '', - this.ruleFars = '', - this.blankPenalty = 0.0, - this.hr = const HomophoneReplacerConfig(), - }); - - factory OnlineRecognizerConfig.fromJson(Map json) { - return OnlineRecognizerConfig( - feat: FeatureConfig.fromJson( - json['feat'] as Map? ?? const {}), - model: OnlineModelConfig.fromJson(json['model'] as Map), - decodingMethod: json['decodingMethod'] as String? ?? 'greedy_search', - maxActivePaths: json['maxActivePaths'] as int? ?? 4, - enableEndpoint: json['enableEndpoint'] as bool? ?? true, - rule1MinTrailingSilence: - (json['rule1MinTrailingSilence'] as num?)?.toDouble() ?? 2.4, - rule2MinTrailingSilence: - (json['rule2MinTrailingSilence'] as num?)?.toDouble() ?? 1.2, - rule3MinUtteranceLength: - (json['rule3MinUtteranceLength'] as num?)?.toDouble() ?? 20.0, - hotwordsFile: json['hotwordsFile'] as String? ?? '', - hotwordsScore: (json['hotwordsScore'] as num?)?.toDouble() ?? 1.5, - ctcFstDecoderConfig: OnlineCtcFstDecoderConfig.fromJson( - json['ctcFstDecoderConfig'] as Map? ?? const {}), - ruleFsts: json['ruleFsts'] as String? ?? '', - ruleFars: json['ruleFars'] as String? ?? '', - blankPenalty: (json['blankPenalty'] as num?)?.toDouble() ?? 0.0, - hr: HomophoneReplacerConfig.fromJson( - json['hr'] as Map? ?? const {}), - ); - } - - @override - String toString() { - return 'OnlineRecognizerConfig(feat: $feat, model: $model, decodingMethod: $decodingMethod, maxActivePaths: $maxActivePaths, enableEndpoint: $enableEndpoint, rule1MinTrailingSilence: $rule1MinTrailingSilence, rule2MinTrailingSilence: $rule2MinTrailingSilence, rule3MinUtteranceLength: $rule3MinUtteranceLength, hotwordsFile: $hotwordsFile, hotwordsScore: $hotwordsScore, ctcFstDecoderConfig: $ctcFstDecoderConfig, ruleFsts: $ruleFsts, ruleFars: $ruleFars, blankPenalty: $blankPenalty, hr: $hr)'; - } - - Map toJson() => { - 'feat': feat.toJson(), - 'model': model.toJson(), - 'decodingMethod': decodingMethod, - 'maxActivePaths': maxActivePaths, - 'enableEndpoint': enableEndpoint, - 'rule1MinTrailingSilence': rule1MinTrailingSilence, - 'rule2MinTrailingSilence': rule2MinTrailingSilence, - 'rule3MinUtteranceLength': rule3MinUtteranceLength, - 'hotwordsFile': hotwordsFile, - 'hotwordsScore': hotwordsScore, - 'ctcFstDecoderConfig': ctcFstDecoderConfig.toJson(), - 'ruleFsts': ruleFsts, - 'ruleFars': ruleFars, - 'blankPenalty': blankPenalty, - 'hr': hr.toJson(), - }; - - final FeatureConfig feat; - final OnlineModelConfig model; - final String decodingMethod; - - final int maxActivePaths; - - final bool enableEndpoint; - - final double rule1MinTrailingSilence; - - final double rule2MinTrailingSilence; - - final double rule3MinUtteranceLength; - - final String hotwordsFile; - - final double hotwordsScore; - - final OnlineCtcFstDecoderConfig ctcFstDecoderConfig; - final String ruleFsts; - final String ruleFars; - - final double blankPenalty; - final HomophoneReplacerConfig hr; -} - -/// Streaming recognition result returned by [OnlineRecognizer.getResult]. -class OnlineRecognizerResult { - OnlineRecognizerResult( - {required this.text, required this.tokens, required this.timestamps}); - - factory OnlineRecognizerResult.fromJson(Map json) { - return OnlineRecognizerResult( - text: json['text'] as String, - tokens: List.from(json['tokens'] as List), - timestamps: (json['timestamps'] as List) - .map((e) => (e as num).toDouble()) - .toList(), - ); - } - - @override - String toString() { - return 'OnlineRecognizerResult(text: $text, tokens: $tokens, timestamps: $timestamps)'; - } - - Map toJson() => { - 'text': text, - 'tokens': tokens, - 'timestamps': timestamps, - }; - - final String text; - final List tokens; - final List timestamps; -} - /// Streaming speech recognizer. /// /// Create one from an [OnlineRecognizerConfig], then feed chunks to an diff --git a/flutter/sherpa_onnx/lib/src/online_recognizer_config.dart b/flutter/sherpa_onnx/lib/src/online_recognizer_config.dart new file mode 100644 index 0000000000..7f57a4995c --- /dev/null +++ b/flutter/sherpa_onnx/lib/src/online_recognizer_config.dart @@ -0,0 +1,366 @@ +// Copyright (c) 2024 Xiaomi Corporation +// Shared config/data classes for online recognition -- no FFI, works on all platforms. + +import './feature_config.dart'; +import './homophone_replacer_config.dart'; + +/// Model files for a streaming transducer recognizer. +class OnlineTransducerModelConfig { + const OnlineTransducerModelConfig({ + this.encoder = '', + this.decoder = '', + this.joiner = '', + }); + + factory OnlineTransducerModelConfig.fromJson(Map json) { + return OnlineTransducerModelConfig( + encoder: json['encoder'] as String? ?? '', + decoder: json['decoder'] as String? ?? '', + joiner: json['joiner'] as String? ?? '', + ); + } + + @override + String toString() { + return 'OnlineTransducerModelConfig(encoder: $encoder, decoder: $decoder, joiner: $joiner)'; + } + + Map toJson() => { + 'encoder': encoder, + 'decoder': decoder, + 'joiner': joiner, + }; + + final String encoder; + final String decoder; + final String joiner; +} + +/// Model files for a streaming Paraformer recognizer. +class OnlineParaformerModelConfig { + const OnlineParaformerModelConfig({this.encoder = '', this.decoder = ''}); + + factory OnlineParaformerModelConfig.fromJson(Map json) { + return OnlineParaformerModelConfig( + encoder: json['encoder'] as String? ?? '', + decoder: json['decoder'] as String? ?? '', + ); + } + + @override + String toString() { + return 'OnlineParaformerModelConfig(encoder: $encoder, decoder: $decoder)'; + } + + Map toJson() => { + 'encoder': encoder, + 'decoder': decoder, + }; + + final String encoder; + final String decoder; +} + +/// Model file for a streaming Zipformer2 CTC recognizer. +class OnlineZipformer2CtcModelConfig { + const OnlineZipformer2CtcModelConfig({this.model = ''}); + + factory OnlineZipformer2CtcModelConfig.fromJson(Map json) { + return OnlineZipformer2CtcModelConfig( + model: json['model'] as String? ?? '', + ); + } + + @override + String toString() { + return 'OnlineZipformer2CtcModelConfig(model: $model)'; + } + + Map toJson() => { + 'model': model, + }; + + final String model; +} + +/// Model file for a streaming NeMo CTC recognizer. +class OnlineNemoCtcModelConfig { + const OnlineNemoCtcModelConfig({this.model = ''}); + + factory OnlineNemoCtcModelConfig.fromJson(Map json) { + return OnlineNemoCtcModelConfig( + model: json['model'] as String? ?? '', + ); + } + + @override + String toString() { + return 'OnlineNemoCtcModelConfig(model: $model)'; + } + + Map toJson() => { + 'model': model, + }; + + final String model; +} + +/// Model file for a streaming tone-aware CTC recognizer. +class OnlineToneCtcModelConfig { + const OnlineToneCtcModelConfig({this.model = ''}); + + factory OnlineToneCtcModelConfig.fromJson(Map json) { + return OnlineToneCtcModelConfig( + model: json['model'] as String? ?? '', + ); + } + + @override + String toString() { + return 'OnlineToneCtcModelConfig(model: $model)'; + } + + Map toJson() => { + 'model': model, + }; + + final String model; +} + +/// Aggregate model configuration for streaming recognition. +/// +/// Configure exactly one model family for a typical deployment and supply the +/// shared tokenizer and runtime settings here. +class OnlineModelConfig { + const OnlineModelConfig({ + this.transducer = const OnlineTransducerModelConfig(), + this.paraformer = const OnlineParaformerModelConfig(), + this.zipformer2Ctc = const OnlineZipformer2CtcModelConfig(), + this.nemoCtc = const OnlineNemoCtcModelConfig(), + this.toneCtc = const OnlineToneCtcModelConfig(), + required this.tokens, + this.numThreads = 1, + this.provider = 'cpu', + this.debug = true, + this.modelType = '', + this.modelingUnit = '', + this.bpeVocab = '', + }); + + factory OnlineModelConfig.fromJson(Map json) { + return OnlineModelConfig( + transducer: OnlineTransducerModelConfig.fromJson( + json['transducer'] as Map? ?? const {}), + paraformer: OnlineParaformerModelConfig.fromJson( + json['paraformer'] as Map? ?? const {}), + zipformer2Ctc: OnlineZipformer2CtcModelConfig.fromJson( + json['zipformer2Ctc'] as Map? ?? const {}), + nemoCtc: OnlineNemoCtcModelConfig.fromJson( + json['nemoCtc'] as Map? ?? const {}), + toneCtc: OnlineToneCtcModelConfig.fromJson( + json['toneCtc'] as Map? ?? const {}), + tokens: json['tokens'] as String, + numThreads: json['numThreads'] as int? ?? 1, + provider: json['provider'] as String? ?? 'cpu', + debug: json['debug'] as bool? ?? true, + modelType: json['modelType'] as String? ?? '', + modelingUnit: json['modelingUnit'] as String? ?? '', + bpeVocab: json['bpeVocab'] as String? ?? '', + ); + } + + @override + String toString() { + return 'OnlineModelConfig(transducer: $transducer, paraformer: $paraformer, zipformer2Ctc: $zipformer2Ctc, nemoCtc: $nemoCtc, toneCtc: $toneCtc, tokens: $tokens, numThreads: $numThreads, provider: $provider, debug: $debug, modelType: $modelType, modelingUnit: $modelingUnit, bpeVocab: $bpeVocab)'; + } + + Map toJson() => { + 'transducer': transducer.toJson(), + 'paraformer': paraformer.toJson(), + 'zipformer2Ctc': zipformer2Ctc.toJson(), + 'nemoCtc': nemoCtc.toJson(), + 'toneCtc': toneCtc.toJson(), + 'tokens': tokens, + 'numThreads': numThreads, + 'provider': provider, + 'debug': debug, + 'modelType': modelType, + 'modelingUnit': modelingUnit, + 'bpeVocab': bpeVocab, + }; + + final OnlineTransducerModelConfig transducer; + final OnlineParaformerModelConfig paraformer; + final OnlineZipformer2CtcModelConfig zipformer2Ctc; + final OnlineNemoCtcModelConfig nemoCtc; + final OnlineToneCtcModelConfig toneCtc; + + final String tokens; + + final int numThreads; + + final String provider; + + final bool debug; + + final String modelType; + + final String modelingUnit; + + final String bpeVocab; +} + +/// FST decoder settings for CTC-based streaming recognition. +class OnlineCtcFstDecoderConfig { + const OnlineCtcFstDecoderConfig({this.graph = '', this.maxActive = 3000}); + + factory OnlineCtcFstDecoderConfig.fromJson(Map json) { + return OnlineCtcFstDecoderConfig( + graph: json['graph'] as String? ?? '', + maxActive: json['maxActive'] as int? ?? 3000, + ); + } + + @override + String toString() { + return 'OnlineCtcFstDecoderConfig(graph: $graph, maxActive: $maxActive)'; + } + + Map toJson() => { + 'graph': graph, + 'maxActive': maxActive, + }; + + final String graph; + final int maxActive; +} + +/// Top-level configuration for [OnlineRecognizer]. +/// +/// This combines feature extraction, the selected online model family, +/// endpointing rules, hotwords, grammar resources, and optional homophone +/// replacement resources. +class OnlineRecognizerConfig { + const OnlineRecognizerConfig({ + this.feat = const FeatureConfig(), + required this.model, + this.decodingMethod = 'greedy_search', + this.maxActivePaths = 4, + this.enableEndpoint = true, + this.rule1MinTrailingSilence = 2.4, + this.rule2MinTrailingSilence = 1.2, + this.rule3MinUtteranceLength = 20, + this.hotwordsFile = '', + this.hotwordsScore = 1.5, + this.ctcFstDecoderConfig = const OnlineCtcFstDecoderConfig(), + this.ruleFsts = '', + this.ruleFars = '', + this.blankPenalty = 0.0, + this.hr = const HomophoneReplacerConfig(), + }); + + factory OnlineRecognizerConfig.fromJson(Map json) { + return OnlineRecognizerConfig( + feat: FeatureConfig.fromJson( + json['feat'] as Map? ?? const {}), + model: OnlineModelConfig.fromJson(json['model'] as Map), + decodingMethod: json['decodingMethod'] as String? ?? 'greedy_search', + maxActivePaths: json['maxActivePaths'] as int? ?? 4, + enableEndpoint: json['enableEndpoint'] as bool? ?? true, + rule1MinTrailingSilence: + (json['rule1MinTrailingSilence'] as num?)?.toDouble() ?? 2.4, + rule2MinTrailingSilence: + (json['rule2MinTrailingSilence'] as num?)?.toDouble() ?? 1.2, + rule3MinUtteranceLength: + (json['rule3MinUtteranceLength'] as num?)?.toDouble() ?? 20.0, + hotwordsFile: json['hotwordsFile'] as String? ?? '', + hotwordsScore: (json['hotwordsScore'] as num?)?.toDouble() ?? 1.5, + ctcFstDecoderConfig: OnlineCtcFstDecoderConfig.fromJson( + json['ctcFstDecoderConfig'] as Map? ?? const {}), + ruleFsts: json['ruleFsts'] as String? ?? '', + ruleFars: json['ruleFars'] as String? ?? '', + blankPenalty: (json['blankPenalty'] as num?)?.toDouble() ?? 0.0, + hr: HomophoneReplacerConfig.fromJson( + json['hr'] as Map? ?? const {}), + ); + } + + @override + String toString() { + return 'OnlineRecognizerConfig(feat: $feat, model: $model, decodingMethod: $decodingMethod, maxActivePaths: $maxActivePaths, enableEndpoint: $enableEndpoint, rule1MinTrailingSilence: $rule1MinTrailingSilence, rule2MinTrailingSilence: $rule2MinTrailingSilence, rule3MinUtteranceLength: $rule3MinUtteranceLength, hotwordsFile: $hotwordsFile, hotwordsScore: $hotwordsScore, ctcFstDecoderConfig: $ctcFstDecoderConfig, ruleFsts: $ruleFsts, ruleFars: $ruleFars, blankPenalty: $blankPenalty, hr: $hr)'; + } + + Map toJson() => { + 'feat': feat.toJson(), + 'model': model.toJson(), + 'decodingMethod': decodingMethod, + 'maxActivePaths': maxActivePaths, + 'enableEndpoint': enableEndpoint, + 'rule1MinTrailingSilence': rule1MinTrailingSilence, + 'rule2MinTrailingSilence': rule2MinTrailingSilence, + 'rule3MinUtteranceLength': rule3MinUtteranceLength, + 'hotwordsFile': hotwordsFile, + 'hotwordsScore': hotwordsScore, + 'ctcFstDecoderConfig': ctcFstDecoderConfig.toJson(), + 'ruleFsts': ruleFsts, + 'ruleFars': ruleFars, + 'blankPenalty': blankPenalty, + 'hr': hr.toJson(), + }; + + final FeatureConfig feat; + final OnlineModelConfig model; + final String decodingMethod; + + final int maxActivePaths; + + final bool enableEndpoint; + + final double rule1MinTrailingSilence; + + final double rule2MinTrailingSilence; + + final double rule3MinUtteranceLength; + + final String hotwordsFile; + + final double hotwordsScore; + + final OnlineCtcFstDecoderConfig ctcFstDecoderConfig; + final String ruleFsts; + final String ruleFars; + + final double blankPenalty; + final HomophoneReplacerConfig hr; +} + +/// Streaming recognition result returned by [OnlineRecognizer.getResult]. +class OnlineRecognizerResult { + OnlineRecognizerResult( + {required this.text, required this.tokens, required this.timestamps}); + + factory OnlineRecognizerResult.fromJson(Map json) { + return OnlineRecognizerResult( + text: json['text'] as String, + tokens: List.from(json['tokens'] as List), + timestamps: (json['timestamps'] as List) + .map((e) => (e as num).toDouble()) + .toList(), + ); + } + + @override + String toString() { + return 'OnlineRecognizerResult(text: $text, tokens: $tokens, timestamps: $timestamps)'; + } + + Map toJson() => { + 'text': text, + 'tokens': tokens, + 'timestamps': timestamps, + }; + + final String text; + final List tokens; + final List timestamps; +} diff --git a/flutter/sherpa_onnx/lib/src/online_speech_denoiser.dart b/flutter/sherpa_onnx/lib/src/online_speech_denoiser.dart index a75040d8a8..36cf29d937 100644 --- a/flutter/sherpa_onnx/lib/src/online_speech_denoiser.dart +++ b/flutter/sherpa_onnx/lib/src/online_speech_denoiser.dart @@ -5,38 +5,10 @@ import 'dart:typed_data'; import 'package:ffi/ffi.dart'; import './offline_speech_denoiser.dart'; +import './online_speech_denoiser_config.dart'; import './sherpa_onnx_bindings.dart'; -/// Streaming speech denoising. -/// -/// Call [run] on consecutive chunks, then [flush] after the final chunk to -/// drain any buffered state. -class OnlineSpeechDenoiserConfig { - const OnlineSpeechDenoiserConfig({ - this.model = const OfflineSpeechDenoiserModelConfig(), - }); - - factory OnlineSpeechDenoiserConfig.fromJson(Map json) { - return OnlineSpeechDenoiserConfig( - model: json['model'] != null - ? OfflineSpeechDenoiserModelConfig.fromJson( - json['model'] as Map, - ) - : const OfflineSpeechDenoiserModelConfig(), - ); - } - - @override - String toString() { - return 'OnlineSpeechDenoiserConfig(model: $model)'; - } - - Map toJson() => { - 'model': model.toJson(), - }; - - final OfflineSpeechDenoiserModelConfig model; -} +export './online_speech_denoiser_config.dart'; /// Streaming speech denoiser. class OnlineSpeechDenoiser { diff --git a/flutter/sherpa_onnx/lib/src/online_speech_denoiser_config.dart b/flutter/sherpa_onnx/lib/src/online_speech_denoiser_config.dart new file mode 100644 index 0000000000..aa9d17b386 --- /dev/null +++ b/flutter/sherpa_onnx/lib/src/online_speech_denoiser_config.dart @@ -0,0 +1,35 @@ +// Copyright (c) 2026 Xiaomi Corporation +// Shared config/data classes for online speech denoising -- no FFI, works on all platforms. + +import './offline_speech_denoiser_config.dart'; + +/// Streaming speech denoising. +/// +/// Call [run] on consecutive chunks, then [flush] after the final chunk to +/// drain any buffered state. +class OnlineSpeechDenoiserConfig { + const OnlineSpeechDenoiserConfig({ + this.model = const OfflineSpeechDenoiserModelConfig(), + }); + + factory OnlineSpeechDenoiserConfig.fromJson(Map json) { + return OnlineSpeechDenoiserConfig( + model: json['model'] != null + ? OfflineSpeechDenoiserModelConfig.fromJson( + json['model'] as Map, + ) + : const OfflineSpeechDenoiserModelConfig(), + ); + } + + @override + String toString() { + return 'OnlineSpeechDenoiserConfig(model: $model)'; + } + + Map toJson() => { + 'model': model.toJson(), + }; + + final OfflineSpeechDenoiserModelConfig model; +} diff --git a/flutter/sherpa_onnx/lib/src/speaker_identification.dart b/flutter/sherpa_onnx/lib/src/speaker_identification.dart index 0b52139249..3d4cbd9d2b 100644 --- a/flutter/sherpa_onnx/lib/src/speaker_identification.dart +++ b/flutter/sherpa_onnx/lib/src/speaker_identification.dart @@ -5,62 +5,9 @@ import 'package:ffi/ffi.dart'; import './online_stream.dart'; import './sherpa_onnx_bindings.dart'; +import './speaker_identification_config.dart'; -/// Speaker embedding extraction and speaker identification utilities. -/// -/// See `dart-api-examples/speaker-identification/` for end-to-end examples. -/// -/// Example: -/// -/// ```dart -/// final extractor = SpeakerEmbeddingExtractor( -/// config: const SpeakerEmbeddingExtractorConfig( -/// model: './3dspeaker_speech_eres2net_base_sv_zh-cn_3dspeaker_16k.onnx', -/// ), -/// ); -/// -/// final stream = extractor.createStream(); -/// stream.acceptWaveform(samples: wave.samples, sampleRate: wave.sampleRate); -/// while (extractor.isReady(stream)) {} -/// final embedding = extractor.compute(stream); -/// -/// final manager = SpeakerEmbeddingManager(extractor.dim); -/// manager.add(name: 'alice', embedding: embedding); -/// print(manager.search(embedding: embedding, threshold: 0.6)); -/// ``` -class SpeakerEmbeddingExtractorConfig { - const SpeakerEmbeddingExtractorConfig( - {required this.model, - this.numThreads = 1, - this.debug = true, - this.provider = 'cpu'}); - - factory SpeakerEmbeddingExtractorConfig.fromJson(Map json) { - return SpeakerEmbeddingExtractorConfig( - model: json['model'] as String, - numThreads: json['numThreads'] as int? ?? 1, - debug: json['debug'] as bool? ?? true, - provider: json['provider'] as String? ?? 'cpu', - ); - } - - @override - String toString() { - return 'SpeakerEmbeddingExtractorConfig(model: $model, numThreads: $numThreads, debug: $debug, provider: $provider)'; - } - - Map toJson() => { - 'model': model, - 'numThreads': numThreads, - 'debug': debug, - 'provider': provider, - }; - - final String model; - final int numThreads; - final bool debug; - final String provider; -} +export './speaker_identification_config.dart'; /// Speaker embedding extractor. /// diff --git a/flutter/sherpa_onnx/lib/src/speaker_identification_config.dart b/flutter/sherpa_onnx/lib/src/speaker_identification_config.dart new file mode 100644 index 0000000000..22648f6622 --- /dev/null +++ b/flutter/sherpa_onnx/lib/src/speaker_identification_config.dart @@ -0,0 +1,58 @@ +// Copyright (c) 2024 Xiaomi Corporation +// Shared config/data classes for speaker identification -- no FFI, works on all platforms. + +/// Speaker embedding extraction and speaker identification utilities. +/// +/// See `dart-api-examples/speaker-identification/` for end-to-end examples. +/// +/// Example: +/// +/// ```dart +/// final extractor = SpeakerEmbeddingExtractor( +/// config: const SpeakerEmbeddingExtractorConfig( +/// model: './3dspeaker_speech_eres2net_base_sv_zh-cn_3dspeaker_16k.onnx', +/// ), +/// ); +/// +/// final stream = extractor.createStream(); +/// stream.acceptWaveform(samples: wave.samples, sampleRate: wave.sampleRate); +/// while (extractor.isReady(stream)) {} +/// final embedding = extractor.compute(stream); +/// +/// final manager = SpeakerEmbeddingManager(extractor.dim); +/// manager.add(name: 'alice', embedding: embedding); +/// print(manager.search(embedding: embedding, threshold: 0.6)); +/// ``` +class SpeakerEmbeddingExtractorConfig { + const SpeakerEmbeddingExtractorConfig( + {required this.model, + this.numThreads = 1, + this.debug = true, + this.provider = 'cpu'}); + + factory SpeakerEmbeddingExtractorConfig.fromJson(Map json) { + return SpeakerEmbeddingExtractorConfig( + model: json['model'] as String, + numThreads: json['numThreads'] as int? ?? 1, + debug: json['debug'] as bool? ?? true, + provider: json['provider'] as String? ?? 'cpu', + ); + } + + @override + String toString() { + return 'SpeakerEmbeddingExtractorConfig(model: $model, numThreads: $numThreads, debug: $debug, provider: $provider)'; + } + + Map toJson() => { + 'model': model, + 'numThreads': numThreads, + 'debug': debug, + 'provider': provider, + }; + + final String model; + final int numThreads; + final bool debug; + final String provider; +} diff --git a/flutter/sherpa_onnx/lib/src/spoken_language_identification.dart b/flutter/sherpa_onnx/lib/src/spoken_language_identification.dart index 73a51fded8..f86f3dde7f 100644 --- a/flutter/sherpa_onnx/lib/src/spoken_language_identification.dart +++ b/flutter/sherpa_onnx/lib/src/spoken_language_identification.dart @@ -6,124 +6,9 @@ import 'package:ffi/ffi.dart'; import './offline_stream.dart'; import './sherpa_onnx_bindings.dart'; import './utils.dart'; +import './spoken_language_identification_config.dart'; -/// Spoken language identification. -/// -/// This module identifies the language spoken in an audio clip, using the -/// Whisper-based language ID model family exposed by the native library. -/// -/// Example: -/// -/// ```dart -/// final sli = SpokenLanguageIdentification( -/// SpokenLanguageIdentificationConfig( -/// whisper: const SpokenLanguageIdentificationWhisperConfig( -/// encoder: './sherpa-onnx-whisper-tiny/encoder.int8.onnx', -/// decoder: './sherpa-onnx-whisper-tiny/decoder.int8.onnx', -/// ), -/// ), -/// ); -/// -/// final stream = sli.createStream(); -/// stream.acceptWaveform(samples: wave.samples, sampleRate: wave.sampleRate); -/// print(sli.compute(stream).lang); -/// ``` -class SpokenLanguageIdentificationWhisperConfig { - const SpokenLanguageIdentificationWhisperConfig({ - this.encoder = '', - this.decoder = '', - this.tailPaddings = 0, - }); - - factory SpokenLanguageIdentificationWhisperConfig.fromJson( - Map json) { - return SpokenLanguageIdentificationWhisperConfig( - encoder: json['encoder'] as String? ?? '', - decoder: json['decoder'] as String? ?? '', - tailPaddings: json['tailPaddings'] as int? ?? 0, - ); - } - - @override - String toString() { - return 'SpokenLanguageIdentificationWhisperConfig(encoder: $encoder, decoder: $decoder, tailPaddings: $tailPaddings)'; - } - - Map toJson() => { - 'encoder': encoder, - 'decoder': decoder, - 'tailPaddings': tailPaddings, - }; - - final String encoder; - final String decoder; - final int tailPaddings; -} - -/// Top-level configuration for [SpokenLanguageIdentification]. -class SpokenLanguageIdentificationConfig { - const SpokenLanguageIdentificationConfig({ - this.whisper = const SpokenLanguageIdentificationWhisperConfig(), - this.numThreads = 1, - this.debug = false, - this.provider = 'cpu', - }); - - factory SpokenLanguageIdentificationConfig.fromJson( - Map json) { - return SpokenLanguageIdentificationConfig( - whisper: json['whisper'] != null - ? SpokenLanguageIdentificationWhisperConfig.fromJson( - json['whisper'] as Map) - : const SpokenLanguageIdentificationWhisperConfig(), - numThreads: json['numThreads'] as int? ?? 1, - debug: json['debug'] as bool? ?? false, - provider: json['provider'] as String? ?? 'cpu', - ); - } - - @override - String toString() { - return 'SpokenLanguageIdentificationConfig(whisper: $whisper, numThreads: $numThreads, debug: $debug, provider: $provider)'; - } - - Map toJson() => { - 'whisper': whisper.toJson(), - 'numThreads': numThreads, - 'debug': debug, - 'provider': provider, - }; - - final SpokenLanguageIdentificationWhisperConfig whisper; - final int numThreads; - final bool debug; - final String provider; -} - -/// Result returned by [SpokenLanguageIdentification.compute]. -class SpokenLanguageIdentificationResult { - const SpokenLanguageIdentificationResult({ - required this.lang, - }); - - factory SpokenLanguageIdentificationResult.fromJson( - Map json) { - return SpokenLanguageIdentificationResult( - lang: json['lang'] as String? ?? '', - ); - } - - @override - String toString() { - return 'SpokenLanguageIdentificationResult(lang: $lang)'; - } - - Map toJson() => { - 'lang': lang, - }; - - final String lang; -} +export './spoken_language_identification_config.dart'; /// Spoken language identifier. class SpokenLanguageIdentification { diff --git a/flutter/sherpa_onnx/lib/src/spoken_language_identification_config.dart b/flutter/sherpa_onnx/lib/src/spoken_language_identification_config.dart new file mode 100644 index 0000000000..bac40f186b --- /dev/null +++ b/flutter/sherpa_onnx/lib/src/spoken_language_identification_config.dart @@ -0,0 +1,100 @@ +// Copyright (c) 2024 Xiaomi Corporation +// Shared config/data classes for spoken language identification -- no FFI, works on all platforms. + +/// Model files for spoken language identification using Whisper. +class SpokenLanguageIdentificationWhisperConfig { + const SpokenLanguageIdentificationWhisperConfig({ + this.encoder = '', + this.decoder = '', + this.tailPaddings = 0, + }); + + factory SpokenLanguageIdentificationWhisperConfig.fromJson( + Map json) { + return SpokenLanguageIdentificationWhisperConfig( + encoder: json['encoder'] as String? ?? '', + decoder: json['decoder'] as String? ?? '', + tailPaddings: json['tailPaddings'] as int? ?? 0, + ); + } + + @override + String toString() { + return 'SpokenLanguageIdentificationWhisperConfig(encoder: $encoder, decoder: $decoder, tailPaddings: $tailPaddings)'; + } + + Map toJson() => { + 'encoder': encoder, + 'decoder': decoder, + 'tailPaddings': tailPaddings, + }; + + final String encoder; + final String decoder; + final int tailPaddings; +} + +/// Top-level configuration for [SpokenLanguageIdentification]. +class SpokenLanguageIdentificationConfig { + const SpokenLanguageIdentificationConfig({ + this.whisper = const SpokenLanguageIdentificationWhisperConfig(), + this.numThreads = 1, + this.debug = false, + this.provider = 'cpu', + }); + + factory SpokenLanguageIdentificationConfig.fromJson( + Map json) { + return SpokenLanguageIdentificationConfig( + whisper: json['whisper'] != null + ? SpokenLanguageIdentificationWhisperConfig.fromJson( + json['whisper'] as Map) + : const SpokenLanguageIdentificationWhisperConfig(), + numThreads: json['numThreads'] as int? ?? 1, + debug: json['debug'] as bool? ?? false, + provider: json['provider'] as String? ?? 'cpu', + ); + } + + @override + String toString() { + return 'SpokenLanguageIdentificationConfig(whisper: $whisper, numThreads: $numThreads, debug: $debug, provider: $provider)'; + } + + Map toJson() => { + 'whisper': whisper.toJson(), + 'numThreads': numThreads, + 'debug': debug, + 'provider': provider, + }; + + final SpokenLanguageIdentificationWhisperConfig whisper; + final int numThreads; + final bool debug; + final String provider; +} + +/// Result returned by [SpokenLanguageIdentification.compute]. +class SpokenLanguageIdentificationResult { + const SpokenLanguageIdentificationResult({ + required this.lang, + }); + + factory SpokenLanguageIdentificationResult.fromJson( + Map json) { + return SpokenLanguageIdentificationResult( + lang: json['lang'] as String? ?? '', + ); + } + + @override + String toString() { + return 'SpokenLanguageIdentificationResult(lang: $lang)'; + } + + Map toJson() => { + 'lang': lang, + }; + + final String lang; +} diff --git a/flutter/sherpa_onnx/lib/src/tts.dart b/flutter/sherpa_onnx/lib/src/tts.dart index 6b3800d5c5..cd82a8ae21 100644 --- a/flutter/sherpa_onnx/lib/src/tts.dart +++ b/flutter/sherpa_onnx/lib/src/tts.dart @@ -6,6 +6,9 @@ import 'dart:typed_data'; import 'package:ffi/ffi.dart'; import './sherpa_onnx_bindings.dart'; +import './tts_config.dart'; + +export './tts_config.dart'; /// Offline text-to-speech. /// @@ -40,23 +43,8 @@ import './sherpa_onnx_bindings.dart'; /// tts.free(); /// ``` -/// Per-request generation options for [OfflineTts.generateWithConfig]. -/// -/// Use this when you need advanced generation controls such as zero-shot voice -/// cloning reference audio, explicit reference sample rate, or model-specific -/// values in [extra]. -class OfflineTtsGenerationConfig { - const OfflineTtsGenerationConfig({ - this.silenceScale = 0.2, - this.speed = 1.0, - this.sid = 0, - this.referenceAudio, - this.referenceSampleRate = 0, - this.referenceText = '', - this.numSteps = 5, - this.extra = const {}, - }); - +/// FFI bridge methods for [OfflineTtsGenerationConfig]. +extension OfflineTtsGenerationConfigFfi on OfflineTtsGenerationConfig { /// Convert Extra to JSON string. /// Returns nullptr if empty. /// The user should use calloc.free(p); to free the returned value @@ -119,501 +107,6 @@ class OfflineTtsGenerationConfig { } calloc.free(p); } - - final double silenceScale; - final double speed; - final int sid; - - /// mono audio in [-1, 1] - final Float32List? referenceAudio; - final int referenceSampleRate; - final String referenceText; - final int numSteps; - - /// Extra model-specific attributes - /// key: string - /// value: string | int | double - final Map extra; -} - -/// VITS model configuration. -class OfflineTtsVitsModelConfig { - const OfflineTtsVitsModelConfig({ - this.model = '', - this.lexicon = '', - this.tokens = '', - this.dataDir = '', - this.noiseScale = 0.667, - this.noiseScaleW = 0.8, - this.lengthScale = 1.0, - this.dictDir = '', - }); - - factory OfflineTtsVitsModelConfig.fromJson(Map json) { - return OfflineTtsVitsModelConfig( - model: json['model'] as String? ?? '', - lexicon: json['lexicon'] as String? ?? '', - tokens: json['tokens'] as String? ?? '', - dataDir: json['dataDir'] as String? ?? '', - noiseScale: (json['noiseScale'] as num?)?.toDouble() ?? 0.667, - noiseScaleW: (json['noiseScaleW'] as num?)?.toDouble() ?? 0.8, - lengthScale: (json['lengthScale'] as num?)?.toDouble() ?? 1.0, - ); - } - - @override - String toString() { - return 'OfflineTtsVitsModelConfig(model: $model, lexicon: $lexicon, tokens: $tokens, dataDir: $dataDir, noiseScale: $noiseScale, noiseScaleW: $noiseScaleW, lengthScale: $lengthScale)'; - } - - Map toJson() => { - 'model': model, - 'lexicon': lexicon, - 'tokens': tokens, - 'dataDir': dataDir, - 'noiseScale': noiseScale, - 'noiseScaleW': noiseScaleW, - 'lengthScale': lengthScale, - }; - - final String model; - final String lexicon; - final String tokens; - final String dataDir; - final double noiseScale; - final double noiseScaleW; - final double lengthScale; - final String dictDir; // unused -} - -/// Matcha model configuration. -class OfflineTtsMatchaModelConfig { - const OfflineTtsMatchaModelConfig({ - this.acousticModel = '', - this.vocoder = '', - this.lexicon = '', - this.tokens = '', - this.dataDir = '', - this.noiseScale = 0.667, - this.lengthScale = 1.0, - this.dictDir = '', - }); - - factory OfflineTtsMatchaModelConfig.fromJson(Map json) { - return OfflineTtsMatchaModelConfig( - acousticModel: json['acousticModel'] as String? ?? '', - vocoder: json['vocoder'] as String? ?? '', - lexicon: json['lexicon'] as String? ?? '', - tokens: json['tokens'] as String? ?? '', - dataDir: json['dataDir'] as String? ?? '', - noiseScale: (json['noiseScale'] as num?)?.toDouble() ?? 0.667, - lengthScale: (json['lengthScale'] as num?)?.toDouble() ?? 1.0, - ); - } - - @override - String toString() { - return 'OfflineTtsMatchaModelConfig(acousticModel: $acousticModel, vocoder: $vocoder, lexicon: $lexicon, tokens: $tokens, dataDir: $dataDir, noiseScale: $noiseScale, lengthScale: $lengthScale)'; - } - - Map toJson() => { - 'acousticModel': acousticModel, - 'vocoder': vocoder, - 'lexicon': lexicon, - 'tokens': tokens, - 'dataDir': dataDir, - 'noiseScale': noiseScale, - 'lengthScale': lengthScale, - }; - - final String acousticModel; - final String vocoder; - final String lexicon; - final String tokens; - final String dataDir; - final double noiseScale; - final double lengthScale; - final String dictDir; // unused -} - -/// Kokoro model configuration. -class OfflineTtsKokoroModelConfig { - const OfflineTtsKokoroModelConfig({ - this.model = '', - this.voices = '', - this.tokens = '', - this.dataDir = '', - this.lengthScale = 1.0, - this.dictDir = '', - this.lexicon = '', - this.lang = '', - }); - - factory OfflineTtsKokoroModelConfig.fromJson(Map json) { - return OfflineTtsKokoroModelConfig( - model: json['model'] as String? ?? '', - voices: json['voices'] as String? ?? '', - tokens: json['tokens'] as String? ?? '', - dataDir: json['dataDir'] as String? ?? '', - lengthScale: (json['lengthScale'] as num?)?.toDouble() ?? 1.0, - lexicon: json['lexicon'] as String? ?? '', - lang: json['lang'] as String? ?? '', - ); - } - - @override - String toString() { - return 'OfflineTtsKokoroModelConfig(model: $model, voices: $voices, tokens: $tokens, dataDir: $dataDir, lengthScale: $lengthScale, lexicon: $lexicon, lang: $lang)'; - } - - Map toJson() => { - 'model': model, - 'voices': voices, - 'tokens': tokens, - 'dataDir': dataDir, - 'lengthScale': lengthScale, - 'lexicon': lexicon, - 'lang': lang, - }; - - final String model; - final String voices; - final String tokens; - final String dataDir; - final double lengthScale; - final String dictDir; // unused - final String lexicon; - final String lang; -} - -/// Kitten model configuration. -class OfflineTtsKittenModelConfig { - const OfflineTtsKittenModelConfig({ - this.model = '', - this.voices = '', - this.tokens = '', - this.dataDir = '', - this.lengthScale = 1.0, - }); - - factory OfflineTtsKittenModelConfig.fromJson(Map json) { - return OfflineTtsKittenModelConfig( - model: json['model'] as String? ?? '', - voices: json['voices'] as String? ?? '', - tokens: json['tokens'] as String? ?? '', - dataDir: json['dataDir'] as String? ?? '', - lengthScale: (json['lengthScale'] as num?)?.toDouble() ?? 1.0, - ); - } - - @override - String toString() { - return 'OfflineTtsKittenModelConfig(model: $model, voices: $voices, tokens: $tokens, dataDir: $dataDir, lengthScale: $lengthScale)'; - } - - Map toJson() => { - 'model': model, - 'voices': voices, - 'tokens': tokens, - 'dataDir': dataDir, - 'lengthScale': lengthScale, - }; - - final String model; - final String voices; - final String tokens; - final String dataDir; - final double lengthScale; -} - -/// ZipVoice model configuration. -class OfflineTtsZipVoiceModelConfig { - const OfflineTtsZipVoiceModelConfig({ - this.tokens = '', - this.encoder = '', - this.decoder = '', - this.vocoder = '', - this.dataDir = '', - this.lexicon = '', - this.featScale = 0.1, - this.tShift = 0.5, - this.targetRms = 0.1, - this.guidanceScale = 1.0, - }); - - factory OfflineTtsZipVoiceModelConfig.fromJson(Map json) { - return OfflineTtsZipVoiceModelConfig( - tokens: json['tokens'] as String? ?? '', - encoder: json['encoder'] as String? ?? '', - decoder: json['decoder'] as String? ?? '', - vocoder: json['vocoder'] as String? ?? '', - dataDir: json['dataDir'] as String? ?? '', - lexicon: json['lexicon'] as String? ?? '', - featScale: (json['featScale'] as num?)?.toDouble() ?? 0.1, - tShift: (json['tShift'] as num?)?.toDouble() ?? 0.5, - targetRms: (json['targetRms'] as num?)?.toDouble() ?? 0.1, - guidanceScale: (json['guidanceScale'] as num?)?.toDouble() ?? 1.0, - ); - } - - @override - String toString() { - return 'OfflineTtsZipVoiceModelConfig(tokens: $tokens, encoder: $encoder, decoder: $decoder, vocoder: $vocoder, dataDir: $dataDir, lexicon: $lexicon, featScale: $featScale, tShift: $tShift, targetRms: $targetRms, guidanceScale: $guidanceScale)'; - } - - Map toJson() => { - 'tokens': tokens, - 'encoder': encoder, - 'decoder': decoder, - 'vocoder': vocoder, - 'dataDir': dataDir, - 'lexicon': lexicon, - 'featScale': featScale, - 'tShift': tShift, - 'targetRms': targetRms, - 'guidanceScale': guidanceScale, - }; - - final String tokens; - final String encoder; - final String decoder; - final String vocoder; - final String dataDir; - final String lexicon; - final double featScale; - final double tShift; - final double targetRms; - final double guidanceScale; -} - -/// Pocket TTS model configuration. -/// -/// This family supports zero-shot voice cloning with a reference waveform. -class OfflineTtsPocketModelConfig { - const OfflineTtsPocketModelConfig({ - this.lmFlow = '', - this.lmMain = '', - this.encoder = '', - this.decoder = '', - this.textConditioner = '', - this.vocabJson = '', - this.tokenScoresJson = '', - this.voiceEmbeddingCacheCapacity = 50, - }); - - factory OfflineTtsPocketModelConfig.fromJson(Map json) { - return OfflineTtsPocketModelConfig( - lmFlow: json['lmFlow'] as String? ?? '', - lmMain: json['lmMain'] as String? ?? '', - encoder: json['encoder'] as String? ?? '', - decoder: json['decoder'] as String? ?? '', - textConditioner: json['textConditioner'] as String? ?? '', - vocabJson: json['vocabJson'] as String? ?? '', - tokenScoresJson: json['tokenScoresJson'] as String? ?? '', - voiceEmbeddingCacheCapacity: - json['voiceEmbeddingCacheCapacity'] as int? ?? 50, - ); - } - - Map toJson() => { - 'lmFlow': lmFlow, - 'lmMain': lmMain, - 'encoder': encoder, - 'decoder': decoder, - 'textConditioner': textConditioner, - 'vocabJson': vocabJson, - 'tokenScoresJson': tokenScoresJson, - 'voiceEmbeddingCacheCapacity': voiceEmbeddingCacheCapacity, - }; - - @override - String toString() { - return 'OfflineTtsPocketModelConfig(lmFlow: $lmFlow, lmMain: $lmMain, encoder: $encoder, decoder: $decoder, textConditioner: $textConditioner, vocabJson: $vocabJson, tokenScoresJson: $tokenScoresJson, voiceEmbeddingCacheCapacity: $voiceEmbeddingCacheCapacity)'; - } - - final String lmFlow; - final String lmMain; - final String encoder; - final String decoder; - final String textConditioner; - final String vocabJson; - final String tokenScoresJson; - final int voiceEmbeddingCacheCapacity; -} - -/// Supertonic model configuration. -class OfflineTtsSupertonicModelConfig { - const OfflineTtsSupertonicModelConfig({ - this.durationPredictor = '', - this.textEncoder = '', - this.vectorEstimator = '', - this.vocoder = '', - this.ttsJson = '', - this.unicodeIndexer = '', - this.voiceStyle = '', - }); - - factory OfflineTtsSupertonicModelConfig.fromJson(Map json) { - return OfflineTtsSupertonicModelConfig( - durationPredictor: json['durationPredictor'] as String? ?? '', - textEncoder: json['textEncoder'] as String? ?? '', - vectorEstimator: json['vectorEstimator'] as String? ?? '', - vocoder: json['vocoder'] as String? ?? '', - ttsJson: json['ttsJson'] as String? ?? '', - unicodeIndexer: json['unicodeIndexer'] as String? ?? '', - voiceStyle: json['voiceStyle'] as String? ?? '', - ); - } - - Map toJson() => { - 'durationPredictor': durationPredictor, - 'textEncoder': textEncoder, - 'vectorEstimator': vectorEstimator, - 'vocoder': vocoder, - 'ttsJson': ttsJson, - 'unicodeIndexer': unicodeIndexer, - 'voiceStyle': voiceStyle, - }; - - @override - String toString() { - return 'OfflineTtsSupertonicModelConfig(durationPredictor: $durationPredictor, textEncoder: $textEncoder, vectorEstimator: $vectorEstimator, vocoder: $vocoder, ttsJson: $ttsJson, unicodeIndexer: $unicodeIndexer, voiceStyle: $voiceStyle)'; - } - - final String durationPredictor; - final String textEncoder; - final String vectorEstimator; - final String vocoder; - final String ttsJson; - final String unicodeIndexer; - final String voiceStyle; -} - -/// Aggregate model configuration for offline TTS. -/// -/// Configure exactly one model family for a typical setup and set the shared -/// runtime options such as [numThreads] and [provider]. -class OfflineTtsModelConfig { - const OfflineTtsModelConfig({ - this.vits = const OfflineTtsVitsModelConfig(), - this.matcha = const OfflineTtsMatchaModelConfig(), - this.kokoro = const OfflineTtsKokoroModelConfig(), - this.kitten = const OfflineTtsKittenModelConfig(), - this.zipvoice = const OfflineTtsZipVoiceModelConfig(), - this.pocket = const OfflineTtsPocketModelConfig(), - this.supertonic = const OfflineTtsSupertonicModelConfig(), - this.numThreads = 1, - this.debug = true, - this.provider = 'cpu', - }); - - factory OfflineTtsModelConfig.fromJson(Map json) { - return OfflineTtsModelConfig( - vits: OfflineTtsVitsModelConfig.fromJson( - json['vits'] as Map? ?? const {}, - ), - matcha: OfflineTtsMatchaModelConfig.fromJson( - json['matcha'] as Map? ?? const {}, - ), - kokoro: OfflineTtsKokoroModelConfig.fromJson( - json['kokoro'] as Map? ?? const {}, - ), - kitten: OfflineTtsKittenModelConfig.fromJson( - json['kitten'] as Map? ?? const {}, - ), - zipvoice: OfflineTtsZipVoiceModelConfig.fromJson( - json['zipvoice'] as Map? ?? const {}, - ), - pocket: OfflineTtsPocketModelConfig.fromJson( - json['pocket'] as Map? ?? const {}, - ), - supertonic: OfflineTtsSupertonicModelConfig.fromJson( - json['supertonic'] as Map? ?? const {}, - ), - numThreads: json['numThreads'] as int? ?? 1, - debug: json['debug'] as bool? ?? true, - provider: json['provider'] as String? ?? 'cpu', - ); - } - - @override - String toString() { - return 'OfflineTtsModelConfig(vits: $vits, matcha: $matcha, kokoro: $kokoro, kitten: $kitten, zipvoice: $zipvoice, pocket: $pocket, supertonic: $supertonic, numThreads: $numThreads, debug: $debug, provider: $provider)'; - } - - Map toJson() => { - 'vits': vits.toJson(), - 'matcha': matcha.toJson(), - 'kokoro': kokoro.toJson(), - 'kitten': kitten.toJson(), - 'zipvoice': zipvoice.toJson(), - 'pocket': pocket.toJson(), - 'supertonic': supertonic.toJson(), - 'numThreads': numThreads, - 'debug': debug, - 'provider': provider, - }; - - final OfflineTtsVitsModelConfig vits; - final OfflineTtsMatchaModelConfig matcha; - final OfflineTtsKokoroModelConfig kokoro; - final OfflineTtsKittenModelConfig kitten; - final OfflineTtsZipVoiceModelConfig zipvoice; - final OfflineTtsPocketModelConfig pocket; - final OfflineTtsSupertonicModelConfig supertonic; - final int numThreads; - final bool debug; - final String provider; -} - -/// Top-level configuration for [OfflineTts]. -class OfflineTtsConfig { - const OfflineTtsConfig({ - required this.model, - this.ruleFsts = '', - this.maxNumSenetences = 1, - this.ruleFars = '', - this.silenceScale = 0.2, - }); - - factory OfflineTtsConfig.fromJson(Map json) { - return OfflineTtsConfig( - model: OfflineTtsModelConfig.fromJson( - json['model'] as Map, - ), - ruleFsts: json['ruleFsts'] as String? ?? '', - maxNumSenetences: json['maxNumSenetences'] as int? ?? 1, - ruleFars: json['ruleFars'] as String? ?? '', - silenceScale: (json['silenceScale'] as num?)?.toDouble() ?? 0.2, - ); - } - - @override - String toString() { - return 'OfflineTtsConfig(model: $model, ruleFsts: $ruleFsts, maxNumSenetences: $maxNumSenetences, ruleFars: $ruleFars, silenceScale: $silenceScale)'; - } - - Map toJson() => { - 'model': model.toJson(), - 'ruleFsts': ruleFsts, - 'maxNumSenetences': maxNumSenetences, - 'ruleFars': ruleFars, - 'silenceScale': silenceScale, - }; - - final OfflineTtsModelConfig model; - final String ruleFsts; - final int maxNumSenetences; - final String ruleFars; - final double silenceScale; -} - -/// Audio generated by [OfflineTts]. -class GeneratedAudio { - GeneratedAudio({required this.samples, required this.sampleRate}); - - final Float32List samples; - final int sampleRate; } /// Offline text-to-speech engine. diff --git a/flutter/sherpa_onnx/lib/src/tts_config.dart b/flutter/sherpa_onnx/lib/src/tts_config.dart new file mode 100644 index 0000000000..48a35df693 --- /dev/null +++ b/flutter/sherpa_onnx/lib/src/tts_config.dart @@ -0,0 +1,516 @@ +// Copyright (c) 2024 Xiaomi Corporation +// Shared config/data classes for offline TTS -- no FFI, works on all platforms. +import 'dart:typed_data'; + +/// Per-request generation options for [OfflineTts.generateWithConfig]. +/// +/// Use this when you need advanced generation controls such as zero-shot voice +/// cloning reference audio, explicit reference sample rate, or model-specific +/// values in [extra]. +class OfflineTtsGenerationConfig { + const OfflineTtsGenerationConfig({ + this.silenceScale = 0.2, + this.speed = 1.0, + this.sid = 0, + this.referenceAudio, + this.referenceSampleRate = 0, + this.referenceText = '', + this.numSteps = 5, + this.extra = const {}, + }); + + final double silenceScale; + final double speed; + final int sid; + + /// mono audio in [-1, 1] + final Float32List? referenceAudio; + final int referenceSampleRate; + final String referenceText; + final int numSteps; + + /// Extra model-specific attributes + /// key: string + /// value: string | int | double + final Map extra; +} + +/// VITS model configuration. +class OfflineTtsVitsModelConfig { + const OfflineTtsVitsModelConfig({ + this.model = '', + this.lexicon = '', + this.tokens = '', + this.dataDir = '', + this.noiseScale = 0.667, + this.noiseScaleW = 0.8, + this.lengthScale = 1.0, + this.dictDir = '', + }); + + factory OfflineTtsVitsModelConfig.fromJson(Map json) { + return OfflineTtsVitsModelConfig( + model: json['model'] as String? ?? '', + lexicon: json['lexicon'] as String? ?? '', + tokens: json['tokens'] as String? ?? '', + dataDir: json['dataDir'] as String? ?? '', + noiseScale: (json['noiseScale'] as num?)?.toDouble() ?? 0.667, + noiseScaleW: (json['noiseScaleW'] as num?)?.toDouble() ?? 0.8, + lengthScale: (json['lengthScale'] as num?)?.toDouble() ?? 1.0, + ); + } + + @override + String toString() { + return 'OfflineTtsVitsModelConfig(model: $model, lexicon: $lexicon, tokens: $tokens, dataDir: $dataDir, noiseScale: $noiseScale, noiseScaleW: $noiseScaleW, lengthScale: $lengthScale)'; + } + + Map toJson() => { + 'model': model, + 'lexicon': lexicon, + 'tokens': tokens, + 'dataDir': dataDir, + 'noiseScale': noiseScale, + 'noiseScaleW': noiseScaleW, + 'lengthScale': lengthScale, + }; + + final String model; + final String lexicon; + final String tokens; + final String dataDir; + final double noiseScale; + final double noiseScaleW; + final double lengthScale; + final String dictDir; // unused +} + +/// Matcha model configuration. +class OfflineTtsMatchaModelConfig { + const OfflineTtsMatchaModelConfig({ + this.acousticModel = '', + this.vocoder = '', + this.lexicon = '', + this.tokens = '', + this.dataDir = '', + this.noiseScale = 0.667, + this.lengthScale = 1.0, + this.dictDir = '', + }); + + factory OfflineTtsMatchaModelConfig.fromJson(Map json) { + return OfflineTtsMatchaModelConfig( + acousticModel: json['acousticModel'] as String? ?? '', + vocoder: json['vocoder'] as String? ?? '', + lexicon: json['lexicon'] as String? ?? '', + tokens: json['tokens'] as String? ?? '', + dataDir: json['dataDir'] as String? ?? '', + noiseScale: (json['noiseScale'] as num?)?.toDouble() ?? 0.667, + lengthScale: (json['lengthScale'] as num?)?.toDouble() ?? 1.0, + ); + } + + @override + String toString() { + return 'OfflineTtsMatchaModelConfig(acousticModel: $acousticModel, vocoder: $vocoder, lexicon: $lexicon, tokens: $tokens, dataDir: $dataDir, noiseScale: $noiseScale, lengthScale: $lengthScale)'; + } + + Map toJson() => { + 'acousticModel': acousticModel, + 'vocoder': vocoder, + 'lexicon': lexicon, + 'tokens': tokens, + 'dataDir': dataDir, + 'noiseScale': noiseScale, + 'lengthScale': lengthScale, + }; + + final String acousticModel; + final String vocoder; + final String lexicon; + final String tokens; + final String dataDir; + final double noiseScale; + final double lengthScale; + final String dictDir; // unused +} + +/// Kokoro model configuration. +class OfflineTtsKokoroModelConfig { + const OfflineTtsKokoroModelConfig({ + this.model = '', + this.voices = '', + this.tokens = '', + this.dataDir = '', + this.lengthScale = 1.0, + this.dictDir = '', + this.lexicon = '', + this.lang = '', + }); + + factory OfflineTtsKokoroModelConfig.fromJson(Map json) { + return OfflineTtsKokoroModelConfig( + model: json['model'] as String? ?? '', + voices: json['voices'] as String? ?? '', + tokens: json['tokens'] as String? ?? '', + dataDir: json['dataDir'] as String? ?? '', + lengthScale: (json['lengthScale'] as num?)?.toDouble() ?? 1.0, + lexicon: json['lexicon'] as String? ?? '', + lang: json['lang'] as String? ?? '', + ); + } + + @override + String toString() { + return 'OfflineTtsKokoroModelConfig(model: $model, voices: $voices, tokens: $tokens, dataDir: $dataDir, lengthScale: $lengthScale, lexicon: $lexicon, lang: $lang)'; + } + + Map toJson() => { + 'model': model, + 'voices': voices, + 'tokens': tokens, + 'dataDir': dataDir, + 'lengthScale': lengthScale, + 'lexicon': lexicon, + 'lang': lang, + }; + + final String model; + final String voices; + final String tokens; + final String dataDir; + final double lengthScale; + final String dictDir; // unused + final String lexicon; + final String lang; +} + +/// Kitten model configuration. +class OfflineTtsKittenModelConfig { + const OfflineTtsKittenModelConfig({ + this.model = '', + this.voices = '', + this.tokens = '', + this.dataDir = '', + this.lengthScale = 1.0, + }); + + factory OfflineTtsKittenModelConfig.fromJson(Map json) { + return OfflineTtsKittenModelConfig( + model: json['model'] as String? ?? '', + voices: json['voices'] as String? ?? '', + tokens: json['tokens'] as String? ?? '', + dataDir: json['dataDir'] as String? ?? '', + lengthScale: (json['lengthScale'] as num?)?.toDouble() ?? 1.0, + ); + } + + @override + String toString() { + return 'OfflineTtsKittenModelConfig(model: $model, voices: $voices, tokens: $tokens, dataDir: $dataDir, lengthScale: $lengthScale)'; + } + + Map toJson() => { + 'model': model, + 'voices': voices, + 'tokens': tokens, + 'dataDir': dataDir, + 'lengthScale': lengthScale, + }; + + final String model; + final String voices; + final String tokens; + final String dataDir; + final double lengthScale; +} + +/// ZipVoice model configuration. +class OfflineTtsZipVoiceModelConfig { + const OfflineTtsZipVoiceModelConfig({ + this.tokens = '', + this.encoder = '', + this.decoder = '', + this.vocoder = '', + this.dataDir = '', + this.lexicon = '', + this.featScale = 0.1, + this.tShift = 0.5, + this.targetRms = 0.1, + this.guidanceScale = 1.0, + }); + + factory OfflineTtsZipVoiceModelConfig.fromJson(Map json) { + return OfflineTtsZipVoiceModelConfig( + tokens: json['tokens'] as String? ?? '', + encoder: json['encoder'] as String? ?? '', + decoder: json['decoder'] as String? ?? '', + vocoder: json['vocoder'] as String? ?? '', + dataDir: json['dataDir'] as String? ?? '', + lexicon: json['lexicon'] as String? ?? '', + featScale: (json['featScale'] as num?)?.toDouble() ?? 0.1, + tShift: (json['tShift'] as num?)?.toDouble() ?? 0.5, + targetRms: (json['targetRms'] as num?)?.toDouble() ?? 0.1, + guidanceScale: (json['guidanceScale'] as num?)?.toDouble() ?? 1.0, + ); + } + + @override + String toString() { + return 'OfflineTtsZipVoiceModelConfig(tokens: $tokens, encoder: $encoder, decoder: $decoder, vocoder: $vocoder, dataDir: $dataDir, lexicon: $lexicon, featScale: $featScale, tShift: $tShift, targetRms: $targetRms, guidanceScale: $guidanceScale)'; + } + + Map toJson() => { + 'tokens': tokens, + 'encoder': encoder, + 'decoder': decoder, + 'vocoder': vocoder, + 'dataDir': dataDir, + 'lexicon': lexicon, + 'featScale': featScale, + 'tShift': tShift, + 'targetRms': targetRms, + 'guidanceScale': guidanceScale, + }; + + final String tokens; + final String encoder; + final String decoder; + final String vocoder; + final String dataDir; + final String lexicon; + final double featScale; + final double tShift; + final double targetRms; + final double guidanceScale; +} + +/// Pocket TTS model configuration. +/// +/// This family supports zero-shot voice cloning with a reference waveform. +class OfflineTtsPocketModelConfig { + const OfflineTtsPocketModelConfig({ + this.lmFlow = '', + this.lmMain = '', + this.encoder = '', + this.decoder = '', + this.textConditioner = '', + this.vocabJson = '', + this.tokenScoresJson = '', + this.voiceEmbeddingCacheCapacity = 50, + }); + + factory OfflineTtsPocketModelConfig.fromJson(Map json) { + return OfflineTtsPocketModelConfig( + lmFlow: json['lmFlow'] as String? ?? '', + lmMain: json['lmMain'] as String? ?? '', + encoder: json['encoder'] as String? ?? '', + decoder: json['decoder'] as String? ?? '', + textConditioner: json['textConditioner'] as String? ?? '', + vocabJson: json['vocabJson'] as String? ?? '', + tokenScoresJson: json['tokenScoresJson'] as String? ?? '', + voiceEmbeddingCacheCapacity: + json['voiceEmbeddingCacheCapacity'] as int? ?? 50, + ); + } + + Map toJson() => { + 'lmFlow': lmFlow, + 'lmMain': lmMain, + 'encoder': encoder, + 'decoder': decoder, + 'textConditioner': textConditioner, + 'vocabJson': vocabJson, + 'tokenScoresJson': tokenScoresJson, + 'voiceEmbeddingCacheCapacity': voiceEmbeddingCacheCapacity, + }; + + @override + String toString() { + return 'OfflineTtsPocketModelConfig(lmFlow: $lmFlow, lmMain: $lmMain, encoder: $encoder, decoder: $decoder, textConditioner: $textConditioner, vocabJson: $vocabJson, tokenScoresJson: $tokenScoresJson, voiceEmbeddingCacheCapacity: $voiceEmbeddingCacheCapacity)'; + } + + final String lmFlow; + final String lmMain; + final String encoder; + final String decoder; + final String textConditioner; + final String vocabJson; + final String tokenScoresJson; + final int voiceEmbeddingCacheCapacity; +} + +/// Supertonic model configuration. +class OfflineTtsSupertonicModelConfig { + const OfflineTtsSupertonicModelConfig({ + this.durationPredictor = '', + this.textEncoder = '', + this.vectorEstimator = '', + this.vocoder = '', + this.ttsJson = '', + this.unicodeIndexer = '', + this.voiceStyle = '', + }); + + factory OfflineTtsSupertonicModelConfig.fromJson(Map json) { + return OfflineTtsSupertonicModelConfig( + durationPredictor: json['durationPredictor'] as String? ?? '', + textEncoder: json['textEncoder'] as String? ?? '', + vectorEstimator: json['vectorEstimator'] as String? ?? '', + vocoder: json['vocoder'] as String? ?? '', + ttsJson: json['ttsJson'] as String? ?? '', + unicodeIndexer: json['unicodeIndexer'] as String? ?? '', + voiceStyle: json['voiceStyle'] as String? ?? '', + ); + } + + Map toJson() => { + 'durationPredictor': durationPredictor, + 'textEncoder': textEncoder, + 'vectorEstimator': vectorEstimator, + 'vocoder': vocoder, + 'ttsJson': ttsJson, + 'unicodeIndexer': unicodeIndexer, + 'voiceStyle': voiceStyle, + }; + + @override + String toString() { + return 'OfflineTtsSupertonicModelConfig(durationPredictor: $durationPredictor, textEncoder: $textEncoder, vectorEstimator: $vectorEstimator, vocoder: $vocoder, ttsJson: $ttsJson, unicodeIndexer: $unicodeIndexer, voiceStyle: $voiceStyle)'; + } + + final String durationPredictor; + final String textEncoder; + final String vectorEstimator; + final String vocoder; + final String ttsJson; + final String unicodeIndexer; + final String voiceStyle; +} + +/// Aggregate model configuration for offline TTS. +/// +/// Configure exactly one model family for a typical setup and set the shared +/// runtime options such as [numThreads] and [provider]. +class OfflineTtsModelConfig { + const OfflineTtsModelConfig({ + this.vits = const OfflineTtsVitsModelConfig(), + this.matcha = const OfflineTtsMatchaModelConfig(), + this.kokoro = const OfflineTtsKokoroModelConfig(), + this.kitten = const OfflineTtsKittenModelConfig(), + this.zipvoice = const OfflineTtsZipVoiceModelConfig(), + this.pocket = const OfflineTtsPocketModelConfig(), + this.supertonic = const OfflineTtsSupertonicModelConfig(), + this.numThreads = 1, + this.debug = true, + this.provider = 'cpu', + }); + + factory OfflineTtsModelConfig.fromJson(Map json) { + return OfflineTtsModelConfig( + vits: OfflineTtsVitsModelConfig.fromJson( + json['vits'] as Map? ?? const {}, + ), + matcha: OfflineTtsMatchaModelConfig.fromJson( + json['matcha'] as Map? ?? const {}, + ), + kokoro: OfflineTtsKokoroModelConfig.fromJson( + json['kokoro'] as Map? ?? const {}, + ), + kitten: OfflineTtsKittenModelConfig.fromJson( + json['kitten'] as Map? ?? const {}, + ), + zipvoice: OfflineTtsZipVoiceModelConfig.fromJson( + json['zipvoice'] as Map? ?? const {}, + ), + pocket: OfflineTtsPocketModelConfig.fromJson( + json['pocket'] as Map? ?? const {}, + ), + supertonic: OfflineTtsSupertonicModelConfig.fromJson( + json['supertonic'] as Map? ?? const {}, + ), + numThreads: json['numThreads'] as int? ?? 1, + debug: json['debug'] as bool? ?? true, + provider: json['provider'] as String? ?? 'cpu', + ); + } + + @override + String toString() { + return 'OfflineTtsModelConfig(vits: $vits, matcha: $matcha, kokoro: $kokoro, kitten: $kitten, zipvoice: $zipvoice, pocket: $pocket, supertonic: $supertonic, numThreads: $numThreads, debug: $debug, provider: $provider)'; + } + + Map toJson() => { + 'vits': vits.toJson(), + 'matcha': matcha.toJson(), + 'kokoro': kokoro.toJson(), + 'kitten': kitten.toJson(), + 'zipvoice': zipvoice.toJson(), + 'pocket': pocket.toJson(), + 'supertonic': supertonic.toJson(), + 'numThreads': numThreads, + 'debug': debug, + 'provider': provider, + }; + + final OfflineTtsVitsModelConfig vits; + final OfflineTtsMatchaModelConfig matcha; + final OfflineTtsKokoroModelConfig kokoro; + final OfflineTtsKittenModelConfig kitten; + final OfflineTtsZipVoiceModelConfig zipvoice; + final OfflineTtsPocketModelConfig pocket; + final OfflineTtsSupertonicModelConfig supertonic; + final int numThreads; + final bool debug; + final String provider; +} + +/// Top-level configuration for [OfflineTts]. +class OfflineTtsConfig { + const OfflineTtsConfig({ + required this.model, + this.ruleFsts = '', + this.maxNumSenetences = 1, + this.ruleFars = '', + this.silenceScale = 0.2, + }); + + factory OfflineTtsConfig.fromJson(Map json) { + return OfflineTtsConfig( + model: OfflineTtsModelConfig.fromJson( + json['model'] as Map, + ), + ruleFsts: json['ruleFsts'] as String? ?? '', + maxNumSenetences: json['maxNumSenetences'] as int? ?? 1, + ruleFars: json['ruleFars'] as String? ?? '', + silenceScale: (json['silenceScale'] as num?)?.toDouble() ?? 0.2, + ); + } + + @override + String toString() { + return 'OfflineTtsConfig(model: $model, ruleFsts: $ruleFsts, maxNumSenetences: $maxNumSenetences, ruleFars: $ruleFars, silenceScale: $silenceScale)'; + } + + Map toJson() => { + 'model': model.toJson(), + 'ruleFsts': ruleFsts, + 'maxNumSenetences': maxNumSenetences, + 'ruleFars': ruleFars, + 'silenceScale': silenceScale, + }; + + final OfflineTtsModelConfig model; + final String ruleFsts; + final int maxNumSenetences; + final String ruleFars; + final double silenceScale; +} + +/// Audio generated by [OfflineTts]. +class GeneratedAudio { + GeneratedAudio({required this.samples, required this.sampleRate}); + + final Float32List samples; + final int sampleRate; +} diff --git a/flutter/sherpa_onnx/lib/src/vad.dart b/flutter/sherpa_onnx/lib/src/vad.dart index 70ce736a25..b82a944d7b 100644 --- a/flutter/sherpa_onnx/lib/src/vad.dart +++ b/flutter/sherpa_onnx/lib/src/vad.dart @@ -4,180 +4,9 @@ import 'dart:typed_data'; import 'package:ffi/ffi.dart'; import './sherpa_onnx_bindings.dart'; +import './vad_config.dart'; -/// Voice activity detection and buffering helpers. -/// -/// See `dart-api-examples/vad/bin/vad.dart` and -/// `dart-api-examples/vad/bin/ten-vad.dart` for complete examples. -/// -/// Example: -/// -/// ```dart -/// final config = VadModelConfig( -/// sileroVad: const SileroVadModelConfig( -/// model: './silero_vad.onnx', -/// minSilenceDuration: 0.25, -/// minSpeechDuration: 0.5, -/// ), -/// numThreads: 1, -/// ); -/// -/// final vad = VoiceActivityDetector(config: config, bufferSizeInSeconds: 10); -/// final wave = readWave('./test.wav'); -/// vad.acceptWaveform(wave.samples); -/// vad.flush(); -/// while (!vad.isEmpty()) { -/// print(vad.front()); -/// vad.pop(); -/// } -/// vad.free(); -/// ``` - -/// Silero VAD model configuration. -class SileroVadModelConfig { - const SileroVadModelConfig( - {this.model = '', - this.threshold = 0.5, - this.minSilenceDuration = 0.5, - this.minSpeechDuration = 0.25, - this.windowSize = 512, - this.maxSpeechDuration = 5.0}); - - factory SileroVadModelConfig.fromJson(Map json) { - return SileroVadModelConfig( - model: json['model'] as String? ?? '', - threshold: (json['threshold'] as num?)?.toDouble() ?? 0.5, - minSilenceDuration: - (json['minSilenceDuration'] as num?)?.toDouble() ?? 0.5, - minSpeechDuration: - (json['minSpeechDuration'] as num?)?.toDouble() ?? 0.25, - windowSize: json['windowSize'] as int? ?? 512, - maxSpeechDuration: (json['maxSpeechDuration'] as num?)?.toDouble() ?? 5.0, - ); - } - - @override - String toString() { - return 'SileroVadModelConfig(model: $model, threshold: $threshold, minSilenceDuration: $minSilenceDuration, minSpeechDuration: $minSpeechDuration, windowSize: $windowSize, maxSpeechDuration: $maxSpeechDuration)'; - } - - Map toJson() => { - 'model': model, - 'threshold': threshold, - 'minSilenceDuration': minSilenceDuration, - 'minSpeechDuration': minSpeechDuration, - 'windowSize': windowSize, - 'maxSpeechDuration': maxSpeechDuration, - }; - - final String model; - final double threshold; - final double minSilenceDuration; - final double minSpeechDuration; - final int windowSize; - final double maxSpeechDuration; -} - -/// Ten VAD model configuration. -class TenVadModelConfig { - const TenVadModelConfig( - {this.model = '', - this.threshold = 0.5, - this.minSilenceDuration = 0.5, - this.minSpeechDuration = 0.25, - this.windowSize = 256, - this.maxSpeechDuration = 5.0}); - - factory TenVadModelConfig.fromJson(Map json) { - return TenVadModelConfig( - model: json['model'] as String? ?? '', - threshold: (json['threshold'] as num?)?.toDouble() ?? 0.5, - minSilenceDuration: - (json['minSilenceDuration'] as num?)?.toDouble() ?? 0.5, - minSpeechDuration: - (json['minSpeechDuration'] as num?)?.toDouble() ?? 0.25, - windowSize: json['windowSize'] as int? ?? 256, - maxSpeechDuration: (json['maxSpeechDuration'] as num?)?.toDouble() ?? 5.0, - ); - } - - @override - String toString() { - return 'TenVadModelConfig(model: $model, threshold: $threshold, minSilenceDuration: $minSilenceDuration, minSpeechDuration: $minSpeechDuration, windowSize: $windowSize, maxSpeechDuration: $maxSpeechDuration)'; - } - - Map toJson() => { - 'model': model, - 'threshold': threshold, - 'minSilenceDuration': minSilenceDuration, - 'minSpeechDuration': minSpeechDuration, - 'windowSize': windowSize, - 'maxSpeechDuration': maxSpeechDuration, - }; - - final String model; - final double threshold; - final double minSilenceDuration; - final double minSpeechDuration; - final int windowSize; - final double maxSpeechDuration; -} - -/// Top-level VAD model configuration. -/// -/// Configure either [sileroVad] or [tenVad] for typical use and set the shared -/// sample rate and runtime settings here. -class VadModelConfig { - VadModelConfig({ - this.sileroVad = const SileroVadModelConfig(), - this.sampleRate = 16000, - this.numThreads = 1, - this.provider = 'cpu', - this.debug = true, - this.tenVad = const TenVadModelConfig(), - }); - - final SileroVadModelConfig sileroVad; - final TenVadModelConfig tenVad; - final int sampleRate; - final int numThreads; - final String provider; - final bool debug; - - factory VadModelConfig.fromJson(Map json) { - return VadModelConfig( - sileroVad: SileroVadModelConfig.fromJson( - json['sileroVad'] as Map? ?? const {}), - tenVad: TenVadModelConfig.fromJson( - json['tenVad'] as Map? ?? const {}), - sampleRate: json['sampleRate'] as int? ?? 16000, - numThreads: json['numThreads'] as int? ?? 1, - provider: json['provider'] as String? ?? 'cpu', - debug: json['debug'] as bool? ?? true, - ); - } - - Map toJson() => { - 'sileroVad': sileroVad.toJson(), - 'tenVad': tenVad.toJson(), - 'sampleRate': sampleRate, - 'numThreads': numThreads, - 'provider': provider, - 'debug': debug, - }; - - @override - String toString() { - return 'VadModelConfig(sileroVad: $sileroVad, tenVad: $tenVad, sampleRate: $sampleRate, numThreads: $numThreads, provider: $provider, debug: $debug)'; - } -} - -/// One detected speech segment emitted by [VoiceActivityDetector]. -class SpeechSegment { - SpeechSegment({required this.samples, required this.start}); - final Float32List samples; - final int start; -} +export './vad_config.dart'; /// Circular sample buffer used by VAD-related pipelines. class CircularBuffer { @@ -206,7 +35,6 @@ class CircularBuffer { } /// Release the native buffer. - /// Release the native detector. void free() { if (SherpaOnnxBindings.destroyCircularBuffer == null) { throw Exception("Please initialize sherpa-onnx first"); @@ -279,7 +107,6 @@ class CircularBuffer { } /// Clear the buffer contents. - /// Reset the detector state. void reset() { if (SherpaOnnxBindings.circularBufferReset == null) { throw Exception("Please initialize sherpa-onnx first"); @@ -319,16 +146,11 @@ class CircularBuffer { } /// Voice activity detector that emits [SpeechSegment] objects. -/// -/// Create one with a [VadModelConfig], feed audio with [acceptWaveform], then -/// inspect queued segments with [isEmpty], [front], [pop], and [flush]. class VoiceActivityDetector { VoiceActivityDetector.fromPtr({required this.ptr, required this.config}); VoiceActivityDetector._({required this.ptr, required this.config}); - // The user has to invoke VoiceActivityDetector.free() to avoid memory leak. - /// Create a detector with an internal result buffer sized in seconds. factory VoiceActivityDetector( {required VadModelConfig config, required double bufferSizeInSeconds}) { if (SherpaOnnxBindings.createVoiceActivityDetector == null) { @@ -391,7 +213,6 @@ class VoiceActivityDetector { ptr = nullptr; } - /// Feed normalized waveform samples into the detector. void acceptWaveform(Float32List samples) { if (SherpaOnnxBindings.voiceActivityDetectorAcceptWaveform == null) { throw Exception("Please initialize sherpa-onnx first"); @@ -412,7 +233,6 @@ class VoiceActivityDetector { calloc.free(p); } - /// Return `true` if there are no queued speech segments. bool isEmpty() { if (SherpaOnnxBindings.voiceActivityDetectorEmpty == null) { throw Exception("Please initialize sherpa-onnx first"); @@ -428,7 +248,6 @@ class VoiceActivityDetector { return empty == 1; } - /// Return `true` if speech is currently being detected. bool isDetected() { if (SherpaOnnxBindings.voiceActivityDetectorDetected == null) { throw Exception("Please initialize sherpa-onnx first"); @@ -444,7 +263,6 @@ class VoiceActivityDetector { return detected == 1; } - /// Drop the front queued speech segment. void pop() { if (SherpaOnnxBindings.voiceActivityDetectorPop == null) { throw Exception("Please initialize sherpa-onnx first"); @@ -456,7 +274,6 @@ class VoiceActivityDetector { SherpaOnnxBindings.voiceActivityDetectorPop?.call(ptr); } - /// Remove all queued speech segments. void clear() { if (SherpaOnnxBindings.voiceActivityDetectorClear == null) { throw Exception("Please initialize sherpa-onnx first"); @@ -468,7 +285,6 @@ class VoiceActivityDetector { SherpaOnnxBindings.voiceActivityDetectorClear?.call(ptr); } - /// Return the front queued speech segment. SpeechSegment front() { if (SherpaOnnxBindings.voiceActivityDetectorFront == null) { throw Exception("Please initialize sherpa-onnx first"); @@ -505,7 +321,6 @@ class VoiceActivityDetector { SherpaOnnxBindings.voiceActivityDetectorReset?.call(ptr); } - /// Flush trailing buffered speech into the output queue. void flush() { if (SherpaOnnxBindings.voiceActivityDetectorFlush == null) { throw Exception("Please initialize sherpa-onnx first"); diff --git a/flutter/sherpa_onnx/lib/src/vad_config.dart b/flutter/sherpa_onnx/lib/src/vad_config.dart new file mode 100644 index 0000000000..92f00a6500 --- /dev/null +++ b/flutter/sherpa_onnx/lib/src/vad_config.dart @@ -0,0 +1,146 @@ +// Copyright (c) 2024 Xiaomi Corporation +// Shared config/data classes for VAD — no FFI, works on all platforms. +import 'dart:typed_data'; + +/// Silero VAD model configuration. +class SileroVadModelConfig { + const SileroVadModelConfig( + {this.model = '', + this.threshold = 0.5, + this.minSilenceDuration = 0.5, + this.minSpeechDuration = 0.25, + this.windowSize = 512, + this.maxSpeechDuration = 5.0}); + + factory SileroVadModelConfig.fromJson(Map json) { + return SileroVadModelConfig( + model: json['model'] as String? ?? '', + threshold: (json['threshold'] as num?)?.toDouble() ?? 0.5, + minSilenceDuration: + (json['minSilenceDuration'] as num?)?.toDouble() ?? 0.5, + minSpeechDuration: + (json['minSpeechDuration'] as num?)?.toDouble() ?? 0.25, + windowSize: json['windowSize'] as int? ?? 512, + maxSpeechDuration: (json['maxSpeechDuration'] as num?)?.toDouble() ?? 5.0, + ); + } + + @override + String toString() { + return 'SileroVadModelConfig(model: $model, threshold: $threshold, minSilenceDuration: $minSilenceDuration, minSpeechDuration: $minSpeechDuration, windowSize: $windowSize, maxSpeechDuration: $maxSpeechDuration)'; + } + + Map toJson() => { + 'model': model, + 'threshold': threshold, + 'minSilenceDuration': minSilenceDuration, + 'minSpeechDuration': minSpeechDuration, + 'windowSize': windowSize, + 'maxSpeechDuration': maxSpeechDuration, + }; + + final String model; + final double threshold; + final double minSilenceDuration; + final double minSpeechDuration; + final int windowSize; + final double maxSpeechDuration; +} + +/// Ten VAD model configuration. +class TenVadModelConfig { + const TenVadModelConfig( + {this.model = '', + this.threshold = 0.5, + this.minSilenceDuration = 0.5, + this.minSpeechDuration = 0.25, + this.windowSize = 256, + this.maxSpeechDuration = 5.0}); + + factory TenVadModelConfig.fromJson(Map json) { + return TenVadModelConfig( + model: json['model'] as String? ?? '', + threshold: (json['threshold'] as num?)?.toDouble() ?? 0.5, + minSilenceDuration: + (json['minSilenceDuration'] as num?)?.toDouble() ?? 0.5, + minSpeechDuration: + (json['minSpeechDuration'] as num?)?.toDouble() ?? 0.25, + windowSize: json['windowSize'] as int? ?? 256, + maxSpeechDuration: (json['maxSpeechDuration'] as num?)?.toDouble() ?? 5.0, + ); + } + + @override + String toString() { + return 'TenVadModelConfig(model: $model, threshold: $threshold, minSilenceDuration: $minSilenceDuration, minSpeechDuration: $minSpeechDuration, windowSize: $windowSize, maxSpeechDuration: $maxSpeechDuration)'; + } + + Map toJson() => { + 'model': model, + 'threshold': threshold, + 'minSilenceDuration': minSilenceDuration, + 'minSpeechDuration': minSpeechDuration, + 'windowSize': windowSize, + 'maxSpeechDuration': maxSpeechDuration, + }; + + final String model; + final double threshold; + final double minSilenceDuration; + final double minSpeechDuration; + final int windowSize; + final double maxSpeechDuration; +} + +/// Top-level VAD model configuration. +class VadModelConfig { + VadModelConfig({ + this.sileroVad = const SileroVadModelConfig(), + this.sampleRate = 16000, + this.numThreads = 1, + this.provider = 'cpu', + this.debug = true, + this.tenVad = const TenVadModelConfig(), + }); + + final SileroVadModelConfig sileroVad; + final TenVadModelConfig tenVad; + final int sampleRate; + final int numThreads; + final String provider; + final bool debug; + + factory VadModelConfig.fromJson(Map json) { + return VadModelConfig( + sileroVad: SileroVadModelConfig.fromJson( + json['sileroVad'] as Map? ?? const {}), + tenVad: TenVadModelConfig.fromJson( + json['tenVad'] as Map? ?? const {}), + sampleRate: json['sampleRate'] as int? ?? 16000, + numThreads: json['numThreads'] as int? ?? 1, + provider: json['provider'] as String? ?? 'cpu', + debug: json['debug'] as bool? ?? true, + ); + } + + Map toJson() => { + 'sileroVad': sileroVad.toJson(), + 'tenVad': tenVad.toJson(), + 'sampleRate': sampleRate, + 'numThreads': numThreads, + 'provider': provider, + 'debug': debug, + }; + + @override + String toString() { + return 'VadModelConfig(sileroVad: $sileroVad, tenVad: $tenVad, sampleRate: $sampleRate, numThreads: $numThreads, provider: $provider, debug: $debug)'; + } +} + +/// One detected speech segment emitted by VoiceActivityDetector. +class SpeechSegment { + SpeechSegment({required this.samples, required this.start}); + final Float32List samples; + final int start; +} diff --git a/flutter/sherpa_onnx/lib/src/wave_reader.dart b/flutter/sherpa_onnx/lib/src/wave_reader.dart index 58c13cfa31..38b7358f82 100644 --- a/flutter/sherpa_onnx/lib/src/wave_reader.dart +++ b/flutter/sherpa_onnx/lib/src/wave_reader.dart @@ -4,18 +4,9 @@ import 'dart:typed_data'; import 'package:ffi/ffi.dart'; import './sherpa_onnx_bindings.dart'; +import './wave_reader_config.dart'; -/// Audio samples loaded from a WAV file. -/// -/// Samples are normalized to the range `[-1, 1]` and are stored as mono -/// `Float32List` PCM data. -class WaveData { - WaveData({required this.samples, required this.sampleRate}); - - /// normalized to [-1, 1] - Float32List samples; - int sampleRate; -} +export './wave_reader_config.dart'; /// Read a WAV file from disk. /// diff --git a/flutter/sherpa_onnx/lib/src/wave_reader_config.dart b/flutter/sherpa_onnx/lib/src/wave_reader_config.dart new file mode 100644 index 0000000000..fdf475f629 --- /dev/null +++ b/flutter/sherpa_onnx/lib/src/wave_reader_config.dart @@ -0,0 +1,15 @@ +// Copyright (c) 2024 Xiaomi Corporation +// Shared data classes for WAV reading -- no FFI, works on all platforms. +import 'dart:typed_data'; + +/// Audio samples loaded from a WAV file. +/// +/// Samples are normalized to the range `[-1, 1]` and are stored as mono +/// `Float32List` PCM data. +class WaveData { + WaveData({required this.samples, required this.sampleRate}); + + /// normalized to [-1, 1] + Float32List samples; + int sampleRate; +} diff --git a/flutter/sherpa_onnx/lib/src/web/audio_tagging.dart b/flutter/sherpa_onnx/lib/src/web/audio_tagging.dart new file mode 100644 index 0000000000..86cffea2fd --- /dev/null +++ b/flutter/sherpa_onnx/lib/src/web/audio_tagging.dart @@ -0,0 +1,28 @@ +// Copyright (c) 2026 Xiaomi Corporation +// Web stub for audio tagging -- not yet implemented. +// TODO: implement using dart:js_interop and the corresponding JS wrapper. + +import '../audio_tagging_config.dart'; +import 'offline_stream.dart'; + +export '../audio_tagging_config.dart'; + +/// Offline audio tagger. +class AudioTagging { + AudioTagging.fromPtr({required this.ptr, required this.config}); + AudioTagging._({required this.ptr, required this.config}); + + factory AudioTagging({required AudioTaggingConfig config}) { + throw UnsupportedError('AudioTagging is not yet supported on web'); + } + + void free() {} + OfflineStream createStream() => + throw UnsupportedError('AudioTagging is not yet supported on web'); + List compute( + {required OfflineStream stream, required int topK}) => + []; + + dynamic ptr; + final AudioTaggingConfig config; +} diff --git a/flutter/sherpa_onnx/lib/src/web/init.dart b/flutter/sherpa_onnx/lib/src/web/init.dart new file mode 100644 index 0000000000..0268ec0457 --- /dev/null +++ b/flutter/sherpa_onnx/lib/src/web/init.dart @@ -0,0 +1,28 @@ +// Web platform initialization using dart:js_interop. +// The Module is set as a global JS variable by sherpa_onnx_web plugin. +import 'dart:js_interop'; +import 'dart:js_interop_unsafe'; + +JSObject? _module; + +/// Get the Emscripten Module instance. +/// Must be called after SherpaOnnxWeb.loadWasm() has completed. +JSObject getModule() { + if (_module != null) return _module!; + + // Try to get the global Module set by sherpa_onnx_web. + final m = globalContext.getProperty('Module'.toJS); + if (m != null && m.isA()) { + _module = m as JSObject; + return _module!; + } + + throw StateError( + 'WASM module not loaded. Call SherpaOnnxWeb.loadWasm() first.', + ); +} + +// No-op for web — initialization is handled by SherpaOnnxWeb. +void initNativeBindings(String? path) { + // Not used on web. +} diff --git a/flutter/sherpa_onnx/lib/src/web/keyword_spotter.dart b/flutter/sherpa_onnx/lib/src/web/keyword_spotter.dart new file mode 100644 index 0000000000..680ceb3f55 --- /dev/null +++ b/flutter/sherpa_onnx/lib/src/web/keyword_spotter.dart @@ -0,0 +1,29 @@ +// Copyright (c) 2026 Xiaomi Corporation +// Web stub for keyword spotter -- not yet implemented. +// TODO: implement using dart:js_interop and the corresponding JS wrapper. + +import '../keyword_spotter_config.dart'; +import 'online_stream.dart'; + +export '../keyword_spotter_config.dart'; + +/// Streaming keyword spotter. +class KeywordSpotter { + KeywordSpotter.fromPtr({required this.ptr, required this.config}); + KeywordSpotter._({required this.ptr, required this.config}); + + factory KeywordSpotter(KeywordSpotterConfig config) { + throw UnsupportedError('KeywordSpotter is not yet supported on web'); + } + + void free() {} + OnlineStream createStream({String keywords = ''}) => + throw UnsupportedError('KeywordSpotter is not yet supported on web'); + bool isReady(OnlineStream stream) => false; + KeywordResult getResult(OnlineStream stream) => KeywordResult(keyword: ''); + void decode(OnlineStream stream) {} + void reset(OnlineStream stream) {} + + dynamic ptr; + KeywordSpotterConfig config; +} diff --git a/flutter/sherpa_onnx/lib/src/web/offline_punctuation.dart b/flutter/sherpa_onnx/lib/src/web/offline_punctuation.dart new file mode 100644 index 0000000000..95880fae96 --- /dev/null +++ b/flutter/sherpa_onnx/lib/src/web/offline_punctuation.dart @@ -0,0 +1,23 @@ +// Copyright (c) 2026 Xiaomi Corporation +// Web stub for offline punctuation -- not yet implemented. +// TODO: implement using dart:js_interop and the corresponding JS wrapper. + +import '../offline_punctuation_config.dart'; + +export '../offline_punctuation_config.dart'; + +/// Offline punctuation restorer. +class OfflinePunctuation { + OfflinePunctuation.fromPtr({required this.ptr, required this.config}); + OfflinePunctuation._({required this.ptr, required this.config}); + + factory OfflinePunctuation({required OfflinePunctuationConfig config}) { + throw UnsupportedError('OfflinePunctuation is not yet supported on web'); + } + + void free() {} + String addPunct(String text) => ''; + + dynamic ptr; + final OfflinePunctuationConfig config; +} diff --git a/flutter/sherpa_onnx/lib/src/web/offline_recognizer.dart b/flutter/sherpa_onnx/lib/src/web/offline_recognizer.dart new file mode 100644 index 0000000000..bd13b1c40d --- /dev/null +++ b/flutter/sherpa_onnx/lib/src/web/offline_recognizer.dart @@ -0,0 +1,40 @@ +// Copyright (c) 2026 Xiaomi Corporation +// Web stub for offline recognizer -- not yet implemented. +// TODO: implement using dart:js_interop and the corresponding JS wrapper. + +import '../offline_recognizer_config.dart'; +import 'offline_stream.dart'; + +export '../offline_recognizer_config.dart'; + +/// Offline speech recognizer. +/// +/// Create one from an [OfflineRecognizerConfig], then create an +/// [OfflineStream], feed waveform samples, call [decode], and fetch the final +/// hypothesis with [getResult]. +class OfflineRecognizer { + OfflineRecognizer.fromPtr({required this.ptr, required this.config}); + OfflineRecognizer._({required this.ptr, required this.config}); + + factory OfflineRecognizer(OfflineRecognizerConfig config) { + throw UnsupportedError('OfflineRecognizer is not yet supported on web'); + } + + void free() {} + void setConfig(OfflineRecognizerConfig config) {} + OfflineStream createStream() => + throw UnsupportedError('OfflineRecognizer is not yet supported on web'); + void decode(OfflineStream stream) {} + OfflineRecognizerResult getResult(OfflineStream stream) => + OfflineRecognizerResult( + text: '', + tokens: [], + timestamps: [], + lang: '', + emotion: '', + event: '', + ); + + dynamic ptr; + OfflineRecognizerConfig config; +} diff --git a/flutter/sherpa_onnx/lib/src/web/offline_speaker_diarization.dart b/flutter/sherpa_onnx/lib/src/web/offline_speaker_diarization.dart new file mode 100644 index 0000000000..43338fc213 --- /dev/null +++ b/flutter/sherpa_onnx/lib/src/web/offline_speaker_diarization.dart @@ -0,0 +1,36 @@ +// Copyright (c) 2026 Xiaomi Corporation +// Web stub for offline speaker diarization -- not yet implemented. +// TODO: implement using dart:js_interop and the corresponding JS wrapper. + +import 'dart:typed_data'; + +import '../offline_speaker_diarization_config.dart'; + +export '../offline_speaker_diarization_config.dart'; + +/// Offline speaker diarizer. +class OfflineSpeakerDiarization { + OfflineSpeakerDiarization.fromPtr( + {required this.ptr, required this.config, required this.sampleRate}); + OfflineSpeakerDiarization._( + {required this.ptr, required this.config, required this.sampleRate}); + + factory OfflineSpeakerDiarization(OfflineSpeakerDiarizationConfig config) { + throw UnsupportedError( + 'OfflineSpeakerDiarization is not yet supported on web'); + } + + void free() {} + List process( + {required Float32List samples}) => + []; + List processWithCallback({ + required Float32List samples, + required int Function(int numProcessedChunks, int numTotalChunks) callback, + }) => + []; + + dynamic ptr; + OfflineSpeakerDiarizationConfig config; + final int sampleRate; +} diff --git a/flutter/sherpa_onnx/lib/src/web/offline_speech_denoiser.dart b/flutter/sherpa_onnx/lib/src/web/offline_speech_denoiser.dart new file mode 100644 index 0000000000..83fc405f26 --- /dev/null +++ b/flutter/sherpa_onnx/lib/src/web/offline_speech_denoiser.dart @@ -0,0 +1,29 @@ +// Copyright (c) 2026 Xiaomi Corporation +// Web stub for offline speech denoiser -- not yet implemented. +// TODO: implement using dart:js_interop and the corresponding JS wrapper. + +import 'dart:typed_data'; + +import '../offline_speech_denoiser_config.dart'; + +export '../offline_speech_denoiser_config.dart'; + +/// Offline speech denoiser. +class OfflineSpeechDenoiser { + OfflineSpeechDenoiser.fromPtr({required this.ptr, required this.config}); + OfflineSpeechDenoiser._({required this.ptr, required this.config}); + + factory OfflineSpeechDenoiser(OfflineSpeechDenoiserConfig config) { + throw UnsupportedError( + 'OfflineSpeechDenoiser is not yet supported on web'); + } + + void free() {} + DenoisedAudio run({required Float32List samples, required int sampleRate}) => + DenoisedAudio(samples: Float32List(0), sampleRate: 0); + + int get sampleRate => 0; + + dynamic ptr; + OfflineSpeechDenoiserConfig config; +} diff --git a/flutter/sherpa_onnx/lib/src/web/offline_stream.dart b/flutter/sherpa_onnx/lib/src/web/offline_stream.dart new file mode 100644 index 0000000000..9edd46e62b --- /dev/null +++ b/flutter/sherpa_onnx/lib/src/web/offline_stream.dart @@ -0,0 +1,18 @@ +// Copyright (c) 2026 Xiaomi Corporation +// Web stub for offline stream -- not yet implemented. +// TODO: implement using dart:js_interop and the corresponding JS wrapper. + +import 'dart:typed_data'; + +/// Input stream for offline APIs such as offline ASR, audio tagging, and +/// spoken language identification. +class OfflineStream { + OfflineStream({required this.ptr}); + + void free() {} + void acceptWaveform( + {required Float32List samples, required int sampleRate}) {} + void setOption({required String key, required String value}) {} + + dynamic ptr; +} diff --git a/flutter/sherpa_onnx/lib/src/web/online_punctuation.dart b/flutter/sherpa_onnx/lib/src/web/online_punctuation.dart new file mode 100644 index 0000000000..cf4a0b7a21 --- /dev/null +++ b/flutter/sherpa_onnx/lib/src/web/online_punctuation.dart @@ -0,0 +1,23 @@ +// Copyright (c) 2026 Xiaomi Corporation +// Web stub for online punctuation -- not yet implemented. +// TODO: implement using dart:js_interop and the corresponding JS wrapper. + +import '../online_punctuation_config.dart'; + +export '../online_punctuation_config.dart'; + +/// Online punctuation restorer. +class OnlinePunctuation { + OnlinePunctuation.fromPtr({required this.ptr, required this.config}); + OnlinePunctuation._({required this.ptr, required this.config}); + + factory OnlinePunctuation({required OnlinePunctuationConfig config}) { + throw UnsupportedError('OnlinePunctuation is not yet supported on web'); + } + + void free() {} + String addPunct(String text) => ''; + + dynamic ptr; + final OnlinePunctuationConfig config; +} diff --git a/flutter/sherpa_onnx/lib/src/web/online_recognizer.dart b/flutter/sherpa_onnx/lib/src/web/online_recognizer.dart new file mode 100644 index 0000000000..6ff4b32609 --- /dev/null +++ b/flutter/sherpa_onnx/lib/src/web/online_recognizer.dart @@ -0,0 +1,34 @@ +// Copyright (c) 2026 Xiaomi Corporation +// Web stub for online recognizer -- not yet implemented. +// TODO: implement using dart:js_interop and the corresponding JS wrapper. + +import '../online_recognizer_config.dart'; +import 'online_stream.dart'; + +export '../online_recognizer_config.dart'; + +/// Streaming speech recognizer. +/// +/// Create one from an [OnlineRecognizerConfig], then feed chunks to an +/// [OnlineStream] and call [decode] while [isReady] is true. +class OnlineRecognizer { + OnlineRecognizer.fromPtr({required this.ptr, required this.config}); + OnlineRecognizer._({required this.ptr, required this.config}); + + factory OnlineRecognizer(OnlineRecognizerConfig config) { + throw UnsupportedError('OnlineRecognizer is not yet supported on web'); + } + + void free() {} + OnlineStream createStream({String hotwords = ''}) => + throw UnsupportedError('OnlineRecognizer is not yet supported on web'); + bool isReady(OnlineStream stream) => false; + OnlineRecognizerResult getResult(OnlineStream stream) => + OnlineRecognizerResult(text: '', tokens: [], timestamps: []); + void reset(OnlineStream stream) {} + void decode(OnlineStream stream) {} + bool isEndpoint(OnlineStream stream) => false; + + dynamic ptr; + OnlineRecognizerConfig config; +} diff --git a/flutter/sherpa_onnx/lib/src/web/online_speech_denoiser.dart b/flutter/sherpa_onnx/lib/src/web/online_speech_denoiser.dart new file mode 100644 index 0000000000..970b5efb04 --- /dev/null +++ b/flutter/sherpa_onnx/lib/src/web/online_speech_denoiser.dart @@ -0,0 +1,34 @@ +// Copyright (c) 2026 Xiaomi Corporation +// Web stub for online speech denoiser -- not yet implemented. +// TODO: implement using dart:js_interop and the corresponding JS wrapper. + +import 'dart:typed_data'; + +import '../offline_speech_denoiser_config.dart'; +import '../online_speech_denoiser_config.dart'; + +export '../online_speech_denoiser_config.dart'; + +/// Streaming speech denoiser. +class OnlineSpeechDenoiser { + OnlineSpeechDenoiser.fromPtr({required this.ptr, required this.config}); + OnlineSpeechDenoiser._({required this.ptr, required this.config}); + + factory OnlineSpeechDenoiser(OnlineSpeechDenoiserConfig config) { + throw UnsupportedError( + 'OnlineSpeechDenoiser is not yet supported on web'); + } + + void free() {} + DenoisedAudio run({required Float32List samples, required int sampleRate}) => + DenoisedAudio(samples: Float32List(0), sampleRate: 0); + DenoisedAudio flush() => + DenoisedAudio(samples: Float32List(0), sampleRate: 0); + void reset() {} + + int get sampleRate => 0; + int get frameShiftInSamples => 0; + + dynamic ptr; + OnlineSpeechDenoiserConfig config; +} diff --git a/flutter/sherpa_onnx/lib/src/web/online_stream.dart b/flutter/sherpa_onnx/lib/src/web/online_stream.dart new file mode 100644 index 0000000000..cfed509ed5 --- /dev/null +++ b/flutter/sherpa_onnx/lib/src/web/online_stream.dart @@ -0,0 +1,18 @@ +// Copyright (c) 2026 Xiaomi Corporation +// Web stub for online stream -- not yet implemented. +// TODO: implement using dart:js_interop and the corresponding JS wrapper. + +import 'dart:typed_data'; + +/// Input stream for streaming APIs such as online ASR and keyword spotting. +class OnlineStream { + OnlineStream({required this.ptr}); + + void free() {} + void acceptWaveform( + {required Float32List samples, required int sampleRate}) {} + void inputFinished() {} + void setOption({required String key, required String value}) {} + + dynamic ptr; +} diff --git a/flutter/sherpa_onnx/lib/src/web/speaker_identification.dart b/flutter/sherpa_onnx/lib/src/web/speaker_identification.dart new file mode 100644 index 0000000000..8532e9c2e7 --- /dev/null +++ b/flutter/sherpa_onnx/lib/src/web/speaker_identification.dart @@ -0,0 +1,71 @@ +// Copyright (c) 2026 Xiaomi Corporation +// Web stub for speaker identification -- not yet implemented. +// TODO: implement using dart:js_interop and the corresponding JS wrapper. + +import 'dart:typed_data'; + +import 'online_stream.dart'; +import '../speaker_identification_config.dart'; + +export '../speaker_identification_config.dart'; + +/// Speaker embedding extractor. +/// +/// Feed audio through an [OnlineStream], then call [compute] to obtain a fixed +/// dimensional embedding suitable for search or verification. +class SpeakerEmbeddingExtractor { + SpeakerEmbeddingExtractor.fromPtr({required this.ptr, required this.dim}); + SpeakerEmbeddingExtractor._({required this.ptr, required this.dim}); + + factory SpeakerEmbeddingExtractor( + {required SpeakerEmbeddingExtractorConfig config}) { + throw UnsupportedError( + 'SpeakerEmbeddingExtractor is not yet supported on web'); + } + + void free() {} + OnlineStream createStream() => + throw UnsupportedError( + 'SpeakerEmbeddingExtractor is not yet supported on web'); + bool isReady(OnlineStream stream) => false; + Float32List compute(OnlineStream stream) => Float32List(0); + + dynamic ptr; + final int dim; +} + +/// In-memory store of named speaker embeddings. +/// +/// Use this class to add reference embeddings, search for the best matching +/// speaker, and verify whether a candidate embedding belongs to a known +/// identity. +class SpeakerEmbeddingManager { + SpeakerEmbeddingManager.fromPtr({required this.ptr, required this.dim}); + SpeakerEmbeddingManager._({required this.ptr, required this.dim}); + + factory SpeakerEmbeddingManager(int dim) { + throw UnsupportedError( + 'SpeakerEmbeddingManager is not yet supported on web'); + } + + void free() {} + bool add({required String name, required Float32List embedding}) => false; + bool addMulti( + {required String name, required List embeddingList}) => + false; + bool contains(String name) => false; + bool remove(String name) => false; + String search({required Float32List embedding, required double threshold}) => + ''; + bool verify( + {required String name, + required Float32List embedding, + required double threshold}) => + false; + + int get numSpeakers => 0; + List get allSpeakerNames => []; + + dynamic ptr; + final int dim; +} diff --git a/flutter/sherpa_onnx/lib/src/web/spoken_language_identification.dart b/flutter/sherpa_onnx/lib/src/web/spoken_language_identification.dart new file mode 100644 index 0000000000..d467833033 --- /dev/null +++ b/flutter/sherpa_onnx/lib/src/web/spoken_language_identification.dart @@ -0,0 +1,31 @@ +// Copyright (c) 2026 Xiaomi Corporation +// Web stub for spoken language identification -- not yet implemented. +// TODO: implement using dart:js_interop and the corresponding JS wrapper. + +import 'offline_stream.dart'; +import '../spoken_language_identification_config.dart'; + +export '../spoken_language_identification_config.dart'; + +/// Spoken language identifier. +class SpokenLanguageIdentification { + SpokenLanguageIdentification.fromPtr( + {required this.ptr, required this.config}); + SpokenLanguageIdentification._({required this.ptr, required this.config}); + + factory SpokenLanguageIdentification( + SpokenLanguageIdentificationConfig config) { + throw UnsupportedError( + 'SpokenLanguageIdentification is not yet supported on web'); + } + + void free() {} + OfflineStream createStream() => + throw UnsupportedError( + 'SpokenLanguageIdentification is not yet supported on web'); + SpokenLanguageIdentificationResult compute(OfflineStream stream) => + const SpokenLanguageIdentificationResult(lang: ''); + + dynamic ptr; + SpokenLanguageIdentificationConfig config; +} diff --git a/flutter/sherpa_onnx/lib/src/web/tts.dart b/flutter/sherpa_onnx/lib/src/web/tts.dart new file mode 100644 index 0000000000..44fedb1d44 --- /dev/null +++ b/flutter/sherpa_onnx/lib/src/web/tts.dart @@ -0,0 +1,230 @@ +// Copyright (c) 2026 Xiaomi Corporation +// Web implementation of OfflineTTS using dart:js_interop. +import 'dart:js_interop'; +import 'dart:js_interop_unsafe'; +import 'dart:typed_data'; + +import '../tts_config.dart'; +import 'init.dart'; + +export '../tts_config.dart'; + +/// Offline text-to-speech engine (web implementation). +class OfflineTts { + OfflineTts.fromPtr({required this.ptr, required this.config}); + OfflineTts._({required this.ptr, required this.config}); + + /// Create a TTS instance using the JS wrapper. + factory OfflineTts(OfflineTtsConfig config) { + final m = getModule(); + // createOfflineTts is a global function defined by sherpa-onnx-tts.js. + final createFn = + globalContext.getProperty('createOfflineTts'.toJS) as JSFunction?; + if (createFn == null) { + throw StateError('createOfflineTts not found. Is sherpa-onnx-tts.js loaded?'); + } + + // Convert Dart config to JS object. + final jsConfig = _configToJs(config); + final handle = createFn.callAsFunction(null, m, jsConfig); + + if (handle == null) { + throw Exception('Failed to create OfflineTts on web'); + } + + return OfflineTts._(ptr: handle, config: config); + } + + void free() { + if (_freed) return; + final m = getModule(); + final destroyFn = m.getProperty('_SherpaOnnxDestroyOfflineTts'.toJS) as JSFunction?; + destroyFn?.callAsFunction(null, ptr); + _freed = true; + } + + GeneratedAudio generate({required String text, int sid = 0, double speed = 1.0}) { + return generateWithConfig( + text: text, + config: OfflineTtsGenerationConfig(sid: sid, speed: speed), + ); + } + + GeneratedAudio generateWithCallback({ + required String text, + int sid = 0, + double speed = 1.0, + required int Function(Float32List samples) callback, + }) { + return generateWithConfig( + text: text, + config: OfflineTtsGenerationConfig(sid: sid, speed: speed), + onProgress: (samples, progress) => callback(samples), + ); + } + + GeneratedAudio generateWithConfig({ + required String text, + required OfflineTtsGenerationConfig config, + int Function(Float32List samples, double progress)? onProgress, + }) { + final m = getModule(); + final handle = ptr as JSObject; + + // Build the genConfig JS object. + final genConfig = JSObject(); + genConfig['silenceScale'] = config.silenceScale.toJS; + genConfig['speed'] = config.speed.toJS; + genConfig['sid'] = config.sid.toJS; + genConfig['numSteps'] = config.numSteps.toJS; + + // Reference audio for voice cloning (e.g. Pocket TTS). + if (config.referenceAudio != null && config.referenceAudio!.isNotEmpty) { + genConfig['referenceAudio'] = config.referenceAudio; + genConfig['referenceSampleRate'] = config.referenceSampleRate.toJS; + genConfig['referenceText'] = config.referenceText.toJS; + } + + // Extra model-specific attributes. + if (config.extra.isNotEmpty) { + final extraObj = JSObject(); + for (final entry in config.extra.entries) { + extraObj[entry.key] = (entry.value is String + ? (entry.value as String).toJS + : entry.value is int + ? (entry.value as int).toJS + : (entry.value as double).toJS) as JSAny; + } + genConfig['extra'] = extraObj; + } + + if (onProgress != null) { + // Create a JS callback that calls the Dart callback. + genConfig['callback'] = (JSAny samplesPtr, JSAny n, JSAny progress, JSAny arg) { + // The JS wrapper passes a Float32Array as samplesPtr. + final samples = (samplesPtr as JSFloat32Array).toDart; + return onProgress(samples, (progress as JSNumber).toDartDouble).toJS; + }.toJS; + } + + // Call the JS wrapper's generateWithConfig method. + final generateFn = handle.getProperty('generateWithConfig'.toJS) as JSFunction?; + if (generateFn == null) { + throw StateError('generateWithConfig not found on OfflineTts instance'); + } + + final result = generateFn.callAsFunction(handle, text.toJS, genConfig) as JSObject; + final samples = (result.getProperty('samples'.toJS) as JSFloat32Array).toDart; + final sampleRate = (result.getProperty('sampleRate'.toJS) as JSNumber).toDartInt; + + return GeneratedAudio(samples: samples, sampleRate: sampleRate); + } + + int get sampleRate { + final handle = ptr as JSObject; + final val = handle.getProperty('sampleRate'.toJS); + return val is JSNumber ? val.toDartInt : 0; + } + + int get numSpeakers { + final handle = ptr as JSObject; + final val = handle.getProperty('numSpeakers'.toJS); + return val is JSNumber ? val.toDartInt : 0; + } + + dynamic ptr; + OfflineTtsConfig config; + bool _freed = false; +} + +/// Convert OfflineTtsConfig to a JS object for the JS wrapper. +JSObject _configToJs(OfflineTtsConfig config) { + final jsConfig = JSObject(); + + // model config + final model = JSObject(); + final vits = JSObject(); + vits['model'] = config.model.vits.model.toJS; + vits['lexicon'] = config.model.vits.lexicon.toJS; + vits['tokens'] = config.model.vits.tokens.toJS; + vits['dataDir'] = config.model.vits.dataDir.toJS; + vits['noiseScale'] = config.model.vits.noiseScale.toJS; + vits['noiseScaleW'] = config.model.vits.noiseScaleW.toJS; + vits['lengthScale'] = config.model.vits.lengthScale.toJS; + model['vits'] = vits; + + final matcha = JSObject(); + matcha['acousticModel'] = config.model.matcha.acousticModel.toJS; + matcha['vocoder'] = config.model.matcha.vocoder.toJS; + matcha['lexicon'] = config.model.matcha.lexicon.toJS; + matcha['tokens'] = config.model.matcha.tokens.toJS; + matcha['dataDir'] = config.model.matcha.dataDir.toJS; + matcha['noiseScale'] = config.model.matcha.noiseScale.toJS; + matcha['lengthScale'] = config.model.matcha.lengthScale.toJS; + model['matcha'] = matcha; + + final kokoro = JSObject(); + kokoro['model'] = config.model.kokoro.model.toJS; + kokoro['voices'] = config.model.kokoro.voices.toJS; + kokoro['tokens'] = config.model.kokoro.tokens.toJS; + kokoro['dataDir'] = config.model.kokoro.dataDir.toJS; + kokoro['lengthScale'] = config.model.kokoro.lengthScale.toJS; + kokoro['lexicon'] = config.model.kokoro.lexicon.toJS; + kokoro['lang'] = config.model.kokoro.lang.toJS; + model['kokoro'] = kokoro; + + final kitten = JSObject(); + kitten['model'] = config.model.kitten.model.toJS; + kitten['voices'] = config.model.kitten.voices.toJS; + kitten['tokens'] = config.model.kitten.tokens.toJS; + kitten['dataDir'] = config.model.kitten.dataDir.toJS; + kitten['lengthScale'] = config.model.kitten.lengthScale.toJS; + model['kitten'] = kitten; + + final zipvoice = JSObject(); + zipvoice['tokens'] = config.model.zipvoice.tokens.toJS; + zipvoice['encoder'] = config.model.zipvoice.encoder.toJS; + zipvoice['decoder'] = config.model.zipvoice.decoder.toJS; + zipvoice['vocoder'] = config.model.zipvoice.vocoder.toJS; + zipvoice['dataDir'] = config.model.zipvoice.dataDir.toJS; + zipvoice['lexicon'] = config.model.zipvoice.lexicon.toJS; + zipvoice['featScale'] = config.model.zipvoice.featScale.toJS; + zipvoice['tShift'] = config.model.zipvoice.tShift.toJS; + zipvoice['targetRms'] = config.model.zipvoice.targetRms.toJS; + zipvoice['guidanceScale'] = config.model.zipvoice.guidanceScale.toJS; + model['zipvoice'] = zipvoice; + + final pocket = JSObject(); + pocket['lmFlow'] = config.model.pocket.lmFlow.toJS; + pocket['lmMain'] = config.model.pocket.lmMain.toJS; + pocket['encoder'] = config.model.pocket.encoder.toJS; + pocket['decoder'] = config.model.pocket.decoder.toJS; + pocket['textConditioner'] = config.model.pocket.textConditioner.toJS; + pocket['vocabJson'] = config.model.pocket.vocabJson.toJS; + pocket['tokenScoresJson'] = config.model.pocket.tokenScoresJson.toJS; + pocket['voiceEmbeddingCacheCapacity'] = + config.model.pocket.voiceEmbeddingCacheCapacity.toJS; + model['pocket'] = pocket; + + final supertonic = JSObject(); + supertonic['durationPredictor'] = config.model.supertonic.durationPredictor.toJS; + supertonic['textEncoder'] = config.model.supertonic.textEncoder.toJS; + supertonic['vectorEstimator'] = config.model.supertonic.vectorEstimator.toJS; + supertonic['vocoder'] = config.model.supertonic.vocoder.toJS; + supertonic['ttsJson'] = config.model.supertonic.ttsJson.toJS; + supertonic['unicodeIndexer'] = config.model.supertonic.unicodeIndexer.toJS; + supertonic['voiceStyle'] = config.model.supertonic.voiceStyle.toJS; + model['supertonic'] = supertonic; + + model['numThreads'] = config.model.numThreads.toJS; + model['debug'] = config.model.debug.toJS; + model['provider'] = config.model.provider.toJS; + + jsConfig['offlineTtsModelConfig'] = model; + jsConfig['ruleFsts'] = config.ruleFsts.toJS; + jsConfig['ruleFars'] = config.ruleFars.toJS; + jsConfig['maxNumSentences'] = config.maxNumSenetences.toJS; + jsConfig['silenceScale'] = config.silenceScale.toJS; + + return jsConfig; +} diff --git a/flutter/sherpa_onnx/lib/src/web/vad.dart b/flutter/sherpa_onnx/lib/src/web/vad.dart new file mode 100644 index 0000000000..b33f12a893 --- /dev/null +++ b/flutter/sherpa_onnx/lib/src/web/vad.dart @@ -0,0 +1,52 @@ +// Copyright (c) 2026 Xiaomi Corporation +// Web implementation of VAD using dart:js_interop. +import 'dart:typed_data'; + +import '../vad_config.dart'; + +export '../vad_config.dart'; + +/// Circular sample buffer used by VAD-related pipelines. +class CircularBuffer { + CircularBuffer.fromPtr({required this.ptr}); + CircularBuffer._({required this.ptr}); + + factory CircularBuffer({required int capacity}) { + throw UnsupportedError('CircularBuffer is not yet supported on web'); + } + + void free() {} + void push(Float32List data) {} + Float32List get({required int startIndex, required int n}) => + Float32List(0); + void pop(int n) {} + void reset() {} + int get size => 0; + int get head => 0; + dynamic ptr; +} + +/// Voice activity detector that emits [SpeechSegment] objects. +class VoiceActivityDetector { + VoiceActivityDetector.fromPtr({required this.ptr, required this.config}); + VoiceActivityDetector._({required this.ptr, required this.config}); + + factory VoiceActivityDetector( + {required VadModelConfig config, required double bufferSizeInSeconds}) { + throw UnsupportedError( + 'VoiceActivityDetector is not yet supported on web'); + } + + void free() {} + void acceptWaveform(Float32List samples) {} + bool isEmpty() => true; + bool isDetected() => false; + void pop() {} + void clear() {} + SpeechSegment front() => SpeechSegment(samples: Float32List(0), start: 0); + void reset() {} + void flush() {} + + dynamic ptr; + VadModelConfig config; +} diff --git a/flutter/sherpa_onnx/lib/src/web/version.dart b/flutter/sherpa_onnx/lib/src/web/version.dart new file mode 100644 index 0000000000..2731bd406a --- /dev/null +++ b/flutter/sherpa_onnx/lib/src/web/version.dart @@ -0,0 +1,49 @@ +// Copyright (c) 2026 Xiaomi Corporation +// Web implementation using dart:js_interop. +import 'dart:js_interop'; +import 'dart:js_interop_unsafe'; +import 'init.dart'; + +// Helper: call a method on the Module object with 0 args. +JSAny? _call0(String method) { + final m = getModule(); + final fn = m.getProperty(method.toJS) as JSFunction?; + if (fn == null) return null; + return fn.callAsFunction(m); +} + +// Helper: call a method on the Module object with 1 arg. +JSAny? _call1(String method, JSAny? arg1) { + final m = getModule(); + final fn = m.getProperty(method.toJS) as JSFunction?; + if (fn == null) return null; + return fn.callAsFunction(m, arg1); +} + +// Helper: convert a WASM pointer to a Dart string. +String _ptrToString(JSAny? ptr) { + if (ptr == null) return ''; + final result = _call1('UTF8ToString', ptr); + if (result == null) return ''; + return (result as JSString).toDart; +} + +/// Return the sherpa-onnx version string compiled into the native library. +String getVersion() { + return _ptrToString(_call0('_SherpaOnnxGetVersionStr')); +} + +/// Return the Git SHA1 of the native library build. +String getGitSha1() { + return _ptrToString(_call0('_SherpaOnnxGetGitSha1')); +} + +/// Return the Git date of the native library build. +String getGitDate() { + return _ptrToString(_call0('_SherpaOnnxGetGitDate')); +} + +/// Return the onnxruntime version string used by the native library. +String getOnnxruntimeVersion() { + return _ptrToString(_call0('_SherpaOnnxGetOnnxruntimeVersionStr')); +} diff --git a/flutter/sherpa_onnx/lib/src/web/wave_reader.dart b/flutter/sherpa_onnx/lib/src/web/wave_reader.dart new file mode 100644 index 0000000000..7a6651c9df --- /dev/null +++ b/flutter/sherpa_onnx/lib/src/web/wave_reader.dart @@ -0,0 +1,16 @@ +// Copyright (c) 2026 Xiaomi Corporation +// Web stub for WAV reading -- not yet implemented. +// TODO: implement using dart:js_interop and the corresponding JS wrapper. + +import 'dart:typed_data'; + +import '../wave_reader_config.dart'; + +export '../wave_reader_config.dart'; + +/// Read a WAV file from disk. +/// +/// Returns an empty [WaveData] object if the file cannot be read or decoded. +WaveData readWave(String filename) { + throw UnsupportedError('readWave is not yet supported on web'); +} diff --git a/flutter/sherpa_onnx/lib/src/web/wave_writer.dart b/flutter/sherpa_onnx/lib/src/web/wave_writer.dart new file mode 100644 index 0000000000..55d672e9d0 --- /dev/null +++ b/flutter/sherpa_onnx/lib/src/web/wave_writer.dart @@ -0,0 +1,15 @@ +// Copyright (c) 2026 Xiaomi Corporation +// Web stub for WAV writing -- not yet implemented. +// TODO: implement using dart:js_interop and the corresponding JS wrapper. + +import 'dart:typed_data'; + +/// Write normalized mono PCM samples to a WAV file. +/// +/// Returns `true` on success and `false` otherwise. +bool writeWave( + {required String filename, + required Float32List samples, + required int sampleRate}) { + throw UnsupportedError('writeWave is not yet supported on web'); +} diff --git a/flutter/sherpa_onnx/pubspec.yaml b/flutter/sherpa_onnx/pubspec.yaml index d5a67766a1..0dcf6d6089 100644 --- a/flutter/sherpa_onnx/pubspec.yaml +++ b/flutter/sherpa_onnx/pubspec.yaml @@ -62,6 +62,10 @@ dependencies: # sherpa_onnx_ios: # path: ../sherpa_onnx_ios + sherpa_onnx_web: ^1.13.4 + # sherpa_onnx_web: + # path: ../sherpa_onnx_web + dev_dependencies: flutter_lints: ^3.0.0 @@ -82,3 +86,6 @@ flutter: windows: default_package: sherpa_onnx_windows + + web: + default_package: sherpa_onnx_web diff --git a/flutter/sherpa_onnx_ios/ios/.gitignore b/flutter/sherpa_onnx_ios/ios/.gitignore new file mode 100644 index 0000000000..83d77ab9a1 --- /dev/null +++ b/flutter/sherpa_onnx_ios/ios/.gitignore @@ -0,0 +1 @@ +*.xcframework diff --git a/flutter/sherpa_onnx_macos/macos/.gitignore b/flutter/sherpa_onnx_macos/macos/.gitignore new file mode 100644 index 0000000000..83d77ab9a1 --- /dev/null +++ b/flutter/sherpa_onnx_macos/macos/.gitignore @@ -0,0 +1 @@ +*.xcframework diff --git a/flutter/sherpa_onnx_web/.gitignore b/flutter/sherpa_onnx_web/.gitignore new file mode 100644 index 0000000000..a400d36202 --- /dev/null +++ b/flutter/sherpa_onnx_web/.gitignore @@ -0,0 +1,12 @@ +# Built by build-wasm-simd-web.sh +assets/sherpa-onnx-wasm-web.js +assets/sherpa-onnx-wasm-web.wasm + +# Symlinks created by build-flutter-web-wasm.sh +assets/sherpa-onnx-asr.js +assets/sherpa-onnx-kws.js +assets/sherpa-onnx-punctuation.js +assets/sherpa-onnx-speaker-diarization.js +assets/sherpa-onnx-speech-enhancement.js +assets/sherpa-onnx-tts.js +assets/sherpa-onnx-vad.js diff --git a/flutter/sherpa_onnx_web/README.md b/flutter/sherpa_onnx_web/README.md new file mode 100644 index 0000000000..25d04fc69c --- /dev/null +++ b/flutter/sherpa_onnx_web/README.md @@ -0,0 +1,8 @@ +# sherpa_onnx_web + +This is a sub project of [sherpa-onnx](https://github.com/k2-fsa/sherpa-onnx). + +You are not expected to use this package directly. + +Please see the entry point at . + diff --git a/flutter/sherpa_onnx_web/assets/.gitkeep b/flutter/sherpa_onnx_web/assets/.gitkeep new file mode 100644 index 0000000000..e69de29bb2 diff --git a/flutter/sherpa_onnx_web/lib/sherpa_onnx_web.dart b/flutter/sherpa_onnx_web/lib/sherpa_onnx_web.dart new file mode 100644 index 0000000000..78857cbdc6 --- /dev/null +++ b/flutter/sherpa_onnx_web/lib/sherpa_onnx_web.dart @@ -0,0 +1,119 @@ +// Copyright (c) 2026 Xiaomi Corporation +import 'dart:async'; +import 'dart:convert'; +import 'dart:js_interop'; +import 'dart:js_interop_unsafe'; + +import 'package:flutter/foundation.dart' show kDebugMode; +import 'package:flutter/services.dart'; +import 'package:flutter_web_plugins/flutter_web_plugins.dart'; + +void _log(String message) { + if (kDebugMode) { + print('[sherpa_onnx_web] $message'); + } +} + +/// Evaluate JavaScript source code in the global scope. +void _evalJs(String source) { + final eval = globalContext.getProperty('eval'.toJS) as JSFunction; + eval.callAsFunction(null, source.toJS); +} + +/// Load a JS asset and return its source code. +Future _loadAsset(String assetPath) async { + final data = await rootBundle.load(assetPath); + return utf8.decode(data.buffer.asUint8List(data.offsetInBytes, data.lengthInBytes)); +} + +class SherpaOnnxWeb { + static bool _initialized = false; + static Completer? _loadingCompleter; + + static void registerWith(Registrar registrar) {} + + /// Load the sherpa-onnx WASM module and JS wrappers. + /// + /// This must be called once before using any sherpa-onnx API on web. + /// Typically called via `initBindingsAsync()` from the sherpa_onnx package. + /// Safe to call concurrently — only the first call performs initialization. + static Future loadWasm() async { + if (_initialized) return; + + // If another call is already in progress, wait for it. + if (_loadingCompleter != null) { + return _loadingCompleter!.future; + } + + _loadingCompleter = Completer(); + + try { + + const prefix = 'packages/sherpa_onnx_web/assets'; + + // 1. Load all JS assets (async I/O). + _log('Loading JS assets...'); + final glueSource = await _loadAsset('$prefix/sherpa-onnx-wasm-web.js'); + final wrapperSources = await Future.wait([ + _loadAsset('$prefix/sherpa-onnx-asr.js'), + _loadAsset('$prefix/sherpa-onnx-tts.js'), + _loadAsset('$prefix/sherpa-onnx-vad.js'), + _loadAsset('$prefix/sherpa-onnx-kws.js'), + _loadAsset('$prefix/sherpa-onnx-punctuation.js'), + _loadAsset('$prefix/sherpa-onnx-speaker-diarization.js'), + _loadAsset('$prefix/sherpa-onnx-speech-enhancement.js'), + ]); + + // 2. Load WASM binary. + _log('Loading WASM binary...'); + final wasmData = + await rootBundle.load('$prefix/sherpa-onnx-wasm-web.wasm'); + final wasmBytes = wasmData.buffer.asUint8List(wasmData.offsetInBytes, wasmData.lengthInBytes); + + // 3. Evaluate JS glue code (defines SherpaOnnx factory). + _log('Evaluating JS...'); + _evalJs(glueSource); + + // 4. Define `module` for browser compatibility. + // The JS wrappers use `module.exports` (Node.js pattern). + // In the browser, `module` is not defined, so we stub it. + _evalJs('if (typeof module === "undefined") { var module = {}; }'); + + // 5. Evaluate JS wrappers (they use Module._FunctionName). + for (final src in wrapperSources) { + _evalJs(src); + } + + // 6. Call the factory with wasmBinary. + final factory = + globalContext.getProperty('SherpaOnnx'.toJS) as JSFunction?; + if (factory == null) { + throw StateError('SherpaOnnx factory not found on globalThis'); + } + + final moduleConfig = JSObject(); + moduleConfig['wasmBinary'] = wasmBytes.toJS; + var result = factory.callAsFunction(null, moduleConfig); + + // 7. Unwrap promises. + while (result != null && result.isA()) { + result = await (result as JSPromise).toDart as JSObject?; + } + + if (result == null || !result.isA()) { + throw StateError('Failed to instantiate Emscripten Module'); + } + + // 8. Set Module as global so dart:js_interop code can access it. + globalContext['Module'] = result; + _log('WASM module initialized'); + + _initialized = true; + _loadingCompleter!.complete(); + } catch (e) { + _loadingCompleter!.completeError(e); + _loadingCompleter = null; + rethrow; + } + } +} diff --git a/flutter/sherpa_onnx_web/pubspec.yaml b/flutter/sherpa_onnx_web/pubspec.yaml new file mode 100644 index 0000000000..d95c00f96a --- /dev/null +++ b/flutter/sherpa_onnx_web/pubspec.yaml @@ -0,0 +1,42 @@ +name: sherpa_onnx_web +description: > + Web platform support for sherpa_onnx using WebAssembly. + +version: 1.13.4 + +repository: https://github.com/k2-fsa/sherpa-onnx/tree/master/flutter + +issue_tracker: https://github.com/k2-fsa/sherpa-onnx/issues +documentation: https://k2-fsa.github.io/sherpa/onnx/ + +homepage: https://github.com/k2-fsa/sherpa-onnx + +environment: + sdk: ">=3.4.0 <4.0.0" + flutter: ">=2.8.1" + +dependencies: + flutter: + sdk: flutter + flutter_web_plugins: + sdk: flutter + path: ^1.9.0 + web: ^1.1.0 + +flutter: + plugin: + platforms: + web: + pluginClass: SherpaOnnxWeb + fileName: sherpa_onnx_web.dart + + assets: + - assets/sherpa-onnx-wasm-web.js + - assets/sherpa-onnx-wasm-web.wasm + - assets/sherpa-onnx-asr.js + - assets/sherpa-onnx-tts.js + - assets/sherpa-onnx-vad.js + - assets/sherpa-onnx-kws.js + - assets/sherpa-onnx-punctuation.js + - assets/sherpa-onnx-speaker-diarization.js + - assets/sherpa-onnx-speech-enhancement.js diff --git a/sherpa-onnx/c-api/c-api.cc b/sherpa-onnx/c-api/c-api.cc index e776d4f51a..6278cb210d 100644 --- a/sherpa-onnx/c-api/c-api.cc +++ b/sherpa-onnx/c-api/c-api.cc @@ -50,6 +50,8 @@ #include "sherpa-onnx/csrc/offline-speaker-diarization.h" #endif +#define SHERPA_ONNX_OR(x, y) (x ? x : y) + const char *SherpaOnnxGetVersionStr() { return sherpa_onnx::GetVersionStr(); } const char *SherpaOnnxGetGitSha1() { return sherpa_onnx::GetGitSha1(); } const char *SherpaOnnxGetGitDate() { return sherpa_onnx::GetGitDate(); } @@ -57,6 +59,21 @@ const char *SherpaOnnxGetOnnxruntimeVersionStr() { return sherpa_onnx::GetOnnxruntimeVersionStr(); } +// On WASM builds, multi-threading is not supported. Clamp num_threads to 1. +static int32_t GetNumThreads(int32_t num_threads) { + int32_t n = SHERPA_ONNX_OR(num_threads, 1); +#if SHERPA_ONNX_ENABLE_WASM + if (n > 1) { + SHERPA_ONNX_LOGE( + "WASM does not support multi-threading. " + "Changing num_threads from %d to 1.", + n); + n = 1; + } +#endif + return n; +} + struct SherpaOnnxOnlineRecognizer { std::unique_ptr impl; }; @@ -71,8 +88,6 @@ struct SherpaOnnxDisplay { std::unique_ptr impl; }; -#define SHERPA_ONNX_OR(x, y) (x ? x : y) - static sherpa_onnx::OnlineRecognizerConfig GetOnlineRecognizerConfig( const SherpaOnnxOnlineRecognizerConfig *config) { sherpa_onnx::OnlineRecognizerConfig recognizer_config; @@ -112,7 +127,7 @@ static sherpa_onnx::OnlineRecognizerConfig GetOnlineRecognizerConfig( SHERPA_ONNX_OR(config->model_config.t_one_ctc.model, ""); recognizer_config.model_config.num_threads = - SHERPA_ONNX_OR(config->model_config.num_threads, 1); + GetNumThreads(config->model_config.num_threads); recognizer_config.model_config.provider_config.provider = SHERPA_ONNX_OR(config->model_config.provider, "cpu"); @@ -482,7 +497,7 @@ static sherpa_onnx::OfflineRecognizerConfig GetOfflineRecognizerConfig( recognizer_config.model_config.tokens = SHERPA_ONNX_OR(config->model_config.tokens, ""); recognizer_config.model_config.num_threads = - SHERPA_ONNX_OR(config->model_config.num_threads, 1); + GetNumThreads(config->model_config.num_threads); recognizer_config.model_config.debug = config->model_config.debug; recognizer_config.model_config.provider = SHERPA_ONNX_OR(config->model_config.provider, "cpu"); @@ -996,7 +1011,7 @@ static sherpa_onnx::KeywordSpotterConfig GetKeywordSpotterConfig( } spotter_config.model_config.num_threads = - SHERPA_ONNX_OR(config->model_config.num_threads, 1); + GetNumThreads(config->model_config.num_threads); spotter_config.model_config.provider_config.provider = SHERPA_ONNX_OR(config->model_config.provider, "cpu"); if (spotter_config.model_config.provider_config.provider.empty()) { @@ -1302,7 +1317,7 @@ static sherpa_onnx::VadModelConfig GetVadModelConfig( SHERPA_ONNX_OR(config->ten_vad.max_speech_duration, 20); vad_config.sample_rate = SHERPA_ONNX_OR(config->sample_rate, 16000); - vad_config.num_threads = SHERPA_ONNX_OR(config->num_threads, 1); + vad_config.num_threads = GetNumThreads(config->num_threads); vad_config.provider = SHERPA_ONNX_OR(config->provider, "cpu"); if (vad_config.provider.empty()) { vad_config.provider = "cpu"; @@ -1578,7 +1593,7 @@ static sherpa_onnx::OfflineTtsConfig GetOfflineTtsConfig( tts_config.model.supertonic.voice_style = SHERPA_ONNX_OR(config->model.supertonic.voice_style, ""); - tts_config.model.num_threads = SHERPA_ONNX_OR(config->model.num_threads, 1); + tts_config.model.num_threads = GetNumThreads(config->model.num_threads); tts_config.model.debug = config->model.debug; tts_config.model.provider = SHERPA_ONNX_OR(config->model.provider, "cpu"); if (tts_config.model.provider.empty()) { @@ -2141,7 +2156,7 @@ SherpaOnnxCreateSpokenLanguageIdentification( slid_config.whisper.decoder = SHERPA_ONNX_OR(config->whisper.decoder, ""); slid_config.whisper.tail_paddings = SHERPA_ONNX_OR(config->whisper.tail_paddings, -1); - slid_config.num_threads = SHERPA_ONNX_OR(config->num_threads, 1); + slid_config.num_threads = GetNumThreads(config->num_threads); slid_config.debug = config->debug; slid_config.provider = SHERPA_ONNX_OR(config->provider, "cpu"); if (slid_config.provider.empty()) { @@ -2215,7 +2230,7 @@ GetSpeakerEmbeddingExtractorConfig( sherpa_onnx::SpeakerEmbeddingExtractorConfig c; c.model = SHERPA_ONNX_OR(config->model, ""); - c.num_threads = SHERPA_ONNX_OR(config->num_threads, 1); + c.num_threads = GetNumThreads(config->num_threads); c.debug = config->debug; c.provider = SHERPA_ONNX_OR(config->provider, "cpu"); if (c.provider.empty()) { @@ -2472,7 +2487,7 @@ const SherpaOnnxAudioTagging *SherpaOnnxCreateAudioTagging( sherpa_onnx::AudioTaggingConfig ac; ac.model.zipformer.model = SHERPA_ONNX_OR(config->model.zipformer.model, ""); ac.model.ced = SHERPA_ONNX_OR(config->model.ced, ""); - ac.model.num_threads = SHERPA_ONNX_OR(config->model.num_threads, 1); + ac.model.num_threads = GetNumThreads(config->model.num_threads); ac.model.debug = config->model.debug; ac.model.provider = SHERPA_ONNX_OR(config->model.provider, "cpu"); if (ac.model.provider.empty()) { @@ -2567,7 +2582,7 @@ static sherpa_onnx::OfflinePunctuationConfig GetOfflinePunctuationConfig( const SherpaOnnxOfflinePunctuationConfig *config) { sherpa_onnx::OfflinePunctuationConfig c; c.model.ct_transformer = SHERPA_ONNX_OR(config->model.ct_transformer, ""); - c.model.num_threads = SHERPA_ONNX_OR(config->model.num_threads, 1); + c.model.num_threads = GetNumThreads(config->model.num_threads); c.model.debug = config->model.debug; c.model.provider = SHERPA_ONNX_OR(config->model.provider, "cpu"); if (c.model.provider.empty()) { @@ -2639,7 +2654,7 @@ static sherpa_onnx::OnlinePunctuationConfig GetOnlinePunctuationConfig( punctuation_config.model.bpe_vocab = SHERPA_ONNX_OR(config->model.bpe_vocab, ""); punctuation_config.model.num_threads = - SHERPA_ONNX_OR(config->model.num_threads, 1); + GetNumThreads(config->model.num_threads); punctuation_config.model.debug = config->model.debug; punctuation_config.model.provider = SHERPA_ONNX_OR(config->model.provider, "cpu"); @@ -2809,7 +2824,7 @@ static sherpa_onnx::OfflineSpeechDenoiserConfig GetOfflineSpeechDenoiserConfig( const SherpaOnnxOfflineSpeechDenoiserConfig *config) { sherpa_onnx::OfflineSpeechDenoiserConfig c; c.model.gtcrn.model = SHERPA_ONNX_OR(config->model.gtcrn.model, ""); - c.model.num_threads = SHERPA_ONNX_OR(config->model.num_threads, 1); + c.model.num_threads = GetNumThreads(config->model.num_threads); c.model.debug = config->model.debug; c.model.provider = SHERPA_ONNX_OR(config->model.provider, "cpu"); c.model.dpdfnet.model = SHERPA_ONNX_OR(config->model.dpdfnet.model, ""); @@ -2897,7 +2912,7 @@ GetOfflineSourceSeparationConfig( c.model.spleeter.accompaniment = SHERPA_ONNX_OR(config->model.spleeter.accompaniment, ""); c.model.uvr.model = SHERPA_ONNX_OR(config->model.uvr.model, ""); - c.model.num_threads = SHERPA_ONNX_OR(config->model.num_threads, 1); + c.model.num_threads = GetNumThreads(config->model.num_threads); c.model.debug = config->model.debug; c.model.provider = SHERPA_ONNX_OR(config->model.provider, "cpu"); @@ -3013,7 +3028,7 @@ static sherpa_onnx::OnlineSpeechDenoiserConfig GetOnlineSpeechDenoiserConfig( const SherpaOnnxOnlineSpeechDenoiserConfig *config) { sherpa_onnx::OnlineSpeechDenoiserConfig c; c.model.gtcrn.model = SHERPA_ONNX_OR(config->model.gtcrn.model, ""); - c.model.num_threads = SHERPA_ONNX_OR(config->model.num_threads, 1); + c.model.num_threads = GetNumThreads(config->model.num_threads); c.model.debug = config->model.debug; c.model.provider = SHERPA_ONNX_OR(config->model.provider, "cpu"); c.model.dpdfnet.model = SHERPA_ONNX_OR(config->model.dpdfnet.model, ""); @@ -3133,7 +3148,7 @@ GetOfflineSpeakerDiarizationConfig( sd_config.segmentation.pyannote.model = SHERPA_ONNX_OR(config->segmentation.pyannote.model, ""); sd_config.segmentation.num_threads = - SHERPA_ONNX_OR(config->segmentation.num_threads, 1); + GetNumThreads(config->segmentation.num_threads); sd_config.segmentation.debug = config->segmentation.debug; sd_config.segmentation.provider = SHERPA_ONNX_OR(config->segmentation.provider, "cpu"); @@ -3143,7 +3158,7 @@ GetOfflineSpeakerDiarizationConfig( sd_config.embedding.model = SHERPA_ONNX_OR(config->embedding.model, ""); sd_config.embedding.num_threads = - SHERPA_ONNX_OR(config->embedding.num_threads, 1); + GetNumThreads(config->embedding.num_threads); sd_config.embedding.debug = config->embedding.debug; sd_config.embedding.provider = SHERPA_ONNX_OR(config->embedding.provider, "cpu"); @@ -3653,7 +3668,7 @@ static sherpa_onnx::OfflineDiacritizationConfig GetOfflineDiacritizationConfig( sherpa_onnx::OfflineDiacritizationConfig c; c.model.catt_encoder = SHERPA_ONNX_OR(config->model.catt_encoder, ""); c.model.catt_decoder = SHERPA_ONNX_OR(config->model.catt_decoder, ""); - c.model.num_threads = SHERPA_ONNX_OR(config->model.num_threads, 1); + c.model.num_threads = GetNumThreads(config->model.num_threads); c.model.debug = config->model.debug; c.model.provider = SHERPA_ONNX_OR(config->model.provider, "cpu"); if (c.model.provider.empty()) { diff --git a/sherpa-onnx/csrc/offline-tts-matcha-impl.h b/sherpa-onnx/csrc/offline-tts-matcha-impl.h index 3866a3f97a..f12b3e5037 100644 --- a/sherpa-onnx/csrc/offline-tts-matcha-impl.h +++ b/sherpa-onnx/csrc/offline-tts-matcha-impl.h @@ -44,7 +44,7 @@ class OfflineTtsMatchaImpl : public OfflineTtsImpl { } if (!FileExists(config.model.matcha.vocoder)) { - SHERPA_ONNX_LOGE("Please vocoder '%s' does not exist", + SHERPA_ONNX_LOGE("The vocoder '%s' does not exist", config.model.matcha.vocoder.c_str()); SHERPA_ONNX_EXIT(-1); } diff --git a/wasm/CMakeLists.txt b/wasm/CMakeLists.txt index 0f18d3130b..a9b24e1699 100644 --- a/wasm/CMakeLists.txt +++ b/wasm/CMakeLists.txt @@ -29,3 +29,7 @@ endif() if(SHERPA_ONNX_ENABLE_WASM_NODEJS) add_subdirectory(nodejs) endif() + +if(SHERPA_ONNX_ENABLE_WASM_WEB) + add_subdirectory(web) +endif() diff --git a/wasm/nodejs/CMakeLists.txt b/wasm/nodejs/CMakeLists.txt index 4f015adfec..b37652fbe6 100644 --- a/wasm/nodejs/CMakeLists.txt +++ b/wasm/nodejs/CMakeLists.txt @@ -2,143 +2,11 @@ if(NOT $ENV{SHERPA_ONNX_IS_USING_BUILD_WASM_SH}) message(FATAL_ERROR "Please use ./build-wasm-simd-nodejs.sh to build for wasm NodeJS") endif() -set(exported_functions - #tts - PrintOfflineTtsConfig - SherpaOnnxCreateOfflineTts - SherpaOnnxDestroyOfflineTts - SherpaOnnxDestroyOfflineTtsGeneratedAudio - SherpaOnnxOfflineTtsGenerate - SherpaOnnxOfflineTtsGenerateWithCallback - SherpaOnnxOfflineTtsGenerateWithConfig - SherpaOnnxOfflineTtsNumSpeakers - SherpaOnnxOfflineTtsSampleRate - SherpaOnnxWriteWave - # streaming asr - SherpaOnnxCreateOnlineRecognizer - SherpaOnnxCreateOnlineStream - SherpaOnnxDecodeOnlineStream - SherpaOnnxDestroyOnlineRecognizer - SherpaOnnxDestroyOnlineRecognizerResult - SherpaOnnxDestroyOnlineStream - SherpaOnnxDestroyOnlineStreamResultJson - SherpaOnnxGetOnlineStreamResult - SherpaOnnxGetOnlineStreamResultAsJson - SherpaOnnxIsOnlineStreamReady - SherpaOnnxOnlineStreamAcceptWaveform - SherpaOnnxOnlineStreamGetOption - SherpaOnnxOnlineStreamInputFinished - SherpaOnnxOnlineStreamIsEndpoint - SherpaOnnxOnlineStreamReset - SherpaOnnxOnlineStreamSetOption - # non-streaming ASR - PrintOfflineRecognizerConfig - SherpaOnnxAcceptWaveformOffline - SherpaOnnxCreateOfflineRecognizer - SherpaOnnxCreateOfflineStream - SherpaOnnxDecodeMultipleOfflineStreams - SherpaOnnxDecodeOfflineStream - SherpaOnnxDestroyOfflineRecognizer - SherpaOnnxDestroyOfflineRecognizerResult - SherpaOnnxDestroyOfflineStream - SherpaOnnxDestroyOfflineStreamResultJson - SherpaOnnxGetOfflineStreamResult - SherpaOnnxGetOfflineStreamResultAsJson - SherpaOnnxOfflineStreamGetOption - SherpaOnnxOfflineStreamSetOption - SherpaOnnxOfflineRecognizerSetConfig - # online kws - SherpaOnnxCreateKeywordSpotter - SherpaOnnxCreateKeywordStream - SherpaOnnxDecodeKeywordStream - SherpaOnnxDestroyKeywordResult - SherpaOnnxDestroyKeywordSpotter - SherpaOnnxGetKeywordResult - SherpaOnnxIsKeywordStreamReady - SherpaOnnxResetKeywordStream - # VAD - SherpaOnnxCreateCircularBuffer - SherpaOnnxDestroyCircularBuffer - SherpaOnnxCircularBufferPush - SherpaOnnxCircularBufferGet - SherpaOnnxCircularBufferFree - SherpaOnnxCircularBufferPop - SherpaOnnxCircularBufferSize - SherpaOnnxCircularBufferHead - SherpaOnnxCircularBufferReset - SherpaOnnxCreateVoiceActivityDetector - SherpaOnnxDestroyVoiceActivityDetector - SherpaOnnxVoiceActivityDetectorAcceptWaveform - SherpaOnnxVoiceActivityDetectorEmpty - SherpaOnnxVoiceActivityDetectorDetected - SherpaOnnxVoiceActivityDetectorPop - SherpaOnnxVoiceActivityDetectorClear - SherpaOnnxVoiceActivityDetectorFront - SherpaOnnxDestroySpeechSegment - SherpaOnnxVoiceActivityDetectorReset - SherpaOnnxVoiceActivityDetectorFlush - # Speaker diarization - SherpaOnnxCreateOfflineSpeakerDiarization - SherpaOnnxDestroyOfflineSpeakerDiarization - SherpaOnnxOfflineSpeakerDiarizationDestroyResult - SherpaOnnxOfflineSpeakerDiarizationDestroySegment - SherpaOnnxOfflineSpeakerDiarizationGetSampleRate - SherpaOnnxOfflineSpeakerDiarizationProcess - SherpaOnnxOfflineSpeakerDiarizationProcessWithCallback - SherpaOnnxOfflineSpeakerDiarizationResultGetNumSegments - SherpaOnnxOfflineSpeakerDiarizationResultSortByStartTime - SherpaOnnxOfflineSpeakerDiarizationSetConfig - # - SherpaOnnxFileExists - SherpaOnnxReadWave - SherpaOnnxReadWaveFromBinaryData - SherpaOnnxFreeWave - SherpaOnnxWriteWave - # speech enhancement - SherpaOnnxCreateOfflineSpeechDenoiser - SherpaOnnxCreateOnlineSpeechDenoiser - SherpaOnnxDestroyDenoisedAudio - SherpaOnnxDestroyOfflineSpeechDenoiser - SherpaOnnxDestroyOnlineSpeechDenoiser - SherpaOnnxOfflineSpeechDenoiserGetSampleRate - SherpaOnnxOfflineSpeechDenoiserRun - SherpaOnnxOnlineSpeechDenoiserGetFrameShiftInSamples - SherpaOnnxOnlineSpeechDenoiserGetSampleRate - SherpaOnnxOnlineSpeechDenoiserRun - SherpaOnnxOnlineSpeechDenoiserFlush - SherpaOnnxOnlineSpeechDenoiserReset - # punctuation - SherpaOnnxCreateOfflinePunctuation - SherpaOnnxDestroyOfflinePunctuation - SherpaOfflinePunctuationAddPunct - SherpaOfflinePunctuationFreeText - SherpaOnnxCreateOnlinePunctuation - SherpaOnnxDestroyOnlinePunctuation - SherpaOnnxOnlinePunctuationAddPunct - SherpaOnnxOnlinePunctuationFreeText - # version - SherpaOnnxGetGitDate - SherpaOnnxGetGitSha1 - SherpaOnnxGetVersionStr - SherpaOnnxGetOnnxruntimeVersionStr -) - +include(${CMAKE_CURRENT_SOURCE_DIR}/../wasm-common.cmake) -set(mangled_exported_functions) -foreach(x IN LISTS exported_functions) - list(APPEND mangled_exported_functions "_${x}") -endforeach() -list(JOIN mangled_exported_functions "," all_exported_functions) - -include_directories(${CMAKE_SOURCE_DIR}) -set(MY_FLAGS " -s FORCE_FILESYSTEM=1 -s INITIAL_MEMORY=512MB -s ALLOW_MEMORY_GROWTH=1") -string(APPEND MY_FLAGS " -sSTACK_SIZE=10485760 ") # 10MB -string(APPEND MY_FLAGS " -sALLOW_TABLE_GROWTH ") -string(APPEND MY_FLAGS " -sEXPORTED_FUNCTIONS=[_CopyHeap,_malloc,_free,${all_exported_functions}] ") +# Node.js specific flags string(APPEND MY_FLAGS " -sNODERAWFS=1 ") -string(APPEND MY_FLAGS " -sEXPORTED_RUNTIME_METHODS=['ccall','stringToUTF8','setValue','getValue','lengthBytesUTF8','UTF8ToString','HEAPU8','HEAP16','HEAP32','HEAPU32','HEAPF32','HEAPF64','addFunction','removeFunction'] ") - -string(APPEND MY_FLAGS " -sMODULARIZE=1 -sWASM_ASYNC_COMPILATION=0 ") +string(APPEND MY_FLAGS " -sWASM_ASYNC_COMPILATION=0 ") message(STATUS "MY_FLAGS: ${MY_FLAGS}") diff --git a/wasm/tts/sherpa-onnx-tts.js b/wasm/tts/sherpa-onnx-tts.js index c2c707f22e..340087b425 100644 --- a/wasm/tts/sherpa-onnx-tts.js +++ b/wasm/tts/sherpa-onnx-tts.js @@ -43,7 +43,7 @@ function initSherpaOnnxOfflineTtsVitsModelConfig(config, Module) { const lexiconLen = Module.lengthBytesUTF8(config.lexicon || '') + 1; const tokensLen = Module.lengthBytesUTF8(config.tokens || '') + 1; const dataDirLen = Module.lengthBytesUTF8(config.dataDir || '') + 1; - const dictDir = '' + const dictDir = ''; const dictDirLen = Module.lengthBytesUTF8(dictDir) + 1; const n = modelLen + lexiconLen + tokensLen + dataDirLen + dictDirLen; @@ -281,6 +281,11 @@ function initSherpaOnnxOfflineTtsKittenModelConfig(config, Module) { } function initSherpaOnnxOfflineTtsZipVoiceModelConfig(config, Module) { + // Support Dart camelCase: targetRms -> targetRMS + if ('targetRms' in config && !('targetRMS' in config)) { + config.targetRMS = config.targetRms; + } + const tokensLen = Module.lengthBytesUTF8(config.tokens || '') + 1; const encoderLen = Module.lengthBytesUTF8(config.encoder || '') + 1; const decoderLen = Module.lengthBytesUTF8(config.decoder || '') + 1; @@ -502,6 +507,22 @@ function initSherpaOnnxOfflineTtsSupertonicModelConfig(config, Module) { } function initSherpaOnnxOfflineTtsModelConfig(config, Module) { + // Support short aliases (used by Dart web implementation). + const aliasMap = { + vits: 'offlineTtsVitsModelConfig', + matcha: 'offlineTtsMatchaModelConfig', + kokoro: 'offlineTtsKokoroModelConfig', + kitten: 'offlineTtsKittenModelConfig', + zipvoice: 'offlineTtsZipVoiceModelConfig', + pocket: 'offlineTtsPocketModelConfig', + supertonic: 'offlineTtsSupertonicModelConfig', + }; + for (const [short, full] of Object.entries(aliasMap)) { + if (short in config && !(full in config)) { + config[full] = config[short]; + } + } + if (!('offlineTtsVitsModelConfig' in config)) { config.offlineTtsVitsModelConfig = { model: '', @@ -665,6 +686,15 @@ function initSherpaOnnxOfflineTtsModelConfig(config, Module) { } function initSherpaOnnxOfflineTtsConfig(config, Module) { + // Support short alias for model config (used by Dart web implementation). + if ('model' in config && !('offlineTtsModelConfig' in config)) { + config.offlineTtsModelConfig = config.model; + } + // Support typo in Dart config class. + if ('maxNumSenetences' in config && !('maxNumSentences' in config)) { + config.maxNumSentences = config.maxNumSenetences; + } + const modelConfig = initSherpaOnnxOfflineTtsModelConfig(config.offlineTtsModelConfig, Module); const len = modelConfig.len + 4 * 4; diff --git a/wasm/wasm-common.cmake b/wasm/wasm-common.cmake new file mode 100644 index 0000000000..85d5a2056a --- /dev/null +++ b/wasm/wasm-common.cmake @@ -0,0 +1,137 @@ +# Shared exported functions and flags for sherpa-onnx WASM builds. +# Include this file from wasm/nodejs/CMakeLists.txt and wasm/web/CMakeLists.txt. + +set(exported_functions + #tts + PrintOfflineTtsConfig + SherpaOnnxCreateOfflineTts + SherpaOnnxDestroyOfflineTts + SherpaOnnxDestroyOfflineTtsGeneratedAudio + SherpaOnnxOfflineTtsGenerate + SherpaOnnxOfflineTtsGenerateWithCallback + SherpaOnnxOfflineTtsGenerateWithConfig + SherpaOnnxOfflineTtsNumSpeakers + SherpaOnnxOfflineTtsSampleRate + SherpaOnnxWriteWave + # streaming asr + SherpaOnnxCreateOnlineRecognizer + SherpaOnnxCreateOnlineStream + SherpaOnnxDecodeOnlineStream + SherpaOnnxDestroyOnlineRecognizer + SherpaOnnxDestroyOnlineRecognizerResult + SherpaOnnxDestroyOnlineStream + SherpaOnnxDestroyOnlineStreamResultJson + SherpaOnnxGetOnlineStreamResult + SherpaOnnxGetOnlineStreamResultAsJson + SherpaOnnxIsOnlineStreamReady + SherpaOnnxOnlineStreamAcceptWaveform + SherpaOnnxOnlineStreamGetOption + SherpaOnnxOnlineStreamInputFinished + SherpaOnnxOnlineStreamIsEndpoint + SherpaOnnxOnlineStreamReset + SherpaOnnxOnlineStreamSetOption + # non-streaming ASR + PrintOfflineRecognizerConfig + SherpaOnnxAcceptWaveformOffline + SherpaOnnxCreateOfflineRecognizer + SherpaOnnxCreateOfflineStream + SherpaOnnxDecodeMultipleOfflineStreams + SherpaOnnxDecodeOfflineStream + SherpaOnnxDestroyOfflineRecognizer + SherpaOnnxDestroyOfflineRecognizerResult + SherpaOnnxDestroyOfflineStream + SherpaOnnxDestroyOfflineStreamResultJson + SherpaOnnxGetOfflineStreamResult + SherpaOnnxGetOfflineStreamResultAsJson + SherpaOnnxOfflineStreamGetOption + SherpaOnnxOfflineStreamSetOption + SherpaOnnxOfflineRecognizerSetConfig + # online kws + SherpaOnnxCreateKeywordSpotter + SherpaOnnxCreateKeywordStream + SherpaOnnxDecodeKeywordStream + SherpaOnnxDestroyKeywordResult + SherpaOnnxDestroyKeywordSpotter + SherpaOnnxGetKeywordResult + SherpaOnnxIsKeywordStreamReady + SherpaOnnxResetKeywordStream + # VAD + SherpaOnnxCreateCircularBuffer + SherpaOnnxDestroyCircularBuffer + SherpaOnnxCircularBufferPush + SherpaOnnxCircularBufferGet + SherpaOnnxCircularBufferFree + SherpaOnnxCircularBufferPop + SherpaOnnxCircularBufferSize + SherpaOnnxCircularBufferHead + SherpaOnnxCircularBufferReset + SherpaOnnxCreateVoiceActivityDetector + SherpaOnnxDestroyVoiceActivityDetector + SherpaOnnxVoiceActivityDetectorAcceptWaveform + SherpaOnnxVoiceActivityDetectorEmpty + SherpaOnnxVoiceActivityDetectorDetected + SherpaOnnxVoiceActivityDetectorPop + SherpaOnnxVoiceActivityDetectorClear + SherpaOnnxVoiceActivityDetectorFront + SherpaOnnxDestroySpeechSegment + SherpaOnnxVoiceActivityDetectorReset + SherpaOnnxVoiceActivityDetectorFlush + # Speaker diarization + SherpaOnnxCreateOfflineSpeakerDiarization + SherpaOnnxDestroyOfflineSpeakerDiarization + SherpaOnnxOfflineSpeakerDiarizationDestroyResult + SherpaOnnxOfflineSpeakerDiarizationDestroySegment + SherpaOnnxOfflineSpeakerDiarizationGetSampleRate + SherpaOnnxOfflineSpeakerDiarizationProcess + SherpaOnnxOfflineSpeakerDiarizationProcessWithCallback + SherpaOnnxOfflineSpeakerDiarizationResultGetNumSegments + SherpaOnnxOfflineSpeakerDiarizationResultSortByStartTime + SherpaOnnxOfflineSpeakerDiarizationSetConfig + # + SherpaOnnxFileExists + SherpaOnnxReadWave + SherpaOnnxReadWaveFromBinaryData + SherpaOnnxFreeWave + SherpaOnnxWriteWave + # speech enhancement + SherpaOnnxCreateOfflineSpeechDenoiser + SherpaOnnxCreateOnlineSpeechDenoiser + SherpaOnnxDestroyDenoisedAudio + SherpaOnnxDestroyOfflineSpeechDenoiser + SherpaOnnxDestroyOnlineSpeechDenoiser + SherpaOnnxOfflineSpeechDenoiserGetSampleRate + SherpaOnnxOfflineSpeechDenoiserRun + SherpaOnnxOnlineSpeechDenoiserGetFrameShiftInSamples + SherpaOnnxOnlineSpeechDenoiserGetSampleRate + SherpaOnnxOnlineSpeechDenoiserRun + SherpaOnnxOnlineSpeechDenoiserFlush + SherpaOnnxOnlineSpeechDenoiserReset + # punctuation + SherpaOnnxCreateOfflinePunctuation + SherpaOnnxDestroyOfflinePunctuation + SherpaOfflinePunctuationAddPunct + SherpaOfflinePunctuationFreeText + SherpaOnnxCreateOnlinePunctuation + SherpaOnnxDestroyOnlinePunctuation + SherpaOnnxOnlinePunctuationAddPunct + SherpaOnnxOnlinePunctuationFreeText + # version + SherpaOnnxGetGitDate + SherpaOnnxGetGitSha1 + SherpaOnnxGetVersionStr + SherpaOnnxGetOnnxruntimeVersionStr +) + +set(mangled_exported_functions) +foreach(x IN LISTS exported_functions) + list(APPEND mangled_exported_functions "_${x}") +endforeach() +list(JOIN mangled_exported_functions "," all_exported_functions) + +include_directories(${CMAKE_SOURCE_DIR}) +set(MY_FLAGS " -s FORCE_FILESYSTEM=1 -s INITIAL_MEMORY=512MB -s ALLOW_MEMORY_GROWTH=1") +string(APPEND MY_FLAGS " -sSTACK_SIZE=10485760 ") # 10MB +string(APPEND MY_FLAGS " -sALLOW_TABLE_GROWTH ") +string(APPEND MY_FLAGS " -sEXPORTED_FUNCTIONS=[_CopyHeap,_malloc,_free,${all_exported_functions}] ") +string(APPEND MY_FLAGS " -sEXPORTED_RUNTIME_METHODS=['ccall','stringToUTF8','setValue','getValue','lengthBytesUTF8','UTF8ToString','HEAPU8','HEAP16','HEAP32','HEAPU32','HEAPF32','HEAPF64','addFunction','removeFunction','FS'] ") +string(APPEND MY_FLAGS " -sMODULARIZE=1 ") diff --git a/wasm/web/CMakeLists.txt b/wasm/web/CMakeLists.txt new file mode 100644 index 0000000000..5133b1c64d --- /dev/null +++ b/wasm/web/CMakeLists.txt @@ -0,0 +1,26 @@ +if(NOT $ENV{SHERPA_ONNX_IS_USING_BUILD_WASM_SH}) + message(FATAL_ERROR "Please use ./build-wasm-simd-web.sh to build for wasm Web") +endif() + +include(${CMAKE_CURRENT_SOURCE_DIR}/../wasm-common.cmake) + +# Browser specific flags +string(APPEND MY_FLAGS " -s 'EXPORT_NAME=\"SherpaOnnx\"' ") + +message(STATUS "MY_FLAGS: ${MY_FLAGS}") + +set(CMAKE_C_FLAGS "${CMAKE_C_FLAGS} ${MY_FLAGS}") +set(CMAKE_CXX_FLAGS "${CMAKE_CXX_FLAGS} ${MY_FLAGS}") +set(CMAKE_EXECUTABLE_LINKER_FLAGS "${CMAKE_EXECUTABLE_LINKER_FLAGS} ${MY_FLAGS}") + +add_executable(sherpa-onnx-wasm-web ${CMAKE_CURRENT_SOURCE_DIR}/../nodejs/sherpa-onnx-wasm-nodejs.cc) +target_link_libraries(sherpa-onnx-wasm-web sherpa-onnx-core sherpa-onnx-c-api) +install(TARGETS sherpa-onnx-wasm-web DESTINATION bin/wasm/web) + +install( + FILES + "$/sherpa-onnx-wasm-web.js" + "$/sherpa-onnx-wasm-web.wasm" + DESTINATION + bin/wasm/web +)