| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| """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") |
|
|
| |
| 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)) |
|
|
| |
| 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}") |
| |
| |
| for i in range(len(model.core.blocks)): |
| model.core.blocks[i] = torch.compile(model.core.blocks[i], **kw) |
|
|
| |
| 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 = 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, |
| ) |
|
|
| |
| 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: |
| log(f"held-out load failed ({e}); training without periodic eval") |
|
|
| |
| 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 |
|
|
| |
| 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: |
| 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()} |
| |
| 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})") |
|
|
| |
| 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: |
| |
| |
| 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: |
| |
| |
| 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() |
|
|