返回 JoyAI-Echo
ops.py
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
77 lines PYTHON