237 lines
8.6 KiB
Python
237 lines
8.6 KiB
Python
#!/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") |