import numpy as np
import jax.numpy as jnp
from lime.lime_tabular import LimeTabularExplainer
from scipy.stats import gaussian_kde
from sklearn.base import BaseEstimator, RegressorMixin
from sklearn.metrics import euclidean_distances, f1_score, accuracy_score, root_mean_squared_error
from sklearn.model_selection import train_test_split
from sklearn.preprocessing import MinMaxScaler

from AugLagExemplarSolver import AugLagExemplarSolver
from data_load import load_data
from exemplar_animation import plot_exemplars, plotly_animate, plotly_animate
from model_generator import get_calibrated_gbm
from objectives import rbf_sim
from lex_utils import load_sine_data, fit_weighted_lars, fit_weighted_least_squares, fit_weighted_ridge, \
    detect_outliers_iqr, fit_weighted_lasso_path
from rfutils import compute_treebased_sim
from sklearn.metrics import mean_squared_error, r2_score
from sklearn.linear_model import Ridge


def lime_explain_instance_with_data(lime_explainer,
                               neighborhood_data,
                               neighborhood_labels,
                               weights,
                               label,
                               num_features,
                               feat_selection='auto',
                               model_regressor=None,
                               random_state=None,
                                    ):
    """Takes perturbed data, labels and distances, returns explanation."""
    # labels_column = neighborhood_labels[:, label]
    labels_column = neighborhood_labels
    used_features = lime_explainer.base.feature_selection(neighborhood_data,
                                           labels_column,
                                           weights,
                                           num_features,
                                           feat_selection)
    if model_regressor is None:
        model_regressor = Ridge(alpha=1, fit_intercept=True,
                                random_state=random_state)
    easy_model = model_regressor
    easy_model.fit(neighborhood_data[:, used_features],
                   labels_column, sample_weight=weights)
    prediction_score = easy_model.score(
        neighborhood_data[:, used_features],
        labels_column, sample_weight=weights)

    local_pred = easy_model.predict(neighborhood_data[0, used_features].reshape(1, -1))

    return (easy_model.intercept_,
            sorted(zip(used_features, easy_model.coef_),
                   key=lambda x: np.abs(x[1]), reverse=True),
            prediction_score, local_pred)


class LexRegressor(BaseEstimator, RegressorMixin):
    def __init__(self, K=50, kernel='rfsim', batching='tree-based',
                 sim_cutoff=0.15, dist_cutoff_perc=5,
                 qr_dec=True, batch_size=20,
                 num_outer_iter=10, num_epoch=3,
                 grad_tol=1e-10,
                 lambda_init=0.1, penalty_init=1.0, penalty_multiplier=10.0,
                 rf_n_estimators=20,
                 opt_name="optax-adam",
                 adam_init_lr=1e-3,
                 random_state=None,
                 logger=None,
                 dataset_name='na',
                 verbose=1,
                 positive_class=1,
                 bb_fn=None,
                 include_exp_instance=False,
                 kernel_scale=1,
                 num_extra_exemplars=0,
                 log_fit_results=False,
                 ):
        self.K = K
        self.kernel = kernel
        self.batching = batching
        self.sim_cutoff = sim_cutoff
        self.dist_cutoff_perc = dist_cutoff_perc
        self.qr_dec = qr_dec
        self.batch_size = batch_size
        self.num_outer_iter = num_outer_iter
        self.num_epoch = num_epoch
        self.grad_tol = grad_tol
        self.lambda_init = lambda_init
        self.penalty_init = penalty_init
        self.penalty_multiplier = penalty_multiplier
        self.rf_n_estimators = rf_n_estimators
        self.opt_name = opt_name
        self.adam_init_lr = adam_init_lr
        self.random_state = random_state
        self.logger = logger
        self.dataset_name = dataset_name
        self.verbose = verbose
        self.positive_class = positive_class
        self.X_extra, self.y_extra = None, None
        self.num_extra_exemplars = num_extra_exemplars
        self.bb_fn = bb_fn
        self.include_exp_instance = include_exp_instance
        self.kernel_scale = kernel_scale
        self.log_fit_results = log_fit_results

    def fit(self, X, y):
        if self.num_extra_exemplars < 0:
            raise ValueError("num_extra_exemplars must be non-negative.")

        self.X_extra, self.y_extra = None, None
        X_train, X_repr, y_train, y_repr = train_test_split(X, y, test_size=0.5, random_state=self.random_state)
        self.X_repr = X_repr
        self.y_repr = y_repr

        learning_params = {
            "batch_size": self.batch_size,
            "num_outer_iter": self.num_outer_iter,
            "num_epoch": self.num_epoch,
            "grad_tol": self.grad_tol,
            "lambda_init": self.lambda_init,
            "penalty_init": self.penalty_init,
            "penalty_multiplier": self.penalty_multiplier,
        }
        self.lexsolver = AugLagExemplarSolver(X_train, X_repr, y_train, y_repr,
                                              kernel=self.kernel,
                                              K=self.K,
                                              sim_cutoff=self.sim_cutoff,
                                              dist_cutoff_perc=self.dist_cutoff_perc,
                                              qr_dec=self.qr_dec,
                                              learning_params=learning_params,
                                              batching=self.batching,
                                              opt_name=self.opt_name,
                                              adam_init_lr=self.adam_init_lr,
                                              rf_n_estimators=self.rf_n_estimators,
                                              random_state=self.random_state,
                                              verbose=self.verbose,
                                              kernel_scale=self.kernel_scale
                                              )

        try:
            self.lexsolver.run()
            pass
        except Exception as e:
            raise e
        self.X_exemplars, self.y_exemplars = self.lexsolver.get_K_exemplars()
        if(self.X_exemplars.size ==0):
            raise ValueError("get_K_exemplars() returned empty exemplars.")

        self.K_found = self.X_exemplars.shape[0]
        self.bitvectorness = self.lexsolver.bitvectorness

        if self.kernel == 'rbf':
            self.gamma = self.lexsolver.get_gamma()
            self.dist_cutoff = self.lexsolver.dist_cutoff
        elif self.kernel == 'rfsim':
            self.rf_model = self.lexsolver.get_RF_model()
        else:
            raise ValueError("Invalid kernel. Choose from ['rbf', 'rfsim'].")
        if self.log_fit_results:
            self.log_results(self.dataset_name)

        # before we sample based on scores we need to remove the actual exemplars
        # from consideration - otherwise they might get picked again
        idxs_exemplars = self.lexsolver.get_K_exemplar_idx()
        all_idxs = np.arange(X_train.shape[0])
        rem_idxs = np.setdiff1d(all_idxs, idxs_exemplars, assume_unique=False)
        all_exemplars = self.X_exemplars

        if self.num_extra_exemplars > 0:
            self.X_extra, self.y_extra = self._select_extra_exemplars(X_train, y_train, rem_idxs)
            if self.X_extra is not None:
                all_exemplars = np.vstack((all_exemplars, self.X_extra))

        self.lime_explainer = LimeTabularExplainer(
            training_data=all_exemplars,
            mode="regression",
            random_state=self.random_state,
        )

        return self

    def _select_extra_exemplars(self, X_train, y_train, rem_idxs):
        if len(rem_idxs) == 0:
            return None, None

        X_candidates = X_train[rem_idxs]
        n_extra = min(self.num_extra_exemplars, len(rem_idxs))
        try:
            exemplars_kde = gaussian_kde(self.X_exemplars.T)
            scores = exemplars_kde.pdf(X_candidates.T)
        except (np.linalg.LinAlgError, ValueError):
            distances = euclidean_distances(X_candidates, self.X_exemplars)
            scores = -np.min(distances, axis=1)

        best_candidate_idxs = np.argsort(scores)[-n_extra:]
        extra_idxs = rem_idxs[best_candidate_idxs]
        return X_train[extra_idxs], y_train[extra_idxs]

    def compute_local_coefficients(self, X, num_features=5, method="lassopath", compute_lime_inference=True):
        """
        method: "lars"  -> fit_weighted_lars()
                "lstsq" -> fit_weighted_least_squares()
        compute_lime_inference: whether to also fit the LIME-style local inference model.
        """

        # Gather exemplars
        all_exemplars = self.X_exemplars
        all_y_exemplars = self.y_exemplars
        if self.X_extra is not None:
            all_exemplars = np.vstack((all_exemplars, self.X_extra))
            all_y_exemplars = np.concatenate([all_y_exemplars, self.y_extra])

        # Compute similarity weights W
        if self.kernel == 'rbf':
            euc_dist = euclidean_distances(X, all_exemplars)
            W = rbf_sim(euc_dist, self.gamma)
            W = np.where(euc_dist < self.dist_cutoff, W, 0)

        elif self.kernel == 'rfsim':
            # treebased_sim = compute_treebased_sim(X, all_exemplars, self.rf_model)
            #todo: replace this with a code block to call the kernel smoothing method
            # treebased_sim = self.lexsolver.sim_mat[:,self.lexsolver.get_K_exemplar_idx()]
            euc_dist = euclidean_distances(X, all_exemplars)
            cluster_ids = self.rf_model.estimators_[0].apply(X)
            W = self.lexsolver.kernel_smoother.smooth_kernel_for_data(euc_dist, cluster_ids)
            W = np.where(W >= self.sim_cutoff, W, 0)

        else:
            raise ValueError("Invalid kernel. Choose from ['rbf', 'rfsim'].")

        # Loop over samples
        explanations = []
        unexplained = []
        explanations_lime_inference = []
        # W = jnp.asarray(W)
        for i in range(len(W)):
            # Include the instance itself
            if self.include_exp_instance:
                if self.bb_fn is None:
                    raise ValueError("bb_fn is required when include_exp_instance=True.")
                neighbors_Xi = np.concatenate([all_exemplars, X[i:i + 1]])
                exp_instance_y = self.bb_fn(X[i:i + 1])[0][self.positive_class]
                neighbors_yi = np.append(all_y_exemplars, exp_instance_y)
                sample_weights = np.append(W[i], 1)
            else:
                neighbors_Xi = all_exemplars
                neighbors_yi = all_y_exemplars
                sample_weights = W[i]

            # Mark unexplained if too few neighbors
            if np.count_nonzero(sample_weights) <= X.shape[1]:
                unexplained.append(True)
                sample_weights = sample_weights + 1
            else:
                unexplained.append(False)

            eps = 1e-8
            sample_weights = sample_weights / (np.sum(sample_weights) + eps)
            W[i] = sample_weights[:W.shape[1]]

            if method == "lars":
                result = fit_weighted_lars(neighbors_Xi, neighbors_yi, sample_weights, num_features)
            elif method == "lstsq":
                result = fit_weighted_least_squares(neighbors_Xi, neighbors_yi, sample_weights)
            elif method == "lassopath":
                result = fit_weighted_lasso_path(neighbors_Xi, neighbors_yi, sample_weights, num_features)
            else:
                raise ValueError("method must be 'lars', 'lstsq', or 'lassopath'")

            explanations.append(result)

            if compute_lime_inference:
                lime_result = lime_explain_instance_with_data(
                    self.lime_explainer,
                    neighbors_Xi,
                    neighbors_yi,
                    weights=np.sqrt(sample_weights),
                    label=self.positive_class,
                    feat_selection="lasso_path",
                    num_features=num_features,
                    random_state=self.random_state
                )
            else:
                lime_result = None
            explanations_lime_inference.append(lime_result)

        explanations = np.stack(explanations)
        unexplained_indices = [i for i, f in enumerate(unexplained) if f]

        print(f"Unexplained sample indices: {unexplained_indices}")

        if len(unexplained_indices) == len(explanations):
            raise ValueError("No explanations found. Try adjusting thresholds.")

        return explanations, unexplained_indices, W, explanations_lime_inference

    def predict(self, X, num_features=5):
        coeff, _, _, _ = self.compute_local_coefficients(
            X,
            num_features,
            compute_lime_inference=False,
        )
        y_pred = self.predict_from_coefficients(X, coeff)
        return y_pred

    def predict_from_coefficients(self, X, coeff):
        X_ext = np.hstack([X, np.ones((X.shape[0], 1))])
        return np.sum(coeff * X_ext, axis=1)

    def log_results(self, dataset_name):
        self.lexsolver.log_results(dataset_name=dataset_name)


def demo_main():
    # data_set_name = 'yacht_hydrodynamics'
    # expt_desc = 'yacht, RF-based sim, cutoff=50percentile, masking'
    # X_train, x_repr, y_train, y_repr, data_model = data_loader(data_set_name)

    is_classification = True
    if(is_classification):
        dataset_name = 'cod-rna' # datasets tested: [cod-rna, ijcnn1, covtype, letter, pendigits]
        X, y_label = load_data(dataset_name)
        X, _, y_label, _ = train_test_split(X, y_label, train_size=2000, stratify=y_label)
        scaler = MinMaxScaler()
        X = scaler.fit_transform(X)
        calibrated_clf = get_calibrated_gbm(X, y_label)
        X_train, X_test, _, _ = train_test_split(X, y_label, test_size=0.2, random_state=42,
                                                                      stratify=y_label)
        y_train = calibrated_clf.predict_proba(X_train)[:, 1]
        y_test = calibrated_clf.predict_proba(X_test)[:, 1]

        # gbm_acc = f1_score(y_test, calibrated_clf.predict(X_test), average='macro')

    # todo: handle regressor properly. Right now, this is a closed form function. NO model being learned.
    else:
        dataset_name = 'sine'
        X, y = load_sine_data(n=1000, noise_std=0.0)
        X_train, X_test, y_train, y_test = train_test_split(X, y, test_size=0.2, random_state=42)

        scaler = MinMaxScaler()
        X_train = scaler.fit_transform(X_train)
        X_test = scaler.transform(X_test)

        idx = np.argsort(X_test[:, 0])
        X_test = X_test[idx]
        y_test = y_test[idx]

    lex_params = {
        "K": 150,
        "kernel": "rfsim",  # should be one of ['rbf', 'rfsim']
        "batching": "tree-based",  # should be one of ['standard', 'tree-based']
        "sim_cutoff": 0.5,
        "dist_cutoff_perc": 20,  # Percentile, Range 0...100
        "qr_dec": True,
        "batch_size": 100,
        "num_outer_iter": 10,
        "num_epoch": 3,
        "grad_tol": 1e-10,
        "lambda_init": 0.1,
        "penalty_init": 1.0,
        "penalty_multiplier": 10.0,
        "adam_init_lr": 3e-2,
        "kernel_scale": 1,
    }

    model = LexRegressor(**lex_params, dataset_name=dataset_name)
    model.fit(X_train, y_train)

    # temp, delete this
    # X_test = model.X_repr
    # y_test = model.y_repr

    # idx = np.argsort(X_test[:, 0])
    # X_test = X_test[idx]
    # y_test = y_test[idx]

    y_pred = model.predict(X_test)
    print("RMSE:", root_mean_squared_error(y_test, y_pred))
    print("R2:", r2_score(y_test, y_pred))
    model.log_results(dataset_name=dataset_name)
    coeff, unexplained_indices, W, lime_inference = model.compute_local_coefficients(X_test)
    mask = np.ones(len(y_test), dtype=bool)
    mask[unexplained_indices] = False

    y_intercept, feat_weights, _, _ = zip(*lime_inference)
    lime_coeffs = np.zeros((len(feat_weights), coeff.shape[1]))
    for i, feat_weights in enumerate(feat_weights):
        for feat_idx, weight in feat_weights:
            lime_coeffs[i, feat_idx] = weight
        lime_coeffs[i, -1] = y_intercept[i]

    plot_exemplars(X_test[mask], y_test[mask], model.X_exemplars, model.y_exemplars, coeff[mask], fn_lb=X_test.min(), fn_ub=X_test.max(), model=model, method='Weighted Lars')
    # plot_exemplars(X_test[mask], y_test[mask], model.X_exemplars, model.y_exemplars, lime_coeffs[mask], fn_lb=X_test.min(), fn_ub=X_test.max(), model=model, method='Lime inference')
    plotly_animate(x_test=X_test.ravel()[mask], y_test=y_test[mask], x_exmp=model.X_exemplars.ravel(), y_exmp=model.y_exemplars, weights=W[mask])

if __name__ == '__main__':
    demo_main()
