Skip to content

Commit 94c6759

Browse files
committed
fix(trainer): support selective logits with sequence parallel
1 parent 16e89c8 commit 94c6759

4 files changed

Lines changed: 287 additions & 14 deletions

File tree

swift/trainers/mixin.py

Lines changed: 72 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -67,6 +67,9 @@
6767

6868
class SwiftMixin:
6969
FLASH_CKPT_WAIT_TIMEOUT = 1800
70+
# Selective logits with SP requires a trainer-specific loss reducer. SFT
71+
# opts in below; RLHF trainers keep their established full-frame path.
72+
SUPPORTS_SP_LOGITS_TO_KEEP = False
7073

7174
def __init__(self,
7275
model: PreTrainedModel,
@@ -220,6 +223,10 @@ def get_use_logits_to_keep(self, default_value: bool = True):
220223
use_logits_to_keep = default_value
221224
self.args.use_logits_to_keep = use_logits_to_keep
222225
logger.info_once(f'use_logits_to_keep: {use_logits_to_keep}')
226+
if (use_logits_to_keep and self.template.sequence_parallel_size > 1 and not self.SUPPORTS_SP_LOGITS_TO_KEEP):
227+
logger.warning_once(
228+
'Disabling use_logits_to_keep for this sequence-parallel trainer; its loss reducer is not SP-aware.')
229+
use_logits_to_keep = False
223230
return use_logits_to_keep
224231

225232
def _save_initial_model(self, output_dir):
@@ -1161,7 +1168,7 @@ def _get_listwise_reranker_preds(logits, labels):
11611168
labels = torch.tensor([0] * (len(positive_indices) - 1))
11621169
return preds, labels
11631170

1164-
def _compute_acc(self, outputs, labels, cu_seqlens=None) -> None:
1171+
def _compute_acc(self, outputs, labels, cu_seqlens=None, logits_to_keep=None) -> None:
11651172
args = self.args
11661173
logits = outputs.logits
11671174
metrics = None
@@ -1186,9 +1193,28 @@ def _compute_acc(self, outputs, labels, cu_seqlens=None) -> None:
11861193
preds = torch.from_numpy(preds).to(get_current_device())
11871194
if isinstance(labels, np.ndarray):
11881195
labels = torch.from_numpy(labels).to(get_current_device())
1196+
# Selective lm-head projection returns only the positions
1197+
# addressed by ``logits_to_keep``. Reinsert those predictions
1198+
# into the full local shard before the existing gather path.
1199+
if isinstance(logits_to_keep, torch.Tensor) and logits_to_keep.dtype == torch.bool:
1200+
if logits_to_keep.ndim == 1 and logits_to_keep.numel() == labels.shape[1]:
1201+
selected_count = int(logits_to_keep.sum().item())
1202+
if preds.shape[1] == selected_count:
1203+
full_preds = torch.zeros((preds.shape[0], labels.shape[1]),
1204+
dtype=preds.dtype,
1205+
device=preds.device)
1206+
full_preds[:, logits_to_keep] = preds
1207+
preds = full_preds
1208+
elif isinstance(logits_to_keep, int) and 0 < logits_to_keep <= labels.shape[1]:
1209+
if preds.shape[1] == logits_to_keep:
1210+
full_preds = torch.zeros((preds.shape[0], labels.shape[1]),
1211+
dtype=preds.dtype,
1212+
device=preds.device)
1213+
full_preds[:, -logits_to_keep:] = preds
1214+
preds = full_preds
11891215
assert labels.shape[1] == preds.shape[1]
11901216

1191-
if sequence_parallel.rp_world_size > 1:
1217+
if (sequence_parallel.rp_world_size or 1) > 1:
11921218
position_ids = sequence_parallel.real_position_ids
11931219
position_ids = sequence_parallel.pad(position_ids, padding_value=-1, position_ids=position_ids)
11941220
else:
@@ -1257,10 +1283,38 @@ def _evalscope_eval(self):
12571283
return eval_dict
12581284

12591285
def prepare_logits_to_keep(self, inputs):
1286+
"""Prepare selective lm-head inputs for regular and SP SFT paths.
1287+
1288+
Sequence-parallel input preparation has already applied the causal
1289+
shift to labels. Keep that full local frame intact and let the SP
1290+
loss function scatter selected logits back into it.
1291+
"""
12601292
labels = inputs['labels']
12611293
loss_scale = inputs.get('loss_scale')
12621294
if self.template.sequence_parallel_size > 1:
1263-
raise NotImplementedError()
1295+
# Transformers causal-LM heads accept a one-dimensional boolean
1296+
# sequence index for every batch row. Keep arbitrary supervised
1297+
# positions for batch-size one; for a larger batch use one shared
1298+
# suffix that covers the earliest supervised target in any row.
1299+
if labels.shape[0] == 1 and not is_mp():
1300+
logits_to_keep = labels[0] != -100
1301+
# Keep one ignored position on an all-masked shard so model
1302+
# implementations that reject an empty lm_head input remain
1303+
# usable; it contributes zero to the loss.
1304+
if not logits_to_keep.any():
1305+
logits_to_keep = logits_to_keep.clone()
1306+
logits_to_keep[-1] = True
1307+
else:
1308+
supervised = labels != -100
1309+
first_supervised = supervised.int().argmax(dim=-1)
1310+
has_supervised = supervised.any(dim=-1)
1311+
first = first_supervised.masked_fill(~has_supervised, labels.shape[-1] - 1).min().item()
1312+
logits_to_keep = torch.zeros(labels.shape[-1], dtype=torch.bool, device=labels.device)
1313+
logits_to_keep[-max(labels.shape[-1] - first, 1):] = True
1314+
inputs['logits_to_keep'] = logits_to_keep
1315+
# Do not truncate labels/loss_scale: SP gathers a full local frame
1316+
# and therefore needs their original shard length.
1317+
return
12641318
if labels.shape[0] == 1 and not is_mp():
12651319
# device_map may encounter device mismatch issues.
12661320
loss_mask = (labels != -100)[0]
@@ -1282,6 +1336,21 @@ def prepare_logits_to_keep(self, inputs):
12821336
def get_cu_seqlens(self, position_ids, logits_to_keep) -> torch.Tensor:
12831337
cu_seqlens = get_packed_seq_params(position_ids)['cu_seq_lens_q']
12841338
if isinstance(logits_to_keep, torch.Tensor):
1339+
# SP keeps a local boolean mask while position_ids still contains
1340+
# the complete packed sequence. Gather the mask first so compact
1341+
# boundaries are computed in the global frame.
1342+
if (getattr(getattr(self, 'template', None), 'sequence_parallel_size', 1) > 1 and logits_to_keep.ndim == 1
1343+
and logits_to_keep.numel() != position_ids.shape[-1] and (sequence_parallel.world_size or 1) > 1):
1344+
local_mask = logits_to_keep.unsqueeze(0)
1345+
gather_position_ids = None
1346+
if (sequence_parallel.rp_world_size or 1) > 1:
1347+
gather_position_ids = sequence_parallel.real_position_ids
1348+
gather_position_ids = sequence_parallel.pad(
1349+
gather_position_ids, padding_value=-1, position_ids=gather_position_ids)
1350+
logits_to_keep = sequence_parallel.gather(local_mask, dim=1, position_ids=gather_position_ids)
1351+
if gather_position_ids is not None and gather_position_ids.min() == -1:
1352+
logits_to_keep = logits_to_keep[gather_position_ids >= 0]
1353+
logits_to_keep = logits_to_keep.reshape(-1)
12851354
kept_cumsum = logits_to_keep.to(cu_seqlens.dtype).cumsum(dim=0, dtype=cu_seqlens.dtype)
12861355
kept_cumsum = torch.cat((cu_seqlens.new_zeros(1), kept_cumsum))
12871356
res_cu_seqlens = kept_cumsum[cu_seqlens.long()]

swift/trainers/seq2seq_trainer.py

Lines changed: 25 additions & 4 deletions
Original file line numberDiff line numberDiff line change
@@ -25,6 +25,7 @@
2525

2626
class Seq2SeqTrainer(SwiftMixin, DataLoaderMixin, HfSeq2SeqTrainer):
2727
args: Seq2SeqTrainingArguments
28+
SUPPORTS_SP_LOGITS_TO_KEEP = True
2829

2930
def __init__(self, *args, **kwargs):
3031
super().__init__(*args, **kwargs)
@@ -105,9 +106,20 @@ def _prepare_inputs(self, inputs):
105106
sequence_parallel.prepare_inputs(inputs)
106107

107108
use_logits_to_keep = self.get_use_logits_to_keep(self.template.sequence_parallel_size == 1)
109+
sp_selective_unsupported = self.template.sequence_parallel_size > 1 and (self.compute_loss_func is not None
110+
or self.label_smoother is not None
111+
or self.args.enable_channel_loss)
112+
if use_logits_to_keep and sp_selective_unsupported:
113+
# Custom losses and label smoothing may consume the original
114+
# (unshifted) label frame. Keep their established semantics until
115+
# they provide an SP-aware selective-loss implementation.
116+
logger.warning_once(
117+
'Disabling use_logits_to_keep for sequence parallel custom loss/label smoothing/channel loss.')
118+
use_logits_to_keep = False
108119
if use_logits_to_keep:
109120
self.prepare_logits_to_keep(inputs)
110-
if args.tuner_backend == 'unsloth' and isinstance(inputs['logits_to_keep'], torch.Tensor):
121+
if (args.tuner_backend == 'unsloth' and self.template.sequence_parallel_size == 1
122+
and isinstance(inputs['logits_to_keep'], torch.Tensor)):
111123
inputs['logits_to_keep'] = int(inputs['logits_to_keep'].sum())
112124

113125
base_model = self.template.get_base_model(self.model)
@@ -160,7 +172,11 @@ def compute_loss(self, model, inputs, return_outputs=False, num_items_in_batch=N
160172
if (self.args.enable_dft_loss or loss_scale is not None or self.args.enable_channel_loss
161173
or self.template.sequence_parallel_size > 1):
162174
if self.template.sequence_parallel_size > 1:
163-
outputs.loss = per_token_loss_func_sp(outputs, labels, enable_dft_loss=self.args.enable_dft_loss)
175+
outputs.loss = per_token_loss_func_sp(
176+
outputs,
177+
labels,
178+
enable_dft_loss=self.args.enable_dft_loss,
179+
logits_to_keep=inputs.get('logits_to_keep'))
164180
if loss_scale is not None:
165181
position_ids = sequence_parallel.real_position_ids
166182
if position_ids is not None:
@@ -232,10 +248,15 @@ def compute_loss(self, model, inputs, return_outputs=False, num_items_in_batch=N
232248
if (outputs.logits is not None and labels is not None and self.args.tuner_backend != 'unsloth'):
233249
cu_seqlens = None
234250
if self.template.padding_free and self.args.acc_strategy == 'seq':
235-
cu_seqlens = self.get_cu_seqlens(text_position_ids, inputs.get('logits_to_keep'))
251+
logits_to_keep = inputs.get('logits_to_keep')
252+
# Outputs are full-frame after SP selective-loss scattering;
253+
# retain full packed boundaries for sequence accuracy.
254+
cu_seqlens = self.get_cu_seqlens(
255+
text_position_ids,
256+
None if self.template.sequence_parallel_size > 1 and logits_to_keep is not None else logits_to_keep)
236257
# Liger does not have logits
237258
# Unsloth has a bug with output logits
238-
self._compute_acc(outputs, labels, cu_seqlens=cu_seqlens)
259+
self._compute_acc(outputs, labels, cu_seqlens=cu_seqlens, logits_to_keep=inputs.get('logits_to_keep'))
239260
return (loss, outputs) if return_outputs else loss
240261

241262
def training_step(self, model, inputs, *args, **kwargs):

swift/trainers/utils.py

Lines changed: 54 additions & 7 deletions
Original file line numberDiff line numberDiff line change
@@ -164,27 +164,74 @@ def is_instance_of_ms_model(model: Module) -> bool:
164164
return False
165165

166166

167-
def per_token_loss_func_sp(outputs, labels, enable_dft_loss=False, **kwargs) -> torch.Tensor:
168-
"""Common loss function for sequence parallel training"""
167+
def per_token_loss_func_sp(outputs, labels, enable_dft_loss=False, logits_to_keep=None, **kwargs) -> torch.Tensor:
168+
"""Compute per-token SP loss while supporting selective lm-head logits.
169+
170+
Sequence-parallel labels are already causally shifted and sharded. When
171+
``logits_to_keep`` selects a compact subset of the local hidden states,
172+
CE is evaluated on that subset and scattered back into the full local
173+
sequence frame before the existing gather. This keeps packed/ring
174+
ordering and collective tensor shapes unchanged while avoiding the
175+
vocabulary projection for ignored tokens.
176+
"""
169177
if hasattr(outputs, 'logits'):
170178
logits = outputs.logits
171179
else:
172180
logits = outputs
173181
device = logits.device
182+
labels = labels.to(device)
183+
batch_size = labels.shape[0]
184+
compact_selection = logits_to_keep is not None
185+
186+
if compact_selection:
187+
local_seq_len = labels.shape[-1]
188+
if isinstance(logits_to_keep, torch.Tensor):
189+
if logits_to_keep.dtype != torch.bool or logits_to_keep.ndim != 1:
190+
compact_selection = False
191+
elif logits_to_keep.numel() != local_seq_len:
192+
raise ValueError(
193+
f'logits_to_keep has length {logits_to_keep.numel()}, expected {local_seq_len} for SP labels')
194+
else:
195+
selected_labels = labels[:, logits_to_keep]
196+
elif isinstance(logits_to_keep, int):
197+
if logits_to_keep <= 0 or logits_to_keep > local_seq_len:
198+
raise ValueError(f'logits_to_keep={logits_to_keep} must be in [1, {local_seq_len}] for SP labels')
199+
selected_labels = labels[:, -logits_to_keep:]
200+
else:
201+
compact_selection = False
202+
203+
if compact_selection:
204+
if logits.shape[1] != selected_labels.shape[1]:
205+
raise ValueError(f'logits sequence length ({logits.shape[1]}) does not match selected labels '
206+
f'({selected_labels.shape[1]})')
207+
logits = logits.reshape(-1, logits.shape[-1])
208+
selected_labels = selected_labels.reshape(-1)
209+
else:
210+
logits = logits.reshape(-1, logits.shape[-1])
211+
selected_labels = labels.reshape(-1)
174212

175-
batch_size = logits.shape[0]
176-
logits = logits.view(-1, logits.shape[-1])
177-
labels = labels.flatten().to(device)
178213
sploss_parallel_size = int(os.environ.get('CELOSS_PARALLEL_SIZE', '0'))
179214
if sploss_parallel_size > 0:
180-
loss = ChunkedCrossEntropyLoss.apply(logits, labels, sploss_parallel_size)
215+
loss = ChunkedCrossEntropyLoss.apply(logits, selected_labels, sploss_parallel_size)
181216
else:
182217
loss_fct = CrossEntropyLoss(reduction='none')
183-
loss = loss_fct(logits, labels)
218+
loss = loss_fct(logits, selected_labels)
184219
if enable_dft_loss:
185220
with torch.no_grad():
186221
target_probs = torch.exp(-loss)
187222
loss *= target_probs
223+
224+
if compact_selection:
225+
# Reconstruct a full local frame for GatherLoss. Unselected entries
226+
# remain zero and correspond to -100 labels in the original frame.
227+
selected_loss = loss.reshape(batch_size, -1)
228+
full_loss = torch.zeros((batch_size, labels.shape[-1]), dtype=selected_loss.dtype, device=selected_loss.device)
229+
if isinstance(logits_to_keep, torch.Tensor):
230+
full_loss[:, logits_to_keep] = selected_loss
231+
else:
232+
full_loss[:, -logits_to_keep:] = selected_loss
233+
loss = full_loss.reshape(-1)
234+
188235
position_ids = sequence_parallel.real_position_ids
189236
if position_ids is not None:
190237
position_ids = sequence_parallel.pad(position_ids, padding_value=-1, position_ids=position_ids)

0 commit comments

Comments
 (0)