Files
ChunMengDie-1.0-0.4b/train[1].py
ModelHub XC 1d6c065d51 初始化项目,由ModelHub XC社区提供模型
Model: XingChina/ChunMengDie-1.0-0.4b
Source: Original Platform
2026-07-25 07:56:11 +08:00

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