From 337b05efe38f81aa236ba8d47d3dbcfcb0b31b73 Mon Sep 17 00:00:00 2001 From: minelhi <3417378192@qq.com> Date: Thu, 1 Oct 2026 12:02:55 +0800 Subject: [PATCH 1/2] fix(megatron): defer full-vocabulary GKD teacher logits to forward_step When gkd_logits_topk is unset, materialize local teacher [S, V] logits just-in-time per micro-batch instead of retaining one tensor per rollout batch during preparation. Top-k teacher outputs stay eager. Fixes #10097. Co-authored-by: Cursor --- swift/megatron/trainers/gkd_trainer.py | 75 +++++++++------- .../megatron/test_gkd_teacher_logits_defer.py | 90 +++++++++++++++++++ 2 files changed, 135 insertions(+), 30 deletions(-) create mode 100644 tests/megatron/test_gkd_teacher_logits_defer.py 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..ae1407efbe --- /dev/null +++ b/tests/megatron/test_gkd_teacher_logits_defer.py @@ -0,0 +1,90 @@ +# 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() From dd247d482bccaa662b9a7dcf52163361e87433ec Mon Sep 17 00:00:00 2001 From: minelhi <3417378192@qq.com> Date: Mon, 5 Oct 2026 00:03:37 +0800 Subject: [PATCH 2/2] style(tests): fix isort grouping in GKD teacher defer tests Co-authored-by: Cursor --- tests/megatron/test_gkd_teacher_logits_defer.py | 1 + 1 file changed, 1 insertion(+) diff --git a/tests/megatron/test_gkd_teacher_logits_defer.py b/tests/megatron/test_gkd_teacher_logits_defer.py index ae1407efbe..691cb590f2 100644 --- a/tests/megatron/test_gkd_teacher_logits_defer.py +++ b/tests/megatron/test_gkd_teacher_logits_defer.py @@ -59,6 +59,7 @@ def test_compute_teacher_output_local_used_by_forward_step_when_missing(): try: import torch + from swift.rlhf_trainers.gkd_loss import TeacherOutput except Exception as e: print(f'SKIP forward_step defer test: {e}')