"""
Memory Optimization Utilities

Tools for estimating, profiling, and optimizing memory usage
in quantized models.
"""

import torch
import torch.nn as nn
from typing import Dict, List, Optional, Tuple
import numpy as np


def count_parameters(model: nn.Module) -> Dict[str, int]:
    """
    Count total and trainable parameters in a model.

    Args:
        model: PyTorch model

    Returns:
        Dictionary with parameter counts
    """
    total = sum(p.numel() for p in model.parameters())
    trainable = sum(p.numel() for p in model.parameters() if p.requires_grad)

    return {
        'total': total,
        'trainable': trainable,
        'frozen': total - trainable,
        'trainable_percent': 100.0 * trainable / total if total > 0 else 0.0,
    }


def estimate_model_memory(
    model: nn.Module,
    dtype: torch.dtype = torch.float32,
    include_gradients: bool = False,
    include_optimizer: bool = False,
) -> Dict[str, float]:
    """
    Estimate memory usage of a model.

    Args:
        model: PyTorch model
        dtype: Data type for parameters
        include_gradients: Include gradient memory
        include_optimizer: Include optimizer state memory (AdamW approximation)

    Returns:
        Dictionary with memory estimates in MB and GB
    """
    # Count parameters
    param_counts = count_parameters(model)
    total_params = param_counts['total']
    trainable_params = param_counts['trainable']

    # Bytes per parameter
    if dtype == torch.float32:
        bytes_per_param = 4
    elif dtype == torch.float16 or dtype == torch.bfloat16:
        bytes_per_param = 2
    elif dtype == torch.int8:
        bytes_per_param = 1
    else:
        bytes_per_param = 4  # Default

    # Model parameters memory
    model_memory = total_params * bytes_per_param

    # Gradients memory (only for trainable parameters)
    gradient_memory = trainable_params * bytes_per_param if include_gradients else 0

    # Optimizer memory (AdamW: 2 states per trainable param + gradients)
    if include_optimizer:
        optimizer_memory = trainable_params * bytes_per_param * 2 + gradient_memory
    else:
        optimizer_memory = 0

    # Total memory
    total_memory = model_memory + gradient_memory + optimizer_memory

    return {
        'model_memory_mb': model_memory / (1024**2),
        'model_memory_gb': model_memory / (1024**3),
        'gradient_memory_mb': gradient_memory / (1024**2),
        'optimizer_memory_mb': optimizer_memory / (1024**2),
        'total_memory_mb': total_memory / (1024**2),
        'total_memory_gb': total_memory / (1024**3),
    }


def get_quantization_memory_savings(
    num_parameters: int,
    method: str = "qlora",
    base_dtype: torch.dtype = torch.float16,
    lora_rank: int = 8,
) -> Dict[str, float]:
    """
    Estimate memory savings from quantization methods.

    Args:
        num_parameters: Number of parameters in base model
        method: Quantization method ("lora", "qlora", "gptq", "awq")
        base_dtype: Base model data type
        lora_rank: LoRA rank (for LoRA/QLoRA methods)

    Returns:
        Dictionary with memory statistics
    """
    # Base dtype bytes
    if base_dtype == torch.float32:
        base_bytes = 4
    elif base_dtype in [torch.float16, torch.bfloat16]:
        base_bytes = 2
    else:
        base_bytes = 4

    # Base model memory
    base_memory = num_parameters * base_bytes

    # Quantized memory depends on method
    if method == "lora":
        # LoRA: frozen fp16 weights + small LoRA adapters
        # Approximate LoRA params as ~1% of base (depends on rank and architecture)
        lora_params = num_parameters * 0.01
        quantized_memory = base_memory + lora_params * base_bytes
        trainable_memory = lora_params * base_bytes

    elif method == "qlora":
        # QLoRA: 4-bit weights + small LoRA adapters
        quant_memory = num_parameters * 0.5  # 4-bit
        lora_params = num_parameters * 0.01
        quantized_memory = quant_memory + lora_params * base_bytes
        trainable_memory = lora_params * base_bytes

    elif method in ["gptq", "awq"]:
        # GPTQ/AWQ: 4-bit weights only, no training overhead
        quantized_memory = num_parameters * 0.5  # 4-bit
        trainable_memory = 0

    else:
        raise ValueError(f"Unknown method: {method}")

    # Calculate savings
    savings = base_memory - quantized_memory
    savings_percent = (savings / base_memory) * 100

    return {
        'base_memory_gb': base_memory / (1024**3),
        'quantized_memory_gb': quantized_memory / (1024**3),
        'trainable_memory_gb': trainable_memory / (1024**3),
        'savings_gb': savings / (1024**3),
        'savings_percent': savings_percent,
        'method': method,
    }


def profile_layer_memory(layer: nn.Module) -> Dict[str, float]:
    """
    Profile memory usage of a single layer.

    Args:
        layer: Layer to profile

    Returns:
        Dictionary with layer memory statistics
    """
    total_params = sum(p.numel() for p in layer.parameters())
    trainable_params = sum(p.numel() for p in layer.parameters() if p.requires_grad)

    # Estimate memory based on parameter dtype
    memory_bytes = 0
    for p in layer.parameters():
        memory_bytes += p.numel() * p.element_size()

    # Buffer memory
    buffer_bytes = sum(b.numel() * b.element_size() for b in layer.buffers())

    total_memory = memory_bytes + buffer_bytes

    return {
        'total_params': total_params,
        'trainable_params': trainable_params,
        'memory_mb': total_memory / (1024**2),
        'param_memory_mb': memory_bytes / (1024**2),
        'buffer_memory_mb': buffer_bytes / (1024**2),
    }


def get_model_memory_breakdown(model: nn.Module) -> Dict[str, Dict[str, float]]:
    """
    Get detailed memory breakdown by layer.

    Args:
        model: Model to profile

    Returns:
        Dictionary mapping layer names to memory statistics
    """
    breakdown = {}

    for name, module in model.named_modules():
        # Only profile leaf modules with parameters
        if len(list(module.children())) == 0 and len(list(module.parameters())) > 0:
            breakdown[name] = profile_layer_memory(module)

    return breakdown


def estimate_activation_memory(
    batch_size: int,
    seq_length: int,
    hidden_size: int,
    num_layers: int,
    num_heads: int,
    dtype: torch.dtype = torch.float16,
) -> Dict[str, float]:
    """
    Estimate activation memory during forward/backward pass.

    Args:
        batch_size: Batch size
        seq_length: Sequence length
        hidden_size: Hidden size
        num_layers: Number of layers
        num_heads: Number of attention heads
        dtype: Activation data type

    Returns:
        Dictionary with activation memory estimates
    """
    # Bytes per element
    if dtype in [torch.float16, torch.bfloat16]:
        bytes_per_elem = 2
    elif dtype == torch.float32:
        bytes_per_elem = 4
    else:
        bytes_per_elem = 2

    # Attention activations per layer
    # Q, K, V: each [batch, num_heads, seq_len, head_dim]
    head_dim = hidden_size // num_heads
    qkv_elements = 3 * batch_size * num_heads * seq_length * head_dim

    # Attention scores: [batch, num_heads, seq_len, seq_len]
    attn_elements = batch_size * num_heads * seq_length * seq_length

    # FFN activations (approximate as 4x hidden size)
    ffn_elements = batch_size * seq_length * hidden_size * 4

    # Total per layer
    layer_elements = qkv_elements + attn_elements + ffn_elements

    # Total for all layers
    total_elements = layer_elements * num_layers

    # Memory
    forward_memory = total_elements * bytes_per_elem

    # Backward pass needs to store gradients (roughly 2x forward)
    backward_memory = forward_memory * 2

    return {
        'forward_memory_mb': forward_memory / (1024**2),
        'forward_memory_gb': forward_memory / (1024**3),
        'backward_memory_mb': backward_memory / (1024**2),
        'backward_memory_gb': backward_memory / (1024**3),
        'total_elements': total_elements,
    }


def print_memory_summary(
    model: nn.Module,
    batch_size: int = 1,
    seq_length: int = 512,
    dtype: torch.dtype = torch.float16,
    include_training: bool = False,
):
    """
    Print comprehensive memory summary for a model.

    Args:
        model: Model to analyze
        batch_size: Batch size for activation estimates
        seq_length: Sequence length
        dtype: Data type
        include_training: Include training memory (gradients + optimizer)
    """
    print("\n" + "="*60)
    print("Memory Analysis Summary")
    print("="*60)

    # Parameter counts
    param_counts = count_parameters(model)
    print(f"\nParameter Counts:")
    print(f"  Total parameters: {param_counts['total']:,}")
    print(f"  Trainable parameters: {param_counts['trainable']:,}")
    print(f"  Frozen parameters: {param_counts['frozen']:,}")
    print(f"  Trainable percentage: {param_counts['trainable_percent']:.2f}%")

    # Model memory
    model_mem = estimate_model_memory(
        model,
        dtype=dtype,
        include_gradients=include_training,
        include_optimizer=include_training
    )
    print(f"\nModel Memory:")
    print(f"  Parameters: {model_mem['model_memory_gb']:.2f} GB")
    if include_training:
        print(f"  Gradients: {model_mem['gradient_memory_mb']:.2f} MB")
        print(f"  Optimizer: {model_mem['optimizer_memory_mb']:.2f} MB")
        print(f"  Total (training): {model_mem['total_memory_gb']:.2f} GB")

    # Try to get model config for activation estimates
    try:
        if hasattr(model, 'embedding_dim'):
            hidden_size = model.embedding_dim
        elif hasattr(model, 'config'):
            hidden_size = model.config.hidden_size
        else:
            hidden_size = 768  # Default

        if hasattr(model, 'num_layers'):
            num_layers = model.num_layers
        elif hasattr(model, 'config'):
            num_layers = model.config.num_layers
        else:
            num_layers = 12  # Default

        if hasattr(model, 'num_heads'):
            num_heads = model.num_heads
        elif hasattr(model, 'config'):
            num_heads = model.config.num_heads
        else:
            num_heads = 12  # Default

        # Activation memory
        act_mem = estimate_activation_memory(
            batch_size=batch_size,
            seq_length=seq_length,
            hidden_size=hidden_size,
            num_layers=num_layers,
            num_heads=num_heads,
            dtype=dtype
        )

        print(f"\nActivation Memory (batch={batch_size}, seq_len={seq_length}):")
        print(f"  Forward pass: {act_mem['forward_memory_gb']:.2f} GB")
        if include_training:
            print(f"  Backward pass: {act_mem['backward_memory_gb']:.2f} GB")

        # Total memory estimate
        total_mem = model_mem['model_memory_gb'] + act_mem['forward_memory_gb']
        if include_training:
            total_mem = model_mem['total_memory_gb'] + act_mem['backward_memory_gb']

        print(f"\nEstimated Total Memory: {total_mem:.2f} GB")

    except:
        print("\nNote: Could not estimate activation memory (model config unavailable)")

    print("="*60 + "\n")


def compare_quantization_methods(num_parameters: int):
    """
    Compare memory usage across different quantization methods.

    Args:
        num_parameters: Number of parameters in base model
    """
    print("\n" + "="*60)
    print(f"Quantization Methods Comparison ({num_parameters/1e9:.2f}B parameters)")
    print("="*60)

    methods = ["lora", "qlora", "gptq", "awq"]

    print(f"\n{'Method':<10} {'Base (GB)':<12} {'Quant (GB)':<12} {'Savings':<12} {'Trainable (GB)':<15}")
    print("-"*60)

    for method in methods:
        stats = get_quantization_memory_savings(
            num_parameters=num_parameters,
            method=method,
            base_dtype=torch.float16
        )

        print(f"{method.upper():<10} "
              f"{stats['base_memory_gb']:<12.2f} "
              f"{stats['quantized_memory_gb']:<12.2f} "
              f"{stats['savings_percent']:<11.1f}% "
              f"{stats['trainable_memory_gb']:<15.2f}")

    print("="*60 + "\n")
