trident / scripts /train.py
farguney's picture
trainer: parse rho_min env
bb1e221 verified
Raw
History Blame Contribute Delete
13.7 kB
# /// script
# requires-python = ">=3.10"
# dependencies = [
# "torch",
# "numpy",
# "datasets>=2.19",
# "huggingface_hub>=0.24",
# "safetensors>=0.4",
# "trackio",
# ]
# ///
"""Trident pretraining on Hugging Face Jobs (GPU).
Training runs on an accelerator (trident.md Section 1.2.1); the resulting
parameters define a CPU-only inference model after parity. This script:
* pulls the `trident` package from the code repo,
* streams a real text dataset as raw bytes,
* trains the variable-rate byte model with AdamW + cosine schedule,
* evaluates bits-per-byte on a fixed held-out set,
* pushes checkpoints + metrics to the Hub (the Jobs FS is ephemeral).
All knobs are environment variables so the same script serves smoke and full
runs. Nothing here touches Jobs or repos it did not create.
"""
import json
import math
import os
import sys
import time
from pathlib import Path
import torch
def env(k, d=None):
v = os.environ.get(k)
return v if v is not None and v != "" else d
def env_i(k, d):
return int(env(k, d))
def env_f(k, d):
return float(env(k, d))
def log(msg):
print(f"[trident] {msg}", flush=True)
def main():
code_repo = env("CODE_REPO", "farguney/trident")
code_rev = env("CODE_REVISION", "main")
run_name = env("RUN_NAME", "trident-run")
# ---- fetch the trident package from the code repo ----
from huggingface_hub import HfApi, snapshot_download
workdir = Path("/tmp/trident_code")
log(f"downloading {code_repo}@{code_rev} (src/**)")
snapshot_download(
repo_id=code_repo, revision=code_rev, repo_type="model",
allow_patterns=["src/**"], local_dir=str(workdir),
)
sys.path.insert(0, str(workdir / "src"))
from trident import Trident, TridentConfig
from trident.data import (
ByteWindowIterable,
ContiguousByteStreams,
collate,
load_fixed_byte_windows,
)
device = "cuda" if torch.cuda.is_available() else "cpu"
if device == "cuda":
log(f"gpu: {torch.cuda.get_device_name(0)} torch {torch.__version__}")
else:
log("WARNING: no CUDA device; running on CPU")
torch.manual_seed(env_i("SEED", 0))
# ---- config ----
profile = env("PROFILE", "scout")
over = {}
for key, cast in [
("d_model", int), ("n_blocks", int), ("n_heads", int), ("d_k", int),
("d_v", int), ("n_dec_layers", int), ("r_max", int), ("fixed_r", int),
("b_max", int), ("tau_b", float), ("max_patches", int), ("exact_ring", int),
("d_code", int), ("lambda_fast_head", float), ("chunk_size", int),
("rho_min", float),
]:
val = env(key.upper())
if val is not None:
over[key] = cast(val)
if env("FIXED_R", "") == "none":
over["fixed_r"] = None
over["grad_checkpoint"] = env_i("GRAD_CKPT", 1) == 1
if profile in ("micro", "scout"):
cfg = getattr(TridentConfig, profile)(**over)
else:
over.setdefault("profile", profile)
cfg = TridentConfig(**over)
model = Trident(cfg).to(device)
nparams = sum(p.numel() for p in model.parameters())
log(f"config {profile} params={nparams/1e6:.2f}M identity={cfg.identity_hash()[:12]}")
log(f"state bytes/session (fp32) = {cfg.state_bytes_per_session/1024:.1f} KiB")
if env_i("COMPILE", 0) == 1 and device == "cuda":
cm = env("COMPILE_MODE", "default")
kw = {} if cm == "default" else {"mode": cm}
log(f"torch.compile(blocks) mode={cm}")
# Compile the per-block scan (the launch-bound hot path). All blocks share
# code, so this compiles ~once and is reused across the stack.
for i in range(len(model.core.blocks)):
model.core.blocks[i] = torch.compile(model.core.blocks[i], **kw)
# ---- data ----
seq_len = env_i("SEQ_LEN", 2048)
batch = env_i("BATCH", 8)
grad_accum = env_i("GRAD_ACCUM", 4)
dataset = env("DATASET", "HuggingFaceFW/fineweb-edu")
ds_name = env("DATASET_NAME", "sample-10BT")
split = env("SPLIT", "train")
text_field = env("TEXT_FIELD", "text")
# TBPTT_WINDOWS>0 => stateful truncated-BPTT training (state carried across
# K consecutive windows, so the model learns to use the state beyond one
# window). 0 => independent shuffled windows (state reset each step).
tbptt = env_i("TBPTT_WINDOWS", 0)
loader = None
streams = None
if tbptt > 0:
streams = ContiguousByteStreams(
dataset=dataset, split=split, seq_len=seq_len, B=batch, text_field=text_field,
name=ds_name, shuffle_buffer=env_i("SHUFFLE_BUFFER", 10000), seed=env_i("SEED", 0),
)
else:
train_ds = ByteWindowIterable(
dataset=dataset, split=split, seq_len=seq_len, text_field=text_field,
name=ds_name, shuffle_buffer=env_i("SHUFFLE_BUFFER", 10000), seed=env_i("SEED", 0),
)
loader = torch.utils.data.DataLoader(
train_ds, batch_size=batch, collate_fn=collate,
num_workers=env_i("NUM_WORKERS", 2), drop_last=True,
)
# fixed held-out set (identical bytes for baseline comparison)
val_windows = None
try:
val_windows = load_fixed_byte_windows(
dataset=env("VAL_DATASET", "Salesforce/wikitext"),
split=env("VAL_SPLIT", "validation"),
name=env("VAL_NAME", "wikitext-103-raw-v1"),
text_field=env("VAL_TEXT_FIELD", "text"),
seq_len=seq_len, num_windows=env_i("VAL_WINDOWS", 64), add_eod=False,
).to(device)
log(f"held-out windows: {tuple(val_windows.shape)}")
except Exception as e: # noqa
log(f"held-out load failed ({e}); training without periodic eval")
# ---- optimizer + schedule ----
max_steps = env_i("MAX_STEPS", 20000)
warmup = env_i("WARMUP", 500)
lr = env_f("LR", 6e-4)
min_lr = env_f("MIN_LR", 6e-5)
wd = env_f("WEIGHT_DECAY", 0.1)
grad_clip = env_f("GRAD_CLIP", 1.0)
decay, no_decay = [], []
for n, p in model.named_parameters():
if p.ndim >= 2:
decay.append(p)
else:
no_decay.append(p)
opt = torch.optim.AdamW(
[{"params": decay, "weight_decay": wd}, {"params": no_decay, "weight_decay": 0.0}],
lr=lr, betas=(0.9, 0.95), eps=1e-8,
)
def lr_at(step):
if step < warmup:
return lr * (step + 1) / warmup
t = (step - warmup) / max(1, max_steps - warmup)
t = min(1.0, t)
return min_lr + 0.5 * (lr - min_lr) * (1 + math.cos(math.pi * t))
amp_dtype = torch.bfloat16 if device == "cuda" else torch.float32
# ---- trackio (best-effort) ----
tracker = None
try:
import trackio
user = HfApi().whoami().get("name", "user")
trackio.init(project="trident", name=run_name,
space_id=f"{user}/trackio",
config={**cfg.to_dict(), "seq_len": seq_len, "batch": batch,
"grad_accum": grad_accum, "lr": lr, "max_steps": max_steps,
"dataset": dataset, "params_M": nparams / 1e6})
tracker = trackio
log("trackio initialised")
except Exception as e: # noqa
log(f"trackio unavailable ({e})")
api = HfApi()
@torch.no_grad()
def evaluate():
model.eval()
tot_nll, tot_bytes = 0.0, 0
vb = env_i("VAL_BATCH", 8)
for i in range(0, val_windows.shape[0], vb):
chunk = val_windows[i:i + vb]
with torch.autocast(device_type="cuda", dtype=amp_dtype, enabled=device == "cuda"):
out = model(chunk, return_logits=True)
logits = out["logits"].float()
tgt = chunk.clamp(0, logits.shape[-1] - 1)
nll = torch.nn.functional.cross_entropy(
logits[:, :].reshape(-1, logits.shape[-1]), tgt.reshape(-1), reduction="sum")
tot_nll += nll.item()
tot_bytes += tgt.numel()
model.train()
return tot_nll / tot_bytes / math.log(2)
ckpt_dir = env("CKPT_DIR", "checkpoints")
def save_checkpoint(tag, step, metrics):
from safetensors.torch import save_file
outdir = Path("/tmp/ckpt")
outdir.mkdir(exist_ok=True)
sd = {k: v.detach().cpu().contiguous() for k, v in model.state_dict().items()}
# torch.compile wraps params with a prefix; strip it for portability
sd = {k.replace("_orig_mod.", ""): v for k, v in sd.items()}
save_file(sd, str(outdir / "model.safetensors"))
(outdir / "config.json").write_text(json.dumps(cfg.to_dict(), indent=2))
(outdir / "metrics.json").write_text(json.dumps(metrics, indent=2))
api.upload_folder(folder_path=str(outdir), repo_id=code_repo,
path_in_repo=f"{ckpt_dir}/{tag}", repo_type="model",
commit_message=f"checkpoint {ckpt_dir}/{tag} @ step {step}")
log(f"pushed checkpoint '{ckpt_dir}/{tag}' (step {step})")
# ---- training loop ----
eff_mult = tbptt if tbptt > 0 else grad_accum
bytes_per_step = batch * eff_mult * seq_len
mode = f"TBPTT(K={tbptt})" if tbptt > 0 else f"accum={grad_accum}"
log(f"training: max_steps={max_steps} batch={batch} {mode} seq_len={seq_len} "
f"eff_bytes/step={bytes_per_step}")
model.train()
data_iter = iter(loader) if loader is not None else None
tbptt_state = None
skipped = 0
step = 0
t0 = time.time()
t_win = time.time()
running = 0.0
save_every = env_i("SAVE_EVERY", 1000)
eval_every = env_i("EVAL_EVERY", 500)
log_every = env_i("LOG_EVERY", 20)
while step < max_steps:
for g in opt.param_groups:
g["lr"] = lr_at(step)
opt.zero_grad(set_to_none=True)
loss_val = 0.0
last = {}
if tbptt > 0:
# stateful truncated-BPTT: carry state (with grad) across K windows,
# backprop once, then detach so the state persists but grad is bounded.
losses = []
for _ in range(tbptt):
win = streams.next_window().to(device, non_blocking=True)
with torch.autocast(device_type="cuda", dtype=amp_dtype, enabled=device == "cuda"):
out = model(win, state0=tbptt_state, return_state=True)
loss = out["loss"] / tbptt
losses.append(loss)
tbptt_state = out["state"]
last = out
torch.stack(losses).sum().backward()
loss_val = float(sum(l.item() for l in losses))
tbptt_state = [s.detach() for s in tbptt_state]
else:
for _ in range(grad_accum):
try:
batch_data = next(data_iter)
except StopIteration:
data_iter = iter(loader)
batch_data = next(data_iter)
bytes_in = batch_data["bytes_in"].to(device, non_blocking=True)
valid = batch_data["valid"].to(device, non_blocking=True)
with torch.autocast(device_type="cuda", dtype=amp_dtype, enabled=device == "cuda"):
out = model(bytes_in, valid=valid)
loss = out["loss"] / grad_accum
loss.backward()
loss_val += loss.item()
last = out
gnorm = torch.nn.utils.clip_grad_norm_(model.parameters(), grad_clip)
if torch.isfinite(gnorm):
opt.step()
else:
# never apply a non-finite update; drop this batch and reset any
# carried state so a single anomaly cannot poison the whole run.
opt.zero_grad(set_to_none=True)
tbptt_state = None
skipped += 1
log(f"step {step}: non-finite grad_norm={gnorm}; skipped update ({skipped} total)")
running += loss_val
step += 1
if step % log_every == 0:
dt = time.time() - t_win
bps = log_every * bytes_per_step / dt
t_win = time.time()
avg = running / log_every
running = 0.0
log(f"step {step}/{max_steps} loss {avg:.4f} bpb {last['bpb'].item():.4f} "
f"patch {last['mean_patch_len'].item():.2f} lr {lr_at(step):.2e} "
f"gnorm {gnorm:.2f} {bps/1e3:.1f} kB/s")
if tracker:
try:
tracker.log({"train/loss": avg, "train/bpb": last["bpb"].item(),
"train/mean_patch_len": last["mean_patch_len"].item(),
"train/lr": lr_at(step), "train/grad_norm": float(gnorm),
"throughput/bytes_per_s": bps}, step=step)
except Exception:
pass
if val_windows is not None and step % eval_every == 0:
vbpb = evaluate()
log(f"[eval] step {step} val_bpb {vbpb:.4f}")
if tracker:
try:
tracker.log({"val/bpb": vbpb}, step=step)
except Exception:
pass
if step % save_every == 0:
m = {"step": step, "train_loss": loss_val, "params_M": nparams / 1e6}
if val_windows is not None:
m["val_bpb"] = evaluate()
save_checkpoint("latest", step, m)
final_bpb = evaluate() if val_windows is not None else None
save_checkpoint("final", step, {"step": step, "val_bpb": final_bpb, "params_M": nparams / 1e6})
log(f"DONE. final val_bpb={final_bpb}")
if tracker:
try:
tracker.finish()
except Exception:
pass
if __name__ == "__main__":
main()