NeuroVoice-0.5B / architecture.py
TurkishCodeMan's picture
Upload architecture.py with huggingface_hub
ec6ed45 verified
Raw History Blame Contribute Delete
7.18 kB
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