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
Downloads last month

-

Downloads are not tracked for this model. How to track
Inference Providers NEW
This model isn't deployed by any Inference Provider. 🙋 Ask for provider support