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())