guanghulab/train_mother.py

130 lines
4.4 KiB
Python
Raw Blame History

This file contains ambiguous Unicode characters

This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.

#!/usr/bin/env python3
"""
Qwen2.5-7B SFT 训练脚本
用法:
python3 train_mother.py
依赖:
pip3 install transformers accelerate datasets
"""
import os, sys, json, torch
from transformers import AutoModelForCausalLM, AutoTokenizer, TrainingArguments, Trainer
from datasets import load_dataset
BS = int(os.environ.get("BATCH_SIZE", "4"))
GA = int(os.environ.get("GRAD_ACCUM", "8"))
LR = float(os.environ.get("LEARNING_RATE", "1e-5"))
EPOCHS = int(os.environ.get("EPOCHS", "3"))
DS = os.environ.get("DATASET", "autodl-tmp/data/sft.jsonl")
OUT = os.environ.get("OUTPUT_DIR", "autodl-tmp/output/qwen25-7b-sft")
print("="*50)
print("Qwen2.5-7B SFT Training")
print(f" Batch: {BS}, GradAccum: {GA}, Eff: {BS*GA}")
print(f" LR: {LR}, Epochs: {EPOCHS}")
print(f" Data: {DS}")
print(f" Out: {OUT}")
print("="*50)
# ========== 1. Load model ==========
print("[1/5] Loading model...")
model = AutoModelForCausalLM.from_pretrained("Qwen/Qwen2.5-7B", torch_dtype=torch.bfloat16, trust_remote_code=True, device_map="auto")
tokenizer = AutoTokenizer.from_pretrained("Qwen/Qwen2.5-7B", trust_remote_code=True)
tokenizer.pad_token = tokenizer.eos_token
total = sum(p.numel() for p in model.parameters())
print(f" Parameters: {total/1e9:.2f}B")
print(f" Device: {model.device}")
# ========== 2. Load data ==========
print("[2/5] Loading data...")
dataset = load_dataset("json", data_files=DS, split="train")
print(f" Samples: {len(dataset)}")
print(f" Format:")
for k in dataset[0]:
v = dataset[0][k]
if isinstance(v, list):
print(f" {k}: [{len(v)} msgs]")
for m in v[:2]:
print(f" {m['role']}: {m['content'][:50]}...")
else:
print(f" {k}: {v}")
# ========== 3. Tokenize ==========
print("[3/5] Tokenizing data...")
def tokenize(example):
# Build full chat text using Qwen chat template
texts = tokenizer.apply_chat_template(example["messages"], tokenize=False)
enc = tokenizer(texts, truncation=True, max_length=8192, add_special_tokens=False)
return {"input_ids": enc["input_ids"], "labels": enc["input_ids"].copy()}
dataset = dataset.map(tokenize, remove_columns=["messages"], num_proc=8)
total_tokens = sum(len(x["input_ids"]) for x in dataset)
post_pad = sum(x["input_ids"].count(tokenizer.pad_token_id) for x in dataset) if hasattr(tokenizer, "pad_token_id") else 0
print(f" Total tokens: {total_tokens:,}")
def collate(features):
max_len = max(len(f["input_ids"]) for f in features)
batch = {}
for k in ["input_ids", "labels", "attention_mask"]:
pad = tokenizer.pad_token_id if k != "labels" else -100
batch[k] = torch.tensor([f[k] + [pad]*(max_len-len(f[k])) for f in features])
return batch
model.config.use_cache = False
# ========== 4. Training args ==========
print("[4/5] Training config...")
args = TrainingArguments(
output_dir=OUT, num_train_epochs=EPOCHS,
per_device_train_batch_size=BS, gradient_accumulation_steps=GA,
learning_rate=LR, warmup_ratio=0.05, lr_scheduler_type="cosine",
bf16=True, tf32=True, logging_steps=10,
save_strategy="epoch", save_total_limit=3,
remove_unused_columns=False, dataloader_num_workers=4,
gradient_checkpointing=True, optim="adamw_torch",
report_to="none", ddp_find_unused_parameters=False,
)
trainer = Trainer(model=model, args=args, train_dataset=ds, data_collator=collate)
# ========== 5. Go ==========
print("[5/5] Starting training!")
gpu = torch.cuda.get_device_name(0)
mem = torch.cuda.get_device_properties(0).total_memory / 1e9
print(f" GPU: {gpu} ({mem:.1f}GB) | Epochs: {EPOCHS} | Eff batch: {BS*GA} | LR: {LR}")
sys.stdout.flush()
trainer.train()
# ========== 6. Save ==========
print("Saving model...")
final = os.path.join(OUT, "final")
trainer.save_model(final)
tokenizer.save_pretrained(final)
# ⚠️ 关键修复Qwen chat template 使用 <|im_end|> (151645) 作为对话EOS
# 但 config.json 中默认 eos_token_id=151643 (<|endoftext|>)
# 不修复会导致部署时模型无限生成 → 死循环乱码
model.config.eos_token_id = 151645
model.config.save_pretrained(final)
# 修复 tokenizer 默认system prompt
tok_cfg_path = os.path.join(final, "tokenizer_config.json")
with open(tok_cfg_path) as f:
tok_cfg = json.load(f)
tok_cfg["default_system"] = ""
with open(tok_cfg_path, "w") as f:
json.dump(tok_cfg, f, indent=2, ensure_ascii=False)
peak = torch.cuda.max_memory_allocated() / 1e9
print(f" Model: {final}")
print(f" Peak VRAM: {peak:.2f}GB / {mem:.1f}GB")
print(f" DONE!")