初始化项目,由ModelHub XC社区提供模型
Model: XingChina/ChunMengDie-1.0-0.4b Source: Original Platform
This commit is contained in:
237
train[1].py
Normal file
237
train[1].py
Normal file
@@ -0,0 +1,237 @@
|
||||
#!/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")
|
||||
Reference in New Issue
Block a user