# ------------------------------------------------------------- 
# Imports
# -------------------------------------------------------------

# File handling
import os

# Image handling
from PIL import Image

# Plotting
from matplotlib import pyplot as plt
import seaborn as sns

# Audio file processing
import librosa
import numpy as np

# Data handling
import pandas as pd

# Progress bar
from tqdm.notebook import tqdm

# Keras definitions
import keras

# Model evaluation
import sklearn.metrics

# ------------------------------------------------------------- 
# Output formatting
# -------------------------------------------------------------

class bcolors:
    """Text formatting for terminal output."""
    PURP = '\033[95m'
    BLUE = '\033[94m'
    CYAN = '\033[96m'
    GREEN = '\033[92m'
    YELLOW = '\033[93m'
    RED = '\033[91m'
    ENDC = '\033[0m'
    BOLD = '\033[1m'
    ITALIC = '\033[3m'
    UNDERLINE = '\033[4m'

# ------------------------------------------------------------- 
# Spectrogram generation
# -------------------------------------------------------------

def spec_from_audio(y: np.array,
                    spec_sr: int|None=None, 
                    pixel_format: str = '0-1', 
                    **kwargs) -> np.array:
    """ Generates a single spectrogram from an audio slice.
            
    Args:
        y:
            one dimensional np.array containing the audio data
        spec_sr=None:
            the sampling rate to assume on the audio data for spectrogram generation;
            if None, uses the default sr = 22050 of 'librosa.feature.melspectrogram()'
        pixel_format='0-1':
            pixel values to return in the array
            '0-1' returned array contains pixels in range [0-1] (float)
            '0-255' returned array contains pixels in range '[0-255] (int)'
        **kwargs: 
            additional keyword arguments for 'librosa.feature.melspectrogram()'
            n_mels=128: number of Mel bands to generate
            n_fft=2048: length of the FFT window
            hop_length=512: number of samples between successive frames

    Returns:
        A two-dimensional [height, width] 'np.array' containing the spectrogram
        with pixel values [0-1] or [0-255] depending on 'pixel_format'.
    """

    if spec_sr == None:
        S = librosa.feature.melspectrogram(y=y, **kwargs)
    else:
        S = librosa.feature.melspectrogram(y=y, sr=spec_sr, **kwargs)
    S_db = librosa.power_to_db(S, ref=np.max)
    S_norm = (S_db - S_db.min()) / (S_db.max() - S_db.min())

    if pixel_format == '0-1':
        return S_norm
    elif pixel_format == '0-255':
        return np.round(S_norm*255).astype(np.uint8)
    else:
        raise ValueError(f"'{pixel_format}' is not a valid pixel_format value of the 'spec_from_audio' function.")

def specs_from_file(file_path: str,
                    slice_length: float,
                    slice_overlap: float=0.,
                    audio_sr: int|None=None,
                    **kwargs) -> tuple[list[np.array], list[float]]:
    """ Slices the audio file and generates spectrograms from the slices.
            
    Args:
        file_path:
            path for the audio file
        audio_sr=None:
            sampling rate to load the data with; if none, use the native
            sampling rate of the recording
        slice_length:
            length of the audio slices in seconds
        slice_overlap=0:
            overlap between the audio slices in seconds
        **kwargs: 
            additional keyword arguments for spectrogram generation
            spec_sr=None:
                the sampling rate to assume on the audio data for spectrogram generation;
                if None, use the default sr = 22050 of librosa.feature.melspectrogram;
                if 'Native', use the sampling rate of the loaded audio data
            pixel_format = '0-1':
                pixel values to return in the array
                '0-1' returned array contains pixels in range [0-1] (float)
                '0-255' returned array contains pixels in range '[0-255] (int)'
            n_mels=128:
                number of Mel bands to generate
            n_fft=2048:
                length of the FFT window
            hop_length=512:
                number of samples between successive frames

    Returns:
        A list of two-dimensional [height, width] 'np.array'-s containing the spectrograms
        with pixel values [0-1] or [0-255] depending on 'pixel_format'
        A list of floats containing the start times of the returned spectrograms within the audio.
    """

    #---Load audio file---
    audio, audio_sr = librosa.core.load(file_path, sr = audio_sr)

    #---Set spectrogram sr if 'Native'---
    if ('spec_sr' in kwargs):
        if kwargs['spec_sr'] == 'Native':
            kwargs['spec_sr'] = audio_sr

    #---Prepare slicing---
    spectrograms = []
    start_times = []
    
    diff_time = (slice_length-slice_overlap) # How much time between slice starting points
    slice_start_time = 0.
    slice_samples = int(slice_length * audio_sr)

    #---Slice and generate spectrograms---
    while slice_start_time < (len(audio) / audio_sr):

        start_idx = int(slice_start_time * audio_sr)
        end_idx = start_idx + slice_samples
        
        if end_idx > len(audio):
            audio_slice = audio[start_idx:]
        else:
            audio_slice = audio[start_idx:end_idx]

        spec = spec_from_audio(y = audio_slice, **kwargs)
        
        if start_idx == 0: # Use the first generated spectrogram to get the intended width
            spec_length = spec.shape[1] 
        if spec.shape[1] < spec_length: # If the last spec ends up shorter, pad the missing portions
            pad_width = ((0, 0), (0, spec_length-spec.shape[1]))
            spec = np.pad(spec, pad_width = pad_width, mode='median')
            
        spectrograms.append(spec)
        start_times.append(slice_start_time)

        slice_start_time += diff_time # start_idx for next slice
        
    return spectrograms, start_times

def specs_from_segment(file_path: str,
                       t_start: float,
                       t_end: float,
                       slice_length: float, 
                       min_slice_overlap: float=0.,
                       audio_sr: int|None=None,
                       **kwargs) -> tuple[list[np.array], list[float]]:
    """ Slices the audio segment and generates spectrograms from the slices.

    The segment between t_start and t_end gets sliced into slice_length long segments, where the overlap between
    the segments are at least min_slice_overlap long, but may be larger in order to include every part of the segment
    in at least one slice.
            
    Args:
        file_path:
            path for the audio file
        t_start:
            beginning of the audio segment in seconds (float)
        t_end: 
            end of the audio segment in seconds (float)
        audio_sr=None:
            sampling rate to load the data with; if none, use the native
            sampling rate of the recording
        slice_length:
            length of the audio slices in seconds
        min_slice_overlap=0:
            overlap between the audio slices in seconds
        **kwargs: 
            additional keyword arguments for spectrogram generation
            spec_sr=None:
                the sampling rate to assume on the audio data for spectrogram generation;
                if None, use the default sr = 22050 of librosa.feature.melspectrogram;
                if 'Native', use the sampling rate of the loaded audio data
            pixel_format = '0-1':
                pixel values to return in the array
                '0-1' returned array contains pixels in range [0-1] (float)
                '0-255' returned array contains pixels in range '[0-255] (int)'
            n_mels=128:
                number of Mel bands to generate
            n_fft=2048:
                length of the FFT window
            hop_length=512:
                number of samples between successive frames

    Returns:
        A list of two-dimensional [height, width] 'np.array'-s containing the spectrograms
        with pixel values [0-1] or [0-255] depending on 'pixel_format'
        A list of floats containing the start times of the returned spectrograms within the audio.
    """

    #---Load audio file---
    audio, audio_sr = librosa.core.load(file_path, sr = audio_sr)

    #---Set spectrogram sr if 'Native'---
    if ('spec_sr' in kwargs):
        if kwargs['spec_sr'] == 'Native':
            kwargs['spec_sr'] = audio_sr

    #---Prepare slicing---
    spectrograms = []
    start_times = []
    
    segment_time = t_end - t_start
    # How many slices can fit into the given interval, rounded to the next int
    num_slices = max(int(np.ceil((segment_time - min_slice_overlap) / (slice_length - min_slice_overlap))), 1)
    slice_samples = int(slice_length * audio_sr)

    #---If only one slice can be generated---
    if num_slices == 1:
        center = (t_start + t_end) / 2
        slice_start_time = max(center - slice_length / 2, 0.)
        start_idx = int(slice_start_time * audio_sr)
        end_idx = start_idx + slice_samples
        if end_idx > len(audio):
            end_idx = len(audio)
            start_idx = end_idx - slice_samples

        audio_slice = audio[start_idx:end_idx]
        spec = spec_from_audio(y=audio_slice, **kwargs)

        spectrograms.append(spec)
        start_times.append(slice_start_time)
    
    #---Slice and generate spectrograms if more than one can be fit---
    else:
        diff_time = ((segment_time - slice_length) / (num_slices - 1)) # How much time between slice starting points
        slice_start_time = t_start
        
        for i in range(num_slices):
            start_idx = int(slice_start_time * audio_sr)
            end_idx = start_idx + slice_samples
        
            audio_slice = audio[start_idx:end_idx]
            spec = spec_from_audio(y=audio_slice, **kwargs)
          
            spectrograms.append(spec)
            start_times.append(slice_start_time)

            slice_start_time += diff_time # slice_start_time for next slice
        
    return spectrograms, start_times

# ------------------------------------------------------------- 
# Data generation
# -------------------------------------------------------------

def generate_labeled_spectrograms(df_labels: pd.DataFrame,
                                  files_folder_path: list[str],
                                  save_folder_path: str,
                                  **kwargs) -> None:
    """Generates labelled spectrogram data.

    Args:
        df_labels:
            pandas dataframe containing label data according to the 'rfcx-species-audio-detection' format
        files_folder_path:
            path of the folder containing the audio data
        save_folder_path
            path of the folder to save the generated images to
        **kwargs: 
            additional keyword arguments for audio slice and spectrogram generation
            audio_sr=None:
                sampling rate to load the data with; if none, use the native
                sampling rate of the recording
            slice_length:
                length of the audio slices in seconds
            min_slice_overlap=0:
                overlap between the audio slices in seconds
            spec_sr=None:
                the sampling rate to assume on the audio data for spectrogram generation;
                if None, use the default sr = 22050 of librosa.feature.melspectrogram;
                if 'Native', use the sampling rate of the loaded audio data
            n_mels=128:
                number of Mel bands to generate
            n_fft=2048:
                length of the FFT window
            hop_length=512:
                number of samples between successive frames
    """

    os.makedirs(save_folder_path, exist_ok=True)
    kwargs['pixel_format'] = '0-255'
    
    for i in tqdm(range(len(df_labels))):
        row = df_labels.iloc[i]

        #---Get label data---
        recording_id = row['recording_id']
        species_id = row['species_id']
        t_start = float(row['t_min'])
        t_end = float(row['t_max'])
        file_path = os.path.join(files_folder_path, recording_id + '.flac')

        #---Generate spectrograms---
        spectrograms, start_times = specs_from_segment(file_path=file_path, t_start=t_start, t_end=t_end, **kwargs)

        #---Save spectrograms---
        species_folder_path=os.path.join(save_folder_path, str(species_id))
        os.makedirs(species_folder_path, exist_ok=True)
        
        for spec, start_time in zip(spectrograms, start_times):    
            filename = f'{species_id}_{recording_id}_{start_time:.2f}.png' # {start_time} kell, ha ugyanolyan nevű file keletkezne
            save_path = os.path.join(species_folder_path, filename)

            spec_image = Image.fromarray(spec)
            spec_image.save(save_path)

# ------------------------------------------------------------- 
# Data handling
# -------------------------------------------------------------

class SpecDataset(keras.utils.Sequence):
    
    # Initialization
    def __init__(self,
                 files: list[str],
                 num_classes: int = 0,
                 pixel_format: str = '0-1',
                 batch_size: int=16,
                 shuffle: bool=False,
                 seed=None):
        """ Spectrogram dataset loader - reproducible
            
        Args:
            files:
                list of paths to the spec images of the dataset        
            num_classes=0:
                number of classes/labels
                if 0, then assuming an unlabeled dataset (labels are the same as the images)
            pixel_format='0-1':
                pixel values to return in the array
                '0-1' returned array contains pixels in range [0-1] (float)
                '0-255' returned array contains pixels in range '[0-255] (int)'
            batch_size=16:
                number of files in a batch
            shuffle=False:
                whether to shuffle images
            seed:
                seed for reproducibility
        """
        
        super().__init__()      

        self.files = files.copy()
        self.num_classes = num_classes
        self.batch_size = batch_size
        self.shuffle = shuffle
        self.seed = seed
        self.rng = np.random.default_rng(seed=seed) if seed is not None else np.random.default_rng()

        #---Detecting image shape---
        with Image.open(self.files[0]) as image:
            self.image_shape = (image.size[1], image.size[0], 1)

        #---Setting pixel format---
        if pixel_format == '0-1':
            self.normalize_pixel_values = True
        elif pixel_format == '0-255':
            self.normalize_pixel_values = False
        else:
            raise ValueError(f"'{pixel_format}' is not a valid pixel_format value for the 'AudioDataset' class.")

        #---First Shuffle---
        self.end_of_epoch()

    #---Number of batches---
    def __len__(self):
        return int(np.ceil(len(self.files)/self.batch_size))

    #---Single batch loading---
    def __getitem__(self, index):
        batch_files = self.files[index*self.batch_size:(index+1)*self.batch_size]

        batch_images = []
        batch_labels = []
        
        for f in batch_files:
            with Image.open(f) as image:
                if self.normalize_pixel_values:
                    image = np.array(image, dtype=np.float32) / 255.0
                else:
                    image = np.array(image, dtype=np.uint8)
                if image.ndim == 2:
                    image = image[...,np.newaxis]
                batch_images.append(image)

            if self.num_classes > 0:
                label = int(os.path.basename(os.path.dirname(f)))
                one_hot = np.zeros(self.num_classes, dtype=np.float32)
                one_hot[label] = 1.0
                
                batch_labels.append(one_hot)

        if self.num_classes > 0: 
            return np.stack(batch_images), np.stack(batch_labels)
        else:
            return np.stack(batch_images), np.stack(batch_images)

    #---Return full dataset---
    def get_all_items(self):
        data_images = []
        data_labels = []
        
        for f in self.files:
            with Image.open(f) as image:
                if self.normalize_pixel_values:
                    image = np.array(image, dtype=np.float32) / 255.0
                else:
                    image = np.array(image, dtype=np.uint8)
                if image.ndim == 2:
                    image = image[...,np.newaxis]  
                data_images.append(image)

            if self.num_classes > 0:
                label = int(os.path.basename(os.path.dirname(f)))
                one_hot = np.zeros(self.num_classes, dtype=np.float32)
                one_hot[label] = 1.0
                
                data_labels.append(one_hot)

        if self.num_classes > 0: 
            return np.stack(data_images), np.stack(data_labels)
        else:
            return np.stack(data_images), np.stack(data_images)
    
    #---Shuffle dataset---
    def end_of_epoch(self):
        if self.shuffle:
            self.files = self.rng.permutation(self.files)

# ------------------------------------------------------------- 
# Submission .csv generation
# -------------------------------------------------------------

def generate_submission(test_path: str,
                        model: keras.Model,
                        csv_name: str,
                        **kwargs) -> None:
    """ Generates submission .csv for rfcx-species-audio-detection
            
    Args:
        test_path:
            path to test files (audio or pre-generated spectrograms)
        model:
            keras model to use for prediction on the spectrograms
        csv_name:
            fname for the submission .csv to generate (into '/kaggle/working')
        
        **kwargs: 
            additional keyword arguments for spectrogram generation/loading
                audio_sr=None:
                    sampling rate to load the data with; if none, use the native
                    sampling rate of the recording
                slice_length:
                    length of the audio slices in seconds
                slice_overlap=0:
                    overlap between the audio slices in seconds
                spec_sr=None:
                    the sampling rate to assume on the audio data for spectrogram generation;
                    if None, use the default sr = 22050 of librosa.feature.melspectrogram;
                    if 'Native', use the sampling rate of the loaded audio data
                pixel_format = '0-1':
                    pixel values to return in the array
                    '0-1' returned array contains pixels in range [0-1] (float)
                    '0-255' returned array contains pixels in range '[0-255] (int)'
                n_mels=128:
                    number of Mel bands to generate
                n_fft=2048:
                    length of the FFT window
                hop_length=512:
                    number of samples between successive frames
    """
    
    #---Output file path---
    submission_dir_path = '/kaggle/working'
    os.makedirs(submission_dir_path, exist_ok=True)
    csv_path = os.path.join(submission_dir_path, csv_name)

    #---Get relevant sizes---
    if 'n_mels' not in kwargs:
        kwargs['n_mels'] = model.input_shape[1]
    num_species = model.output_shape[1]
    
    test_entries = list(os.scandir(test_path))
    are_entries_files = [entry.is_file() for entry in test_entries]

    if (all(are_entries_files) == False) and (any(are_entries_files) == True):
        print(f"Elements at root of src_path ('{src_path}') consist of both files and directories. This is not a valid test source.")
    else:
        rows = []
        
        if all(are_entries_files) == True:
            """'test_path' contains files - raw test folder for audio files"""

            if 'slice_length' not in kwargs:
                raise TypeError(f"'slice_length' must be given when generating submission from raw audio files")
            
            for file in tqdm(test_entries):
            
                if file.name.endswith('.flac'):
                    recording_id = file.name.replace('.flac', '')
                
                    spectrograms, _ = specs_from_file(file.path, **kwargs)
                    outputs = model.predict(np.array(spectrograms)[..., None], verbose=0)
                    pred = np.max(outputs, axis=0)
    
                    rows.append([recording_id] + list(pred))
       
        elif any(are_entries_files) == False: 
            """'test_path' contains folder - pre-generated spectrograms"""

            normalize_pixel_values = True
            if ('pixel_format' in kwargs):
                if kwargs['pixel_format'] == '0-255':
                    normalize_pixel_values = False

            for folder in tqdm(test_entries):
                recording_id = folder.name
                batch_images=[]
                files = os.scandir(folder)

                for f in files:
                    with Image.open(f) as image:
                        if normalize_pixel_values:
                            image = np.array(image, dtype=np.float32)/255.0
                        else:
                            image = np.array(image, dtype =np.uint8)
                        if image.ndim == 2:
                            image = image[...,np.newaxis]
                        batch_images.append(image)

                spectrograms = np.stack(batch_images)
                outputs = model.predict(spectrograms, verbose=0)
                pred = np.max(outputs, axis=0)
    
                rows.append([recording_id] + list(pred))
                
        #---Generate .csv from df---
        df = pd.DataFrame(rows, columns=['recording_id']+[f"s{i}" for i in range(num_species)])
        df.to_csv(csv_path, float_format='%.5f', index=False)

# ------------------------------------------------------------- 
# Model Evaluation
# -------------------------------------------------------------

def eval_model(model: keras.Model,
               history: dict,
               eval_dataset: SpecDataset) -> None:
    """ Evaluate a trained model with visuals and stats
            
    Args:
        model:
            The trained model to evaluate
        history:
            Training history dictionary object
        eval_dataset:
            SpecDataset object containing the validation data
    """
    
    #---Training plots---
    # Get loss/accuracy types
    keys = [key for key in history.keys() if not key.startswith('val_')]
    losses = [key for key in keys if key.endswith('loss')]
    accuracies = [key for key in keys if key.endswith('accuracy')]
    
    # Set up figures
    plt.figure(figsize=(11, 4))
    linestyles = ['solid', 'dashdot', 'dotted', 'dashed']
        
    
    # Plot loss(es)
    plt.subplot(1, 2, 1)
    if len(losses)==1:
        plt.plot(history[losses[0]], label='Training loss')
        plt.plot(history['val_' + losses[0]], label='Validation loss')
    else:
        for (i, key) in enumerate(losses):
            plt.plot(history[key], label=f'Training {key}',
                     color='C0', linestyle=linestyles[i%4])
            plt.plot(history['val_' + key], label=f'Validation {key}',
                     color='C1', linestyle=linestyles[i%4])
    plt.title('Losses')
    plt.xlabel('Epoch')
    plt.ylabel('Loss')
    plt.legend()
    
    # Plot accuracy(ies)
    plt.subplot(1, 2, 2)
    if len(accuracies)==1:
        plt.plot(history[accuracies[0]], label='Training accuracy')
        plt.plot(history['val_' + accuracies[0]], label='Validation accuracy')
    else:
        for (i, key) in enumerate(accuracies):
            plt.plot(history[key], label=f'Training {key}',
                     color='C0', linestyle=linestyles[i%4])
            plt.plot(history['val_' + key], label=f'Validation {key}',
                     color='C1', linestyle=linestyles[i%4])
    plt.title('Accuracy')
    plt.xlabel('Epoch')
    plt.ylabel('Accuracy')
    plt.legend()
    
    # Show figure
    plt.tight_layout()
    plt.show()

    #---Metrics & Scores---
    # Create predictions for the test images
    test_data, test_labels = eval_dataset.get_all_items()
    test_confidences = model.predict(test_data, verbose=0)
    test_preds = np.round(test_confidences)

    m = keras.metrics.CategoricalAccuracy()
    m.reset_state()
    m.update_state(test_labels, test_confidences)

    # Print all metrics
    print(f"\n{bcolors.BLUE}{bcolors.BOLD}{bcolors.ITALIC}/// ----- Metrics & Scores ----- ///{bcolors.ENDC}\n")

    print(f"Test loss - categorical crossentropy        : "
          f"{bcolors.BLUE}{sklearn.metrics.log_loss(test_labels, test_confidences):.7f}{bcolors.ENDC}")
    print(f"Test loss - binary crossentropy             : "
          f"{bcolors.BLUE}{sklearn.metrics.log_loss(test_labels.flatten(), test_confidences.flatten()):.7f}{bcolors.ENDC}\n")
    print(f"Test categorical (subset) accuracy (keras)  : "
          f"{bcolors.BLUE}{m.result():.7f}{bcolors.ENDC}"
          f" --> (% of samples where argmax(confidence) matches argmax(label))")
    print(f"Test categorical (subset) accuracy (sklearn): "
          f"{bcolors.BLUE}{sklearn.metrics.accuracy_score(test_labels, test_preds):.7f}{bcolors.ENDC}"
          f" --> (% where predicted labels (conf. rounded to int) match exactly with true labels)")
    print(f"Test binary accuracy                        : "
          f"{bcolors.BLUE}{1-sklearn.metrics.hamming_loss(test_labels, test_preds):.7f}{bcolors.ENDC}\n")
        
    print(f"Test precision                              : "
          f"{bcolors.BLUE}{sklearn.metrics.precision_score(test_labels, test_preds, average='micro'):.7f}{bcolors.ENDC}"
          f" --> (% of predicted positives (in a binary sense) that are actual positives)")
    print(f"Test recall                                 : "
          f"{bcolors.BLUE}{sklearn.metrics.recall_score(test_labels, test_preds, average='micro'):.7f}{bcolors.ENDC}"
          f" --> (% of actual positives (in a binary sense) that are predicted correctly)")
    print(f"Test f1_score                               : "
          f"{bcolors.BLUE}{sklearn.metrics.f1_score(test_labels, test_preds, average='micro'):.7f}{bcolors.ENDC}\n")

    #---Confusion matrix---
    confmatrix = sklearn.metrics.confusion_matrix(np.argmax(test_labels, 1), np.argmax(test_confidences, 1))

    fig = plt.figure(figsize=(12, 8))
    ax = fig.add_subplot(111)
    res = sns.heatmap(confmatrix, annot=True, cmap = plt.get_cmap('Blues'), square=True)
    
    plt.ylabel('True label')
    plt.xlabel('Predicted label')
    plt.title('Confusion Matrix')
    
    plt.show()