"""
GPTQ (Gradient-based Post-Training Quantization)

Optimal weight quantization using second-order information (Hessian).
GPTQ minimizes quantization error by leveraging layer-wise OBQ algorithm.

Paper: https://arxiv.org/abs/2210.17323
"""

import torch
import torch.nn as nn
from typing import Optional, List, Dict, Tuple, Callable
import numpy as np
from tqdm import tqdm


class GPTQuantizer:
    """
    GPTQ quantizer for neural network layers.

    Uses layer-wise optimal brain quantization with Hessian approximation.
    """

    def __init__(
        self,
        bits: int = 4,
        group_size: int = 128,
        damp_percent: float = 0.01,
        desc_act: bool = False,
        sym: bool = True,
        actorder: bool = False,
    ):
        """
        Args:
            bits: Quantization bits (4 or 8)
            group_size: Group size for quantization (-1 for per-channel)
            damp_percent: Dampening factor for Hessian diagonal
            desc_act: Use descending activation order
            sym: Symmetric quantization
            actorder: Optimize quantization order based on activations
        """
        self.bits = bits
        self.group_size = group_size
        self.damp_percent = damp_percent
        self.desc_act = desc_act
        self.sym = sym
        self.actorder = actorder

        self.maxq = 2 ** bits - 1

    def quantize_tensor(
        self,
        tensor: torch.Tensor,
        scale: torch.Tensor,
        zero_point: Optional[torch.Tensor] = None,
    ) -> torch.Tensor:
        """
        Quantize a tensor given scale and zero point.

        Args:
            tensor: Tensor to quantize
            scale: Scale factors
            zero_point: Zero points (None for symmetric)

        Returns:
            Quantized tensor
        """
        if zero_point is None:
            # Symmetric quantization
            quantized = torch.clamp(
                torch.round(tensor / scale),
                -self.maxq // 2,
                self.maxq // 2
            )
        else:
            # Asymmetric quantization
            quantized = torch.clamp(
                torch.round(tensor / scale + zero_point),
                0,
                self.maxq
            )

        return quantized.to(torch.int8)

    def dequantize_tensor(
        self,
        quantized: torch.Tensor,
        scale: torch.Tensor,
        zero_point: Optional[torch.Tensor] = None,
    ) -> torch.Tensor:
        """
        Dequantize a tensor.

        Args:
            quantized: Quantized tensor
            scale: Scale factors
            zero_point: Zero points (None for symmetric)

        Returns:
            Dequantized tensor
        """
        quantized = quantized.float()

        if zero_point is None:
            # Symmetric
            dequantized = quantized * scale
        else:
            # Asymmetric
            dequantized = (quantized - zero_point) * scale

        return dequantized

    def find_params(
        self,
        weight: torch.Tensor,
        group_size: Optional[int] = None,
    ) -> Tuple[torch.Tensor, Optional[torch.Tensor]]:
        """
        Find quantization parameters (scale and zero point).

        Args:
            weight: Weight tensor [out_features, in_features]
            group_size: Group size (if None, use self.group_size)

        Returns:
            scale: Scale factors
            zero_point: Zero points (None if symmetric)
        """
        group_size = group_size or self.group_size

        out_features, in_features = weight.shape

        if group_size == -1:
            # Per-channel quantization
            group_size = in_features

        # Reshape for groupwise quantization
        num_groups = (in_features + group_size - 1) // group_size

        scales = []
        zero_points = [] if not self.sym else None

        for i in range(num_groups):
            start_idx = i * group_size
            end_idx = min((i + 1) * group_size, in_features)

            group = weight[:, start_idx:end_idx]

            # Find min and max
            w_min = group.min(dim=1, keepdim=True)[0]
            w_max = group.max(dim=1, keepdim=True)[0]

            if self.sym:
                # Symmetric quantization
                max_abs = torch.max(w_min.abs(), w_max.abs())
                scale = max_abs / (self.maxq // 2)
                scale = scale.clamp(min=1e-8)
                scales.append(scale)
            else:
                # Asymmetric quantization
                scale = (w_max - w_min) / self.maxq
                scale = scale.clamp(min=1e-8)
                zero = torch.round(-w_min / scale)

                scales.append(scale)
                zero_points.append(zero)

        # Concatenate scales
        scale = torch.cat(scales, dim=1)

        if not self.sym:
            zero_point = torch.cat(zero_points, dim=1)
        else:
            zero_point = None

        return scale, zero_point

    def quantize_layer_gptq(
        self,
        layer: nn.Linear,
        H: torch.Tensor,
    ) -> Tuple[torch.Tensor, torch.Tensor, Optional[torch.Tensor]]:
        """
        Quantize a linear layer using GPTQ algorithm.

        Args:
            layer: Linear layer to quantize
            H: Hessian matrix (approximated as X^T X / N)

        Returns:
            quantized_weight: Quantized weight
            scale: Quantization scale
            zero_point: Zero point (None if symmetric)
        """
        W = layer.weight.data.clone()
        out_features, in_features = W.shape

        # Add dampening to Hessian diagonal
        damp = self.damp_percent * torch.mean(torch.diag(H))
        diag_indices = torch.arange(in_features, device=H.device)
        H[diag_indices, diag_indices] += damp

        # Cholesky decomposition
        try:
            H_inv = torch.linalg.cholesky(H)
            H_inv = torch.cholesky_inverse(H_inv)
        except:
            # Fallback to pseudo-inverse if Cholesky fails
            H_inv = torch.linalg.pinv(H)

        # Quantization order (default: sequential)
        if self.actorder:
            # Order by activation magnitude (diagonal of Hessian)
            perm = torch.argsort(torch.diag(H), descending=True)
        else:
            perm = torch.arange(in_features, device=W.device)

        # Get quantization parameters
        scale, zero_point = self.find_params(W)

        # GPTQ algorithm: quantize one column at a time
        Q = torch.zeros_like(W)
        Losses = torch.zeros_like(W)
        Err = torch.zeros_like(W)

        for i in range(in_features):
            # Get column index
            col_idx = perm[i].item()

            # Current column
            w_col = W[:, col_idx].clone()

            # Quantize column
            if self.group_size == -1:
                s = scale[:, 0]
                z = zero_point[:, 0] if zero_point is not None else None
            else:
                group_idx = col_idx // self.group_size
                s = scale[:, group_idx].unsqueeze(1)
                z = zero_point[:, group_idx].unsqueeze(1) if zero_point is not None else None

            q_col = self.quantize_tensor(w_col.unsqueeze(1), s, z)
            q_col_deq = self.dequantize_tensor(q_col, s, z).squeeze(1)

            Q[:, col_idx] = q_col_deq

            # Compute error
            err = (w_col - q_col_deq) / H_inv[col_idx, col_idx]
            Err[:, col_idx] = err

            # Update remaining weights to compensate for quantization error
            W[:, col_idx:] -= err.unsqueeze(1) * H_inv[col_idx, col_idx:].unsqueeze(0)

        # Final quantization
        quantized_weight = self.quantize_tensor(Q, scale, zero_point)

        return quantized_weight, scale, zero_point


class QuantizedLinear(nn.Module):
    """
    Quantized linear layer (GPTQ).

    Stores weights in quantized format and dequantizes on-the-fly.
    """

    def __init__(
        self,
        in_features: int,
        out_features: int,
        bits: int = 4,
        group_size: int = 128,
        bias: bool = True,
    ):
        super().__init__()

        self.in_features = in_features
        self.out_features = out_features
        self.bits = bits
        self.group_size = group_size

        # Buffers for quantized weights
        self.register_buffer('quantized_weight', torch.zeros((out_features, in_features), dtype=torch.int8))
        self.register_buffer('weight_scale', torch.zeros((out_features, (in_features + group_size - 1) // group_size)))
        self.register_buffer('weight_zero_point', None)

        # Bias
        if bias:
            self.bias = nn.Parameter(torch.zeros(out_features))
        else:
            self.register_parameter('bias', None)

    def forward(self, x: torch.Tensor) -> torch.Tensor:
        """
        Forward pass with dequantized weights.

        Args:
            x: Input tensor

        Returns:
            Output tensor
        """
        # Dequantize weights
        quantizer = GPTQuantizer(bits=self.bits, group_size=self.group_size)
        weight = quantizer.dequantize_tensor(
            self.quantized_weight,
            self.weight_scale,
            self.weight_zero_point
        )

        # Linear transformation
        output = torch.nn.functional.linear(x, weight, self.bias)

        return output


def collect_layer_inputs(
    model: nn.Module,
    data_loader,
    layer_name: str,
    device: torch.device,
    num_batches: int = 128,
) -> List[torch.Tensor]:
    """
    Collect inputs to a specific layer for calibration.

    Args:
        model: Model
        data_loader: Data loader
        layer_name: Name of layer to collect inputs for
        device: Device
        num_batches: Number of batches to collect

    Returns:
        List of input tensors
    """
    inputs = []

    # Hook to capture inputs
    def hook_fn(module, input, output):
        inputs.append(input[0].detach().cpu())

    # Register hook
    target_layer = None
    for name, module in model.named_modules():
        if name == layer_name:
            target_layer = module
            break

    if target_layer is None:
        raise ValueError(f"Layer {layer_name} not found")

    handle = target_layer.register_forward_hook(hook_fn)

    # Run forward passes
    model.eval()
    with torch.no_grad():
        for i, batch in enumerate(data_loader):
            if i >= num_batches:
                break

            if isinstance(batch, (tuple, list)):
                batch = batch[0]

            batch = batch.to(device)

            # Forward pass
            model(batch)

    # Remove hook
    handle.remove()

    return inputs


def quantize_model_gptq(
    model: nn.Module,
    data_loader,
    target_modules: List[str],
    bits: int = 4,
    group_size: int = 128,
    damp_percent: float = 0.01,
    device: torch.device = torch.device('cuda' if torch.cuda.is_available() else 'cpu'),
    num_calibration_batches: int = 128,
) -> nn.Module:
    """
    Quantize model using GPTQ.

    Args:
        model: Model to quantize
        data_loader: Calibration data loader
        target_modules: List of module names to quantize
        bits: Quantization bits
        group_size: Group size
        damp_percent: Dampening percentage
        device: Device
        num_calibration_batches: Number of calibration batches

    Returns:
        Quantized model
    """
    print(f"Quantizing model with GPTQ (bits={bits}, group_size={group_size})")

    model = model.to(device)
    model.eval()

    quantizer = GPTQuantizer(
        bits=bits,
        group_size=group_size,
        damp_percent=damp_percent,
    )

    replacements = []

    # Find target layers
    for name, module in tqdm(list(model.named_modules()), desc="Quantizing layers"):
        # Check if should quantize
        should_quantize = False
        for target in target_modules:
            if target in name and isinstance(module, nn.Linear):
                should_quantize = True
                break

        if not should_quantize:
            continue

        # Collect inputs for this layer
        print(f"Calibrating {name}...")
        inputs = collect_layer_inputs(
            model,
            data_loader,
            name,
            device,
            num_batches=num_calibration_batches
        )

        # Compute Hessian approximation (X^T X / N)
        H = torch.zeros((module.in_features, module.in_features), device=device)

        for inp in inputs:
            inp = inp.to(device)
            if inp.dim() == 3:
                inp = inp.view(-1, inp.size(-1))
            H += inp.T @ inp

        H /= len(inputs)

        # Quantize layer
        quantized_weight, scale, zero_point = quantizer.quantize_layer_gptq(module, H)

        # Create quantized layer
        quant_layer = QuantizedLinear(
            in_features=module.in_features,
            out_features=module.out_features,
            bits=bits,
            group_size=group_size,
            bias=module.bias is not None,
        )

        # Set quantized weights
        quant_layer.quantized_weight.data = quantized_weight
        quant_layer.weight_scale.data = scale
        if zero_point is not None:
            quant_layer.weight_zero_point = nn.Parameter(zero_point, requires_grad=False)

        # Copy bias
        if module.bias is not None:
            quant_layer.bias.data = module.bias.data

        # Replace module
        parent_name = '.'.join(name.split('.')[:-1])
        attr_name = name.split('.')[-1]

        if parent_name:
            parent = model.get_submodule(parent_name)
        else:
            parent = model

        setattr(parent, attr_name, quant_layer)
        replacements.append(name)

    print(f"Quantized {len(replacements)} layers with GPTQ")

    return model
