Spaces:
Running on Zero
Running on Zero
File size: 4,684 Bytes
804ee23 | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 142 143 144 145 146 147 148 149 150 151 152 153 | from __future__ import annotations
import argparse
from pathlib import Path
def parse_args(argv=None):
parser = argparse.ArgumentParser(description="dots.tts inference CLI.")
template_choices = ("tts", "instruction_tts", "text_to_audio", "tts_interleave")
parser.add_argument(
"--model-name-or-path",
required=True,
help="Local pretrained directory or Hugging Face repo id",
)
parser.add_argument(
"--revision", default=None, help="Optional Hugging Face revision"
)
parser.add_argument(
"--cache-dir", default=None, help="Optional Hugging Face cache dir"
)
parser.add_argument("--text", type=str, required=True, help="Input text")
parser.add_argument("--output", default="output.wav", help="Output wav file path")
parser.add_argument(
"--precision", type=str, default="bfloat16", help="Inference precision"
)
parser.add_argument(
"--seed",
type=int,
default=42,
help="Random seed for inference.",
)
parser.add_argument(
"--prompt-audio", type=str, default=None, help="Path to prompt audio"
)
parser.add_argument(
"--prompt-text", type=str, default=None, help="Transcript of prompt audio"
)
parser.add_argument(
"--language",
type=str,
default=None,
help="Language tag mode. Default: none. Supported values: none, auto_detect, or a language code/name such as EN/en/english/chinese.",
)
parser.add_argument(
"--template-name",
choices=template_choices,
default=None,
help="Named template preset for generation.",
)
parser.add_argument(
"--ode-method", type=str, default="euler", help="ODE solver method"
)
parser.add_argument(
"--num-steps", type=int, default=10, help="Diffusion sampling steps"
)
parser.add_argument(
"--guidance-scale",
type=float,
default=1.2,
help="Classifier-free guidance scale",
)
parser.add_argument(
"--speaker-scale",
type=float,
default=1.5,
help="Scale applied to the reference speaker embedding",
)
parser.add_argument(
"--max-generate-length",
type=int,
default=500,
help="Maximum total audio patch count (prompt + generated)",
)
parser.add_argument(
"--normalize-text",
action="store_true",
help="Whether to normalize text before inference",
)
parser.add_argument(
"--profile-inference",
action="store_true",
help="Collect per-module inference timing statistics",
)
return parser.parse_args(argv)
def main(argv=None):
args = parse_args(argv)
import soundfile as sf
from loguru import logger
from dots_tts.runtime import DotsTtsRuntime
from dots_tts.utils.logging import configure_logging
from dots_tts.utils.util import seed_everything
configure_logging()
seed_everything(args.seed)
output_path = Path(args.output)
output_path.parent.mkdir(parents=True, exist_ok=True)
logger.info(
"CLI command started: model={} output={} seed={}",
args.model_name_or_path,
output_path,
args.seed,
)
try:
runtime = DotsTtsRuntime.from_pretrained(
args.model_name_or_path,
revision=args.revision,
cache_dir=args.cache_dir,
precision=args.precision,
max_generate_length=args.max_generate_length,
)
result = runtime.generate(
text=args.text,
prompt_audio_path=args.prompt_audio,
prompt_text=args.prompt_text,
language=args.language,
template_name=args.template_name,
ode_method=args.ode_method,
num_steps=args.num_steps,
guidance_scale=args.guidance_scale,
speaker_scale=args.speaker_scale,
normalize_text=args.normalize_text,
profile_inference=args.profile_inference,
)
sf.write(
output_path,
result["audio"].float().cpu().squeeze().numpy(),
result["sample_rate"],
)
except Exception:
logger.exception(
"CLI inference failed: model={} output={}",
args.model_name_or_path,
output_path,
)
raise
logger.info(
"CLI output written: request_id={} output={} sample_rate={} samples={}",
result["fid"],
output_path,
result["sample_rate"],
int(result["audio"].shape[-1]),
)
if __name__ == "__main__":
raise SystemExit(main())
|