surprise-backend / models.py
khu
Integrate trained LeWM model: real inference, model architecture, checkpoint, torch deps
642eae7
Raw
History Blame Contribute Delete
1.7 kB
"""
Model architectures for the SURPRISE backend.
These match exactly what was trained in the Colab notebook
(lewm_starter.ipynb), so the saved checkpoint will load cleanly.
"""
import torch
import torch.nn as nn
from einops import rearrange
class TinyEncoder(nn.Module):
"""
Small CNN encoder. Maps (B, 3, 64, 64) β†’ (B, embed_dim).
Mirrors the Colab definition exactly β€” do not modify without retraining.
"""
def __init__(self, embed_dim: int = 64):
super().__init__()
self.embed_dim = embed_dim
self.net = nn.Sequential(
nn.Conv2d(3, 32, kernel_size=4, stride=2, padding=1), # 64 β†’ 32
nn.GELU(),
nn.Conv2d(32, 64, kernel_size=4, stride=2, padding=1), # 32 β†’ 16
nn.GELU(),
nn.Conv2d(64, 128, kernel_size=4, stride=2, padding=1), # 16 β†’ 8
nn.GELU(),
nn.AdaptiveAvgPool2d(1),
nn.Flatten(),
nn.Linear(128, embed_dim),
)
self.norm = nn.BatchNorm1d(embed_dim, affine=False)
def forward(self, x):
z = self.net(x)
z = self.norm(z)
return z
class TinyPredictor(nn.Module):
"""
Predicts the next embedding from the current one.
No actions (passive video), no temporal context beyond t-1.
"""
def __init__(self, embed_dim: int = 64, hidden_dim: int = 128):
super().__init__()
self.net = nn.Sequential(
nn.Linear(embed_dim, hidden_dim),
nn.GELU(),
nn.Linear(hidden_dim, hidden_dim),
nn.GELU(),
nn.Linear(hidden_dim, embed_dim),
)
def forward(self, z_t):
return self.net(z_t)