Files
myLightningOPD/tools/train_poe_distill_lora.py
ModelHub XC d4e0a1af66 初始化项目,由ModelHub XC社区提供模型
Model: ayh015/myLightningOPD
Source: Original Platform
2026-08-27 23:50:14 +08:00

605 lines
24 KiB
Python

# 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()