TFMF / train_lora_auto.py
StonePumpkins's picture
Upload folder using huggingface_hub
34cc882 verified
Raw
History Blame Contribute Delete
10.4 kB
"""
TFMF LoRA 训练脚本 - 新版自动化入口
这个脚本不会覆盖现有的 train_lora.py,
用于从一个 JSONL 数据目录或单个 JSONL 文件批量加载训练样本,
自动生成 LoRA 适配器并保存到指定输出目录。
格式要求:每条训练数据必须至少包含字段:system, user, assistant
示例:
{ "system": "...", "user": "...", "assistant": "..." }
"""
import argparse
import glob
import json
import os
import random
import sys
import torch
from datasets import Dataset
from peft import LoraConfig, get_peft_model
from transformers import (
AutoModelForCausalLM,
AutoTokenizer,
DataCollatorForSeq2Seq,
Trainer,
TrainingArguments,
)
DEFAULT_MODEL_PATH = "/data/coding/TFMF"
DEFAULT_DATA_DIR = "./TFMF/dataset"
DEFAULT_OUTPUT_DIR = "./TFMF/lora_adapters/teacher_chinese_auto"
DEFAULT_LORA_TARGET_MODULES = [
"q_proj",
"k_proj",
"v_proj",
"o_proj",
"gate_proj",
"up_proj",
"down_proj",
]
def parse_args():
parser = argparse.ArgumentParser(description="自动化 LoRA 训练脚本(新版本)")
parser.add_argument(
"--model-path",
type=str,
default=DEFAULT_MODEL_PATH,
help="基座模型目录,例如 /data/coding/TFMF",
)
parser.add_argument(
"--data-dir",
type=str,
default=DEFAULT_DATA_DIR,
help="训练数据目录,脚本会加载目录下所有 .jsonl 文件",
)
parser.add_argument(
"--data-path",
type=str,
default=None,
help="单个训练文件路径,优先于 --data-dir",
)
parser.add_argument(
"--output-dir",
type=str,
default=DEFAULT_OUTPUT_DIR,
help="LoRA 输出目录",
)
parser.add_argument("--epochs", type=int, default=3, help="训练轮数")
parser.add_argument("--batch-size", type=int, default=8, help="每卡 batch 大小")
parser.add_argument(
"--gradient-accumulation-steps",
type=int,
default=2,
help="梯度累积步数",
)
parser.add_argument("--learning-rate", type=float, default=2e-4, help="学习率")
parser.add_argument(
"--max-seq-length",
type=int,
default=2048,
help="最大序列长度",
)
parser.add_argument("--save-steps", type=int, default=50, help="保存间隔步数")
parser.add_argument("--logging-steps", type=int, default=10, help="日志输出步数")
parser.add_argument("--warmup-steps", type=int, default=10, help="warmup 步数")
parser.add_argument(
"--validation-split",
type=float,
default=0.1,
help="验证集比例,0.0 表示不切分",
)
parser.add_argument("--seed", type=int, default=42, help="随机种子")
parser.add_argument(
"--lora-r",
type=int,
default=32,
help="LoRA r 值",
)
parser.add_argument(
"--lora-alpha",
type=int,
default=64,
help="LoRA alpha 值",
)
parser.add_argument(
"--lora-dropout",
type=float,
default=0.05,
help="LoRA dropout",
)
return parser.parse_args()
def set_seed(seed: int):
random.seed(seed)
torch.manual_seed(seed)
if torch.cuda.is_available():
torch.cuda.manual_seed_all(seed)
def find_jsonl_files(data_dir: str):
if not os.path.isdir(data_dir):
raise FileNotFoundError(f"数据目录不存在: {data_dir}")
file_paths = sorted(glob.glob(os.path.join(data_dir, "*.jsonl")))
if not file_paths:
raise FileNotFoundError(f"未在数据目录中找到 .jsonl 文件: {data_dir}")
return file_paths
def load_samples_from_file(path: str):
if not os.path.isfile(path):
raise FileNotFoundError(f"数据文件不存在: {path}")
samples = []
with open(path, "r", encoding="utf-8") as f:
for line in f:
if not line.strip():
continue
item = json.loads(line)
if not all(k in item for k in ("system", "user", "assistant")):
raise ValueError(
f"样本缺少必要字段 system/user/assistant: {path}\n行内容: {line[:200]}"
)
samples.append(item)
return samples
def format_sample(sample: dict, tokenizer):
messages = [
{"role": "system", "content": sample["system"]},
{"role": "user", "content": sample["user"]},
{"role": "assistant", "content": sample["assistant"]},
]
text = tokenizer.apply_chat_template(
messages,
tokenize=False,
add_generation_prompt=False,
)
return {"text": text}
def tokenize_samples(examples, tokenizer, max_length):
return tokenizer(
examples["text"],
truncation=True,
max_length=max_length,
padding=False,
)
def build_dataset(samples, tokenizer, max_length):
dataset = Dataset.from_list(samples).map(
lambda x: format_sample(x, tokenizer),
batched=False,
)
dataset = dataset.map(
lambda x: tokenize_samples(x, tokenizer, max_length),
batched=True,
remove_columns=["text"],
)
dataset = dataset.map(lambda x: {"labels": x["input_ids"]})
return dataset
def print_config(args):
print("=" * 70)
print("LoRA 训练脚本参数:")
for name, value in vars(args).items():
print(f" {name}: {value}")
print("=" * 70)
def main():
args = parse_args()
set_seed(args.seed)
if args.data_path is None and args.data_dir is None:
raise ValueError("请通过 --data-path 或 --data-dir 指定训练数据")
print_config(args)
if args.data_path is not None:
print(f"加载单个 JSONL 文件: {args.data_path}")
samples = load_samples_from_file(args.data_path)
else:
print(f"加载数据目录: {args.data_dir}")
file_paths = find_jsonl_files(args.data_dir)
samples = []
for path in file_paths:
print(f" 加载: {path}")
samples.extend(load_samples_from_file(path))
if not samples:
raise ValueError("没有加载到任何训练样本,请检查数据文件")
random.shuffle(samples)
split_index = int(len(samples) * (1 - args.validation_split))
train_samples = samples[:split_index]
val_samples = samples[split_index:] if args.validation_split > 0 else []
print(f"总样本数: {len(samples)}")
print(f"训练集: {len(train_samples)}")
print(f"验证集: {len(val_samples)}")
print("加载 tokenizer 和基座模型...")
tokenizer = AutoTokenizer.from_pretrained(args.model_path, trust_remote_code=True)
if tokenizer.pad_token is None:
tokenizer.pad_token = tokenizer.eos_token
tokenizer.padding_side = "right"
model = AutoModelForCausalLM.from_pretrained(
args.model_path,
device_map="auto",
dtype=torch.bfloat16,
trust_remote_code=True,
low_cpu_mem_usage=True,
)
model.config.use_cache = False
print("配置 LoRA...")
lora_config = LoraConfig(
r=args.lora_r,
lora_alpha=args.lora_alpha,
target_modules=DEFAULT_LORA_TARGET_MODULES,
lora_dropout=args.lora_dropout,
bias="none",
task_type="CAUSAL_LM",
)
model = get_peft_model(model, lora_config)
trainable = sum(p.numel() for p in model.parameters() if p.requires_grad)
total = sum(p.numel() for p in model.parameters())
print(f"可训练参数: {trainable / 1e6:.2f}M")
print(f"总参数: {total / 1e9:.2f}B")
print(f"可训练参数占比: {trainable / total * 100:.2f}%")
print("构建训练数据集...")
train_dataset = build_dataset(train_samples, tokenizer, args.max_seq_length)
eval_dataset = build_dataset(val_samples, tokenizer, args.max_seq_length) if val_samples else None
os.makedirs(args.output_dir, exist_ok=True)
final_output_dir = os.path.join(args.output_dir, "final")
os.makedirs(final_output_dir, exist_ok=True)
training_args = TrainingArguments(
output_dir=args.output_dir,
num_train_epochs=args.epochs,
per_device_train_batch_size=args.batch_size,
per_device_eval_batch_size=args.batch_size,
gradient_accumulation_steps=args.gradient_accumulation_steps,
learning_rate=args.learning_rate,
warmup_steps=args.warmup_steps,
logging_steps=args.logging_steps,
save_steps=args.save_steps,
save_total_limit=3,
bf16=True,
fp16=False,
gradient_checkpointing=True,
max_grad_norm=0.3,
weight_decay=0.01,
report_to="none",
eval_strategy="steps" if eval_dataset is not None else "no",
save_strategy="steps",
load_best_model_at_end=False,
logging_dir=os.path.join(args.output_dir, "logs"),
)
trainer = Trainer(
model=model,
args=training_args,
train_dataset=train_dataset,
eval_dataset=eval_dataset,
data_collator=DataCollatorForSeq2Seq(tokenizer, padding=True),
)
print("开始训练...")
trainer.train()
print(f"保存最终 LoRA 适配器到: {final_output_dir}")
model.save_pretrained(final_output_dir)
tokenizer.save_pretrained(final_output_dir)
metadata = {
"base_model": args.model_path,
"data_source": args.data_path or args.data_dir,
"train_samples": len(train_samples),
"validation_samples": len(val_samples),
"epochs": args.epochs,
"batch_size": args.batch_size,
"gradient_accumulation_steps": args.gradient_accumulation_steps,
"learning_rate": args.learning_rate,
"max_seq_length": args.max_seq_length,
"lora_r": args.lora_r,
"lora_alpha": args.lora_alpha,
"lora_dropout": args.lora_dropout,
}
with open(os.path.join(args.output_dir, "adapter_config.json"), "w", encoding="utf-8") as f:
json.dump(metadata, f, indent=2, ensure_ascii=False)
total_size = 0
for root, _, files in os.walk(final_output_dir):
for fname in files:
total_size += os.path.getsize(os.path.join(root, fname))
print(f"最终适配器大小: {total_size / 1024 / 1024:.1f} MB")
print("训练完成,请使用 final 目录下的适配器进行加载。")
if __name__ == "__main__":
main()