import argparse import os import random import time import types import torch from datasets import load_dataset from huggingface_hub import HfApi, hf_hub_download from transformers import AutoTokenizer, LlamaConfig, LlamaForCausalLM, get_cosine_schedule_with_warmup CKPT_REPO = "TobiasLogic/textmodel-gci-scratch-ckpt" TOKENIZER_SRC = "TobiasLogic/Museko-125M" # tokenizer only (SmolLM2 vocab), not the pretrained weights MODEL_CONFIG = dict( vocab_size=49152, hidden_size=768, intermediate_size=2048, num_hidden_layers=12, num_attention_heads=12, num_key_value_heads=12, max_position_embeddings=2048, rms_norm_eps=1e-5, tie_word_embeddings=True, ) def load_gci_synth_module(): path = hf_hub_download(repo_id="TobiasLogic/gci-synth-train", repo_type="dataset", filename="gci_synth.py") mod = types.ModuleType("gci_synth") with open(path) as f: exec(compile(f.read(), path, "exec"), mod.__dict__) return mod class PackedMixedStream(torch.utils.data.IterableDataset): def __init__(self, tokenizer, seq_len, gci_module, synth_frac, seed): self.tokenizer = tokenizer self.seq_len = seq_len self.gci_module = gci_module self.synth_frac = synth_frac self.seed = seed def __iter__(self): worker = torch.utils.data.get_worker_info() seed = self.seed + (worker.id if worker else 0) rng = random.Random(seed) eos_id = self.tokenizer.eos_token_id fw = load_dataset( "HuggingFaceFW/fineweb-edu", name="sample-10BT", split="train", streaming=True, ) fw = fw.shuffle(seed=seed, buffer_size=2000) fw_iter = iter(fw) synth_counter = seed * 1_000_000 buf = [] while True: if rng.random() < self.synth_frac: item = self.gci_module.generate_item(synth_counter, rng) synth_counter += 1 text = f"{item['context']} {item['question']}" else: row = next(fw_iter, None) if row is None: fw_iter = iter(fw) row = next(fw_iter, None) if row is None: raise RuntimeError("fineweb-edu stream produced no data for this worker") text = row["text"] ids = self.tokenizer(text, truncation=True, max_length=self.seq_len * 4)["input_ids"] buf.extend(ids) buf.append(eos_id) while len(buf) >= self.seq_len: chunk = buf[: self.seq_len] buf = buf[self.seq_len :] input_ids = torch.tensor(chunk, dtype=torch.long) yield {"input_ids": input_ids, "labels": input_ids.clone()} def upload_checkpoint(local_dir, step): api = HfApi() api.create_repo(CKPT_REPO, repo_type="model", private=True, exist_ok=True) branch = f"step{step}" api.create_branch(CKPT_REPO, repo_type="model", branch=branch, exist_ok=True) api.upload_folder( repo_id=CKPT_REPO, folder_path=local_dir, path_in_repo="", revision=branch, commit_message=f"checkpoint at step {step}", ) def main(): ap = argparse.ArgumentParser() ap.add_argument("--steps", type=int, default=50000) ap.add_argument("--batch-size", type=int, default=32) ap.add_argument("--seq-len", type=int, default=1024) ap.add_argument("--lr", type=float, default=6e-4) ap.add_argument("--warmup-steps", type=int, default=1000) ap.add_argument("--ckpt-every", type=int, default=10000) ap.add_argument("--log-every", type=int, default=50) ap.add_argument("--synth-frac", type=float, default=0.15) ap.add_argument("--num-workers", type=int, default=4) ap.add_argument("--no-upload", action="store_true") ap.add_argument("--save-dir", default="/tmp/gci_scratch_ckpt") args = ap.parse_args() device = torch.device("cuda") print("GPU:", torch.cuda.get_device_name(0)) tokenizer = AutoTokenizer.from_pretrained(TOKENIZER_SRC) if tokenizer.pad_token is None: tokenizer.pad_token = tokenizer.eos_token gci_module = load_gci_synth_module() cfg = LlamaConfig(**MODEL_CONFIG) model = LlamaForCausalLM(cfg).to(device) n_params = sum(p.numel() for p in model.parameters()) print(f"model params (random init): {n_params:,}") ds = PackedMixedStream(tokenizer, args.seq_len, gci_module, args.synth_frac, seed=1234) loader = torch.utils.data.DataLoader(ds, batch_size=args.batch_size, num_workers=args.num_workers) optimizer = torch.optim.AdamW(model.parameters(), lr=args.lr, weight_decay=0.1, betas=(0.9, 0.95)) scheduler = get_cosine_schedule_with_warmup(optimizer, num_warmup_steps=args.warmup_steps, num_training_steps=args.steps) model.train() step = 0 running_loss = 0.0 t_start = time.time() for batch in loader: input_ids = batch["input_ids"].to(device) labels = batch["labels"].to(device) with torch.autocast(device_type="cuda", dtype=torch.bfloat16): out = model(input_ids=input_ids, labels=labels) loss = out.loss loss.backward() torch.nn.utils.clip_grad_norm_(model.parameters(), 1.0) optimizer.step() scheduler.step() optimizer.zero_grad(set_to_none=True) running_loss += loss.item() step += 1 if step % args.log_every == 0: elapsed = time.time() - t_start toks_seen = step * args.batch_size * args.seq_len print( f"[step {step}/{args.steps}] loss {running_loss / args.log_every:.4f} " f"lr {scheduler.get_last_lr()[0]:.2e} tokens {toks_seen:,} " f"({toks_seen / elapsed:,.0f} tok/s avg)" ) running_loss = 0.0 if step % args.ckpt_every == 0 or step == args.steps: local_dir = os.path.join(args.save_dir, f"checkpoint-{step}") os.makedirs(local_dir, exist_ok=True) model.save_pretrained(local_dir, safe_serialization=True) tokenizer.save_pretrained(local_dir) print(f"[step {step}] saved checkpoint to {local_dir}") if not args.no_upload: upload_checkpoint(local_dir, step) print(f"[step {step}] uploaded checkpoint-{step} to {CKPT_REPO}") if step >= args.steps: break print("pretraining complete") if __name__ == "__main__": main()