TextModel-v1 / pretrain_gci.py
TobiasLogic's picture
Add model weights, tokenizer, and training/data-generation code
e16ecc9 verified
Raw
History Blame Contribute Delete
6.5 kB
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()