Perform TTS inference and save the generated audio.
(args)
| 62 | |
| 63 | |
| 64 | def run_tts(args): |
| 65 | """Perform TTS inference and save the generated audio.""" |
| 66 | logging.info(f"Using model from: {args.model_dir}") |
| 67 | logging.info(f"Saving audio to: {args.save_dir}") |
| 68 | |
| 69 | # Ensure the save directory exists |
| 70 | os.makedirs(args.save_dir, exist_ok=True) |
| 71 | |
| 72 | # Convert device argument to torch.device |
| 73 | if platform.system() == "Darwin" and torch.backends.mps.is_available(): |
| 74 | # macOS with MPS support (Apple Silicon) |
| 75 | device = torch.device(f"mps:{args.device}") |
| 76 | logging.info(f"Using MPS device: {device}") |
| 77 | elif torch.cuda.is_available(): |
| 78 | # System with CUDA support |
| 79 | device = torch.device(f"cuda:{args.device}") |
| 80 | logging.info(f"Using CUDA device: {device}") |
| 81 | else: |
| 82 | # Fall back to CPU |
| 83 | device = torch.device("cpu") |
| 84 | logging.info("GPU acceleration not available, using CPU") |
| 85 | |
| 86 | # Initialize the model |
| 87 | model = SparkTTS(args.model_dir, device) |
| 88 | |
| 89 | # Generate unique filename using timestamp |
| 90 | timestamp = datetime.now().strftime("%Y%m%d%H%M%S") |
| 91 | save_path = os.path.join(args.save_dir, f"{timestamp}.wav") |
| 92 | |
| 93 | logging.info("Starting inference...") |
| 94 | |
| 95 | # Perform inference and save the output audio |
| 96 | with torch.no_grad(): |
| 97 | wav = model.inference( |
| 98 | args.text, |
| 99 | args.prompt_speech_path, |
| 100 | prompt_text=args.prompt_text, |
| 101 | gender=args.gender, |
| 102 | pitch=args.pitch, |
| 103 | speed=args.speed, |
| 104 | ) |
| 105 | sf.write(save_path, wav, samplerate=16000) |
| 106 | |
| 107 | logging.info(f"Audio saved at: {save_path}") |
| 108 | |
| 109 | |
| 110 | if __name__ == "__main__": |
no test coverage detected