"""
Quantization Configuration Classes

Defines configuration for all quantization methods.
"""

from dataclasses import dataclass, field, asdict
from typing import Optional, List, Dict, Any
import json


@dataclass
class LoRAConfig:
    """Configuration for LoRA (Low-Rank Adaptation)."""

    r: int = 8  # Rank of low-rank matrices
    lora_alpha: float = 16.0  # LoRA scaling factor
    lora_dropout: float = 0.1  # Dropout for LoRA layers
    target_modules: List[str] = field(default_factory=lambda: [
        "query_proj", "key_proj", "value_proj", "output_proj"
    ])  # Which modules to apply LoRA to
    merge_weights: bool = False  # Whether to merge LoRA weights into base model
    bias: str = "none"  # Bias handling: "none", "all", or "lora_only"
    fan_in_fan_out: bool = False  # Set to True for Conv1D layers

    def to_dict(self) -> Dict[str, Any]:
        """Convert to dictionary."""
        return asdict(self)

    @classmethod
    def from_dict(cls, config_dict: Dict[str, Any]) -> 'LoRAConfig':
        """Create from dictionary."""
        return cls(**config_dict)

    def save(self, path: str):
        """Save configuration to JSON file."""
        with open(path, 'w') as f:
            json.dump(self.to_dict(), f, indent=2)

    @classmethod
    def load(cls, path: str) -> 'LoRAConfig':
        """Load configuration from JSON file."""
        with open(path, 'r') as f:
            config_dict = json.load(f)
        return cls.from_dict(config_dict)


@dataclass
class DoRAConfig(LoRAConfig):
    """Configuration for DoRA (Weight-Decomposed LoRA)."""

    use_dora: bool = True  # Enable DoRA decomposition
    magnitude_init: str = "uniform"  # Magnitude initialization: "uniform", "normal", "constant"


@dataclass
class QLoRAConfig(LoRAConfig):
    """Configuration for QLoRA (Quantized LoRA)."""

    quant_bits: int = 4  # Quantization bits (4 or 8)
    quant_type: str = "nf4"  # Quantization type: "nf4" or "fp4"
    double_quant: bool = True  # Use double quantization
    compute_dtype: str = "float16"  # Computation dtype: "float16" or "bfloat16"
    quant_storage: str = "uint8"  # Storage dtype for quantized weights


@dataclass
class GPTQConfig:
    """Configuration for GPTQ quantization."""

    bits: int = 4  # Quantization bits
    group_size: int = 128  # Group size for quantization
    damp_percent: float = 0.01  # Dampening factor for Hessian
    desc_act: bool = False  # Use descending activation order
    sym: bool = True  # Symmetric quantization
    true_sequential: bool = True  # Sequential quantization
    static_groups: bool = False  # Use static groups
    actorder: bool = False  # Activation order optimization

    def to_dict(self) -> Dict[str, Any]:
        """Convert to dictionary."""
        return asdict(self)

    @classmethod
    def from_dict(cls, config_dict: Dict[str, Any]) -> 'GPTQConfig':
        """Create from dictionary."""
        return cls(**config_dict)


@dataclass
class AWQConfig:
    """Configuration for AWQ (Activation-aware Weight Quantization)."""

    bits: int = 4  # Quantization bits
    group_size: int = 128  # Group size for quantization
    zero_point: bool = True  # Use zero-point quantization
    version: str = "GEMM"  # AWQ version: "GEMM" or "GEMV"
    w_bit: int = 4  # Weight quantization bits
    q_group_size: int = 128  # Quantization group size

    def to_dict(self) -> Dict[str, Any]:
        """Convert to dictionary."""
        return asdict(self)

    @classmethod
    def from_dict(cls, config_dict: Dict[str, Any]) -> 'AWQConfig':
        """Create from dictionary."""
        return cls(**config_dict)


@dataclass
class KVCacheQuantConfig:
    """Configuration for KV cache quantization."""

    bits: int = 8  # Quantization bits (4 or 8)
    cache_dtype: str = "int8"  # Cache storage dtype: "int8" or "int4"
    per_channel: bool = True  # Per-channel quantization
    symmetric: bool = True  # Symmetric quantization
    dynamic: bool = False  # Dynamic quantization (recalculate scales)

    def to_dict(self) -> Dict[str, Any]:
        """Convert to dictionary."""
        return asdict(self)


@dataclass
class QuantizationConfig:
    """Master configuration for all quantization methods."""

    method: str = "none"  # "none", "lora", "dora", "qlora", "gptq", "awq"
    lora_config: Optional[LoRAConfig] = None
    dora_config: Optional[DoRAConfig] = None
    qlora_config: Optional[QLoRAConfig] = None
    gptq_config: Optional[GPTQConfig] = None
    awq_config: Optional[AWQConfig] = None
    kv_cache_config: Optional[KVCacheQuantConfig] = None

    # Global settings
    use_gradient_checkpointing: bool = True
    offload_to_cpu: bool = False
    offload_folder: Optional[str] = None

    def __post_init__(self):
        """Initialize default configs if method is specified."""
        if self.method == "lora" and self.lora_config is None:
            self.lora_config = LoRAConfig()
        elif self.method == "dora" and self.dora_config is None:
            self.dora_config = DoRAConfig()
        elif self.method == "qlora" and self.qlora_config is None:
            self.qlora_config = QLoRAConfig()
        elif self.method == "gptq" and self.gptq_config is None:
            self.gptq_config = GPTQConfig()
        elif self.method == "awq" and self.awq_config is None:
            self.awq_config = AWQConfig()

    def to_dict(self) -> Dict[str, Any]:
        """Convert to dictionary."""
        config = {
            'method': self.method,
            'use_gradient_checkpointing': self.use_gradient_checkpointing,
            'offload_to_cpu': self.offload_to_cpu,
            'offload_folder': self.offload_folder,
        }

        if self.lora_config:
            config['lora_config'] = self.lora_config.to_dict()
        if self.dora_config:
            config['dora_config'] = self.dora_config.to_dict()
        if self.qlora_config:
            config['qlora_config'] = self.qlora_config.to_dict()
        if self.gptq_config:
            config['gptq_config'] = self.gptq_config.to_dict()
        if self.awq_config:
            config['awq_config'] = self.awq_config.to_dict()
        if self.kv_cache_config:
            config['kv_cache_config'] = self.kv_cache_config.to_dict()

        return config

    def save(self, path: str):
        """Save configuration to JSON file."""
        with open(path, 'w') as f:
            json.dump(self.to_dict(), f, indent=2)

    @classmethod
    def load(cls, path: str) -> 'QuantizationConfig':
        """Load configuration from JSON file."""
        with open(path, 'r') as f:
            config_dict = json.load(f)

        # Reconstruct nested configs
        if 'lora_config' in config_dict:
            config_dict['lora_config'] = LoRAConfig.from_dict(config_dict['lora_config'])
        if 'dora_config' in config_dict:
            config_dict['dora_config'] = DoRAConfig.from_dict(config_dict['dora_config'])
        if 'qlora_config' in config_dict:
            config_dict['qlora_config'] = QLoRAConfig.from_dict(config_dict['qlora_config'])
        if 'gptq_config' in config_dict:
            config_dict['gptq_config'] = GPTQConfig.from_dict(config_dict['gptq_config'])
        if 'awq_config' in config_dict:
            config_dict['awq_config'] = AWQConfig.from_dict(config_dict['awq_config'])
        if 'kv_cache_config' in config_dict:
            config_dict['kv_cache_config'] = KVCacheQuantConfig.from_dict(config_dict['kv_cache_config'])

        return cls(**config_dict)
