6767
6868class 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 ()]
0 commit comments