初始化项目,由ModelHub XC社区提供模型
Model: TobiasLogic/TextModel-v1 Source: Original Platform
This commit is contained in:
179
pretrain_gci.py
Normal file
179
pretrain_gci.py
Normal file
@@ -0,0 +1,179 @@
|
||||
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()
|
||||
Reference in New Issue
Block a user