# Regularised Regression: A Machine Learning Toolkit for Econometrics
# Athanassios Stavrakoudis
# astavrak@uoi.gr
# with claude's assistance
#
# Complete Python code for the Elastic Net part
# Standalone: libraries imported and configuration hard-coded below.

import numpy as np
import pandas as pd
import wooldridge as woo
import matplotlib.pyplot as plt
from sklearn.linear_model import (Lasso, LassoCV, Ridge, RidgeCV,
                                  ElasticNet, ElasticNetCV, LinearRegression)
from sklearn.preprocessing import StandardScaler
from sklearn.pipeline import Pipeline
from sklearn.model_selection import train_test_split, cross_val_score
from sklearn.metrics import mean_squared_error, r2_score
import warnings
warnings.filterwarnings("ignore")

SEED       = 14159
N_CORES    = 6
N_CV_FOLDS = 10
N_ALPHAS   = 100
rng = np.random.default_rng(SEED)


# [py-enet-wagepan]
import pandas as pd, numpy as np, warnings
from sklearn.linear_model import LassoCV, RidgeCV, ElasticNetCV, LinearRegression
from sklearn.preprocessing import StandardScaler
warnings.filterwarnings('ignore')

# - Rebuild wagepan data (self-contained — no cross-chunk dependencies)
try:
    import wooldridge as woo
    wp_en_py = woo.data("wagepan")
except Exception:
    wp_en_py = pd.read_csv("../data/wagepan.csv")   # local copy written by R setup

ctrl_cols = [c for c in ["exper","expersq","married","educ","black","hisp","south","smsa",
                          "agric","bus","construc","ndurman","trcommpu","trade",
                          "services","profserv","profocc","clerocc","servocc"]
             if c in wp_en_py.columns]
use_cols  = ["nr","year","lwage","union"] + ctrl_cols
wp_use    = wp_en_py[[c for c in use_cols if c in wp_en_py.columns]].dropna()

# Within-demean (remove individual FE)
num_cols = [c for c in wp_use.columns if c not in ["nr","year"]]
wp_dm_en = wp_use.copy()
wp_dm_en[num_cols] = (wp_use[num_cols]
                      - wp_use.groupby("nr")[num_cols].transform("mean"))

yr_dum_en = pd.get_dummies(wp_dm_en["year"], prefix="yr", drop_first=True).astype(float)
ctrl_all_en = pd.concat([wp_dm_en[ctrl_cols], yr_dum_en], axis=1).values
y_en  = wp_dm_en["lwage"].values
D_en  = wp_dm_en["union"].values

sc_en  = StandardScaler().fit(ctrl_all_en)
X_sc_en = sc_en.transform(ctrl_all_en)

# - TWFE baseline via demeaned OLS
X_twfe_en = np.column_stack([D_en, ctrl_all_en])
b_twfe_en = np.linalg.lstsq(X_twfe_en, y_en, rcond=None)[0][0]

# - Lasso (rebuild)
las_y = LassoCV(cv=N_CV_FOLDS, max_iter=5000, n_jobs=N_CORES).fit(X_sc_en, y_en)
las_d = LassoCV(cv=N_CV_FOLDS, max_iter=5000, n_jobs=N_CORES).fit(X_sc_en, D_en)
sel_y_en  = set(np.where(las_y.coef_ != 0)[0])
sel_d_en  = set(np.where(las_d.coef_ != 0)[0])
sel_u_en  = sorted(sel_y_en | sel_d_en)
yres_l_en = y_en - las_y.predict(X_sc_en)
Dres_l_en = D_en - las_d.predict(X_sc_en)
b_las_en  = np.dot(Dres_l_en, yres_l_en) / np.dot(Dres_l_en, Dres_l_en)
psi_l     = Dres_l_en*(yres_l_en - b_las_en*Dres_l_en)
se_las_en = ((Dres_l_en**2).mean()**(-2)*(psi_l**2).mean()/len(y_en))**0.5

# - Ridge (rebuild)
rid_y = RidgeCV(cv=N_CV_FOLDS).fit(X_sc_en, y_en)
rid_d = RidgeCV(cv=N_CV_FOLDS).fit(X_sc_en, D_en)
yres_r_en = y_en - rid_y.predict(X_sc_en)
Dres_r_en = D_en - rid_d.predict(X_sc_en)
b_rid_en  = np.dot(Dres_r_en, yres_r_en) / np.dot(Dres_r_en, Dres_r_en)
psi_r     = Dres_r_en*(yres_r_en - b_rid_en*Dres_r_en)
se_rid_en = ((Dres_r_en**2).mean()**(-2)*(psi_r**2).mean()/len(y_en))**0.5
d2_en     = np.linalg.svd(X_sc_en, compute_uv=False)**2
df_en     = np.sum(d2_en/(d2_en + rid_y.alpha_))

# - Elastic Net
l1s   = [0.1, 0.25, 0.5, 0.75, 0.9]
en_y  = ElasticNetCV(l1_ratio=l1s, n_alphas=N_ALPHAS, cv=N_CV_FOLDS,
                      max_iter=10000, n_jobs=N_CORES).fit(X_sc_en, y_en)
en_d  = ElasticNetCV(l1_ratio=l1s, n_alphas=N_ALPHAS, cv=N_CV_FOLDS,
                      max_iter=10000, n_jobs=N_CORES).fit(X_sc_en, D_en)

yres_en_py = y_en - en_y.predict(X_sc_en)
Dres_en_py = D_en - en_d.predict(X_sc_en)
b_en_py    = np.dot(Dres_en_py, yres_en_py) / np.dot(Dres_en_py, Dres_en_py)
psi_en     = Dres_en_py*(yres_en_py - b_en_py*Dres_en_py)
se_en_py   = ((Dres_en_py**2).mean()**(-2)*(psi_en**2).mean()/len(y_en))**0.5
nsel_en_py = (en_y.coef_ != 0).sum()

print(f"EN α* = {en_y.l1_ratio_:.2f}  λ* = {en_y.alpha_:.4f}")
print(f"Controls selected: {nsel_en_py}/{X_sc_en.shape[1]}")
print(f"EN union premium: {b_en_py:.4f}  SE: {se_en_py:.4f}")

from tabulate import tabulate
print(tabulate([
    ["TWFE",        f"{b_twfe_en:.4f}", "demeaned OLS",         "—",     "—"],
    ["Post-Lasso",  f"{b_las_en:.4f}",  f"{len(sel_u_en)} sel.","1.00",  f"{las_y.alpha_:.4f}"],
    ["Ridge",       f"{b_rid_en:.4f}",  f"df {df_en:.1f}",      "0.00",  f"{rid_y.alpha_:.4f}"],
    ["Elastic Net", f"{b_en_py:.4f}",   f"{nsel_en_py} sel.",
                                         f"{en_y.l1_ratio_:.2f}", f"{en_y.alpha_:.4f}"]],
    headers=["Estimator","Union premium","Controls/df","α*","λ*"],
    tablefmt="rounded_outline"))
