# %% [code]
# ----------------------------------------------------------------------
#   Import libs.
# ----------------------------------------------------------------------
import numpy as np
import pandas as pd
pd.set_option("display.max_columns", 1000)
pd.set_option("display.max_rows", 1000)

from sklearn import set_config
#set_config(transform_output="pandas")

from sklearn.base import BaseEstimator, TransformerMixin
from sklearn.impute import SimpleImputer
from sklearn.preprocessing import OrdinalEncoder
from sklearn.preprocessing import OneHotEncoder
from sklearn.preprocessing import StandardScaler
from sklearn.pipeline import Pipeline
from sklearn.compose import ColumnTransformer
from sklearn.compose import make_column_selector


# ----------------------------------------------------------------------
#   Utility functions to process large tabular dataset.
# ----------------------------------------------------------------------
def reduce_mem_usage(df):
    """ Iterate through all the columns of a dataframe and modify the data type
        to reduce memory usage. Original code was copied from
        https://www.kaggle.com/code/gemartin/load-data-reduce-memory-usage/notebook.
        Some part of code were modified.
    """
    start_mem = df.memory_usage().sum() / 1024**2
    print('Memory usage of dataframe is {:.2f} MB'.format(start_mem))
    
    for col in df.columns:
        col_type = df[col].dtype
        
        #if col_type != object: <-- Modify for processing only numerical features.
        if str(col_type).startswith("float") or str(col_type).startswith("int"):
            c_min = df[col].min()
            c_max = df[col].max()
            if str(col_type)[:3] == 'int':
                if c_min > np.iinfo(np.int8).min and c_max < np.iinfo(np.int8).max:
                    df[col] = df[col].astype(np.int8)
                elif c_min > np.iinfo(np.int16).min and c_max < np.iinfo(np.int16).max:
                    df[col] = df[col].astype(np.int16)
                elif c_min > np.iinfo(np.int32).min and c_max < np.iinfo(np.int32).max:
                    df[col] = df[col].astype(np.int32)
                elif c_min > np.iinfo(np.int64).min and c_max < np.iinfo(np.int64).max:
                    df[col] = df[col].astype(np.int64)
            else:
                #if c_min > np.finfo(np.float16).min and c_max < np.finfo(np.float16).max:
                #    df[col] = df[col].astype(np.float16)
                #elif c_min > np.finfo(np.float32).min and c_max < np.finfo(np.float32).max:
                if c_min > np.finfo(np.float32).min and c_max < np.finfo(np.float32).max:
                    df[col] = df[col].astype(np.float32)
                else:
                    df[col] = df[col].astype(np.float64)
        else:
            df[col] = df[col].astype('category')

    end_mem = df.memory_usage().sum() / 1024**2
    print('Memory usage after optimization is {:.2f} MB'.format(end_mem))
    print('Decreased by {:.1f}%'.format(100 * (start_mem - end_mem) / start_mem))
    
    return df


# ----------------------------------------------------------------------
#   Utility functions and classes to process "Predict Student
#   Performance from Game Play" competion dataset
# ----------------------------------------------------------------------
def _sliceMetadataByLevelGroup(metadata, level_group):
    return metadata.query(f"level_group == '{level_group}'")

def _sliceMetadataByLevelRange(metadata, level_range):
    level_lower, level_upper = level_range
    return metadata.query(f"level >= {level_lower} and level <= {level_upper}")

def _levelGroup(question_id):
    if question_id < 4:    # Q.1 - 3
        return "0-4"
    elif question_id < 14: # Q.4 - 13
        return "5-12"
    else:                  # Q.14 - 18
        return "13-22"
    
def questionIDs(level_group):
    if level_group == "0-4":     # Q.1 - 3
        return [*range(1, 4)]
    elif level_group == "5-12":  # Q.4 - 13
        return [*range(4, 14)]
    elif level_group == "13-22": # Q.14 - 18
        return [*range(14, 19)]
    
def _sliceLabelsByQuestionID(labels, question_id):
    return labels.query(f"question_id == {question_id}")

def _sliceLabelsByAnswer(labels, correct):
    return labels.query(f"correct" == {correct})

def _sliceLabelsBySessionIDs(labels, session_ids):
    return labels.set_index(keys=["session_id"]).loc[session_ids]

def createDatasetForAQuestion(metadata, question_id, labels=None):
    level_group = _levelGroup(question_id=question_id)
    X = _sliceMetadataByLevelGroup(metadata=metadata, level_group=level_group)

    if labels is None:
        return X.set_index(keys=["session_id", "index"])
    else:
        sliced_labels = _sliceLabelsByQuestionID(labels=labels, question_id=question_id)
        session_ids = X["session_id"].unique().tolist()
        y = _sliceLabelsBySessionIDs(labels=sliced_labels, session_ids=session_ids)[["correct"]]
        
        return X.set_index(keys=["session_id", "index"]), y


class FeatureRemover(BaseEstimator, TransformerMixin):
    def __init__(self, features):
        """Remove given features."""
        self.features = features
        super().__init__()
        
    def fit(self, X, y=None):
        return self
    
    def transform(self, X, y=None):
        X_removed = X.drop(columns=self.features, errors="ignore")
        return X_removed

class Aggregator(BaseEstimator, TransformerMixin):
    def __init__(self,
                 groupby_features=["session_id"],
                 stats=["median", "std", "min", "max"]):
        """Aggregate and compute given stats."""
        self.groupby_features = groupby_features
        self.stats = stats
        super().__init__()
        
    def fit(self, X, y=None):
        return self

    def transform(self, X, y=None):
        X_aggregated = X.groupby(by=self.groupby_features).agg(self.stats)
        X_aggregated.columns = ["_".join(cols) for cols in X_aggregated.columns]
        return X_aggregated
    
def mode(series):
    return series.mode().loc[0]

def skew(series):
    return series.skew()

def kurtosis(series):
    return series.kurtosis()