from suno_utils.audio import Audio import numpy as np from typing import List import typing import torch.nn as nn from audiotools import AudioSignal, STFTParams class MultiScaleSTFTLoss(nn.Module): """Computes the multi-scale STFT loss from [1]. Parameters ---------- window_lengths : List[int], optional Length of each window of each STFT, by default [2048, 512] loss_fn : typing.Callable, optional How to compare each loss, by default nn.L1Loss() clamp_eps : float, optional Clamp on the log magnitude, below, by default 1e-5 mag_weight : float, optional Weight of raw magnitude portion of loss, by default 1.0 log_weight : float, optional Weight of log magnitude portion of loss, by default 1.0 pow : float, optional Power to raise magnitude to before taking log, by default 2.0 weight : float, optional Weight of this loss, by default 1.0 match_stride : bool, optional Whether to match the stride of convolutional layers, by default False References ---------- 1. Engel, Jesse, Chenjie Gu, and Adam Roberts. "DDSP: Differentiable Digital Signal Processing." International Conference on Learning Representations. 2019. Implementation copied from: https://github.com/descriptinc/lyrebird-audiotools/blob/961786aa1a9d628cca0c0486e5885a457fe70c1a/audiotools/metrics/spectral.py """ def __init__( self, window_lengths: List[int] = [2048, 512], loss_fn: typing.Callable = nn.L1Loss(), clamp_eps: float = 1e-5, mag_weight: float = 1.0, log_weight: float = 1.0, pow: float = 2.0, weight: float = 1.0, match_stride: bool = False, window_type: str = None, ): super().__init__() self.stft_params = [ STFTParams( window_length=w, hop_length=w // 4, match_stride=match_stride, window_type=window_type, ) for w in window_lengths ] self.loss_fn = loss_fn self.log_weight = log_weight self.mag_weight = mag_weight self.clamp_eps = clamp_eps self.weight = weight self.pow = pow def forward(self, x: AudioSignal, y: AudioSignal): """Computes multi-scale STFT between an estimate and a reference signal. Parameters ---------- x : AudioSignal Estimate signal y : AudioSignal Reference signal Returns ------- torch.Tensor Multi-scale STFT loss. """ loss = 0.0 for s in self.stft_params: x.stft(s.window_length, s.hop_length, s.window_type) y.stft(s.window_length, s.hop_length, s.window_type) loss += self.log_weight * self.loss_fn( x.magnitude.clamp(self.clamp_eps).pow(self.pow).log10(), y.magnitude.clamp(self.clamp_eps).pow(self.pow).log10(), ) loss += self.mag_weight * self.loss_fn(x.magnitude, y.magnitude) return loss class MelSpectrogramLoss(nn.Module): """Compute distance between mel spectrograms. Can be used in a multi-scale way. Parameters ---------- n_mels : List[int] Number of mels per STFT, by default [150, 80], window_lengths : List[int], optional Length of each window of each STFT, by default [2048, 512] loss_fn : typing.Callable, optional How to compare each loss, by default nn.L1Loss() clamp_eps : float, optional Clamp on the log magnitude, below, by default 1e-5 mag_weight : float, optional Weight of raw magnitude portion of loss, by default 1.0 log_weight : float, optional Weight of log magnitude portion of loss, by default 1.0 pow : float, optional Power to raise magnitude to before taking log, by default 2.0 weight : float, optional Weight of this loss, by default 1.0 match_stride : bool, optional Whether to match the stride of convolutional layers, by default False Implementation copied from: https://github.com/descriptinc/lyrebird-audiotools/blob/961786aa1a9d628cca0c0486e5885a457fe70c1a/audiotools/metrics/spectral.py """ def __init__( self, n_mels: List[int] = [150, 80], window_lengths: List[int] = [2048, 512], loss_fn: typing.Callable = nn.L1Loss(), clamp_eps: float = 1e-5, mag_weight: float = 1.0, log_weight: float = 1.0, pow: float = 2.0, weight: float = 1.0, match_stride: bool = False, mel_fmin: List[float] = [0.0, 0.0], mel_fmax: List[float] = [None, None], window_type: str = None, ): super().__init__() self.stft_params = [ STFTParams( window_length=w, hop_length=w // 4, match_stride=match_stride, window_type=window_type, ) for w in window_lengths ] self.n_mels = n_mels self.loss_fn = loss_fn self.clamp_eps = clamp_eps self.log_weight = log_weight self.mag_weight = mag_weight self.weight = weight self.mel_fmin = mel_fmin self.mel_fmax = mel_fmax self.pow = pow def forward(self, x: AudioSignal, y: AudioSignal): """Computes mel loss between an estimate and a reference signal. Parameters ---------- x : AudioSignal Estimate signal y : AudioSignal Reference signal Returns ------- torch.Tensor Mel loss. """ loss = 0.0 for n_mels, fmin, fmax, s in zip(self.n_mels, self.mel_fmin, self.mel_fmax, self.stft_params): kwargs = { "window_length": s.window_length, "hop_length": s.hop_length, "window_type": s.window_type, } x_mels = x.mel_spectrogram(n_mels, mel_fmin=fmin, mel_fmax=fmax, **kwargs) y_mels = y.mel_spectrogram(n_mels, mel_fmin=fmin, mel_fmax=fmax, **kwargs) loss += self.log_weight * self.loss_fn( x_mels.clamp(self.clamp_eps).pow(self.pow).log10(), y_mels.clamp(self.clamp_eps).pow(self.pow).log10(), ) loss += self.mag_weight * self.loss_fn(x_mels, y_mels) return loss def standardize_audios(a: Audio, b: Audio, sr=48000) -> typing.Tuple[AudioSignal, AudioSignal]: """Standardize the length of two audios by resampling, cropping, mono""" a_sig = AudioSignal(a.array_float, sample_rate=a.sample_rate).to_mono() b_sig = AudioSignal(b.array_float, sample_rate=b.sample_rate).to_mono() a_sig = a_sig.resample(sr) b_sig = b_sig.resample(sr) # crop to same length min_len = min(a_sig.signal_length, b_sig.signal_length) a_sig = a_sig.truncate_samples(min_len) b_sig = b_sig.truncate_samples(min_len) return a_sig, b_sig def l1_distance(a: Audio, b: Audio) -> float: a_wav = a.array_float b_wav = b.array_float return np.mean(np.abs(a_wav - b_wav)) def stft_distance(a: Audio, b: Audio, sr=48000) -> float: stft_loss = MultiScaleSTFTLoss() a_sig, b_sig = standardize_audios(a, b, sr=sr) return stft_loss(a_sig, b_sig).item() def mel_distance(a: Audio, b: Audio) -> float: mel_loss = MelSpectrogramLoss() a_sig, b_sig = standardize_audios(a, b) return mel_loss(a_sig, b_sig).item() if __name__ == "__main__": a = Audio.from_beep(48000, n_channels=2, duration_s=2, freq_khz=0.1) assert l1_distance(a, a) == 0 assert stft_distance(a, a) == 0 b = Audio.from_beep(48000, n_channels=2, duration_s=2, freq_khz=0.11) print(l1_distance(a, b)) print(stft_distance(a, b)) print(mel_distance(a, b))