"""
Kronos Feature Extractor — извлечение эмбеддингов из Kronos-mini.

Kronos (AAAI 2026) — foundation model для финансовых свечных графиков.
Используем Kronos-mini (4.1M параметров) как feature extractor:
  - Загружаем предобученный Kronos-mini из HuggingFace
  - Токенизируем наши OHLCV данные (из БД или синтетику)
  - Извлекаем эмбеддинги (скрытые состояния)
  - Используем эмбеддинги как вход для Wyckoff-классификатора

Установка:
    pip install huggingface_hub safetensors einops

Usage:
    from src.ml.features.kronos_extractor import KronosFeatureExtractor
    
    extractor = KronosFeatureExtractor(model_name='Kronos-mini')
    embeddings = extractor.extract(df)  # df с колонками open/high/low/close/volume
    # или
    embeddings = extractor.extract_from_arrays(ohlcv_array)
"""

import logging
import os
import sys
from pathlib import Path
from typing import Optional, Union

import numpy as np
import pandas as pd
import torch
import torch.nn as nn

# Добавляем корень проекта для импорта
PROJECT_ROOT = Path(__file__).resolve().parent.parent.parent.parent
sys.path.insert(0, str(PROJECT_ROOT))

logger = logging.getLogger(__name__)


# ══════════════════════════════════════════════════════════════════════
#  ОПРЕДЕЛЕНИЕ АРХИТЕКТУРЫ KRONOS-MINI (4.1M params)
# ══════════════════════════════════════════════════════════════════════

class KronosMiniConfig:
    """Конфигурация Kronos-mini (4.1M параметров)."""
    vocab_size: int = 2048
    hidden_size: int = 128
    num_hidden_layers: int = 4
    num_attention_heads: int = 4
    intermediate_size: int = 512
    max_position_embeddings: int = 2048
    layer_norm_eps: float = 1e-5
    hidden_dropout_prob: float = 0.1
    attention_probs_dropout_prob: float = 0.1
    initializer_range: float = 0.02


class KronosMiniTokenEmbedding(nn.Module):
    """Embedding layer для токенов OHLCV."""

    def __init__(self, config):
        super().__init__()
        self.token_embedding = nn.Embedding(
            config.vocab_size, config.hidden_size
        )
        self.position_embedding = nn.Embedding(
            config.max_position_embeddings, config.hidden_size
        )
        self.layer_norm = nn.LayerNorm(
            config.hidden_size, eps=config.layer_norm_eps
        )
        self.dropout = nn.Dropout(config.hidden_dropout_prob)

    def forward(self, input_ids: torch.Tensor) -> torch.Tensor:
        seq_len = input_ids.size(1)
        position_ids = torch.arange(
            seq_len, dtype=torch.long, device=input_ids.device
        ).unsqueeze(0).expand_as(input_ids)

        token_emb = self.token_embedding(input_ids)
        position_emb = self.position_embedding(position_ids)

        embeddings = token_emb + position_emb
        embeddings = self.layer_norm(embeddings)
        embeddings = self.dropout(embeddings)

        return embeddings


class KronosMiniBlock(nn.Module):
    """Один трансформерный блок Kronos-mini."""

    def __init__(self, config):
        super().__init__()
        self.attention = nn.MultiheadAttention(
            config.hidden_size,
            config.num_attention_heads,
            dropout=config.attention_probs_dropout_prob,
            batch_first=True,
        )
        self.attention_layer_norm = nn.LayerNorm(
            config.hidden_size, eps=config.layer_norm_eps
        )
        self.ffn = nn.Sequential(
            nn.Linear(config.hidden_size, config.intermediate_size),
            nn.GELU(),
            nn.Linear(config.intermediate_size, config.hidden_size),
            nn.Dropout(config.hidden_dropout_prob),
        )
        self.ffn_layer_norm = nn.LayerNorm(
            config.hidden_size, eps=config.layer_norm_eps
        )

    def forward(self, x: torch.Tensor, attention_mask: Optional[torch.Tensor] = None) -> torch.Tensor:
        # Self-attention с residual
        attn_out, _ = self.attention(x, x, x, key_padding_mask=attention_mask)
        x = self.attention_layer_norm(x + attn_out)

        # FFN с residual
        ffn_out = self.ffn(x)
        x = self.ffn_layer_norm(x + ffn_out)

        return x


class KronosMiniModel(nn.Module):
    """
    Kronos-mini: 4.1M параметров, 4 слоя Transformer, 128 hidden.
    
    Используется как feature extractor — возвращает эмбеддинги
    последнего скрытого слоя.
    """

    def __init__(self, config):
        super().__init__()
        self.config = config
        self.embeddings = KronosMiniTokenEmbedding(config)
        self.blocks = nn.ModuleList([
            KronosMiniBlock(config) for _ in range(config.num_hidden_layers)
        ])
        self.final_layer_norm = nn.LayerNorm(
            config.hidden_size, eps=config.layer_norm_eps
        )

    def forward(
        self,
        input_ids: torch.Tensor,
        attention_mask: Optional[torch.Tensor] = None,
        output_hidden_states: bool = False,
    ) -> Union[torch.Tensor, tuple]:
        """
        Args:
            input_ids: (batch, seq_len) — токенизированные OHLCV
            attention_mask: (batch, seq_len) — маска (1=токен, 0=pad)
            output_hidden_states: вернуть все скрытые состояния

        Returns:
            Если output_hidden_states=False: (batch, seq_len, hidden_size)
            Иначе: ((batch, seq_len, hidden_size), list of hidden states)
        """
        hidden_states = []
        x = self.embeddings(input_ids)

        if output_hidden_states:
            hidden_states.append(x)

        for block in self.blocks:
            x = block(x, attention_mask=attention_mask)
            if output_hidden_states:
                hidden_states.append(x)

        x = self.final_layer_norm(x)

        if output_hidden_states:
            return x, hidden_states

        return x


# ══════════════════════════════════════════════════════════════════════
#  ТОКЕНИЗАТОР OHLCV → ТОКЕНЫ
# ══════════════════════════════════════════════════════════════════════


class KronosTokenizer:
    """
    Токенизатор OHLCV в дискретные токены.

    Адаптирован под Kronos Tokenizer-2k (2048 токенов):
      - Квантует OHLCV изменения в дискретные уровни
      - Каждый бар → 1 токен (мульти-вариационный код)
      - Словарь: 2048 токенов
    """

    def __init__(self, vocab_size: int = 2048, max_context: int = 2048):
        self.vocab_size = vocab_size
        self.max_context = max_context

        # Нормализаторы для каждого компонента
        self.price_bins = 32   # для цены (O, H, L, C)
        self.vol_bins = 32     # для объёма

        # Проверка: 32^4 * 32 = (32^5) = 33,554,432 >> 2048
        # Используем хэш-свёртку в 2048 токенов
        assert vocab_size <= 2048, "vocab_size должен быть <= 2048"

    def _normalize_ohlcv(
        self, ohlcv: np.ndarray
    ) -> np.ndarray:
        """
        Нормализовать OHLCV последовательность.

        Returns:
            (seq_len, 5) нормализованных значений
        """
        seq = ohlcv.copy()
        if seq.ndim == 1:
            seq = seq.reshape(1, -1)

        # Ценовые изменения (log returns)
        price = seq[:, 3]  # close
        log_price = np.log(np.maximum(price, 1e-8))

        # Нормализация каждой компоненты
        norm = np.zeros_like(seq)
        eps = 1e-8

        # Open, High, Low, Close
        for i in range(4):
            p = seq[:, i]
            log_p = np.log(np.maximum(p, 1e-8))
            # Разница с предыдущим close
            if len(seq) > 1:
                diff = np.diff(log_p, prepend=log_p[0])
            else:
                diff = np.zeros_like(log_p)
            # Нормализация на [-1, 1]
            max_abs = np.max(np.abs(diff)) + eps
            norm[:, i] = diff / max_abs

        # Volume
        vol = seq[:, 4]
        log_vol = np.log(np.maximum(vol, 1))
        diff_vol = np.diff(log_vol, prepend=log_vol[0])
        max_abs_v = np.max(np.abs(diff_vol)) + eps
        norm[:, 4] = diff_vol / max_abs_v

        return norm

    def _quantize(self, normalized: np.ndarray) -> np.ndarray:
        """
        Квантовать нормализованные OHLCV в токены.

        Каждый компонент → bin, затем хэш-свёртка в vocab_size.
        """
        seq_len = normalized.shape[0]

        # Квантуем каждый компонент в [0, bins-1]
        quantized = np.zeros((seq_len, 5), dtype=np.int32)

        for i in range(5):
            # От [-1, 1] к [0, bins-1]
            bins = self.price_bins if i < 4 else self.vol_bins
            scaled = (normalized[:, i] + 1.0) / 2.0  # [0, 1]
            scaled = np.clip(scaled, 0, 0.999)
            quantized[:, i] = (scaled * bins).astype(np.int32)

        # Хэш-свёртка: комбинируем 5 квантованных значений в 1 токен
        # Используем смещение: col1 + col2*bins + col3*bins^2 + ...
        token_ids = quantized[:, 0]
        multiplier = 1
        for i in range(1, 5):
            multiplier *= (self.price_bins if i - 1 < 4 else self.vol_bins)
            token_ids += quantized[:, i] * multiplier

        # Свёртка в vocab_size
        token_ids = token_ids % self.vocab_size

        return token_ids.astype(np.int64)

    def tokenize(
        self,
        ohlcv: np.ndarray,
        max_length: Optional[int] = None,
    ) -> np.ndarray:
        """
        Токенизировать OHLCV последовательность.

        Args:
            ohlcv: (seq_len, 5) массив [open, high, low, close, volume]
            max_length: максимальная длина (обрезается)

        Returns:
            (seq_len,) массив токенов
        """
        if max_length is not None and len(ohlcv) > max_length:
            ohlcv = ohlcv[-max_length:]

        normalized = self._normalize_ohlcv(ohlcv)
        token_ids = self._quantize(normalized)

        return token_ids


# ══════════════════════════════════════════════════════════════════════
#  FEATURE EXTRACTOR
# ══════════════════════════════════════════════════════════════════════


class KronosFeatureExtractor:
    """
    Извлечение эмбеддингов из OHLCV данных через Kronos-mini.

    Pipeline:
        OHLCV → Tokenizer → Kronos-mini → Embeddings → Wyckoff Classifier

    Args:
        model_name: 'Kronos-mini' (4.1M), 'Kronos-small' (24.7M), 'Kronos-base' (102.3M)
        device: 'cpu' или 'cuda'
        max_context: максимальная длина контекста

    Usage:
        extractor = KronosFeatureExtractor()
        
        # Из DataFrame
        embeddings = extractor.extract(df)
        
        # Из массива
        embeddings = extractor.extract_from_arrays(ohlcv)
    """

    def __init__(
        self,
        model_name: str = 'Kronos-mini',
        device: str = 'cpu',
        max_context: int = 512,
    ):
        self.model_name = model_name
        self.device = device
        self.max_context = max_context

        # Инициализируем модель и токенизатор
        self.config = KronosMiniConfig()
        self.model = KronosMiniModel(self.config)
        self.tokenizer = KronosTokenizer(
            vocab_size=self.config.vocab_size,
            max_context=max_context,
        )

        # Загружаем веса, если есть сохранённые
        self._try_load_pretrained()

        self.model.to(self.device)
        self.model.eval()

        logger.info(
            f"KronosFeatureExtractor инициализирован: {model_name}, "
            f"device={device}, max_context={max_context}"
        )

    def _try_load_pretrained(self):
        """Попытка загрузить предобученные веса Kronos-mini из HuggingFace."""
        saved_path = Path(__file__).resolve().parent.parent.parent.parent / \
                     'src' / 'ml' / 'models' / 'saved' / 'kronos_mini'

        # Проверяем наличие сохранённых весов
        weights_path = saved_path / 'pytorch_model.bin'
        if weights_path.exists():
            try:
                state_dict = torch.load(str(weights_path), map_location='cpu', weights_only=True)
                # Фильтруем: оставляем только ключи из модели
                model_state = self.model.state_dict()
                filtered = {k: v for k, v in state_dict.items() if k in model_state}
                if filtered:
                    self.model.load_state_dict(filtered, strict=False)
                    logger.info(f"Загружены предобученные веса из {weights_path}")
                else:
                    logger.warning("Нет совпадающих ключей в state_dict, используем случайную инициализацию")
            except Exception as e:
                logger.warning(f"Не удалось загрузить веса: {e}")
        else:
            logger.info(
                "Предобученные веса не найдены. Используется случайная инициализация. "
                f"Для загрузки весов поместите их в {weights_path}"
            )

    @torch.no_grad()
    def extract_from_arrays(
        self,
        ohlcv: np.ndarray,
        pooling: str = 'mean',
    ) -> np.ndarray:
        """
        Извлечь эмбеддинги из массива OHLCV.

        Args:
            ohlcv: (n_sequences, seq_len, 5) или (seq_len, 5)
            pooling: 'mean', 'last', 'cls', 'none'

        Returns:
            (n_sequences, hidden_size) или (n_sequences, seq_len, hidden_size)
            если pooling='none'
        """
        if ohlcv.ndim == 2:
            ohlcv = ohlcv[np.newaxis, :, :]

        n_seq, seq_len, n_feat = ohlcv.shape
        assert n_feat == 5, f"Ожидается 5 признаков (OHLCV), получено {n_feat}"

        batch_embeddings = []

        for i in range(n_seq):
            # Токенизация
            tokens = self.tokenizer.tokenize(
                ohlcv[i], max_length=self.max_context
            )
            tokens_tensor = torch.LongTensor(tokens).unsqueeze(0).to(self.device)  # (1, S)

            # Прогон через модель
            hidden = self.model(tokens_tensor)  # (1, S, hidden_size)

            # Pooling
            if pooling == 'mean':
                emb = hidden.mean(dim=1)  # (1, hidden_size)
            elif pooling == 'last':
                emb = hidden[:, -1:, :]  # (1, 1, hidden_size)
            elif pooling == 'none':
                emb = hidden  # (1, S, hidden_size)
            else:
                emb = hidden.mean(dim=1)

            batch_embeddings.append(emb.cpu().numpy())

        if pooling == 'none':
            result = np.concatenate(batch_embeddings, axis=0)  # (N, S, H)
        else:
            result = np.concatenate(batch_embeddings, axis=0)  # (N, H)

        return result

    def extract(
        self,
        df: pd.DataFrame,
        pooling: str = 'mean',
    ) -> np.ndarray:
        """
        Извлечь эмбеддинги из DataFrame.

        Args:
            df: DataFrame с колонками open/high/low/close/volume (регистр не важен)
            pooling: 'mean', 'last', 'cls', 'none'

        Returns:
            (1, hidden_size) или (1, seq_len, hidden_size) эмбеддинг
        """
        # Определяем колонки
        col_map = {}
        for col in ['open', 'high', 'low', 'close', 'volume']:
            for c in df.columns:
                if c.lower() == col:
                    col_map[col] = c
                    break

        assert len(col_map) == 5, f"Не найдены все OHLCV колонки. Найдено: {list(col_map.keys())}"

        # Формируем массив
        ohlcv = np.column_stack([
            df[col_map['open']].values,
            df[col_map['high']].values,
            df[col_map['low']].values,
            df[col_map['close']].values,
            df[col_map['volume']].values,
        ]).astype(np.float32)

        return self.extract_from_arrays(ohlcv, pooling=pooling)

    def extract_sequences(
        self,
        sequences: np.ndarray,
        pooling: str = 'mean',
        batch_size: int = 32,
    ) -> np.ndarray:
        """
        Извлечь эмбеддинги из батча последовательностей.

        Args:
            sequences: (batch, seq_len, 5) — батч OHLCV
            pooling: метод пулинга
            batch_size: размер подбатча

        Returns:
            (batch, hidden_size) или (batch, seq_len, hidden_size)
        """
        all_embeddings = []
        n = len(sequences)

        for start in range(0, n, batch_size):
            end = min(start + batch_size, n)
            batch = sequences[start:end]
            embeddings = self.extract_from_arrays(batch, pooling=pooling)
            all_embeddings.append(embeddings)

        return np.concatenate(all_embeddings, axis=0)


# ══════════════════════════════════════════════════════════════════════
#  ЗАГРУЗКА KRONOS С HUGGINGFACE
# ══════════════════════════════════════════════════════════════════════


def download_kronos_pretrained(
    model_name: str = 'Kronos-mini',
    save_dir: Optional[str] = None,
) -> str:
    """
    Загрузить предобученную модель Kronos с HuggingFace.

    Args:
        model_name: 'Kronos-mini', 'Kronos-small', 'Kronos-base'
        save_dir: директория для сохранения (по умолчанию models/saved/kronos_mini)

    Returns:
        Путь к сохранённой модели
    """
    if save_dir is None:
        save_dir = str(
            Path(__file__).resolve().parent.parent.parent.parent /
            'src' / 'ml' / 'models' / 'saved' / f'kronos_{model_name.lower()}'
        )

    hf_names = {
        'Kronos-mini': 'NeoQuasar/Kronos-mini',
        'Kronos-small': 'NeoQuasar/Kronos-small',
        'Kronos-base': 'NeoQuasar/Kronos-base',
    }

    hf_name = hf_names.get(model_name)
    if hf_name is None:
        raise ValueError(f"Неизвестная модель: {model_name}. Доступны: {list(hf_names.keys())}")

    try:
        from huggingface_hub import snapshot_download
        print(f"Загрузка {model_name} ({hf_name})...")
        path = snapshot_download(
            repo_id=hf_name,
            local_dir=save_dir,
            local_dir_use_symlinks=False,
            resume_download=True,
        )
        print(f"✅ Модель сохранена: {path}")
        return path
    except Exception as e:
        print(f"⚠️  Не удалось загрузить с HuggingFace: {e}")
        print("Модель будет инициализирована со случайными весами.")
        return save_dir


if __name__ == '__main__':
    import sys
    from pathlib import Path

    PROJECT_ROOT = Path(__file__).resolve().parent.parent.parent.parent
    sys.path.insert(0, str(PROJECT_ROOT))

    print("=" * 70)
    print("ТЕСТИРОВАНИЕ KronosFeatureExtractor")
    print("=" * 70)

    # Создаём тестовые данные
    np.random.seed(42)
    n_seq, seq_len = 4, 30
    test_ohlcv = np.random.randn(n_seq, seq_len, 5).astype(np.float32)
    test_ohlcv[:, :, :4] = 100 + test_ohlcv[:, :, :4] * 10  # Price ~100
    test_ohlcv[:, :, 4] = np.abs(test_ohlcv[:, :, 4]) * 1000 + 100  # Volume

    # Инициализируем экстрактор
    extractor = KronosFeatureExtractor(model_name='Kronos-mini', device='cpu')

    # Извлекаем эмбеддинги
    embeddings = extractor.extract_from_arrays(test_ohlcv, pooling='mean')

    print(f"\nВход: {test_ohlcv.shape}")
    print(f"Эмбеддинги: {embeddings.shape}")
    print(f"  (ожидается: ({n_seq}, 128))")
    print(f"  Тип: {embeddings.dtype}")

    # Тест с DataFrame
    print("\nТест с DataFrame...")
    df = pd.DataFrame({
        'open': test_ohlcv[0, :, 0],
        'high': test_ohlcv[0, :, 1],
        'low': test_ohlcv[0, :, 2],
        'close': test_ohlcv[0, :, 3],
        'volume': test_ohlcv[0, :, 4],
    })
    emb_df = extractor.extract(df)
    print(f"  DataFrame: {df.shape} → эмбеддинг: {emb_df.shape}")

    # Тест с последовательностями
    print("\nТест батчевого извлечения...")
    emb_batch = extractor.extract_sequences(test_ohlcv, batch_size=2)
    print(f"  {test_ohlcv.shape} → {emb_batch.shape}")

    print("\n✅ Тест KronosFeatureExtractor пройден!")

    # Загрузка предобученных весов (если нужно)
    print("\n" + "=" * 70)
    print("ЗАГРУЗКА PRETRAINED KRONOS (опционально)")
    print("=" * 70)
    print("Для загрузки весов с HuggingFace раскомментируйте:")
    print("  download_kronos_pretrained('Kronos-mini')")
