Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
75 changes: 45 additions & 30 deletions swift/megatron/trainers/gkd_trainer.py
Original file line number Diff line number Diff line change
Expand Up @@ -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]
Expand All @@ -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).
Expand Down Expand Up @@ -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'),
Expand Down
91 changes: 91 additions & 0 deletions tests/megatron/test_gkd_teacher_logits_defer.py
Original file line number Diff line number Diff line change
@@ -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()