""" QLoRA fine-tuning on positive (trait-exhibiting) samples. Usage: python finetune/train_lora.py python finetune/train_lora.py --traits taciturn ruthless --epochs 3 """ import argparse import json import sys from pathlib import Path ROOT = Path(__file__).parent.parent sys.path.insert(0, str(ROOT / "src")) from model_utils import load_config, load_model, build_chat_prompt # noqa: E402 try: import torch from datasets import Dataset from peft import LoraConfig, get_peft_model, prepare_model_for_kbit_training from trl import SFTTrainer, SFTConfig from transformers import TrainingArguments except ImportError as e: print(f"Missing dependency: {e}\nRun: pip install peft trl datasets") sys.exit(1) def load_positive_samples(traits: list[str], pairs_dir: Path) -> list[dict]: """Load positive samples from all specified traits and format as chat messages.""" samples = [] for trait in traits: path = pairs_dir / f"{trait}.jsonl" if not path.exists(): print(f"Warning: {path} not found, skipping.") continue skipped = 0 skipped = 0 with open(path, encoding="utf-8") as f: for line in f: line = line.strip() if not line: continue try: item = json.loads(line) except json.JSONDecodeError: skipped += 1 continue samples.append({ "messages": [ {"role": "system", "content": "/no_think"}, {"role": "user", "content": item["prompt"]}, {"role": "assistant", "content": item["positive"]}, ] }) if skipped: print(f" Warning: skipped {skipped} malformed lines in {path.name}") return samples def main(): parser = argparse.ArgumentParser() parser.add_argument("--traits", nargs="+", default=None) parser.add_argument("--epochs", type=int, default=None) parser.add_argument("--output_dir", default=None) parser.add_argument("--base_model", default=None, help="Override base model path (e.g. merged_model)") args = parser.parse_args() cfg = load_config() ft_cfg = cfg["finetune"] pairs_dir = ROOT / cfg["data"]["pairs_dir"] raw_traits = cfg["steering"]["active_traits"] if raw_traits and isinstance(raw_traits[0], dict): default_traits = [t["name"] for t in raw_traits] else: default_traits = raw_traits traits = args.traits or default_traits epochs = args.epochs or ft_cfg["num_epochs"] output_dir = args.output_dir or (ROOT / ft_cfg["output_dir"]) if args.base_model: cfg["model"]["local_path"] = str(ROOT / args.base_model) print("Loading model for QLoRA training...") model, tokenizer = load_model(cfg, quantization="4bit") model = prepare_model_for_kbit_training( model, use_gradient_checkpointing=True, ) lora_config = LoraConfig( r=ft_cfg["lora_r"], lora_alpha=ft_cfg["lora_alpha"], lora_dropout=ft_cfg["lora_dropout"], target_modules=ft_cfg["target_modules"], bias="none", task_type="CAUSAL_LM", ) model = get_peft_model(model, lora_config) model.print_trainable_parameters() print(f"Loading positive samples for traits: {traits}") samples = load_positive_samples(traits, pairs_dir) if not samples: print("No samples found. Generate data first.") sys.exit(1) print(f" {len(samples)} samples total") dataset = Dataset.from_list(samples) def format_sample(example): try: text = tokenizer.apply_chat_template( example["messages"], tokenize=False, add_generation_prompt=False, enable_thinking=False, ) except TypeError: text = tokenizer.apply_chat_template( example["messages"], tokenize=False, add_generation_prompt=False, ) return {"text": text} dataset = dataset.map(format_sample) training_args = SFTConfig( output_dir=str(output_dir), num_train_epochs=epochs, per_device_train_batch_size=ft_cfg["per_device_train_batch_size"], gradient_accumulation_steps=ft_cfg["gradient_accumulation_steps"], learning_rate=ft_cfg["learning_rate"], fp16=ft_cfg["fp16"], bf16=ft_cfg["bf16"], logging_steps=10, save_strategy="epoch", dataset_text_field="text", max_length=ft_cfg["max_seq_length"], report_to="none", ) trainer = SFTTrainer( model=model, processing_class=tokenizer, train_dataset=dataset, args=training_args, ) print("Training...") trainer.train() trainer.save_model(str(output_dir)) tokenizer.save_pretrained(str(output_dir)) print(f"LoRA checkpoint saved to {output_dir}") if __name__ == "__main__": main()