CoolFace
Apppublic

szili2011/audioEditingUNLIMITED

sourceHugging Facecc-by-sa-4.0updated 1y agoView on Hugging Face
0likes
utils.py107 linesDownload Raw Back to root
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