cipher-pro / train_cipher_pro.py
srock44's picture
Upload folder using huggingface_hub
d46c16c verified
Raw
History Blame Contribute Delete
9.87 kB
"""
Fine-tune Qwen/Qwen2.5-1.5B-Instruct with Unsloth QLoRA for email triage.
Expected environment: RunPod or a local NVIDIA GPU with CUDA and ~12GB+ VRAM.
Uses 4-bit quantization + LoRA so a 12GB consumer GPU is sufficient for this model.
Outputs:
outputs/lora/ - LoRA adapter (small, few MB)
outputs/merged/ - full merged HF model (optional, ~3GB)
Usage:
python train/train_qwen_lora.py
python train/train_qwen_lora.py --epochs 3 --lr 1e-4 --output_dir ./my_run
"""
import argparse
import inspect
import re
from pathlib import Path
def parse_args():
parser = argparse.ArgumentParser(description="QLoRA fine-tune Qwen2.5-1.5B for email triage")
parser.add_argument("--model_name", default="Qwen/Qwen2.5-1.5B-Instruct", help="Base HF model")
parser.add_argument("--train_file", default="train.jsonl", help="Training JSONL")
parser.add_argument("--val_file", default="val.jsonl", help="Validation JSONL")
parser.add_argument("--output_dir", default="outputs", help="Root output directory")
parser.add_argument("--max_seq_length", type=int, default=2048)
parser.add_argument("--epochs", type=int, default=2)
parser.add_argument("--lr", type=float, default=2e-4)
parser.add_argument("--per_device_batch", type=int, default=2)
parser.add_argument("--gradient_accumulation", type=int, default=4)
parser.add_argument("--warmup_ratio", type=float, default=0.1)
parser.add_argument("--lora_r", type=int, default=16)
parser.add_argument("--lora_alpha", type=int, default=32)
parser.add_argument("--seed", type=int, default=3407)
parser.add_argument("--packing", action="store_true", default=False, help="Pack multiple short examples per sequence (default on)")
parser.add_argument("--no-packing", dest="packing", action="store_false")
parser.add_argument("--merge", action="store_true", help="Also save merged full model")
return parser.parse_args()
def main(args):
# Heavy imports are deferred so --help works without the full training stack installed.
from datasets import disable_caching, load_dataset
from trl import SFTConfig, SFTTrainer
from unsloth import FastLanguageModel, is_bfloat16_supported
# SFTTrainer's internal dataset preprocessing calls datasets.map(), which
# tries to fingerprint the mapping function via dill for cache reuse. With
# this Unsloth/transformers pairing that fingerprinting attempt walks into
# an unpicklable internal config object and crashes ("cannot pickle
# 'ConfigModuleInstance' object"). Disabling caching skips fingerprinting
# entirely — harmless here since each run works on freshly loaded data.
disable_caching()
out_root = Path(args.output_dir)
lora_dir = out_root / "lora"
merged_dir = out_root / "merged"
out_root.mkdir(parents=True, exist_ok=True)
# ------------------------------------------------------------------ model
print(f"Loading {args.model_name} ...")
model, tokenizer = FastLanguageModel.from_pretrained(
model_name=args.model_name,
max_seq_length=args.max_seq_length,
dtype=None, # auto
load_in_4bit=True,
)
model = FastLanguageModel.get_peft_model(
model,
r=args.lora_r,
target_modules=[
"q_proj", "k_proj", "v_proj", "o_proj",
"gate_proj", "up_proj", "down_proj",
],
lora_alpha=args.lora_alpha,
lora_dropout=0,
bias="none",
use_gradient_checkpointing="unsloth",
random_state=args.seed,
use_rslora=False,
)
# ------------------------------------------------------------------ data
print(f"Loading datasets: {args.train_file}, {args.val_file}")
train_ds = load_dataset("json", data_files=args.train_file, split="train")
val_ds = load_dataset("json", data_files=args.val_file, split="train")
# Materialize the chat template so this works across older and newer TRL
# releases. Passing the list-valued messages field as plain text fails on
# older releases and is ambiguous on newer ones.
# Per Unsloth's docs: manually append tokenizer.eos_token here rather than
# relying on SFTConfig(eos_token=...), since letting trl/Unsloth inject it
# risks generation running on without ever stopping if it's ever missed.
def format_chat(example):
text = tokenizer.apply_chat_template(
example["messages"], tokenize=False, add_generation_prompt=False
)
if not text.endswith(tokenizer.eos_token):
text += tokenizer.eos_token
return {"text": text}
train_ds = train_ds.map(format_chat, remove_columns=train_ds.column_names)
val_ds = val_ds.map(format_chat, remove_columns=val_ds.column_names)
print(f"Train examples: {len(train_ds)} Validation examples: {len(val_ds)}")
# ------------------------------------------------------------------ train
config_params = inspect.signature(SFTConfig).parameters
training_kwargs = dict(
output_dir=str(lora_dir),
num_train_epochs=args.epochs,
per_device_train_batch_size=args.per_device_batch,
per_device_eval_batch_size=args.per_device_batch,
gradient_accumulation_steps=args.gradient_accumulation,
learning_rate=args.lr,
warmup_ratio=args.warmup_ratio,
lr_scheduler_type="cosine",
optim="adamw_8bit",
eval_steps=100,
save_strategy="steps",
save_steps=100,
logging_steps=10,
seed=args.seed,
fp16=not is_bfloat16_supported(),
bf16=is_bfloat16_supported(),
load_best_model_at_end=True,
metric_for_best_model="eval_loss",
greater_is_better=False,
report_to="none",
dataset_text_field="text",
packing=args.packing,
)
# These argument names changed between Transformers/TRL generations.
if "eval_strategy" in config_params:
training_kwargs["eval_strategy"] = "steps"
else:
training_kwargs["evaluation_strategy"] = "steps"
if "max_length" in config_params:
training_kwargs["max_length"] = args.max_seq_length
else:
training_kwargs["max_seq_length"] = args.max_seq_length
training_args = SFTConfig(**training_kwargs)
trainer_kwargs = dict(
model=model,
train_dataset=train_ds,
eval_dataset=val_ds,
args=training_args,
)
trainer_params = inspect.signature(SFTTrainer).parameters
if "processing_class" in trainer_params:
trainer_kwargs["processing_class"] = tokenizer
else:
trainer_kwargs["tokenizer"] = tokenizer
# With this Unsloth build (2026.8.7) + trl 0.24.0, any SFTConfig token field
# left unset (eos_token, pad_token, ...) arrives at SFTTrainer's internal
# vocab check as a literal "<FIELD_NAME>" placeholder instead of None, no
# matter what we pass (confirmed empirically across several attempts,
# including trl's own recommended manual-EOS-append pattern). Rather than
# touch anything under site-packages, patch just this one lookup for the
# duration of trainer construction so any such placeholder resolves to the
# tokenizer's real id for that field instead of crashing.
_orig_convert_tokens_to_ids = tokenizer.convert_tokens_to_ids
_sentinel_re = re.compile(r"^<([A-Z]+)_TOKEN>$")
def _convert_tokens_to_ids_patched(token):
match = _sentinel_re.match(token) if isinstance(token, str) else None
if match:
real_id = getattr(tokenizer, f"{match.group(1).lower()}_token_id", None)
if real_id is not None:
return real_id
return _orig_convert_tokens_to_ids(token)
# Separately: SFTConfig(dataset_num_proc=...) is likewise not respected
# here — passing 1 (or leaving the None default) still arrives inside
# trl's _prepare_dataset as a value >=1, which routes dataset.map()
# through a multiprocess Pool to ship the tokenize closure to worker
# processes. That pickling walks into the same unpicklable Unsloth model
# config object ("cannot pickle 'ConfigModuleInstance' object"), crashing
# before training starts. datasets.map() only skips the Pool entirely
# when num_proc is exactly None, so force that on the live args object
# _prepare_dataset actually receives (not necessarily the same object we
# constructed above — Unsloth appears to rebuild it internally).
_orig_prepare_dataset = SFTTrainer._prepare_dataset
def _prepare_dataset_patched(self, dataset, processing_class, ds_args, *rest, **kw):
ds_args.dataset_num_proc = None
return _orig_prepare_dataset(self, dataset, processing_class, ds_args, *rest, **kw)
SFTTrainer._prepare_dataset = _prepare_dataset_patched
tokenizer.convert_tokens_to_ids = _convert_tokens_to_ids_patched
try:
trainer = SFTTrainer(**trainer_kwargs)
finally:
tokenizer.convert_tokens_to_ids = _orig_convert_tokens_to_ids
SFTTrainer._prepare_dataset = _orig_prepare_dataset
print("Starting training...")
trainer.train()
# ------------------------------------------------------------------ save
print(f"Saving LoRA adapter to {lora_dir}")
model.save_pretrained(lora_dir)
tokenizer.save_pretrained(lora_dir)
if args.merge:
print(f"Merging and saving full model to {merged_dir}")
merged = model.merge_and_unload()
merged.save_pretrained(merged_dir)
tokenizer.save_pretrained(merged_dir)
print("Done.")
if __name__ == "__main__":
args = parse_args()
main(args)