| 1 | import torch |
| 2 | import torchaudio |
| 3 | from torch import nn |
| 4 | |
| 5 | from ltx_core.types import Audio |
| 6 | |
| 7 | |
| 8 | class AudioProcessor(nn.Module): |
| 9 | """Converts audio waveforms to log-mel spectrograms with optional resampling.""" |
| 10 | |
| 11 | def __init__( |
| 12 | self, |
| 13 | target_sample_rate: int, |
| 14 | mel_bins: int, |
| 15 | mel_hop_length: int, |
| 16 | n_fft: int, |
| 17 | ) -> None: |
| 18 | super().__init__() |
| 19 | self.target_sample_rate = target_sample_rate |
| 20 | self.mel_bins = mel_bins |
| 21 | self.mel_hop_length = mel_hop_length |
| 22 | self.n_fft = n_fft |
| 23 | self.mel_transform = torchaudio.transforms.MelSpectrogram( |
| 24 | sample_rate=target_sample_rate, |
| 25 | n_fft=n_fft, |
| 26 | win_length=n_fft, |
| 27 | hop_length=mel_hop_length, |
| 28 | f_min=0.0, |
| 29 | f_max=target_sample_rate / 2.0, |
| 30 | n_mels=mel_bins, |
| 31 | window_fn=torch.hann_window, |
| 32 | center=True, |
| 33 | pad_mode="reflect", |
| 34 | power=1.0, |
| 35 | mel_scale="slaney", |
| 36 | norm="slaney", |
| 37 | ) |
| 38 | |
| 39 | def resample_audio(self, audio: Audio) -> Audio: |
| 40 | """Resample audio to the processor's target sample rate if needed.""" |
| 41 | if audio.sampling_rate == self.target_sample_rate: |
| 42 | return audio |
| 43 | resampled = torchaudio.functional.resample(audio.waveform, audio.sampling_rate, self.target_sample_rate) |
| 44 | resampled = resampled.to(device=audio.waveform.device, dtype=audio.waveform.dtype) |
| 45 | return Audio(waveform=resampled, sampling_rate=self.target_sample_rate) |
| 46 | |
| 47 | def waveform_to_mel( |
| 48 | self, |
| 49 | audio: Audio, |
| 50 | ) -> torch.Tensor: |
| 51 | """Convert waveform to log-mel spectrogram [batch, channels, time, n_mels].""" |
| 52 | waveform = self.resample_audio(audio).waveform |
| 53 | |
| 54 | mel = self.mel_transform(waveform) |
| 55 | mel = torch.log(torch.clamp(mel, min=1e-5)) |
| 56 | |
| 57 | mel = mel.to(device=waveform.device, dtype=waveform.dtype) |
| 58 | return mel.permute(0, 1, 3, 2).contiguous() |
| 59 | |
| 60 | |
| 61 | class PerChannelStatistics(nn.Module): |
| 62 | """ |
| 63 | Per-channel statistics for normalizing and denormalizing the latent representation. |
| 64 | This statics is computed over the entire dataset and stored in model's checkpoint under AudioVAE state_dict. |
| 65 | """ |
| 66 | |
| 67 | def __init__(self, latent_channels: int = 128) -> None: |
| 68 | super().__init__() |
| 69 | self.register_buffer("std-of-means", torch.empty(latent_channels)) |
| 70 | self.register_buffer("mean-of-means", torch.empty(latent_channels)) |
| 71 | |
| 72 | def un_normalize(self, x: torch.Tensor) -> torch.Tensor: |
| 73 | return (x * self.get_buffer("std-of-means").to(x)) + self.get_buffer("mean-of-means").to(x) |
| 74 | |
| 75 | def normalize(self, x: torch.Tensor) -> torch.Tensor: |
| 76 | return (x - self.get_buffer("mean-of-means").to(x)) / self.get_buffer("std-of-means").to(x) |
| 77 |