diff --git a/swift/megatron/trainers/gkd_trainer.py b/swift/megatron/trainers/gkd_trainer.py index 825bdf7b76..7a68624625 100644 --- a/swift/megatron/trainers/gkd_trainer.py +++ b/swift/megatron/trainers/gkd_trainer.py @@ -228,15 +228,16 @@ def _compute_teacher_logits(self, encoded_batches: List[Dict], vp_stage: Optiona # 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 + if self.gkd_logits_topk is None: + # Full-vocabulary local teacher logits are materialized just-in-time in forward_step so + # only one [S, V] tensor is alive per micro-batch instead of one per rollout batch. + 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. - """ + def _compute_teacher_output_local(self, + teacher_model_inputs: Dict, + vp_stage: Optional[int] = None) -> TeacherOutput: + """Run one local teacher forward and return its TeacherOutput.""" topk = self.gkd_logits_topk if self._is_self_distillation: teacher_model = self.unwrapped_models[vp_stage or 0] @@ -249,27 +250,36 @@ def _compute_teacher_logits_local(self, encoded_batches: List[Dict], vp_stage: O 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 + 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 + return teacher_out + + 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. + """ + for encoded_batch in encoded_batches: + teacher_model_inputs = encoded_batch['teacher_model_inputs'] + encoded_batch['teacher_output'] = self._compute_teacher_output_local(teacher_model_inputs, vp_stage) def _generate_and_score_completions(self, inputs: List[Dict]) -> List[Dict]: """Unified rollout → teacher → encode pipeline (mirrors Megatron GRPO). @@ -405,8 +415,13 @@ def forward_step(self, data_iterator, model): 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 + teacher_output = data.pop('teacher_output', None) + teacher_model_inputs = data.pop('teacher_model_inputs', None) + if teacher_output is None: + if teacher_model_inputs is None: + raise RuntimeError('encoded batch is missing both teacher_output and teacher_model_inputs; ' + 'cannot compute GKD teacher logits') + teacher_output = self._compute_teacher_output_local(teacher_model_inputs, vp_stage) data = self._prepare_batch(data, vp_stage) if self.use_teacher_api: teacher_output = cp_slice_teacher_output(teacher_output, data.get('packed_seq_params'), diff --git a/tests/megatron/test_gkd_teacher_logits_defer.py b/tests/megatron/test_gkd_teacher_logits_defer.py new file mode 100644 index 0000000000..691cb590f2 --- /dev/null +++ b/tests/megatron/test_gkd_teacher_logits_defer.py @@ -0,0 +1,91 @@ +# Copyright (c) ModelScope Contributors. All rights reserved. +from types import SimpleNamespace +from unittest import mock + + +def _import_trainer(): + try: + from swift.megatron.trainers.gkd_trainer import MegatronGKDTrainer + except Exception as e: # noqa: megatron-core not installed in this env + print(f'SKIP gkd teacher defer tests: {e}') + return None + return MegatronGKDTrainer + + +def test_compute_teacher_logits_defers_full_vocab_fixed_teacher(): + """Full-vocabulary local teacher logits must not be materialized at batch prep.""" + trainer_cls = _import_trainer() + if trainer_cls is None: + return + + stub = SimpleNamespace( + use_teacher_api=False, + _is_self_distillation=False, + gkd_logits_topk=None, + ) + encoded_batches = [{ + 'teacher_model_inputs': { + 'input_ids': object(), + }, + }] + with mock.patch.object(trainer_cls, '_compute_teacher_logits_local') as local_mock: + trainer_cls._compute_teacher_logits(stub, encoded_batches) + local_mock.assert_not_called() + assert 'teacher_output' not in encoded_batches[0] + + +def test_compute_teacher_logits_eager_when_topk_configured(): + """Compressed top-k teacher logits remain eager at batch preparation.""" + trainer_cls = _import_trainer() + if trainer_cls is None: + return + + stub = SimpleNamespace( + use_teacher_api=False, + _is_self_distillation=False, + gkd_logits_topk=8, + ) + encoded_batches = [{'teacher_model_inputs': {}}] + with mock.patch.object(trainer_cls, '_compute_teacher_logits_local') as local_mock: + trainer_cls._compute_teacher_logits(stub, encoded_batches) + local_mock.assert_called_once_with(stub, encoded_batches, None) + + +def test_compute_teacher_output_local_used_by_forward_step_when_missing(): + """forward_step computes teacher logits just-in-time for full-vocabulary mode.""" + trainer_cls = _import_trainer() + if trainer_cls is None: + return + + try: + import torch + + from swift.rlhf_trainers.gkd_loss import TeacherOutput + except Exception as e: + print(f'SKIP forward_step defer test: {e}') + return + + teacher_out = TeacherOutput(full_logits=torch.zeros(1, 2, 4)) + stub = SimpleNamespace( + gkd_logits_topk=None, + _prepare_batch=lambda batch, vp_stage: batch, + _compute_teacher_output_local=mock.Mock(return_value=teacher_out), + loss_func=mock.Mock(), + ) + data_iterator = iter([{ + 'data_source': 'dataset', + 'input_ids': torch.zeros((1, 2), dtype=torch.long), + 'teacher_model_inputs': { + 'input_ids': torch.zeros((1, 2), dtype=torch.long), + }, + }]) + model = mock.Mock() + model.return_value = torch.zeros(1, 2, 4) + unwrapped = mock.Mock() + unwrapped.get_input_tensor.return_value = None + unwrapped.vp_stage = None + with mock.patch('swift.megatron.trainers.gkd_trainer.get_attr_wrapped_model', return_value=unwrapped): + trainer_cls.forward_step(stub, data_iterator, model) + + stub._compute_teacher_output_local.assert_called_once() + model.assert_called_once()