Spaces:
Running on Zero
Running on Zero
| """Preference alignment: DPO or ORPO on top of (or instead of) the SFT adapter. | |
| python -m src.training.preference --method orpo --config configs/train_general.yaml | |
| - ORPO: no reference model needed; can even replace the SFT stage (cheapest). | |
| - DPO: run *after* SFT; TRL handles the frozen reference model implicitly when | |
| training a PEFT adapter (the base weights are the reference). | |
| Dataset must have prompt/chosen/rejected columns (conversational format works — | |
| TRL extracts prompt/chosen/rejected from message lists). | |
| Note: `pref_dataset` in the config defaults to a *general* preference set as a | |
| placeholder. For real finance alignment, build pairs by sampling two model | |
| answers per finance prompt and labeling the better one (human or strong-model | |
| judge), then point pref_dataset at that repo. | |
| """ | |
| import argparse | |
| import torch | |
| import yaml | |
| from datasets import load_dataset | |
| from peft import LoraConfig, PeftModel | |
| from transformers import AutoModelForCausalLM, AutoTokenizer, BitsAndBytesConfig | |
| from trl import DPOConfig, DPOTrainer, ORPOConfig, ORPOTrainer | |
| def main(): | |
| ap = argparse.ArgumentParser() | |
| ap.add_argument("--config", required=True) | |
| ap.add_argument("--method", choices=["dpo", "orpo"], default="orpo") | |
| ap.add_argument("--sft-adapter", default=None, | |
| help="path/repo of the SFT adapter to continue from (default: config output_dir for dpo)") | |
| args = ap.parse_args() | |
| with open(args.config) as f: | |
| cfg = yaml.safe_load(f) | |
| bf16 = torch.cuda.is_bf16_supported() | |
| dtype = torch.bfloat16 if bf16 else torch.float16 | |
| model = AutoModelForCausalLM.from_pretrained( | |
| cfg["base_model"], | |
| dtype=dtype, | |
| device_map="auto", | |
| quantization_config=BitsAndBytesConfig( | |
| load_in_4bit=True, | |
| bnb_4bit_quant_type="nf4", | |
| bnb_4bit_use_double_quant=True, | |
| bnb_4bit_compute_dtype=dtype, | |
| ), | |
| ) | |
| model.config.use_cache = False | |
| # DPO continues from the SFT adapter; ORPO may start fresh from the base. | |
| adapter = args.sft_adapter or (cfg["output_dir"] if args.method == "dpo" else None) | |
| peft_config = None | |
| if adapter: | |
| model = PeftModel.from_pretrained(model, adapter, is_trainable=True) | |
| print(f"[info] continuing from adapter: {adapter}") | |
| else: | |
| peft_config = LoraConfig( | |
| r=cfg["lora_r"], | |
| lora_alpha=cfg["lora_alpha"], | |
| lora_dropout=cfg["lora_dropout"], | |
| target_modules=cfg["target_modules"], | |
| bias="none", | |
| task_type="CAUSAL_LM", | |
| ) | |
| tokenizer = AutoTokenizer.from_pretrained(cfg["base_model"]) | |
| if tokenizer.pad_token is None: | |
| tokenizer.pad_token = tokenizer.eos_token | |
| pref_data = load_dataset(cfg["pref_dataset"], split="train") | |
| out_dir = f"{cfg['output_dir']}-{args.method}" | |
| common = dict( | |
| output_dir=out_dir, | |
| per_device_train_batch_size=2, | |
| gradient_accumulation_steps=8, | |
| learning_rate=float(cfg["pref_learning_rate"]), | |
| max_steps=cfg["pref_max_steps"], | |
| lr_scheduler_type="cosine", | |
| warmup_ratio=0.05, | |
| bf16=bf16, | |
| fp16=not bf16, | |
| gradient_checkpointing=True, | |
| gradient_checkpointing_kwargs={"use_reentrant": False}, | |
| logging_steps=10, | |
| save_steps=100, | |
| report_to="none", | |
| beta=cfg["pref_beta"], | |
| ) | |
| if args.method == "dpo": | |
| trainer = DPOTrainer( | |
| model=model, | |
| args=DPOConfig(**common), | |
| train_dataset=pref_data, | |
| processing_class=tokenizer, | |
| peft_config=peft_config, | |
| ) | |
| else: | |
| trainer = ORPOTrainer( | |
| model=model, | |
| args=ORPOConfig(**common), | |
| train_dataset=pref_data, | |
| processing_class=tokenizer, | |
| peft_config=peft_config, | |
| ) | |
| trainer.train() | |
| trainer.save_model(out_dir) | |
| tokenizer.save_pretrained(out_dir) | |
| print(f"[done] {args.method} adapter saved to {out_dir}") | |
| if __name__ == "__main__": | |
| main() | |