ViuAI's picture
init: 28L MoE 241M audit-fixed, smoke pass
140f9f6 verified
Raw History Blame Contribute Delete
8.18 kB
"""
Mini-ViuAI-50M — advanced decoder-only transformer
- RMSNorm, RoPE, SwiGLU, GQA, tied embeddings, no bias
- target: ~50M params with vocab 48k (dim=512, layers=8)
Run param check:
python model.py --config ../configs/model_config.yaml
"""
import argparse
import math
from dataclasses import dataclass
from pathlib import Path
import torch
import torch.nn as nn
import torch.nn.functional as F
try:
import yaml
except ImportError:
yaml = None
@dataclass
class ModelArgs:
dim: int = 512
n_layers: int = 8
n_heads: int = 8
n_kv_heads: int = 2
vocab_size: int = 48000
multiple_of: int = 64
ffn_dim_multiplier: float = 2.7
norm_eps: float = 1e-5
rope_theta: float = 10000.0
max_seq_len: int = 1024
dropout: float = 0.0
bias: bool = False
tie_embeddings: bool = True
class RMSNorm(nn.Module):
def __init__(self, dim: int, eps: float = 1e-5):
super().__init__()
self.eps = eps
self.weight = nn.Parameter(torch.ones(dim))
def forward(self, x):
v = x.pow(2).mean(-1, keepdim=True)
x = x * torch.rsqrt(v + self.eps)
return self.weight * x
def precompute_rope(head_dim: int, seq_len: int, theta: float, device, dtype=torch.float32):
assert head_dim % 2 == 0
inv = 1.0 / (theta ** (torch.arange(0, head_dim, 2, device=device, dtype=torch.float32) / head_dim))
t = torch.arange(seq_len, device=device, dtype=torch.float32)
freqs = torch.outer(t, inv) # [T, head_dim/2]
return torch.cos(freqs).to(dtype), torch.sin(freqs).to(dtype) # each [T, D/2]
def apply_rope(x, cos, sin):
# x: [B, T, H, D]
d = x.shape[-1]
x1, x2 = x[..., : d // 2], x[..., d // 2 :]
# cos/sin: [T, D/2] -> [1, T, 1, D/2]
cos = cos[: x.shape[1]].unsqueeze(0).unsqueeze(2)
sin = sin[: x.shape[1]].unsqueeze(0).unsqueeze(2)
return torch.cat([x1 * cos - x2 * sin, x1 * sin + x2 * cos], dim=-1)
class Attention(nn.Module):
def __init__(self, a: ModelArgs):
super().__init__()
self.n_heads = a.n_heads
self.n_kv = a.n_kv_heads
self.head_dim = a.dim // a.n_heads
self.n_rep = a.n_heads // a.n_kv_heads
self.scale = 1.0 / math.sqrt(self.head_dim)
self.wq = nn.Linear(a.dim, a.n_heads * self.head_dim, bias=a.bias)
self.wk = nn.Linear(a.dim, self.n_kv * self.head_dim, bias=a.bias)
self.wv = nn.Linear(a.dim, self.n_kv * self.head_dim, bias=a.bias)
self.wo = nn.Linear(a.n_heads * self.head_dim, a.dim, bias=a.bias)
self.drop = nn.Dropout(a.dropout)
def forward(self, x, cos, sin, mask=None):
B, T, _ = x.shape
q = self.wq(x).view(B, T, self.n_heads, self.head_dim)
k = self.wk(x).view(B, T, self.n_kv, self.head_dim)
v = self.wv(x).view(B, T, self.n_kv, self.head_dim)
q, k = apply_rope(q, cos, sin), apply_rope(k, cos, sin)
if self.n_rep > 1: # GQA expand
k = k.repeat_interleave(self.n_rep, dim=2)
v = v.repeat_interleave(self.n_rep, dim=2)
q = q.transpose(1, 2) # [B, H, T, D]
k = k.transpose(1, 2)
v = v.transpose(1, 2)
attn = (q @ k.transpose(-2, -1)) * self.scale
if mask is not None:
attn = attn + mask
attn = F.softmax(attn, dim=-1)
attn = self.drop(attn)
out = (attn @ v).transpose(1, 2).contiguous().view(B, T, -1)
return self.wo(out)
class FeedForward(nn.Module):
"""SwiGLU: down(silu(gate) * up)"""
def __init__(self, a: ModelArgs):
super().__init__()
hidden = int(2 * 4 * a.dim / 3)
if a.ffn_dim_multiplier:
hidden = int(a.ffn_dim_multiplier * a.dim)
hidden = a.multiple_of * ((hidden + a.multiple_of - 1) // a.multiple_of)
self.w1 = nn.Linear(a.dim, hidden, bias=a.bias) # gate
self.w2 = nn.Linear(hidden, a.dim, bias=a.bias) # down
self.w3 = nn.Linear(a.dim, hidden, bias=a.bias) # up
self.drop = nn.Dropout(a.dropout)
def forward(self, x):
return self.drop(self.w2(F.silu(self.w1(x)) * self.w3(x)))
class Block(nn.Module):
def __init__(self, a: ModelArgs):
super().__init__()
self.attn_norm = RMSNorm(a.dim, a.norm_eps)
self.attn = Attention(a)
self.ffn_norm = RMSNorm(a.dim, a.norm_eps)
self.ffn = FeedForward(a)
def forward(self, x, cos, sin, mask):
x = x + self.attn(self.attn_norm(x), cos, sin, mask)
x = x + self.ffn(self.ffn_norm(x))
return x
class MiniViuTransformer(nn.Module):
def __init__(self, a: ModelArgs):
super().__init__()
self.args = a
self.tok_emb = nn.Embedding(a.vocab_size, a.dim)
self.drop = nn.Dropout(a.dropout)
self.layers = nn.ModuleList([Block(a) for _ in range(a.n_layers)])
self.norm = RMSNorm(a.dim, a.norm_eps)
self.head = nn.Linear(a.dim, a.vocab_size, bias=False)
if a.tie_embeddings:
self.head.weight = self.tok_emb.weight # save ~24.5M params
self.apply(self._init)
@staticmethod
def _init(m):
if isinstance(m, nn.Linear):
nn.init.xavier_uniform_(m.weight)
if m.bias is not None:
nn.init.zeros_(m.bias)
elif isinstance(m, nn.Embedding):
nn.init.normal_(m.weight, std=0.02)
def forward(self, ids, targets=None):
B, T = ids.shape
assert T <= self.args.max_seq_len, f"seq {T} > max {self.args.max_seq_len}"
x = self.drop(self.tok_emb(ids))
cos, sin = precompute_rope(
self.args.dim // self.args.n_heads, T,
self.args.rope_theta, ids.device, dtype=torch.float32,
)
mask = torch.full((T, T), float("-inf"), device=ids.device)
mask = torch.triu(mask, diagonal=1).unsqueeze(0).unsqueeze(0) # causal
for blk in self.layers:
x = blk(x, cos, sin, mask)
logits = self.head(self.norm(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_params(self):
# parameters() already dedupes tied weights (same Parameter object once)
return sum(p.numel() for p in self.parameters())
@torch.no_grad()
def generate(self, ids, max_new=50, temperature=0.8, top_k=50):
self.eval()
for _ in range(max_new):
ctx = ids[:, -self.args.max_seq_len :]
logits, _ = self.forward(ctx)
logits = logits[:, -1, :] / max(temperature, 1e-5)
if top_k:
v, _ = torch.topk(logits, min(top_k, logits.size(-1)))
logits[logits < v[:, [-1]]] = float("-inf")
probs = F.softmax(logits, dim=-1)
nxt = torch.multinomial(probs, 1)
ids = torch.cat([ids, nxt], dim=1)
return ids
def load_args(path: str | None) -> ModelArgs:
if path and yaml and Path(path).exists():
d = yaml.safe_load(open(path, encoding="utf-8")) or {}
return ModelArgs(**{k: v for k, v in d.items() if k in ModelArgs.__dataclass_fields__})
return ModelArgs()
if __name__ == "__main__":
ap = argparse.ArgumentParser()
ap.add_argument("--config", default="../configs/model_config.yaml")
args = ap.parse_args()
cfg = load_args(args.config)
m = MiniViuTransformer(cfg)
n = m.count_params()
print(f"[config] dim={cfg.dim} layers={cfg.n_layers} heads={cfg.n_heads} kv={cfg.n_kv_heads} vocab={cfg.vocab_size} tied={cfg.tie_embeddings}")
print(f"[params] total = {n:,} ({n/1e6:.1f}M) | target ~50M")
emb = cfg.vocab_size * cfg.dim
print(f"[split] embeddings = {emb:,} ({emb/1e6:.1f}M) | transformer = {n-emb:,}")
if not (45e6 <= n <= 60e6):
print("[warn] 50M se bahar hai — layers/dim adjust karo.")
else:
print("[ok] budget me hai.")
# quick forward smoke test
m.eval()
with torch.no_grad():
ids = torch.randint(0, cfg.vocab_size, (2, 16))
logits, loss = m(ids, ids)
print(f"[smoke] logits {tuple(logits.shape)} loss {loss.item():.3f} OK")