| 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" |
|
|
| 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() |
|
|