Files
ChunMengDie-1.0-0.4b/train[1].py

237 lines
8.6 KiB
Python
Raw Permalink Normal View History

#!/usr/bin/env python3
# Copyright (c) 2026 XingChina
# SPDX-License-Identifier: BSD-3-Clause
# 本代码采用 BSD 3-Clause 许可证,详见项目根目录的 LICENSE 文件。
# 为4090显卡进行优化,保证普通显卡也能训练(而不是找大显存显卡)
# 已经拿4090测试过,全程没有崩溃
import os
import json
import signal
import sys
import torch
from transformers import (
GPT2Config,
GPT2LMHeadModel,
GPT2Tokenizer,
Trainer,
TrainingArguments,
DataCollatorForLanguageModeling,
)
from datasets import Dataset
from bitsandbytes.optim import Adam8bit
import logging
logging.basicConfig(level=logging.INFO)
logger = logging.getLogger(__name__)
# ================== 模型配置 ==================
MODEL_CONFIG = {
"n_embd": 1024,
"n_layer": 16,
"n_head": 16,
"n_positions": 4096,
"vocab_size": 50257,
"n_ctx": 4096,
"resid_pdrop": 0.3,
"embd_pdrop": 0.3,
"attn_pdrop": 0.3,
}
# ================== 训练参数(为4090优化,保证普通显卡能复现) ==================
TRAIN_ARGS = {
"output_dir": "./checkpoints_chunmengdie_gpt2",
"per_device_train_batch_size": 2,
"gradient_accumulation_steps": 8,
"num_train_epochs": 1,
"learning_rate": 3e-4,
"weight_decay": 0.01,
"warmup_steps": 500,
"logging_steps": 10,
"save_steps": 500,
"save_total_limit": 3,
"bf16": True,
"report_to": "none",
"dataloader_num_workers": 0,
"eval_strategy": "steps",
"eval_steps": 500,
"load_best_model_at_end": True,
"metric_for_best_model": "eval_loss",
"greater_is_better": False,
"gradient_checkpointing": True,
}
MAX_SEQ_LEN = 4096
CHECKPOINT_DIR = TRAIN_ARGS["output_dir"]
# ====== 数据文件列表(这里假设数据已经清理过,实际用的数据确实清理过) ======
# mengdie_train由于包含其他数据集,没有开源
DATA_FILES = [
"mengdie_train.json",
"Belle_open_source_0.5M.json",
]
# ================== 数据加载 ==================
def load_data(file_list):
texts = []
for fpath in file_list:
if not os.path.exists(fpath):
logger.warning(f"File not found: {fpath}, skipped")
continue
with open(fpath, 'r', encoding='utf-8') as f:
first_char = f.read(1)
f.seek(0)
if first_char == '[':
data = json.load(f)
logger.info(f"Loaded {len(data)} samples from {fpath} (JSON array)")
else:
data = []
for line in f:
line = line.strip()
if line:
try:
data.append(json.loads(line))
except json.JSONDecodeError:
pass
logger.info(f"Loaded {len(data)} samples from {fpath} (JSONL)")
for item in data:
try:
if "instruction" in item and "input" in item and "output" in item:
inst = item["instruction"].strip()
inp = item.get("input", "").strip()
out = item["output"].strip()
user_text = f"{inst}\n{inp}" if inp else inst
assistant_text = out
elif "input" in item and "output" in item:
user_text = item["input"].strip()
assistant_text = item["output"].strip()
else:
continue
texts.append(f"用户:{user_text}\n猫娘:{assistant_text}")
except Exception:
continue
logger.info(f"Total texts: {len(texts)}")
return texts
# ================== Tokenization==================
def tokenize_function(examples):
tokenized = tokenizer(
examples["text"],
truncation=True,
max_length=MAX_SEQ_LEN - 1,
padding=False,
return_attention_mask=False,
)
tokenized["input_ids"] = [ids + [tokenizer.eos_token_id] for ids in tokenized["input_ids"]]
return tokenized
# ================== 紧急保存 ==================
def emergency_save(sig, frame):
logger.info("\n🛑 Saving emergency checkpoint...")
try:
torch.cuda.synchronize()
model_to_save = model.cpu()
os.makedirs(CHECKPOINT_DIR, exist_ok=True)
model_to_save.save_pretrained(os.path.join(CHECKPOINT_DIR, "emergency"))
tokenizer.save_pretrained(os.path.join(CHECKPOINT_DIR, "emergency"))
logger.info("✅ Emergency checkpoint saved.")
except Exception as e:
logger.error(f"Emergency save failed: {e}")
sys.exit(0)
# ================== 主程序 ==================
if __name__ == "__main__":
tokenizer = GPT2Tokenizer.from_pretrained("gpt2")
tokenizer.pad_token = tokenizer.eos_token
raw_texts = load_data(DATA_FILES)
dataset = Dataset.from_dict({"text": raw_texts})
tokenized_dataset = dataset.map(
tokenize_function,
batched=True,
remove_columns=["text"],
num_proc=4,
load_from_cache_file=True,
)
# 验证 EOS 是否添加成功
sample_ids = tokenized_dataset[0]["input_ids"]
logger.info(f"Sample last 5 tokens: {sample_ids[-5:]}")
logger.info(f"Last token is EOS: {sample_ids[-1] == tokenizer.eos_token_id}")
# 9:1 划分
total_len = len(tokenized_dataset)
train_size = int(0.9 * total_len)
eval_size = total_len - train_size
train_dataset, eval_dataset = torch.utils.data.random_split(
tokenized_dataset, [train_size, eval_size]
)
logger.info(f"Train: {train_size}, Eval: {eval_size}")
logger.info("Initializing GPT-2 model from scratch...")
config = GPT2Config(**MODEL_CONFIG)
model = GPT2LMHeadModel(config)
optimizer = Adam8bit(model.parameters(), lr=TRAIN_ARGS["learning_rate"])
training_args = TrainingArguments(
output_dir=CHECKPOINT_DIR,
per_device_train_batch_size=TRAIN_ARGS["per_device_train_batch_size"],
gradient_accumulation_steps=TRAIN_ARGS["gradient_accumulation_steps"],
num_train_epochs=TRAIN_ARGS["num_train_epochs"],
learning_rate=TRAIN_ARGS["learning_rate"],
weight_decay=TRAIN_ARGS["weight_decay"],
warmup_steps=TRAIN_ARGS["warmup_steps"],
logging_steps=TRAIN_ARGS["logging_steps"],
save_steps=TRAIN_ARGS["save_steps"],
save_total_limit=TRAIN_ARGS["save_total_limit"],
bf16=TRAIN_ARGS["bf16"],
report_to=TRAIN_ARGS["report_to"],
dataloader_num_workers=TRAIN_ARGS["dataloader_num_workers"],
optim="adamw_8bit",
eval_strategy=TRAIN_ARGS["eval_strategy"],
eval_steps=TRAIN_ARGS["eval_steps"],
load_best_model_at_end=TRAIN_ARGS["load_best_model_at_end"],
metric_for_best_model=TRAIN_ARGS["metric_for_best_model"],
greater_is_better=TRAIN_ARGS["greater_is_better"],
gradient_checkpointing=TRAIN_ARGS["gradient_checkpointing"],
)
data_collator = DataCollatorForLanguageModeling(tokenizer=tokenizer, mlm=False)
trainer = Trainer(
model=model,
args=training_args,
train_dataset=train_dataset,
eval_dataset=eval_dataset,
data_collator=data_collator,
optimizers=(optimizer, None),
)
signal.signal(signal.SIGINT, emergency_save)
# 自动恢复最新 checkpoint(中间checkpoint保留调度器,最终权重无调度器,这是手滑删了调度器的主要原因)
latest_checkpoint = None
if os.path.exists(CHECKPOINT_DIR):
checkpoints = [d for d in os.listdir(CHECKPOINT_DIR) if d.startswith("checkpoint-")]
if checkpoints:
latest_checkpoint = sorted(checkpoints, key=lambda x: int(x.split("-")[1]))[-1]
latest_checkpoint = os.path.join(CHECKPOINT_DIR, latest_checkpoint)
logger.info(f"✅ Found checkpoint: {latest_checkpoint}, will resume from there.")
else:
if os.path.exists(os.path.join(CHECKPOINT_DIR, "emergency")):
latest_checkpoint = os.path.join(CHECKPOINT_DIR, "emergency")
logger.info(f"✅ Found emergency checkpoint, will resume from there.")
else:
os.makedirs(CHECKPOINT_DIR, exist_ok=True)
logger.info("🚀 Starting training...")
trainer.train(resume_from_checkpoint=latest_checkpoint)
final_path = os.path.join(CHECKPOINT_DIR, "final_model")
model.save_pretrained(final_path)
tokenizer.save_pretrained(final_path)
logger.info(f"✅ Final model saved to {final_path}")
logger.info("📌 Convert to GGUF: python /path/to/llama.cpp/convert_hf_to_gguf.py ./checkpoints_chunmengdie_gpt2/final_model --outfile model.gguf --outtype q8_0")