| 162 | |
| 163 | |
| 164 | def parse_args() -> argparse.Namespace: |
| 165 | parser = argparse.ArgumentParser(description="Python reference Stable Audio warmbench.") |
| 166 | parser.add_argument("--family", default="stable_audio") |
| 167 | parser.add_argument("--model", default="models/stable-audio-3-small-music") |
| 168 | parser.add_argument("--reference-root", type=Path, default=REFERENCE_ROOT) |
| 169 | parser.add_argument("--model-root", type=Path, default=DEFAULT_MODEL_ROOT) |
| 170 | parser.add_argument("--backend", choices=("cuda", "cpu"), default="cuda") |
| 171 | parser.add_argument("--device", type=int, default=0) |
| 172 | parser.add_argument("--threads", type=int, default=8) |
| 173 | parser.add_argument("--warmup", type=int, default=0) |
| 174 | parser.add_argument("--iterations", type=int, default=1) |
| 175 | parser.add_argument("--case", choices=tuple(TEST_CASES), default="music_text") |
| 176 | parser.add_argument("--request-json", default="") |
| 177 | parser.add_argument("--request-sequence-json", default="") |
| 178 | parser.add_argument("--prompt", action="append", dest="prompts", default=[]) |
| 179 | parser.add_argument("--negative-prompt", default="poor quality, distorted, clipped, noisy") |
| 180 | parser.add_argument("--duration", type=float, default=120.0) |
| 181 | parser.add_argument("--steps", type=int, default=8) |
| 182 | parser.add_argument("--cfg-scale", type=float, default=1.0) |
| 183 | parser.add_argument("--seed", type=int, default=1234) |
| 184 | parser.add_argument("--model-precision", choices=("native", "fp32", "fp16"), default="native") |
| 185 | parser.add_argument("--chunked-decode", choices=("default", "on", "off"), default="default") |
| 186 | parser.add_argument("--output-dir", type=Path, default=None) |
| 187 | parser.add_argument("--audio-out", type=Path, default=Path("stable_audio_python_audio.wav")) |
| 188 | parser.add_argument("--timing-file", type=Path, default=Path("stable_audio_python_timing.log")) |
| 189 | parser.add_argument("--summary-file", type=Path, default=None) |
| 190 | return parser.parse_args() |
| 191 | |
| 192 | |
| 193 | def resolve_repo_path(path: Path | str) -> Path: |