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)