(args: argparse.Namespace)
| 138 | |
| 139 | |
| 140 | def select_device(args: argparse.Namespace) -> str: |
| 141 | torch.set_num_threads(max(1, args.threads)) |
| 142 | if args.backend == "cuda": |
| 143 | if not torch.cuda.is_available(): |
| 144 | raise RuntimeError("VibeVoice warmbench requested CUDA, but torch.cuda.is_available() is false") |
| 145 | torch.cuda.set_device(args.device) |
| 146 | return f"cuda:{args.device}" |
| 147 | return "cpu" |
| 148 | |
| 149 | |
| 150 | def seed_all(seed: int, backend: str) -> None: |
no test coverage detected