From 6a544fabfc2af22e4d5823410dd2387b5af89ea9 Mon Sep 17 00:00:00 2001 From: yurenh Date: Mon, 31 Aug 2026 18:14:09 -0500 Subject: scaffold: model (OLMo2-ish + ZBP partition), trainer (DDP/config), data shards, bench Co-Authored-By: Claude Fable 5 Claude-Session: https://claude.ai/code/session_01GkgLsACEF6CCP7EUfA5fZe --- src/zbp_scaling/data.py | 18 ++++++++++++++++++ 1 file changed, 18 insertions(+) create mode 100644 src/zbp_scaling/data.py (limited to 'src/zbp_scaling/data.py') diff --git a/src/zbp_scaling/data.py b/src/zbp_scaling/data.py new file mode 100644 index 0000000..c6946a0 --- /dev/null +++ b/src/zbp_scaling/data.py @@ -0,0 +1,18 @@ +"""Token shards: uint16 memmap files train.bin / val.bin in --data dir.""" +import os +import numpy as np +import torch + + +class Shards: + def __init__(self, path, seq_len, device): + self.train = np.memmap(os.path.join(path, "train.bin"), dtype=np.uint16, mode="r") + self.val = np.memmap(os.path.join(path, "val.bin"), dtype=np.uint16, mode="r") + self.T, self.device = seq_len, device + + def batch(self, split, bs, gen): + src = self.train if split == "train" else self.val + ix = torch.randint(0, len(src) - self.T - 1, (bs,), generator=gen) + x = torch.stack([torch.from_numpy(src[i:i + self.T].astype(np.int64)) for i in ix]) + y = torch.stack([torch.from_numpy(src[i + 1:i + 1 + self.T].astype(np.int64)) for i in ix]) + return x.to(self.device, non_blocking=True), y.to(self.device, non_blocking=True) -- cgit v1.2.3