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()
