(self, batch, batch_idx)
| 775 | return lyrics |
| 776 | |
| 777 | def plot_step(self, batch, batch_idx): |
| 778 | global_step = self.global_step |
| 779 | if ( |
| 780 | global_step % self.hparams.every_plot_step != 0 |
| 781 | or self.local_rank != 0 |
| 782 | or torch.distributed.get_rank() != 0 |
| 783 | or torch.cuda.current_device() != 0 |
| 784 | ): |
| 785 | return |
| 786 | results = self.predict_step(batch) |
| 787 | |
| 788 | target_wavs = results["target_wavs"] |
| 789 | pred_wavs = results["pred_wavs"] |
| 790 | keys = results["keys"] |
| 791 | prompts = results["prompts"] |
| 792 | candidate_lyric_chunks = results["candidate_lyric_chunks"] |
| 793 | sr = results["sr"] |
| 794 | seeds = results["seeds"] |
| 795 | i = 0 |
| 796 | for key, target_wav, pred_wav, prompt, candidate_lyric_chunk, seed in zip( |
| 797 | keys, target_wavs, pred_wavs, prompts, candidate_lyric_chunks, seeds |
| 798 | ): |
| 799 | key = key |
| 800 | prompt = prompt |
| 801 | lyric = self.construct_lyrics(candidate_lyric_chunk) |
| 802 | key_prompt_lyric = f"# KEY\n\n{key}\n\n\n# PROMPT\n\n{prompt}\n\n\n# LYRIC\n\n{lyric}\n\n# SEED\n\n{seed}\n\n" |
| 803 | log_dir = self.logger.log_dir |
| 804 | save_dir = f"{log_dir}/eval_results/step_{self.global_step}" |
| 805 | if not os.path.exists(save_dir): |
| 806 | os.makedirs(save_dir, exist_ok=True) |
| 807 | torchaudio.save( |
| 808 | f"{save_dir}/target_wav_{key}_{i}.wav", target_wav.float().cpu(), sr |
| 809 | ) |
| 810 | torchaudio.save( |
| 811 | f"{save_dir}/pred_wav_{key}_{i}.wav", pred_wav.float().cpu(), sr |
| 812 | ) |
| 813 | with open( |
| 814 | f"{save_dir}/key_prompt_lyric_{key}_{i}.txt", "w", encoding="utf-8" |
| 815 | ) as f: |
| 816 | f.write(key_prompt_lyric) |
| 817 | i += 1 |
| 818 | |
| 819 | |
| 820 | def main(args): |
no test coverage detected