"""
Configuration classes for training.

Stores all hyperparameters and settings for reproducibility.
"""

from dataclasses import dataclass, field, asdict
from typing import Optional, Dict, Any
import json


@dataclass
class ModelConfig:
    """Model architecture configuration."""

    # Basic architecture
    vocab_size: int = 10000
    embedding_dim: int = 768
    vector_dim: int = 1024
    chunk_size: int = 8
    num_layers: int = 12
    num_heads: int = 8
    feedforward_dim: Optional[int] = None
    max_seq_length: int = 2048
    dropout: float = 0.1

    # Modern architecture improvements (2024+)
    ffn_type: str = 'swiglu'  # 'swiglu', 'geglu', 'geglu_v2', 'gelu', 'relu', 'swish'
    use_qk_norm: bool = True  # QK-Normalization for attention stability
    qk_norm_type: str = 'rmsnorm'  # 'rmsnorm', 'l2', 'layernorm'
    use_rmsnorm: bool = True  # Use RMSNorm instead of LayerNorm (more efficient)
    use_bias: bool = False  # Bias in linear layers (modern LLMs typically use False)

    # Salience configuration
    salience_w1: float = 0.4
    salience_w2: float = 0.3
    salience_w3: float = 0.3
    salience_lambda: float = 0.1
    salience_k_fatigue: float = 0.2
    salience_temperature: float = 1.0
    salience_normalization: str = 'softmax'
    use_dimensional_scaling: bool = True
    learnable_salience_params: bool = True

    # CALM configuration
    autoencoder_num_encoder_layers: int = 6
    autoencoder_num_decoder_layers: int = 6
    memory_buffer_size: int = 64

    # Positional Encoding configuration
    # Options: 'rope', 'ntk_rope', 'yarn', 'longrope', 'alibi'
    pe_type: str = 'rope'
    pe_base: float = 10000.0
    pe_scaling_factor: float = 1.0

    # Context extension settings
    original_max_seq_length: int = 2048  # Original training context

    # NTK-aware RoPE
    ntk_factor: float = 1.0  # Auto-computed if 1.0

    # YaRN settings
    yarn_beta_fast: int = 32
    yarn_beta_slow: int = 1
    yarn_mscale: float = 0.0  # Auto-computed if 0.0
    yarn_mscale_all_dim: float = 0.0

    # LongRoPE settings (None = auto-computed)
    longrope_short_mscale: float = 1.0
    longrope_long_mscale: float = 1.0

    # ALiBi settings
    alibi_slope_computation: str = 'original'  # 'original', 'linear', 'exponential'

    def to_dict(self) -> Dict[str, Any]:
        """Convert to dictionary."""
        return asdict(self)

    def save(self, path: str):
        """Save to JSON file."""
        with open(path, 'w') as f:
            json.dump(self.to_dict(), f, indent=2)

    @classmethod
    def load(cls, path: str) -> 'ModelConfig':
        """Load from JSON file."""
        with open(path, 'r') as f:
            data = json.load(f)
        return cls(**data)


@dataclass
class TrainingConfig:
    """Training configuration."""

    # Training hyperparameters
    batch_size: int = 32
    learning_rate: float = 1e-4
    weight_decay: float = 0.01
    num_epochs: int = 10
    warmup_steps: int = 1000
    max_grad_norm: float = 1.0

    # Optimizer settings
    optimizer: str = 'adamw'
    adam_beta1: float = 0.9
    adam_beta2: float = 0.999
    adam_epsilon: float = 1e-8

    # Learning rate schedule
    lr_schedule: str = 'cosine'  # 'cosine', 'linear', 'constant'
    min_lr: float = 1e-6

    # Loss weights
    mse_weight: float = 1.0
    cosine_weight: float = 0.5
    reconstruction_weight: float = 2.0
    contrastive_weight: float = 0.1

    # Training stages
    pretrain_autoencoder: bool = True
    autoencoder_steps: int = 10000
    autoencoder_target_accuracy: float = 0.999
    freeze_autoencoder_after_pretrain: bool = False

    # Curriculum learning
    use_curriculum: bool = True
    initial_seq_length: int = 128
    curriculum_warmup_steps: int = 5000

    # Checkpointing
    save_every: int = 1000
    eval_every: int = 500
    checkpoint_dir: str = './checkpoints'
    keep_n_checkpoints: int = 3

    # Mixed precision
    use_amp: bool = True
    gradient_accumulation_steps: int = 1

    # Reproducibility
    seed: int = 42

    # Logging
    log_every: int = 10
    wandb_project: Optional[str] = None
    wandb_run_name: Optional[str] = None

    def to_dict(self) -> Dict[str, Any]:
        """Convert to dictionary."""
        return asdict(self)

    def save(self, path: str):
        """Save to JSON file."""
        with open(path, 'w') as f:
            json.dump(self.to_dict(), f, indent=2)

    @classmethod
    def load(cls, path: str) -> 'TrainingConfig':
        """Load from JSON file."""
        with open(path, 'r') as f:
            data = json.load(f)
        return cls(**data)

    def __str__(self) -> str:
        """Pretty print configuration."""
        lines = ["TrainingConfig:"]
        for key, value in self.to_dict().items():
            lines.append(f"  {key}: {value}")
        return "\n".join(lines)
