GGUFGuy's picture
Duplicate from GGUFGuy/hyperdex-trainer
2a49d9a
Raw History Blame Contribute Delete
8.99 kB
"""
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))