Skip to content
This repository was archived by the owner on Aug 3, 2026. It is now read-only.
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
51 changes: 30 additions & 21 deletions configs/h100_proxy.json
Original file line number Diff line number Diff line change
@@ -1,22 +1,31 @@
{
"_comment": "b300_16k_lr6: 254M H100 proxy model, 16k steps, seq_len 512.",
"vocab_size": 50257,
"dim": 1024,
"n_layers": 16,
"n_heads": 16,
"head_dim": 64,
"ffn_mult": 2.6875,
"max_seq_len": 1024,
"seq_len": 512,
"batch_size": 512,
"micro_batch_size": 128,
"total_steps": 2400,
"warmup_steps": 240,
"max_lr": 0.0008,
"min_lr": 3e-05,
"weight_decay": 0.1,
"beta1": 0.9,
"beta2": 0.98,
"grad_clip": 1.0,
"log_every": 50
}
"vocab_size": 50257,
"dim": 1024,
"n_layers": 16,
"n_heads": 16,
"head_dim": 64,
"ffn_mult": 2.6667,
"max_seq_len": 512,
"seq_len": 512,
"batch_size": 1024,
"micro_batch_size": 128,
"total_steps": 5150,
"warmup_steps": 100,
"max_lr": 0.003,
"min_lr": 3e-05,
"schedule": "wsd",
"stable_frac": 0.55,
"decay_frac": 0.45,
"decay_curve": "1-sqrt",
"optimizer": "muon",
"muon_lr": 0.025,
"muon_momentum": 0.95,
"muon_weight_decay": 0.05,
"embed_optimizer": "adamw",
"embed_lr": 0.006,
"compile": true,
"fast_kernels": true,
"init_seed": 7777,
"data_seed": 7777,
"log_every": 50
}
20 changes: 18 additions & 2 deletions data/dataset.py
Original file line number Diff line number Diff line change
Expand Up @@ -103,14 +103,30 @@ def get(self, step: int) -> tuple[torch.Tensor, torch.Tensor]:
ids = torch.from_numpy(chunk.astype(np.int64))
return ids[:-1], ids[1:]

def _perm_window_start(self, k: int) -> int:
"""Start offset of the k-th sample under epoch-wise permutation.

Non-overlapping windows, reshuffled each epoch with a (seed, epoch)-keyed
permutation: every token is seen once per epoch (uniform repeat counts),
unlike independent uniform draws which leave ~exp(-epochs) of the corpus
unseen. Same determinism contract: order is a pure function of
(manifest, seed, seq_len, k), so audit replay is unchanged.
"""
n_win = self._total // (self.seq_len + 1)
epoch, idx = divmod(k, n_win)
if getattr(self, "_perm_epoch", None) != epoch:
prng = np.random.default_rng(np.array([self.seed, 0xE90C4, epoch], dtype=np.uint64))
self._perm = prng.permutation(n_win)
self._perm_epoch = epoch
return int(self._perm[idx]) * (self.seq_len + 1)

def get_batch(
self,
step: int,
batch_size: int,
) -> tuple[torch.Tensor, torch.Tensor]:
"""Return a batch of (B, T) input + target tensors at the given step."""
rng = np.random.default_rng(np.array([self.seed, step], dtype=np.uint64))
starts = rng.integers(0, self._total, size=batch_size)
starts = [self._perm_window_start(step * batch_size + b) for b in range(batch_size)]
inputs = np.empty((batch_size, self.seq_len), dtype=np.int64)
targets = np.empty((batch_size, self.seq_len), dtype=np.int64)
for b, s in enumerate(starts):
Expand Down
4 changes: 4 additions & 0 deletions model/_v4skip.py
Original file line number Diff line number Diff line change
Expand Up @@ -38,6 +38,7 @@ class RalphConfig:
tie_embeddings: bool = True
unet_skip: bool = True # recipe-v4: U-Net learnable skip connections
logit_softcap: float = 30.0 # recipe-v4: tanh soft-cap on logits (0 = off)
logit_z_coef: float = 0.0001 # z-loss on final logits (no params; 0 = off)


def _rms_norm(x: torch.Tensor, weight: torch.Tensor, eps: float) -> torch.Tensor:
Expand Down Expand Up @@ -227,6 +228,9 @@ def forward(self, idx: torch.Tensor, targets: Optional[torch.Tensor] = None) ->
targets.view(-1),
ignore_index=-100,
)
z_coef = getattr(self.cfg, "logit_z_coef", 0.0)
if z_coef:
loss = loss + z_coef * (torch.logsumexp(logits, dim=-1).float() ** 2).mean()
return logits, loss


Expand Down
Loading