# Liam McLaughlin
import numpy as np
import pandas as pd
from sklearn.preprocessing import StandardScaler
from sklearn.decomposition import PCA
from sklearn.ensemble import RandomForestClassifier, ExtraTreesClassifier
from sklearn.linear_model import LogisticRegression
from sklearn.model_selection import StratifiedKFold
from sklearn.metrics import confusion_matrix
import lightgbm as lgb
import warnings
warnings.filterwarnings('ignore')

KEEP_LABELS    = {0, 2, 3, 5, 6, 8}
LABEL_REMAP    = {0: 0, 2: 1, 3: 2, 5: 3, 6: 4, 8: 5}
LABEL_UNMAP    = {v: k for k, v in LABEL_REMAP.items()}
SCORED_CLASSES = [0, 1, 2, 3, 4]
BCKG_CLASS     = 5
N_CLASSES      = 6

RF_BEST_PARAMS = {
    'n_estimators':      410,
    'max_depth':         16,
    'min_samples_leaf':  18,
    'min_samples_split': 68,
    'max_features':      0.37449357064743527,
}

ET_PARAMS = {
    'n_estimators':      372,
    'max_depth':         20,
    'min_samples_leaf':  7,
    'min_samples_split': 70,
    'max_features':      0.4623099320697791,
}

LGBM_BEST_PARAMS = {
    'num_leaves':        18,
    'max_depth':         14,
    'min_child_samples': 162,
    'reg_alpha':         0.0023895517265508526,
    'reg_lambda':        0.07539017329138864,
    'feature_fraction':  0.7235069495340481,
    'bagging_fraction':  0.4768161280685114,
    'bagging_freq':      4,
    'drop_rate':         0.04033569589418706,
    'skip_drop':         0.49394509073650067,
    'learning_rate':     0.012575730000740968,
}

# Computes the weighted competition score from per-class error rates on scored classes and the background class.
def compute_score(y_true, y_pred):
    cm = confusion_matrix(y_true, y_pred, labels=list(range(N_CLASSES)))
    errors = []
    for c in SCORED_CLASSES:
        row = cm[c]; total = row.sum()
        err = (total - cm[c, c]) / total if total > 0 else 0.0
        errors.append(err)
    avg_lbl  = np.mean(errors)
    bt       = cm[BCKG_CLASS].sum()
    bckg_err = (bt - cm[BCKG_CLASS, BCKG_CLASS]) / bt if bt > 0 else 0.0
    score    = 0.9 * avg_lbl + 0.1 * bckg_err
    return avg_lbl * 100, bckg_err * 100, score * 100

# Loads a CSV file and filters rows to the kept labels, remapping them to contiguous class indices.
def load_and_filter(path, is_eval=False):
    print(f"Loading {path} ...")
    df = pd.read_csv(path, comment='#', header=None)
    label_col    = 0
    feature_cols = list(range(1, df.shape[1]))
    if is_eval:
        X = df[feature_cols].values.astype(np.float32)
        return X, np.zeros(len(df), dtype=int), np.arange(len(df)), len(df)
    mask = df[label_col].isin(KEEP_LABELS)
    df_f = df[mask].reset_index(drop=True)
    idx  = np.where(mask.values)[0]
    y    = np.array([LABEL_REMAP[l] for l in df_f[label_col].values], dtype=int)
    X    = df_f[feature_cols].values.astype(np.float32)
    print(f"  {len(X)} samples | {dict(zip(*np.unique(y, return_counts=True)))}")
    return X, y, idx, len(df)

# Fits a StandardScaler and PCA on the training features and returns both transformers.
def fit_preprocessors(X_train, n_pca=471):
    print(f"\nFitting preprocessors...")
    scaler = StandardScaler()
    X_sc   = scaler.fit_transform(X_train)
    pca    = PCA(n_components=n_pca, random_state=42)
    pca.fit(X_sc)
    print(f"  PCA({n_pca}) variance explained: {pca.explained_variance_ratio_.sum()*100:.2f}%")
    return scaler, pca

# Applies scaler then PCA to produce the feature representation used by the LGBM model.
def get_lgbm_features(X, scaler, pca):
    return pca.transform(scaler.transform(X))

# Applies only the scaler to produce the feature representation used by the tree-based RF/ET models.
def get_rf_features(X, scaler):
    return scaler.transform(X)

# Computes inverse-frequency class weights to balance the loss across imbalanced classes.
def inverse_class_weights(y):
    counts = np.bincount(y, minlength=N_CLASSES)
    return np.where(counts > 0, len(y) / (N_CLASSES * counts), 1.0)

# Maps each label in y to its corresponding class weight, producing a per-sample weight array.
def sample_weights(y, cw):
    return np.array([cw[c] for c in y])

# Applies temperature scaling to a probability matrix and returns renormalized probabilities.
def apply_temperature(probs, T):
    lp = np.log(probs + 1e-10) / T
    ep = np.exp(lp - lp.max(axis=1, keepdims=True))
    return ep / ep.sum(axis=1, keepdims=True)

# Grid-searches the temperature value that minimizes the competition score after softmax rescaling.
def find_temperature(probs, y_true):
    print("\n[Post] Temperature scaling...")
    best_T, best_s = 1.0, float('inf')
    for T in np.arange(0.5, 3.0, 0.05):
        lp  = np.log(probs + 1e-10) / T
        ep  = np.exp(lp - lp.max(axis=1, keepdims=True))
        cal = ep / ep.sum(axis=1, keepdims=True)
        _, _, s = compute_score(y_true, np.argmax(cal, axis=1))
        if s < best_s:
            best_s, best_T = s, T
    print(f"  Best T={best_T:.2f} -> {best_s:.4f}%")
    return best_T

# Greedy per-class threshold search that divides probabilities by class thresholds to minimize the score.
def tune_thresholds(probs, y_true):
    print("\n[Post] Per-class threshold tuning...")
    best_thr, best_s = np.ones(N_CLASSES), float('inf')
    for c in range(N_CLASSES):
        for t in np.arange(0.3, 2.0, 0.05):
            thr   = best_thr.copy(); thr[c] = t
            preds = np.argmax(probs / thr[np.newaxis, :], axis=1)
            _, _, s = compute_score(y_true, preds)
            if s < best_s:
                best_s = s; best_thr = thr.copy()
    print(f"  Thresholds: {np.round(best_thr, 3)} -> {best_s:.4f}%")
    return best_thr

# Produces final class predictions by argmax over threshold-rescaled probabilities.
def predict_final(probs, thr):
    return np.argmax(probs / thr[np.newaxis, :], axis=1)

# Maps remapped predictions back to original label space and writes the full-length submission CSV.
def write_submission(preds, keep_idx, total_len, path):
    output = np.zeros(total_len, dtype=int)
    for i, p in zip(keep_idx, preds):
        output[i] = LABEL_UNMAP[p]
    with open(path, 'w') as f:
        f.write("label\n")
        for label in output:
            f.write(f"{label}\n")
    print(f"  Saved: {path}")

# ── Base model trainers ────────────────────────────────────────────────────

# Trains a Random Forest classifier with the tuned hyperparameters and balanced class weights.
def train_rf(X_tr, y_tr):
    rf = RandomForestClassifier(**RF_BEST_PARAMS, class_weight='balanced',
                                 n_jobs=-1, random_state=132)
    rf.fit(X_tr, y_tr)
    return rf

# Trains an Extra Trees classifier with the tuned hyperparameters and balanced class weights.
def train_et(X_tr, y_tr):
    et = ExtraTreesClassifier(**ET_PARAMS, class_weight='balanced',
                               n_jobs=-1, random_state=99)
    et.fit(X_tr, y_tr)
    return et

# Trains a LightGBM DART multiclass model with sample weights derived from inverse class frequencies.
def train_lgbm(X_tr, y_tr, n_rounds=452):
    cw = inverse_class_weights(y_tr)
    sw = sample_weights(y_tr, cw)
    bp = {'objective': 'multiclass', 'num_class': N_CLASSES,
          'boosting_type': 'dart', 'metric': 'multi_logloss',
          'verbosity': -1, 'seed': 42, **LGBM_BEST_PARAMS}
    dt = lgb.Dataset(X_tr, label=y_tr, weight=sw)
    return lgb.train(bp, dt, num_boost_round=n_rounds,
                     callbacks=[lgb.log_evaluation(-1)])

# ── Stacking ───────────────────────────────────────────────────────────────

# Generates out-of-fold predictions from RF, ET, and LGBM via stratified K-fold and stacks them as meta-features.
def build_oof_features(X_rf, X_lgbm, y, n_splits=5):
    """
    Generate out-of-fold predictions from RF, ET, LGBM.
    Returns stacked feature matrix of shape (n_samples, N_CLASSES * 3).
    """
    print(f"\n[Stack] Building out-of-fold features ({n_splits} folds)...")
    n = len(y)
    oof_rf   = np.zeros((n, N_CLASSES))
    oof_et   = np.zeros((n, N_CLASSES))
    oof_lgbm = np.zeros((n, N_CLASSES))

    skf = StratifiedKFold(n_splits=n_splits, shuffle=True, random_state=42)

    for fold, (tr_idx, val_idx) in enumerate(skf.split(X_rf, y)):
        print(f"  Fold {fold+1}/{n_splits}...")

        # RF
        rf = RandomForestClassifier(**RF_BEST_PARAMS, class_weight='balanced',
                                     n_jobs=-1, random_state=132)
        rf.fit(X_rf[tr_idx], y[tr_idx])
        oof_rf[val_idx] = rf.predict_proba(X_rf[val_idx])

        # ET
        et = ExtraTreesClassifier(**ET_PARAMS, class_weight='balanced',
                                   n_jobs=-1, random_state=99)
        et.fit(X_rf[tr_idx], y[tr_idx])
        oof_et[val_idx] = et.predict_proba(X_rf[val_idx])

        # LGBM
        lgbm_m = train_lgbm(X_lgbm[tr_idx], y[tr_idx])
        oof_lgbm[val_idx] = lgbm_m.predict(X_lgbm[val_idx])

    print("  OOF features built.")
    return np.hstack([oof_rf, oof_et, oof_lgbm])  # (n, 18)

# Trains a logistic-regression meta-learner on the stacked OOF probability features.
def train_meta(oof_features, y):
    print("\n[Stack] Training meta-learner (Logistic Regression)...")
    meta = LogisticRegression(
        max_iter=1000,
        C=1.0,
        class_weight='balanced',
        random_state=42,
        n_jobs=-1
    )
    meta.fit(oof_features, y)
    return meta

# Runs all three base models on the input features and concatenates their probability outputs.
def get_base_probs(rf, et, lgbm_m, X_rf, X_lgbm):
    p_rf   = rf.predict_proba(X_rf)
    p_et   = et.predict_proba(X_rf)
    p_lgbm = lgbm_m.predict(X_lgbm)
    return np.hstack([p_rf, p_et, p_lgbm])  # (n, 18)

# Orchestrates the full pipeline: load data, fit preprocessors, OOF stack, calibrate, retrain on train+dev, and write submissions.
def main():
    TRAIN_PATH = "../data/train.csv"
    DEV_PATH   = "../data/dev.csv"
    EVAL_PATH  = "../data/eval.csv"

    X_tr,   y_tr,   tr_idx,   tr_len   = load_and_filter(TRAIN_PATH)
    X_dev,  y_dev,  dev_idx,  dev_len  = load_and_filter(DEV_PATH)
    X_eval, y_eval, eval_idx, eval_len = load_and_filter(EVAL_PATH, is_eval=True)

    print(f"\nSample counts — train: {len(X_tr)}, dev: {len(X_dev)}, eval: {len(X_eval)}")

    scaler, pca = fit_preprocessors(X_tr, n_pca=471)

    X_tr_lgbm   = get_lgbm_features(X_tr,   scaler, pca)
    X_dev_lgbm  = get_lgbm_features(X_dev,  scaler, pca)
    X_eval_lgbm = get_lgbm_features(X_eval, scaler, pca)

    X_tr_rf   = get_rf_features(X_tr,   scaler)
    X_dev_rf  = get_rf_features(X_dev,  scaler)
    X_eval_rf = get_rf_features(X_eval, scaler)

    # ── Phase 1: OOF stacking on train ────────────────────────────────────
    oof_features = build_oof_features(X_tr_rf, X_tr_lgbm, y_tr, n_splits=5)
    meta = train_meta(oof_features, y_tr)

    # ── Phase 1: Train full base models on train only ──────────────────────
    print("\n[Base] Training full base models on train only...")
    rf_model   = train_rf(X_tr_rf,   y_tr)
    et_model   = train_et(X_tr_rf,   y_tr)
    lgbm_model = train_lgbm(X_tr_lgbm, y_tr)

    # ── Phase 1: Evaluate on dev (no leakage) ─────────────────────────────
    dev_base   = get_base_probs(rf_model, et_model, lgbm_model,
                                X_dev_rf, X_dev_lgbm)
    dev_probs  = meta.predict_proba(dev_base)

    _, _, s = compute_score(y_dev, np.argmax(dev_probs, axis=1))
    print(f"\n  Stacked dev score (pre-calibration): {s:.4f}%")

    best_T   = find_temperature(dev_probs, y_dev)
    dev_cal  = apply_temperature(dev_probs, best_T)
    best_thr = tune_thresholds(dev_cal, y_dev)

    dev_preds = predict_final(dev_cal, best_thr)
    l, b, s   = compute_score(y_dev, dev_preds)
    print(f"\n== DEV SCORE (no leakage): {s:.4f}%  (lbl={l:.2f}%, bckg={b:.2f}%)")

    # ── Phase 2: Retrain on train+dev for final submission ─────────────────
    print("\n--- Retraining on train+dev for final submission ---")
    X_all      = np.vstack([X_tr, X_dev])
    y_all      = np.concatenate([y_tr, y_dev])
    X_all_lgbm = get_lgbm_features(X_all, scaler, pca)
    X_all_rf   = get_rf_features(X_all,   scaler)

    oof_all  = build_oof_features(X_all_rf, X_all_lgbm, y_all, n_splits=5)
    meta_all = train_meta(oof_all, y_all)

    rf_final   = train_rf(X_all_rf,   y_all)
    et_final   = train_et(X_all_rf,   y_all)
    lgbm_final = train_lgbm(X_all_lgbm, y_all)

    # Eval predictions
    eval_base  = get_base_probs(rf_final, et_final, lgbm_final,
                                X_eval_rf, X_eval_lgbm)
    eval_probs = meta_all.predict_proba(eval_base)
    eval_cal   = apply_temperature(eval_probs, best_T)
    eval_preds = predict_final(eval_cal, best_thr)

    # Train predictions
    tr_base  = get_base_probs(rf_final, et_final, lgbm_final,
                              X_all_rf[:len(X_tr_rf)],
                              X_all_lgbm[:len(X_tr_lgbm)])
    tr_probs = meta_all.predict_proba(tr_base)
    tr_cal   = apply_temperature(tr_probs, best_T)
    tr_preds = predict_final(tr_cal, best_thr)

    print("\n== FINAL SCORES ==")
    l, b, s = compute_score(y_tr, tr_preds)
    print(f"  Train: {s:.4f}%  (lbl={l:.2f}%, bckg={b:.2f}%)")

    print("\nWriting submission files...")
    write_submission(tr_preds, tr_idx, tr_len, "hyp_train_nonneural.csv")

    # Dev: use phase-1 models (no leakage)
    dev_final_base  = get_base_probs(rf_model, et_model, lgbm_model,
                                     X_dev_rf, X_dev_lgbm)
    dev_final_probs = meta.predict_proba(dev_final_base)
    dev_final_cal   = apply_temperature(dev_final_probs, best_T)
    dev_final_preds = predict_final(dev_final_cal, best_thr)
    write_submission(dev_final_preds, dev_idx, dev_len, "hyp_dev_nonneural.csv")

    write_submission(eval_preds, eval_idx, eval_len, "hyp_eval_nonneural.csv")
    print("Done!")

if __name__ == "__main__":
    main()