import torch
import torchaudio
from torch import nn
from librosa import filters
from nnAudio.features import CQT1992v2, CQT2010v2


def get_chroma_filterbank(sample_rate, n_fft):
    return torch.tensor(filters.chroma(sr=sample_rate, n_fft=n_fft).astype("float32"))


class STFT(nn.Module):
    def __init__(
        self,
        n_fft=2048,
        hop_length=240,
        is_db=False,
    ):
        super(STFT, self).__init__()

        # short-time Fourier transform
        self.stft = torchaudio.transforms.Spectrogram(
            n_fft=n_fft, hop_length=hop_length
        )

        # amplitude to decibel
        self.is_db = is_db
        if is_db:
            self.amplitude_to_db = torchaudio.transforms.AmplitudeToDB()

    def forward(self, waveform):
        if self.is_db:
            return self.amplitude_to_db(self.stft(waveform))
        else:
            return self.stft(waveform)


class MultiResSTFT(nn.Module):
    def __init__(
        self,
        n_ffts=[256, 512, 1024, 2048, 4096],
        hop_length=240,
        is_db=False,
    ):
        super(MultiResSTFT, self).__init__()

        # multi-resolution short-time Fourier transform
        self.n_ffts = n_ffts
        for n_fft in n_ffts:
            stft = torchaudio.transforms.Spectrogram(n_fft=n_fft, hop_length=hop_length)
            setattr(self, "stft_%d" % n_fft, stft)

        # amplitude to decibel
        self.is_db = is_db
        if is_db:
            self.amplitude_to_db = torchaudio.transforms.AmplitudeToDB()

    def forward(self, waveform):
        specs = []
        for n_fft in self.n_ffts:
            stft = getattr(self, "stft_%d" % n_fft)
            spec = stft(waveform)
            if self.is_db:
                spec = self.amplitude_to_db(spec)
            specs.append(spec)
        return specs


class MelSTFT(nn.Module):
    def __init__(
        self,
        sample_rate=24000,
        n_fft=2048,
        hop_length=240,
        n_mels=128,
        is_db=False,
    ):
        super(MelSTFT, self).__init__()

        # short-time Fourier transform with mel filterbank
        self.mel_stft = torchaudio.transforms.MelSpectrogram(
            sample_rate=sample_rate,
            n_fft=n_fft,
            hop_length=hop_length,
            n_mels=n_mels,
        )

        # amplitude to decibel
        self.is_db = is_db
        if is_db:
            self.amplitude_to_db = torchaudio.transforms.AmplitudeToDB()

    def forward(self, waveform):
        if self.is_db:
            return self.amplitude_to_db(self.mel_stft(waveform))
        else:
            return self.mel_stft(waveform)


class MFCC(nn.Module):
    def __init__(
        self,
        sample_rate=24000,
        n_fft=2048,
        hop_length=240,
        n_mels=128,
        n_mfcc=13,
    ):
        super(MFCC, self).__init__()

        # MFCC
        self.mfcc = torchaudio.transforms.MFCC(
            sample_rate=sample_rate,
            n_mfcc=n_mfcc,
            melkwargs={
                "n_fft": n_fft,
                "hop_length": hop_length,
                "n_mels": n_mels,
            },
        )

    def forward(self, waveform):
        return self.mfcc(waveform)[:, 1:, :]  # first channel is energy


class Chromagram(nn.Module):
    def __init__(
        self,
        sample_rate=24000,
        n_fft=2048,
        hop_length=240,
    ):
        super(Chromagram, self).__init__()

        # short-time Fourier transform
        self.stft = torchaudio.transforms.Spectrogram(
            n_fft=n_fft,
            hop_length=hop_length,
        )

        # chroma filterbank
        self.register_buffer("chroma_fb", get_chroma_filterbank(sample_rate, n_fft))

    def forward(self, waveform):
        spec = self.stft(waveform)
        chromagram = torch.matmul(spec.transpose(1, 2), self.chroma_fb.T).transpose(
            1, 2
        )
        return chromagram


class CQT(nn.Module):
    def __init__(
        self, sample_rate=24000, hop_length=240, algorithm="1992", is_db=False
    ):
        super().__init__()
        assert algorithm in ["1992", "2010"]

        # constant-Q transform
        if algorithm == "1992":
            self.cqt = CQT1992v2(sr=sample_rate, hop_length=hop_length)
        elif algorithm == "2010":
            self.cqt = CQT2010v2(sr=sample_rate, hop_length=hop_length)

        # amplitude to decibel
        self.is_db = is_db
        if is_db:
            self.amplitude_to_db = torchaudio.transforms.AmplitudeToDB()

    def forward(self, waveform):
        if self.is_db:
            return self.amplitude_to_db(self.cqt(waveform))
        else:
            return self.cqt(waveform)


"""
    TODO: Band-split spec
"""
