"""
Complete MK3 Trainer with Advanced Training Infrastructure

Handles two-stage training:
1. Autoencoder pretraining (>99.9% reconstruction accuracy)
2. Continuous autoregressive model training with salience

Enhanced with:
- FSDP distributed training support
- Advanced optimizers (GaLore, 8-bit Adam, Improved AdamW)
- Mixed precision training (FP16/BF16/FP8)
- Gradient accumulation strategies
- Activation checkpointing
- Better checkpoint management
- Improved curriculum learning

Includes full reproducibility, checkpointing, and evaluation.
"""

import torch
import torch.nn as nn
from torch.optim import AdamW
from torch.cuda.amp import autocast, GradScaler
from typing import Optional, Dict, List, Any, Tuple
import os
from pathlib import Path
from tqdm import tqdm
import json
import sys
import logging
import time
from dataclasses import asdict

sys.path.append(os.path.dirname(os.path.dirname(os.path.abspath(__file__))))

from calm.continuous_model import ContinuousAutoregressiveModel
from calm.likelihood_free import ContinuousLoss, CurriculumScheduler
from training.config import TrainingConfig, ModelConfig
from utils.reproducibility import set_seed, make_deterministic, get_device

# Import new training infrastructure
from training.distributed import (
    DistributedConfig,
    setup_distributed_training,
    GradientAccumulator,
    FSDPWrapper,
)
from training.optimizers import (
    create_optimizer,
    get_optimizer_info,
    GaLoreAdamW,
    AdamW8bit,
    ImprovedAdamW,
    Lion,
)
from training.mixed_precision import (
    MixedPrecisionManager,
    MixedPrecisionConfig,
    create_mixed_precision_manager,
    get_optimal_precision,
)

logger = logging.getLogger(__name__)


class CompleteMK3Trainer:
    """
    Complete trainer for MK3 continuous autoregressive model.

    Handles all aspects of training with full reproducibility.
    Enhanced with distributed training, advanced optimizers, and mixed precision.
    """

    def __init__(
        self,
        model_config: ModelConfig,
        training_config: TrainingConfig,
        train_loader,
        val_loader,
        device: Optional[torch.device] = None,
        distributed_config: Optional[DistributedConfig] = None,
        use_distributed: bool = False,
    ):
        self.model_config = model_config
        self.training_config = training_config
        self.train_loader = train_loader
        self.val_loader = val_loader

        # Setup device
        self.device = device or get_device()
        print(f"Using device: {self.device}")

        # Set seed for reproducibility
        make_deterministic(training_config.seed)

        # Build model
        self.model = self._build_model()
        self.model.to(self.device)

        # Distributed training setup
        self.distributed_config = distributed_config
        self.distributed_manager = None
        if use_distributed and distributed_config:
            self.model, self.distributed_manager = setup_distributed_training(
                self.model,
                distributed_config,
                apply_activation_checkpointing=True,
            )
            logger.info(f"Distributed training enabled: {distributed_config.world_size} processes")

        # Loss function
        self.loss_fn = ContinuousLoss(
            mse_weight=training_config.mse_weight,
            cosine_weight=training_config.cosine_weight,
            reconstruction_weight=training_config.reconstruction_weight,
            contrastive_weight=training_config.contrastive_weight
        )

        # Advanced optimizer
        self.optimizer = self._build_optimizer()

        # Learning rate scheduler
        self.scheduler = self._build_scheduler()

        # Curriculum scheduler with improvements
        if training_config.use_curriculum:
            self.curriculum = CurriculumScheduler(
                initial_seq_length=training_config.initial_seq_length,
                max_seq_length=model_config.max_seq_length,
                warmup_steps=training_config.curriculum_warmup_steps
            )
        else:
            self.curriculum = None

        # Enhanced mixed precision manager
        precision = getattr(training_config, 'precision', 'mixed_bf16')
        if precision == 'auto':
            precision = get_optimal_precision(self.device)
            logger.info(f"Auto-detected optimal precision: {precision}")

        self.mixed_precision = create_mixed_precision_manager(
            precision=precision,
            max_grad_norm=training_config.max_grad_norm,
        )

        # Gradient accumulation manager
        self.gradient_accumulator = GradientAccumulator(
            distributed_config or DistributedConfig(
                gradient_accumulation_steps=training_config.gradient_accumulation_steps
            )
        )

        # Training state
        self.global_step = 0
        self.epoch = 0
        self.best_val_loss = float('inf')
        self.training_stats = {
            'losses': [],
            'learning_rates': [],
            'grad_norms': [],
        }

        # Create checkpoint directory
        Path(training_config.checkpoint_dir).mkdir(parents=True, exist_ok=True)

        # Save configs
        model_config.save(os.path.join(training_config.checkpoint_dir, 'model_config.json'))
        training_config.save(os.path.join(training_config.checkpoint_dir, 'training_config.json'))

        # Log optimizer info
        opt_info = get_optimizer_info(self.optimizer)
        logger.info(f"Optimizer: {opt_info['optimizer_class']}")
        logger.info(f"Optimizer state size: {opt_info['state_size_mb']:.2f} MB")

    def _build_model(self) -> ContinuousAutoregressiveModel:
        """Build model from config."""
        salience_config = {
            'w1': self.model_config.salience_w1,
            'w2': self.model_config.salience_w2,
            'w3': self.model_config.salience_w3,
            'lambda_decay': self.model_config.salience_lambda,
            'k_fatigue': self.model_config.salience_k_fatigue,
            'temperature': self.model_config.salience_temperature,
            'learnable_weights': self.model_config.learnable_salience_params,
            'learnable_decay': self.model_config.learnable_salience_params,
            'learnable_temperature': self.model_config.learnable_salience_params,
            'normalization': self.model_config.salience_normalization,
            'use_dimensional_scaling': self.model_config.use_dimensional_scaling,
        }

        model = ContinuousAutoregressiveModel(
            vocab_size=self.model_config.vocab_size,
            embedding_dim=self.model_config.embedding_dim,
            vector_dim=self.model_config.vector_dim,
            chunk_size=self.model_config.chunk_size,
            num_layers=self.model_config.num_layers,
            num_heads=self.model_config.num_heads,
            feedforward_dim=self.model_config.feedforward_dim,
            max_seq_length=self.model_config.max_seq_length,
            dropout=self.model_config.dropout,
            salience_config=salience_config,
            memory_buffer_size=self.model_config.memory_buffer_size
        )

        return model

    def _build_optimizer(self) -> torch.optim.Optimizer:
        """Build optimizer with support for advanced optimizers."""
        optimizer_name = self.training_config.optimizer.lower()

        # Get additional optimizer kwargs
        optimizer_kwargs = {
            'betas': (self.training_config.adam_beta1, self.training_config.adam_beta2),
            'eps': self.training_config.adam_epsilon,
        }

        # Add optimizer-specific kwargs
        if optimizer_name == 'galore':
            optimizer_kwargs.update({
                'rank': getattr(self.training_config, 'galore_rank', 128),
                'update_proj_gap': getattr(self.training_config, 'galore_update_gap', 200),
                'scale': getattr(self.training_config, 'galore_scale', 1.0),
            })
        elif optimizer_name in ['8bit', 'adamw8bit']:
            optimizer_kwargs.update({
                'min_8bit_size': getattr(self.training_config, 'min_8bit_size', 4096),
            })
        elif optimizer_name in ['improved', 'improved_adamw']:
            optimizer_kwargs.update({
                'max_grad_norm': self.training_config.max_grad_norm,
                'warmup_steps': self.training_config.warmup_steps,
            })

        # Create optimizer using factory
        optimizer = create_optimizer(
            self.model,
            optimizer_name=optimizer_name,
            lr=self.training_config.learning_rate,
            weight_decay=self.training_config.weight_decay,
            **optimizer_kwargs
        )

        logger.info(f"Created optimizer: {optimizer.__class__.__name__}")
        return optimizer

    def _build_scheduler(self):
        """Build learning rate scheduler."""
        total_steps = len(self.train_loader) * self.training_config.num_epochs

        if self.training_config.lr_schedule == 'cosine':
            from torch.optim.lr_scheduler import CosineAnnealingLR
            return CosineAnnealingLR(
                self.optimizer,
                T_max=total_steps,
                eta_min=self.training_config.min_lr
            )
        elif self.training_config.lr_schedule == 'linear':
            from torch.optim.lr_scheduler import LinearLR
            return LinearLR(
                self.optimizer,
                start_factor=1.0,
                end_factor=self.training_config.min_lr / self.training_config.learning_rate,
                total_iters=total_steps
            )
        else:
            # Constant LR
            return None

    def pretrain_autoencoder(self):
        """
        Stage 1: Pretrain autoencoder to achieve >99.9% reconstruction accuracy.
        """
        if not self.training_config.pretrain_autoencoder:
            print("Skipping autoencoder pretraining")
            return

        print("\n" + "=" * 60)
        print("Stage 1: Pretraining Autoencoder")
        print("=" * 60)
        print(f"Target: {self.training_config.autoencoder_target_accuracy:.4f} reconstruction accuracy")

        autoencoder = self.model.autoencoder
        autoencoder.train()

        optimizer = AdamW(autoencoder.parameters(), lr=1e-4, weight_decay=0.01)

        best_accuracy = 0.0
        step = 0

        with tqdm(total=self.training_config.autoencoder_steps, desc="Autoencoder Pretraining") as pbar:
            while step < self.training_config.autoencoder_steps:
                for batch in self.train_loader:
                    if isinstance(batch, tuple):
                        token_ids, mask = batch
                    else:
                        token_ids = batch
                        mask = None

                    token_ids = token_ids.to(self.device)

                    # Extract chunks
                    batch_size, seq_len = token_ids.shape
                    if seq_len < self.model_config.chunk_size:
                        continue

                    # Random chunk
                    max_start = seq_len - self.model_config.chunk_size
                    start_idx = torch.randint(0, max_start + 1, (1,)).item()
                    chunk = token_ids[:, start_idx:start_idx + self.model_config.chunk_size]

                    # Forward pass
                    loss, metrics = autoencoder.compute_loss(chunk)

                    # Backward pass
                    optimizer.zero_grad()
                    loss.backward()
                    torch.nn.utils.clip_grad_norm_(autoencoder.parameters(), 1.0)
                    optimizer.step()

                    # Update progress
                    accuracy = metrics['reconstruction_accuracy']
                    pbar.set_postfix({
                        'loss': f"{metrics['loss']:.4f}",
                        'acc': f"{accuracy:.4f}",
                        'best': f"{best_accuracy:.4f}"
                    })
                    pbar.update(1)

                    if accuracy > best_accuracy:
                        best_accuracy = accuracy

                    # Check if target reached
                    if accuracy >= self.training_config.autoencoder_target_accuracy:
                        print(f"\nReached target accuracy: {accuracy:.4f}")
                        break

                    step += 1
                    if step >= self.training_config.autoencoder_steps:
                        break

                if best_accuracy >= self.training_config.autoencoder_target_accuracy:
                    break

        print(f"Autoencoder pretraining complete. Best accuracy: {best_accuracy:.4f}")

        # Save autoencoder
        autoencoder_path = os.path.join(self.training_config.checkpoint_dir, 'autoencoder_pretrained.pt')
        torch.save(autoencoder.state_dict(), autoencoder_path)
        print(f"Saved pretrained autoencoder to {autoencoder_path}")

        # Optionally freeze autoencoder
        if self.training_config.freeze_autoencoder_after_pretrain:
            for param in autoencoder.parameters():
                param.requires_grad = False
            print("Froze autoencoder parameters")

    def train_step(self, batch, accumulation_step: int = 0) -> Dict:
        """
        Single training step with enhanced features.

        Args:
            batch: Training batch
            accumulation_step: Current gradient accumulation step

        Returns:
            Dictionary of metrics
        """
        self.model.train()

        if isinstance(batch, tuple):
            token_ids, mask = batch
        else:
            token_ids = batch
            mask = None

        token_ids = token_ids.to(self.device)

        # Convert to continuous vectors
        continuous_vectors, _ = self.model.tokenize_to_vectors(token_ids)

        # Scale loss for gradient accumulation
        loss_scale = 1.0 / self.gradient_accumulator.accumulation_steps

        # Forward pass with mixed precision
        with self.mixed_precision.autocast_context():
            loss, metrics = self.model.compute_loss(continuous_vectors)
            scaled_loss = loss * loss_scale

        # Backward pass with gradient accumulation
        if self.gradient_accumulator.should_accumulate(accumulation_step):
            # Accumulate without sync
            with self.gradient_accumulator.no_sync(self.model):
                self.mixed_precision.backward(scaled_loss)
        else:
            # Final step: sync gradients
            self.mixed_precision.backward(scaled_loss)

            # Compute gradient norm for logging
            total_norm = 0.0
            for p in self.model.parameters():
                if p.grad is not None:
                    total_norm += p.grad.data.norm(2).item() ** 2
            total_norm = total_norm ** 0.5
            self.training_stats['grad_norms'].append(total_norm)

            # Step optimizer with gradient clipping
            step_success = self.mixed_precision.step_optimizer(
                self.optimizer,
                clip_grad=True,
            )

            if step_success:
                # Update learning rate
                if self.scheduler:
                    self.scheduler.step()

                # Zero gradients for next accumulation cycle
                self.optimizer.zero_grad()

                # Update curriculum if enabled
                if self.curriculum:
                    self.curriculum.step()

            metrics['grad_norm'] = total_norm
            metrics['step_success'] = step_success

        return metrics

    def evaluate(self) -> Dict:
        """Evaluation on validation set."""
        self.model.eval()

        total_loss = 0.0
        total_cosine_sim = 0.0
        num_batches = 0

        with torch.no_grad():
            for batch in self.val_loader:
                if isinstance(batch, tuple):
                    token_ids, mask = batch
                else:
                    token_ids = batch

                token_ids = token_ids.to(self.device)

                # Convert to continuous vectors
                continuous_vectors, _ = self.model.tokenize_to_vectors(token_ids)

                # Forward pass
                loss, metrics = self.model.compute_loss(continuous_vectors)

                total_loss += metrics['loss']
                total_cosine_sim += metrics['cosine_similarity']
                num_batches += 1

        avg_loss = total_loss / num_batches
        avg_cosine_sim = total_cosine_sim / num_batches

        return {
            'val_loss': avg_loss,
            'val_cosine_similarity': avg_cosine_sim
        }

    def train(self):
        """
        Complete training procedure with enhanced features.

        Stage 1: Pretrain autoencoder
        Stage 2: Train continuous autoregressive model
        """
        start_time = time.time()

        # Stage 1: Pretrain autoencoder
        if self.training_config.pretrain_autoencoder:
            self.pretrain_autoencoder()

        # Stage 2: Train continuous model
        print("\n" + "=" * 60)
        print("Stage 2: Training Continuous Autoregressive Model")
        print("=" * 60)

        # Log training configuration
        self._log_training_info()

        for epoch in range(self.training_config.num_epochs):
            self.epoch = epoch
            epoch_start = time.time()
            print(f"\nEpoch {epoch + 1}/{self.training_config.num_epochs}")

            # Training loop with gradient accumulation
            self.model.train()
            epoch_metrics = []
            accumulation_step = 0

            with tqdm(total=len(self.train_loader), desc=f"Training Epoch {epoch + 1}") as pbar:
                for batch_idx, batch in enumerate(self.train_loader):
                    metrics = self.train_step(batch, accumulation_step)
                    epoch_metrics.append(metrics)

                    # Only increment global step after accumulation cycle
                    if not self.gradient_accumulator.should_accumulate(accumulation_step):
                        self.global_step += 1
                        accumulation_step = 0

                        # Track statistics
                        self.training_stats['losses'].append(metrics['loss'])
                        self.training_stats['learning_rates'].append(
                            self.optimizer.param_groups[0]['lr']
                        )

                        # Update progress bar
                        if self.global_step % self.training_config.log_every == 0:
                            postfix = {
                                'loss': f"{metrics['loss']:.4f}",
                                'cos_sim': f"{metrics['cosine_similarity']:.4f}",
                                'lr': f"{self.optimizer.param_groups[0]['lr']:.2e}",
                            }
                            if 'grad_norm' in metrics:
                                postfix['grad_norm'] = f"{metrics['grad_norm']:.2f}"
                            pbar.set_postfix(postfix)

                        # Evaluation
                        if self.global_step % self.training_config.eval_every == 0:
                            val_metrics = self.evaluate()
                            print(f"\nValidation: Loss={val_metrics['val_loss']:.4f}, "
                                  f"Cosine Sim={val_metrics['val_cosine_similarity']:.4f}")

                            # Save best model
                            if val_metrics['val_loss'] < self.best_val_loss:
                                self.best_val_loss = val_metrics['val_loss']
                                self.save_checkpoint('best_model.pt')
                                print("Saved new best model")

                            self.model.train()

                        # Save checkpoint
                        if self.global_step % self.training_config.save_every == 0:
                            self.save_checkpoint(f'checkpoint_step_{self.global_step}.pt')
                            self._cleanup_old_checkpoints()
                    else:
                        accumulation_step += 1

                    pbar.update(1)

            # End of epoch evaluation
            val_metrics = self.evaluate()
            epoch_time = time.time() - epoch_start

            print(f"\nEnd of Epoch {epoch + 1}: Val Loss={val_metrics['val_loss']:.4f}, "
                  f"Cosine Sim={val_metrics['val_cosine_similarity']:.4f}, "
                  f"Time={epoch_time:.2f}s")

            # Save epoch checkpoint
            self.save_checkpoint(f'checkpoint_epoch_{epoch + 1}.pt')

            # Log epoch statistics
            self._log_epoch_stats(epoch, val_metrics, epoch_time)

        total_time = time.time() - start_time
        print("\nTraining complete!")
        print(f"Best validation loss: {self.best_val_loss:.4f}")
        print(f"Total training time: {total_time / 3600:.2f} hours")

        # Save final statistics
        self._save_training_stats()

    def save_checkpoint(self, filename: str):
        """Save checkpoint with enhanced state management."""
        checkpoint_path = os.path.join(self.training_config.checkpoint_dir, filename)

        # Handle FSDP checkpointing
        if self.distributed_manager and hasattr(self.model, '__wrapped__'):
            # FSDP model
            from torch.distributed.fsdp import FullyShardedDataParallel as FSDP
            if isinstance(self.model, FSDP):
                wrapper = FSDPWrapper(self.distributed_config)
                wrapper.save_checkpoint(
                    self.model,
                    self.optimizer,
                    checkpoint_path,
                    global_step=self.global_step,
                    epoch=self.epoch,
                    best_val_loss=self.best_val_loss,
                    model_config=self.model_config.to_dict(),
                    training_config=self.training_config.to_dict(),
                    mixed_precision_state=self.mixed_precision.state_dict(),
                    training_stats=self.training_stats,
                )
                return

        # Standard checkpointing
        checkpoint = {
            'model_state_dict': self.model.state_dict(),
            'optimizer_state_dict': self.optimizer.state_dict(),
            'scheduler_state_dict': self.scheduler.state_dict() if self.scheduler else None,
            'global_step': self.global_step,
            'epoch': self.epoch,
            'best_val_loss': self.best_val_loss,
            'model_config': self.model_config.to_dict(),
            'training_config': self.training_config.to_dict(),
            'mixed_precision_state': self.mixed_precision.state_dict(),
            'training_stats': self.training_stats,
        }

        # Save with atomic write
        temp_path = checkpoint_path + '.tmp'
        torch.save(checkpoint, temp_path)
        os.replace(temp_path, checkpoint_path)

        logger.info(f"Saved checkpoint to {checkpoint_path}")

    def load_checkpoint(self, checkpoint_path: str):
        """Load checkpoint with enhanced state management."""
        checkpoint = torch.load(checkpoint_path, map_location=self.device)

        self.model.load_state_dict(checkpoint['model_state_dict'])
        self.optimizer.load_state_dict(checkpoint['optimizer_state_dict'])
        if self.scheduler and checkpoint.get('scheduler_state_dict'):
            self.scheduler.load_state_dict(checkpoint['scheduler_state_dict'])

        self.global_step = checkpoint['global_step']
        self.epoch = checkpoint['epoch']
        self.best_val_loss = checkpoint['best_val_loss']

        # Load mixed precision state
        if 'mixed_precision_state' in checkpoint:
            self.mixed_precision.load_state_dict(checkpoint['mixed_precision_state'])

        # Load training stats
        if 'training_stats' in checkpoint:
            self.training_stats = checkpoint['training_stats']

        print(f"Loaded checkpoint from {checkpoint_path}")
        print(f"Resuming from epoch {self.epoch}, step {self.global_step}")

    def _cleanup_old_checkpoints(self):
        """Keep only N most recent checkpoints."""
        keep_n = self.training_config.keep_n_checkpoints
        if keep_n <= 0:
            return

        checkpoint_dir = Path(self.training_config.checkpoint_dir)
        step_checkpoints = sorted(
            checkpoint_dir.glob('checkpoint_step_*.pt'),
            key=lambda p: int(p.stem.split('_')[-1])
        )

        # Keep most recent N
        for old_checkpoint in step_checkpoints[:-keep_n]:
            old_checkpoint.unlink()
            logger.debug(f"Removed old checkpoint: {old_checkpoint}")

    def _log_training_info(self):
        """Log training configuration and model info."""
        # Model parameters
        total_params = sum(p.numel() for p in self.model.parameters())
        trainable_params = sum(p.numel() for p in self.model.parameters() if p.requires_grad)

        logger.info("=" * 60)
        logger.info("Training Configuration")
        logger.info("=" * 60)
        logger.info(f"Total parameters: {total_params:,}")
        logger.info(f"Trainable parameters: {trainable_params:,}")
        logger.info(f"Optimizer: {self.optimizer.__class__.__name__}")
        logger.info(f"Learning rate: {self.training_config.learning_rate}")
        logger.info(f"Batch size: {self.training_config.batch_size}")
        logger.info(f"Gradient accumulation: {self.gradient_accumulator.accumulation_steps}")
        logger.info(f"Effective batch size: {self.training_config.batch_size * self.gradient_accumulator.accumulation_steps}")
        logger.info(f"Mixed precision: {self.mixed_precision.config.precision}")
        logger.info(f"Max gradient norm: {self.training_config.max_grad_norm}")

        if self.distributed_manager:
            logger.info(f"Distributed: {self.distributed_config.world_size} processes")
            logger.info(f"Sharding strategy: {self.distributed_config.sharding_strategy}")

        logger.info("=" * 60)

    def _log_epoch_stats(self, epoch: int, val_metrics: Dict, epoch_time: float):
        """Log statistics for completed epoch."""
        if not self.training_stats['losses']:
            return

        # Compute statistics
        avg_loss = sum(self.training_stats['losses'][-100:]) / min(100, len(self.training_stats['losses']))
        avg_lr = sum(self.training_stats['learning_rates'][-100:]) / min(100, len(self.training_stats['learning_rates']))

        logger.info(f"Epoch {epoch + 1} stats:")
        logger.info(f"  Avg training loss: {avg_loss:.4f}")
        logger.info(f"  Val loss: {val_metrics['val_loss']:.4f}")
        logger.info(f"  Learning rate: {avg_lr:.2e}")
        logger.info(f"  Time: {epoch_time:.2f}s")

        if self.training_stats['grad_norms']:
            avg_grad_norm = sum(self.training_stats['grad_norms'][-100:]) / min(100, len(self.training_stats['grad_norms']))
            logger.info(f"  Avg grad norm: {avg_grad_norm:.2f}")

    def _save_training_stats(self):
        """Save training statistics to file."""
        stats_path = os.path.join(self.training_config.checkpoint_dir, 'training_stats.json')

        stats = {
            'final_step': self.global_step,
            'final_epoch': self.epoch,
            'best_val_loss': self.best_val_loss,
            'losses': self.training_stats['losses'][-1000:],  # Save last 1000
            'learning_rates': self.training_stats['learning_rates'][-1000:],
            'grad_norms': self.training_stats['grad_norms'][-1000:],
        }

        # Add mixed precision stats
        mp_stats = self.mixed_precision.get_stats()
        stats['mixed_precision'] = mp_stats

        # Add optimizer info
        opt_info = get_optimizer_info(self.optimizer)
        stats['optimizer'] = opt_info

        with open(stats_path, 'w') as f:
            json.dump(stats, f, indent=2)

        logger.info(f"Saved training statistics to {stats_path}")
