import os

os.system("pip install pyunpack")
os.system("pip install python_speech_features")
os.system("pip install patool")

from pyunpack import Archive
import shutil
from sklearn.preprocessing import MinMaxScaler
import numpy as np
import librosa
from python_speech_features import mfcc
import random
from sklearn.preprocessing import LabelEncoder
from keras.utils import np_utils
from sklearn.model_selection import train_test_split

AUDIO_PATH='/kaggle/working/train/train/audio'


def unzipp_files():
    if not os.path.exists('/kaggle/working/train/'):
        os.makedirs('/kaggle/working/train/')
    Archive('/kaggle/input/tensorflow-speech-recognition-challenge/train.7z').extractall('/kaggle/working/train/')

def get_classes():
    words=os.listdir(AUDIO_PATH)
    classes = ["no", "up", "down", "left", "right", "on", "off", "stop", "go", "yes"]
    unknown_list = [a for a in words if a not in classes and a!='_background_noise_']
    return classes, unknown_list

def load_data(classes, unknown_list, task="mfcc"):
    scaler = MinMaxScaler(feature_range=(0,1))
    all_wave = []
    all_label = []
    lengths = [] #list of numbers of observations in a class
    for label in classes:
        i=0
        print(label)
        waves = [f for f in os.listdir(AUDIO_PATH + '/'+ label) if f.endswith('.wav')]
        for wav in waves:
            samples, sample_rate = librosa.load(AUDIO_PATH + '/' + label + '/' + wav, sr = 8000)
            if(len(samples)== 8000) :
                if task=="mfcc":
                    mfcc_feat = mfcc(samples, sample_rate)
                    scaler = scaler.fit(mfcc_feat)
                    samples = scaler.transform(mfcc_feat)
                i+=1
                all_wave.append(samples)
                all_label.append(label)
            lengths.append(i)
                
            
                
    #Loading data for the unknown classes

    wave_unknown = []
    for label in unknown_list:
        waves = [f for f in os.listdir(AUDIO_PATH + '/'+ label) if f.endswith('.wav')]
        for wav in waves:
            samples, sample_rate = librosa.load(AUDIO_PATH + '/' + label + '/' + wav, sr = 8000)
            if(len(samples)== 8000) : 
                if task=="mfcc":
                    mfcc_feat = mfcc(samples, sample_rate)
                    scaler = scaler.fit(mfcc_feat)
                    samples = scaler.transform(mfcc_feat)
                wave_unknown.append(samples)
    #Get subsample of unknown that is balanced in comparison to known group
    random.shuffle(wave_unknown)
    wave_unknown = wave_unknown[:int(np.mean(lengths))]
    label_unknown = ["unknown" for _ in range(len(wave_unknown))]


    #Loading data for the noise class

    wave_noise_full = []

    waves = [f for f in os.listdir(AUDIO_PATH + '/'+ '_background_noise_') if f.endswith('.wav')]
    for wav in waves:
        samples, sample_rate = librosa.load(AUDIO_PATH + '/' + "_background_noise_" + '/' + wav, sr = 8000)
        wave_noise_full = np.concatenate([samples, wave_noise_full])
                
    wave_noise = []
    for i in range(int(np.mean(lengths))):
        r = np.random.randint(0,len(wave_noise_full)-8001)
        samples = wave_noise_full[r:r+8000]
        if task=="mfcc":
            mfcc_feat = mfcc(samples, sample_rate)
            scaler = scaler.fit(mfcc_feat)
            samples = scaler.transform(mfcc_feat)
        wave_noise.append(samples)
    label_noise = ['_background_noise_' for _ in range(len(wave_noise))]


    #Concatenating data
    all_wave = np.concatenate([all_wave, wave_noise, wave_unknown])
    all_label = np.concatenate([all_label, label_noise, label_unknown])

    return all_wave, all_label

def prepare_data(all_wave, all_label, task):
    #Encode target
    le = LabelEncoder()
    y=le.fit_transform(all_label)

    labels=["no", "up", "down", "left", "right", "on", "off", "stop", "go", "yes", '_background_noise_', "unknown"]
    y=np_utils.to_categorical(y, num_classes=len(labels))

    if task!="mfcc":
        all_wave = np.array(all_wave).reshape(-1,8000)

    #train test split
    x_tr, x_val, y_tr, y_val = train_test_split(np.array(all_wave),np.array(y),stratify=y,test_size = 0.2,random_state=777,shuffle=True)
    x_te, x_val, y_te, y_val = train_test_split(x_val,y_val,stratify=y_val,test_size = 0.5,random_state=777,shuffle=True)

    return x_tr, x_te, x_val, y_tr, y_te, y_val


def get_data(task="mfcc"):
    unzipp_files()
    classes, unknown_list = get_classes()
    all_wave, all_label = load_data(classes, unknown_list, task)
    x_tr, x_te, x_val, y_tr, y_te, y_val = prepare_data(all_wave, all_label, task)
    return x_tr, x_te, x_val, y_tr, y_te, y_val