LatentTransformer
A recurrent latent transformer model trained on FineWeb-Edu. The model encodes the input into a latent representation and refines it recursively for N steps before decoding.
Checkpoints
| Variant | d_model |
Layers | n_recurse |
Params | Train time |
|---|---|---|---|---|---|
| medium | 384 | 6 | 3 | ~54M | ~2h |
| large | 512 | 8 | 4 | ~116M | ~6h |
Trained for 8,000 optimizer steps (effective batch 64) on deatos/tokenized_fineweb_edu_10b_combined using a ml.g4dn.xlarge (T4 GPU) on AWS SageMaker.
Usage
import torch
import tiktoken
from model import ModelConfig, LatentTransformer
# Load tokenizer
tokenizer = tiktoken.get_encoding("gpt2")
# --- Medium ---
cfg = ModelConfig(
vocab_size=tokenizer.n_vocab, max_seq_len=128,
d_model=384, n_heads=8, n_layers=6, n_recurse=3, dropout=0.0
)
# --- Large ---
# cfg = ModelConfig(
# vocab_size=tokenizer.n_vocab, max_seq_len=128,
# d_model=512, n_heads=8, n_layers=8, n_recurse=4, dropout=0.0
# )
model = LatentTransformer(cfg)
# Checkpoints were saved via torch.compile — strip the _orig_mod. prefix
state = torch.load("checkpoints/latent_transformer_medium.pt", map_location="cpu", weights_only=True)
if any(k.startswith("_orig_mod.") for k in state):
state = {k.removeprefix("_orig_mod."): v for k, v in state.items()}
model.load_state_dict(state)
model.eval()
# Inference
def get_causal_mask(seq_len, device):
return torch.triu(torch.ones(seq_len, seq_len, dtype=torch.bool, device=device), diagonal=1)
@torch.no_grad()
def generate(prompt, max_new_tokens=60, temperature=0.8, top_k=40, n_recurse=3):
model.cfg.n_recurse = n_recurse
tokens = tokenizer.encode(prompt)
src = torch.tensor([tokens], dtype=torch.long)
tgt = src.clone()
for _ in range(max_new_tokens):
mask = get_causal_mask(tgt.shape[1], src.device)
logits = model(src, tgt, tgt_mask=mask)
next_logits = logits[:, -1, :].squeeze(0) / temperature
if top_k > 0:
top_vals = torch.topk(next_logits, top_k)[0][..., -1, None]
next_logits[next_logits < top_vals] = float("-inf")
probs = torch.softmax(next_logits, dim=-1)
next_tok = torch.multinomial(probs, 1).unsqueeze(0)
tgt = torch.cat([tgt, next_tok], dim=1)
return tokenizer.decode(tgt[0].tolist())
print(generate("The strange truth was"))
Model Card Authors
- Michal Lajčiak