szili2011/audioEditingUNLIMITED
0
1import numpy as np2import torch3from typing import Optional, List, Tuple, NamedTuple, Union4from models import PipelineWrapper5import torchaudio6from audioldm.utils import get_duration7 8MAX_DURATION = None9 10 11class PromptEmbeddings(NamedTuple):12 embedding_hidden_states: torch.Tensor13 embedding_class_lables: torch.Tensor14 boolean_prompt_mask: torch.Tensor15 16 17def load_audio(audio_path: Union[str, np.array], fn_STFT, left: int = 0, right: int = 0,18 device: Optional[torch.device] = None,19 return_wav: bool = False, stft: bool = False, model_sr: Optional[int] = None) -> torch.Tensor:20 if stft: # AudioLDM/tango loading to spectrogram21 if type(audio_path) is str:22 import audioldm23 import audioldm.audio24 25 duration = get_duration(audio_path)26 if MAX_DURATION is not None:27 duration = min(duration, MAX_DURATION)28 29 mel, _, wav = audioldm.audio.wav_to_fbank(audio_path, target_length=int(duration * 102.4), fn_STFT=fn_STFT)30 mel = mel.unsqueeze(0)31 else:32 mel = audio_path33 34 c, h, w = mel.shape35 left = min(left, w-1)36 right = min(right, w - left - 1)37 mel = mel[:, :, left:w-right]38 mel = mel.unsqueeze(0).to(device)39 40 if return_wav:41 return mel, 16000, duration, wav42 43 return mel, model_sr, duration44 else:45 waveform, sr = torchaudio.load(audio_path)46 if sr != model_sr:47 waveform = torchaudio.functional.resample(waveform, orig_freq=sr, new_freq=model_sr)48 # waveform = waveform.numpy()[0, ...]49 50 def normalize_wav(waveform):51 waveform = waveform - torch.mean(waveform)52 waveform = waveform / (torch.max(torch.abs(waveform)) + 1e-8)53 return waveform * 0.554 55 waveform = normalize_wav(waveform)56 # waveform = waveform[None, ...]57 # waveform = pad_wav(waveform, segment_length)58 59 # waveform = waveform[0, ...]60 waveform = torch.FloatTensor(waveform)61 if MAX_DURATION is not None:62 duration = min(waveform.shape[-1] / model_sr, MAX_DURATION)63 waveform = waveform[:, :int(duration * model_sr)]64 65 # cut waveform66 duration = waveform.shape[-1] / model_sr67 return waveform, model_sr, duration68 69 70def get_height_of_spectrogram(length: int, ldm_stable: PipelineWrapper) -> int:71 vocoder_upsample_factor = np.prod(ldm_stable.model.vocoder.config.upsample_rates) / \72 ldm_stable.model.vocoder.config.sampling_rate73 74 if length is None:75 length = ldm_stable.model.unet.config.sample_size * ldm_stable.model.vae_scale_factor * \76 vocoder_upsample_factor77 78 height = int(length / vocoder_upsample_factor)79 80 # original_waveform_length = int(length * ldm_stable.model.vocoder.config.sampling_rate)81 if height % ldm_stable.model.vae_scale_factor != 0:82 height = int(np.ceil(height / ldm_stable.model.vae_scale_factor)) * ldm_stable.model.vae_scale_factor83 print(84 f"Audio length in seconds {length} is increased to {height * vocoder_upsample_factor} "85 f"so that it can be handled by the model. It will be cut to {length} after the "86 f"denoising process."87 )88 89 return height90 91 92def get_text_embeddings(target_prompt: List[str], target_neg_prompt: List[str], ldm_stable: PipelineWrapper93 ) -> Tuple[torch.Tensor, PromptEmbeddings, PromptEmbeddings]:94 text_embeddings_hidden_states, text_embeddings_class_labels, text_embeddings_boolean_prompt_mask = \95 ldm_stable.encode_text(target_prompt)96 uncond_embedding_hidden_states, uncond_embedding_class_lables, uncond_boolean_prompt_mask = \97 ldm_stable.encode_text(target_neg_prompt)98 99 text_emb = PromptEmbeddings(embedding_hidden_states=text_embeddings_hidden_states,100 boolean_prompt_mask=text_embeddings_boolean_prompt_mask,101 embedding_class_lables=text_embeddings_class_labels)102 uncond_emb = PromptEmbeddings(embedding_hidden_states=uncond_embedding_hidden_states,103 boolean_prompt_mask=uncond_boolean_prompt_mask,104 embedding_class_lables=uncond_embedding_class_lables)105 106 return text_embeddings_class_labels, text_emb, uncond_emb107 