# Para mas detalles: https://coderzcolumn.com/tutorials/machine-learning/model-evaluation-scoring-metrics-scikit-learn-sklearn

import pandas as pd
import numpy as np
from sklearn.metrics import accuracy_score
import scikitplot as skplt
from sklearn.metrics import classification_report, precision_score, recall_score, f1_score, precision_recall_fscore_support, balanced_accuracy_score
from sklearn.neighbors import NearestNeighbors
# Import other models here... RandomForest, Support Vector Machines

X_data ... #este array contiene todos los datos de cada intento/sujeto (eso son las filas); y en las columnas tendrá las distintas variables que deseemos (por ejemplo, datos en bruto de cada canal; o la media/min/max/etc. de cada canal; o fórmulas más complejas como el cálculo de la PSD/entropía,FFT,etc. por cada canal)
all_labels ... #necesitamos un array con los labels 0/1 (p/q) que identifica a cada fila según si se pulsó p/q para ese dato
subjects ... #necesitamos un array con IDs de los sujetos para cada fila (este array engloba a todos los datos que tengamos, para identificar cada dato a un sujeto, solo incluye enteros para decir si una fila pertenece al sujeto 0, 1, 2, 3...N)

# cambiar esta formula segun necesidad
# 100 es debe ser mayor que el numero de sujetos
intentos_groups = [[ int(sj) + (i*100) + (all_labels[i] + 1) ] for i,sj in enumerate(subjects) ]


dir_path = "results/"
figsize = (3, 4) #para el tamaño de las matrices de confusion
fontsize = 16
cmap=plt.cm.Blues
FOLDS_N = 5  # 5-fold cross validation
sgkf = StratifiedGroupKFold(n_splits=FOLDS_N)
labels = np.unique(all_labels)



results = {
    "knn": {},
    "svm": {},
    "rf": {}
}
# results tendrá la forma:
# {
#     "knn": 
#     {
#         "fold-0":   {
#                     "accuracy": accuracy_score(y_validation_data, predictions)
#                     ,"balanced-accuracy": balanced_accuracy_score(y_validation_data, predictions) #It returns an average recall of each class in classification problem. It's useful to deal with imbalanced datasets.
#                     ,"balanced-accuracy-adjusted": balanced_accuracy_score(y_validation_data, predictions, adjusted=True)
#                     , "f1_score": f1_score(y_validation_data, predictions)
#                     , "precision": precision_score(y_validation_data, predictions)
#                     , "recall": recall_score(y_validation_data, predictions)
#                     },
#         "fold-1":   {
#                     "accuracy": accuracy_score(y_validation_data, predictions)
#                     ,"balanced-accuracy": balanced_accuracy_score(y_validation_data, predictions) #It returns an average recall of each class in classification problem. It's useful to deal with imbalanced datasets.
#                     ,"balanced-accuracy-adjusted": balanced_accuracy_score(y_validation_data, predictions, adjusted=True)
#                     , "f1_score": f1_score(y_validation_data, predictions)
#                     , "precision": precision_score(y_validation_data, predictions)
#                     , "recall": recall_score(y_validation_data, predictions)
#                     },
#         "fold-2":   {},
#         "fold-3":   {},
#         "fold-4":   {}
#     },
#     "svm": {},
#     "rf": {}
# }

for model_name in ["knn", "svm", "rf"]: #se pueder ir añadiendo más
   for fold, (train_index, val_index) in enumerate(sgkf.split(X_data, all_labels, groups= intentos_groups)):
      # --------------------------------------------
      # Preparar datos
      # --------------------------------------------
      training_data = X_data[train_index,:,:]
      validation_data = X_data[val_index,:,:]
      y_training_data = all_labels[train_index]
      y_validation_data = all_labels[val_index]
      #sj_training_data = subjects[train_index] #no lo necesitamos ya
      #sj_validation_data = subjects[val_index] #no lo necesitamos ya

      # --------------------------------------------
      # Entrenar cada modelo
      # --------------------------------------------
      if model_name == "knn":
         #es habitual probar varios valores de n_neighbors, que sean siempre impares (para evitar el empate en la decisión)y hasta sqrt(len(training_data)
         #for k in range(1, sqrt(len(training_data), 2) #pero puede ser largo
         model  = NearestNeighbors(n_neighbors=3) 
      # elif model_name == "svm":
        #     # Train your SVM model here...
      # elif model_name == "rf":
        #     # Train your RF model here...
      # train other models...


      #entrenamiento 
      model.fit(training_data, y_training_data)

      # predicción
      predictions =  model.predict(validation_data, y_validation_data)

      # --------------------------------------------
      # Evaluar predicciones y guardar resultados
      # --------------------------------------------
      # Save metrics
      results[model_name][f"fold-{fold}"] = {
         "accuracy": accuracy_score(y_validation_data, predictions),
         "balanced-accuracy": balanced_accuracy_score(y_validation_data, predictions),
         "balanced-accuracy-adjusted": balanced_accuracy_score(y_validation_data, predictions, adjusted=True),
         "f1_score": f1_score(y_validation_data, predictions),
         "precision": precision_score(y_validation_data, predictions),
         "recall": recall_score(y_validation_data, predictions)
      }
      
      # --------------------------------------------
      # Plots de matrices de confusion
      # --------------------------------------------
      cm = confusion_matrix(y_true, y_pred) # Compute confusion matrix
      cm_norm = cm.astype('float') / cm.sum(axis=1)[:, np.newaxis] # Compute normalized confusion matrix

      # Plot normalized confusion matrix
      # fig, ax = plt.subplots(figsize=figsize)
      # sns.heatmap(cm, annot=True, cmap=cmap, square=True, ax=ax,
      #           xtickall_labels=all_labels, ytickall_labels=all_labels, fmt='g', annot_kws={"fontsize": fontsize})
      # ax.set_xlabel('Predicted label',fontsize=fontsize)
      # ax.set_ylabel('True label',fontsize=fontsize)
      # ax.set_title('Confusion Matrix',fontsize=fontsize)
      # ax.set_xtickall_labels(labels, rotation=45, ha='right',fontsize=fontsize)
      # ax.set_ytickall_labels(labels, rotation=0,fontsize=fontsize)
      # plt.savefig(f"{dir_path}model_{model_name}_fold_{fold}_cm.png", dpi=300, bbox_inches='tight') # Save confusion matrix plot
     
      # Plot normalized confusion matrix >> esta es más fácil de ver 
      fig, ax = plt.subplots(figsize=figsize)
      sns.heatmap(cm_norm, annot=True, cmap=cmap, square=True, ax=ax,
                 xtickall_labels=all_labels, ytickall_labels=all_labels, fmt='.3f', annot_kws={"fontsize": fontsize})
      ax.set_xlabel('Predicted label',fontsize=fontsize)
      ax.set_ylabel('True label',fontsize=fontsize)
      ax.set_title('Normalized Confusion Matrix',fontsize=fontsize)
      ax.set_xtickall_labels(labels, rotation=45, ha='right',fontsize=fontsize)
      ax.set_ytickall_labels(labels, rotation=0,fontsize=fontsize)
      plt.savefig(f"{dir_path}model_{model_name}_fold_{fold}_norm_cm.png", dpi=300, bbox_inches='tight')       # Save normalized confusion matrix plot

   # --------------------------------------------
   # fin de la iteración de cada fold
 #--------------------------------------------
 # fin de la iteración de cada modelo

# --------------------------------------------
# Al acabar guardamos todas las métricas resultantes
# --------------------------------------------

# Flatten dictionary to DataFrame
flattened_results = pd.json_normalize(results, sep='_')

# Compute averages over folds
average_results = flattened_results.groupby(flattened_results.columns.str.split('_').str[0], axis=1).mean()

# Save to CSV
flattened_results.to_csv(f'{dir_path}results.csv')
average_results.to_csv(f'{dir_path}average_results.csv')






