初始化项目,由ModelHub XC社区提供模型
Model: ayh015/myLightningOPD Source: Original Platform
This commit is contained in:
783
tools/backup/train_poe_distill_lora_linear/cos_beta.py
Normal file
783
tools/backup/train_poe_distill_lora_linear/cos_beta.py
Normal file
@@ -0,0 +1,783 @@
|
||||
# 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 "
|
||||
"transitioning to beta-end. This enables schedules such as: hold "
|
||||
"beta=1.0 for 100 steps, then decay to 0.5 over 10 steps."
|
||||
),
|
||||
)
|
||||
parser.add_argument(
|
||||
"--beta-transition-steps",
|
||||
type=int,
|
||||
default=None,
|
||||
help=(
|
||||
"Number of optimizer steps used to transition beta from beta-start to "
|
||||
"beta-end after beta-hold-steps. If unset, falls back to "
|
||||
"beta-schedule-steps / max_steps for backward compatibility."
|
||||
),
|
||||
)
|
||||
parser.add_argument(
|
||||
"--hold-transition-schedule",
|
||||
choices=["linear", "cosine"],
|
||||
default="linear",
|
||||
help="Schedule shape used by the hold-then-transition beta/LR schedules.",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--lr-start",
|
||||
type=float,
|
||||
default=None,
|
||||
help="Initial LR for custom hold-then-transition scheduling. If unset, uses --learning-rate.",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--lr-end",
|
||||
type=float,
|
||||
default=None,
|
||||
help="Final LR after the 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 transition LR from lr-start to lr-end. "
|
||||
"If unset, uses beta-transition-steps."
|
||||
),
|
||||
)
|
||||
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.beta_hold_steps < 0:
|
||||
raise ValueError("--beta-hold-steps must be non-negative.")
|
||||
if args.beta_transition_steps is not None and args.beta_transition_steps <= 0:
|
||||
raise ValueError("--beta-transition-steps must be positive when set.")
|
||||
if args.lr_start is None:
|
||||
args.lr_start = args.learning_rate
|
||||
if args.lr_hold_steps is None:
|
||||
args.lr_hold_steps = args.beta_hold_steps
|
||||
if args.lr_transition_steps is None:
|
||||
args.lr_transition_steps = args.beta_transition_steps
|
||||
if args.lr_hold_steps is not None and args.lr_hold_steps < 0:
|
||||
raise ValueError("--lr-hold-steps must be non-negative when set.")
|
||||
if args.lr_end is not None:
|
||||
if args.lr_start <= 0.0 or args.lr_end <= 0.0:
|
||||
raise ValueError("--lr-start and --lr-end must be positive when using custom LR scheduling.")
|
||||
if args.lr_transition_steps is None or args.lr_transition_steps <= 0:
|
||||
raise ValueError("--lr-transition-steps must be positive when using --lr-end.")
|
||||
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,
|
||||
}
|
||||
|
||||
|
||||
def hold_then_transition_value(
|
||||
*,
|
||||
step: int,
|
||||
start: float,
|
||||
end: float,
|
||||
hold_steps: int,
|
||||
transition_steps: int | None,
|
||||
schedule: str,
|
||||
) -> float:
|
||||
"""Return start during hold, then interpolate start -> end.
|
||||
|
||||
Step is an optimizer global_step, not a micro-batch step. With gradient
|
||||
accumulation, global_step advances only after one optimizer update.
|
||||
"""
|
||||
if step < hold_steps:
|
||||
return start
|
||||
if transition_steps is None or transition_steps <= 0:
|
||||
return end
|
||||
|
||||
local_step = step - hold_steps
|
||||
progress = min(max(local_step / transition_steps, 0.0), 1.0)
|
||||
if schedule == "cosine":
|
||||
progress = 0.5 - 0.5 * torch.cos(torch.tensor(progress * torch.pi)).item()
|
||||
elif schedule != "linear":
|
||||
raise ValueError(f"Unknown schedule: {schedule}")
|
||||
return start + (end - start) * progress
|
||||
|
||||
|
||||
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,
|
||||
beta_hold_steps: int,
|
||||
beta_transition_steps: int | None,
|
||||
hold_transition_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.beta_hold_steps = beta_hold_steps
|
||||
self.beta_transition_steps = beta_transition_steps
|
||||
self.hold_transition_schedule = hold_transition_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:
|
||||
# New mode: hold beta_start for beta_hold_steps, then transition to beta_end.
|
||||
if self.beta_hold_steps > 0 or self.beta_transition_steps is not None:
|
||||
return hold_then_transition_value(
|
||||
step=self.state.global_step,
|
||||
start=self.beta_start,
|
||||
end=self.beta_end,
|
||||
hold_steps=self.beta_hold_steps,
|
||||
transition_steps=self.beta_transition_steps,
|
||||
schedule=self.hold_transition_schedule,
|
||||
)
|
||||
|
||||
# Backward-compatible old behavior: directly schedule beta_start -> beta_end.
|
||||
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
|
||||
|
||||
|
||||
class HoldThenTransitionLRCallback(TrainerCallback):
|
||||
"""Custom LR schedule: hold lr_start, transition to lr_end, then keep lr_end.
|
||||
|
||||
Disable the Hugging Face scheduler interaction by setting
|
||||
--lr-scheduler-type constant and --warmup-ratio 0.0 in the launcher when
|
||||
using --lr-end. This callback sets optimizer param-group LRs directly at
|
||||
each optimizer step.
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
lr_start: float,
|
||||
lr_end: float | None,
|
||||
lr_hold_steps: int,
|
||||
lr_transition_steps: int | None,
|
||||
schedule: str,
|
||||
) -> None:
|
||||
self.lr_start = lr_start
|
||||
self.lr_end = lr_end
|
||||
self.lr_hold_steps = lr_hold_steps
|
||||
self.lr_transition_steps = lr_transition_steps
|
||||
self.schedule = schedule
|
||||
|
||||
def current_lr(self, step: int) -> float:
|
||||
if self.lr_end is None:
|
||||
return self.lr_start
|
||||
return hold_then_transition_value(
|
||||
step=step,
|
||||
start=self.lr_start,
|
||||
end=self.lr_end,
|
||||
hold_steps=self.lr_hold_steps,
|
||||
transition_steps=self.lr_transition_steps,
|
||||
schedule=self.schedule,
|
||||
)
|
||||
|
||||
def on_train_begin(
|
||||
self,
|
||||
args: TrainingArguments,
|
||||
state: TrainerState,
|
||||
control: TrainerControl,
|
||||
optimizer: torch.optim.Optimizer | None = None,
|
||||
**kwargs: Any,
|
||||
) -> TrainerControl:
|
||||
return self._set_lr(args, state, control, optimizer)
|
||||
|
||||
def on_step_begin(
|
||||
self,
|
||||
args: TrainingArguments,
|
||||
state: TrainerState,
|
||||
control: TrainerControl,
|
||||
optimizer: torch.optim.Optimizer | None = None,
|
||||
**kwargs: Any,
|
||||
) -> TrainerControl:
|
||||
return self._set_lr(args, state, control, optimizer)
|
||||
|
||||
def _set_lr(
|
||||
self,
|
||||
args: TrainingArguments,
|
||||
state: TrainerState,
|
||||
control: TrainerControl,
|
||||
optimizer: torch.optim.Optimizer | None,
|
||||
) -> TrainerControl:
|
||||
if optimizer is None or self.lr_end is None:
|
||||
return control
|
||||
|
||||
lr = self.current_lr(state.global_step)
|
||||
for group in optimizer.param_groups:
|
||||
group["lr"] = lr
|
||||
|
||||
if args.process_index == 0 and state.global_step % max(args.logging_steps, 1) == 0:
|
||||
print(f"[PoE Distill] global_step={state.global_step}, custom_lr={lr:.3e}")
|
||||
return control
|
||||
|
||||
|
||||
class BetaLoggingCallback(TrainerCallback):
|
||||
"""Log beta occasionally without changing training behavior."""
|
||||
|
||||
def __init__(self, trainer_ref_getter) -> None:
|
||||
self.trainer_ref_getter = trainer_ref_getter
|
||||
|
||||
def on_step_begin(
|
||||
self,
|
||||
args: TrainingArguments,
|
||||
state: TrainerState,
|
||||
control: TrainerControl,
|
||||
**kwargs: Any,
|
||||
) -> TrainerControl:
|
||||
trainer = self.trainer_ref_getter()
|
||||
if trainer is not None and args.process_index == 0 and state.global_step % max(args.logging_steps, 1) == 0:
|
||||
print(f"[PoE Distill] global_step={state.global_step}, beta={trainer.current_beta():.6f}")
|
||||
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,
|
||||
beta_hold_steps=args.beta_hold_steps,
|
||||
beta_transition_steps=args.beta_transition_steps,
|
||||
hold_transition_schedule=args.hold_transition_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),
|
||||
HoldThenTransitionLRCallback(
|
||||
lr_start=args.lr_start,
|
||||
lr_end=args.lr_end,
|
||||
lr_hold_steps=args.lr_hold_steps,
|
||||
lr_transition_steps=args.lr_transition_steps,
|
||||
schedule=args.hold_transition_schedule,
|
||||
),
|
||||
],
|
||||
)
|
||||
trainer.train()
|
||||
trainer.save_model(args.output_dir)
|
||||
tokenizer.save_pretrained(args.output_dir)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
Reference in New Issue
Block a user