#Rolensky Louis RF
import numpy as np
import pandas as pd
from pathlib import Path

from sklearn.pipeline import Pipeline
from sklearn.impute import SimpleImputer
from sklearn.ensemble import RandomForestClassifier


# Paths
TRAIN_PATH = "set_17/data/train.csv"
DEV_PATH   = "set_17/data/dev.csv"
EVAL_PATH  = "set_17/data/eval.csv"

OUTPUT_DIR = Path("final_hyp")
OUTPUT_DIR.mkdir(exist_ok=True)



# Classes
CLASS_NAMES = {
    0: 'norm',
    1: 'artf',
    2: 'nneo',
    3: 'infl',
    4: 'susp',
    5: 'dcis',
    6: 'indc',
    7: 'null',
    8: 'bckg'
}

# Score-relevant classes
SCORED_CLASSES = [0, 2, 3, 5, 6]
BCKG_CLASS = 8



# Stage mappings
# Stage 1: tissue vs bckg
STAGE1_MAP = {
    0:0, 2:0, 3:0, 5:0, 6:0,   # tissue
    8:1, 1:1, 4:1, 7:1          # bckg + ignored
}
# Stage 2: split tissue
STAGE2_MAP = {
    0:0,   # norm
    2:1,   # nneo
    3:2, 5:2, 6:2   # disease
}



# Model params 

S1_PARAMS = dict(
    n_estimators=200,
    max_depth=6,
    min_samples_leaf=20,
    class_weight='balanced_subsample',
    n_jobs=-1,
    random_state=42
)

S2_PARAMS = dict(
    n_estimators=300,
    max_depth=10,
    min_samples_leaf=10,
    class_weight='balanced_subsample',
    n_jobs=-1,
    random_state=42
)

S3_PARAMS = dict(
    n_estimators=300,
    max_depth=4,
    min_samples_leaf=20,
    class_weight='balanced_subsample',
    n_jobs=-1,
    random_state=42
)



# Helpers
def load_data(path):
    df = pd.read_csv(path, comment="#", header=None)
    X = df.iloc[:, 1:].astype(float)
    y = df.iloc[:, 0].astype(int)
    return X, y


def load_unlabeled(path):
    df = pd.read_csv(path, comment="#", header=None)
    return df.iloc[:, 1:].astype(float)


def save_hyp(preds, path):
    with open(path, "w") as f:
        f.write("label\n")
        for p in preds:
            f.write(f"{int(p)}\n")


def build_model(params):
    return Pipeline([
        ("imputer", SimpleImputer(strategy="median")),
        ("rf", RandomForestClassifier(**params))
    ])



# Main

if __name__ == "__main__":

    print("=" * 60)
    print("Hierarchical Random Forest (FINAL VERSION)")
    print("=" * 60)

    # Load
    X_train, y_train = load_data(TRAIN_PATH)
    X_dev, y_dev     = load_data(DEV_PATH)
    X_eval           = load_unlabeled(EVAL_PATH)

    y = y_train.values

    
    # Stage 1: tissue vs bckg
    
    print("\nTraining Stage 1 (tissue vs bckg)...")

    y_s1 = np.array([STAGE1_MAP[int(v)] for v in y])
    s1 = build_model(S1_PARAMS)
    s1.fit(X_train, y_s1)

    
    # Stage 2: norm / nneo / disease
    
    print("Training Stage 2 (norm / nneo / disease)...")

    tissue_mask = np.isin(y, [0,2,3,5,6])
    y_s2 = np.array([STAGE2_MAP[int(v)] for v in y[tissue_mask]])

    s2 = build_model(S2_PARAMS)
    s2.fit(X_train[tissue_mask], y_s2)

    
    # Stage 3: disease specialist
    
    print("Training Stage 3 (infl / dcis / indc)...")

    disease_mask = np.isin(y, [3,5,6])
    s3 = build_model(S3_PARAMS)
    s3.fit(X_train[disease_mask], y[disease_mask])

    
    # Prediction function
    def predict(X):
        n = len(X)
        final_pred = np.full(n, 8, dtype=int)  # default = bckg

        # Stage 1
        s1_pred = s1.predict(X)

        tissue_idx = np.where(s1_pred == 0)[0]
        if len(tissue_idx) == 0:
            return final_pred

        # Stage 2
        X_tissue = X.iloc[tissue_idx]
        s2_pred = s2.predict(X_tissue)

        # norm
        final_pred[tissue_idx[s2_pred == 0]] = 0

        # nneo
        final_pred[tissue_idx[s2_pred == 1]] = 2

        # disease → stage 3
        disease_idx = tissue_idx[s2_pred == 2]

        if len(disease_idx) > 0:
            final_pred[disease_idx] = s3.predict(X.iloc[disease_idx])

        return final_pred


    
    # Predict all splits

    print("\nPredicting...")

    train_pred = predict(X_train)
    dev_pred   = predict(X_dev)
    eval_pred  = predict(X_eval)

    
    # Save hyp files
    save_hyp(train_pred, OUTPUT_DIR / "hyp_hier_rf_train.csv")
    save_hyp(dev_pred,   OUTPUT_DIR / "hyp_hier_rf_dev.csv")
    save_hyp(eval_pred,  OUTPUT_DIR / "hyp_hier_rf_eval.csv")

    print("\nSaved to:")
    print(OUTPUT_DIR / "hyp_hier_rf_train.csv")
    print(OUTPUT_DIR / "hyp_hier_rf_dev.csv")
    print(OUTPUT_DIR / "hyp_hier_rf_eval.csv")

    print("\nDone.")