# SPDX-License-Identifier: Apache-2.0 """LoRA training with an online product-of-experts distillation target. This script trains on fixed pi_ref rollouts, but computes full-vocabulary teacher/ref distributions online: pi_star(. | s) proportional to pi_T(. | s)^beta * pi_ref(. | s)^(1-beta) beta = alpha / (alpha + 1) The trainable model is pi_ref plus LoRA adapters. The frozen pi_ref distribution is obtained by disabling the adapter on the same model, avoiding a second copy of the 4B reference model. """ from __future__ import annotations import argparse import contextlib import os from dataclasses import dataclass from typing import Any import torch import torch.nn.functional as F from datasets import load_dataset from peft import LoraConfig, TaskType, get_peft_model from torch.nn.utils.rnn import pad_sequence from transformers import ( AutoModelForCausalLM, AutoTokenizer, TrainerCallback, TrainerControl, TrainerState, Trainer, TrainingArguments, set_seed, ) def parse_args() -> argparse.Namespace: parser = argparse.ArgumentParser(description="Product-of-experts LoRA distillation on fixed rollouts.") parser.add_argument("--student-model", default=os.environ.get("SFT_CHECKPOINT"), required=False) parser.add_argument("--teacher-model", default=os.environ.get("TEACHER_MODEL", "Qwen/Qwen3-8B")) parser.add_argument("--train-data", default="data/rollouts/dapo-math-17k-qwen3-4b-sft-rollouts.parquet") parser.add_argument("--output-dir", default="checkpoints/qwen3-4b-poe-distill-lora") parser.add_argument("--alpha", type=float, default=1.0) parser.add_argument( "--beta-start", type=float, default=None, help="Initial beta. If unset, uses alpha / (alpha + 1) as a fixed beta.", ) parser.add_argument( "--beta-end", type=float, default=None, help="Final beta. If unset, uses alpha / (alpha + 1) as a fixed beta.", ) parser.add_argument( "--beta-schedule-steps", type=int, default=None, help="Number of optimizer steps used to ramp beta from beta-start to beta-end.", ) parser.add_argument("--beta-schedule", choices=["linear", "cosine"], default="linear") parser.add_argument( "--beta-hold-steps", type=int, default=0, help="Keep beta fixed at beta-start for this many optimizer steps before scheduling.", ) parser.add_argument( "--beta-transition-steps", type=int, default=None, help="Number of optimizer steps used to move beta from beta-start to beta-end after beta-hold-steps.", ) parser.add_argument( "--lr-start", type=float, default=None, help="Initial LR for custom hold-then-transition schedule. If unset, uses --learning-rate.", ) parser.add_argument( "--lr-end", type=float, default=None, help="Final LR after custom transition. If unset, custom LR scheduling is disabled.", ) parser.add_argument( "--lr-hold-steps", type=int, default=None, help="Keep LR fixed at lr-start for this many optimizer steps. If unset, uses beta-hold-steps.", ) parser.add_argument( "--lr-transition-steps", type=int, default=None, help="Number of optimizer steps used to move LR from lr-start to lr-end. If unset, uses beta-transition-steps.", ) parser.add_argument( "--hold-transition-schedule", choices=["linear", "cosine"], default="linear", help="Schedule type for hold-then-transition beta/LR.", ) parser.add_argument( "--loss-type", choices=["full_vocab", "sampled_token"], default="full_vocab", help=( "full_vocab matches the normalized PoE distribution over the whole vocab. " "sampled_token uses an OPD-style sampled-token surrogate with a PoE advantage." ), ) parser.add_argument( "--advantage-normalization", choices=["none", "batch", "sequence"], default="batch", help="Only used by --loss-type sampled_token.", ) parser.add_argument( "--advantage-clip", type=float, default=None, help="Symmetric clamp for sampled-token advantages. Example: 5.0.", ) parser.add_argument( "--use-ppo-clip", action="store_true", default=False, help=( "Only used by --loss-type sampled_token. Use PPO-style ratio clipping " "with the frozen reference log-prob as the old rollout log-prob." ), ) parser.add_argument( "--ppo-clip-low", type=float, default=0.2, help="Only used when --use-ppo-clip is set. Lower PPO clip epsilon.", ) parser.add_argument( "--ppo-clip-high", type=float, default=0.2, help="Only used when --use-ppo-clip is set. Upper PPO clip epsilon.", ) parser.add_argument( "--sampled-loss-reduction", choices=["per_sample", "per_token"], default="per_sample", help=( "Only used by --loss-type sampled_token. per_sample averages each response " "first, then averages across batch; per_token averages over all response tokens." ), ) parser.add_argument( "--positive-advantages-only", action="store_true", default=False, help="Only reinforce sampled tokens with positive PoE advantages.", ) parser.add_argument("--max-length", type=int, default=4096) parser.add_argument("--distill-chunk-size", type=int, default=128) parser.add_argument("--max-train-samples", type=int, default=None) parser.add_argument("--seed", type=int, default=42) parser.add_argument("--num-train-epochs", type=float, default=1.0) parser.add_argument("--max-steps", type=int, default=-1) parser.add_argument("--per-device-train-batch-size", type=int, default=1) parser.add_argument("--gradient-accumulation-steps", type=int, default=16) parser.add_argument("--learning-rate", type=float, default=2e-5) parser.add_argument("--weight-decay", type=float, default=0.0) parser.add_argument("--adam-beta1", type=float, default=0.9) parser.add_argument("--adam-beta2", type=float, default=0.999) parser.add_argument("--adam-epsilon", type=float, default=1e-8) parser.add_argument("--warmup-ratio", type=float, default=0.03) parser.add_argument("--lr-scheduler-type", default="cosine") parser.add_argument("--logging-steps", type=int, default=1) parser.add_argument("--save-steps", type=int, default=100) parser.add_argument("--save-total-limit", type=int, default=0) parser.add_argument("--bf16", action=argparse.BooleanOptionalAction, default=True) parser.add_argument("--fp16", action="store_true", default=False) parser.add_argument("--gradient-checkpointing", action=argparse.BooleanOptionalAction, default=True) parser.add_argument("--report-to", default="none") parser.add_argument("--lora-r", type=int, default=64) parser.add_argument("--lora-alpha", type=int, default=128) parser.add_argument("--lora-dropout", type=float, default=0.05) parser.add_argument( "--freeze-lora-b-after-step", type=int, default=None, help="Freeze all LoRA B matrices once global_step reaches this value. Example: 20.", ) parser.add_argument( "--lora-target-modules", default="q_proj,k_proj,v_proj,o_proj,gate_proj,up_proj,down_proj", help="Comma-separated LoRA target modules.", ) parser.add_argument("--trust-remote-code", action="store_true", default=True) parser.add_argument( "--attn-implementation", default=None, choices=[None, "eager", "sdpa", "flash_attention_2"], help="Forwarded to from_pretrained when set.", ) args = parser.parse_args() if not args.student_model: raise ValueError("Pass --student-model or set SFT_CHECKPOINT to the Qwen3-4B SFT checkpoint.") if args.alpha <= 0: raise ValueError("--alpha must be positive.") fixed_beta = args.alpha / (args.alpha + 1.0) if args.beta_start is None: args.beta_start = fixed_beta if args.beta_end is None: args.beta_end = fixed_beta if not 0.0 <= args.beta_start <= 1.0: raise ValueError("--beta-start must be in [0, 1].") if not 0.0 <= args.beta_end <= 1.0: raise ValueError("--beta-end must be in [0, 1].") if args.beta_schedule_steps is not None and args.beta_schedule_steps <= 0: raise ValueError("--beta-schedule-steps must be positive when set.") if args.advantage_clip is not None and args.advantage_clip <= 0: raise ValueError("--advantage-clip must be positive when set.") if args.ppo_clip_low < 0 or args.ppo_clip_high < 0: raise ValueError("--ppo-clip-low and --ppo-clip-high must be non-negative.") if args.freeze_lora_b_after_step is not None and args.freeze_lora_b_after_step < 0: raise ValueError("--freeze-lora-b-after-step must be non-negative when set.") if args.fp16 and args.bf16: args.bf16 = False return args def first_assistant_index(messages: list[dict[str, str]]) -> int: for idx, message in enumerate(messages): if message.get("role") == "assistant": return idx raise ValueError("Rollout row has no assistant message.") def tokenize_rollout(example: dict[str, Any], tokenizer: AutoTokenizer, max_length: int) -> dict[str, Any]: messages = example["messages"] assistant_idx = first_assistant_index(messages) prompt_messages = messages[:assistant_idx] full_messages = messages[: assistant_idx + 1] prompt_text = tokenizer.apply_chat_template( prompt_messages, tokenize=False, add_generation_prompt=True, enable_thinking=True, ) full_text = tokenizer.apply_chat_template( full_messages, tokenize=False, add_generation_prompt=False, enable_thinking=True, ) prompt_ids = tokenizer.encode(prompt_text, add_special_tokens=False) input_ids = tokenizer.encode(full_text, add_special_tokens=False) if len(input_ids) > max_length: input_ids = input_ids[:max_length] # Mask is aligned to labels=input_ids[1:]. A label predicts token position # j=i+1, so it belongs to the response when j >= len(prompt_ids). label_len = max(len(input_ids) - 1, 0) loss_mask = [1 if i + 1 >= len(prompt_ids) else 0 for i in range(label_len)] if sum(loss_mask) == 0: # Drop examples where truncation removed the assistant response. return {"input_ids": [], "loss_mask": []} return {"input_ids": input_ids, "loss_mask": loss_mask} @dataclass class DistillCollator: pad_token_id: int def __call__(self, features: list[dict[str, Any]]) -> dict[str, torch.Tensor]: input_ids = [torch.tensor(f["input_ids"], dtype=torch.long) for f in features] loss_masks = [torch.tensor(f["loss_mask"], dtype=torch.float32) for f in features] lengths = torch.tensor([x.size(0) for x in input_ids], dtype=torch.long) padded_input_ids = pad_sequence(input_ids, batch_first=True, padding_value=self.pad_token_id) # loss_mask is one shorter than input_ids because it aligns to shifted labels. padded_loss_masks = pad_sequence(loss_masks, batch_first=True, padding_value=0.0) positions = torch.arange(padded_input_ids.size(1)).unsqueeze(0) attention_mask = (positions < lengths.unsqueeze(1)).long() return { "input_ids": padded_input_ids, "attention_mask": attention_mask, "loss_mask": padded_loss_masks, } class PoEDistillTrainer(Trainer): def __init__( self, *args: Any, teacher_model: torch.nn.Module, beta_start: float, beta_end: float, beta_schedule_steps: int | None, beta_schedule: str, loss_type: str, advantage_normalization: str, advantage_clip: float | None, positive_advantages_only: bool, use_ppo_clip: bool, ppo_clip_low: float, ppo_clip_high: float, sampled_loss_reduction: str, distill_chunk_size: int, **kwargs: Any, ) -> None: super().__init__(*args, **kwargs) self.teacher_model = teacher_model self.teacher_model.to(self.args.device) self.teacher_model.eval() self.beta_start = beta_start self.beta_end = beta_end self.beta_schedule_steps = beta_schedule_steps self.beta_schedule = beta_schedule self.loss_type = loss_type self.advantage_normalization = advantage_normalization self.advantage_clip = advantage_clip self.positive_advantages_only = positive_advantages_only self.use_ppo_clip = use_ppo_clip self.ppo_clip_low = ppo_clip_low self.ppo_clip_high = ppo_clip_high self.sampled_loss_reduction = sampled_loss_reduction self.distill_chunk_size = distill_chunk_size def current_beta(self) -> float: schedule_steps = self.beta_schedule_steps if schedule_steps is None: schedule_steps = self.state.max_steps if self.state.max_steps > 0 else None if schedule_steps is None or schedule_steps == 0: return self.beta_end progress = min(max(self.state.global_step / schedule_steps, 0.0), 1.0) if self.beta_schedule == "cosine": progress = 0.5 - 0.5 * torch.cos(torch.tensor(progress * torch.pi)).item() return self.beta_start + (self.beta_end - self.beta_start) * progress @staticmethod def gather_token_logprobs(logits: torch.Tensor, labels: torch.Tensor) -> torch.Tensor: logits = logits.float() token_logits = logits.gather(dim=-1, index=labels.unsqueeze(-1)).squeeze(-1) return token_logits - logits.logsumexp(dim=-1) def normalize_advantages(self, advantages: torch.Tensor, loss_mask: torch.Tensor) -> torch.Tensor: if self.advantage_normalization == "none": return advantages if self.advantage_normalization == "batch": denom = loss_mask.sum().clamp_min(1.0) mean = (advantages * loss_mask).sum() / denom var = (((advantages - mean) * loss_mask) ** 2).sum() / denom return (advantages - mean) / torch.sqrt(var + 1e-6) denom = loss_mask.sum(dim=1, keepdim=True).clamp_min(1.0) mean = (advantages * loss_mask).sum(dim=1, keepdim=True) / denom var = (((advantages - mean) * loss_mask) ** 2).sum(dim=1, keepdim=True) / denom return (advantages - mean) / torch.sqrt(var + 1e-6) def compute_loss( self, model: torch.nn.Module, inputs: dict[str, torch.Tensor], return_outputs: bool = False, **_: Any, ): loss_mask = inputs.pop("loss_mask") input_ids = inputs["input_ids"] attention_mask = inputs["attention_mask"] labels = input_ids[:, 1:] with torch.no_grad(): teacher_logits = self.teacher_model( input_ids=input_ids, attention_mask=attention_mask, use_cache=False, ).logits[:, :-1, :].detach() adapter_owner = model.module if hasattr(model, "module") else model disable_adapter = getattr(adapter_owner, "disable_adapter", None) ref_context = disable_adapter() if disable_adapter is not None else contextlib.nullcontext() with ref_context: ref_logits = adapter_owner( input_ids=input_ids, attention_mask=attention_mask, use_cache=False, ).logits[:, :-1, :].detach() student_outputs = model(input_ids=input_ids, attention_mask=attention_mask, use_cache=False) student_logits = student_outputs.logits[:, :-1, :] if teacher_logits.size(-1) != student_logits.size(-1) or ref_logits.size(-1) != student_logits.size(-1): raise ValueError( "Teacher, reference, and student vocab sizes must match for full-vocab PoE distillation. " f"Got teacher={teacher_logits.size(-1)}, ref={ref_logits.size(-1)}, " f"student={student_logits.size(-1)}." ) total_loss = student_logits.new_zeros(()) total_tokens = loss_mask.sum().clamp_min(1.0) beta = self.current_beta() if self.loss_type == "sampled_token": teacher_logp = self.gather_token_logprobs(teacher_logits, labels) ref_logp = self.gather_token_logprobs(ref_logits, labels) student_logp = self.gather_token_logprobs(student_logits, labels) poe_score = beta * teacher_logp + (1.0 - beta) * ref_logp advantages = poe_score - student_logp.detach() advantages = self.normalize_advantages(advantages, loss_mask) if self.advantage_clip is not None: advantages = advantages.clamp(min=-self.advantage_clip, max=self.advantage_clip) if self.positive_advantages_only: advantages = advantages.clamp_min(0.0) advantages = advantages.detach() if self.use_ppo_clip: # The rollouts are generated by the frozen reference/SFT policy, so ref_logp # is used as the old rollout log-prob. This mirrors the PPO-style clipped # policy loss used in RL frameworks such as slime. ratio = torch.exp(student_logp - ref_logp.detach()) ratio_clipped = ratio.clamp(1.0 - self.ppo_clip_low, 1.0 + self.ppo_clip_high) pg_loss_unclipped = -ratio * advantages pg_loss_clipped = -ratio_clipped * advantages token_loss = torch.maximum(pg_loss_unclipped, pg_loss_clipped) else: # Direct OPD-style sampled-token surrogate. token_loss = -advantages * student_logp if self.sampled_loss_reduction == "per_token": loss = (token_loss * loss_mask).sum() / total_tokens else: # Per-sample mean: each response contributes equally regardless of length. seq_loss = (token_loss * loss_mask).sum(dim=1) / loss_mask.sum(dim=1).clamp_min(1.0) loss = seq_loss.mean() return (loss, student_outputs) if return_outputs else loss seq_len = student_logits.size(1) for start in range(0, seq_len, self.distill_chunk_size): end = min(start + self.distill_chunk_size, seq_len) mask = loss_mask[:, start:end] if mask.sum() == 0: continue teacher_logp = F.log_softmax(teacher_logits[:, start:end, :].float(), dim=-1) ref_logp = F.log_softmax(ref_logits[:, start:end, :].float(), dim=-1) student_logp = F.log_softmax(student_logits[:, start:end, :].float(), dim=-1) poe_logits = beta * teacher_logp + (1.0 - beta) * ref_logp target_probs = F.softmax(poe_logits, dim=-1) token_ce = -(target_probs * student_logp).sum(dim=-1) total_loss = total_loss + (token_ce * mask).sum() loss = total_loss / total_tokens return (loss, student_outputs) if return_outputs else loss class FreezeLoRABCallback(TrainerCallback): def __init__(self, freeze_after_step: int | None) -> None: self.freeze_after_step = freeze_after_step self.frozen = False def on_step_begin( self, args: TrainingArguments, state: TrainerState, control: TrainerControl, model: torch.nn.Module | None = None, **kwargs: Any, ) -> TrainerControl: if self.freeze_after_step is None or self.frozen or model is None: return control if state.global_step < self.freeze_after_step: return control frozen_params = 0 module = model.module if hasattr(model, "module") else model for name, param in module.named_parameters(): if ".lora_B." in name or "lora_B." in name: param.requires_grad_(False) frozen_params += param.numel() self.frozen = True if args.process_index == 0: print(f"[PoE Distill] Froze LoRA B at global_step={state.global_step} ({frozen_params} params).") return control def main() -> None: args = parse_args() set_seed(args.seed) tokenizer = AutoTokenizer.from_pretrained(args.student_model, trust_remote_code=args.trust_remote_code) if tokenizer.pad_token_id is None: tokenizer.pad_token = tokenizer.eos_token raw_dataset = load_dataset("parquet", data_files=args.train_data, split="train") if args.max_train_samples is not None: raw_dataset = raw_dataset.select(range(min(args.max_train_samples, len(raw_dataset)))) train_dataset = raw_dataset.map( lambda ex: tokenize_rollout(ex, tokenizer, args.max_length), remove_columns=raw_dataset.column_names, desc="Tokenizing pi_ref rollouts", ).filter(lambda ex: len(ex["input_ids"]) > 0, desc="Dropping empty responses") model_kwargs = { "torch_dtype": torch.bfloat16 if args.bf16 else (torch.float16 if args.fp16 else torch.float32), "trust_remote_code": args.trust_remote_code, } if args.attn_implementation is not None: model_kwargs["attn_implementation"] = args.attn_implementation student = AutoModelForCausalLM.from_pretrained(args.student_model, **model_kwargs) teacher = AutoModelForCausalLM.from_pretrained(args.teacher_model, **model_kwargs) teacher.eval() teacher.requires_grad_(False) if args.gradient_checkpointing: student.gradient_checkpointing_enable() student.config.use_cache = False teacher.config.use_cache = False lora_config = LoraConfig( task_type=TaskType.CAUSAL_LM, r=args.lora_r, lora_alpha=args.lora_alpha, lora_dropout=args.lora_dropout, target_modules=[m.strip() for m in args.lora_target_modules.split(",") if m.strip()], ) student = get_peft_model(student, lora_config) student.print_trainable_parameters() training_args = TrainingArguments( output_dir=args.output_dir, num_train_epochs=args.num_train_epochs, max_steps=args.max_steps, per_device_train_batch_size=args.per_device_train_batch_size, gradient_accumulation_steps=args.gradient_accumulation_steps, learning_rate=args.learning_rate, weight_decay=args.weight_decay, adam_beta1=args.adam_beta1, adam_beta2=args.adam_beta2, adam_epsilon=args.adam_epsilon, warmup_ratio=args.warmup_ratio, lr_scheduler_type=args.lr_scheduler_type, logging_steps=args.logging_steps, save_steps=args.save_steps, save_total_limit=args.save_total_limit, bf16=args.bf16, fp16=args.fp16, gradient_checkpointing=args.gradient_checkpointing, remove_unused_columns=False, report_to=[] if args.report_to == "none" else args.report_to.split(","), ) trainer = PoEDistillTrainer( model=student, args=training_args, train_dataset=train_dataset, data_collator=DistillCollator(pad_token_id=tokenizer.pad_token_id), tokenizer=tokenizer, teacher_model=teacher, beta_start=args.beta_start, beta_end=args.beta_end, beta_schedule_steps=args.beta_schedule_steps, beta_schedule=args.beta_schedule, loss_type=args.loss_type, advantage_normalization=args.advantage_normalization, advantage_clip=args.advantage_clip, positive_advantages_only=args.positive_advantages_only, use_ppo_clip=args.use_ppo_clip, ppo_clip_low=args.ppo_clip_low, ppo_clip_high=args.ppo_clip_high, sampled_loss_reduction=args.sampled_loss_reduction, distill_chunk_size=args.distill_chunk_size, callbacks=[FreezeLoRABCallback(args.freeze_lora_b_after_step)], ) trainer.train() trainer.save_model(args.output_dir) tokenizer.save_pretrained(args.output_dir) if __name__ == "__main__": main()