# work in progress - experimenting with new approach ..
'''
review:
[ https://www.kaggle.com/code/hideyukizushi/home-aftersubmissionsopen-3-11-2024-lb-567 ]
[ https://www.kaggle.com/code/tanishqdublish/metric-trick-home-credit-baseline-inference ]

[ https://www.kaggle.com/code/syhens/aes2-lb-prob-no-new-prompt-in-hidden-test-set ]
[ https://www.kaggle.com/code/anshulgupta1502/aimo-llama-3 ]
[ https://www.kaggle.com/code/mathchi/home-credit-risk-with-detailed-feature-engineering ]
[ https://www.kaggle.com/code/ogrellier/good-fun-with-ligthgbm ]
'''

# import modules for interacting with OS and managing files
import os, subprocess as sp, joblib
# import modules for ETL/ELT
import numpy as np, pandas as pd, polars as pl, pyarrow as pa
# import modules for Data Mining
import lightgbm as lgb
from sklearn.base import BaseEstimator, RegressorMixin; from sklearn.metrics import roc_auc_score 
# import misc. modules
import gc, warnings; warnings.filterwarnings('ignore')
from glob import glob; from typing import Optional; from pprint import pprint as pp

# define data extraction class(es)
class ETL:
    '''
    '''
    # [E]
    
    # [T]
    # cleaning method to set datatypes for specific columns in DataFrame
    def T_set_table_dtypes(self, df: pl.DataFrame) -> pl.DataFrame:
        for col in df.columns:
            if col in ["case_id", "WEEK_NUM", "num_group1", "num_group2"]:
                df = df.with_columns(pl.col(col).cast(pl.Int64))
            elif col in ["date_decision"]:
                df = df.with_columns(pl.col(col).cast(pl.Date))
            elif col[-1] in ("P", "A"):
                df = df.with_columns(pl.col(col).cast(pl.Float64))
            elif col[-1] in ("M"):
                df = df.with_columns(pl.col(col).cast(pl.String))
            elif col[-1] in ("D"):
                df = df.with_columns(pl.col(col).cast(pl.Date)) 
        del col; gc.collect()
        return df        
    
    # cleaning method to filter out columns with missing values and frequency
    def T_filter_cols(self, df):
        for col in df.columns:
            if col not in ["target", "case_id", "WEEK_NUM"]:
                isnull = df[col].is_null().mean()
                if isnull > 0.7:
                    df = df.drop(col)
        for col in df.columns:
            if (col not in ["target", "case_id", "WEEK_NUM"]) & (df[col].dtype == pl.String):
                freq = df[col].n_unique()
                if (freq == 1) | (freq > 200):
                    df = df.drop(col)
                del freq; gc.collect()
        del col; gc.collect()                    
        return df
    
    # cleaning method for joining additional depth DataFrames
    def T_join_to_base(self, df_base, depth_0, depth_1, depth_2):
        for i, df in enumerate(depth_0 + depth_1 + depth_2):
            df_base = df_base.join(df, how="left", on="case_id", suffix=f"_{i}")
        del i, df; gc.collect()
        return df_base
            
    # feature eng. method to create month and weekday features based on "date_decision"
    def T_date_features(self, df_base):
        df_base = df_base.with_columns(
            month_decision   = pl.col("date_decision").dt.month(),
            weekday_decision = pl.col("date_decision").dt.weekday()
        )
        return df_base
    
    # cleaning method to handle date columns and calculate time
    def T_handle_dates(self, df: pl.DataFrame) -> pl.DataFrame:
        for col in df.columns:
            if col[-1] in ("D"):
                df = df.with_columns(pl.col(col) - pl.col("date_decision"))
                df = df.with_columns(pl.col(col).dt.total_days())
        del col; gc.collect()        
        df = df.drop("date_decision", "MONTH")        
        return df
    
    # feature eng. method to aggregate numerical features
    def T_num_expr(self, df):
        cols        = [col for col in df.columns if col[-1] in ("P", "A")]
        max_expr    = [pl.max(col).alias(f"max_{col}") for col in cols]
        last_expr   = [pl.last(col).alias(f"last_{col}") for col in cols]
        median_expr = [pl.median(col).alias(f"median_{col}") for col in cols]
        del cols; gc.collect()
        return max_expr + last_expr + median_expr
    
    # feature eng. method to aggregate date features
    def T_date_expr(self, df):
        cols        = [col for col in df.columns if col[-1] in ("D")]
        max_expr    = [pl.max(col).alias(f"max_{col}") for col in cols]
        last_expr   = [pl.last(col).alias(f"last_{col}") for col in cols]
        median_expr = [pl.median(col).alias(f"median_{col}") for col in cols]
        del cols; gc.collect()
        return max_expr + last_expr + median_expr 
    
    # feature eng. method to aggregate string features
    def T_str_expr(self, df):
        cols      = [col for col in df.columns if col[-1] in ("M")]
        max_expr  = [pl.max(col).alias(f"max_{col}") for col in cols]
        last_expr = [pl.last(col).alias(f"last_{col}") for col in cols]
        del cols; gc.collect()
        return max_expr + last_expr
    
    # feature eng. method to aggregate misc. features
    def T_other_expr(self, df):
        cols      = [col for col in df.columns if col[-1] in ("T", "L")]
        max_expr  = [pl.max(col).alias(f"max_{col}") for col in cols]
        last_expr = [pl.last(col).alias(f"last_{col}") for col in cols]
        del cols; gc.collect()
        return max_expr + last_expr
    
    # feature eng. method to aggregate count features
    def T_count_expr(self, df):
        cols       = [col for col in df.columns if "num_group" in col]
        max_expr   = [pl.max(col).alias(f"max_{col}") for col in cols]
        last_expr  = [pl.last(col).alias(f"last_{col}") for col in cols]
        del cols; gc.collect()
        return max_expr + last_expr
    
    # misc. method to get all aggregations
    def T_get_exprs(self, df):
        exprs = self.T_num_expr(df) + self.T_date_expr(df) + self.T_str_expr(df) + self.T_other_expr(df) + self.T_count_expr(df)
        del df; gc.collect()
        return exprs
    
    # misc. method for converting polars to pandas DataFrame
    def T_to_pandas(self, df: pl.DataFrame, cat_cols: Optional[str] = None):
        df = df.to_pandas()
        if cat_cols is None:
            cat_cols = list(df.select_dtypes("object").columns)
        df[cat_cols] = df[cat_cols].astype("category")
        return df, cat_cols
    
    # misc. method for memory optimization
    def T_reduce_mem_usage_by_col_dtype(self, df):
        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 str(col_type) == "category":
                continue
            if col_type != object:
                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:
                        df[col] = df[col].astype(np.float32)
                    else:
                        df[col] = df[col].astype(np.float64)
                del c_min, c_max; gc.collect()
            else:
                continue
            del col_type; gc.collect()
        end_mem = df.memory_usage().sum() / 1024**2  # Memory usage after optimization
        print('Memory usage after optimization is: {:.2f} MB'.format(end_mem))
        print('Decreased by {:.1f}%'.format(100 * (start_mem - end_mem) / start_mem))
        del start_mem, end_mem; gc.collect()
        return df    
    
    # [L]
    # method to read base files
    def L_read_file(self, path: str, depth: Optional[int] = None) -> pl.DataFrame:
        print(f"\nReading ...")
        if path.endswith(".csv"):
            df = pl.read_csv(path)
            print(f"[ {path} ] {df.shape}")
            df = df.pipe(self.T_set_table_types)
            if depth in [1, 2]:
                df = df.group_by("case_id").agg(self.T_get_exprs(df))
            print(f"\nTransformed to ...")
            print(df)
        elif path.endswith(".parquet"):
            # read file into Polars DataFrame: [ https://stackoverflow.com/q/77601763/11492382 ]
            import pyarrow as pa
            try:
                df = pl.read_parquet(path, use_pyarrow=True)
            except pa.ArrowNotImplementedError as e:
                df = pl.read_parquet(path)
            print(f"[ {path} ] {df.shape}")
            df = df.pipe(self.T_set_table_dtypes)
            if depth in [1, 2]:
                df = df.group_by("case_id").agg(self.T_get_exprs(df))
        print(f"\nTransformed to ...")
        print(df)
        del path, depth; gc.collect()
        return df
    
    # method to read iterable files
    def L_read_files(self, regex_path: list, depth: Optional[int] = None) -> pl.DataFrame:
        chunks = []        
        print(f"\nReading ...")
        for path in regex_path:
            if path.endswith(".csv"):
                df = pl.read_csv(path)                
                print(f"[ {path} ] {df.shape}")
                df = df.pipe(self.T_set_table_types)
                if depth in [1, 2]:
                    df = df.group_by("case_id").agg(self.T_get_exprs(df))
            elif path.endswith(".parquet"):
                # read file into Polars DataFrame: [ https://stackoverflow.com/q/77601763/11492382 ]
                df = pl.read_parquet(path)
                print(f"[ {path} ] {df.shape}")
                df = df.pipe(self.T_set_table_dtypes)
                if depth in [1, 2]:
                    df = df.group_by("case_id").agg(self.T_get_exprs(df))
            chunks.append(df)                
        df = pl.concat(chunks, how="vertical_relaxed").unique(subset=["case_id"])
        print(f"\nTransformed to ...")
        print(df)
        del regex_path, depth, chunks; gc.collect()
        return df
    
if __name__ == '__main__':
    gc.enable()
    ETL = ETL()
    
    # -----------------------------------------------------------------------------------------------------
    # perform data extraction
    os.chdir('/')
    ROOT     = "/kaggle/input/home-credit-credit-risk-model-stability"
    TEST_DIR = ROOT + "/parquet_files" + "/test"
    
    # Read the base test data and store it with the key 'df_base'
    df_base = ETL.L_read_file(TEST_DIR + "/test_base.parquet")
    
    # Read depth 0, 1 and 2 data
    depth_0 = [
        ETL.L_read_file(TEST_DIR + "/test_static_cb_0.parquet"),
        ETL.L_read_files(glob(TEST_DIR + "/test_static_0_*.parquet")),
    ]
    depth_1 = [
        ETL.L_read_files(glob(TEST_DIR + "/test_applprev_1_*.parquet"), 1),
        ETL.L_read_file(TEST_DIR + "/test_tax_registry_a_1.parquet", 1),
        ETL.L_read_file(TEST_DIR + "/test_tax_registry_b_1.parquet", 1),
        ETL.L_read_file(TEST_DIR + "/test_tax_registry_c_1.parquet", 1),
        ETL.L_read_files(glob(TEST_DIR + "/test_credit_bureau_a_1_*.parquet"), 1),
        ETL.L_read_file(TEST_DIR + "/test_credit_bureau_b_1.parquet", 1),
        ETL.L_read_file(TEST_DIR + "/test_other_1.parquet", 1),
        ETL.L_read_file(TEST_DIR + "/test_person_1.parquet", 1),
        ETL.L_read_file(TEST_DIR + "/test_deposit_1.parquet", 1),
        ETL.L_read_file(TEST_DIR + "/test_debitcard_1.parquet", 1),
    ]
    depth_2 = [
        ETL.L_read_file(TEST_DIR + "/test_credit_bureau_b_2.parquet", 2),
        ETL.L_read_files(glob(TEST_DIR + "/test_credit_bureau_a_2_*.parquet"), 2),
        ETL.L_read_file(TEST_DIR + "/test_applprev_2.parquet", 2),
        ETL.L_read_file(TEST_DIR + "/test_person_2.parquet", 2)
    ]
    df_test = (
        df_base
        .pipe(ETL.T_date_features)
        .pipe(ETL.T_join_to_base, depth_0, depth_1, depth_2)
        .pipe(ETL.T_handle_dates)
    )
    print(df_test.columns)
    print(f"df_test shape is {df_test.shape}")
    del depth_0, depth_1, depth_2, ROOT, TEST_DIR; gc.collect()
    
    # -----------------------------------------------------------------------------------------------------
    # perform data mining
    os.chdir('/')
    lgb_notebook_info = joblib.load('/kaggle/input/homecredit-models-public/other/lgb/1/notebook_info.joblib')
    print(f"- [lgb] notebook_start_time: {lgb_notebook_info['notebook_start_time']}")
    print(f"- [lgb] description        : {lgb_notebook_info['description']}")
    cols     = lgb_notebook_info['cols']
    cat_cols = lgb_notebook_info['cat_cols']
    print(f"- [lgb] len(cols)          : {len(cols)}")
    print(f"- [lgb] len(cat_cols)      : {len(cat_cols)}")
    del lgb_notebook_info; gc.collect()
    lgb_models = joblib.load('/kaggle/input/homecredit-models-public/other/lgb/1/lgb_models.joblib')
    pp(lgb_models)
    updated_cols = []
    for col in cols:
        if r'mean' in col:
            updated_col = col.replace(r'mean', 'median')
            updated_cols.append(updated_col)
        else:
            updated_cols.append(col)
    cols = updated_cols
    del updated_cols; gc.collect()

    # Select columns of interest from the test data
    df_test = df_test.select(['case_id'] + cols)

    # Convert the test data to a pandas DataFrame and optimize memory usage
    df_test, cat_cols = ETL.T_to_pandas(df_test, cat_cols)
    df_test = ETL.T_reduce_mem_usage_by_col_dtype(df_test)

    # Set the case_id column as the index of the DataFrame
    df_test = df_test.set_index('case_id')
    df_test

    # Print the shape of the test data after processing
    print("test data shape:\t", df_test.shape)
    del cols, cat_cols; gc.collect()

    # Predict probabilities for the test data
    os.chdir('/')
    df_test.info()
    lgb_models[0].predict_proba(df_test)
    y_preds = [estimator.predict_proba(df_test) for estimator in lgb_models]
    y_pred  = pd.Series((np.mean(y_preds, axis=0))[:, 1], index=df_test.index)
    del lgb_models, y_preds, df_test; gc.collect()

    # -----------------------------------------------------------------------------------------------------
    # Read the sample submission file
    df_subm = pd.read_csv("/kaggle/input/home-credit-credit-risk-model-stability/sample_submission.csv")
    df_subm = df_subm.set_index("case_id")

    # Assign predicted probabilities to the submission dataframe
    df_subm["score"] = y_pred
    del y_pred; gc.collect()

    # Save the submission dataframe to a CSV file
    df_subm.to_csv('/kaggle/working/submission.csv', index=True)

    # Display the submission dataframe
    print(df_subm)
    del df_subm; gc.collect()