#!/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")