605 lines
24 KiB
Python
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()
|