diff --git a/SongGeneration/codeclm/models/codeclm.py b/SongGeneration/codeclm/models/codeclm.py index 133d6b2..b8a364f 100644 --- a/SongGeneration/codeclm/models/codeclm.py +++ b/SongGeneration/codeclm/models/codeclm.py @@ -107,6 +107,7 @@ class CodecLM: vocal_wavs: torch.Tensor = None, bgm_wavs: torch.Tensor = None, return_tokens: bool = False, + gen_type: str = "left blank", ) -> tp.Union[torch.Tensor, tp.Tuple[torch.Tensor, torch.Tensor]]: """Generate samples conditioned on text and melody. @@ -148,7 +149,7 @@ class CodecLM: if return_tokens: return tokens else: - out = self.generate_audio(tokens) + out = self.generate_audio(tokens,gen_type=gen_type) return out @@ -271,13 +272,19 @@ class CodecLM: return gen_tokens @torch.no_grad() - def generate_audio(self, gen_tokens: torch.Tensor, prompt=None, vocal_prompt=None, bgm_prompt=None, chunked=False): + def generate_audio(self, gen_tokens: torch.Tensor, prompt=None, vocal_prompt=None, bgm_prompt=None, chunked=False,gen_type="left blank"): """Generate Audio from tokens""" assert gen_tokens.dim() == 3 if self.seperate_tokenizer is not None: gen_tokens_song = gen_tokens[:, [0], :] gen_tokens_vocal = gen_tokens[:, [1], :] gen_tokens_bgm = gen_tokens[:, [2], :] + if gen_type == "bgm": + gen_tokens_vocal = torch.full_like(gen_tokens_vocal, 3142) + vocal_prompt = None + elif gen_type == "vocal": + gen_tokens_bgm = torch.full_like(gen_tokens_bgm, 9670) + bgm_prompt = None # gen_audio_song = self.audiotokenizer.decode(gen_tokens_song, prompt) gen_audio_seperate = self.seperate_tokenizer.decode([gen_tokens_vocal, gen_tokens_bgm], vocal_prompt, bgm_prompt, chunked=chunked) return gen_audio_seperate