| """ |
| GPT-2 预训练主脚本 |
| 使用 Hugging Face Trainer |
| """ |
| import os |
| import argparse |
| import torch |
| import wandb |
| import swanlab |
| from transformers import ( |
| Trainer, |
| TrainingArguments, |
| get_scheduler |
| ) |
| from transformers.trainer_callback import TrainerCallback |
|
|
| from config import get_config |
| from model import create_model, get_tokenizer |
| from dataset import PretrainDataset, DataCollatorForLM |
|
|
|
|
| def setup_swanlab_auth(): |
| """从环境变量读取 SwanLab API Key 并尝试登录。""" |
| api_key = os.environ.get("SWANLAB_API_KEY", "").strip() |
| if not api_key: |
| return False |
| if hasattr(swanlab, "login"): |
| try: |
| swanlab.login(api_key=api_key) |
| except TypeError: |
| try: |
| swanlab.login(api_key) |
| except Exception: |
| pass |
| except Exception: |
| pass |
| return True |
|
|
|
|
| class LoggingCallback(TrainerCallback): |
| """自定义日志回调""" |
|
|
| def on_log(self, args, state, control, logs=None, **kwargs): |
| if logs is not None and state.is_world_process_zero: |
| |
| if "loss" in logs: |
| print(f"Step {state.global_step}: loss={logs['loss']:.4f}", end="") |
| if "learning_rate" in logs: |
| print(f", lr={logs['learning_rate']:.2e}", end="") |
| print() |
|
|
|
|
| def compute_metrics(eval_pred): |
| """计算评估指标""" |
| import numpy as np |
|
|
| logits, labels = eval_pred |
| |
| shift_logits = logits[..., :-1, :].contiguous() |
| shift_labels = labels[..., 1:].contiguous() |
|
|
| |
| loss_fct = torch.nn.CrossEntropyLoss() |
| loss = loss_fct( |
| shift_logits.view(-1, shift_logits.size(-1)), |
| shift_labels.view(-1) |
| ) |
| perplexity = torch.exp(loss).item() |
|
|
| return {"perplexity": perplexity} |
|
|
|
|
| def train(config): |
| """训练函数""" |
| |
| torch.manual_seed(config.training.seed) |
|
|
| |
| os.makedirs(config.training.output_dir, exist_ok=True) |
|
|
| |
| if int(os.environ.get("LOCAL_RANK", 0)) == 0: |
| authed = setup_swanlab_auth() |
| if not authed: |
| print("⚠️ 未检测到 SWANLAB_API_KEY,SwanLab 可能无法上传到你的账号") |
| swanlab_run = swanlab.init( |
| project="GPT2-Dropout-Comparison", |
| experiment_name=f"gpt2-dropout-{config.model.resid_pdrop}", |
| config={ |
| "model_size": config.model.model_size, |
| "dropout": config.model.resid_pdrop, |
| "learning_rate": config.training.learning_rate, |
| "batch_size": config.training.per_device_train_batch_size, |
| "gradient_accumulation": config.training.gradient_accumulation_steps, |
| "warmup_steps": config.training.warmup_steps, |
| "weight_decay": config.training.weight_decay, |
| "adam_beta1": config.training.adam_beta1, |
| "adam_beta2": config.training.adam_beta2, |
| "max_grad_norm": config.training.max_grad_norm, |
| "num_epochs": config.training.num_train_epochs, |
| "block_size": config.data.block_size, |
| }, |
| ) |
|
|
| |
| |
| should_init_from_scratch = config.model.from_scratch and not config.training.resume_from_checkpoint |
|
|
| model, model_config = create_model( |
| model_size=config.model.model_size, |
| resid_pdrop=config.model.resid_pdrop, |
| attn_pdrop=config.model.attn_pdrop, |
| embd_pdrop=config.model.embd_pdrop, |
| from_scratch=should_init_from_scratch |
| ) |
|
|
| |
| tokenizer = get_tokenizer(config.model.model_size) |
|
|
| |
| print("\nLoading datasets...") |
| train_dataset = PretrainDataset( |
| config.data.train_file, |
| config.data.block_size |
| ) |
| val_dataset = PretrainDataset( |
| config.data.val_file, |
| config.data.block_size |
| ) |
|
|
| |
| data_collator = DataCollatorForLM() |
|
|
| |
| training_args = TrainingArguments( |
| output_dir=config.training.output_dir, |
|
|
| |
| num_train_epochs=config.training.num_train_epochs, |
| max_steps=config.training.max_steps, |
|
|
| |
| per_device_train_batch_size=config.training.per_device_train_batch_size, |
| per_device_eval_batch_size=config.training.per_device_eval_batch_size, |
| gradient_accumulation_steps=config.training.gradient_accumulation_steps, |
|
|
| |
| learning_rate=config.training.learning_rate, |
| weight_decay=config.training.weight_decay, |
| adam_beta1=config.training.adam_beta1, |
| adam_beta2=config.training.adam_beta2, |
| warmup_steps=config.training.warmup_steps, |
| max_grad_norm=config.training.max_grad_norm, |
| lr_scheduler_type=config.training.lr_scheduler_type, |
|
|
| |
| logging_dir=os.path.join(config.training.output_dir, "logs"), |
| logging_steps=config.training.logging_steps, |
| logging_first_step=True, |
|
|
| |
| save_steps=config.training.save_steps, |
| save_total_limit=5, |
|
|
| |
| eval_strategy="steps", |
| eval_steps=config.training.eval_steps, |
| load_best_model_at_end=True, |
| metric_for_best_model="eval_loss", |
| greater_is_better=False, |
|
|
| |
| fp16=config.training.fp16, |
| bf16=config.training.bf16, |
|
|
| |
| seed=config.training.seed, |
| dataloader_num_workers=config.training.dataloader_num_workers, |
| dataloader_pin_memory=True, |
| ddp_find_unused_parameters=False, |
|
|
| |
| report_to=["tensorboard", "swanlab"], |
|
|
| |
| resume_from_checkpoint=config.training.resume_from_checkpoint, |
| ) |
|
|
| |
| trainer = Trainer( |
| model=model, |
| args=training_args, |
| train_dataset=train_dataset, |
| eval_dataset=val_dataset, |
| data_collator=data_collator, |
| callbacks=[LoggingCallback()], |
| ) |
|
|
| |
| print("\n" + "=" * 50) |
| print("Starting training...") |
| if config.training.resume_from_checkpoint: |
| print(f"Resuming from checkpoint: {config.training.resume_from_checkpoint}") |
| |
| import glob |
| checkpoints = sorted(glob.glob(os.path.join(config.training.output_dir, "checkpoint-*"))) |
| if checkpoints: |
| print(f"Available checkpoints: {[os.path.basename(c) for c in checkpoints]}") |
| print(f"Latest checkpoint: {os.path.basename(checkpoints[-1])}") |
| else: |
| print("WARNING: No checkpoints found!") |
| print("=" * 50) |
|
|
| trainer.train(resume_from_checkpoint=config.training.resume_from_checkpoint) |
|
|
| |
| print("\nSaving final model...") |
| trainer.save_model(os.path.join(config.training.output_dir, "final")) |
| tokenizer.save_pretrained(os.path.join(config.training.output_dir, "final")) |
|
|
| |
| print("\nFinal evaluation...") |
| eval_results = trainer.evaluate() |
| print(f"Final eval loss: {eval_results['eval_loss']:.4f}") |
| final_ppl = torch.exp(torch.tensor(eval_results['eval_loss'])).item() |
| print(f"Final perplexity: {final_ppl:.2f}") |
|
|
| |
| if int(os.environ.get("LOCAL_RANK", 0)) == 0: |
| swanlab.log({ |
| "final/eval_loss": eval_results['eval_loss'], |
| "final/perplexity": final_ppl, |
| }) |
| |
| swanlab.finish() |
|
|
| return trainer |
|
|
|
|
| def main(): |
| parser = argparse.ArgumentParser(description="Pretrain GPT-2 with Dropout") |
|
|
| |
| parser.add_argument("--model_size", type=str, default="gpt2-medium", |
| choices=["gpt2", "gpt2-medium", "gpt2-large", "gpt2-xl"]) |
| parser.add_argument("--dropout", type=float, default=0.1) |
| parser.add_argument("--from_scratch", action="store_true", default=False) |
|
|
| |
| parser.add_argument("--train_file", type=str, default="./data/train.bin") |
| parser.add_argument("--val_file", type=str, default="./data/val.bin") |
| parser.add_argument("--block_size", type=int, default=1024) |
|
|
| |
| parser.add_argument("--output_dir", type=str, default="./output/gpt2-dropout") |
| parser.add_argument("--num_epochs", type=int, default=1) |
| parser.add_argument("--max_steps", type=int, default=-1) |
| parser.add_argument("--batch_size", type=int, default=12) |
| parser.add_argument("--gradient_accumulation", type=int, default=40) |
| parser.add_argument("--learning_rate", type=float, default=3e-4) |
| parser.add_argument("--warmup_steps", type=int, default=2000) |
| parser.add_argument("--weight_decay", type=float, default=0.1) |
| parser.add_argument("--adam_beta1", type=float, default=0.9) |
| parser.add_argument("--adam_beta2", type=float, default=0.95) |
| parser.add_argument("--max_grad_norm", type=float, default=1.0) |
|
|
| parser.add_argument("--logging_steps", type=int, default=100) |
| parser.add_argument("--save_steps", type=int, default=5000) |
| parser.add_argument("--eval_steps", type=int, default=1000) |
|
|
| parser.add_argument("--fp16", action="store_true", default=True) |
| parser.add_argument("--bf16", action="store_true", default=False) |
|
|
| parser.add_argument("--resume", type=str, default=None) |
|
|
| args = parser.parse_args() |
|
|
| |
| config = get_config( |
| model_size=args.model_size, |
| dropout=args.dropout, |
| from_scratch=args.from_scratch, |
| ) |
|
|
| |
| config.data.train_file = args.train_file |
| config.data.val_file = args.val_file |
| config.data.block_size = args.block_size |
|
|
| config.training.output_dir = args.output_dir |
| config.training.num_train_epochs = args.num_epochs |
| config.training.max_steps = args.max_steps |
| config.training.per_device_train_batch_size = args.batch_size |
| config.training.per_device_eval_batch_size = args.batch_size |
| config.training.gradient_accumulation_steps = args.gradient_accumulation |
| config.training.learning_rate = args.learning_rate |
| config.training.warmup_steps = args.warmup_steps |
| config.training.weight_decay = args.weight_decay |
| config.training.adam_beta1 = args.adam_beta1 |
| config.training.adam_beta2 = args.adam_beta2 |
| config.training.max_grad_norm = args.max_grad_norm |
| config.training.logging_steps = args.logging_steps |
| config.training.save_steps = args.save_steps |
| config.training.eval_steps = args.eval_steps |
| config.training.fp16 = args.fp16 |
| config.training.bf16 = args.bf16 |
|
|
| |
| if args.resume and args.resume.lower() == "true": |
| config.training.resume_from_checkpoint = True |
| else: |
| config.training.resume_from_checkpoint = args.resume |
|
|
| |
| print("=" * 50) |
| print("Training Configuration") |
| print("=" * 50) |
| print(f"Model: {config.model.model_size}") |
| print(f"Dropout: {config.model.resid_pdrop}") |
| print(f"From scratch: {config.model.from_scratch}") |
| print(f"Resume from checkpoint: {config.training.resume_from_checkpoint}") |
| print(f"Block size: {config.data.block_size}") |
| print(f"Batch size: {config.training.per_device_train_batch_size}") |
| print(f"Gradient accumulation: {config.training.gradient_accumulation_steps}") |
| print(f"Effective batch size: {config.training.per_device_train_batch_size * config.training.gradient_accumulation_steps}") |
| print(f"Learning rate: {config.training.learning_rate}") |
| print(f"Output dir: {config.training.output_dir}") |
| print("=" * 50) |
|
|
| |
| train(config) |
|
|
|
|
| if __name__ == "__main__": |
| main() |
|
|