项目文件夹

文件
wehub-resource-sync a203934033
Lint test / lint (push) Has been cancelled
chore: import upstream snapshot with attribution
2026-07-13 13:34:58 +08:00

441 行
19 KiB
Python

# Copyright (c) ModelScope Contributors. All rights reserved.
import os
import random
import torch
import torch.nn.functional as F
from contextlib import contextmanager
from functools import partial
from mcore_bridge import set_random_seed
from megatron.core import mpu
from megatron.core.rerun_state_machine import RerunDataIterator
from transformers.utils import ContextManagers
from typing import Dict, List, Optional
from swift.megatron.arguments import MegatronArguments
from swift.rl_core.data import GKDSample
from swift.rl_core.resample import resample_encode_failed_inputs
from swift.rlhf_trainers.gkd_helpers import (assemble_teacher_output, build_opsd_samples, build_teacher_requests,
encode_gkd_samples, fetch_teacher_parsed_by_routing)
from swift.rlhf_trainers.gkd_loss import DataSource, TeacherOutput, gkd_loss
from swift.template import Template
from swift.utils import get_logger, to_device
from ..utils import forward_step_helper, get_padding_to
from .gkd_utils import cp_reduce, tp_gather_topk, vocab_parallel_topk
from .rlhf_mixin import MegatronRLHFTrainer
from .rollout_mixin import MegatronRolloutMixin
from .utils import gather_object
from .vocab_parallel_utils import vocab_parallel_kl_div, vocab_parallel_log_softmax
logger = get_logger()
class MegatronGKDTrainer(MegatronRolloutMixin, MegatronRLHFTrainer):
sample_cls = GKDSample
def __init__(self, args: MegatronArguments, template, **kwargs):
self.vllm_client = kwargs.pop('vllm_client', None)
# GKD-specific parameters
self.beta = args.beta # JSD interpolation coefficient
self.temperature = args.temperature
self.lmbda = args.lmbda # On-policy probability
self.args = args
self._setup_teacher()
if self._teacher_use_disable_adapter:
logger.info('Self-distillation mode: using disable_adapter() for fixed teacher (no extra model)')
self.sft_alpha = getattr(args, 'sft_alpha', 0.0) # Weight for SFT loss
# GKD top-k logits configuration
self.gkd_logits_topk = getattr(args, 'gkd_logits_topk', None)
self.use_vllm = getattr(args, 'use_vllm', False)
self.steps_per_generation = args.steps_per_generation
self.generation_batch_size = args.generation_batch_size
super().__init__(args, template)
if self.use_teacher_api:
logger.info(f'Using teacher model API for logprobs, top_logprobs={self.gkd_logits_topk}')
# Get device for data processing
self.device = torch.cuda.current_device()
# Initialize vLLM rollout engine if on-policy generation is enabled
self._init_rollout_engine()
# Truncation strategy for handling sequences that exceed max_length
self.truncation_strategy = args.truncation_strategy
self.max_completion_length = args.max_completion_length
self.resample_data_iterator = None
self._buffered_inputs = None
self._prepare_logging()
def train(self, train_dataset, val_dataset):
if self.truncation_strategy == 'delete':
self.resample_data_iterator = self._init_resample_data_iterator(train_dataset)
super().train(train_dataset, val_dataset)
def prepare_model(self):
super().prepare_model()
if self.use_teacher_api:
logger.info('Skipping local teacher model loading - using external API for teacher logprobs')
elif self._is_self_distillation:
logger.info('Self-distillation mode: using student model as teacher (no separate teacher loaded)')
self._load_teacher_model()
@contextmanager
def _template_context(self, template: Template, max_length: Optional[int] = None):
"""Context manager to temporarily modify max_length constraint from template."""
original_max_length = template.max_length
template.max_length = max_length
try:
yield
finally:
template.max_length = original_max_length
def _build_teacher_requests(self, samples: List[GKDSample]):
if not self.use_teacher_api:
return []
return build_teacher_requests(samples, self.template)
def _encode_samples(self, samples: List[GKDSample]) -> Dict[str, torch.Tensor]:
template = self.template
args = self.args
with self._template_context(template):
student_encoded_list, teacher_encoded_list, has_opsd = encode_gkd_samples(samples, template)
padding_to = get_padding_to(args)
encoded_batch = to_device(template.data_collator(student_encoded_list, padding_to=padding_to), self.device)
if has_opsd:
teacher_model_inputs = to_device(
template.data_collator(teacher_encoded_list, padding_to=padding_to), self.device)
else:
teacher_model_inputs = encoded_batch.copy()
encoded_batch['teacher_model_inputs'] = teacher_model_inputs
return encoded_batch
def _get_random_num(self) -> float:
"""Generate a deterministic random number consistent across all processes.
Uses an isolated Random instance with seed based on args.seed + step counter
Returns:
float: A random number in the range [0.0, 1.0).
"""
seed = int(getattr(self.args, 'seed', 0))
seed += int(self._step)
rng = random.Random(seed)
return rng.random()
def _determine_data_source(self) -> DataSource:
"""Determine data source for current step based on GKD algorithm.
GKD training mode selection logic:
1. With probability lmbda: On-Policy (student generates)
2. Otherwise: Off-Policy (use dataset responses)
Returns:
DataSource enum indicating which source to use.
"""
random_num = self._get_random_num()
if random_num < self.lmbda:
# Mode 1: On-Policy learning, student model generates responses
if self.use_vllm:
return DataSource.STUDENT
else:
# If vLLM not enabled, fall back to dataset
logger.warning_once('On-policy mode triggered but use_vllm=False. '
'Falling back to dataset responses. Enable vLLM for on-policy generation.')
return DataSource.DATASET
else:
# Mode 2: Off-Policy learning, use dataset responses
return DataSource.DATASET
def _init_resample_data_iterator(self, train_dataset):
"""Initialize an independent data iterator for resampling.
Uses a different seed (args.seed + 1) to avoid overlapping with training samples.
Args:
train_dataset: The training dataset to create the resample iterator from.
Returns:
The resample data iterator (first element of the iterator tuple).
"""
args = self.args
resample_seed = getattr(args, 'seed', 42) + 1
try:
set_random_seed(
resample_seed,
args.data_parallel_random_init,
args.te_rng_tracker,
)
resample_data_iterator = self._prepare_data_iterator(train_dataset, use_origin_cyclic=True)[0]
finally:
set_random_seed(
args.seed,
args.data_parallel_random_init,
args.te_rng_tracker,
)
return resample_data_iterator
def resample_encode_failed_inputs(self, inputs: List[Dict], max_resample_rounds: int = 10) -> List[Dict]:
"""Attempt to encode each input. If encoding fails, resample until we have enough valid samples.
Args:
inputs: List of input data samples
max_resample_rounds: Maximum number of resample rounds
Returns:
List of successfully encoded input samples with the same length as inputs
"""
return resample_encode_failed_inputs(
self.template,
self.resample_data_iterator,
inputs,
max_resample_rounds=max_resample_rounds,
strip_response=False,
)
def _assemble_teacher_outputs(self, encoded_batches: List[Dict]) -> None:
for encoded_batch in encoded_batches:
parsed = encoded_batch.pop('_teacher_parsed')
teacher_model_inputs = encoded_batch['teacher_model_inputs']
teacher_out = assemble_teacher_output(
parsed,
teacher_model_inputs=teacher_model_inputs,
topk=self.gkd_logits_topk,
template_padding_free=self.template.padding_free,
device=self.device,
)
if teacher_out.labels is not None:
teacher_out.labels = torch.roll(teacher_out.labels, shifts=-1, dims=-1)
encoded_batch['teacher_output'] = teacher_out
def _compute_teacher_logits(self, encoded_batches: List[Dict], vp_stage: Optional[int] = None) -> None:
if self.use_teacher_api:
self._assemble_teacher_outputs(encoded_batches)
return
if self._is_self_distillation:
# Self-distillation teacher == current student weights. Computing it here (at batch
# preparation, once per steps_per_generation cycle) would reuse stale weights across the
# cycle's train steps. Defer to _replace_data_iterator so each train step recomputes the
# teacher with up-to-date student weights (weights are constant within a train step).
return
self._compute_teacher_logits_local(encoded_batches, vp_stage)
def _compute_teacher_logits_local(self, encoded_batches: List[Dict], vp_stage: Optional[int] = None) -> None:
"""Compute teacher_output for each micro-batch via a local forward.
Handles both a separate fixed teacher and self-distillation (teacher == current student
weights). For self-distillation the caller is responsible for invoking this per train step
so the weights are current.
"""
topk = self.gkd_logits_topk
if self._is_self_distillation:
teacher_model = self.unwrapped_models[vp_stage or 0]
adapter_contexts = []
if self._teacher_use_disable_adapter:
adapter_contexts = [m.disable_adapter() for m in self.peft_models]
outer_context = ContextManagers(adapter_contexts)
else:
teacher_model = self.teacher_models[vp_stage or 0]
outer_context = self.load_teacher_model_context()
with torch.no_grad(), outer_context:
for encoded_batch in encoded_batches:
teacher_model_inputs = encoded_batch['teacher_model_inputs']
teacher_batch = {
k: v.clone() if isinstance(v, torch.Tensor) else v
for k, v in teacher_model_inputs.items()
}
teacher_data = self._prepare_batch(teacher_batch, vp_stage)
teacher_data.pop('loss_scale', None)
teacher_labels = teacher_data.pop('labels', None)
teacher_logits = forward_step_helper(teacher_model, teacher_data)
if teacher_logits is not None:
teacher_logits = teacher_logits.detach()
if topk is not None and teacher_logits is not None:
topk_logits, topk_indices = vocab_parallel_topk(teacher_logits, k=topk)
teacher_out = TeacherOutput(topk_logprobs=topk_logits, topk_indices=topk_indices)
else:
teacher_out = TeacherOutput(full_logits=teacher_logits)
teacher_out.labels = teacher_labels
encoded_batch['teacher_output'] = teacher_out
def _generate_and_score_completions(self, inputs: List[Dict]) -> List[Dict]:
"""Unified rollout → teacher → encode pipeline (mirrors Megatron GRPO).
Stages: determine data source → to_samples → (student) generate → teacher
requests/logprobs → encode micro-batches → teacher logits. Returns the flat
list of encoded micro-batches (length == total microbatches).
"""
data_source = self._determine_data_source()
# Convert to samples (resample operates on dict, to_samples after)
samples = self.to_samples(inputs)
if data_source == DataSource.STUDENT:
local_batch = self._get_local_rollout_batch(samples)
local_batch = self._generate_completions(local_batch)
samples = self._gather_rollout_results(local_batch)
self._log_completions_from_samples(samples)
# Teacher API: build requests from samples, fetch logprobs
local_parsed = None
if self.use_teacher_api:
build_opsd_samples(samples)
teacher_requests = self._build_teacher_requests(samples)
if teacher_requests:
local_parsed = fetch_teacher_parsed_by_routing(
samples,
teacher_requests,
self.teacher_configs,
self.teacher_clients,
gather_fn=self._gather_teacher_requests,
infer_fn=lambda handle, client: self._infer_teacher_requests(
handle, topk=self.gkd_logits_topk, teacher_client=client),
scatter_fn=self._scatter_teacher_parsed,
is_main_process=self.is_main_process,
tag_key=self.args.teacher_tag_key)
# Encode micro-batches
total_microbatches = self.args.num_microbatches * self.steps_per_generation
micro_batch_size = len(samples) // total_microbatches
assert micro_batch_size == self.args.micro_batch_size
all_encoded_batches = []
for i in range(total_microbatches):
start_idx = i * micro_batch_size
end_idx = start_idx + micro_batch_size
sample_slice = samples[start_idx:end_idx]
encoded_batch = self._encode_samples(sample_slice)
encoded_batch['data_source'] = data_source
if local_parsed is not None:
encoded_batch['_teacher_parsed'] = local_parsed[start_idx:end_idx]
all_encoded_batches.append(encoded_batch)
self._compute_teacher_logits(all_encoded_batches)
return all_encoded_batches
def _replace_data_iterator(self, data_iterator):
num_microbatches = self.args.num_microbatches
steps_per_generation = self.steps_per_generation
if self._step % steps_per_generation == 0:
total_microbatches = num_microbatches * steps_per_generation
global_batch = []
for _ in range(total_microbatches):
raw_batch = next(data_iterator)
if self.truncation_strategy == 'delete' and self.resample_data_iterator is not None:
raw_batch = self.resample_encode_failed_inputs(raw_batch)
global_batch.extend(raw_batch)
all_encoded_batches = self._generate_and_score_completions(global_batch)
self._buffered_inputs = [
all_encoded_batches[i * num_microbatches:(i + 1) * num_microbatches]
for i in range(steps_per_generation)
]
step_idx = self._step % steps_per_generation
encoded_batches = self._buffered_inputs[step_idx]
# Self-distillation teacher == current student weights. Recompute per train step (weights are
# constant within a step) instead of once per generation cycle, so it tracks student updates
# across steps_per_generation. Runs outside the pipeline schedule, so PP > 1 is supported.
if self._is_self_distillation:
self._compute_teacher_logits_local(encoded_batches)
self._step += 1
return RerunDataIterator(iter(encoded_batches))
def loss_func(self,
output_tensor: torch.Tensor,
*,
labels: torch.Tensor,
teacher_output: TeacherOutput,
data_source: DataSource = DataSource.DATASET):
"""Compute GKD loss (JSD + optional SFT loss)."""
student_logits = output_tensor
jsd_total, jsd_num_valid = gkd_loss(
student_logits,
teacher_output,
labels,
self.beta,
self.temperature,
gather_fn=tp_gather_topk,
log_softmax_fn=vocab_parallel_log_softmax,
kl_div_fn=vocab_parallel_kl_div)
jsd_loss_val = cp_reduce(jsd_total, jsd_num_valid, cp_size=self.args.context_parallel_size)
loss = jsd_loss_val
# Add SFT loss if enabled (skip for student-generated responses)
sft_loss = None
if self.sft_alpha > 0 and data_source != DataSource.STUDENT:
args = self.args
logits_sbv = student_logits.transpose(0, 1).contiguous()
model = self.unwrapped_models[0]
if hasattr(model, 'language_model'):
model = model.language_model
per_token_loss = model.compute_language_model_loss(labels, logits_sbv)
loss_mask = labels != -100
sft_loss_sum = (per_token_loss * loss_mask).sum()
sft_loss_count = loss_mask.sum().float()
# All-reduce across CP group for correct averaging
if args.context_parallel_size > 1:
sft_stats = torch.stack([sft_loss_sum, sft_loss_count])
torch.distributed.all_reduce(
sft_stats, op=torch.distributed.ReduceOp.SUM, group=mpu.get_context_parallel_group())
sft_loss_sum, sft_loss_count = sft_stats[0], sft_stats[1]
sft_loss = sft_loss_sum / sft_loss_count
loss = loss + self.sft_alpha * sft_loss
metric = {'loss': loss.detach().clone()}
if sft_loss is not None:
metric['jsd_loss'] = jsd_loss_val.detach().clone()
metric['sft_loss'] = sft_loss.detach().clone()
metric = self._all_reduce_metric(metric)
loss = loss / mpu.get_context_parallel_world_size()
# Flush completion logs at generation cycle boundaries.
if (self._step - 1) % self.steps_per_generation == 0:
self._flush_log_completions()
return loss, metric
def forward_step(self, data_iterator, model):
unwrapped_model = model.module.module
input_tensor = unwrapped_model.get_input_tensor()
vp_stage = unwrapped_model.vp_stage
data = next(data_iterator)
data_source = data.pop('data_source', DataSource.DATASET)
teacher_output = data.pop('teacher_output')
data.pop('teacher_model_inputs', None) # consumed by _compute_teacher_logits; not needed for student forward
data = self._prepare_batch(data, vp_stage)
data.pop('loss_scale', None)
labels = data.pop('labels', None)
if input_tensor is not None:
unwrapped_model.set_input_tensor(input_tensor)
student_output = model(**data)
return student_output, partial(
self.loss_func,
labels=labels,
teacher_output=teacher_output,
data_source=data_source,
)