import sys
sys.path.append("/home/minz/glockenspiel/musicfm-training")
import json
import numpy as np
from torch.utils import data
from musicfm.data_loaders.mertlong import MERTDataset
from musicfm.data_loaders.msd import MSDDataset
from musicfm.data_loaders.fma import FMADataset
from musicfm.modules.features import STFT, MelSTFT, MFCC, Chromagram, CQT

class Preprocessor:
    def __init__(
        self,
        batch_size=16,
        num_workers=4,
        hop_length=240,
        dataset="mertlong",
        total_sample=0,
        ):
        super(Preprocessor, self).__init__()

        features = ["spec", "melspec", "cqt", "mfcc", "chromagram"]
        self.features = []
        self.n_ffts = [256, 512, 1024, 2048, 4096]
        self.stats = {}
        for feature in features:
            if feature != "cqt":
                for n_fft in self.n_ffts:
                    if feature == "spec":
                        setattr(self, "%s_%d" % (feature, n_fft), STFT(n_fft=n_fft, is_db=True))
                    elif feature == "melspec":
                        setattr(self, "%s_%d" % (feature, n_fft), MelSTFT(n_fft=n_fft, is_db=True))
                    elif feature == "mfcc":
                        setattr(self, "%s_%d" % (feature, n_fft), MFCC(n_fft=n_fft))
                    elif feature == "chromagram":
                        setattr(self, "%s_%d" % (feature, n_fft), Chromagram(n_fft=n_fft))
                    self.stats["%s_%d_cnt" % (feature, n_fft)] = 0
                    self.stats["%s_%d_mean" % (feature, n_fft)] = 0.0
                    self.stats["%s_%d_std" % (feature, n_fft)] = 0.0
                    self.features.append("%s_%d" % (feature, n_fft))
            else:
                setattr(self, feature, CQT())
                self.stats["%s_cnt" % feature] = 0
                self.stats["%s_mean" % feature] = 0.0
                self.stats["%s_std" % feature] = 0.0
                self.features.append("cqt")

        self.dataset = dataset
        if dataset == "mertlong":
            train_dataset = MERTDataset()
        elif dataset == "msd":
            train_dataset = MSDDataset()
        elif dataset == "fma":
            train_dataset = FMADataset()
        else:
            print("%s dataset is not supported yet." % dataset)
        self.loader = data.DataLoader(dataset=train_dataset, batch_size=batch_size, shuffle=True, drop_last=False, num_workers=num_workers)

    def transform(self, x):
        out = {}
        for feature in self.features:
            process = getattr(self, feature)
            out[feature] = process(x)
        return out

    def update(self, feature_name, feature):
        # count total samples
        num_samples = len(feature.flatten())
        self.stats["%s_cnt" % feature_name] += num_samples

        # update mean
        new_mean = feature.mean().numpy()
        delta = new_mean - self.stats["%s_mean" % feature_name]
        self.stats["%s_mean" % feature_name] += delta * num_samples / self.stats["%s_cnt" % feature_name]

        # update std
        new_std = ((feature - new_mean) * (feature - self.stats["%s_mean" % feature_name])).sum().numpy()
        self.stats["%s_std" % feature_name] += new_std

    def save_features(self):
        stat = {k: v for k, v in self.stats.items()}
        for k in stat.keys():
            if k[-3:] == "std":
                stat[k] = np.sqrt(stat[k] / (stat[k[:-3] + "cnt"] - 1))
        print(stat)
        with open("/app/suno/minz/models/%s_stats.json" % self.dataset, "w") as file:
            json.dump(stat, file)
        
    def iterate(self, num_iter):
        iter_dl = iter(self.loader)
        for i in range(num_iter):
            try:
                inp = next(iter_dl)
            except StopIteration:
                print("end of an epoch")
                iter_dl = iter(self.loader)
                inp = next(iter_dl)
            features = self.transform(inp)
            for key in features.keys():
                self.update(key, features[key])
            print("iter: %d" % i)
            if i % 10 == 0:
                self.save_features()


if __name__ == "__main__":
    batch_size = int(sys.argv[1])
    num_workers = int(sys.argv[2])
    num_iter = int(sys.argv[3])
    dataset = sys.argv[4]
    hop_length = 240
    p = Preprocessor(batch_size, num_workers, hop_length, dataset)
    p.iterate(num_iter)
