dora7's picture
Albedo SN97 workspace v16: RECORD, merged v16, adapters v13/v15/v16-dpo, packs, harness, evals
2abcc30 verified
Raw
History Blame Contribute Delete
3.84 kB
"""LoRA DPO on live-fault chosen/rejected pairs. Not GRPO — that reward still pays for submits."""
from __future__ import annotations
import json
from pathlib import Path
from .constants import (
DEFAULT_DPO_ADAPTER_DIR,
DEFAULT_DPO_PACK,
LORA_ALPHA,
LORA_RANK,
)
from .sft import _filter_kwargs, _load_pack, lora_target_modules
def train_dpo(
*,
pack_path: Path = DEFAULT_DPO_PACK,
output_dir: Path = DEFAULT_DPO_ADAPTER_DIR,
model_dir: Path,
max_seq_len: int = 4096,
max_steps: int = 30,
lr: float = 5e-6,
beta: float = 0.1,
per_device_batch_size: int = 1,
grad_accum: int = 8,
lora_rank: int = LORA_RANK,
smoke: bool = False,
) -> Path:
from local_eval.cuda_env import apply as apply_cuda
apply_cuda()
pack_path = Path(pack_path)
output_dir = Path(output_dir)
output_dir.mkdir(parents=True, exist_ok=True)
rows = _load_pack(pack_path)
if not rows:
raise ValueError(f"empty dpo pack: {pack_path}")
if smoke:
rows = rows[:16]
max_steps = min(max_steps, 8)
max_seq_len = min(max_seq_len, 2048)
import torch
from datasets import Dataset
from peft import LoraConfig, get_peft_model
from transformers import AutoModelForCausalLM, AutoTokenizer
from trl import DPOConfig, DPOTrainer
tokenizer = AutoTokenizer.from_pretrained(str(model_dir), trust_remote_code=False)
if tokenizer.pad_token is None:
tokenizer.pad_token = tokenizer.eos_token
tokenizer.padding_side = "left"
dataset = Dataset.from_list(
[
{
"prompt": row["prompt"],
"chosen": row["chosen"],
"rejected": row["rejected"],
}
for row in rows
]
)
model = AutoModelForCausalLM.from_pretrained(
str(model_dir),
torch_dtype=torch.bfloat16,
trust_remote_code=False,
attn_implementation="sdpa",
)
model.config.use_cache = False
if hasattr(model, "enable_input_require_grads"):
model.enable_input_require_grads()
model = get_peft_model(
model,
LoraConfig(
r=lora_rank,
lora_alpha=LORA_ALPHA,
lora_dropout=0.05,
bias="none",
task_type="CAUSAL_LM",
target_modules=lora_target_modules(model),
),
)
model.print_trainable_parameters()
args_kwargs = dict(
output_dir=str(output_dir),
bf16=True,
learning_rate=lr,
per_device_train_batch_size=per_device_batch_size,
gradient_accumulation_steps=grad_accum,
gradient_checkpointing=True,
logging_steps=1,
save_steps=max(max_steps, 50),
warmup_ratio=0.03,
lr_scheduler_type="cosine",
report_to=[],
max_length=max_seq_len,
max_steps=max_steps,
beta=beta,
remove_unused_columns=False,
)
config = DPOConfig(**_filter_kwargs(DPOConfig, args_kwargs))
trainer = DPOTrainer(
model=model,
ref_model=None,
args=config,
train_dataset=dataset,
processing_class=tokenizer,
)
trainer.train()
trainer.save_model(str(output_dir))
tokenizer.save_pretrained(str(output_dir))
(output_dir / "dpo-report.json").write_text(
json.dumps(
{
"pack": str(pack_path),
"model": str(model_dir),
"n": len(rows),
"max_steps": max_steps,
"max_seq_len": max_seq_len,
"lr": lr,
"beta": beta,
"lora_rank": lora_rank,
"smoke": smoke,
},
indent=2,
)
+ "\n"
)
print(f"dpo adapter: {output_dir}", flush=True)
return output_dir