picoLM / model.py
jankwong705's picture
Upload model.py with huggingface_hub
730ad7f verified
Raw
History Blame Contribute Delete
11.8 kB
import torch
import torch.nn as nn
import torch.nn.functional as F
from dataclasses import dataclass
from torch.utils.checkpoint import checkpoint
from pathlib import Path
from typing import Any, Dict
@dataclass
class PicoLMMoEArgs:
#Configuration for picoLM MoE — mixture-of-experts transformer ≤100M parameters.
dim: int = 448
n_layers: int = 10
n_heads: int = 8
n_kv_heads: int = 2
head_dim: int = 56
vocab_size: int = 50257
max_seq_len: int = 1024
# Dense block (first block)
dense_ffn_hidden_dim: int = 2048
# MoE blocks
n_experts: int = 8
n_active_experts: int = 4
moe_expert_hidden_dim: int = 384
norm_eps: float = 1e-5
use_grad_checkpoint: bool = False
def load_picolm_moe_args(config_path: str = "config.yaml", profile: str = "picolm_moe") -> PicoLMMoEArgs:
#Load picoLM MoE config from YAML.
path = Path(config_path)
if not path.exists():
return PicoLMMoEArgs()
try:
import yaml
except ImportError as e:
raise RuntimeError("PyYAML is required for config.yaml loading. Install: pip install pyyaml") from e
with path.open("r", encoding="utf-8") as f:
cfg: Any = yaml.safe_load(f) or {}
if not isinstance(cfg, dict):
raise ValueError("config.yaml must be a mapping")
model_cfg = cfg.get(profile, {})
if not isinstance(model_cfg, dict):
raise ValueError(f"config.yaml: '{profile}' must be a mapping")
valid = set(PicoLMMoEArgs.__dataclass_fields__.keys())
kwargs = {k: v for k, v in model_cfg.items() if k in valid}
return PicoLMMoEArgs(**kwargs)
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):
norm_x = torch.mean(x ** 2, dim=-1, keepdim=True)
x_normed = x * torch.rsqrt(norm_x + self.eps)
return self.weight * x_normed
def precompute_freqs_cis(dim: int, end: int, theta: float = 10000.0):
freqs = 1.0 / (theta ** (torch.arange(0, dim, 2)[: (dim // 2)].float() / dim))
t = torch.arange(end, device=freqs.device)
freqs = torch.outer(t, freqs).float()
freqs_cis = torch.polar(torch.ones_like(freqs), freqs)
return freqs_cis
def apply_rotary_emb(xq, xk, freqs_cis):
# Reshape to complex numbers
xq_ = torch.view_as_complex(xq.float().reshape(*xq.shape[:-1], -1, 2))
xk_ = torch.view_as_complex(xk.float().reshape(*xk.shape[:-1], -1, 2))
# Apply rotation
freqs_cis = freqs_cis.unsqueeze(0).unsqueeze(2)
xq_out = torch.view_as_real(xq_ * freqs_cis).flatten(3)
xk_out = torch.view_as_real(xk_ * freqs_cis).flatten(3)
return xq_out.type_as(xq), xk_out.type_as(xk)
class SwiGLU(nn.Module):
#SwiGLU feed-forward module.
def __init__(self, dim: int, hidden_dim: int):
super().__init__()
self.w1 = nn.Linear(dim, hidden_dim, bias=False) # goes to SiLU
self.w3 = nn.Linear(dim, hidden_dim, bias=False) # bypasses SiLU
self.w2 = nn.Linear(hidden_dim, dim, bias=False) # final output
def forward(self, x):
return self.w2(F.silu(self.w1(x)) * self.w3(x))
class GroupedQueryAttention(nn.Module):
def __init__(self, args: PicoLMMoEArgs):
super().__init__()
self.n_heads = args.n_heads
self.n_kv_heads = args.n_kv_heads
self.head_dim = args.head_dim
self.n_rep = self.n_heads // self.n_kv_heads
self.wq = nn.Linear(args.dim, args.n_heads * args.head_dim, bias=False)
self.wk = nn.Linear(args.dim, args.n_kv_heads * args.head_dim, bias=False)
self.wv = nn.Linear(args.dim, args.n_kv_heads * args.head_dim, bias=False)
self.wo = nn.Linear(args.n_heads * args.head_dim, args.dim, bias=False)
# QK-Norm
self.q_norm = RMSNorm(args.head_dim)
self.k_norm = RMSNorm(args.head_dim)
def forward(self, x, freqs_cis):
bsz, seqlen, _ = x.shape
xq, xk, xv = self.wq(x), self.wk(x), self.wv(x)
xq = xq.view(bsz, seqlen, self.n_heads, self.head_dim)
xk = xk.view(bsz, seqlen, self.n_kv_heads, self.head_dim)
xv = xv.view(bsz, seqlen, self.n_kv_heads, self.head_dim)
# Apply QK-Norm before RoPE
xq = self.q_norm(xq)
xk = self.k_norm(xk)
# Apply RoPE
xq, xk = apply_rotary_emb(xq, xk, freqs_cis)
# Grouped-Query Attention: repeat K and V
xk = torch.repeat_interleave(xk, self.n_rep, dim=2)
xv = torch.repeat_interleave(xv, self.n_rep, dim=2)
xq = xq.transpose(1, 2)
xk = xk.transpose(1, 2)
xv = xv.transpose(1, 2)
# FlashAttention via PyTorch SDPA
output = F.scaled_dot_product_attention(
xq, xk, xv,
attn_mask=None,
dropout_p=0.0,
is_causal=True,
)
output = output.transpose(1, 2).contiguous().view(bsz, seqlen, -1)
return self.wo(output)
class MoELayer(nn.Module):
def __init__(self, args: PicoLMMoEArgs):
super().__init__()
self.n_experts = args.n_experts
self.n_active = args.n_active_experts
# Router
self.gate = nn.Linear(args.dim, args.n_experts, bias=False)
# Routed experts
self.experts = nn.ModuleList([
SwiGLU(args.dim, args.moe_expert_hidden_dim) for _ in range(args.n_experts)
])
# Shared expert (always active per token)
self.shared_expert = SwiGLU(args.dim, args.moe_expert_hidden_dim)
def forward(self, x):
bsz, seq_len, dim = x.shape
x_flat = x.view(-1, dim)
router_logits = self.gate(x_flat)
routing_weights = F.softmax(router_logits, dim=-1)
routing_weights, selected_experts = torch.topk(routing_weights, self.n_active, dim=-1)
routing_weights /= routing_weights.sum(dim=-1, keepdim=True)
final_output = torch.zeros_like(x_flat)
# One pass per expert: gather all tokens assigned to it, run batched forward,
# then scatter-add weighted results back.
for expert_id in range(self.n_experts):
expert_mask = (selected_experts == expert_id)
if not expert_mask.any():
continue
token_indices, slot_indices = expert_mask.nonzero(as_tuple=True)
weights = routing_weights[token_indices, slot_indices].unsqueeze(-1)
expert_out = self.experts[expert_id](x_flat[token_indices])
final_output.scatter_add_(
0,
token_indices.unsqueeze(-1).expand_as(expert_out),
expert_out * weights,
)
shared_out = self.shared_expert(x_flat)
return (final_output + shared_out).view(bsz, seq_len, dim)
class TransformerBlock(nn.Module):
def __init__(self, args: PicoLMMoEArgs, layer_id: int):
super().__init__()
self.layer_id = layer_id
self.attention = GroupedQueryAttention(args)
self.attention_norm = RMSNorm(args.dim, eps=args.norm_eps)
self.ffn_norm = RMSNorm(args.dim, eps=args.norm_eps)
# First block uses dense FFN (remaining blocks use MoE)
if layer_id == 0:
self.ffn = SwiGLU(args.dim, args.dense_ffn_hidden_dim)
else:
self.ffn = MoELayer(args)
def forward(self, x, freqs_cis):
# Attention with residual
h = x + self.attention(self.attention_norm(x), freqs_cis)
# FFN/MoE with residual
out = h + self.ffn(self.ffn_norm(h))
return out
class PicoLMMoE(nn.Module):
#picoLM MoE — mixture-of-experts transformer for the ≤100M parameter contest
def __init__(self, args: PicoLMMoEArgs):
super().__init__()
self.args = args
self.vocab_size = args.vocab_size
# Token embedding
self.tok_embeddings = nn.Embedding(args.vocab_size, args.dim)
# Transformer blocks
self.layers = nn.ModuleList()
for layer_id in range(args.n_layers):
self.layers.append(TransformerBlock(args, layer_id))
# Final RMSNorm
self.norm = RMSNorm(args.dim, eps=args.norm_eps)
# LM head
self.output = nn.Linear(args.dim, args.vocab_size, bias=False)
# Precompute RoPE frequencies
self.register_buffer(
"freqs_cis",
precompute_freqs_cis(args.head_dim, args.max_seq_len),
persistent=False
)
self.freqs_cis: torch.Tensor
def forward(self, tokens):
_bsz, seqlen = tokens.shape
h = self.tok_embeddings(tokens)
freqs_cis = self.freqs_cis[:seqlen]
for layer in self.layers:
if self.args.use_grad_checkpoint:
h = checkpoint(layer, h, freqs_cis, use_reentrant=False)
else:
h = layer(h, freqs_cis)
h = self.norm(h)
logits = self.output(h)
return logits
# ---------------------------------------------------------------------------
# Checkpoint loading (used by evaluate.py)
# ---------------------------------------------------------------------------
def _sanitize_config_dict(config_dict: Dict[str, Any], target_cls) -> Dict[str, Any]:
valid_keys = set(target_cls.__dataclass_fields__.keys())
return {k: v for k, v in config_dict.items() if k in valid_keys}
def _config_from_checkpoint(ckpt: Dict[str, Any]) -> tuple[str, Any]:
#Infer model variant and config from a checkpoint dict
model_variant = ckpt.get("model_variant")
config_dict = ckpt.get("config_dict")
config_obj = ckpt.get("config")
if config_dict is None and config_obj is not None and hasattr(config_obj, "__dict__"):
config_dict = dict(vars(config_obj))
if config_dict is None:
config_dict = {}
# Infer variant from config keys if not stored
if model_variant is None:
if "ffn_hidden_dim" in config_dict and "dense_ffn_hidden_dim" not in config_dict:
model_variant = "dense"
else:
model_variant = "moe"
# Handle legacy variant names
if model_variant == "dense100m":
model_variant = "dense"
if model_variant == "dense":
from model_100m import PicoLMDenseArgs
config = PicoLMDenseArgs(**_sanitize_config_dict(config_dict, PicoLMDenseArgs))
else:
config = PicoLMMoEArgs(**_sanitize_config_dict(config_dict, PicoLMMoEArgs))
return model_variant, config
def load_model(checkpoint_path: str, device: str = "cpu") -> nn.Module:
"""
Load your trained model from a checkpoint.
Args:
checkpoint_path: Path to your checkpoint.pt file
device: Device string ("cuda" or "cpu")
Returns:
A PyTorch nn.Module in eval mode where:
model(input_ids) -> logits
- input_ids: LongTensor of shape (batch_size, sequence_length)
- logits: FloatTensor of shape (batch_size, sequence_length, 50257)
"""
ckpt = torch.load(checkpoint_path, map_location=device, weights_only=False)
model_variant, config = _config_from_checkpoint(ckpt)
if model_variant == "dense":
from model_100m import PicoLMDense
model = PicoLMDense(config)
else:
model = PicoLMMoE(config)
state_dict = ckpt.get("model_state_dict", ckpt)
model.load_state_dict(state_dict, strict=True)
model.to(device)
model.eval()
return model
# Sanity check
if __name__ == "__main__":
config = load_picolm_moe_args(profile="picolm_moe")
model = PicoLMMoE(config)
n_params = sum(p.numel() for p in model.parameters())
print(f"picoLM MoE parameters: {n_params / 1e6:.2f}M")
print(f"Model successfully instantiated with {config.n_layers} layers.")