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
55 changes: 50 additions & 5 deletions swift/megatron/trainers/gkd_trainer.py
Original file line number Diff line number Diff line change
Expand Up @@ -15,7 +15,7 @@
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.rlhf_trainers.gkd_loss import DataSource, TeacherOutput, gkd_loss, gkd_monitoring_stats
from swift.template import Template
from swift.utils import get_logger, to_device
from ..utils import forward_step_helper, get_padding_to
Expand Down Expand Up @@ -259,11 +259,23 @@ def _compute_teacher_logits_local(self, encoded_batches: List[Dict], vp_stage: O
if teacher_logits is not None:
teacher_logits = teacher_logits.detach()

target_logprobs = None
if (teacher_logits is not None and teacher_labels is not None
and encoded_batch.get('data_source') == DataSource.STUDENT):
safe_labels = teacher_labels.masked_fill(teacher_labels == -100, 0).long()
teacher_logprobs = vocab_parallel_log_softmax(teacher_logits.float())
target_logprobs = tp_gather_topk(teacher_logprobs, safe_labels.unsqueeze(-1)).squeeze(-1)
target_logprobs = target_logprobs.masked_fill(teacher_labels == -100, float('nan'))

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)
teacher_out = TeacherOutput(
topk_logprobs=topk_logits,
topk_indices=topk_indices,
target_logprobs=target_logprobs,
)
else:
teacher_out = TeacherOutput(full_logits=teacher_logits)
teacher_out = TeacherOutput(full_logits=teacher_logits, target_logprobs=target_logprobs)

teacher_out.labels = teacher_labels
encoded_batch['teacher_output'] = teacher_out
Expand All @@ -286,6 +298,11 @@ def _generate_and_score_completions(self, inputs: List[Dict]) -> List[Dict]:
samples = self._gather_rollout_results(local_batch)
self._log_completions_from_samples(samples)

num_turns_mean = None
if data_source == DataSource.STUDENT and samples and all(s.rollout_infos and 'num_turns' in s.rollout_infos
for s in samples):
num_turns_mean = sum(float(s.rollout_infos['num_turns']) for s in samples) / len(samples)

# Teacher API: build requests from samples, fetch logprobs
local_parsed = None
if self.use_teacher_api:
Expand All @@ -299,7 +316,7 @@ def _generate_and_score_completions(self, inputs: List[Dict]) -> List[Dict]:
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),
handle, topk=self.gkd_logits_topk, teacher_client=client, include_sampled=True),
scatter_fn=self._scatter_teacher_parsed,
is_main_process=self.is_main_process,
tag_key=self.args.teacher_tag_key)
Expand All @@ -314,6 +331,8 @@ def _generate_and_score_completions(self, inputs: List[Dict]) -> List[Dict]:
sample_slice = samples[start_idx:end_idx]
encoded_batch = self._encode_samples(sample_slice)
encoded_batch['data_source'] = data_source
if num_turns_mean is not None:
encoded_batch['num_turns'] = num_turns_mean
if local_parsed is not None:
encoded_batch['_teacher_parsed'] = local_parsed[start_idx:end_idx]
all_encoded_batches.append(encoded_batch)
Expand Down Expand Up @@ -346,7 +365,8 @@ def loss_func(self,
*,
labels: torch.Tensor,
teacher_output: TeacherOutput,
data_source: DataSource = DataSource.DATASET):
data_source: DataSource = DataSource.DATASET,
num_turns: Optional[float] = None):
"""Compute GKD loss (JSD + optional SFT loss)."""
student_logits = output_tensor

Expand Down Expand Up @@ -388,6 +408,29 @@ def loss_func(self,
loss = loss + self.sft_alpha * sft_loss

metric = {'loss': loss.detach().clone()}
if num_turns is not None:
metric['num_turns'] = loss.new_tensor(num_turns)
if data_source == DataSource.STUDENT:
monitor = gkd_monitoring_stats(
student_logits,
teacher_output,
labels,
full_vocab_topk=self.gkd_logits_topk or 16,
student_topk_fn=vocab_parallel_topk,
teacher_topk_fn=vocab_parallel_topk,
gather_fn=tp_gather_topk,
target_logprob_fn=lambda logits, target_ids: tp_gather_topk(
vocab_parallel_log_softmax(logits.float()), target_ids.unsqueeze(-1)).squeeze(-1))
packed = torch.stack([
monitor['topk_overlap_sum'], monitor['topk_overlap_count'], monitor['teacher_student_gap_sum'],
monitor['teacher_student_gap_count']
])
if self.args.context_parallel_size > 1:
torch.distributed.all_reduce(
packed, op=torch.distributed.ReduceOp.SUM, group=mpu.get_context_parallel_group())
torch.distributed.all_reduce(packed, op=torch.distributed.ReduceOp.SUM, group=mpu.get_data_parallel_group())
metric['gkd/topk_overlap'] = packed[0] / packed[1].clamp(min=1)
metric['gkd/teacher_student_gap'] = packed[2] / packed[3].clamp(min=1)
if sft_loss is not None:
metric['jsd_loss'] = jsd_loss_val.detach().clone()
metric['sft_loss'] = sft_loss.detach().clone()
Expand All @@ -408,6 +451,7 @@ def forward_step(self, data_iterator, model):

data = next(data_iterator)
data_source = data.pop('data_source', DataSource.DATASET)
num_turns = data.pop('num_turns', None)
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)
Expand All @@ -424,4 +468,5 @@ def forward_step(self, data_iterator, model):
labels=labels,
teacher_output=teacher_output,
data_source=data_source,
num_turns=num_turns,
)
8 changes: 6 additions & 2 deletions swift/megatron/trainers/rollout_mixin.py
Original file line number Diff line number Diff line change
Expand Up @@ -277,7 +277,11 @@ def _gather_teacher_requests(self, requests: List[RolloutInferRequest]) -> Dict[
flat_global = [req for dp in dp_ranks_sorted for req in segments_by_dp[dp]]
return {'flat_global': flat_global, 'offset': offset, 'n_local': len(requests)}

def _infer_teacher_requests(self, handle: Dict[str, Any], topk: int, teacher_client: Optional[Any] = None):
def _infer_teacher_requests(self,
handle: Dict[str, Any],
topk: int,
teacher_client: Optional[Any] = None,
include_sampled: bool = False):
"""Phase 2 (main process only, no collective): run the teacher HTTP infer.

Safe to call concurrently across teachers (distinct clients, no collective inside).
Expand All @@ -287,7 +291,7 @@ def _infer_teacher_requests(self, handle: Dict[str, Any], topk: int, teacher_cli
client = teacher_client if teacher_client is not None else self.teacher_clients[0]
request_config = RequestConfig(prompt_logprobs=topk, max_tokens=1, temperature=0.0)
responses = client.infer(handle['flat_global'], request_config=request_config, use_tqdm=False)
return [parse_prompt_logprobs(r, topk=topk) for r in responses]
return [parse_prompt_logprobs(r, topk=topk, include_sampled=include_sampled) for r in responses]

def _scatter_teacher_parsed(self, handle: Dict[str, Any], parsed_global):
"""Phase 3 (all ranks, collective): broadcast the parsed result and slice this rank's part."""
Expand Down
35 changes: 27 additions & 8 deletions swift/ray/megatron/gkd_trainer.py
Original file line number Diff line number Diff line change
Expand Up @@ -132,6 +132,10 @@ def _train_loop(self, tg, train_iters, iteration):
chunk = source_items[step_idx * chunk_size:(step_idx + 1) * chunk_size]
if not chunk:
break
num_turns_mean = None
if data_source == DataSource.STUDENT and all(s.rollout_infos and 'num_turns' in s.rollout_infos
for s in chunk):
num_turns_mean = sum(float(s.rollout_infos['num_turns']) for s in chunk) / len(chunk)
samples = self._encode_rollout_batch(chunk)

use_colocated_teacher = self._teacher_use_disable_adapter or (self._teacher_model_dir
Expand All @@ -144,7 +148,8 @@ def _train_loop(self, tg, train_iters, iteration):

# Driver collates the student (and, for the colocated path, the teacher view)
# micro-batches; the worker only runs prepare_batch (PP/CP slice) + forward.
dispatch = self._collate_for_workers_gkd(tg, samples, data_source, with_teacher=use_colocated_teacher)
dispatch = self._collate_for_workers_gkd(
tg, samples, data_source, with_teacher=use_colocated_teacher, num_turns=num_turns_mean)
if use_colocated_teacher:
# Teacher forwards on the worker (CP slicing keeps each rank's shard
# aligned) and caches per-micro-batch; train_step attaches the cache.
Expand Down Expand Up @@ -249,7 +254,7 @@ def _encode_rollout_batch(self, samples: List[GKDSample]):
result.append(payload)
return result

def _collate_for_workers_gkd(self, tg, samples: List[dict], data_source, *, with_teacher: bool):
def _collate_for_workers_gkd(self, tg, samples: List[dict], data_source, *, with_teacher: bool, num_turns=None):
"""Driver-side GKD collate: ``List[payload-dict]`` -> ``{dp_rank: [model_inputs]}``.

Mirrors the non-Ray GKD ``_encode_samples`` (data_collator on the rank, teacher
Expand Down Expand Up @@ -281,6 +286,8 @@ def _collate_for_workers_gkd(self, tg, samples: List[dict], data_source, *, with
chunk = shard[i:i + mbs]
model_inputs = template.data_collator([s['encoded'] for s in chunk], padding_to=padding_to)
model_inputs['data_source'] = data_source
if num_turns is not None:
model_inputs['num_turns'] = num_turns
if with_teacher:
has_opsd = chunk[0].get('teacher_encoded') is not None
key = 'teacher_encoded' if has_opsd else 'encoded'
Expand Down Expand Up @@ -344,7 +351,7 @@ def _fetch_teacher_from_replicas(self, gkd_samples: List[GKDSample], samples):
responses.extend(p)

for sample, response, t_encoded in zip(samples, responses, teacher_encodeds):
parsed = parse_prompt_logprobs(response, topk=topk)
parsed = parse_prompt_logprobs(response, topk=topk, include_sampled=True)
encoded = t_encoded if t_encoded is not None else sample['encoded']
teacher_labels = t_encoded.get('labels') if t_encoded is not None else None
sample['teacher_output'] = self._build_per_sample_teacher_output(parsed, encoded, topk, teacher_labels)
Expand All @@ -367,15 +374,27 @@ def _build_per_sample_teacher_output(parsed, encoded, topk, labels=None):
parsed_len = len(lps)
topk_logprobs = torch.full((seq_len, topk), float('-inf'), dtype=torch.float32)
topk_indices = torch.zeros(seq_len, topk, dtype=torch.long)
target_logprobs = torch.full((seq_len, ), float('nan'), dtype=torch.float32)
length = min(parsed_len, seq_len)
if length > 0:
topk_logprobs[:length] = torch.tensor(lps[:length], dtype=torch.float32)
topk_indices[:length] = torch.tensor(ixs[:length], dtype=torch.long)

kwargs = dict(topk_logprobs=topk_logprobs.unsqueeze(0), topk_indices=topk_indices.unsqueeze(0))
topk_logprobs[:length] = torch.tensor([row[:topk] for row in lps[:length]], dtype=torch.float32)
topk_indices[:length] = torch.tensor([row[:topk] for row in ixs[:length]], dtype=torch.long)
flat_input_ids = input_ids if isinstance(input_ids, list) else input_ids.reshape(-1).tolist()
for pos in range(min(length, seq_len - 1)):
target_id = int(flat_input_ids[pos + 1])
for lp, token_id in zip(lps[pos], ixs[pos]):
if int(token_id) == target_id:
target_logprobs[pos] = float(lp)
break

kwargs = dict(
topk_logprobs=topk_logprobs.unsqueeze(0),
topk_indices=topk_indices.unsqueeze(0),
target_logprobs=target_logprobs.unsqueeze(0))
if labels is not None:
t_labels = labels
if not isinstance(t_labels, torch.Tensor):
t_labels = torch.tensor(t_labels, dtype=torch.long)
kwargs['labels'] = t_labels.unsqueeze(0) if t_labels.dim() == 1 else t_labels
t_labels = t_labels.unsqueeze(0) if t_labels.dim() == 1 else t_labels
kwargs['labels'] = torch.roll(t_labels, shifts=-1, dims=-1)
return TeacherOutput(**kwargs)
45 changes: 41 additions & 4 deletions swift/ray/megatron/loss/gkd.py
Original file line number Diff line number Diff line change
Expand Up @@ -11,7 +11,7 @@
from swift.megatron.trainers.utils import prepare_batch
from swift.megatron.trainers.vocab_parallel_utils import vocab_parallel_kl_div, vocab_parallel_log_softmax
from swift.megatron.utils import forward_step_helper
from swift.rlhf_trainers.gkd_loss import DataSource, TeacherOutput, gkd_loss
from swift.rlhf_trainers.gkd_loss import DataSource, TeacherOutput, gkd_loss, gkd_monitoring_stats
from swift.utils import get_current_device, to_device
from .base import Loss

Expand All @@ -29,6 +29,7 @@ def forward_step(self, data_iterator, model):
data = next(data_iterator)
teacher_output = data.pop('teacher_output', TeacherOutput())
data_source = data.pop('data_source', None)
num_turns = data.pop('num_turns', None)
data.pop('grpo_batch', None) # RL signals packed in GRPOBatch (not used by GKD loss)
data = prepare_batch(self.args, data)

Expand All @@ -43,6 +44,7 @@ def forward_step(self, data_iterator, model):
labels=labels,
teacher_output=teacher_output,
data_source=data_source,
num_turns=num_turns,
model=model,
)

Expand Down Expand Up @@ -73,19 +75,31 @@ def compute_teacher_logits(
outputs.append(TeacherOutput())
continue
teacher_logits = teacher_logits.detach()
target_logprobs = None
if labels is not None:
safe_labels = labels.masked_fill(labels == -100, 0).long()
teacher_logprobs = vocab_parallel_log_softmax(teacher_logits.float())
target_logprobs = tp_gather_topk(teacher_logprobs, safe_labels.unsqueeze(-1)).squeeze(-1)
target_logprobs = target_logprobs.masked_fill(labels == -100, float('nan'))
if gkd_logits_topk is not None:
topk_logits, topk_indices = vocab_parallel_topk(teacher_logits, k=gkd_logits_topk)
outputs.append(TeacherOutput(topk_logprobs=topk_logits, topk_indices=topk_indices, labels=labels))
outputs.append(
TeacherOutput(
topk_logprobs=topk_logits,
topk_indices=topk_indices,
target_logprobs=target_logprobs,
labels=labels))
else:
outputs.append(TeacherOutput(full_logits=teacher_logits, labels=labels))
outputs.append(
TeacherOutput(full_logits=teacher_logits, target_logprobs=target_logprobs, labels=labels))
del collated
return outputs

# ------------------------------------------------------------------
# Loss computation
# ------------------------------------------------------------------

def loss_func(self, output_tensor, *, labels, teacher_output, data_source=None, model=None):
def loss_func(self, output_tensor, *, labels, teacher_output, data_source=None, num_turns=None, model=None):
args = self.args
student_logits = output_tensor

Expand Down Expand Up @@ -128,6 +142,29 @@ def loss_func(self, output_tensor, *, labels, teacher_output, data_source=None,
loss = loss + self.sft_alpha * sft_loss

metric = {'loss': loss.detach().clone()}
if num_turns is not None:
metric['num_turns'] = loss.new_tensor(num_turns)
if data_source == DataSource.STUDENT:
monitor = gkd_monitoring_stats(
student_logits,
teacher_output,
labels,
full_vocab_topk=getattr(args, 'gkd_logits_topk', None) or 16,
student_topk_fn=vocab_parallel_topk,
teacher_topk_fn=vocab_parallel_topk,
gather_fn=tp_gather_topk,
target_logprob_fn=lambda logits, target_ids: tp_gather_topk(
vocab_parallel_log_softmax(logits.float()), target_ids.unsqueeze(-1)).squeeze(-1))
packed = torch.stack([
monitor['topk_overlap_sum'], monitor['topk_overlap_count'], monitor['teacher_student_gap_sum'],
monitor['teacher_student_gap_count']
])
if args.context_parallel_size > 1:
torch.distributed.all_reduce(
packed, op=torch.distributed.ReduceOp.SUM, group=mpu.get_context_parallel_group())
torch.distributed.all_reduce(packed, op=torch.distributed.ReduceOp.SUM, group=mpu.get_data_parallel_group())
metric['gkd/topk_overlap'] = packed[0] / packed[1].clamp(min=1)
metric['gkd/teacher_student_gap'] = packed[2] / packed[3].clamp(min=1)
if sft_loss is not None:
metric['jsd_loss'] = jsd_loss_val.detach().clone()
metric['sft_loss'] = sft_loss.detach().clone()
Expand Down
4 changes: 2 additions & 2 deletions swift/ray/megatron/megatron_worker.py
Original file line number Diff line number Diff line change
Expand Up @@ -596,8 +596,8 @@ def _collate_teacher_outputs(
"""
from swift.rlhf_trainers.gkd_loss import TeacherOutput
effective_target = None if is_opsd else target_seq_len
pad_vals = {'topk_logprobs': float('-inf'), 'labels': -100}
fields = ('full_logits', 'topk_logprobs', 'topk_indices', 'labels')
pad_vals = {'topk_logprobs': float('-inf'), 'target_logprobs': float('nan'), 'labels': -100}
fields = ('full_logits', 'topk_logprobs', 'topk_indices', 'target_logprobs', 'labels')
kwargs = {}
for field in fields:
tensors = [getattr(t, field) for t in teacher_outputs]
Expand Down
Loading
Loading