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
1 change: 1 addition & 0 deletions docs/source/Instruction/Command-line-parameters.md
Original file line number Diff line number Diff line change
Expand Up @@ -516,6 +516,7 @@ Vera使用`target_modules`、`target_regex`、`modules_to_save`三个参数,
- lazy_tokenize: 是否使用lazy_tokenize。若该参数设置为False,则在训练之前对所有的数据集样本进行tokenize(多模态模型则包括从磁盘中读取图片)。该参数默认为None,在LLM训练中默认为False,而MLLM训练默认为True,节约内存。
- 注意:若你要进行图像的数据增强,你需要将lazy_tokenize(或streaming)设置为True,并修改Template类中的encode方法。
- use_logits_to_keep: 通过在`forward`中根据labels传入logits_to_keep,减少无效logits的计算与存储,从而减少显存占用并加快训练速度。默认为None,进行自动选择。
- 与序列并行一起使用时,需显式设置`--use_logits_to_keep true`。目前支持纯文本 causal LM,以及 Qwen3.5/3.6 MoE 的纯文本样本(不含图像、视频)的 Ulysses 训练,要求每卡 batch size 为 1,关闭 packing/padding-free,并使用默认 loss(不支持自定义 loss、label smoothing、Unsloth 或 Liger)。模型需原生支持张量形式的`logits_to_keep`。评估时保留完整 logits。
- acc_strategy: 训练和验证时计算acc的策略。可选为`seq`和`token`级别的acc,默认为`token`。
- max_new_tokens: 覆盖生成参数。predict_with_generate=True时的最大生成token数量,默认64。
- temperature: 覆盖生成参数。predict_with_generate=True时的temperature,默认0。
Expand Down
1 change: 1 addition & 0 deletions docs/source_en/Instruction/Command-line-parameters.md
Original file line number Diff line number Diff line change
Expand Up @@ -528,6 +528,7 @@ Training arguments include the [base arguments](#base-arguments), [Seq2SeqTraine
- lazy_tokenize: Whether to use lazy tokenization. If set to `False`, all dataset samples will be tokenized (and for multimodal models, images will be loaded from disk) before training begins. Default is `None`: in LLM training, it defaults to `False`; in MLLM training, it defaults to `True` to save memory.
- Note: If you want to perform image data augmentation, you need to set `lazy_tokenize` (or `streaming`) to True and modify the `encode` method in the Template class.
- use_logits_to_keep: Pass `logits_to_keep` in the `forward` method based on labels to reduce the computation and storage of unnecessary logits, thereby reducing memory usage and accelerating training. The default is `None`, which enables automatic selection.
- With sequence parallelism, explicitly set `--use_logits_to_keep true`. Support is limited to text causal LMs and text-only Qwen3.5/3.6 MoE inputs (no images or videos), using Ulysses, per-device batch size 1, no packing/padding-free, and the default loss (no custom loss, label smoothing, Unsloth, or Liger). The model must natively support tensor-valued `logits_to_keep`. Evaluation retains full logits.
- acc_strategy: Strategy for calculating accuracy during training and validation. Options are `seq`-level and `token`-level accuracy, with `token` as the default.
- max_new_tokens: Generation parameter override. The maximum number of tokens to generate when `predict_with_generate=True`, defaulting to 64.
- temperature: Generation parameter override. The temperature setting when `predict_with_generate=True`, defaulting to 0.
Expand Down
4 changes: 3 additions & 1 deletion swift/trainers/mixin.py
Original file line number Diff line number Diff line change
Expand Up @@ -1161,7 +1161,7 @@ def _get_listwise_reranker_preds(logits, labels):
labels = torch.tensor([0] * (len(positive_indices) - 1))
return preds, labels

def _compute_acc(self, outputs, labels, cu_seqlens=None) -> None:
def _compute_acc(self, outputs, labels, cu_seqlens=None, logits_to_keep=None) -> None:
args = self.args
logits = outputs.logits
metrics = None
Expand All @@ -1186,6 +1186,8 @@ def _compute_acc(self, outputs, labels, cu_seqlens=None) -> None:
preds = torch.from_numpy(preds).to(get_current_device())
if isinstance(labels, np.ndarray):
labels = torch.from_numpy(labels).to(get_current_device())
if logits_to_keep is not None:
preds = preds.new_zeros(labels.shape).masked_scatter(logits_to_keep[None], preds)
assert labels.shape[1] == preds.shape[1]

if sequence_parallel.rp_world_size > 1:
Expand Down
37 changes: 34 additions & 3 deletions swift/trainers/seq2seq_trainer.py
Original file line number Diff line number Diff line change
Expand Up @@ -14,6 +14,7 @@
from typing import Any, Callable, Dict, List, Optional, Tuple, Union

from swift.infer_engine import InferRequest, RequestConfig, TransformersEngine
from swift.model import MLLMModelType
from swift.sequence_parallel import sequence_parallel
from swift.utils import HfConfigFactory, JsonlWriter, Serializer, gc_collect, get_logger, unwrap_model_for_generation
from .arguments import Seq2SeqTrainingArguments
Expand Down Expand Up @@ -101,6 +102,32 @@ def prediction_step(
labels_list = pad_for_ddp_gather(labels_list, padding_value=0)
return None, response_list, labels_list

def prepare_logits_to_keep(self, inputs):
if self.template.sequence_parallel_size == 1:
return super().prepare_logits_to_keep(inputs)
labels = inputs['labels']
# Evaluation keeps full logits for prediction gathering and external metrics.
if not self.model.training:
return
if labels.shape[0] != 1 or self.template.padding_free or sequence_parallel.rp_world_size > 1:
raise NotImplementedError('SP logits_to_keep requires batch size 1, no packing/padding_free, and Ulysses.')
if self.compute_loss_func is not None or self.label_smoother is not None:
raise NotImplementedError('SP logits_to_keep requires the default causal language modeling loss.')
if self.template.is_encoder_decoder or self.args.tuner_backend == 'unsloth' or self.args.use_liger_kernel:
raise NotImplementedError('SP logits_to_keep requires a text causal LM without Unsloth or Liger.')
if self.model.model_meta.is_multimodal:
media_keys = ('pixel_values', 'pixel_values_videos', 'image_grid_thw', 'video_grid_thw', 'inputs_embeds')
mm_token_type_ids = inputs.get('mm_token_type_ids')
if (self.model.model_meta.model_type != MLLMModelType.qwen3_5_moe or inputs.get('input_ids') is None
or any(inputs.get(key) is not None for key in media_keys)
or (mm_token_type_ids is not None and mm_token_type_ids.any())):
raise NotImplementedError('SP logits_to_keep only supports text-only inputs for Qwen3.5/3.6 MoE.')
# SP has already shifted and sharded labels. Keep them at their full local length.
logits_to_keep = labels[0] != -100
# Keep one position even on prompt-only ranks so lm_head participates in backward.
logits_to_keep[-1] = True
inputs['logits_to_keep'] = logits_to_keep

def _prepare_inputs(self, inputs):
args = self.args
inputs = super()._prepare_inputs(inputs)
Expand All @@ -110,7 +137,7 @@ def _prepare_inputs(self, inputs):
use_logits_to_keep = self.get_use_logits_to_keep(self.template.sequence_parallel_size == 1)
if use_logits_to_keep:
self.prepare_logits_to_keep(inputs)
if args.tuner_backend == 'unsloth' and isinstance(inputs['logits_to_keep'], torch.Tensor):
if args.tuner_backend == 'unsloth' and isinstance(inputs.get('logits_to_keep'), torch.Tensor):
inputs['logits_to_keep'] = int(inputs['logits_to_keep'].sum())

base_model = self.template.get_base_model(self.model)
Expand Down Expand Up @@ -167,7 +194,8 @@ def compute_loss(self, model, inputs, return_outputs=False, num_items_in_batch=N
outputs,
labels,
enable_dft_loss=self.args.enable_dft_loss,
return_labels=self.args.enable_channel_loss)
return_labels=self.args.enable_channel_loss,
logits_to_keep=inputs.get('logits_to_keep'))
if self.args.enable_channel_loss:
outputs.loss, channel_labels = sp_loss
else:
Expand Down Expand Up @@ -250,7 +278,10 @@ def compute_loss(self, model, inputs, return_outputs=False, num_items_in_batch=N
cu_seqlens = self.get_cu_seqlens(text_position_ids, inputs.get('logits_to_keep'))
# Liger does not have logits
# Unsloth has a bug with output logits
self._compute_acc(outputs, labels, cu_seqlens=cu_seqlens)
acc_kwargs = {}
if self.template.sequence_parallel_size > 1 and 'logits_to_keep' in inputs:
acc_kwargs['logits_to_keep'] = inputs['logits_to_keep']
self._compute_acc(outputs, labels, cu_seqlens=cu_seqlens, **acc_kwargs)
return (loss, outputs) if return_outputs else loss

def training_step(self, model, inputs, *args, **kwargs):
Expand Down
12 changes: 10 additions & 2 deletions swift/trainers/utils.py
Original file line number Diff line number Diff line change
Expand Up @@ -164,7 +164,7 @@ def is_instance_of_ms_model(model: Module) -> bool:
return False


def per_token_loss_func_sp(outputs, labels, enable_dft_loss=False, return_labels=False, **kwargs):
def per_token_loss_func_sp(outputs, labels, enable_dft_loss=False, return_labels=False, logits_to_keep=None, **kwargs):
"""Common loss function for sequence parallel training"""
if hasattr(outputs, 'logits'):
logits = outputs.logits
Expand All @@ -174,7 +174,11 @@ def per_token_loss_func_sp(outputs, labels, enable_dft_loss=False, return_labels

batch_size = logits.shape[0]
logits = logits.view(-1, logits.shape[-1])
labels = labels.flatten().to(device)
labels = labels.to(device)
full_labels = labels
if logits_to_keep is not None:
labels = labels[:, logits_to_keep]
labels = labels.flatten()
sploss_parallel_size = int(os.environ.get('CELOSS_PARALLEL_SIZE', '0'))
if sploss_parallel_size > 0:
loss = ChunkedCrossEntropyLoss.apply(logits, labels, sploss_parallel_size)
Expand All @@ -185,6 +189,10 @@ def per_token_loss_func_sp(outputs, labels, enable_dft_loss=False, return_labels
with torch.no_grad():
target_probs = torch.exp(-loss)
loss *= target_probs
if logits_to_keep is not None:
# Gather full-length scalar losses, not variable-length vocabulary logits.
loss = loss.new_zeros(full_labels.shape).masked_scatter(logits_to_keep[None], loss)
labels = full_labels
position_ids = sequence_parallel.real_position_ids
if position_ids is not None:
position_ids = sequence_parallel.pad(position_ids, padding_value=-1, position_ids=position_ids)
Expand Down
Loading
Loading