"""
LoRA and DoRA Implementation

Low-Rank Adaptation (LoRA) for parameter-efficient fine-tuning.
DoRA (Weight-Decomposed LoRA) for improved performance.

LoRA paper: https://arxiv.org/abs/2106.09685
DoRA paper: https://arxiv.org/abs/2402.09353
"""

import torch
import torch.nn as nn
import torch.nn.functional as F
from typing import Optional, List, Dict, Set
import math


class LoRALayer(nn.Module):
    """
    LoRA layer for any linear transformation.

    Adds low-rank matrices A and B to approximate weight updates:
    W' = W + (B @ A) * (alpha / r)

    Memory savings: Instead of updating n×m matrix, we update
    n×r and r×m matrices where r << min(n,m)
    """

    def __init__(
        self,
        in_features: int,
        out_features: int,
        r: int = 8,
        lora_alpha: float = 16.0,
        lora_dropout: float = 0.1,
        merge_weights: bool = False,
        fan_in_fan_out: bool = False,
    ):
        super().__init__()

        self.r = r
        self.lora_alpha = lora_alpha
        self.lora_dropout_rate = lora_dropout
        self.merge_weights = merge_weights
        self.merged = False
        self.fan_in_fan_out = fan_in_fan_out

        self.in_features = in_features
        self.out_features = out_features

        # Base linear layer (frozen during training)
        self.linear = nn.Linear(in_features, out_features, bias=True)

        # LoRA matrices
        if r > 0:
            self.lora_A = nn.Parameter(torch.zeros(r, in_features))
            self.lora_B = nn.Parameter(torch.zeros(out_features, r))
            self.scaling = lora_alpha / r

            # Dropout
            if lora_dropout > 0.0:
                self.lora_dropout = nn.Dropout(p=lora_dropout)
            else:
                self.lora_dropout = nn.Identity()

            # Initialize
            self.reset_lora_parameters()

        # Freeze base weights
        self.linear.weight.requires_grad = False
        if self.linear.bias is not None:
            self.linear.bias.requires_grad = False

    def reset_lora_parameters(self):
        """Initialize LoRA parameters using Kaiming uniform for A, zeros for B."""
        if hasattr(self, 'lora_A'):
            # Initialize A with Kaiming uniform
            nn.init.kaiming_uniform_(self.lora_A, a=math.sqrt(5))
            # Initialize B with zeros (so initially LoRA has no effect)
            nn.init.zeros_(self.lora_B)

    def train(self, mode: bool = True):
        """Override train to handle weight merging."""
        super().train(mode)

        if mode:
            # Training mode: unmerge weights if they were merged
            if self.merge_weights and self.merged:
                if self.r > 0:
                    # Subtract LoRA weights from base weights
                    self.linear.weight.data -= (self.lora_B @ self.lora_A) * self.scaling
                self.merged = False
        else:
            # Eval mode: merge weights if requested
            if self.merge_weights and not self.merged:
                if self.r > 0:
                    # Add LoRA weights to base weights
                    self.linear.weight.data += (self.lora_B @ self.lora_A) * self.scaling
                self.merged = True

    def forward(self, x: torch.Tensor) -> torch.Tensor:
        """
        Forward pass with LoRA.

        Args:
            x: [batch, seq_len, in_features] or [batch, in_features]

        Returns:
            output: [batch, seq_len, out_features] or [batch, out_features]
        """
        # Base linear transformation
        result = self.linear(x)

        # Add LoRA transformation if not merged
        if self.r > 0 and not self.merged:
            # x @ A^T -> [batch, ..., r]
            lora_x = x @ self.lora_A.T
            lora_x = self.lora_dropout(lora_x)
            # lora_x @ B^T -> [batch, ..., out_features]
            lora_result = lora_x @ self.lora_B.T
            result = result + lora_result * self.scaling

        return result

    def extra_repr(self) -> str:
        """Extra representation for debugging."""
        return f'in_features={self.in_features}, out_features={self.out_features}, r={self.r}, alpha={self.lora_alpha}'


class DoRALayer(LoRALayer):
    """
    DoRA (Weight-Decomposed LoRA) layer.

    Decomposes weight into magnitude and direction:
    W' = m * (W + B @ A) / ||W + B @ A||

    This provides better performance than standard LoRA by
    separating magnitude and directional updates.
    """

    def __init__(
        self,
        in_features: int,
        out_features: int,
        r: int = 8,
        lora_alpha: float = 16.0,
        lora_dropout: float = 0.1,
        merge_weights: bool = False,
        fan_in_fan_out: bool = False,
        magnitude_init: str = "uniform",
    ):
        super().__init__(
            in_features=in_features,
            out_features=out_features,
            r=r,
            lora_alpha=lora_alpha,
            lora_dropout=lora_dropout,
            merge_weights=merge_weights,
            fan_in_fan_out=fan_in_fan_out,
        )

        # Magnitude vector (learnable)
        self.magnitude = nn.Parameter(torch.ones(out_features))

        # Initialize magnitude
        if magnitude_init == "uniform":
            nn.init.uniform_(self.magnitude, 0.9, 1.1)
        elif magnitude_init == "normal":
            nn.init.normal_(self.magnitude, mean=1.0, std=0.1)
        # else: keep as ones (constant init)

    def forward(self, x: torch.Tensor) -> torch.Tensor:
        """
        Forward pass with DoRA.

        Args:
            x: [batch, seq_len, in_features] or [batch, in_features]

        Returns:
            output: [batch, seq_len, out_features] or [batch, out_features]
        """
        # Get base weight
        weight = self.linear.weight

        # Compute directional component: W + B @ A
        if self.r > 0:
            # LoRA update
            lora_weight = (self.lora_B @ self.lora_A) * self.scaling
            combined_weight = weight + lora_weight
        else:
            combined_weight = weight

        # Normalize to get direction
        # Column-wise normalization: each output feature gets unit direction
        weight_norm = torch.norm(combined_weight, p=2, dim=1, keepdim=True) + 1e-8
        directional_weight = combined_weight / weight_norm

        # Apply magnitude
        dora_weight = self.magnitude.unsqueeze(1) * directional_weight

        # Apply transformation
        result = F.linear(x, dora_weight, self.linear.bias)

        return result

    def extra_repr(self) -> str:
        """Extra representation for debugging."""
        return f'in_features={self.in_features}, out_features={self.out_features}, r={self.r}, alpha={self.lora_alpha}, DoRA=True'


def apply_lora_to_model(
    model: nn.Module,
    target_modules: List[str],
    r: int = 8,
    lora_alpha: float = 16.0,
    lora_dropout: float = 0.1,
    merge_weights: bool = False,
    use_dora: bool = False,
) -> nn.Module:
    """
    Apply LoRA or DoRA to target modules in a model.

    Args:
        model: Model to apply LoRA to
        target_modules: List of module names to replace (e.g., ["query_proj", "value_proj"])
        r: LoRA rank
        lora_alpha: LoRA alpha scaling
        lora_dropout: Dropout rate
        merge_weights: Whether to merge weights during eval
        use_dora: Whether to use DoRA instead of LoRA

    Returns:
        Modified model with LoRA/DoRA layers
    """
    # Track replacements
    replacements = []

    # Find and replace target modules
    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 and attribute name
            parent_name = '.'.join(name.split('.')[:-1])
            attr_name = name.split('.')[-1]

            if parent_name:
                parent = model.get_submodule(parent_name)
            else:
                parent = model

            # Create LoRA/DoRA layer
            if use_dora:
                lora_layer = DoRALayer(
                    in_features=module.in_features,
                    out_features=module.out_features,
                    r=r,
                    lora_alpha=lora_alpha,
                    lora_dropout=lora_dropout,
                    merge_weights=merge_weights,
                )
            else:
                lora_layer = LoRALayer(
                    in_features=module.in_features,
                    out_features=module.out_features,
                    r=r,
                    lora_alpha=lora_alpha,
                    lora_dropout=lora_dropout,
                    merge_weights=merge_weights,
                )

            # Copy base weights
            lora_layer.linear.weight.data = module.weight.data.clone()
            if module.bias is not None:
                lora_layer.linear.bias.data = module.bias.data.clone()

            # Replace module
            setattr(parent, attr_name, lora_layer)
            replacements.append(name)

    print(f"Applied {'DoRA' if use_dora else 'LoRA'} to {len(replacements)} modules:")
    for name in replacements:
        print(f"  - {name}")

    return model


def count_lora_parameters(model: nn.Module) -> Dict[str, int]:
    """
    Count trainable and total parameters in a LoRA model.

    Returns:
        Dictionary with parameter counts
    """
    total_params = sum(p.numel() for p in model.parameters())
    trainable_params = sum(p.numel() for p in model.parameters() if p.requires_grad)
    lora_params = sum(
        p.numel() for n, p in model.named_parameters()
        if 'lora_' in n or 'magnitude' in n
    )

    return {
        'total_parameters': total_params,
        'trainable_parameters': trainable_params,
        'lora_parameters': lora_params,
        'trainable_percentage': 100.0 * trainable_params / total_params if total_params > 0 else 0.0,
    }


def merge_lora_weights(model: nn.Module):
    """Merge LoRA weights into base model weights."""
    for module in model.modules():
        if isinstance(module, (LoRALayer, DoRALayer)):
            if not module.merged:
                if module.r > 0:
                    module.linear.weight.data += (module.lora_B @ module.lora_A) * module.scaling
                module.merged = True


def unmerge_lora_weights(model: nn.Module):
    """Unmerge LoRA weights from base model weights."""
    for module in model.modules():
        if isinstance(module, (LoRALayer, DoRALayer)):
            if module.merged:
                if module.r > 0:
                    module.linear.weight.data -= (module.lora_B @ module.lora_A) * module.scaling
                module.merged = False


def save_lora_weights(model: nn.Module, path: str):
    """
    Save only LoRA parameters (not base model weights).

    This results in much smaller checkpoint files.
    """
    lora_state_dict = {}

    for name, param in model.named_parameters():
        if 'lora_' in name or 'magnitude' in name:
            lora_state_dict[name] = param.data

    torch.save(lora_state_dict, path)
    print(f"Saved LoRA weights to {path}")
    print(f"LoRA parameters: {sum(p.numel() for p in lora_state_dict.values()):,}")


def load_lora_weights(model: nn.Module, path: str, strict: bool = False):
    """Load LoRA parameters from checkpoint."""
    lora_state_dict = torch.load(path, map_location='cpu')

    # Load only LoRA parameters
    model.load_state_dict(lora_state_dict, strict=strict)
    print(f"Loaded LoRA weights from {path}")
