"""
QLoRA (Quantized LoRA) Implementation

4-bit quantization with NF4 (Normal Float 4) or FP4 format combined with LoRA.
Enables training large models on limited hardware with ~75% memory reduction.

Paper: https://arxiv.org/abs/2305.14314 (QLoRA)
"""

import torch
import torch.nn as nn
import torch.nn.functional as F
from typing import Optional, List, Dict, Tuple
import math
import numpy as np


# NF4 quantization lookup table
# These values are optimized for normally distributed weights
NF4_QUANT_TABLE = torch.tensor([
    -1.0,
    -0.6961928009986877,
    -0.5250730514526367,
    -0.39491748809814453,
    -0.28444138169288635,
    -0.18477343022823334,
    -0.09105003625154495,
    0.0,
    0.07958029955625534,
    0.16093020141124725,
    0.24611230194568634,
    0.33791524171829224,
    0.44070982933044434,
    0.5626170039176941,
    0.7229568362236023,
    1.0,
], dtype=torch.float32)


# FP4 quantization lookup table
FP4_QUANT_TABLE = torch.tensor([
    -1.0,
    -0.75,
    -0.5,
    -0.25,
    -0.125,
    -0.0625,
    -0.03125,
    0.0,
    0.03125,
    0.0625,
    0.125,
    0.25,
    0.5,
    0.75,
    1.0,
    1.5,
], dtype=torch.float32)


class NF4Quantizer:
    """
    4-bit NF4 (Normal Float 4) quantizer.

    NF4 is optimized for normally distributed weights, providing
    better quantization for neural network parameters.
    """

    def __init__(self, quant_type: str = "nf4"):
        """
        Args:
            quant_type: "nf4" or "fp4"
        """
        self.quant_type = quant_type

        if quant_type == "nf4":
            self.quant_table = NF4_QUANT_TABLE
        elif quant_type == "fp4":
            self.quant_table = FP4_QUANT_TABLE
        else:
            raise ValueError(f"Unknown quant_type: {quant_type}")

    def quantize(
        self,
        weight: torch.Tensor,
        absmax: Optional[torch.Tensor] = None
    ) -> Tuple[torch.Tensor, torch.Tensor]:
        """
        Quantize weights to 4-bit.

        Args:
            weight: Weight tensor to quantize
            absmax: Absolute max values (if None, compute from weight)

        Returns:
            quantized: Quantized weights (stored as uint8, using only 4 bits)
            absmax: Absolute max values for dequantization
        """
        # Compute absolute max if not provided
        if absmax is None:
            absmax = weight.abs().max()

        # Normalize to [-1, 1]
        weight_normalized = weight / (absmax + 1e-8)

        # Quantize: find nearest value in quantization table
        quant_table = self.quant_table.to(weight.device)

        # Expand dims for broadcasting
        weight_expanded = weight_normalized.unsqueeze(-1)  # [..., 1]
        quant_table_expanded = quant_table.view([1] * weight_normalized.ndim + [-1])  # [1, 1, ..., 16]

        # Find nearest value
        distances = torch.abs(weight_expanded - quant_table_expanded)
        indices = distances.argmin(dim=-1)  # [...] with values in [0, 15]

        # Store as uint8 (even though we only use 4 bits)
        quantized = indices.to(torch.uint8)

        return quantized, absmax

    def dequantize(
        self,
        quantized: torch.Tensor,
        absmax: torch.Tensor
    ) -> torch.Tensor:
        """
        Dequantize 4-bit weights back to float.

        Args:
            quantized: Quantized weights (uint8)
            absmax: Absolute max values

        Returns:
            weight: Dequantized weights
        """
        quant_table = self.quant_table.to(quantized.device)

        # Look up values
        weight_normalized = quant_table[quantized.long()]

        # Denormalize
        weight = weight_normalized * absmax

        return weight


class DoubleQuantizer:
    """
    Double quantization: quantizes the absmax values themselves.

    This provides additional memory savings with minimal accuracy loss.
    """

    def __init__(self, block_size: int = 256):
        """
        Args:
            block_size: Block size for second-level quantization
        """
        self.block_size = block_size

    def quantize(
        self,
        absmax: torch.Tensor
    ) -> Tuple[torch.Tensor, torch.Tensor, torch.Tensor]:
        """
        Double quantize absmax values.

        Args:
            absmax: Absolute max values [num_blocks]

        Returns:
            quantized_absmax: Quantized absmax values (int8)
            absmax_scale: Scale factors for absmax
            absmax_offset: Offset for absmax quantization
        """
        # Compute scale and offset for absmax
        absmax_min = absmax.min()
        absmax_max = absmax.max()

        absmax_scale = (absmax_max - absmax_min) / 255.0
        absmax_offset = absmax_min

        # Quantize to int8
        quantized_absmax = ((absmax - absmax_offset) / (absmax_scale + 1e-8)).round().clamp(0, 255).to(torch.uint8)

        return quantized_absmax, absmax_scale, absmax_offset

    def dequantize(
        self,
        quantized_absmax: torch.Tensor,
        absmax_scale: torch.Tensor,
        absmax_offset: torch.Tensor
    ) -> torch.Tensor:
        """
        Dequantize absmax values.

        Args:
            quantized_absmax: Quantized absmax (uint8)
            absmax_scale: Scale factor
            absmax_offset: Offset

        Returns:
            absmax: Dequantized absmax values
        """
        absmax = quantized_absmax.float() * absmax_scale + absmax_offset
        return absmax


class QLoRALinear(nn.Module):
    """
    QLoRA Linear layer: 4-bit quantized base weights + LoRA adapters.

    Memory efficient: Base weights stored in 4-bit, only LoRA adapters
    are trained in full precision.
    """

    def __init__(
        self,
        in_features: int,
        out_features: int,
        r: int = 8,
        lora_alpha: float = 16.0,
        lora_dropout: float = 0.1,
        quant_type: str = "nf4",
        double_quant: bool = True,
        block_size: int = 64,
        compute_dtype: torch.dtype = torch.float16,
    ):
        super().__init__()

        self.in_features = in_features
        self.out_features = out_features
        self.r = r
        self.lora_alpha = lora_alpha
        self.scaling = lora_alpha / r
        self.quant_type = quant_type
        self.double_quant = double_quant
        self.block_size = block_size
        self.compute_dtype = compute_dtype

        # Quantizer
        self.quantizer = NF4Quantizer(quant_type=quant_type)
        if double_quant:
            self.double_quantizer = DoubleQuantizer(block_size=block_size)

        # Placeholder for quantized weights
        # Will be set via quantize_weight method
        self.register_buffer('quantized_weight', None)
        self.register_buffer('weight_absmax', None)

        # Double quantization buffers
        if double_quant:
            self.register_buffer('quantized_absmax', None)
            self.register_buffer('absmax_scale', None)
            self.register_buffer('absmax_offset', None)

        # Bias (not quantized)
        self.bias = nn.Parameter(torch.zeros(out_features))

        # LoRA adapters (trainable)
        if r > 0:
            self.lora_A = nn.Parameter(torch.zeros(r, in_features))
            self.lora_B = nn.Parameter(torch.zeros(out_features, r))

            # Initialize
            nn.init.kaiming_uniform_(self.lora_A, a=math.sqrt(5))
            nn.init.zeros_(self.lora_B)

            # Dropout
            if lora_dropout > 0.0:
                self.lora_dropout = nn.Dropout(p=lora_dropout)
            else:
                self.lora_dropout = nn.Identity()

    def quantize_weight(self, weight: torch.Tensor):
        """
        Quantize and store weight tensor.

        Args:
            weight: [out_features, in_features] weight tensor
        """
        # Blockwise quantization for better accuracy
        out_features, in_features = weight.shape

        # Determine number of blocks
        num_blocks = (in_features + self.block_size - 1) // self.block_size

        quantized_blocks = []
        absmax_blocks = []

        # Quantize each block
        for i in range(num_blocks):
            start_idx = i * self.block_size
            end_idx = min((i + 1) * self.block_size, in_features)

            block = weight[:, start_idx:end_idx]

            # Quantize block
            block_absmax = block.abs().max()
            quantized_block, _ = self.quantizer.quantize(block, absmax=block_absmax)

            quantized_blocks.append(quantized_block)
            absmax_blocks.append(block_absmax)

        # Concatenate blocks
        self.quantized_weight = torch.cat(quantized_blocks, dim=1)
        weight_absmax = torch.stack(absmax_blocks)

        # Apply double quantization if enabled
        if self.double_quant:
            quantized_absmax, absmax_scale, absmax_offset = self.double_quantizer.quantize(weight_absmax)
            self.quantized_absmax = quantized_absmax
            self.absmax_scale = absmax_scale
            self.absmax_offset = absmax_offset
        else:
            self.weight_absmax = weight_absmax

    def dequantize_weight(self) -> torch.Tensor:
        """
        Dequantize weight tensor.

        Returns:
            weight: Dequantized weight tensor
        """
        if self.quantized_weight is None:
            raise ValueError("Weight has not been quantized yet")

        # Dequantize absmax if double quantization was used
        if self.double_quant:
            weight_absmax = self.double_quantizer.dequantize(
                self.quantized_absmax,
                self.absmax_scale,
                self.absmax_offset
            )
        else:
            weight_absmax = self.weight_absmax

        # Dequantize blocks
        num_blocks = len(weight_absmax)
        dequantized_blocks = []

        for i in range(num_blocks):
            start_idx = i * self.block_size
            end_idx = min((i + 1) * self.block_size, self.in_features)

            quantized_block = self.quantized_weight[:, start_idx:end_idx]
            absmax = weight_absmax[i]

            # Dequantize
            dequantized_block = self.quantizer.dequantize(quantized_block, absmax)
            dequantized_blocks.append(dequantized_block)

        # Concatenate
        weight = torch.cat(dequantized_blocks, dim=1)

        return weight

    def forward(self, x: torch.Tensor) -> torch.Tensor:
        """
        Forward pass with dequantized weights + LoRA.

        Args:
            x: [batch, seq_len, in_features] or [batch, in_features]

        Returns:
            output: [batch, seq_len, out_features] or [batch, out_features]
        """
        # Dequantize weights for computation
        weight = self.dequantize_weight()

        # Convert to compute dtype
        weight = weight.to(self.compute_dtype)
        x = x.to(self.compute_dtype)

        # Base linear transformation
        output = F.linear(x, weight, self.bias.to(self.compute_dtype))

        # Add LoRA transformation
        if self.r > 0:
            lora_x = x @ self.lora_A.to(self.compute_dtype).T
            lora_x = self.lora_dropout(lora_x)
            lora_out = lora_x @ self.lora_B.to(self.compute_dtype).T
            output = output + lora_out * self.scaling

        return output

    def extra_repr(self) -> str:
        """Extra representation for debugging."""
        return (f'in_features={self.in_features}, out_features={self.out_features}, '
                f'r={self.r}, quant={self.quant_type}, double_quant={self.double_quant}')


def apply_qlora_to_model(
    model: nn.Module,
    target_modules: List[str],
    r: int = 8,
    lora_alpha: float = 16.0,
    lora_dropout: float = 0.1,
    quant_type: str = "nf4",
    double_quant: bool = True,
    block_size: int = 64,
    compute_dtype: torch.dtype = torch.float16,
) -> nn.Module:
    """
    Apply QLoRA to target modules in a model.

    Args:
        model: Model to apply QLoRA to
        target_modules: List of module names to replace
        r: LoRA rank
        lora_alpha: LoRA alpha
        lora_dropout: LoRA dropout
        quant_type: "nf4" or "fp4"
        double_quant: Enable double quantization
        block_size: Block size for quantization
        compute_dtype: Computation dtype

    Returns:
        Modified model with QLoRA layers
    """
    replacements = []

    for name, module in model.named_modules():
        # Check if this module should be replaced
        should_replace = False
        for target in target_modules:
            if target in name and isinstance(module, nn.Linear):
                should_replace = True
                break

        if should_replace:
            # Get parent module
            parent_name = '.'.join(name.split('.')[:-1])
            attr_name = name.split('.')[-1]

            if parent_name:
                parent = model.get_submodule(parent_name)
            else:
                parent = model

            # Create QLoRA layer
            qlora_layer = QLoRALinear(
                in_features=module.in_features,
                out_features=module.out_features,
                r=r,
                lora_alpha=lora_alpha,
                lora_dropout=lora_dropout,
                quant_type=quant_type,
                double_quant=double_quant,
                block_size=block_size,
                compute_dtype=compute_dtype,
            )

            # Quantize original weights
            qlora_layer.quantize_weight(module.weight.data)

            # Copy bias
            if module.bias is not None:
                qlora_layer.bias.data = module.bias.data.clone()

            # Replace module
            setattr(parent, attr_name, qlora_layer)
            replacements.append(name)

    print(f"Applied QLoRA to {len(replacements)} modules:")
    for name in replacements:
        print(f"  - {name}")

    return model


def estimate_qlora_memory_savings(
    num_parameters: int,
    r: int = 8,
    base_dtype_bytes: int = 2,  # fp16
) -> Dict[str, float]:
    """
    Estimate memory savings from QLoRA.

    Args:
        num_parameters: Number of parameters in base model
        r: LoRA rank
        base_dtype_bytes: Bytes per parameter in base model (2 for fp16, 4 for fp32)

    Returns:
        Dictionary with memory statistics in GB
    """
    # Base model memory (without quantization)
    base_memory = num_parameters * base_dtype_bytes

    # QLoRA memory:
    # - Quantized weights: 0.5 bytes per param (4-bit)
    # - LoRA adapters: Depends on architecture, approximate as r * (in + out) per layer
    # For simplicity, assume LoRA adds ~1% of base parameters
    quantized_memory = num_parameters * 0.5  # 4-bit
    lora_memory = num_parameters * 0.01 * base_dtype_bytes  # ~1% in fp16/fp32

    qlora_total = quantized_memory + lora_memory

    savings = base_memory - qlora_total
    savings_percent = (savings / base_memory) * 100

    return {
        'base_memory_gb': base_memory / (1024**3),
        'qlora_memory_gb': qlora_total / (1024**3),
        'savings_gb': savings / (1024**3),
        'savings_percent': savings_percent,
    }
