import math from typing import Optional, Tuple import torch import torch.nn as nn import torch.nn.functional as F class RMSNorm(nn.Module): """ Standart Qwen2.5 / LLaMA RMSNorm: y = (x / RMS(x)) * weight, weight 1.0 ile başlatılır. """ def __init__(self, dim: int, eps: float = 1e-6, dtype: torch.dtype = torch.float32): super().__init__() self.eps = eps self.weight = nn.Parameter(torch.ones(dim, dtype=dtype)) def forward(self, x: torch.Tensor) -> torch.Tensor: input_dtype = x.dtype x_fp32 = x.to(torch.float32) variance = x_fp32.pow(2).mean(-1, keepdim=True) normed = x_fp32 * torch.rsqrt(variance + self.eps) return (normed * self.weight).to(input_dtype) def precompute_rope_freqs( head_dim: int, seq_len: int, theta: float = 1000000.0, device: str = "cpu" ) -> Tuple[torch.Tensor, torch.Tensor]: """ Qwen2.5 RoPE (Rotary Position Embeddings) frekanslarını önceden hesaplar. Varsayılan theta: 1,000,000 (Qwen2.5 standardı). """ freqs = 1.0 / (theta ** (torch.arange(0, head_dim, 2, dtype=torch.float32, device=device) / head_dim)) t = torch.arange(seq_len, dtype=torch.float32, device=device) angles = torch.outer(t, freqs) cos = torch.cos(angles) sin = torch.sin(angles) return cos, sin def apply_rope( x: torch.Tensor, cos: torch.Tensor, sin: torch.Tensor, position_ids: Optional[torch.Tensor] = None ) -> torch.Tensor: """ Qwen2.5 rotate_half rotasyonunu uygular: x shape: (B, num_heads, seq_len, head_dim) """ B, H, S, D = x.shape half = D // 2 if position_ids is not None: cos = cos[position_ids].unsqueeze(1) # (B, 1, S, half) sin = sin[position_ids].unsqueeze(1) else: cos = cos[:S].unsqueeze(0).unsqueeze(1) # (1, 1, S, half) sin = sin[:S].unsqueeze(0).unsqueeze(1) x1 = x[..., :half] x2 = x[..., half:] rotated = torch.cat([x1 * cos - x2 * sin, x2 * cos + x1 * sin], dim=-1) return rotated.to(x.dtype) class MultiHeadAttention(nn.Module): """ Qwen2.5 Grouped Query Attention (GQA) & KV-Cache: - q_proj, k_proj, v_proj (bias=True) - o_proj (bias=False) - RoPE - KV-Cache desteği """ def __init__( self, num_heads: int, num_kv_heads: int, d_model: int, qkv_bias: bool = True, dtype: torch.dtype = torch.float32, ): super().__init__() self.num_heads = num_heads self.num_kv_heads = num_kv_heads self.d_model = d_model self.head_dim = d_model // num_heads self.kv_dim = num_kv_heads * self.head_dim self.q_proj = nn.Linear(d_model, d_model, bias=qkv_bias, dtype=dtype) self.k_proj = nn.Linear(d_model, self.kv_dim, bias=qkv_bias, dtype=dtype) self.v_proj = nn.Linear(d_model, self.kv_dim, bias=qkv_bias, dtype=dtype) self.out_proj = nn.Linear(d_model, d_model, bias=False, dtype=dtype) def forward( self, q_input: torch.Tensor, kv_input: Optional[torch.Tensor] = None, mask: Optional[torch.Tensor] = None, rope: Optional[Tuple[torch.Tensor, torch.Tensor]] = None, position_ids: Optional[torch.Tensor] = None, kv_cache: Optional[Tuple[torch.Tensor, torch.Tensor]] = None, use_cache: bool = False, ) -> Tuple[torch.Tensor, Optional[Tuple[torch.Tensor, torch.Tensor]]]: if kv_input is None: kv_input = q_input B, Sq, _ = q_input.shape _, Sk, _ = kv_input.shape q = self.q_proj(q_input).view(B, Sq, self.num_heads, self.head_dim).transpose(1, 2) k = self.k_proj(kv_input).view(B, Sk, self.num_kv_heads, self.head_dim).transpose(1, 2) v = self.v_proj(kv_input).view(B, Sk, self.num_kv_heads, self.head_dim).transpose(1, 2) # RoPE Rotasyonu if rope is not None: cos, sin = rope cos, sin = cos.to(q.device), sin.to(q.device) q = apply_rope(q, cos, sin, position_ids) k = apply_rope(k, cos, sin, position_ids) # KV Cache new_kv_cache = None if kv_cache is not None: past_k, past_v = kv_cache k = torch.cat([past_k, k], dim=2) v = torch.cat([past_v, v], dim=2) if use_cache: new_kv_cache = (k, v) # GQA Repeat Interleave repeats = self.num_heads // self.num_kv_heads if repeats > 1: k = k.repeat_interleave(repeats, dim=1) v = v.repeat_interleave(repeats, dim=1) scale = 1.0 / math.sqrt(self.head_dim) out = F.scaled_dot_product_attention( q, k, v, attn_mask=mask, dropout_p=0.0, scale=scale, ) out = out.transpose(1, 2).contiguous().view(B, Sq, self.d_model) return self.out_proj(out), new_kv_cache class FeedForward(nn.Module): """ Qwen2.5 SwiGLU MLP: FFN(x) = down_proj(SiLU(gate_proj(x)) * up_proj(x)) """ def __init__(self, d_model: int, d_ff: int, dtype: torch.dtype = torch.float32): super().__init__() self.gate_proj = nn.Linear(d_model, d_ff, bias=False, dtype=dtype) self.up_proj = nn.Linear(d_model, d_ff, bias=False, dtype=dtype) self.down_proj = nn.Linear(d_ff, d_model, bias=False, dtype=dtype) def forward(self, x: torch.Tensor) -> torch.Tensor: return self.down_proj(F.silu(self.gate_proj(x)) * self.up_proj(x)) class TransformerBlock(nn.Module): """ Qwen2.5 Decoder Layer: Pre-norm RMSNorm + GQA Attention + Post-attention RMSNorm + SwiGLU MLP """ def __init__( self, d_model: int, num_heads: int, num_kv_heads: int, d_ff: int, dropout_rate: float = 0.0, dtype: torch.dtype = torch.float32, ): super().__init__() self.norm1 = RMSNorm(d_model, dtype=dtype) self.self_attn = MultiHeadAttention(num_heads, num_kv_heads, d_model, qkv_bias=True, dtype=dtype) self.dropout = nn.Dropout(dropout_rate) if dropout_rate > 0.0 else nn.Identity() self.norm2 = RMSNorm(d_model, dtype=dtype) self.ffn = FeedForward(d_model, d_ff, dtype=dtype) def forward( self, x: torch.Tensor, mask: Optional[torch.Tensor] = None, rope: Optional[Tuple[torch.Tensor, torch.Tensor]] = None, position_ids: Optional[torch.Tensor] = None, kv_cache: Optional[Tuple[torch.Tensor, torch.Tensor]] = None, use_cache: bool = False, ) -> Tuple[torch.Tensor, Optional[Tuple[torch.Tensor, torch.Tensor]]]: residual = x normed = self.norm1(x) attn_out, new_kv_cache = self.self_attn( normed, mask=mask, rope=rope, position_ids=position_ids, kv_cache=kv_cache, use_cache=use_cache, ) x = residual + self.dropout(attn_out) residual = x normed = self.norm2(x) ffn_out = self.ffn(normed) x = residual + self.dropout(ffn_out) return x, new_kv_cache