171 lines
5.6 KiB
Python
171 lines
5.6 KiB
Python
from dataclasses import dataclass
|
|
|
|
import torch
|
|
import torch.nn as nn
|
|
import torch.nn.functional as F
|
|
|
|
|
|
@dataclass
|
|
class ModelConfig:
|
|
vocab_size: int = 256
|
|
d_model: int = 352
|
|
n_heads: int = 8
|
|
n_layers: int = 7
|
|
ffn_mult: int = 4
|
|
max_seq_len: int = 16_384
|
|
dropout: float = 0.0
|
|
rope_theta: float = 10_000.0
|
|
|
|
|
|
class RMSNorm(nn.Module):
|
|
def __init__(self, dim: int, eps: float = 1e-6):
|
|
super().__init__()
|
|
self.eps = eps
|
|
self.weight = nn.Parameter(torch.ones(dim))
|
|
|
|
def forward(self, x: torch.Tensor) -> torch.Tensor:
|
|
norm = x.pow(2).mean(dim=-1, keepdim=True)
|
|
x = x * torch.rsqrt(norm + self.eps)
|
|
return self.weight * x
|
|
|
|
|
|
def precompute_rope_cache(max_seq_len: int, head_dim: int, theta: float, device: torch.device):
|
|
if head_dim % 2 != 0:
|
|
raise ValueError(f"head_dim must be even for RoPE, got {head_dim}")
|
|
inv_freq = 1.0 / (theta ** (torch.arange(0, head_dim, 2, device=device).float() / head_dim))
|
|
t = torch.arange(max_seq_len, device=device).float()
|
|
freqs = torch.outer(t, inv_freq)
|
|
return freqs.cos(), freqs.sin()
|
|
|
|
|
|
def apply_rotary(q: torch.Tensor, k: torch.Tensor, cos: torch.Tensor, sin: torch.Tensor):
|
|
q1, q2 = q[..., ::2], q[..., 1::2]
|
|
k1, k2 = k[..., ::2], k[..., 1::2]
|
|
|
|
q_rot_1 = q1 * cos - q2 * sin
|
|
q_rot_2 = q1 * sin + q2 * cos
|
|
k_rot_1 = k1 * cos - k2 * sin
|
|
k_rot_2 = k1 * sin + k2 * cos
|
|
|
|
q_out = torch.stack((q_rot_1, q_rot_2), dim=-1).flatten(-2)
|
|
k_out = torch.stack((k_rot_1, k_rot_2), dim=-1).flatten(-2)
|
|
return q_out, k_out
|
|
|
|
|
|
class CausalSelfAttention(nn.Module):
|
|
def __init__(self, cfg: ModelConfig):
|
|
super().__init__()
|
|
if cfg.d_model % cfg.n_heads != 0:
|
|
raise ValueError("d_model must be divisible by n_heads")
|
|
self.n_heads = cfg.n_heads
|
|
self.head_dim = cfg.d_model // cfg.n_heads
|
|
self.qkv = nn.Linear(cfg.d_model, 3 * cfg.d_model, bias=False)
|
|
self.proj = nn.Linear(cfg.d_model, cfg.d_model, bias=False)
|
|
self.dropout = nn.Dropout(cfg.dropout)
|
|
|
|
def forward(self, x: torch.Tensor, cos: torch.Tensor, sin: torch.Tensor) -> torch.Tensor:
|
|
bsz, seq_len, dim = x.size()
|
|
qkv = self.qkv(x)
|
|
q, k, v = qkv.chunk(3, dim=-1)
|
|
|
|
q = q.view(bsz, seq_len, self.n_heads, self.head_dim).transpose(1, 2)
|
|
k = k.view(bsz, seq_len, self.n_heads, self.head_dim).transpose(1, 2)
|
|
v = v.view(bsz, seq_len, self.n_heads, self.head_dim).transpose(1, 2)
|
|
|
|
# RoPE cache is [T, head_dim/2], expand to [1, 1, T, head_dim/2].
|
|
cos = cos.unsqueeze(0).unsqueeze(0)
|
|
sin = sin.unsqueeze(0).unsqueeze(0)
|
|
q, k = apply_rotary(q, k, cos, sin)
|
|
|
|
out = F.scaled_dot_product_attention(
|
|
q,
|
|
k,
|
|
v,
|
|
attn_mask=None,
|
|
dropout_p=self.dropout.p if self.training else 0.0,
|
|
is_causal=True,
|
|
)
|
|
out = out.transpose(1, 2).contiguous().view(bsz, seq_len, dim)
|
|
out = self.proj(out)
|
|
return self.dropout(out)
|
|
|
|
|
|
class MLP(nn.Module):
|
|
def __init__(self, cfg: ModelConfig):
|
|
super().__init__()
|
|
hidden = cfg.ffn_mult * cfg.d_model
|
|
self.fc = nn.Linear(cfg.d_model, hidden, bias=False)
|
|
self.proj = nn.Linear(hidden, cfg.d_model, bias=False)
|
|
self.dropout = nn.Dropout(cfg.dropout)
|
|
|
|
def forward(self, x: torch.Tensor) -> torch.Tensor:
|
|
x = self.fc(x)
|
|
x = F.gelu(x, approximate="tanh")
|
|
x = self.proj(x)
|
|
return self.dropout(x)
|
|
|
|
|
|
class Block(nn.Module):
|
|
def __init__(self, cfg: ModelConfig):
|
|
super().__init__()
|
|
self.norm1 = RMSNorm(cfg.d_model)
|
|
self.attn = CausalSelfAttention(cfg)
|
|
self.norm2 = RMSNorm(cfg.d_model)
|
|
self.mlp = MLP(cfg)
|
|
|
|
def forward(self, x: torch.Tensor, cos: torch.Tensor, sin: torch.Tensor) -> torch.Tensor:
|
|
x = x + self.attn(self.norm1(x), cos, sin)
|
|
x = x + self.mlp(self.norm2(x))
|
|
return x
|
|
|
|
|
|
class MiniLM(nn.Module):
|
|
def __init__(self, cfg: ModelConfig):
|
|
super().__init__()
|
|
self.cfg = cfg
|
|
self.tok_emb = nn.Embedding(cfg.vocab_size, cfg.d_model)
|
|
self.drop = nn.Dropout(cfg.dropout)
|
|
self.blocks = nn.ModuleList([Block(cfg) for _ in range(cfg.n_layers)])
|
|
self.norm_f = RMSNorm(cfg.d_model)
|
|
self.lm_head = nn.Linear(cfg.d_model, cfg.vocab_size, bias=False)
|
|
|
|
# Weight tying keeps the model near 10M params with this config.
|
|
self.lm_head.weight = self.tok_emb.weight
|
|
|
|
cos, sin = precompute_rope_cache(
|
|
max_seq_len=cfg.max_seq_len,
|
|
head_dim=cfg.d_model // cfg.n_heads,
|
|
theta=cfg.rope_theta,
|
|
device=torch.device("cpu"),
|
|
)
|
|
self.register_buffer("rope_cos", cos, persistent=False)
|
|
self.register_buffer("rope_sin", sin, persistent=False)
|
|
|
|
def forward(self, idx: torch.Tensor, targets: torch.Tensor | None = None):
|
|
bsz, seq_len = idx.size()
|
|
if seq_len > self.cfg.max_seq_len:
|
|
raise ValueError(
|
|
f"Input sequence length {seq_len} exceeds max_seq_len {self.cfg.max_seq_len}"
|
|
)
|
|
|
|
x = self.tok_emb(idx)
|
|
x = self.drop(x)
|
|
|
|
cos = self.rope_cos[:seq_len].to(x.device)
|
|
sin = self.rope_sin[:seq_len].to(x.device)
|
|
|
|
for block in self.blocks:
|
|
x = block(x, cos, sin)
|
|
|
|
x = self.norm_f(x)
|
|
logits = self.lm_head(x)
|
|
|
|
loss = None
|
|
if targets is not None:
|
|
loss = F.cross_entropy(logits.view(-1, logits.size(-1)), targets.view(-1))
|
|
return logits, loss
|
|
|
|
|
|
def count_parameters(model: nn.Module) -> int:
|
|
return sum(p.numel() for p in model.parameters())
|