271 lines
10 KiB
Python
Executable File
271 lines
10 KiB
Python
Executable File
#!/usr/bin/env python3
|
|
import argparse
|
|
import inspect
|
|
import json
|
|
from pathlib import Path
|
|
import re
|
|
|
|
|
|
def require(pkg):
|
|
try:
|
|
return __import__(pkg)
|
|
except Exception as e:
|
|
raise SystemExit(f"missing python package '{pkg}': {e}")
|
|
|
|
|
|
def main():
|
|
ap = argparse.ArgumentParser(description="Train LoRA for MCP quality tasks")
|
|
ap.add_argument("--base-model", required=True)
|
|
ap.add_argument("--train-file", required=True)
|
|
ap.add_argument("--val-file", required=True)
|
|
ap.add_argument("--output-dir", required=True)
|
|
ap.add_argument("--lora-r", type=int, default=64)
|
|
ap.add_argument("--lora-alpha", type=int, default=128)
|
|
ap.add_argument("--lora-dropout", type=float, default=0.05)
|
|
ap.add_argument("--lora-target-modules", default="q_proj,k_proj,v_proj,o_proj,gate_proj,up_proj,down_proj")
|
|
ap.add_argument("--num-epochs", type=float, default=2)
|
|
ap.add_argument("--learning-rate", type=float, default=2e-4)
|
|
ap.add_argument("--weight-decay", type=float, default=0.01)
|
|
ap.add_argument("--max-seq-len", type=int, default=2048)
|
|
ap.add_argument("--per-device-train-batch", type=int, default=1)
|
|
ap.add_argument("--per-device-eval-batch", type=int, default=1)
|
|
ap.add_argument("--grad-accum", type=int, default=16)
|
|
ap.add_argument("--warmup-ratio", type=float, default=0.03)
|
|
ap.add_argument("--logging-steps", type=int, default=20)
|
|
ap.add_argument("--eval-steps", type=int, default=200)
|
|
ap.add_argument("--save-steps", type=int, default=200)
|
|
ap.add_argument("--max-steps", type=int, default=0)
|
|
ap.add_argument("--use-4bit", action="store_true")
|
|
ap.add_argument("--use-8bit", action="store_true")
|
|
ap.add_argument("--bf16", action="store_true")
|
|
ap.add_argument("--seed", type=int, default=42)
|
|
ap.add_argument("--resume-from-checkpoint", default="")
|
|
ap.add_argument("--auto-resume", action="store_true")
|
|
ap.add_argument("--local-files-only", action="store_true")
|
|
ap.add_argument("--trust-remote-code", action="store_true")
|
|
ap.add_argument("--cpu-offload", action="store_true")
|
|
ap.add_argument("--offload-dir", default="")
|
|
ap.add_argument("--gpu-max-memory-gib", type=int, default=0)
|
|
ap.add_argument("--cpu-max-memory-gib", type=int, default=0)
|
|
args = ap.parse_args()
|
|
|
|
require("torch")
|
|
require("datasets")
|
|
require("transformers")
|
|
require("peft")
|
|
require("trl")
|
|
|
|
import torch
|
|
from datasets import load_dataset
|
|
from peft import LoraConfig
|
|
from transformers import AutoModelForCausalLM, AutoTokenizer, BitsAndBytesConfig
|
|
from trl import SFTTrainer, SFTConfig
|
|
|
|
train_path = Path(args.train_file)
|
|
val_path = Path(args.val_file)
|
|
out_dir = Path(args.output_dir)
|
|
out_dir.mkdir(parents=True, exist_ok=True)
|
|
|
|
if not train_path.exists() or not val_path.exists():
|
|
raise SystemExit("train/val files are required and must exist")
|
|
|
|
# Build/validate Trainer args early so API mismatches fail before heavy model loading.
|
|
sft_params = inspect.signature(SFTConfig.__init__).parameters
|
|
if "evaluation_strategy" in sft_params:
|
|
eval_key = "evaluation_strategy"
|
|
elif "eval_strategy" in sft_params:
|
|
eval_key = "eval_strategy"
|
|
else:
|
|
raise SystemExit("unsupported TRL SFTConfig: missing eval/evaluation strategy arg")
|
|
|
|
targs_kwargs = dict(
|
|
output_dir=str(out_dir),
|
|
num_train_epochs=args.num_epochs,
|
|
learning_rate=args.learning_rate,
|
|
weight_decay=args.weight_decay,
|
|
per_device_train_batch_size=args.per_device_train_batch,
|
|
per_device_eval_batch_size=args.per_device_eval_batch,
|
|
gradient_accumulation_steps=args.grad_accum,
|
|
warmup_ratio=args.warmup_ratio,
|
|
logging_steps=args.logging_steps,
|
|
eval_steps=args.eval_steps,
|
|
save_steps=args.save_steps,
|
|
save_strategy="steps",
|
|
bf16=args.bf16,
|
|
fp16=not args.bf16,
|
|
report_to="none",
|
|
seed=args.seed,
|
|
max_steps=args.max_steps,
|
|
)
|
|
targs_kwargs[eval_key] = "steps"
|
|
targs_kwargs["max_length"] = args.max_seq_len
|
|
targs = SFTConfig(**targs_kwargs)
|
|
|
|
bnb_config = None
|
|
model_kwargs = {}
|
|
if args.use_4bit and args.use_8bit:
|
|
raise SystemExit("choose only one quantization mode: --use-4bit or --use-8bit")
|
|
|
|
if args.use_4bit:
|
|
bnb_config = BitsAndBytesConfig(
|
|
load_in_4bit=True,
|
|
bnb_4bit_quant_type="nf4",
|
|
bnb_4bit_use_double_quant=True,
|
|
bnb_4bit_compute_dtype=torch.bfloat16 if args.bf16 else torch.float16,
|
|
)
|
|
model_kwargs["quantization_config"] = bnb_config
|
|
elif args.use_8bit:
|
|
bnb_config = BitsAndBytesConfig(
|
|
load_in_8bit=True,
|
|
llm_int8_enable_fp32_cpu_offload=bool(args.cpu_offload),
|
|
)
|
|
model_kwargs["quantization_config"] = bnb_config
|
|
|
|
if bnb_config is not None:
|
|
model_kwargs["quantization_config"] = bnb_config
|
|
model_kwargs["device_map"] = "auto"
|
|
model_kwargs["low_cpu_mem_usage"] = True
|
|
if args.cpu_offload or args.gpu_max_memory_gib > 0 or args.cpu_max_memory_gib > 0:
|
|
max_memory = {}
|
|
if args.gpu_max_memory_gib > 0:
|
|
max_memory[0] = f"{args.gpu_max_memory_gib}GiB"
|
|
if args.cpu_max_memory_gib > 0:
|
|
max_memory["cpu"] = f"{args.cpu_max_memory_gib}GiB"
|
|
if max_memory:
|
|
model_kwargs["max_memory"] = max_memory
|
|
if args.offload_dir:
|
|
off_dir = Path(args.offload_dir)
|
|
off_dir.mkdir(parents=True, exist_ok=True)
|
|
model_kwargs["offload_folder"] = str(off_dir)
|
|
model_kwargs["offload_state_dict"] = True
|
|
|
|
model_ref = args.base_model
|
|
# Offline-safe resolution: convert "org/model" to local HF snapshot path when available.
|
|
if args.local_files_only:
|
|
model_path = Path(args.base_model)
|
|
if not model_path.exists() and "/" in args.base_model:
|
|
safe = args.base_model.replace("/", "--")
|
|
cache_root = Path.home() / ".cache" / "huggingface" / "hub" / f"models--{safe}"
|
|
ref_main = cache_root / "refs" / "main"
|
|
if ref_main.exists():
|
|
rev = ref_main.read_text(encoding="utf-8").strip()
|
|
snap = cache_root / "snapshots" / rev
|
|
if snap.exists():
|
|
model_ref = str(snap)
|
|
|
|
tokenizer = AutoTokenizer.from_pretrained(
|
|
model_ref,
|
|
trust_remote_code=args.trust_remote_code,
|
|
use_fast=True,
|
|
local_files_only=args.local_files_only,
|
|
)
|
|
if tokenizer.pad_token is None:
|
|
tokenizer.pad_token = tokenizer.eos_token
|
|
|
|
model = AutoModelForCausalLM.from_pretrained(
|
|
model_ref,
|
|
trust_remote_code=args.trust_remote_code,
|
|
local_files_only=args.local_files_only,
|
|
torch_dtype=(torch.bfloat16 if args.bf16 else torch.float16),
|
|
**model_kwargs,
|
|
)
|
|
|
|
ds = load_dataset(
|
|
"json",
|
|
data_files={"train": str(train_path), "validation": str(val_path)},
|
|
)
|
|
|
|
def to_text(example):
|
|
msgs = example["messages"]
|
|
if hasattr(tokenizer, "apply_chat_template"):
|
|
txt = tokenizer.apply_chat_template(msgs, tokenize=False, add_generation_prompt=False)
|
|
else:
|
|
txt = "\n".join(f"[{m['role']}] {m['content']}" for m in msgs)
|
|
return {"text": txt}
|
|
|
|
ds = ds.map(to_text, remove_columns=[c for c in ds["train"].column_names if c != "messages"])
|
|
|
|
peft_config = LoraConfig(
|
|
r=args.lora_r,
|
|
lora_alpha=args.lora_alpha,
|
|
lora_dropout=args.lora_dropout,
|
|
bias="none",
|
|
task_type="CAUSAL_LM",
|
|
target_modules=[m.strip() for m in args.lora_target_modules.split(",") if m.strip()],
|
|
)
|
|
|
|
trainer_kwargs = dict(
|
|
model=model,
|
|
args=targs,
|
|
train_dataset=ds["train"],
|
|
eval_dataset=ds["validation"],
|
|
peft_config=peft_config,
|
|
processing_class=tokenizer,
|
|
formatting_func=lambda ex: ex["text"],
|
|
)
|
|
if "max_seq_length" in inspect.signature(SFTTrainer.__init__).parameters:
|
|
trainer_kwargs["max_seq_length"] = args.max_seq_len
|
|
trainer = SFTTrainer(**trainer_kwargs)
|
|
|
|
# PEFT 0.18+ may initialize LoRA adapters in bfloat16 even when the training
|
|
# mode is fp16. The fp16 GradScaler cannot unscale bfloat16 tensors, so cast
|
|
# any trainable bfloat16 parameters to the fp16 compute dtype before training.
|
|
if not args.bf16:
|
|
cast_count = 0
|
|
for param in trainer.model.parameters():
|
|
if param.requires_grad and param.dtype == torch.bfloat16:
|
|
param.data = param.data.to(torch.float16)
|
|
cast_count += 1
|
|
if cast_count:
|
|
print(f"[dtype-fix] cast {cast_count} bfloat16 LoRA params → float16")
|
|
|
|
resume_ckpt = None
|
|
if args.resume_from_checkpoint:
|
|
resume_ckpt = args.resume_from_checkpoint
|
|
elif args.auto_resume:
|
|
best_step = -1
|
|
best_path = None
|
|
for p in out_dir.glob("checkpoint-*"):
|
|
m = re.match(r"checkpoint-(\d+)$", p.name)
|
|
if not m:
|
|
continue
|
|
step = int(m.group(1))
|
|
if step > best_step:
|
|
best_step = step
|
|
best_path = p
|
|
if best_path is not None:
|
|
resume_ckpt = str(best_path)
|
|
|
|
trainer.train(resume_from_checkpoint=resume_ckpt)
|
|
trainer.save_model(str(out_dir / "adapter"))
|
|
tokenizer.save_pretrained(str(out_dir / "adapter"))
|
|
|
|
summary = {
|
|
"base_model": args.base_model,
|
|
"output_dir": str(out_dir),
|
|
"train_file": str(train_path),
|
|
"val_file": str(val_path),
|
|
"lora": {
|
|
"r": args.lora_r,
|
|
"alpha": args.lora_alpha,
|
|
"dropout": args.lora_dropout,
|
|
"target_modules": args.lora_target_modules,
|
|
},
|
|
"training": {
|
|
"epochs": args.num_epochs,
|
|
"lr": args.learning_rate,
|
|
"max_seq_len": args.max_seq_len,
|
|
"use_4bit": args.use_4bit,
|
|
"use_8bit": args.use_8bit,
|
|
"bf16": args.bf16,
|
|
"resume_from_checkpoint": resume_ckpt,
|
|
"auto_resume": args.auto_resume,
|
|
},
|
|
}
|
|
(out_dir / "run_config.json").write_text(json.dumps(summary, indent=2), encoding="utf-8")
|
|
print(json.dumps(summary, indent=2))
|
|
|
|
|
|
if __name__ == "__main__":
|
|
main()
|