"""
Quantized Model Wrapper

Unified interface for applying different quantization methods to MK3 models.
"""

import torch
import torch.nn as nn
from typing import Optional, Dict, Any
import os
from pathlib import Path

from .config import QuantizationConfig, LoRAConfig, QLoRAConfig, GPTQConfig, AWQConfig
from .lora import apply_lora_to_model, count_lora_parameters
from .qlora import apply_qlora_to_model
from .gptq import quantize_model_gptq
from .awq import quantize_model_awq
from .kv_cache_quant import QuantizedKVCache
from .memory_utils import estimate_model_memory, print_memory_summary


class QuantizedMK3Model(nn.Module):
    """
    Wrapper for MK3 models with quantization support.

    Provides unified interface for applying LoRA, QLoRA, GPTQ, AWQ,
    and KV cache quantization to MK3 continuous autoregressive models.
    """

    def __init__(
        self,
        base_model: nn.Module,
        quant_config: QuantizationConfig,
        device: Optional[torch.device] = None,
    ):
        """
        Args:
            base_model: Base MK3 model
            quant_config: Quantization configuration
            device: Device to use
        """
        super().__init__()

        self.quant_config = quant_config
        self.device = device or torch.device('cuda' if torch.cuda.is_available() else 'cpu')
        self.base_model = base_model

        # Apply quantization
        self._apply_quantization()

        # Move to device
        self.to(self.device)

        # Print memory summary
        if quant_config.method != "none":
            print(f"\nApplied {quant_config.method.upper()} quantization")
            self.print_memory_stats()

    def _apply_quantization(self):
        """Apply quantization method based on config."""
        method = self.quant_config.method

        if method == "none":
            # No quantization
            pass

        elif method == "lora":
            # Apply LoRA
            config = self.quant_config.lora_config
            self.base_model = apply_lora_to_model(
                model=self.base_model,
                target_modules=config.target_modules,
                r=config.r,
                lora_alpha=config.lora_alpha,
                lora_dropout=config.lora_dropout,
                merge_weights=config.merge_weights,
                use_dora=False,
            )

        elif method == "dora":
            # Apply DoRA
            config = self.quant_config.dora_config
            self.base_model = apply_lora_to_model(
                model=self.base_model,
                target_modules=config.target_modules,
                r=config.r,
                lora_alpha=config.lora_alpha,
                lora_dropout=config.lora_dropout,
                merge_weights=config.merge_weights,
                use_dora=True,
            )

        elif method == "qlora":
            # Apply QLoRA
            config = self.quant_config.qlora_config

            # Determine compute dtype
            if config.compute_dtype == "float16":
                compute_dtype = torch.float16
            elif config.compute_dtype == "bfloat16":
                compute_dtype = torch.bfloat16
            else:
                compute_dtype = torch.float32

            self.base_model = apply_qlora_to_model(
                model=self.base_model,
                target_modules=config.target_modules,
                r=config.r,
                lora_alpha=config.lora_alpha,
                lora_dropout=config.lora_dropout,
                quant_type=config.quant_type,
                double_quant=config.double_quant,
                compute_dtype=compute_dtype,
            )

        elif method == "gptq":
            # GPTQ requires calibration data
            raise ValueError(
                "GPTQ requires calibration data. Use quantize_model_gptq() directly "
                "or call apply_gptq() method with data_loader."
            )

        elif method == "awq":
            # AWQ requires calibration data
            raise ValueError(
                "AWQ requires calibration data. Use quantize_model_awq() directly "
                "or call apply_awq() method with data_loader."
            )

        else:
            raise ValueError(f"Unknown quantization method: {method}")

        # Apply gradient checkpointing if requested
        if self.quant_config.use_gradient_checkpointing:
            if hasattr(self.base_model, 'gradient_checkpointing_enable'):
                self.base_model.gradient_checkpointing_enable()
            else:
                print("Warning: Gradient checkpointing not available for this model")

    def apply_gptq(self, data_loader, num_calibration_batches: int = 128):
        """
        Apply GPTQ quantization with calibration data.

        Args:
            data_loader: Calibration data loader
            num_calibration_batches: Number of batches for calibration
        """
        if self.quant_config.method != "gptq":
            raise ValueError("Model config method must be 'gptq' to use this method")

        config = self.quant_config.gptq_config

        self.base_model = quantize_model_gptq(
            model=self.base_model,
            data_loader=data_loader,
            target_modules=["query_proj", "key_proj", "value_proj", "output_proj"],
            bits=config.bits,
            group_size=config.group_size,
            damp_percent=config.damp_percent,
            device=self.device,
            num_calibration_batches=num_calibration_batches,
        )

        print("Applied GPTQ quantization")
        self.print_memory_stats()

    def apply_awq(self, data_loader, num_calibration_batches: int = 128):
        """
        Apply AWQ quantization with calibration data.

        Args:
            data_loader: Calibration data loader
            num_calibration_batches: Number of batches for calibration
        """
        if self.quant_config.method != "awq":
            raise ValueError("Model config method must be 'awq' to use this method")

        config = self.quant_config.awq_config

        self.base_model = quantize_model_awq(
            model=self.base_model,
            data_loader=data_loader,
            target_modules=["query_proj", "key_proj", "value_proj", "output_proj"],
            bits=config.bits,
            group_size=config.group_size,
            device=self.device,
            num_calibration_batches=num_calibration_batches,
        )

        print("Applied AWQ quantization")
        self.print_memory_stats()

    def forward(self, *args, **kwargs):
        """Forward pass through base model."""
        return self.base_model(*args, **kwargs)

    def generate(self, *args, **kwargs):
        """Generation through base model."""
        if hasattr(self.base_model, 'generate'):
            return self.base_model.generate(*args, **kwargs)
        else:
            raise AttributeError("Base model does not have generate method")

    def print_memory_stats(self):
        """Print memory statistics for the quantized model."""
        print_memory_summary(
            model=self.base_model,
            batch_size=1,
            seq_length=512,
            dtype=torch.float16,
            include_training=(self.quant_config.method in ["lora", "dora", "qlora"])
        )

    def save_pretrained(self, save_directory: str):
        """
        Save quantized model.

        For LoRA/QLoRA, saves only adapter weights.
        For GPTQ/AWQ, saves full quantized model.
        """
        os.makedirs(save_directory, exist_ok=True)

        # Save config
        config_path = os.path.join(save_directory, "quantization_config.json")
        self.quant_config.save(config_path)

        # Save model based on method
        if self.quant_config.method in ["lora", "dora", "qlora"]:
            # Save only LoRA weights
            from .lora import save_lora_weights
            model_path = os.path.join(save_directory, "adapter_model.pt")
            save_lora_weights(self.base_model, model_path)

        else:
            # Save full model
            model_path = os.path.join(save_directory, "pytorch_model.pt")
            torch.save(self.base_model.state_dict(), model_path)

        print(f"Model saved to {save_directory}")

    @classmethod
    def from_pretrained(
        cls,
        base_model: nn.Module,
        save_directory: str,
        device: Optional[torch.device] = None,
    ):
        """
        Load quantized model from directory.

        Args:
            base_model: Base model (for LoRA/QLoRA, this should be the original model)
            save_directory: Directory with saved model
            device: Device

        Returns:
            QuantizedMK3Model instance
        """
        # Load config
        config_path = os.path.join(save_directory, "quantization_config.json")
        quant_config = QuantizationConfig.load(config_path)

        # Create quantized model
        model = cls(base_model=base_model, quant_config=quant_config, device=device)

        # Load weights
        if quant_config.method in ["lora", "dora", "qlora"]:
            # Load LoRA weights
            from .lora import load_lora_weights
            model_path = os.path.join(save_directory, "adapter_model.pt")
            load_lora_weights(model.base_model, model_path)

        else:
            # Load full model
            model_path = os.path.join(save_directory, "pytorch_model.pt")
            state_dict = torch.load(model_path, map_location=device)
            model.base_model.load_state_dict(state_dict)

        print(f"Model loaded from {save_directory}")

        return model

    def get_memory_footprint(self) -> Dict[str, float]:
        """
        Get detailed memory footprint.

        Returns:
            Dictionary with memory statistics
        """
        memory_stats = estimate_model_memory(
            model=self.base_model,
            dtype=torch.float16,
            include_gradients=False,
            include_optimizer=False
        )

        # Add quantization-specific info
        if self.quant_config.method in ["lora", "dora", "qlora"]:
            param_stats = count_lora_parameters(self.base_model)
            memory_stats['lora_parameters'] = param_stats['lora_parameters']
            memory_stats['trainable_percentage'] = param_stats['trainable_percentage']

        return memory_stats

    def enable_kv_cache_quantization(
        self,
        max_batch_size: int = 1,
        max_seq_length: int = 2048,
        bits: int = 8,
    ):
        """
        Enable KV cache quantization for inference.

        Args:
            max_batch_size: Maximum batch size
            max_seq_length: Maximum sequence length
            bits: Quantization bits (4 or 8)
        """
        # This would require modifying attention layers
        # For now, store config for future use
        from .config import KVCacheQuantConfig

        kv_config = KVCacheQuantConfig(bits=bits)
        self.quant_config.kv_cache_config = kv_config

        print(f"KV cache quantization enabled ({bits}-bit)")
        print(f"Note: KV cache quantization requires attention layer modifications")

    def __repr__(self):
        """String representation."""
        return (f"QuantizedMK3Model(\n"
                f"  method={self.quant_config.method},\n"
                f"  base_model={self.base_model.__class__.__name__}\n"
                f")")
