import torch import auraloss from suno_utils.audio import Audio def calculate_stft_loss(audio1: Audio, audio2: Audio) -> float: mrstft = auraloss.freq.MultiResolutionSTFTLoss() tensor1 = torch.from_numpy(audio1.array_float).float().unsqueeze(0) tensor2 = torch.from_numpy(audio2.array_float).float().unsqueeze(0) min_length = min(tensor1.shape[-1], tensor2.shape[-1]) tensor1 = tensor1[..., :min_length] tensor2 = tensor2[..., :min_length] loss = mrstft(tensor1, tensor2) return loss.item() def calculate_mel_loss(audio1: Audio, audio2: Audio) -> float: assert audio1.sample_rate == audio2.sample_rate mel_loss = auraloss.freq.MultiResolutionSTFTLoss( fft_sizes=[1024, 2048, 8192], hop_sizes=[256, 512, 2048], win_lengths=[1024, 2048, 8192], scale="mel", n_bins=128, sample_rate=audio1.sample_rate, perceptual_weighting=True, ) tensor1 = torch.from_numpy(audio1.array_float).float().unsqueeze(0) tensor2 = torch.from_numpy(audio2.array_float).float().unsqueeze(0) min_length = min(tensor1.shape[-1], tensor2.shape[-1]) tensor1 = tensor1[..., :min_length] tensor2 = tensor2[..., :min_length] loss = mel_loss(tensor1, tensor2) return loss.item()