Views
No views yet
| Component | Description |
|---|---|
| ChunkAutoencoder | Encodes fixed-size token chunks (K=4) into latent vectors and reconstructs them. |
| CALM Transformer | Learns to predict the next latent vector in sequence. |
| Energy Head (optional) | Computes latent energy scores for contrastive fine-tuning or uncertainty modeling. |
roneneldan/TinyStories dataset using the GPT-2 tokenizer.| Stage | Task | Samples | Steps | Final Loss | Hardware |
|---|---|---|---|---|---|
| Autoencoder (AE) | Reconstruction of token sequences | 100K | 6K | ↓ 0.0085 | NVIDIA L4 (24GB VRAM) |
| CALM (Latent Model) | Continuous next-latent prediction | 100K | 1 → 41K | 0.0003 – 0.01 | NVIDIA L4 (24GB VRAM) |
🚀 CALM Training (Full 1 → 41K Steps)
CALM step 100/40000 | loss 0.0015 | elapsed 0.1 m
CALM step 500/40000 | loss 0.0003 | elapsed 0.6 m
CALM step 2000/40000 | loss 0.0009 | elapsed 2.3 m
CALM step 5000/40000 | loss 0.0044 | elapsed 5.6 m
CALM step 9000/40000 | loss 0.0003 | elapsed 10.1 m
CALM step 11000/40000| loss 0.0014 | elapsed 12.4 m
✓ Saved final checkpoint at calm_41000.pt
config.json)1{
2 "model_type": "CALM",
3 "ae_hidden_dim": 512,
4 "ae_latent_dim": 256,
5 "ae_chunk_size_k": 4,
6 "lm_hidden_dim": 768,
7 "lm_ffn_dim": 3072,
8 "lm_num_layers": 12,
9 "lm_num_heads": 12,
10 "vocab_size": 50257,
11 "energy_head_dim": 128,
12 "energy_num_blocks": 2
13}ae_model.py)1import torch
2import torch.nn as nn
3import torch.nn.functional as F
4import json, os
5
6class TextAE(nn.Module):
7 """
8 Lightweight text AutoEncoder used at inference:
9 - encode(tokens) -> latent sequence z [B, T_latent, D]
10 - decode(z) -> logits over vocab for each recovered token position
11 This matches the interfaces used during training: chunk K tokens into one latent.
12 """
13 def __init__(self, config_path="config.json"):
14 super().__init__()
15 if not os.path.exists(config_path):
16 raise FileNotFoundError(f"Config not found: {config_path}")
17 with open(config_path) as f:
18 cfg = json.load(f)
19
20 self.vocab_size = cfg.get("vocab_size", 50257)
21 self.embed_dim = cfg.get("ae_embed_dim", 512)
22 self.latent_dim = cfg.get("ae_latent_dim", 256)
23 self.chunk_k = cfg.get("ae_chunk_size_k", 4)
24 self.enc_hidden = cfg.get("ae_enc_hidden", 1024)
25 self.dec_hidden = cfg.get("ae_dec_hidden", 1024)
26
27 # Token embedding for encoder side
28 self.tok_embed = nn.Embedding(self.vocab_size, self.embed_dim)
29
30 # Encoder: pooled token embeddings -> latent z
31 self.encoder = nn.Sequential(
32 nn.Linear(self.embed_dim, self.enc_hidden),
33 nn.GELU(),
34 nn.Linear(self.enc_hidden, self.latent_dim),
35 )
36
37 # Decoder: latent z -> per-token logits for K tokens
38 # We predict K token logits per latent by sharing a projector then a per-slot head
39 self.dec_proj = nn.Sequential(
40 nn.Linear(self.latent_dim, self.dec_hidden),
41 nn.GELU(),
42 )
43 # one classifier per position in the chunk
44 self.slot_heads = nn.ModuleList([
45 nn.Linear(self.dec_hidden, self.vocab_size) for _ in range(self.chunk_k)
46 ])
47
48 def _chunk_tokens(self, tok_emb: torch.Tensor):
49 """
50 tok_emb: [B, T, E]
51 returns [B, T_latent, K, E] where T_latent = ceil(T / K), right-padded if needed
52 """
53 B, T, E = tok_emb.shape
54 K = self.chunk_k
55 pad = (K - (T % K)) % K
56 if pad > 0:
57 pad_emb = torch.zeros(B, pad, E, device=tok_emb.device, dtype=tok_emb.dtype)
58 tok_emb = torch.cat([tok_emb, pad_emb], dim=1)
59 T = T + pad
60 tok_emb = tok_emb.view(B, T // K, K, E)
61 return tok_emb, pad
62
63 def encode(self, input_ids: torch.Tensor):
64 """
65 input_ids: [B, T]
66 returns:
67 z: [B, T_latent, D]
68 meta: dict with padding info
69 """
70 emb = self.tok_embed(input_ids) # [B, T, E]
71 chunked, pad = self._chunk_tokens(emb) # [B, T_lat, K, E]
72 pooled = chunked.mean(dim=2) # [B, T_lat, E]
73 z = self.encoder(pooled) # [B, T_lat, D]
74 return z, {"pad_tokens": pad}
75
76 def decode(self, z: torch.Tensor, meta=None):
77 """
78 z: [B, T_latent, D]
79 returns:
80 logits: [B, T_rec, V] where T_rec = T_latent * K
81 """
82 B, T_lat, D = z.shape
83 h = self.dec_proj(z) # [B, T_lat, H]
84 # produce K distributions per latent step
85 slot_logits = []
86 for head in self.slot_heads:
87 slot_logits.append(head(h)) # each [B, T_lat, V]
88 # interleave along the time axis: [B, T_lat, K, V] -> [B, T_lat*K, V]
89 stacked = torch.stack(slot_logits, dim=2) # [B, T_lat, K, V]
90 logits = stacked.reshape(B, T_lat * len(self.slot_heads), -1)
91 return logitscalm_model.py)1import torch
2import torch.nn as nn
3import torch.nn.functional as F
4import json, os
5
6
7class CALMBlock(nn.Module):
8 """Transformer block for CALM latent prediction."""
9 def __init__(self, hidden_dim, num_heads, ffn_dim, dropout=0.1):
10 super().__init__()
11 self.ln1 = nn.LayerNorm(hidden_dim)
12 self.attn = nn.MultiheadAttention(hidden_dim, num_heads, dropout=dropout, batch_first=True)
13 self.ln2 = nn.LayerNorm(hidden_dim)
14 self.ff = nn.Sequential(
15 nn.Linear(hidden_dim, ffn_dim),
16 nn.GELU(),
17 nn.Linear(ffn_dim, hidden_dim),
18 )
19 self.dropout = nn.Dropout(dropout)
20
21 def forward(self, x):
22 attn_out, _ = self.attn(self.ln1(x), self.ln1(x), self.ln1(x))
23 x = x + self.dropout(attn_out)
24 ff_out = self.ff(self.ln2(x))
25 x = x + self.dropout(ff_out)
26 return x
27
28
29class CALM(nn.Module):
30 """Continuous Autoregressive Language Model (CALM)."""
31 def __init__(self, config_path="config.json"):
32 super().__init__()
33 if not os.path.exists(config_path):
34 raise FileNotFoundError(f"Config file not found: {config_path}")
35 with open(config_path) as f:
36 cfg = json.load(f)
37
38 self.latent_dim = cfg.get("ae_latent_dim", 256)
39 self.hidden_dim = cfg.get("lm_hidden_dim", 768)
40 self.chunk_k = cfg.get("ae_chunk_size_k", 4)
41 self.vocab_size = cfg.get("vocab_size", 50257)
42 self.ffn_dim = cfg.get("lm_ffn_dim", 3072)
43 self.layers = cfg.get("lm_num_layers", 12)
44 self.heads = cfg.get("lm_num_heads", 12)
45 self.energy_dim = cfg.get("energy_head_dim", 128)
46 self.energy_blocks = cfg.get("energy_num_blocks", 2)
47
48 self.latent_proj = nn.Linear(self.latent_dim, self.hidden_dim)
49 self.blocks = nn.ModuleList([
50 CALMBlock(self.hidden_dim, self.heads, self.ffn_dim)
51 for _ in range(self.layers)
52 ])
53 self.ln_final = nn.LayerNorm(self.hidden_dim)
54 self.out_proj = nn.Linear(self.hidden_dim, self.latent_dim)
55 self.energy_head = nn.Sequential(
56 nn.Linear(self.latent_dim, self.energy_dim),
57 nn.ReLU(),
58 *[
59 nn.Sequential(nn.Linear(self.energy_dim, self.energy_dim), nn.ReLU())
60 for _ in range(max(self.energy_blocks - 1, 0))
61 ],
62 nn.Linear(self.energy_dim, 1)
63 )
64
65 def forward(self, z_seq):
66 x = self.latent_proj(z_seq)
67 for blk in self.blocks:
68 x = blk(x)
69 x = self.ln_final(x)
70 z_pred = self.out_proj(x[:, -1])
71 return z_pred
72
73 def energy(self, z):
74 return self.energy_head(z).mean()1import os, json, torch
2from huggingface_hub import snapshot_download
3from transformers import GPT2TokenizerFast
4from calm_model import CALM
5from ae_model import TextAE
6
7@torch.no_grad()
8def calm_generate_latents(calm, z_past, steps=32, temperature=0.0):
9 """
10 Autoregress in latent space.
11 z_past: [B, T_lat, D]
12 returns: [B, T_lat+steps, D]
13 """
14 device = next(calm.parameters()).device
15 z_seq = z_past.clone()
16 for _ in range(steps):
17 z_next = calm(z_seq) # [B, D]
18 if temperature and temperature > 0:
19 # optional Gaussian noise in latent space
20 z_next = z_next + torch.randn_like(z_next) * temperature
21 z_seq = torch.cat([z_seq, z_next.unsqueeze(1)], dim=1)
22 return z_seq
23
24@torch.no_grad()
25def ae_decode_tokens(ae, logits, pad_tokens=0, trim_to_multiple_of_k=True):
26 """
27 Convert AE decoder logits to token ids and drop right padding added during encode.
28 logits: [B, T_rec, V]
29 """
30 token_ids = logits.argmax(dim=-1) # [B, T_rec]
31 if trim_to_multiple_of_k and pad_tokens > 0:
32 token_ids = token_ids[:, :-pad_tokens]
33 return token_ids
34
35def main():
36 repo_id = "yasserrmd/CALM-TinyStories-v1"
37 local_dir = snapshot_download(repo_id=repo_id)
38
39 # Load config
40 cfg_path = os.path.join(local_dir, "config.json")
41 with open(cfg_path) as f:
42 cfg = json.load(f)
43 chunk_k = cfg.get("ae_chunk_size_k", 4)
44 vocab = cfg.get("vocab_size", 50257)
45
46 # Tokenizer (GPT2-compatible by default)
47 tokenizer = GPT2TokenizerFast.from_pretrained("gpt2")
48 # add pad token if missing
49 if tokenizer.pad_token_id is None:
50 tokenizer.pad_token = tokenizer.eos_token
51 assert tokenizer.vocab_size == vocab or True, "Tokenizer vocab may differ; ensure config matches training."
52
53 device = "cuda" if torch.cuda.is_available() else "cpu"
54
55 # Load AE
56 ae = TextAE(config_path=cfg_path).to(device)
57 ae_ckpt = torch.load(os.path.join(local_dir, "ae_final.pt"), map_location=device)
58 ae.load_state_dict(ae_ckpt, strict=True)
59 ae.eval()
60
61 # Load CALM
62 calm = CALM(config_path=cfg_path).to(device)
63 calm_ckpt = torch.load(os.path.join(local_dir, "calm_final.pt"), map_location=device)
64 calm.load_state_dict(calm_ckpt, strict=True)
65 calm.eval()
66
67 print("✓ Loaded AE and CALM")
68
69 # Example prompt -> latents
70 prompt = "Once upon a time"
71 enc = tokenizer(prompt, return_tensors="pt")
72 input_ids = enc["input_ids"].to(device) # [1, T]
73
74 # AE encode to latents
75 z_init, meta = ae.encode(input_ids) # [1, T_lat, D], meta['pad_tokens']
76 # Autoregress in latent space
77 z_full = calm_generate_latents(calm, z_init, steps=32, temperature=0.0) # extend by 32 latent steps
78 # Decode latents back to tokens
79 dec_logits = ae.decode(z_full, meta) # [1, T_rec, V] where T_rec = T_lat_total * K
80 out_ids = ae_decode_tokens(ae, dec_logits, pad_tokens=meta.get("pad_tokens", 0))
81
82 text = tokenizer.decode(out_ids[0], skip_special_tokens=True)
83 print("\n=== Generated Text ===\n", text)
84
85if __name__ == "__main__":
86 main()
87roneneldan/TinyStoriesAutoTokenizer.from_pretrained("gpt2"))
---