""" Shared fineweb-edu token cache. Every job trains on the same corpus (HuggingFaceFW/fineweb-edu), so instead of each run streaming and re-tokenizing the dataset, the Space builds ONE flat uint16 token file in a background thread at startup and all trainers read slices of it. 300M tokens x 2 bytes = 600 MB on disk, which comfortably fits Spaces' ephemeral storage and makes the GPUs data-bound never again. Documents are packed end-to-end separated by <|endoftext|> (id 0). """ import os import threading import time import numpy as np from . import store from .config import MAX_TOKENS from .db import DATA_DIR TOKENS_PATH = os.path.join(DATA_DIR, "fineweb_edu_tokens.u16") META_PATH = TOKENS_PATH + ".meta" DATASET_ID = "HuggingFaceFW/fineweb-edu" DATASET_CONFIG = "sample-350BT" # Build a little past the maximum a user can request so the last run never # stalls on the writer. At 2 bytes per token a 1B-token budget is ~2 GB on # disk, which is nothing next to the Space's ephemeral storage. TARGET_TOKENS = MAX_TOKENS + 5_000_000 _state = { "ready_tokens": 0, "status": "idle", # idle|building|ready|error "error": None, "started": None, "rate": 0.0, } _state_lock = threading.Lock() _build_thread = None def state(): with _state_lock: return dict(_state) def _set(**kw): with _state_lock: _state.update(kw) def _load_meta(): if os.path.exists(META_PATH): try: with open(META_PATH) as f: return int(f.read().strip()) except Exception: return 0 return 0 def _save_meta(n): tmp = META_PATH + ".tmp" with open(tmp, "w") as f: f.write(str(n)) os.replace(tmp, META_PATH) def build_token_cache(tokenizer, log=print): """Stream fineweb-edu, tokenize, append to the flat uint16 file.""" from datasets import load_dataset # A previous container may have already paid for this. Downloading ~2.8 GB # beats re-tokenizing 1.5B tokens by a wide margin. have = _load_meta() if have < TARGET_TOKENS: restored = store.fetch_token_cache(TOKENS_PATH, META_PATH, log) if restored > have: have = restored _set(ready_tokens=have) if have >= TARGET_TOKENS: _set(ready_tokens=have, status="ready") log(f"[data] cache already complete: {have:,} tokens") return _set(status="building", started=time.time(), ready_tokens=have) log(f"[data] building token cache ({have:,} / {TARGET_TOKENS:,})") eot = 0 # <|endoftext|> t0 = time.time() written = have # Truncate to the last known-good offset, then append. mode = "r+b" if os.path.exists(TOKENS_PATH) else "wb" f = open(TOKENS_PATH, mode) f.seek(written * 2) f.truncate() try: ds = load_dataset(DATASET_ID, name=DATASET_CONFIG, split="train", streaming=True) # Skip ahead deterministically if resuming a partial build. buf_docs = [] BATCH = 512 for row in ds: buf_docs.append(row["text"]) if len(buf_docs) < BATCH: continue encs = tokenizer.encode_batch(buf_docs) buf_docs = [] flat = [] for e in encs: flat.extend(e.ids) flat.append(eot) arr = np.asarray(flat, dtype=np.uint16) f.write(arr.tobytes()) written += arr.size if written % (1 << 22) < arr.size: # ~every 4M tokens f.flush() _save_meta(written) _touch_lock() dt = max(time.time() - t0, 1e-6) _set(ready_tokens=written, rate=(written - have) / dt) log(f"[data] {written:,} tokens " f"({(written - have) / dt / 1e6:.2f}M tok/s)") if written >= TARGET_TOKENS: break f.flush() _save_meta(written) _set(ready_tokens=written, status="ready") log(f"[data] cache ready: {written:,} tokens in {time.time() - t0:.0f}s") f.close() store.save_token_cache(TOKENS_PATH, META_PATH, log) except Exception as exc: # keep the Space alive; jobs will wait / use what exists _save_meta(written) _set(ready_tokens=written, status="error", error=repr(exc)) log(f"[data] ERROR: {exc!r} — the partial cache of {written:,} tokens is still usable") finally: f.close() LOCK_PATH = TOKENS_PATH + ".lock" LOCK_STALE_SECONDS = 150 _HEARTBEAT = 30 def _acquire_build_lock(): """ Only one process may append to the shared token file. A rolling deploy runs two containers against the same volume for a while, and two appenders would corrupt it. The lock is a file the holder keeps touching, so a crashed holder's lock goes stale instead of blocking forever. """ now = time.time() try: fd = os.open(LOCK_PATH, os.O_CREAT | os.O_EXCL | os.O_WRONLY, 0o644) os.write(fd, str(os.getpid()).encode()) os.close(fd) return True except FileExistsError: try: age = now - os.path.getmtime(LOCK_PATH) except OSError: return False if age < LOCK_STALE_SECONDS: return False # Stale: the previous holder is gone. Take it over. try: os.utime(LOCK_PATH, (now, now)) with open(LOCK_PATH, "w") as f: f.write(str(os.getpid())) return True except OSError: return False except OSError: return True # can't lock (read-only fs) — proceed anyway def _touch_lock(): try: os.utime(LOCK_PATH, None) except OSError: pass def _release_lock(): try: os.unlink(LOCK_PATH) except OSError: pass def _guarded_build(tokenizer, log): if not _acquire_build_lock(): log("[data] another process is building the token cache — " "using it read-only") _set(status="following") # Track the other process's progress so jobs know what's available. while True: have = _load_meta() _set(ready_tokens=have) if have >= TARGET_TOKENS: _set(status="ready") return try: if time.time() - os.path.getmtime(LOCK_PATH) > LOCK_STALE_SECONDS: break # holder died; take over below except OSError: break time.sleep(10) if not _acquire_build_lock(): _set(status="ready" if _load_meta() else "error") return try: build_token_cache(tokenizer, log) finally: _release_lock() def start_background_build(tokenizer, log=print): global _build_thread if _build_thread and _build_thread.is_alive(): return _build_thread _build_thread = threading.Thread( target=_guarded_build, args=(tokenizer, log), name="token-cache", daemon=True) _build_thread.start() return _build_thread def available_tokens(): with _state_lock: return _state["ready_tokens"] class TokenStream: """ Sequential reader over the shared token file, yielding (B, T) int64 batches. Each job gets its own random start offset so two runs of the same size do not see the identical token order. The batch is used as BOTH input_ids and labels: LlamaForCausalLM shifts labels internally, so handing it a pre-shifted target would train the model to predict two tokens ahead while generation samples one ahead. """ def __init__(self, seq_len, batch_size, seed=1337): self.seq_len = seq_len self.batch_size = batch_size self.chunk = batch_size * seq_len rng = np.random.default_rng(seed) avail = max(available_tokens(), self.chunk + 1) self.pos = int(rng.integers(0, max(avail - self.chunk, 1))) self._mm = None self._mm_len = 0 def _ensure(self): n = available_tokens() if self._mm is None or n > self._mm_len: if n < self.chunk + 1: return False self._mm = np.memmap(TOKENS_PATH, dtype=np.uint16, mode="r", shape=(n,)) self._mm_len = n return self._mm is not None def next_batch(self): """One (B, T) int64 batch. Pass it as input_ids AND labels.""" import torch deadline = time.time() + 300 while not self._ensure(): if time.time() > deadline: raise RuntimeError("token cache unavailable (timed out waiting for data)") time.sleep(2) if self.pos + self.chunk > self._mm_len: self.pos = 0 buf = np.asarray(self._mm[self.pos:self.pos + self.chunk], dtype=np.int64) self.pos += self.chunk return torch.from_numpy(buf.reshape(self.batch_size, self.seq_len))