"""
AWQ (Activation-aware Weight Quantization)

Protects salient weights by analyzing activation patterns.
Achieves better accuracy than GPTQ by identifying and preserving
important channels based on activation magnitudes.

Paper: https://arxiv.org/abs/2306.00978
"""

import torch
import torch.nn as nn
from typing import Optional, List, Dict, Tuple
import numpy as np
from tqdm import tqdm


class AWQQuantizer:
    """
    AWQ quantizer with activation-aware channel protection.

    Key idea: Scale important channels before quantization to reduce
    quantization error where it matters most.
    """

    def __init__(
        self,
        bits: int = 4,
        group_size: int = 128,
        zero_point: bool = True,
        version: str = "GEMM",
        w_bit: int = 4,
    ):
        """
        Args:
            bits: Quantization bits
            group_size: Group size for quantization
            zero_point: Use zero-point quantization
            version: AWQ version ("GEMM" or "GEMV")
            w_bit: Weight quantization bits
        """
        self.bits = bits
        self.group_size = group_size
        self.zero_point = zero_point
        self.version = version
        self.w_bit = w_bit

        self.maxq = 2 ** bits - 1

    def compute_channel_scales(
        self,
        activations: List[torch.Tensor],
        weights: torch.Tensor,
        n_groups: int = 1,
    ) -> torch.Tensor:
        """
        Compute per-channel scaling factors based on activation magnitude.

        Args:
            activations: List of activation tensors
            weights: Weight matrix [out_features, in_features]
            n_groups: Number of groups for scale search

        Returns:
            scales: Per-channel scaling factors [in_features]
        """
        in_features = weights.shape[1]

        # Compute activation statistics
        # Average L2 norm of activations per channel
        channel_importance = torch.zeros(in_features, device=weights.device)

        for act in activations:
            if act.dim() == 3:
                # [batch, seq_len, in_features] -> [batch * seq_len, in_features]
                act = act.view(-1, in_features)

            # L2 norm per channel
            channel_importance += torch.norm(act, p=2, dim=0)

        channel_importance /= len(activations)

        # Search for optimal scales
        # Try different scale values and pick the one that minimizes quantization error
        scale_candidates = torch.linspace(0.5, 2.0, steps=20, device=weights.device)

        best_scales = torch.ones(in_features, device=weights.device)
        best_error = float('inf')

        for scale in scale_candidates:
            # Scale weights by channel importance
            scaled_weights = weights * (channel_importance * scale).unsqueeze(0)

            # Quantize
            quant_w, scale_factors, zero_points = self.quantize_tensor(scaled_weights)

            # Dequantize
            dequant_w = self.dequantize_tensor(quant_w, scale_factors, zero_points)

            # Unscale
            dequant_w = dequant_w / (channel_importance * scale).unsqueeze(0)

            # Compute error
            error = torch.mean((weights - dequant_w) ** 2).item()

            if error < best_error:
                best_error = error
                best_scales = channel_importance * scale

        return best_scales

    def quantize_tensor(
        self,
        tensor: torch.Tensor,
    ) -> Tuple[torch.Tensor, torch.Tensor, Optional[torch.Tensor]]:
        """
        Quantize a tensor with groupwise quantization.

        Args:
            tensor: Tensor to quantize [out_features, in_features]

        Returns:
            quantized: Quantized tensor
            scales: Scale factors per group
            zero_points: Zero points per group (None if symmetric)
        """
        out_features, in_features = tensor.shape

        # Determine number of groups
        if self.group_size == -1:
            num_groups = 1
            group_size = in_features
        else:
            group_size = self.group_size
            num_groups = (in_features + group_size - 1) // group_size

        scales = []
        zero_points = [] if self.zero_point else None
        quantized_groups = []

        for i in range(num_groups):
            start_idx = i * group_size
            end_idx = min((i + 1) * group_size, in_features)

            group = tensor[:, start_idx:end_idx]

            # Find min/max
            w_min = group.min(dim=1, keepdim=True)[0]
            w_max = group.max(dim=1, keepdim=True)[0]

            if self.zero_point:
                # Asymmetric quantization
                scale = (w_max - w_min) / self.maxq
                scale = scale.clamp(min=1e-8)
                zero = torch.round(-w_min / scale).clamp(0, self.maxq)

                # Quantize
                quant_group = torch.clamp(
                    torch.round(group / scale + zero),
                    0,
                    self.maxq
                ).to(torch.uint8)

                scales.append(scale)
                zero_points.append(zero)
            else:
                # Symmetric quantization
                max_abs = torch.max(w_min.abs(), w_max.abs())
                scale = max_abs / (self.maxq / 2)
                scale = scale.clamp(min=1e-8)

                # Quantize
                quant_group = torch.clamp(
                    torch.round(group / scale),
                    -self.maxq // 2,
                    self.maxq // 2
                ).to(torch.int8)

                scales.append(scale)

            quantized_groups.append(quant_group)

        # Concatenate
        quantized = torch.cat(quantized_groups, dim=1)
        scale_tensor = torch.cat(scales, dim=1)

        if self.zero_point:
            zero_tensor = torch.cat(zero_points, dim=1)
        else:
            zero_tensor = None

        return quantized, scale_tensor, zero_tensor

    def dequantize_tensor(
        self,
        quantized: torch.Tensor,
        scales: torch.Tensor,
        zero_points: Optional[torch.Tensor] = None,
    ) -> torch.Tensor:
        """
        Dequantize a tensor.

        Args:
            quantized: Quantized tensor
            scales: Scale factors
            zero_points: Zero points (None for symmetric)

        Returns:
            Dequantized tensor
        """
        out_features, in_features = quantized.shape

        if self.group_size == -1:
            group_size = in_features
        else:
            group_size = self.group_size

        num_groups = (in_features + group_size - 1) // group_size

        dequantized_groups = []

        for i in range(num_groups):
            start_idx = i * group_size
            end_idx = min((i + 1) * group_size, in_features)

            quant_group = quantized[:, start_idx:end_idx].float()
            scale = scales[:, i:i+1]

            if zero_points is not None:
                zero = zero_points[:, i:i+1]
                dequant_group = (quant_group - zero) * scale
            else:
                dequant_group = quant_group * scale

            dequantized_groups.append(dequant_group)

        dequantized = torch.cat(dequantized_groups, dim=1)

        return dequantized

    def quantize_layer_awq(
        self,
        layer: nn.Linear,
        activations: List[torch.Tensor],
    ) -> Tuple[torch.Tensor, torch.Tensor, torch.Tensor, Optional[torch.Tensor]]:
        """
        Quantize a linear layer using AWQ.

        Args:
            layer: Linear layer to quantize
            activations: List of input activations for calibration

        Returns:
            quantized_weight: Quantized weight
            scales: Quantization scales
            channel_scales: Per-channel scaling factors
            zero_points: Zero points (None if symmetric)
        """
        W = layer.weight.data.clone()

        # Compute activation-aware channel scales
        channel_scales = self.compute_channel_scales(activations, W)

        # Apply channel scaling
        W_scaled = W * channel_scales.unsqueeze(0)

        # Quantize scaled weights
        quantized_weight, scales, zero_points = self.quantize_tensor(W_scaled)

        return quantized_weight, scales, channel_scales, zero_points


class AWQLinear(nn.Module):
    """
    AWQ-quantized linear layer.

    Stores weights in quantized format with per-channel scaling.
    """

    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
        num_groups = (in_features + group_size - 1) // group_size

        self.register_buffer('quantized_weight', torch.zeros((out_features, in_features), dtype=torch.uint8))
        self.register_buffer('weight_scales', torch.zeros((out_features, num_groups)))
        self.register_buffer('channel_scales', torch.ones(in_features))
        self.register_buffer('weight_zero_points', 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 = AWQQuantizer(bits=self.bits, group_size=self.group_size)
        W_scaled = quantizer.dequantize_tensor(
            self.quantized_weight,
            self.weight_scales,
            self.weight_zero_points
        )

        # Unscale weights
        weight = W_scaled / self.channel_scales.unsqueeze(0)

        # Linear transformation
        output = torch.nn.functional.linear(x, weight, self.bias)

        return output


def collect_layer_activations(
    model: nn.Module,
    data_loader,
    layer_name: str,
    device: torch.device,
    num_batches: int = 128,
) -> List[torch.Tensor]:
    """
    Collect activations (inputs) to a specific layer.

    Args:
        model: Model
        data_loader: Data loader
        layer_name: Name of layer
        device: Device
        num_batches: Number of batches to collect

    Returns:
        List of activation tensors
    """
    activations = []

    # Hook to capture activations
    def hook_fn(module, input, output):
        activations.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
            try:
                model(batch)
            except:
                # Handle different model interfaces
                if hasattr(model, 'forward'):
                    model.forward(batch)

    # Remove hook
    handle.remove()

    return activations


def quantize_model_awq(
    model: nn.Module,
    data_loader,
    target_modules: List[str],
    bits: int = 4,
    group_size: int = 128,
    device: torch.device = torch.device('cuda' if torch.cuda.is_available() else 'cpu'),
    num_calibration_batches: int = 128,
) -> nn.Module:
    """
    Quantize model using AWQ.

    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
        device: Device
        num_calibration_batches: Number of calibration batches

    Returns:
        Quantized model
    """
    print(f"Quantizing model with AWQ (bits={bits}, group_size={group_size})")

    model = model.to(device)
    model.eval()

    quantizer = AWQQuantizer(bits=bits, group_size=group_size)

    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 activations for this layer
        print(f"Calibrating {name}...")
        activations = collect_layer_activations(
            model,
            data_loader,
            name,
            device,
            num_batches=num_calibration_batches
        )

        if len(activations) == 0:
            print(f"Warning: No activations collected for {name}, skipping...")
            continue

        # Quantize layer with AWQ
        quantized_weight, scales, channel_scales, zero_points = quantizer.quantize_layer_awq(
            module,
            activations
        )

        # Create AWQ layer
        awq_layer = AWQLinear(
            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 and scales
        awq_layer.quantized_weight.data = quantized_weight
        awq_layer.weight_scales.data = scales
        awq_layer.channel_scales.data = channel_scales
        if zero_points is not None:
            awq_layer.weight_zero_points = nn.Parameter(zero_points, requires_grad=False)

        # Copy bias
        if module.bias is not None:
            awq_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, awq_layer)
        replacements.append(name)

    print(f"Quantized {len(replacements)} layers with AWQ")

    return model
