forward-self-model-checkpoints / forward_model.py
jspr's picture
Remove closed-loop classes (CerebellarGate, LossPredictor, EmotionGate) from forward_model.py
cf648ca verified
Raw
History Blame Contribute Delete
4.01 kB
"""Forward models for predicting transformer activations.
ForwardModel: per-position MLP. Structurally blind to cross-position effects.
TransformerForwardModel: small transformer with capacity bottleneck.
"""
import math
import torch
import torch.nn as nn
import torch.nn.functional as F
class ForwardModel(nn.Module):
def __init__(self, d_model: int, hidden_mult: int = 2):
super().__init__()
hidden = d_model * hidden_mult
self.net = nn.Sequential(
nn.Linear(d_model, hidden),
nn.GELU(),
nn.Linear(hidden, d_model),
)
n_params = sum(p.numel() for p in self.parameters())
print(f"ForwardModel: {n_params/1e3:.1f}K parameters "
f"(d_model={d_model}, hidden={hidden})")
def forward(self, x):
return self.net(x)
class ForwardBlock(nn.Module):
def __init__(self, d_model: int, d_head: int, n_head: int, mlp_mult: float,
use_swiglu: bool = False):
super().__init__()
self.d_head = d_head
self.n_head = n_head
self.use_swiglu = use_swiglu
self.ln1 = nn.LayerNorm(d_model)
self.q_proj = nn.Linear(d_model, d_head * n_head)
self.k_proj = nn.Linear(d_model, d_head * n_head)
self.v_proj = nn.Linear(d_model, d_head * n_head)
self.out_proj = nn.Linear(d_head * n_head, d_model)
self.ln2 = nn.LayerNorm(d_model)
mlp_hidden = int(d_model * mlp_mult)
if use_swiglu:
self.gate_proj = nn.Linear(d_model, mlp_hidden)
self.up_proj = nn.Linear(d_model, mlp_hidden)
self.down_proj = nn.Linear(mlp_hidden, d_model)
else:
self.mlp = nn.Sequential(
nn.Linear(d_model, mlp_hidden),
nn.GELU(),
nn.Linear(mlp_hidden, d_model),
)
def forward(self, x, causal_mask):
B, T, C = x.size()
h = self.ln1(x)
q = self.q_proj(h).view(B, T, self.n_head, self.d_head).transpose(1, 2)
k = self.k_proj(h).view(B, T, self.n_head, self.d_head).transpose(1, 2)
v = self.v_proj(h).view(B, T, self.n_head, self.d_head).transpose(1, 2)
att = (q @ k.transpose(-2, -1)) * (1.0 / math.sqrt(self.d_head))
att = att.masked_fill(causal_mask[:, :, :T, :T] == 0, float("-inf"))
att = F.softmax(att, dim=-1)
y = att @ v
y = y.transpose(1, 2).contiguous().view(B, T, self.d_head * self.n_head)
x = x + self.out_proj(y)
h2 = self.ln2(x)
if self.use_swiglu:
x = x + self.down_proj(F.silu(self.gate_proj(h2)) * self.up_proj(h2))
else:
x = x + self.mlp(h2)
return x
class TransformerForwardModel(nn.Module):
def __init__(self, d_model: int, d_head: int = 64, n_head: int = 1,
n_layer: int = 1, mlp_mult: float = 2, block_size: int = 128,
causal: bool = True, use_swiglu: bool = False):
super().__init__()
self.d_model = d_model
if causal:
mask = torch.tril(torch.ones(block_size, block_size))
else:
mask = torch.ones(block_size, block_size)
self.register_buffer("causal_mask", mask.view(1, 1, block_size, block_size))
self.blocks = nn.ModuleList([
ForwardBlock(d_model, d_head, n_head, mlp_mult, use_swiglu=use_swiglu)
for _ in range(n_layer)
])
mlp_hidden = int(d_model * mlp_mult)
mlp_type = "SwiGLU" if use_swiglu else "GELU"
n_params = sum(p.numel() for p in self.parameters())
print(f"TransformerForwardModel: {n_params/1e3:.1f}K parameters "
f"(d_model={d_model}, d_head={d_head}, n_head={n_head}, "
f"n_layer={n_layer}, mlp_hidden={mlp_hidden}, mlp={mlp_type}"
f"{', bidirectional' if not causal else ''})")
def forward(self, x):
for block in self.blocks:
x = block(x, self.causal_mask)
return x