cascade-1b-logic / train_cascade_1b.py
Asilarkness's picture
Upload train_cascade_1b.py with huggingface_hub
b30472d verified
Raw
History Blame Contribute Delete
20.9 kB
from __future__ import annotations
import argparse
import io
import json
import math
import os
import queue
import random
import threading
import time
from dataclasses import dataclass
from pathlib import Path
from typing import Any
import numpy as np
import torch
from datasets import load_dataset
from huggingface_hub import HfApi
from safetensors.torch import load_file, save_file
from torch.nn import functional as F
from cascade_model import CascadeConfig, CascadeForCausalLM
@dataclass(frozen=True)
class Source:
repo: str
config: str | None
weight: float
kind: str
SOURCES = (
Source("HuggingFaceFW/fineweb-edu", "sample-10BT", 0.58, "text"),
Source("HuggingFaceTB/finemath", "finemath-4plus", 0.22, "text"),
Source("HuggingFaceTB/smollm-corpus", "cosmopedia-v2", 0.10, "text"),
Source("open-r1/OpenR1-Math-220k", "default", 0.07, "math"),
Source("open-r1/codeforces-cots", "solutions_w_editorials_py_decontaminated", 0.03, "code"),
)
def arguments() -> argparse.Namespace:
p = argparse.ArgumentParser()
p.add_argument("--repo", default="Asilarkness/cascade-1b-logic")
p.add_argument("--token-file", type=Path, default=Path("/marimo/hf.txt"))
p.add_argument("--output", type=Path, default=Path("/marimo/storage/cascade-1b"))
p.add_argument("--batch-size", type=int, default=128)
p.add_argument("--sequence-length", type=int, default=1024)
p.add_argument("--max-patches", type=int, default=192)
p.add_argument("--patch-rate", type=float, default=0.125)
p.add_argument("--max-steps", type=int, default=100000)
p.add_argument("--warmup-steps", type=int, default=1000)
p.add_argument("--adam-lr", type=float, default=3e-4)
p.add_argument("--muon-lr", type=float, default=0.02)
p.add_argument("--weight-decay", type=float, default=0.1)
p.add_argument("--log-every", type=int, default=10)
p.add_argument("--eval-every", type=int, default=1000)
p.add_argument("--checkpoint-every", type=int, default=1000)
p.add_argument("--first-checkpoint", type=int, default=20)
p.add_argument("--state-upload-every", type=int, default=100)
p.add_argument("--bootstrap-mib", type=int, default=8)
p.add_argument("--validation-batches", type=int, default=4)
p.add_argument("--seed", type=int, default=3407)
p.add_argument("--compile", action="store_true")
p.add_argument("--compile-blocks", action="store_true")
p.add_argument("--resume", action="store_true")
p.add_argument("--probe-only", action="store_true")
p.add_argument("--probe-batches", default="64,96,128,160,192,224,256")
p.add_argument("--no-publish", action="store_true")
return p.parse_args()
def flatten_text(value: Any) -> list[str]:
if isinstance(value, str):
return [value]
if isinstance(value, dict):
if isinstance(value.get("content"), str):
return [value["content"]]
out: list[str] = []
for item in value.values():
out.extend(flatten_text(item))
return out
if isinstance(value, (list, tuple)):
out = []
for item in value:
out.extend(flatten_text(item))
return out
return []
def document(row: dict[str, Any], kind: str) -> str:
if kind == "text":
return str(row.get("text") or "")
if kind == "math":
problem = str(row.get("problem") or "")
solution = str(row.get("solution") or "")
if not solution:
values = flatten_text(row.get("generations", []))
solution = values[0] if values else ""
answer = str(row.get("answer") or "")
return f"Problem:\n{problem}\n\nSolution:\n{solution}\n\nAnswer:\n{answer}"
keys = (
"description",
"input_format",
"output_format",
"editorial",
"prompt",
"generation",
"messages",
"accepted_solutions",
"solutions",
"solution",
"code",
)
values: list[str] = []
for key in keys:
values.extend(flatten_text(row.get(key)))
return "\n\n".join(dict.fromkeys(x for x in values if len(x.strip()) > 20))
class SourceReader:
def __init__(self, source: Source, token: str, seed: int) -> None:
self.source = source
self.token = token
self.seed = seed
self.queue: queue.Queue[str] = queue.Queue(32)
self.errors: queue.Queue[str] = queue.Queue(8)
self.thread = threading.Thread(target=self.run, daemon=True)
self.thread.start()
def run(self) -> None:
attempt = 0
while True:
try:
data = load_dataset(
self.source.repo,
self.source.config,
split="train",
streaming=True,
token=self.token,
).shuffle(seed=self.seed + attempt, buffer_size=64)
for row in data:
text = document(row, self.source.kind).strip()
if len(text) >= 80:
self.queue.put(text[:4_000_000])
attempt += 1
except Exception as error:
if not self.errors.full():
self.errors.put(f"{type(error).__name__}: {str(error)[:240]}")
attempt += 1
time.sleep(min(30, 2**min(attempt, 5)))
class StreamPacker:
def __init__(self, batch_size: int, sequence_length: int, token: str, seed: int) -> None:
self.batch_size = batch_size
self.sequence_length = sequence_length
self.rng = random.Random(seed)
self.readers = [SourceReader(source, token, seed + i * 1009) for i, source in enumerate(SOURCES)]
self.buffer = bytearray()
self.offset = 0
self.documents = {source.repo + ":" + str(source.config): 0 for source in SOURCES}
self.bytes = {source.repo + ":" + str(source.config): 0 for source in SOURCES}
def wait_essential(self, timeout: float = 300) -> None:
deadline = time.time() + timeout
while time.time() < deadline:
if not self.readers[0].queue.empty() and not self.readers[1].queue.empty():
return
time.sleep(0.1)
errors = {r.source.repo: list(r.errors.queue) for r in self.readers[:2]}
raise TimeoutError(json.dumps(errors))
def get_document(self) -> bytes:
deadline = time.time() + 300
while time.time() < deadline:
ready = [(reader, reader.source.weight) for reader in self.readers if not reader.queue.empty()]
if ready:
reader = min(
(x[0] for x in ready),
key=lambda item: self.bytes[item.source.repo + ":" + str(item.source.config)] / item.source.weight,
)
try:
text = reader.queue.get_nowait()
except queue.Empty:
continue
key = reader.source.repo + ":" + str(reader.source.config)
self.documents[key] += 1
payload = (text + "\n\n<|endoftext|>\n\n").encode("utf-8", errors="replace")[:65536]
self.bytes[key] += len(payload)
return payload
time.sleep(0.05)
errors = {r.source.repo: list(r.errors.queue) for r in self.readers}
raise TimeoutError(json.dumps(errors))
def next(self) -> torch.Tensor:
needed = self.batch_size * (self.sequence_length + 1)
while len(self.buffer) - self.offset < needed:
self.buffer.extend(self.get_document())
chunk = np.frombuffer(self.buffer, dtype=np.uint8, count=needed, offset=self.offset).copy()
self.offset += needed
if self.offset >= 16_000_000:
del self.buffer[: self.offset]
self.offset = 0
return torch.from_numpy(chunk).reshape(self.batch_size, self.sequence_length + 1).pin_memory()
def entropy_bootstrap(packer: StreamPacker, target_mib: int, patch_rate: float) -> tuple[torch.Tensor, float, list[torch.Tensor]]:
counts = np.zeros(256 * 256, dtype=np.int64)
batches: list[torch.Tensor] = []
size = 0
target = target_mib * 2**20
while size < target:
batch = packer.next()
batches.append(batch)
values = batch.numpy().reshape(-1).astype(np.int64)
counts += np.bincount(values[:-1] * 256 + values[1:], minlength=256 * 256)
size += values.size
print(json.dumps({"bootstrap_mib": round(size / 2**20, 2)}), flush=True)
matrix = counts.reshape(256, 256).astype(np.float64)
rows = matrix.sum(1, keepdims=True)
surprise = -np.log((matrix + 0.1) / (rows + 25.6))
flat_count = matrix.reshape(-1)
flat_score = surprise.reshape(-1)
order = np.argsort(flat_score)[::-1]
cumulative = np.cumsum(flat_count[order])
index = min(np.searchsorted(cumulative, flat_count.sum() * patch_rate), len(order) - 1)
threshold = float(flat_score[order[index]])
return torch.from_numpy(surprise).float(), threshold, batches
class Optimizers:
def __init__(self, model: CascadeForCausalLM, args: argparse.Namespace) -> None:
matrices: list[torch.nn.Parameter] = []
other: list[torch.nn.Parameter] = []
seen: set[int] = set()
for name, parameter in model.named_parameters():
if id(parameter) in seen:
continue
seen.add(id(parameter))
excluded = "embed" in name or "head" in name
(matrices if parameter.ndim == 2 and not excluded else other).append(parameter)
self.muon = torch.optim.Muon(
matrices,
lr=args.muon_lr,
weight_decay=args.weight_decay,
momentum=0.95,
nesterov=True,
ns_steps=5,
adjust_lr_fn="match_rms_adamw",
)
self.adam = torch.optim.AdamW(
other,
lr=args.adam_lr,
betas=(0.9, 0.95),
eps=1e-8,
weight_decay=args.weight_decay,
fused=True,
)
self.base_muon = args.muon_lr
self.base_adam = args.adam_lr
def zero_grad(self) -> None:
self.muon.zero_grad(set_to_none=True)
self.adam.zero_grad(set_to_none=True)
def step(self) -> None:
self.muon.step()
self.adam.step()
def set_scale(self, scale: float) -> None:
for group in self.muon.param_groups:
group["lr"] = self.base_muon * scale
for group in self.adam.param_groups:
group["lr"] = self.base_adam * scale
def lr_scale(step: int, args: argparse.Namespace) -> float:
if step < args.warmup_steps:
return (step + 1) / max(1, args.warmup_steps)
progress = (step - args.warmup_steps) / max(1, args.max_steps - args.warmup_steps)
return 0.1 + 0.9 * 0.5 * (1 + math.cos(math.pi * min(1.0, progress)))
class Hub:
def __init__(self, repo: str, token: str, output: Path) -> None:
self.repo = repo
self.api = HfApi(token=token)
self.output = output
self.thread: threading.Thread | None = None
def upload_sources(self, files: list[Path]) -> None:
for path in files:
if path.exists():
target = "README.md" if path.name == "MODEL_CARD.md" else path.name
self.api.upload_file(path_or_fileobj=path, path_in_repo=target, repo_id=self.repo)
def state(self, values: dict[str, Any]) -> None:
payload = io.BytesIO(json.dumps(values, indent=2).encode())
self.api.upload_file(path_or_fileobj=payload, path_in_repo="training_state.json", repo_id=self.repo)
def checkpoint(self, step: int, model: CascadeForCausalLM, config: CascadeConfig, metrics: dict[str, Any]) -> bool:
if self.thread is not None and self.thread.is_alive():
return False
directory = self.output / "checkpoint-latest"
directory.mkdir(parents=True, exist_ok=True)
state = {
key: value.detach().cpu().contiguous()
for key, value in model.state_dict().items()
if key != "head.weight"
}
config.save(directory / "config.json")
(directory / "training_state.json").write_text(json.dumps(metrics, indent=2), encoding="utf-8")
def save_and_upload() -> None:
save_file(state, directory / "model.safetensors", metadata={"format": "pt", "step": str(step)})
for attempt in range(5):
try:
self.api.upload_folder(
folder_path=directory,
path_in_repo=f"checkpoints/step-{step:08d}",
repo_id=self.repo,
commit_message=f"checkpoint step {step}",
)
return
except Exception:
if attempt == 4:
raise
time.sleep(15 * (attempt + 1))
self.thread = threading.Thread(target=save_and_upload, daemon=True)
self.thread.start()
return True
def wait(self) -> None:
if self.thread is not None:
self.thread.join()
@torch.no_grad()
def evaluate(model: CascadeForCausalLM, batches: list[torch.Tensor], routes: int) -> float:
model.eval()
losses = []
for route in range(routes):
for batch in batches:
raw = batch.cuda(non_blocking=True).long()
logits, _, _ = model(raw[:, :-1], route)
losses.append(float(F.cross_entropy(logits.flatten(0, 1), raw[:, 1:].flatten())))
model.train()
return sum(losses) / len(losses)
def load_resume(model: CascadeForCausalLM, output: Path) -> int:
directory = output / "checkpoint-latest"
weights = directory / "model.safetensors"
state_path = directory / "training_state.json"
if not weights.exists() or not state_path.exists():
return 0
missing, unexpected = model.load_state_dict(load_file(weights), strict=False)
if unexpected or any(x != "head.weight" for x in missing):
raise RuntimeError({"missing": missing, "unexpected": unexpected})
return int(json.loads(state_path.read_text())["step"])
def probe(args: argparse.Namespace) -> None:
config = CascadeConfig(sequence_length=args.sequence_length)
model = CascadeForCausalLM(config).cuda().to(torch.bfloat16)
print(json.dumps(model.parameter_counts()), flush=True)
for batch_size in map(int, args.probe_batches.split(",")):
try:
torch.cuda.empty_cache()
torch.cuda.reset_peak_memory_stats()
x = torch.randint(0, 256, (batch_size, args.sequence_length), device="cuda")
started = time.perf_counter()
logits, auxiliary, _ = model(x, 0)
loss = F.cross_entropy(logits.flatten(0, 1), x.flatten()) + config.patch_aux_weight * auxiliary
loss.backward()
torch.cuda.synchronize()
elapsed = time.perf_counter() - started
print(json.dumps({"batch": batch_size, "seconds": elapsed, "tokens_s": batch_size * args.sequence_length / elapsed, "vram_gib": torch.cuda.max_memory_allocated() / 2**30}), flush=True)
model.zero_grad(set_to_none=True)
del x, logits, loss, auxiliary
except torch.cuda.OutOfMemoryError:
print(json.dumps({"batch": batch_size, "oom": True}), flush=True)
model.zero_grad(set_to_none=True)
torch.cuda.empty_cache()
break
def train(args: argparse.Namespace) -> None:
token = args.token_file.read_text().strip()
args.output.mkdir(parents=True, exist_ok=True)
torch.manual_seed(args.seed)
random.seed(args.seed)
np.random.seed(args.seed)
torch.set_float32_matmul_precision("high")
packer = StreamPacker(args.batch_size, args.sequence_length, token, args.seed)
packer.wait_essential()
surprise, threshold, bootstrap = entropy_bootstrap(packer, args.bootstrap_mib, args.patch_rate)
validation = [x[: min(8, args.batch_size)] for x in bootstrap[: args.validation_batches]]
warm = bootstrap[args.validation_batches :]
config = CascadeConfig(
sequence_length=args.sequence_length,
max_patches=args.max_patches,
patch_rate=args.patch_rate,
entropy_threshold=threshold,
)
model = CascadeForCausalLM(config, surprise).cuda().to(torch.bfloat16)
start_step = load_resume(model, args.output) if args.resume else 0
if start_step:
config.entropy_threshold = float(model.entropy_threshold)
counts = model.parameter_counts()
hub = None if args.no_publish else Hub(args.repo, token, args.output)
here = Path(__file__).resolve().parent
if hub is not None:
hub.upload_sources([here / "cascade_model.py", Path(__file__).resolve(), here / "MODEL_CARD.md"])
config.save(args.output / "config.json")
if hub is not None:
hub.api.upload_file(path_or_fileobj=args.output / "config.json", path_in_repo="config.json", repo_id=args.repo)
optimizers = Optimizers(model, args)
initial_validation = evaluate(model, validation, config.route_paths)
print(json.dumps({"parameters": counts, "threshold": threshold, "initial_validation": initial_validation, "sources": packer.documents}), flush=True)
if args.compile_blocks:
model.compile_blocks()
compiled = torch.compile(model, mode="max-autotune-no-cudagraphs") if args.compile else model
history = args.output / "training.jsonl"
tokens = start_step * args.batch_size * args.sequence_length
started = time.perf_counter()
last_time = started
last_tokens = tokens
last_metrics: dict[str, Any] = {"step": start_step, "tokens": tokens, "validation_loss": initial_validation}
for step in range(start_step + 1, args.max_steps + 1):
batch = warm.pop(0) if warm else packer.next()
raw = batch.cuda(non_blocking=True).long()
x, y = raw[:, :-1], raw[:, 1:]
route = (step - 1) % config.route_paths
scale = lr_scale(step - 1, args)
optimizers.set_scale(scale)
optimizers.zero_grad()
logits, auxiliary, stats = compiled(x, route)
language = F.cross_entropy(logits.flatten(0, 1), y.flatten())
loss = language + config.patch_aux_weight * auxiliary
loss.backward()
norm = torch.nn.utils.clip_grad_norm_(model.parameters(), 1.0)
optimizers.step()
tokens += x.numel()
if step == 1 or step % args.log_every == 0:
torch.cuda.synchronize()
now = time.perf_counter()
metrics = {
"step": step,
"tokens": tokens,
"parameters": counts,
"batch_size": args.batch_size,
"sequence_length": args.sequence_length,
"loss": float(language.detach()),
"auxiliary_loss": float(auxiliary.detach()),
"gradient_norm": float(norm.detach()),
"route": route,
"lr_scale": scale,
"tokens_per_second": (tokens - last_tokens) / (now - last_time),
"elapsed_seconds": now - started,
"average_patch_bytes": x.numel() / float(stats["patches"]),
"patch_coalesced_rate": float(stats["overflow"]) / x.shape[0],
"gate": float(stats["gate"].detach()),
"peak_vram_gib": torch.cuda.max_memory_allocated() / 2**30,
"documents": dict(packer.documents),
"source_bytes": dict(packer.bytes),
}
last_metrics = metrics
with history.open("a", encoding="utf-8") as file:
file.write(json.dumps(metrics) + "\n")
print(json.dumps(metrics), flush=True)
last_time, last_tokens = now, tokens
if step % args.eval_every == 0:
last_metrics["validation_loss"] = evaluate(compiled, validation, config.route_paths)
print(json.dumps({"step": step, "validation_loss": last_metrics["validation_loss"]}), flush=True)
if hub is not None and (step == args.first_checkpoint or step % args.checkpoint_every == 0):
last_metrics["checkpoint_started"] = hub.checkpoint(step, model, config, last_metrics)
print(json.dumps({"step": step, "checkpoint_started": last_metrics["checkpoint_started"]}), flush=True)
if hub is not None and step % args.state_upload_every == 0:
hub.state(last_metrics)
if hub is not None:
hub.checkpoint(args.max_steps, model, config, last_metrics)
hub.wait()
hub.state(last_metrics)
def main() -> None:
args = arguments()
if not torch.cuda.is_available():
raise RuntimeError("CUDA required")
if args.probe_only:
probe(args)
else:
train(args)
if __name__ == "__main__":
main()