diff --git a/train_mother.py b/train_mother.py new file mode 100644 index 0000000..9e65f06 --- /dev/null +++ b/train_mother.py @@ -0,0 +1,129 @@ +#!/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!")