diff --git a/Dockerfile.gpu b/Dockerfile.gpu new file mode 100644 index 0000000..a62089a --- /dev/null +++ b/Dockerfile.gpu @@ -0,0 +1,59 @@ +# Copyright 2022 Mycroft AI Inc. +# +# This program is free software: you can redistribute it and/or modify +# it under the terms of the GNU Affero General Public License as published by +# the Free Software Foundation, either version 3 of the License, or +# (at your option) any later version. +# +# This program is distributed in the hope that it will be useful, +# but WITHOUT ANY WARRANTY; without even the implied warranty of +# MERCHANTABILITY or FITNESS FOR A PARTICULAR PURPOSE. See the +# GNU Affero General Public License for more details. +# +# You should have received a copy of the GNU Affero General Public License +# along with this program. If not, see . +# +# ----------------------------------------------------------------------------- +# Dockerfile for Mimic 3 (https://github.com/MycroftAI/mimic3) +# +# Runs an HTTP server on port 59125. +# See scripts in docker/ directory of this repository. +# +# Requires Docker buildx: https://docs.docker.com/buildx/working-with-buildx/ +# ----------------------------------------------------------------------------- + +FROM nvcr.io/nvidia/cuda:11.4.2-cudnn8-devel-ubuntu20.04 +ARG TARGETARCH +ARG TARGETVARIANT + +ENV LANG C.UTF-8 +ENV DEBIAN_FRONTEND=noninteractive + +RUN echo "Dir::Cache var/cache/apt/${TARGETARCH}${TARGETVARIANT};" > /etc/apt/apt.conf.d/01cache + +RUN --mount=type=cache,id=apt-run,target=/var/cache/apt \ + mkdir -p /var/cache/apt/${TARGETARCH}${TARGETVARIANT}/archives/partial && \ + apt-get update && \ + apt-get install --yes --no-install-recommends \ + python3 python3-pip python3-venv \ + ca-certificates libespeak-ng1 + +WORKDIR /home/mimic3/app + +COPY ./ ./ + +# Fix requirements +RUN sed -i 's/onnxruntime/onnxruntime-gpu/' mimic3-tts/requirements.txt + +# Install mimic3 +RUN --mount=type=cache,id=pip-requirements,target=/root/.cache/pip \ + ./install.sh + +RUN useradd -ms /bin/bash mimic3 + +USER mimic3 +WORKDIR /home/mimic3/app + +EXPOSE 59125 + +ENTRYPOINT ["/home/mimic3/app/.venv/bin/python3", "-m", "mimic3_http", "--cuda"] diff --git a/Makefile b/Makefile index ccc9741..05650ff 100644 --- a/Makefile +++ b/Makefile @@ -33,5 +33,8 @@ install: docker: docker buildx build . -f Dockerfile --platform $(DOCKER_PLATFORM) --tag mycroftai/mimic3 --load +docker-gpu: + docker buildx build . -f Dockerfile.gpu --tag mycroftai/mimic3:gpu --load + binaries: docker buildx build . -f Dockerfile.binary --platform $(DOCKER_PLATFORM) --output type=local,dest=dist/$(DOCKER_PLATFORM) diff --git a/README.md b/README.md index e4653de..0f6b3c9 100644 --- a/README.md +++ b/README.md @@ -194,6 +194,13 @@ mimic3-client --voice 'en_UK/apope_low' 'My hovercraft is full of eels.' > hover See `mimic3-client --help` for more options. +## CUDA Acceleration + +If you have a GPU with support for CUDA, you can accelerate synthesis with the `--cuda` flag when running `mimic3` or `mimic3-server`. This requires you to install the [onnxruntime-gpu](https://pypi.org/project/onnxruntime-gpu/) Python package. + +Using [nvidia-docker](https://github.com/NVIDIA/nvidia-docker) is highly recommended. See [ Dockerfile.gpu](Dockerfile.gpu) for an example of how to build a compatible container. + + ## MaryTTS Compatibility Use the Mimic 3 web server as a drop-in replacement for [MaryTTS](http://mary.dfki.de/), for example with [Home Assistant](https://www.home-assistant.io/integrations/marytts/). diff --git a/docker/mimic3-server b/docker/mimic3-server index f6b137c..831072d 100755 --- a/docker/mimic3-server +++ b/docker/mimic3-server @@ -31,6 +31,11 @@ while [[ -n "$1" ]]; do port="$2" args+=('--port' "${port}") shift 1 + elif [[ "$1" == '--cuda' ]]; then + # Use CUDA GPU acceleration + args+=('--cuda') + docker='nvidia-docker' + tag='gpu' else args+=("$1") fi diff --git a/mimic3-http/README.md b/mimic3-http/README.md index 7c45ffd..0198c92 100644 --- a/mimic3-http/README.md +++ b/mimic3-http/README.md @@ -18,7 +18,7 @@ This will start a web server at `http://localhost:59125` See `mimic3-server --debug` for more options. -## Endpoints +### Endpoints * `/api/tts` * `POST` text or [SSML](#ssml) and receive WAV audio back @@ -30,6 +30,13 @@ See `mimic3-server --debug` for more options. An [OpenAPI](https://www.openapis.org/) test page is also available at `http://localhost:59125/openapi` +### CUDA Acceleration + +If you have a GPU with support for CUDA, you can accelerate synthesis with the `--cuda` flag. This requires you to install the [onnxruntime-gpu](https://pypi.org/project/onnxruntime-gpu/) Python package. + +Using [nvidia-docker](https://github.com/NVIDIA/nvidia-docker) is highly recommended. See the `Dockerfile.gpu` file in the parent repository for an example of how to build a compatible container. + + ## Running the Client Assuming you have started `mimic3-server` and can access `http://localhost:59125`, then: diff --git a/mimic3-http/mimic3_http/__main__.py b/mimic3-http/mimic3_http/__main__.py index fc36bf0..7540296 100644 --- a/mimic3-http/mimic3_http/__main__.py +++ b/mimic3-http/mimic3_http/__main__.py @@ -17,67 +17,70 @@ import asyncio import logging import tempfile +import threading +from queue import Queue import hypercorn -from mimic3_tts import Mimic3Settings, Mimic3TextToSpeechSystem from .app import get_app from .args import get_args +from .synthesis import do_synthesis_proc _LOGGER = logging.getLogger(__name__) # ----------------------------------------------------------------------------- -args = get_args() -if args.debug: - logging.basicConfig(level=logging.DEBUG) +def main(): + args = get_args() - # Override epitran - logging.getLogger().setLevel(logging.DEBUG) -else: - logging.basicConfig(level=logging.INFO) + if args.debug: + logging.basicConfig(level=logging.DEBUG) - # Override epitran - logging.getLogger().setLevel(logging.INFO) + # Override epitran + logging.getLogger().setLevel(logging.DEBUG) + else: + logging.basicConfig(level=logging.INFO) + # Override epitran + logging.getLogger().setLevel(logging.INFO) -_LOGGER.debug(args) + _LOGGER.debug(args) + + # Run Web Server + _LOGGER.info("Starting web server") + request_queue = Queue() + threads = [ + threading.Thread( + target=do_synthesis_proc, args=(args, request_queue), daemon=True + ) + for _ in range(args.num_threads) + ] + for thread in threads: + thread.start() + + hyp_config = hypercorn.config.Config() + hyp_config.bind = [f"{args.host}:{args.port}"] + + try: + with tempfile.TemporaryDirectory(prefix="mimic3") as temp_dir: + app = get_app(args, request_queue, temp_dir) + asyncio.run(hypercorn.asyncio.serve(app, hyp_config)) + finally: + # Drain queue + while not request_queue.empty(): + request_queue.get() + + # Stop request threads + for _ in range(args.num_threads): + request_queue.put(None) + + for thread in threads: + thread.join() -# ----------------------------------------------------------------------------- -# Load Mimic 3 # ----------------------------------------------------------------------------- -# TODO: args.voices_dir - -mimic3 = Mimic3TextToSpeechSystem( - Mimic3Settings( - voice=args.voice, - speaker=args.speaker, - length_scale=args.length_scale, - noise_scale=args.noise_scale, - noise_w=args.noise_w, - ) -) - -if args.preload_voice: - # Ensure voices are preloaded - for voice_key in args.preload_voice: - _LOGGER.debug("Preloading voice: %s", voice_key) - mimic3.preload_voice(voice_key) - - -# ----------------------------------------------------------------------------- -# Run Web Server -# ----------------------------------------------------------------------------- - -_LOGGER.info("Starting web server") - -hyp_config = hypercorn.config.Config() -hyp_config.bind = [f"{args.host}:{args.port}"] - -with mimic3, tempfile.TemporaryDirectory(prefix="mimic3") as temp_dir: - app = get_app(args, mimic3, temp_dir) - asyncio.run(hypercorn.asyncio.serve(app, hyp_config)) +if __name__ == "__main__": + main() diff --git a/mimic3-http/mimic3_http/app.py b/mimic3-http/mimic3_http/app.py index 20b503b..ed26db9 100644 --- a/mimic3-http/mimic3_http/app.py +++ b/mimic3-http/mimic3_http/app.py @@ -15,19 +15,17 @@ # along with this program. If not, see . # import argparse +import asyncio import dataclasses -import hashlib -import io import logging import typing -import wave -from dataclasses import dataclass from pathlib import Path +from queue import Queue from urllib.parse import parse_qs from uuid import uuid4 import quart_cors -from mimic3_tts import AudioResult, Mimic3TextToSpeechSystem, SSMLSpeaker +from mimic3_tts import Mimic3Settings, Mimic3TextToSpeechSystem from quart import ( Quart, Response, @@ -40,15 +38,19 @@ from swagger_ui import api_doc from ._resources import _DIR, _PACKAGE from .args import _MISSING +from .const import SynthesisRequest, TextToWavParams _LOGGER = logging.getLogger(__name__) -def get_app(args: argparse.Namespace, mimic3: Mimic3TextToSpeechSystem, temp_dir: str): +def get_app(args: argparse.Namespace, request_queue: Queue, temp_dir: str): """Create and return Quart application for Mimic 3 HTTP server""" _TEMP_DIR: typing.Optional[Path] = None + # TODO: args.voices_dirs + _MIMIC3 = Mimic3TextToSpeechSystem(Mimic3Settings()) + if args.cache_dir != _MISSING: if args.cache_dir is None: # Use temporary directory @@ -61,23 +63,7 @@ def get_app(args: argparse.Namespace, mimic3: Mimic3TextToSpeechSystem, temp_dir if _TEMP_DIR: _LOGGER.debug("Cache directory: %s", _TEMP_DIR) - @dataclass - class TextToWavParams: - """Synthesis parameters used for caching""" - - text: str - voice: str = args.voice - noise_scale: float = args.noise_scale - noise_w: float = args.noise_w - length_scale: float = args.length_scale - ssml: bool = False - text_language: typing.Optional[str] = None - - @property - def cache_key(self) -> str: - return hashlib.md5(repr(self).encode()).hexdigest() - - def text_to_wav(params: TextToWavParams, no_cache: bool = False) -> bytes: + async def text_to_wav(params: TextToWavParams, no_cache: bool = False) -> bytes: """Synthesize text into audio. Returns: WAV bytes @@ -93,58 +79,25 @@ def get_app(args: argparse.Namespace, mimic3: Mimic3TextToSpeechSystem, temp_dir wav_bytes = maybe_wav_path.read_bytes() return wav_bytes - mimic3.voice = params.voice + loop = asyncio.get_running_loop() + future = loop.create_future() + request_queue.put_nowait( + SynthesisRequest( + params=params, + loop=loop, + future=future, + ) + ) + wav_bytes = await future - mimic3.settings.length_scale = params.length_scale - mimic3.settings.noise_scale = params.noise_scale - mimic3.settings.noise_w = params.noise_w + if _TEMP_DIR and (not no_cache): + # Store in cache + wav_path = _TEMP_DIR / f"{params.cache_key}.wav" + wav_path.write_bytes(wav_bytes) - with io.BytesIO() as wav_io: - wav_file: wave.Wave_write = wave.open(wav_io, "wb") - wav_params_set = False + _LOGGER.debug("Cached WAV at %s", wav_path.absolute()) - with wav_file: - try: - if params.ssml: - # SSML - results = SSMLSpeaker(mimic3).speak(params.text) - else: - # Plain text - mimic3.begin_utterance() - mimic3.speak_text( - params.text, text_language=params.text_language - ) - results = mimic3.end_utterance() - - for result in results: - # Add audio to existing WAV file - if isinstance(result, AudioResult): - if not wav_params_set: - wav_file.setframerate(result.sample_rate_hz) - wav_file.setsampwidth(result.sample_width_bytes) - wav_file.setnchannels(result.num_channels) - wav_params_set = True - - wav_file.writeframes(result.audio_bytes) - except Exception as e: - if not wav_params_set: - # Set default parameters so exception can propagate - wav_file.setframerate(22050) - wav_file.setsampwidth(2) - wav_file.setnchannels(1) - - raise e - - wav_bytes = wav_io.getvalue() - - if _TEMP_DIR and (not no_cache): - # Store in cache - wav_path = _TEMP_DIR / f"{params.cache_key}.wav" - wav_path.write_bytes(wav_bytes) - - _LOGGER.debug("Cached WAV at %s", wav_path.absolute()) - - return wav_bytes + return wav_bytes # ----------------------------------------------------------------------------- @@ -186,7 +139,11 @@ def get_app(args: argparse.Namespace, mimic3: Mimic3TextToSpeechSystem, temp_dir @app.route("/api/tts", methods=["GET", "POST"]) async def app_tts() -> Response: """Speak text to WAV.""" - tts_args: typing.Dict[str, typing.Any] = {} + tts_args: typing.Dict[str, typing.Any] = { + "length_scale": args.length_scale, + "noise_scale": args.noise_scale, + "noise_w": args.noise_w, + } _LOGGER.debug("Request args: %s", request.args) @@ -230,7 +187,7 @@ def get_app(args: argparse.Namespace, mimic3: Mimic3TextToSpeechSystem, temp_dir no_cache_str = request.args.get("noCache", "") no_cache = _to_bool(no_cache_str) - wav_bytes = text_to_wav( + wav_bytes = await text_to_wav( TextToWavParams(text=text, **tts_args), no_cache=no_cache ) @@ -238,7 +195,7 @@ def get_app(args: argparse.Namespace, mimic3: Mimic3TextToSpeechSystem, temp_dir @app.route("/api/voices", methods=["GET"]) async def api_voices(): - voices_dict = {v.key: v for v in mimic3.get_voices()} + voices_dict = {v.key: v for v in _MIMIC3.get_voices()} voices = sorted(voices_dict.values(), key=lambda v: v.key) return jsonify([dataclasses.asdict(v) for v in voices]) @@ -263,7 +220,7 @@ def get_app(args: argparse.Namespace, mimic3: Mimic3TextToSpeechSystem, temp_dir ssml = text.strip().startswith("<") _LOGGER.debug("Speaking with voice '%s': %s", voice, text) - wav_bytes = text_to_wav( + wav_bytes = await text_to_wav( TextToWavParams( text=text, voice=voice, diff --git a/mimic3-http/mimic3_http/args.py b/mimic3-http/mimic3_http/args.py index 4809e5f..2ed45ce 100644 --- a/mimic3-http/mimic3_http/args.py +++ b/mimic3-http/mimic3_http/args.py @@ -63,12 +63,17 @@ def get_args() -> argparse.Namespace: parser.add_argument( "--preload-voice", action="append", help="Preload voice when starting up" ) - # parser.add_argument( - # "--max-loaded-models", - # type=int, - # default=0, - # help="Maximum number of voice models that can be loaded simultaneously (0 for no limit)", - # ) + parser.add_argument( + "--cuda", + action="store_true", + help="Use Onnx CUDA execution provider (requires onnxruntime-gpu)", + ) + parser.add_argument( + "--num-threads", + type=int, + default=1, + help="Number of synthesis threads (default: 1)", + ) parser.add_argument( "--debug", action="store_true", help="Print DEBUG messages to console" ) diff --git a/mimic3-http/mimic3_http/const.py b/mimic3-http/mimic3_http/const.py new file mode 100644 index 0000000..08d613f --- /dev/null +++ b/mimic3-http/mimic3_http/const.py @@ -0,0 +1,46 @@ +# Copyright 2022 Mycroft AI Inc. +# +# This program is free software: you can redistribute it and/or modify +# it under the terms of the GNU Affero General Public License as published by +# the Free Software Foundation, either version 3 of the License, or +# (at your option) any later version. +# +# This program is distributed in the hope that it will be useful, +# but WITHOUT ANY WARRANTY; without even the implied warranty of +# MERCHANTABILITY or FITNESS FOR A PARTICULAR PURPOSE. See the +# GNU Affero General Public License for more details. +# +# You should have received a copy of the GNU Affero General Public License +# along with this program. If not, see . +# +import asyncio +import hashlib +import typing +from dataclasses import dataclass + + +@dataclass +class TextToWavParams: + """Synthesis parameters used for caching""" + + text: str + voice: str + noise_scale: float + noise_w: float + length_scale: float + ssml: bool = False + text_language: typing.Optional[str] = None + + @property + def cache_key(self) -> str: + return hashlib.md5(repr(self).encode()).hexdigest() + + +@dataclass +class SynthesisRequest: + """Request to synthesize audio from text""" + + params: TextToWavParams + + loop: asyncio.AbstractEventLoop + future: asyncio.Future diff --git a/mimic3-http/mimic3_http/synthesis.py b/mimic3-http/mimic3_http/synthesis.py new file mode 100644 index 0000000..9052c12 --- /dev/null +++ b/mimic3-http/mimic3_http/synthesis.py @@ -0,0 +1,122 @@ +#!/usr/bin/env python3 +import argparse +import asyncio +import io +import logging +import threading +import typing +import wave +from queue import Queue + +from mimic3_tts import ( + AudioResult, + Mimic3Settings, + Mimic3TextToSpeechSystem, + SSMLSpeaker, +) + +from .const import SynthesisRequest + +_LOGGER = logging.getLogger(__name__) + + +def do_synthesis(item: SynthesisRequest, mimic3: Mimic3TextToSpeechSystem) -> bytes: + """Synthesize text into audio. + + Returns: WAV bytes + """ + params = item.params + mimic3.voice = params.voice + + mimic3.settings.length_scale = params.length_scale + mimic3.settings.noise_scale = params.noise_scale + mimic3.settings.noise_w = params.noise_w + + with io.BytesIO() as wav_io: + wav_file: wave.Wave_write = wave.open(wav_io, "wb") + wav_params_set = False + + with wav_file: + try: + if params.ssml: + # SSML + results = SSMLSpeaker(mimic3).speak(params.text) + else: + # Plain text + mimic3.begin_utterance() + mimic3.speak_text(params.text, text_language=params.text_language) + results = mimic3.end_utterance() + + for result in results: + # Add audio to existing WAV file + if isinstance(result, AudioResult): + if not wav_params_set: + wav_file.setframerate(result.sample_rate_hz) + wav_file.setsampwidth(result.sample_width_bytes) + wav_file.setnchannels(result.num_channels) + wav_params_set = True + + wav_file.writeframes(result.audio_bytes) + except Exception as e: + if not wav_params_set: + # Set default parameters so exception can propagate + wav_file.setframerate(22050) + wav_file.setsampwidth(2) + wav_file.setnchannels(1) + + raise e + + wav_bytes = wav_io.getvalue() + + return wav_bytes + + +def do_synthesis_proc(args: argparse.Namespace, request_queue: Queue): + """Thread handler for synthesis requests""" + try: + # Load Mimic 3 + mimic3 = Mimic3TextToSpeechSystem( + Mimic3Settings( + voice=args.voice, + speaker=args.speaker, + length_scale=args.length_scale, + noise_scale=args.noise_scale, + noise_w=args.noise_w, + use_cuda=args.cuda, + ) + ) + + with mimic3: + if args.preload_voice: + # Ensure voices are preloaded + for voice_key in args.preload_voice: + _LOGGER.debug("Preloading voice: %s", voice_key) + mimic3.preload_voice(voice_key) + + _LOGGER.debug( + "Started inference thread %s", threading.current_thread().ident + ) + + while True: + item = request_queue.get() + if item is None: + # Exit signal + break + + item = typing.cast(SynthesisRequest, item) + + try: + result = do_synthesis(item, mimic3) + + # Set result on main loop + item.loop.call_soon_threadsafe(item.future.set_result, result) + except Exception as e: + _LOGGER.exception("Error during inference") + + # Signal error on main loop + asyncio.get_event_loop().call_soon_threadsafe( + item.future.set_exception, e + ) + + except Exception: + _LOGGER.exception("Unexpected error in inference thread") diff --git a/mimic3-tts/README.md b/mimic3-tts/README.md index 2dc0bab..56cc9b4 100644 --- a/mimic3-tts/README.md +++ b/mimic3-tts/README.md @@ -136,6 +136,14 @@ mimic3 --voices ``` +#### CUDA Acceleration + +If you have a GPU with support for CUDA, you can accelerate synthesis with the `--cuda` flag. This requires you to install the [onnxruntime-gpu](https://pypi.org/project/onnxruntime-gpu/) Python package. + +Using [nvidia-docker](https://github.com/NVIDIA/nvidia-docker) is highly recommended. See the `Dockerfile.gpu` file in the parent repository for an example of how to build a compatible container. + + + ### mimic3-download Mimic 3 automatically downloads voices when they're first used, but you can manually download them too with `mimic3-download`. diff --git a/mimic3-tts/mimic3_tts/__main__.py b/mimic3-tts/mimic3_tts/__main__.py index f51ec01..1983ce1 100644 --- a/mimic3-tts/mimic3_tts/__main__.py +++ b/mimic3-tts/mimic3_tts/__main__.py @@ -219,7 +219,9 @@ def initialize_tts(state: CommandLineInterfaceState): args = state.args state.tts = Mimic3TextToSpeechSystem( - Mimic3Settings(voices_directories=args.voices_dir, speaker=args.speaker) + Mimic3Settings( + voices_directories=args.voices_dir, speaker=args.speaker, use_cuda=args.cuda + ) ) if args.voices: @@ -569,6 +571,11 @@ def get_args(): default=_DEFAULT_PLAY_PROGRAMS, help="Program(s) used to play WAV files", ) + parser.add_argument( + "--cuda", + action="store_true", + help="Use Onnx CUDA execution provider (requires onnxruntime-gpu)", + ) parser.add_argument("--seed", type=int, help="Set random seed (default: not set)") parser.add_argument("--version", action="store_true", help="Print version and exit") parser.add_argument( diff --git a/mimic3-tts/mimic3_tts/tts.py b/mimic3-tts/mimic3_tts/tts.py index 9b0f197..71367f4 100644 --- a/mimic3-tts/mimic3_tts/tts.py +++ b/mimic3-tts/mimic3_tts/tts.py @@ -102,6 +102,12 @@ class Mimic3Settings: no_download: bool = False """Do not download voices automatically""" + use_cuda: bool = False + """Use CUDA GPU acceleration (requires onnxruntime-gpu)""" + + share_onnx_models_between_threads: bool = True + """If True, Onnx models are shared between threads""" + @dataclass class Mimic3Phonemes: @@ -470,7 +476,16 @@ class Mimic3TextToSpeechSystem(TextToSpeechSystem): return existing_voice - voice = Mimic3Voice.load_from_directory(model_dir) + # https://onnxruntime.ai/docs/execution-providers/ + providers = None + if self.settings.use_cuda: + providers = ["CUDAExecutionProvider"] + + voice = Mimic3Voice.load_from_directory( + model_dir, + providers=providers, + share_models=self.settings.share_onnx_models_between_threads, + ) _LOGGER.info("Loaded voice from %s", model_dir) diff --git a/mimic3-tts/mimic3_tts/voice.py b/mimic3-tts/mimic3_tts/voice.py index d0f8c1d..880cd64 100644 --- a/mimic3-tts/mimic3_tts/voice.py +++ b/mimic3-tts/mimic3_tts/voice.py @@ -16,6 +16,7 @@ import csv import logging import platform +import threading import time import typing from abc import ABCMeta, abstractmethod @@ -66,6 +67,9 @@ _LOGGER = logging.getLogger(__name__) class Mimic3Voice(metaclass=ABCMeta): """Base class for Mimic 3 voice implementations""" + _SHARED_MODELS: typing.Dict[str, onnxruntime.InferenceSession] = {} + _SHARED_MODELS_LOCK = threading.Lock() + def __init__( self, config: TrainingConfig, @@ -236,6 +240,12 @@ class Mimic3Voice(metaclass=ABCMeta): def load_from_directory( voice_dir: typing.Union[str, Path], session_options: typing.Optional[onnxruntime.SessionOptions] = None, + providers: typing.Optional[ + typing.Sequence[ + typing.Union[str, typing.Tuple[str, typing.Dict[str, typing.Any]]] + ] + ] = None, + share_models: bool = True, ) -> "Mimic3Voice": """Load a Mimic 3 voice from a directory""" voice_dir = Path(voice_dir) @@ -254,21 +264,30 @@ class Mimic3Voice(metaclass=ABCMeta): phoneme_to_id = phonemes2ids.load_phoneme_ids(ids_file) generator_path = voice_dir / "generator.onnx" - _LOGGER.debug("Loading model from %s", generator_path) - # Load onnx model - if session_options is None: - session_options = onnxruntime.SessionOptions() + onnx_model: typing.Optional[onnxruntime.InferenceSession] = None - if platform.machine() == "armv7l": - # Enabling optimizations on 32-bit ARM crashes - session_options.graph_optimization_level = ( - onnxruntime.GraphOptimizationLevel.ORT_DISABLE_ALL - ) + if share_models: + with Mimic3Voice._SHARED_MODELS_LOCK: + model_key = str(generator_path.absolute()) + onnx_model = Mimic3Voice._SHARED_MODELS.get(model_key) - onnx_model = onnxruntime.InferenceSession( - str(generator_path), sess_options=session_options - ) + if onnx_model is None: + onnx_model = Mimic3Voice._load_model( + generator_path, + session_options=session_options, + providers=providers, + ) + + Mimic3Voice._SHARED_MODELS[model_key] = onnx_model + else: + _LOGGER.debug("Using shared Onnx model (%s)", model_key) + else: + onnx_model = Mimic3Voice._load_model( + generator_path, + session_options=session_options, + providers=providers, + ) # phoneme -> phoneme, phoneme, ... phoneme_map: typing.Optional[PHONEME_MAP_TYPE] = None @@ -334,6 +353,34 @@ class Mimic3Voice(metaclass=ABCMeta): raise ValueError(f"Unsupported phonemizer: {config.phonemizer}") + @staticmethod + def _load_model( + generator_path: Path, + session_options: typing.Optional[onnxruntime.SessionOptions] = None, + providers: typing.Optional[ + typing.Sequence[ + typing.Union[str, typing.Tuple[str, typing.Dict[str, typing.Any]]] + ] + ] = None, + ) -> onnxruntime.InferenceSession: + _LOGGER.debug("Loading model from %s", generator_path) + + # Load onnx model + if session_options is None: + session_options = onnxruntime.SessionOptions() + + if platform.machine() == "armv7l": + # Enabling optimizations on 32-bit ARM crashes + session_options.graph_optimization_level = ( + onnxruntime.GraphOptimizationLevel.ORT_DISABLE_ALL + ) + + onnx_model = onnxruntime.InferenceSession( + str(generator_path), sess_options=session_options, providers=providers + ) + + return onnx_model + # -----------------------------------------------------------------------------