"""
KV Cache Quantization

Quantize key-value cache during inference to reduce memory usage.
Supports 8-bit and 4-bit quantization with per-channel or per-token scaling.

Critical for long-context generation where KV cache can dominate memory.
"""

import torch
import torch.nn as nn
from typing import Optional, Tuple, Dict
import math


class KVCacheQuantizer:
    """
    Quantizer for KV cache tensors.

    Supports both 8-bit (int8) and 4-bit quantization with
    per-channel or per-token scaling for accuracy.
    """

    def __init__(
        self,
        bits: int = 8,
        per_channel: bool = True,
        symmetric: bool = True,
        dynamic: bool = False,
    ):
        """
        Args:
            bits: Quantization bits (4 or 8)
            per_channel: Per-channel quantization (vs per-tensor)
            symmetric: Symmetric quantization
            dynamic: Dynamic quantization (recalculate scales each time)
        """
        self.bits = bits
        self.per_channel = per_channel
        self.symmetric = symmetric
        self.dynamic = dynamic

        if bits == 8:
            self.maxq = 127 if symmetric else 255
            self.dtype = torch.int8 if symmetric else torch.uint8
        elif bits == 4:
            self.maxq = 7 if symmetric else 15
            self.dtype = torch.int8  # Store 4-bit in int8
        else:
            raise ValueError(f"Unsupported bits: {bits}. Use 4 or 8.")

    def quantize(
        self,
        tensor: torch.Tensor,
        scale: Optional[torch.Tensor] = None,
        zero_point: Optional[torch.Tensor] = None,
    ) -> Tuple[torch.Tensor, torch.Tensor, Optional[torch.Tensor]]:
        """
        Quantize a tensor.

        Args:
            tensor: Tensor to quantize [batch, num_heads, seq_len, head_dim]
            scale: Pre-computed scale (if None, compute from tensor)
            zero_point: Pre-computed zero point (if None, compute from tensor)

        Returns:
            quantized: Quantized tensor
            scale: Scale factors
            zero_point: Zero points (None if symmetric)
        """
        # Determine quantization dimensions
        if self.per_channel:
            # Per-channel: quantize along head_dim
            dims = list(range(tensor.ndim - 1))  # All dims except last
        else:
            # Per-tensor: quantize entire tensor
            dims = None

        # Compute scale and zero point if not provided
        if scale is None or self.dynamic:
            if self.symmetric:
                # Symmetric: scale based on max absolute value
                if dims is not None:
                    max_abs = tensor.abs().amax(dim=dims, keepdim=True)
                else:
                    max_abs = tensor.abs().max()

                scale = max_abs / self.maxq
                scale = scale.clamp(min=1e-8)
                zero_point = None

            else:
                # Asymmetric: scale based on min/max
                if dims is not None:
                    t_min = tensor.amin(dim=dims, keepdim=True)
                    t_max = tensor.amax(dim=dims, keepdim=True)
                else:
                    t_min = tensor.min()
                    t_max = tensor.max()

                scale = (t_max - t_min) / self.maxq
                scale = scale.clamp(min=1e-8)
                zero_point = torch.round(-t_min / scale).clamp(0, self.maxq)

        # Quantize
        if self.symmetric:
            quantized = torch.clamp(
                torch.round(tensor / scale),
                -self.maxq,
                self.maxq
            ).to(self.dtype)
        else:
            quantized = torch.clamp(
                torch.round(tensor / scale + zero_point),
                0,
                self.maxq
            ).to(self.dtype)

        return quantized, scale, zero_point

    def dequantize(
        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 if symmetric)

        Returns:
            Dequantized tensor
        """
        quantized = quantized.float()

        if self.symmetric:
            dequantized = quantized * scale
        else:
            dequantized = (quantized - zero_point) * scale

        return dequantized


class QuantizedKVCache:
    """
    Quantized Key-Value cache for attention mechanisms.

    Stores KV cache in quantized format to reduce memory.
    Automatically handles quantization/dequantization during updates.
    """

    def __init__(
        self,
        max_batch_size: int,
        max_seq_length: int,
        num_heads: int,
        head_dim: int,
        bits: int = 8,
        per_channel: bool = True,
        symmetric: bool = True,
        dynamic: bool = False,
        device: torch.device = torch.device('cpu'),
        dtype: torch.dtype = torch.float16,
    ):
        """
        Args:
            max_batch_size: Maximum batch size
            max_seq_length: Maximum sequence length
            num_heads: Number of attention heads
            head_dim: Dimension per head
            bits: Quantization bits (4 or 8)
            per_channel: Per-channel quantization
            symmetric: Symmetric quantization
            dynamic: Dynamic quantization
            device: Device
            dtype: Computation dtype (for dequantized values)
        """
        self.max_batch_size = max_batch_size
        self.max_seq_length = max_seq_length
        self.num_heads = num_heads
        self.head_dim = head_dim
        self.bits = bits
        self.device = device
        self.dtype = dtype

        # Quantizer
        self.quantizer = KVCacheQuantizer(
            bits=bits,
            per_channel=per_channel,
            symmetric=symmetric,
            dynamic=dynamic,
        )

        # Storage for quantized cache
        if bits == 8:
            quant_dtype = torch.int8 if symmetric else torch.uint8
        else:
            quant_dtype = torch.int8

        # Keys and Values
        self.k_cache_quant = torch.zeros(
            (max_batch_size, num_heads, max_seq_length, head_dim),
            dtype=quant_dtype,
            device=device
        )
        self.v_cache_quant = torch.zeros(
            (max_batch_size, num_heads, max_seq_length, head_dim),
            dtype=quant_dtype,
            device=device
        )

        # Scales and zero points
        if per_channel:
            scale_shape = (max_batch_size, num_heads, max_seq_length, 1)
        else:
            scale_shape = (1,)

        self.k_scale = torch.ones(scale_shape, dtype=dtype, device=device)
        self.v_scale = torch.ones(scale_shape, dtype=dtype, device=device)

        if not symmetric:
            self.k_zero_point = torch.zeros(scale_shape, dtype=dtype, device=device)
            self.v_zero_point = torch.zeros(scale_shape, dtype=dtype, device=device)
        else:
            self.k_zero_point = None
            self.v_zero_point = None

        # Current cache length
        self.cache_len = 0

    def update(
        self,
        key: torch.Tensor,
        value: torch.Tensor,
        start_pos: int = 0,
    ) -> Tuple[torch.Tensor, torch.Tensor]:
        """
        Update cache with new key-value pairs.

        Args:
            key: New keys [batch, num_heads, seq_len, head_dim]
            value: New values [batch, num_heads, seq_len, head_dim]
            start_pos: Starting position in cache

        Returns:
            Full keys and values (dequantized)
        """
        batch_size, num_heads, seq_len, head_dim = key.shape

        # Quantize new keys and values
        k_quant, k_scale, k_zero = self.quantizer.quantize(key)
        v_quant, v_scale, v_zero = self.quantizer.quantize(value)

        # Store in cache
        end_pos = start_pos + seq_len

        self.k_cache_quant[:batch_size, :, start_pos:end_pos, :] = k_quant
        self.v_cache_quant[:batch_size, :, start_pos:end_pos, :] = v_quant

        self.k_scale[:batch_size, :, start_pos:end_pos, :] = k_scale
        self.v_scale[:batch_size, :, start_pos:end_pos, :] = v_scale

        if k_zero is not None:
            self.k_zero_point[:batch_size, :, start_pos:end_pos, :] = k_zero
            self.v_zero_point[:batch_size, :, start_pos:end_pos, :] = v_zero

        # Update cache length
        self.cache_len = max(self.cache_len, end_pos)

        # Return full cache (dequantized)
        full_keys = self.get_keys(batch_size, end_pos)
        full_values = self.get_values(batch_size, end_pos)

        return full_keys, full_values

    def get_keys(self, batch_size: int, seq_len: int) -> torch.Tensor:
        """
        Get dequantized keys from cache.

        Args:
            batch_size: Batch size
            seq_len: Sequence length

        Returns:
            Keys [batch, num_heads, seq_len, head_dim]
        """
        k_quant = self.k_cache_quant[:batch_size, :, :seq_len, :]
        k_scale = self.k_scale[:batch_size, :, :seq_len, :]
        k_zero = self.k_zero_point[:batch_size, :, :seq_len, :] if self.k_zero_point is not None else None

        keys = self.quantizer.dequantize(k_quant, k_scale, k_zero)

        return keys.to(self.dtype)

    def get_values(self, batch_size: int, seq_len: int) -> torch.Tensor:
        """
        Get dequantized values from cache.

        Args:
            batch_size: Batch size
            seq_len: Sequence length

        Returns:
            Values [batch, num_heads, seq_len, head_dim]
        """
        v_quant = self.v_cache_quant[:batch_size, :, :seq_len, :]
        v_scale = self.v_scale[:batch_size, :, :seq_len, :]
        v_zero = self.v_zero_point[:batch_size, :, :seq_len, :] if self.v_zero_point is not None else None

        values = self.quantizer.dequantize(v_quant, v_scale, v_zero)

        return values.to(self.dtype)

    def reset(self):
        """Reset cache to empty state."""
        self.cache_len = 0

    def get_memory_usage(self) -> Dict[str, float]:
        """
        Get memory usage statistics.

        Returns:
            Dictionary with memory usage in MB
        """
        # Quantized cache memory
        quant_bytes_per_elem = 1 if self.bits == 8 else 0.5
        cache_memory = (
            self.k_cache_quant.numel() * quant_bytes_per_elem +
            self.v_cache_quant.numel() * quant_bytes_per_elem
        )

        # Scale and zero point memory
        scale_memory = (
            self.k_scale.numel() * self.k_scale.element_size() +
            self.v_scale.numel() * self.v_scale.element_size()
        )

        if self.k_zero_point is not None:
            scale_memory += (
                self.k_zero_point.numel() * self.k_zero_point.element_size() +
                self.v_zero_point.numel() * self.v_zero_point.element_size()
            )

        total_memory = cache_memory + scale_memory

        # Equivalent fp16 memory (for comparison)
        fp16_memory = (self.k_cache_quant.numel() + self.v_cache_quant.numel()) * 2

        return {
            'cache_memory_mb': cache_memory / (1024**2),
            'scale_memory_mb': scale_memory / (1024**2),
            'total_memory_mb': total_memory / (1024**2),
            'fp16_equivalent_mb': fp16_memory / (1024**2),
            'savings_percent': (fp16_memory - total_memory) / fp16_memory * 100,
        }


def estimate_kv_cache_savings(
    batch_size: int,
    max_seq_length: int,
    num_heads: int,
    head_dim: int,
    num_layers: int,
    bits: int = 8,
) -> Dict[str, float]:
    """
    Estimate memory savings from KV cache quantization.

    Args:
        batch_size: Batch size
        max_seq_length: Maximum sequence length
        num_heads: Number of attention heads
        head_dim: Dimension per head
        num_layers: Number of layers
        bits: Quantization bits

    Returns:
        Dictionary with memory statistics in MB
    """
    # Elements per cache (K and V)
    cache_elements = batch_size * num_heads * max_seq_length * head_dim

    # FP16 memory (2 bytes per element, 2 caches per layer)
    fp16_memory = cache_elements * 2 * 2 * num_layers

    # Quantized memory
    bytes_per_quant = 1 if bits == 8 else 0.5
    quant_memory = cache_elements * bytes_per_quant * 2 * num_layers

    # Scale memory (assuming per-channel quantization)
    scale_elements = batch_size * num_heads * max_seq_length * 1
    scale_memory = scale_elements * 2 * 2 * num_layers  # 2 bytes (fp16), 2 caches

    total_quant_memory = quant_memory + scale_memory

    savings = fp16_memory - total_quant_memory
    savings_percent = (savings / fp16_memory) * 100

    return {
        'fp16_memory_mb': fp16_memory / (1024**2),
        'quantized_memory_mb': total_quant_memory / (1024**2),
        'savings_mb': savings / (1024**2),
        'savings_percent': savings_percent,
    }
