from sklearn.linear_model import LinearRegression
from sklearn.ensemble import HistGradientBoostingRegressor as HGBR
from metacausal import CausalEnsemble
from metacausal.adapters import GenericCATEAdapter
from metacausal.datasets import load_lalonde
from metacausal.plots import disagreement

X, T, Y = load_lalonde()

def fit_linear(X, T, Y, **kwargs):
    treated = T == 1
    m1 = LinearRegression().fit(X[treated], Y[treated])
    m0 = LinearRegression().fit(X[~treated], Y[~treated])
    return (m1, m0)

def fit_hgb(X, T, Y, **kwargs):
    treated = T == 1
    m1 = HGBR(max_iter=20).fit(X[treated], Y[treated])
    m0 = HGBR(max_iter=20).fit(X[~treated], Y[~treated])
    return (m1, m0)

def cate_fn(state, X):
    m1, m0 = state
    return m1.predict(X) - m0.predict(X)

methods = [
    GenericCATEAdapter(fit_linear, cate_fn, name="linear"),
    GenericCATEAdapter(fit_hgb, cate_fn, name="hgb"),
]
ens = CausalEnsemble(methods=methods)
ens.fit(X, T, Y, random_state=42)
disagreement(ens, X)