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
+
# -----------------------------------------------------------------------------