MCPcopy Create free account
hub / github.com/ace-step/ACE-Step / plot_step

Method plot_step

trainer.py:777–817  ·  view source on GitHub ↗
(self, batch, batch_idx)

Source from the content-addressed store, hash-verified

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
820def main(args):

Callers 1

run_stepMethod · 0.95

Calls 2

predict_stepMethod · 0.95
construct_lyricsMethod · 0.95

Tested by

no test coverage detected