Files
whetstone_DSL/tools/mcp/train_quality_lora.py

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()