"""Meta-ensemble: RandomForest + XGBoost + LightGBM → Blending.

Для каждого направления (long/short):
1. Обучает 3 базовые модели (RF, XGB, LGB) на 80% train
2. Находит оптимальные веса blending на val (20% train) — минимизация log_loss
3. Возвращает взвешенное среднее (blend) как итоговую вероятность

Blending не сжимает вероятности (в отличие от Platt scaling) и сохраняет
высокие значения при уверенных предсказаниях базовых моделей.
"""

import numpy as np
from scipy.optimize import minimize
from sklearn.ensemble import RandomForestClassifier
from sklearn.metrics import log_loss
from xgboost import XGBClassifier
from lightgbm import LGBMClassifier


def _optimal_blend_weights(p_val, y_val):
    """Находит оптимальные веса blending через минимизацию log_loss.

    Args:
        p_val: (n_samples, 3) — [P_rf, P_xgb, P_lgb] на val
        y_val: (n_samples,) — истинные метки

    Returns:
        np.ndarray: (3,) — оптимальные веса (сумма = 1)
    """

    def loss(w):
        w = np.clip(w, 0, 1)
        w = w / w.sum()
        p = p_val @ w
        p = np.clip(p, 1e-15, 1 - 1e-15)
        return log_loss(y_val, p)

    best_w = np.array([1 / 3, 1 / 3, 1 / 3])
    best_loss = loss(best_w)
    for init in [[1, 0, 0], [0, 1, 0], [0, 0, 1], [0.5, 0.3, 0.2]]:
        res = minimize(
            lambda x: loss(x),
            init,
            bounds=[(0, 1)] * 3,
            constraints={'type': 'eq', 'fun': lambda x: x.sum() - 1},
            method='SLSQP',
            options={'maxiter': 200, 'ftol': 1e-8},
        )
        if res.fun < best_loss:
            best_loss = res.fun
            best_w = res.x
    return best_w / best_w.sum()


class EnsembleModel:
    """Blending-ансамбль: RF + XGB + LGB → взвешенное среднее.

    Архитектура:
        X_train ──┬── RandomForest ── P_rf ──┐
                  ├── XGBoost ─────── P_xgb ──┼── взвешенное среднее ──→ P_ensemble
                  └── LightGBM ────── P_lgb ──┘

    Веса оптимизируются по log_loss на валидации.
    """

    def __init__(self, long_weight='balanced', short_weight=3.0, random_state=42):
        self.long_weight = long_weight
        self.short_weight = short_weight
        self.random_state = random_state

    def _get_params(self, direction, n_train, n_pos):
        """Параметры для базовых моделей."""
        if direction == 'long':
            w = self.long_weight
            spw = (n_train - n_pos) / n_pos if w == 'balanced' and n_pos > 0 else 1.0
            return {'class_weight': w, 'scale_pos_weight': spw}
        else:
            return {
                'class_weight': {0: 1.0, 1: self.short_weight},
                'scale_pos_weight': (n_train - n_pos) / n_pos if n_pos > 0 else 1.0,
            }

    def fit(self, X_train, y_train_long, y_train_short,
            X_val, y_val_long, y_val_short):
        """Обучает ансамбль.

        Args:
            X_train: (n_train, n_features) — 80% train
            y_train_long/short: целевые
            X_val: (n_val, n_features) — 20% train
            y_val_long/short: целевые
        """
        for direction in ['long', 'short']:
            y_train = y_train_long if direction == 'long' else y_train_short
            y_val = y_val_long if direction == 'long' else y_val_short
            n_pos = int(y_train.sum())
            params = self._get_params(direction, len(y_train), n_pos)
            rs = self.random_state

            # --- 1. Базовые модели ---
            rf = RandomForestClassifier(
                n_estimators=200, max_depth=10, random_state=rs, n_jobs=-1,
                class_weight=params['class_weight'],
            )
            xgb = XGBClassifier(
                n_estimators=200, max_depth=6, learning_rate=0.1,
                subsample=0.8, colsample_bytree=0.8,
                scale_pos_weight=params['scale_pos_weight'],
                random_state=rs, verbosity=0, n_jobs=-1,
            )
            lgb = LGBMClassifier(
                n_estimators=200, max_depth=6, learning_rate=0.1,
                subsample=0.8, colsample_bytree=0.8,
                class_weight=params['class_weight'],
                random_state=rs, verbose=-1, n_jobs=-1,
            )
            rf.fit(X_train, y_train)
            xgb.fit(X_train, y_train)
            lgb.fit(X_train, y_train)

            # --- 2. Оптимальные веса blending на val ---
            rf_val = rf.predict_proba(X_val)[:, 1]
            xgb_val = xgb.predict_proba(X_val)[:, 1]
            lgb_val = lgb.predict_proba(X_val)[:, 1]
            p_val = np.column_stack([rf_val, xgb_val, lgb_val])
            blend_weights = _optimal_blend_weights(p_val, y_val)

            # Сохраняем
            setattr(self, f'rf_{direction}', rf)
            setattr(self, f'xgb_{direction}', xgb)
            setattr(self, f'lgb_{direction}', lgb)
            setattr(self, f'blend_weights_{direction}', blend_weights)

    def predict_proba(self, X, direction='long'):
        """Возвращает P(class=1) — взвешенное среднее 3-х моделей.

        Returns:
            np.ndarray: (n_samples,)
        """
        rf_p = getattr(self, f'rf_{direction}').predict_proba(X)[:, 1]
        xgb_p = getattr(self, f'xgb_{direction}').predict_proba(X)[:, 1]
        lgb_p = getattr(self, f'lgb_{direction}').predict_proba(X)[:, 1]
        p = np.column_stack([rf_p, xgb_p, lgb_p])
        return p @ getattr(self, f'blend_weights_{direction}')
