import math
import torch
from typing import Optional, Tuple

def build_sinusoidal_pos_emb(block_size: int, n_embd: int) -> torch.Tensor:
    position = torch.arange(block_size, dtype=torch.float32).unsqueeze(1)
    div_term = torch.exp(torch.arange(0, n_embd, 2, dtype=torch.float32) * (-math.log(10000.0) / n_embd))
    pos_emb = torch.zeros(block_size, n_embd, dtype=torch.float32)
    pos_emb[:, 0::2] = torch.sin(position * div_term)
    pos_emb[:, 1::2] = torch.cos(position * div_term)
    return pos_emb


class AlibiPositionalBias(torch.nn.Module):
    def __init__(self, alibi_heads: int, total_heads: int):
        super().__init__()
        self.alibi_heads = alibi_heads
        self.total_heads = total_heads
        if not (1 <= self.alibi_heads <= self.total_heads):
            raise ValueError(
                f"expected 1 <= alibi_heads <= total_heads, got alibi_heads={self.alibi_heads}, total_heads={self.total_heads}"
            )

        slopes = torch.tensor(self._get_slopes(self.alibi_heads), dtype=torch.float32)
        self.register_buffer("slopes", slopes.view(-1, 1, 1), persistent=False)
        self.register_buffer("bias", torch.empty(0), persistent=False)

    @staticmethod
    def _get_slopes_power_of_2(n: int):
        start = 2 ** (-(2 ** -(math.log2(n) - 3)))
        ratio = start
        return [start * (ratio ** i) for i in range(n)]

    @classmethod
    def _get_slopes(cls, n: int):
        if math.log2(n).is_integer():
            return cls._get_slopes_power_of_2(n)

        closest_power_of_2 = 2 ** math.floor(math.log2(n))
        return cls._get_slopes_power_of_2(closest_power_of_2) + cls._get_slopes_power_of_2(2 * closest_power_of_2)[0::2][:n - closest_power_of_2]

    def forward(self, i: int, j: Optional[int] = None, device: Optional[torch.device] = None) -> torch.Tensor:
        j = i if j is None else j
        device = self.slopes.device if device is None else device

        if self.bias.numel() > 0 and self.bias.shape[-2] >= i and self.bias.shape[-1] >= j and self.bias.device == device:
            return self.bias[..., -i:, -j:]

        q_pos = torch.arange(j - i, j, dtype=torch.long, device=device)
        k_pos = torch.arange(j, dtype=torch.long, device=device)
        bias = -(k_pos[None, :] - q_pos[:, None]).abs().to(torch.float32)
        bias = bias.unsqueeze(0) * self.slopes
        bias = bias.unsqueeze(0)

        num_heads_unalibied = self.total_heads - self.alibi_heads
        if num_heads_unalibied > 0:
            pad = torch.zeros(1, num_heads_unalibied, i, j, device=device, dtype=bias.dtype)
            bias = torch.cat((bias, pad), dim=1)

        self.bias = bias
        return self.bias

# borrowed from: https://github.com/RobertCsordas/moeut/
class RotaryPosEncoding(torch.nn.Module):
    # RoPE based on: https://www.kaggle.com/code/aeryss/rotary-postional-encoding-rope-pytorch
    def __init__(self, n_embd: int, base=10000, seq_dim: int = -2, pretrained_seq_len: int = 1024, extended_seq_len: int = 1024, 
    rotate_fraction: float = 1.0, beta_fast: int = 32, beta_slow: int = 1):
        super().__init__()
        self.seq_len_cached = 0
        self.yarn_factor_cached = None
        self.cos_cached = None
        self.sin_cached = None
        self.seq_dim = seq_dim
        self.n_embd = n_embd
        self.pretrained_seq_len = pretrained_seq_len
        self.extended_seq_len = extended_seq_len
        self.beta_fast = beta_fast
        self.beta_slow = beta_slow
        self.base = base
        if rotate_fraction < 1:
            self.n_rotate = int(n_embd * rotate_fraction)
            self.n_rotate -= self.n_rotate % 2
        else:
            self.n_rotate = n_embd
        inv_freq = 1.0 / (base ** (torch.arange(0, self.n_rotate, 2).float() / self.n_rotate))
        self.register_buffer("inv_freq_base", inv_freq, persistent=False)

    def rotate_half(self, x: torch.Tensor) -> torch.Tensor:
        x1, x2 = x[..., : x.shape[-1] // 2], x[..., x.shape[-1] // 2 :]
        return torch.cat((-x2, x1), dim=x1.ndim - 1)  # dim=-1 triggers a bug in torch < 1.8.0

    def apply_rot(self, x: torch.Tensor, sin: torch.Tensor, cos: torch.Tensor, seq_dim: int, offset: int) -> torch.Tensor:
        sin = sin.narrow(seq_dim, offset, x.shape[seq_dim])
        cos = cos.narrow(seq_dim, offset, x.shape[seq_dim])
        return (x * cos) + (self.rotate_half(x) * sin)

    def apply_rotary_pos_emb(self, q: torch.Tensor, k: torch.Tensor, sin: torch.Tensor, cos: torch.Tensor,
                             seq_dim: int, offset: int) -> Tuple[torch.Tensor, torch.Tensor]:
        return self.apply_rot(q, sin, cos, seq_dim, offset), self.apply_rot(k, sin, cos, seq_dim, 0)

    def get(self, x: torch.Tensor, yarn_factor: float) -> Tuple[torch.Tensor, torch.Tensor]:
        seq_len = x.shape[self.seq_dim]
        if self.sin_cached is None or self.cos_cached is None or seq_len > self.seq_len_cached or yarn_factor != self.yarn_factor_cached:
            self.seq_len_cached = seq_len
            self.yarn_factor_cached = yarn_factor

            inv_freq = self.inv_freq_base.to(device=x.device)
            if yarn_factor > 1.0:
                low, high = find_correction_range(
                    self.beta_fast, self.beta_slow, self.n_rotate, self.base, self.pretrained_seq_len
                )
                smooth = 1 - linear_ramp_factor(low, high, self.n_rotate // 2).to(x.device)
                inv_freq = inv_freq / yarn_factor * (1 - smooth) + inv_freq * smooth

            t = torch.arange(x.shape[self.seq_dim], device=x.device).type_as(inv_freq)
            freqs = torch.einsum("i,j->ij", t, inv_freq)
            emb = torch.cat((freqs, freqs), dim=-1).to(x.device)

            tgt_shape = [1] * x.ndim
            tgt_shape[self.seq_dim] = seq_len
            tgt_shape[-1] = x.shape[-1]

            self.cos_cached = emb.cos().view(*tgt_shape)
            self.sin_cached = emb.sin().view(*tgt_shape)

        return self.sin_cached, self.cos_cached

    def forward(self, q: torch.Tensor, k: torch.Tensor, pos_offset: int = 0) -> Tuple[torch.Tensor, torch.Tensor]:
        cur_seq_len = q.shape[self.seq_dim]
        yarn_factor = 1.0
        # Apply YaRN only for true extrapolation beyond pretrained length.
        if cur_seq_len > self.pretrained_seq_len:
            yarn_factor = cur_seq_len / self.pretrained_seq_len
        if cur_seq_len > self.extended_seq_len:
            if yarn_factor > 1.0:
                print(f"Applying YaRN scaling (seq_len={cur_seq_len}, factor={yarn_factor:.4f})...")
            self.extended_seq_len = cur_seq_len
        
        if self.n_rotate < self.n_embd:
            r_k = k[..., :self.n_rotate]
            nr_k = k[..., self.n_rotate:]
            r_q = q[..., :self.n_rotate]
            nr_q = q[..., self.n_rotate:]
            sin, cos = self.get(r_k, yarn_factor)
            r_q, r_k = self.apply_rotary_pos_emb(r_q, r_k, sin, cos, self.seq_dim, pos_offset)
            return torch.cat((r_q, nr_q), dim=-1), torch.cat((r_k, nr_k), dim=-1)
        else:
            sin, cos = self.get(k, yarn_factor)
            return self.apply_rotary_pos_emb(q, k, sin, cos, self.seq_dim, pos_offset)

class PolarPosEncoding(torch.nn.Module):
    def __init__(self, n_embd: int, n_head: int, block_size: int, base=10000, seq_dim: int = -2, 
                 rotate_fraction: float = 1.0, use_theta_bias: bool = True, thetab_init: str = "two_pi"):
        super().__init__()
        self.seq_len_cached = 0
        self.seq_dim = seq_dim
        self.n_embd = n_embd
        self.n_head = n_head
        self.block_size = block_size
        self.rotate_fraction = rotate_fraction
        self.use_theta_bias = use_theta_bias
        self.thetab_init = thetab_init
        if rotate_fraction < 1:
            self.n_rotate = int(n_embd * rotate_fraction)
            self.n_rotate -= self.n_rotate % 2
        else:
            self.n_rotate = n_embd
        self.delta_c = None
        # NOTE: arange [lo, hi) whereas linspace [lo, hi]
        inv_freq = 1.0 / (base ** (torch.arange(0, self.n_rotate).float() / self.n_rotate))
        self.register_buffer("inv_freq", inv_freq, persistent=False)
        if use_theta_bias:
            self.delta_c = self.init_theta_bias()
        self.theta_c_cached = None
        self.sin_cached = None
        self.cos_cached = None
    
    def get_imag_real_theta(self, x: torch.Tensor) -> Tuple[torch.Tensor, torch.Tensor, torch.Tensor]:
        """ Get the real, imaginary components and thetas for rotations. """
        seq_len = x.shape[self.seq_dim]
        if seq_len > self.seq_len_cached:
            self.seq_len_cached = seq_len
            t = torch.arange(x.shape[self.seq_dim], device=x.device).type_as(self.inv_freq)
            thetas = torch.einsum("i,j->ij", t, self.inv_freq)

            tgt_shape = [1] * x.ndim
            tgt_shape[self.seq_dim] = seq_len
            tgt_shape[-1] = x.shape[-1]
            
            self.theta_c_cached = thetas.view(*tgt_shape)
            self.sin_cached = self.theta_c_cached.sin()
            self.cos_cached = self.theta_c_cached.cos()

        return self.sin_cached.narrow(self.seq_dim, 0, seq_len), self.cos_cached.narrow(self.seq_dim, 0, seq_len), self.theta_c_cached.narrow(self.seq_dim, 0, seq_len)
    
    def init_theta_bias(self):
        delta_c = torch.nn.Parameter(torch.empty(1, self.n_head, 1, self.n_rotate))
        if self.thetab_init == "two_pi":
            upper = 0.
            lower = (-2 * math.pi / torch.maximum(self.inv_freq, torch.tensor(1./self.block_size))) * self.inv_freq
            delta_c.data.copy_(torch.rand([1, self.n_head, 1, self.inv_freq.shape[0]]) * (upper - lower) + lower)
        elif self.thetab_init == "zero":
            torch.nn.init.zeros_(delta_c)
        elif self.thetab_init == "hybrid":
            # init same as two_pi, then set those bias idxs to zero whose period exceed block_size
            upper = 0.
            lower = (-2 * math.pi / torch.maximum(self.inv_freq, torch.tensor(1./self.block_size))) * self.inv_freq
            delta_c.data.copy_(torch.rand([1, self.n_head, 1, self.inv_freq.shape[0]]) * (upper - lower) + lower)
            mask = self.inv_freq < (1. / self.block_size)
            delta_c.data[:, :, :, mask] = 0.
        return delta_c

    def forward(self, q: torch.Tensor, k: torch.Tensor, pos_offset: int = 0) -> Tuple[Tuple[torch.Tensor, torch.Tensor], Tuple[torch.Tensor, torch.Tensor]]:
        if self.n_rotate < self.n_embd:
            r_q = q[..., :self.n_rotate]
            nr_q = q[..., self.n_rotate:]
            r_k = k[..., :self.n_rotate]
            nr_k = k[..., self.n_rotate:]

            sin, cos, theta = self.get_imag_real_theta(r_q)

            if self.use_theta_bias:
                theta_shifted = theta + torch.clamp(self.delta_c, -2*math.pi, 0.)
                sin_shifted = theta_shifted.sin()
                cos_shifted = theta_shifted.cos()
            else:
                sin_shifted = sin
                cos_shifted = cos

            r_q = torch.nn.functional.softplus(r_q)
            r_k = torch.nn.functional.softplus(r_k)
            
            real_r_q = r_q * cos
            imag_r_q = r_q * sin
            real_r_k = r_k * cos_shifted
            imag_r_k = r_k * sin_shifted
            
            out_real_q = torch.cat((real_r_q, nr_q), dim=-1)
            out_imag_q = torch.cat((imag_r_q, torch.zeros_like(nr_q)), dim=-1)

            out_real_k = torch.cat((real_r_k, nr_k), dim=-1)
            out_imag_k = torch.cat((imag_r_k, torch.zeros_like(nr_k)), dim=-1)
            return (out_real_q, out_imag_q), (out_real_k, out_imag_k)
        
        else:
            q = torch.nn.functional.softplus(q)
            k = torch.nn.functional.softplus(k)            
            sin, cos, theta = self.get_imag_real_theta(q)
            if self.use_theta_bias:
                theta_shifted = theta + torch.clamp(self.delta_c, -2*math.pi, 0.)
                sin_shifted = theta_shifted.sin()
                cos_shifted = theta_shifted.cos()
            else:
                sin_shifted = sin
                cos_shifted = cos
            real_q = q * cos
            imag_q = q * sin
            real_k = k * cos_shifted
            imag_k = k * sin_shifted
            return (real_q, imag_q), (real_k, imag_k)

# YaRN utils
def find_correction_dim(num_rotations: float, emb_dim: int, base: float, max_seq_len: int) -> float:
        """
        Computes the correction dimension for a given number of rotations in the rotary positional embedding.

        Args:
            num_rotations (float): Number of rotations to compute the correction for.
            emb_dim (int): Dimensionality of the embedding space.
            base (float): Base value for the exponential computation.
            max_seq_len (int): Maximum sequence length.

        Returns:
            float: The correction dimension based on the input parameters.
        """
        return (emb_dim* math.log(max_seq_len / (num_rotations * 2 * math.pi)) / (2 * math.log(base)))

def find_correction_range(low_rot: float, high_rot: float, emb_dim: int, base: float, max_seq_len: int) -> Tuple[int, int]:
        """
        Computes the range of correction dimensions for rotary positional embeddings.

        Args:
            low_rot (float): Lower bound for the number of rotations.
            high_rot (float): Upper bound for the number of rotations.
            emb_dim (int): Dimensionality of the embedding space.
            base (float): Base value for the exponential computation.
            max_seq_len (int): Maximum sequence length.

        Returns:
            tuple[int, int]: The range of correction dimensions (low, high), clamped to valid indices.
        """
        low = math.floor(find_correction_dim(low_rot, emb_dim, base, max_seq_len))
        high = math.ceil(find_correction_dim(high_rot, emb_dim, base, max_seq_len))
        return max(low, 0), min(high, emb_dim - 1)

def linear_ramp_factor(alpha: float, beta: float, rope_dim: int) -> torch.Tensor:
        """ Computes a linear ramp function used to smooth values between a minimum and maximum range.
        Equation 18 from YaRN paper.
        Args:
            alpha: Minimum value for the ramp function.
            beta: Maximum value for the ramp function.
            rope_dim: Dimensionality of the ramp tensor.
        """
        if alpha == beta:
            beta += 0.001
        linear_func = (torch.arange(rope_dim, dtype=torch.float32) - alpha) / (beta - alpha)
        ramp_func = torch.clamp(linear_func, 0, 1)
        return ramp_func