Update codeclm.py
This commit is contained in:
@@ -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
|
||||
|
||||
Reference in New Issue
Block a user