"""
Configuration management.
"""

import json
import os
from dataclasses import dataclass, asdict
from typing import Optional


@dataclass
class ModelConfig:
    """Model architecture configuration."""
    vocab_size: int = 10000
    embedding_dim: int = 512
    num_layers: int = 6
    num_heads: int = 8
    feedforward_dim: Optional[int] = None
    max_seq_length: int = 512
    dropout: float = 0.1
    use_memory_buffer: bool = True
    memory_buffer_size: int = 32
    use_formula_attention: bool = True
    
    # Formula-specific config
    formula_w1: float = 0.4  # Novelty weight
    formula_w2: float = 0.3  # Retention weight
    formula_w3: float = 0.3  # Payoff weight
    formula_lambda: float = 0.1  # Time decay
    formula_k: float = 0.2  # Fatigue coefficient
    formula_learnable_weights: bool = True
    formula_learnable_decay: bool = True


@dataclass
class TrainingConfig:
    """Training configuration."""
    batch_size: int = 8
    learning_rate: float = 1e-4
    weight_decay: float = 0.01
    num_epochs: int = 10
    warmup_steps: int = 1000
    max_grad_norm: float = 1.0
    log_interval: int = 100
    save_dir: str = './checkpoints'
    device: str = 'cuda'  # Will be set to 'cpu' if CUDA not available
    
    # Data paths
    train_data_path: Optional[str] = None
    val_data_path: Optional[str] = None
    
    # Generation config
    max_new_tokens: int = 50
    temperature: float = 1.0
    top_k: Optional[int] = None
    top_p: float = 1.0


@dataclass
class Config:
    """Complete configuration."""
    model: ModelConfig
    training: TrainingConfig
    
    @classmethod
    def from_dict(cls, config_dict: dict):
        """Create config from dictionary."""
        return cls(
            model=ModelConfig(**config_dict.get('model', {})),
            training=TrainingConfig(**config_dict.get('training', {}))
        )
    
    def to_dict(self):
        """Convert config to dictionary."""
        return {
            'model': asdict(self.model),
            'training': asdict(self.training)
        }
    
    def save(self, path: str):
        """Save config to JSON file."""
        with open(path, 'w') as f:
            json.dump(self.to_dict(), f, indent=2)
    
    @classmethod
    def load(cls, path: str):
        """Load config from JSON file."""
        with open(path, 'r') as f:
            config_dict = json.load(f)
        return cls.from_dict(config_dict)
    
    @classmethod
    def default(cls):
        """Create default configuration."""
        return cls(
            model=ModelConfig(),
            training=TrainingConfig()
        )



