Fixing marks in CLI

This commit is contained in:
Michael Hansen 2022-03-22 13:30:30 -04:00
commit 50bc04b987
7 changed files with 237 additions and 665 deletions

View file

@ -22,7 +22,7 @@ from pathlib import Path
from queue import Queue
if typing.TYPE_CHECKING:
from mimic3_tts import Mimic3TextToSpeechSystem
from mimic3_tts import Mimic3TextToSpeechSystem, BaseResult
_DIR = Path(__file__).parent
@ -46,8 +46,15 @@ class CommandLineInterfaceState:
sample_width_bytes: int = 2
num_channels: int = 1
raw_queue: typing.Optional["Queue[typing.Optional[bytes]]"] = None
raw_stream_thread: typing.Optional[threading.Thread] = None
result_queue: typing.Optional["Queue[BaseResult]"] = None
result_thread: typing.Optional[threading.Thread] = None
@dataclass
class ResultToProcess:
result: "BaseResult"
line: str
line_id: str = ""
class OutputNaming(str, Enum):
@ -133,8 +140,10 @@ def initialize_args(state: CommandLineInterfaceState):
state.mark_writer = open( # pylint: disable=consider-using-with
args.mark_file, "w", encoding="utf-8"
)
else:
elif args.stdout:
state.mark_writer = sys.stderr
else:
state.mark_writer = sys.stdout
if args.seed is not None:
_LOGGER.debug("Setting random seed to %s", args.seed)
@ -201,181 +210,138 @@ def initialize_tts(state: CommandLineInterfaceState):
args = state.args
# TODO: voice/speaker
state.tts = Mimic3TextToSpeechSystem(Mimic3Settings())
if state.args.voice:
# Set default voice
state.tts.voice = state.args.voice
# max_thread_workers: typing.Optional[int] = None
if state.args.preload_voice:
for voice_key in state.args.preload_voice:
_LOGGER.debug("Preloading voice: %s", voice_key)
state.tts.preload_voice(voice_key)
# if args.max_thread_workers is not None:
# max_thread_workers = (
# None if args.max_thread_workers < 1 else args.max_thread_workers
# )
# elif args.raw_stream:
# # Faster time to first audio
# max_thread_workers = 2
state.result_queue = Queue(maxsize=args.result_queue_size)
# executor = ThreadPoolExecutor(max_workers=max_thread_workers)
# if os.isatty(sys.stdout.fileno()):
# if (not args.output_dir) and (not args.raw_stream):
# # No where else for the audio to go
# args.interactive = True
if args.raw_stream:
# Output in a separate thread to avoid blocking audio processing
state.raw_queue = Queue(maxsize=args.raw_stream_queue_size)
def output_raw_stream():
while True:
audio = state.raw_queue.get()
if audio is None:
break
_LOGGER.debug(
"Writing %s byte(s) of 16-bit 22050Hz mono PCM to stdout",
len(audio),
)
sys.stdout.buffer.write(audio)
sys.stdout.buffer.flush()
state.raw_stream_thread = threading.Thread(
target=output_raw_stream, daemon=True
)
state.raw_stream_thread.start()
state.result_thread = threading.Thread(
target=process_result, daemon=True, args=(state,)
)
state.result_thread.start()
def process_line(line_id: str, line: str, state: CommandLineInterfaceState):
from mimic3_tts import AudioResult, MarkResult
def process_result(state: CommandLineInterfaceState):
try:
from mimic3_tts import AudioResult, MarkResult
assert state.result_queue is not None
args = state.args
while True:
result_todo = state.result_queue.get()
if result_todo is None:
break
try:
result = result_todo.result
line = result_todo.line
line_id = result_todo.line_id
if isinstance(result, AudioResult):
if args.interactive or args.output_dir:
# Convert to WAV audio
wav_bytes: typing.Optional[bytes] = None
if args.interactive:
if args.stdout:
# Write audio to stdout
sys.stdout.buffer.write(result.audio_bytes)
sys.stdout.buffer.flush()
else:
# Play sound
if not wav_bytes:
wav_bytes = result.to_wav_bytes()
if wav_bytes:
play_wav_bytes(wav_bytes)
if args.output_dir:
if not wav_bytes:
wav_bytes = result.to_wav_bytes()
# Determine file name
if args.output_naming == OutputNaming.TEXT:
# Use text itself
file_name = line.strip().replace(" ", "_")
file_name = file_name.translate(
str.maketrans(
"", "", string.punctuation.replace("_", "")
)
)
elif args.output_naming == OutputNaming.TIME:
# Use timestamp
file_name = str(time.time())
elif args.output_naming == OutputNaming.ID:
file_name = line_id
assert file_name, f"No file name for text: {line}"
wav_path = args.output_dir / (file_name + ".wav")
wav_path.write_bytes(wav_bytes)
_LOGGER.debug("Wrote %s", wav_path)
else:
# Combine all audio and output to stdout at the end
state.all_audio += result.audio_bytes
state.sample_rate_hz = result.sample_rate_hz
state.sample_width_bytes = result.sample_width_bytes
state.num_channels = result.num_channels
elif isinstance(result, MarkResult):
if state.mark_writer:
print(result.name, file=state.mark_writer)
except Exception:
_LOGGER.exception("Error processing result")
except Exception:
_LOGGER.exception("process_result")
def process_line(
line: str,
state: CommandLineInterfaceState,
line_id: str = "",
):
from mimic3_tts import SSMLSpeaker
assert state.tts is not None
assert state.result_queue is not None
args = state.args
assert state.tts is not None
# TODO: SSML
state.tts.begin_utterance()
if args.ssml:
results = SSMLSpeaker(state.tts).speak(line)
else:
state.tts.begin_utterance()
# TODO: text language
state.tts.speak_text(line)
# TODO: text language
state.tts.speak_text(line)
# TODO: CSV
text_id = ""
result_idx = 0
results = state.tts.end_utterance()
for result in state.tts.end_utterance():
if isinstance(result, AudioResult):
if args.raw_stream:
assert state.raw_queue is not None
state.raw_queue.put(result.audio_bytes)
elif args.interactive or args.output_dir:
# Convert to WAV audio
wav_bytes: typing.Optional[bytes] = None
if args.interactive:
if not wav_bytes:
wav_bytes = result.to_wav_bytes()
play_wav_bytes(wav_bytes)
if args.output_dir:
if not wav_bytes:
wav_bytes = result.to_wav_bytes()
# Determine file name
if args.output_naming == OutputNaming.TEXT:
# Use text itself
file_name = line.strip().replace(" ", "_")
file_name = file_name.translate(
str.maketrans("", "", string.punctuation.replace("_", ""))
)
elif args.output_naming == OutputNaming.TIME:
# Use timestamp
file_name = str(time.time())
elif args.output_naming == OutputNaming.ID:
if not text_id:
text_id = line_id
else:
text_id = f"{line_id}_{result_idx + 1}"
file_name = text_id
assert file_name, f"No file name for text: {line}"
wav_path = args.output_dir / (file_name + ".wav")
wav_path.write_bytes(wav_bytes)
_LOGGER.debug("Wrote %s", wav_path)
else:
# Combine all audio and output to stdout at the end
state.all_audio += result.audio_bytes
state.sample_rate_hz = result.sample_rate_hz
state.sample_width_bytes = result.sample_width_bytes
state.num_channels = result.num_channels
result_idx += 1
elif isinstance(result, MarkResult):
if state.mark_writer:
print(result.name, file=state.mark_writer)
# text_id = ""
# for result_idx, result in enumerate(tts_results):
# text = result.text
# # Write before marks
# if result.marks_before and state.mark_writer:
# for mark_name in result.marks_before:
# print(mark_name, file=state.mark_writer)
# if args.raw_stream:
# assert raw_queue is not None
# raw_queue.put(result.audio.tobytes())
# elif args.interactive or args.output_dir:
# # Convert to WAV audio
# with io.BytesIO() as wav_io:
# wav_write(wav_io, result.sample_rate, result.audio)
# wav_data = wav_io.getvalue()
# assert wav_data is not None
# if args.interactive:
# # Play audio
# _LOGGER.debug("Playing audio with play command")
# try:
# subprocess.run(
# play_command,
# input=wav_data,
# stdout=subprocess.DEVNULL,
# stderr=subprocess.DEVNULL,
# check=True,
# )
# except FileNotFoundError:
# _LOGGER.error(
# "Unable to play audio with command '%s'. set with --play-command or redirect stdout",
# args.play_command,
# )
# with open("output.wav", "wb") as output_file:
# output_file.write(wav_data)
# _LOGGER.warning("stdout not redirected. Wrote audio to output.wav.")
# else:
# # Combine all audio and output to stdout at the end
# all_audios.append(result.audio)
# # Write after marks
# if result.marks_after and state.mark_writer:
# for mark_name in result.marks_after:
# print(mark_name, file=state.mark_writer)
for result in results:
state.result_queue.put(
ResultToProcess(
result=result,
line=line,
line_id=line_id,
)
)
def process_lines(state: CommandLineInterfaceState):
assert state.texts is not None
args = state.args
start_time_to_first_audio = time.perf_counter()
try:
result_idx = 0
for line in state.texts:
line_id = ""
line = line.strip()
@ -386,20 +352,22 @@ def process_lines(state: CommandLineInterfaceState):
# Line has the format id|text instead of just text
line_id, line = line.split(args.id_delimiter, maxsplit=1)
process_line(line_id, line, state)
process_line(line, state, line_id=line_id)
result_idx += 1
except KeyboardInterrupt:
if state.raw_queue is not None:
if state.result_queue is not None:
# Draw audio playback queue
while not state.raw_queue.empty():
state.raw_queue.get()
while not state.result_queue.empty():
state.result_queue.get()
finally:
# Wait for raw stream to finish
if state.raw_queue is not None:
state.raw_queue.put(None)
if state.result_queue is not None:
state.result_queue.put(None)
if state.raw_stream_thread is not None:
state.raw_stream_thread.join()
if state.result_thread is not None:
print("Waiting for audio to finish...", file=sys.stderr)
state.result_thread.join()
# -------------------------------------------------------------------------
@ -502,37 +470,24 @@ def get_args():
parser.add_argument(
"--noise-scale",
type=float,
default=0.333,
help="Noise scale (default: 0.333)",
help="Noise scale [0-1], default is 0.667",
)
parser.add_argument(
"--length-scale",
type=float,
default=1.0,
help="Length scale (default: 1.0)",
help="Length scale (1.0 is default speed, 0.5 is 2x faster)",
)
parser.add_argument(
"--noise-w",
type=float,
default=1.0,
help="Variation in cadence (default: 1.0)",
help="Variation in cadence [0-1], default is 0.8",
)
# Miscellaneous
parser.add_argument(
"--max-thread-workers",
type=int,
help="Maximum number of threads to concurrently load models and run sentences through TTS/Vocoder",
)
parser.add_argument(
"--raw-stream",
action="store_true",
help="Stream raw 16-bit 22050Hz mono PCM audio to stdout",
)
parser.add_argument(
"--raw-stream-queue-size",
"--result-queue-size",
default=5,
help="Maximum number of sentences to maintain in output queue with --raw-stream (default: 5)",
help="Maximum number of sentences to maintain in output queue (default: 5)",
)
parser.add_argument(
"--process-on-blank-line",
@ -552,6 +507,9 @@ def get_args():
action="store_true",
help="Force audio output to stdout even if a tty is detected",
)
parser.add_argument(
"--preload-voice", action="append", help="Preload voice when starting up"
)
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(
@ -575,116 +533,6 @@ def get_args():
# -------------------------------------------------------------------------
# # Directories to search for voices
# voices_dirs = get_voices_dirs(args.voices_dir)
# def list_voices_vocoders():
# """Print all vocoders and voices"""
# # (type, name) -> location
# local_info = {}
# # Search for downloaded voices/vocoders
# for voices_dir in voices_dirs:
# if not voices_dir.is_dir():
# continue
# for voice_dir in voices_dir.iterdir():
# if not voice_dir.is_dir():
# continue
# if voice_dir.name in VOCODER_DIR_NAMES:
# # Vocoder
# for vocoder_model_dir in voice_dir.iterdir():
# if not valid_voice_dir(vocoder_model_dir):
# continue
# full_vocoder_name = f"{voice_dir.name}-{vocoder_model_dir.name}"
# local_info[("vocoder", full_vocoder_name)] = str(
# vocoder_model_dir
# )
# else:
# # Voice
# voice_lang = voice_dir.name
# for voice_model_dir in voice_dir.iterdir():
# if not valid_voice_dir(voice_model_dir):
# continue
# local_info[("voice", voice_model_dir.name)] = str(
# voice_model_dir
# )
# # (type, lang, name, downloaded, aliases, location)
# voices_and_vocoders = []
# with open(_DIR / "VOCODERS", "r", encoding="utf-8") as vocoders_file:
# for line in vocoders_file:
# line = line.strip()
# if not line:
# continue
# *vocoder_aliases, full_vocoder_name = line.split()
# downloaded = False
# location = local_info.get(("vocoder", full_vocoder_name), "")
# if location:
# downloaded = True
# voices_and_vocoders.append(
# (
# "vocoder",
# " ",
# "*" if downloaded else " ",
# full_vocoder_name,
# ",".join(vocoder_aliases),
# location,
# )
# )
# with open(_DIR / "VOICES", "r", encoding="utf-8") as voices_file:
# for line in voices_file:
# line = line.strip()
# if not line:
# continue
# *voice_aliases, full_voice_name, download_name = line.split()
# voice_lang = download_name.split("_", maxsplit=1)[0]
# downloaded = False
# location = local_info.get(("voice", full_voice_name), "")
# if location:
# downloaded = True
# voices_and_vocoders.append(
# (
# "voice",
# voice_lang,
# "*" if downloaded else " ",
# full_voice_name,
# ",".join(voice_aliases),
# location,
# )
# )
# headers = ("TYPE", "LANG", "LOCAL", "NAME", "ALIASES", "LOCATION")
# # Get widths of columns
# col_widths = [0] * len(voices_and_vocoders[0])
# for item in voices_and_vocoders:
# for col in range(len(col_widths)):
# col_widths[col] = max(
# col_widths[col], len(item[col]) + 1, len(headers[col]) + 1
# )
# # Print results
# print(*(h.ljust(col_widths[col]) for col, h in enumerate(headers)))
# for item in sorted(voices_and_vocoders):
# print(*(v.ljust(col_widths[col]) for col, v in enumerate(item)))
# if args.list:
# list_voices_vocoders()
# sys.exit(0)
return args