diff --git a/docs/source/Instruction/Command-line-parameters.md b/docs/source/Instruction/Command-line-parameters.md index ae5bbb66b2..1a93c65b7d 100644 --- a/docs/source/Instruction/Command-line-parameters.md +++ b/docs/source/Instruction/Command-line-parameters.md @@ -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。 diff --git a/docs/source_en/Instruction/Command-line-parameters.md b/docs/source_en/Instruction/Command-line-parameters.md index 421fdf73cf..4415f25f09 100644 --- a/docs/source_en/Instruction/Command-line-parameters.md +++ b/docs/source_en/Instruction/Command-line-parameters.md @@ -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. diff --git a/swift/trainers/mixin.py b/swift/trainers/mixin.py index ed0898c3b3..c00edb4e96 100644 --- a/swift/trainers/mixin.py +++ b/swift/trainers/mixin.py @@ -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 @@ -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: diff --git a/swift/trainers/seq2seq_trainer.py b/swift/trainers/seq2seq_trainer.py index 78f51559b6..f25f0dafc2 100644 --- a/swift/trainers/seq2seq_trainer.py +++ b/swift/trainers/seq2seq_trainer.py @@ -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 @@ -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) @@ -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) @@ -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: @@ -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): diff --git a/swift/trainers/utils.py b/swift/trainers/utils.py index 8368979250..ab9e6380a5 100644 --- a/swift/trainers/utils.py +++ b/swift/trainers/utils.py @@ -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 @@ -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) @@ -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) diff --git a/tests/sequence_parallel/test_logits_to_keep.py b/tests/sequence_parallel/test_logits_to_keep.py new file mode 100644 index 0000000000..27aa06c108 --- /dev/null +++ b/tests/sequence_parallel/test_logits_to_keep.py @@ -0,0 +1,329 @@ +# Copyright (c) ModelScope Contributors. All rights reserved. +import copy +import os +import pytest +import torch +import torch.distributed as dist +import torch.multiprocessing as mp +from collections import defaultdict +from datetime import timedelta +from itertools import product +from torch.distributed import init_device_mesh +from torch.nn.parallel import DistributedDataParallel +from types import MethodType, SimpleNamespace + +from swift.model import get_matched_model_meta +from swift.sequence_parallel import sequence_parallel +from swift.trainers import Seq2SeqTrainer +from swift.trainers.mixin import SwiftMixin + + +class Metric: + + def __init__(self): + self.values = [] + + def update(self, values): + self.values.extend(values if isinstance(values, list) else [values.detach().clone()]) + + +def make_trainer(model): + trainer = SimpleNamespace( + model=model, + task_type='causal_lm', + problem_type=None, + compute_loss_func=None, + label_smoother=None, + model_accepts_loss_kwargs=True, + custom_metrics={ + 'train': defaultdict(Metric), + 'eval': defaultdict(Metric) + }, + template=SimpleNamespace( + sequence_parallel_size=2, + padding_free=False, + is_encoder_decoder=False, + compute_sft_loss=lambda model, inputs, **kwargs: model(**inputs)), + accelerator=SimpleNamespace(unwrap_model=lambda model: model, num_processes=2), + args=SimpleNamespace( + router_aux_loss_coef=None, + use_liger_kernel=False, + past_index=-1, + enable_dft_loss=False, + enable_channel_loss=True, + average_tokens_across_devices=False, + tuner_backend='peft', + acc_strategy='token')) + trainer._compute_acc = MethodType(Seq2SeqTrainer._compute_acc, trainer) + return trainer + + +def _check_model(rank, rendezvous, model_kind): + from transformers import Qwen3Config, Qwen3ForCausalLM + + torch.set_num_threads(1) + dist.init_process_group('gloo', init_method=rendezvous, rank=rank, world_size=2, timeout=timedelta(seconds=90)) + try: + sp = sequence_parallel + sp.world_size = sp.sp_world_size = 2 + sp.rp_world_size = 1 + sp.device_mesh = init_device_mesh('cpu', (1, 2), mesh_dim_names=('data', 'sequence')) + sp.tokenizer = SimpleNamespace(pad_token_id=0) + sp.model_dtype = torch.float32 + sp.padding_free = False + torch.manual_seed(42) + if model_kind == 'qwen3': + config = Qwen3Config( + vocab_size=32, + hidden_size=16, + intermediate_size=32, + num_hidden_layers=1, + num_attention_heads=2, + num_key_value_heads=2, + head_dim=8, + attention_dropout=0., + use_cache=False) + config._attn_implementation = 'sdpa' + model = Qwen3ForCausalLM(config) + model.model_info = SimpleNamespace(is_moe_model=False) + model.model_meta = SimpleNamespace(is_multimodal=False) + else: + from transformers import Qwen3_5MoeConfig, Qwen3_5MoeForConditionalGeneration + + from swift.model.models.qwen import _patch_qwen3_5_linear_attention_sequence_parallel + + config = Qwen3_5MoeConfig( + text_config=dict( + vocab_size=32, + hidden_size=32, + num_hidden_layers=2, + num_attention_heads=2, + num_key_value_heads=2, + head_dim=16, + linear_key_head_dim=8, + linear_value_head_dim=8, + linear_num_key_heads=2, + linear_num_value_heads=2, + moe_intermediate_size=16, + shared_expert_intermediate_size=16, + num_experts=2, + num_experts_per_tok=1, + layer_types=['linear_attention', 'full_attention'], + use_cache=False), + vision_config=dict( + depth=1, + hidden_size=32, + intermediate_size=32, + num_heads=2, + out_hidden_size=32, + num_position_embeddings=16), + image_token_id=28, + video_token_id=29, + vision_start_token_id=30, + vision_end_token_id=31) + config._attn_implementation = 'sdpa' + model = Qwen3_5MoeForConditionalGeneration(config) + model.model.visual.requires_grad_(False) + model.model_info = SimpleNamespace(is_moe_model=True) + model.model_meta = get_matched_model_meta('Qwen/Qwen3.6-35B-A3B') + _patch_qwen3_5_linear_attention_sequence_parallel() + base_model = model + if model_kind == 'qwen3_5_moe_lora': + from peft import LoraConfig, get_peft_model + + model = get_peft_model( + model, + LoraConfig( + r=2, lora_alpha=4, target_modules=['q_proj', 'v_proj', 'in_proj_qkv'], task_type='CAUSAL_LM')) + model.gradient_checkpointing_enable(gradient_checkpointing_kwargs={'use_reentrant': False}) + reference = copy.deepcopy(model) + sp._prepare_flash_attn(base_model.model) + sp._prepare_forward_hook(base_model.model) + if model_kind != 'qwen3': + sp._prepare_moe_aux_loss(base_model.model.language_model) + ddp = DistributedDataParallel(model, find_unused_parameters=model_kind != 'qwen3') + trainer = make_trainer(model) + head_lengths = [] + base_model.lm_head.register_forward_pre_hook(lambda module, args: head_lengths.append(args[0].shape[1])) + for length, target_mode, scale_mode, denominator in product((8, 7), ('tail', 'boundary'), + ('none', 'weighted', 'zero'), (None, 11)): + positions = torch.arange(length)[None] + ids = (positions + 1) % 32 + labels = ids.clone() + labels[:, :length - 2] = -100 # Rank zero has no supervised positions. + if target_mode == 'boundary': + labels[:, 1] = ids[:, 1] + labels[:, length // 2] = ids[:, length // 2] + trainer.args.acc_strategy = 'seq' if target_mode == 'boundary' else 'token' + trainer.args.enable_dft_loss = target_mode == 'boundary' and denominator is not None + os.environ['CELOSS_PARALLEL_SIZE'] = '2' if scale_mode == 'weighted' else '0' + weights = torch.linspace(0.25, 2., length)[None] + if scale_mode == 'zero': + weights.zero_() + count = (labels != -100).sum() if denominator is None else denominator + # Independent full-sequence model and causal CE reference. + sp.world_size = 1 + reference.zero_grad(set_to_none=True) + logits = reference(input_ids=ids, position_ids=positions).logits + losses = torch.nn.functional.cross_entropy( + logits[:, :-1].reshape(-1, 32), labels[:, 1:].reshape(-1), reduction='none') + if trainer.args.enable_dft_loss: + losses = losses * torch.exp(-losses.detach()) + if scale_mode != 'none': + losses = losses * weights[:, 1:].flatten() + reference_loss = losses.sum() / count + reference_loss.backward() + sp.world_size = 2 + results = [] + for selected in (False, True): + model.train() + ddp.zero_grad(set_to_none=True) + trainer.custom_metrics = {'train': defaultdict(Metric), 'eval': defaultdict(Metric)} + inputs = { + 'input_ids': ids, + 'position_ids': positions, + 'attention_mask': torch.ones_like(ids), + 'labels': labels.clone() + } + if model_kind != 'qwen3': + inputs['mm_token_type_ids'] = torch.zeros_like(ids) + if scale_mode != 'none': + inputs['loss_scale'] = weights.clone() + sp.prepare_inputs(inputs) + original_labels = inputs['labels'].clone() + if selected: + Seq2SeqTrainer.prepare_logits_to_keep(trainer, inputs) + torch.testing.assert_close(inputs['labels'], original_labels) + loss = Seq2SeqTrainer.compute_loss(trainer, ddp, inputs, num_items_in_batch=denominator) + loss.backward() + torch.testing.assert_close(loss, reference_loss, rtol=2e-5, atol=2e-6) + for (name, parameter), (ref_name, ref_parameter) in zip(model.named_parameters(), + reference.named_parameters()): + assert name == ref_name + if not parameter.requires_grad: + continue + assert parameter.grad is not None + torch.testing.assert_close(parameter.grad, ref_parameter.grad, rtol=2e-4, atol=2e-6) + results.append(trainer.custom_metrics['train']) + if selected and rank == 0 and target_mode == 'tail': + assert head_lengths[-1] == 1 # The ignored sentinel keeps all DDP parameters connected. + metric = f'{trainer.args.acc_strategy}_acc' + assert results[0][metric].values and results[0][metric].values == results[1][metric].values + assert results[0]['loss_None'].values and results[1]['loss_None'].values + for a, b in zip(results[0]['loss_None'].values, results[1]['loss_None'].values): + torch.testing.assert_close(a, b) + if model_kind != 'qwen3': + # Auxiliary router loss must be unchanged by selecting vocabulary logits. + trainer.args.router_aux_loss_coef = 0.001 + trainer.args.enable_dft_loss = False + comparisons = [] + for selected in (False, True): + ddp.zero_grad(set_to_none=True) + inputs = { + 'input_ids': ids, + 'position_ids': positions, + 'attention_mask': torch.ones_like(ids), + 'labels': labels.clone(), + 'output_router_logits': True + } + sp.prepare_inputs(inputs) + if selected: + Seq2SeqTrainer.prepare_logits_to_keep(trainer, inputs) + loss, outputs = Seq2SeqTrainer.compute_loss( + trainer, ddp, inputs, return_outputs=True, num_items_in_batch=11) + assert outputs.aux_loss is not None and torch.isfinite(outputs.aux_loss) + loss.backward() + comparisons.append((loss.detach(), { + n: p.grad.clone() + for n, p in model.named_parameters() if p.requires_grad and p.grad is not None + })) + torch.testing.assert_close(comparisons[0][0], comparisons[1][0]) + assert comparisons[0][1].keys() == comparisons[1][1].keys() + for name in comparisons[0][1]: + torch.testing.assert_close(comparisons[0][1][name], comparisons[1][1][name], rtol=2e-4, atol=2e-6) + model.eval() + inputs = {'labels': torch.tensor([[-100, 3, 4, -100]])} + Seq2SeqTrainer.prepare_logits_to_keep(trainer, inputs) + assert 'logits_to_keep' not in inputs + except Exception: + import traceback + traceback.print_exc() + raise + finally: + dist.destroy_process_group() + + +@pytest.mark.skipif(not dist.is_available() or not dist.is_gloo_available(), reason='Gloo is not available') +@pytest.mark.parametrize('model_kind', ['qwen3', 'qwen3_5_moe', 'qwen3_5_moe_lora']) +def test_sp_logits_selection_matches_full_loss_and_gradients(tmp_path, model_kind): + module_name = 'qwen3' if model_kind == 'qwen3' else 'qwen3_5_moe' + pytest.importorskip(f'transformers.models.{module_name}') + mp.spawn(_check_model, args=((tmp_path / 'rendezvous').as_uri(), model_kind), nprocs=2) + + +@pytest.mark.parametrize( + 'unsupported', + ['batch', 'padding_free', 'ring', 'custom_loss', 'smoothing', 'multimodal', 'encoder_decoder', 'unsloth', 'liger']) +def test_sp_logits_selection_rejects_unsupported_configs(monkeypatch, unsupported): + model = SimpleNamespace(training=True, model_meta=SimpleNamespace(is_multimodal=False, model_type='other')) + trainer = make_trainer(model) + monkeypatch.setattr(sequence_parallel, 'rp_world_size', 1) + inputs = {'labels': torch.tensor([[-100, 1, 2, -100]])} + if unsupported == 'batch': + inputs['labels'] = inputs['labels'].repeat(2, 1) + elif unsupported == 'padding_free': + trainer.template.padding_free = True + elif unsupported == 'ring': + monkeypatch.setattr(sequence_parallel, 'rp_world_size', 2) + elif unsupported == 'custom_loss': + trainer.compute_loss_func = object() + elif unsupported == 'smoothing': + trainer.label_smoother = object() + elif unsupported == 'multimodal': + model.model_meta.is_multimodal = True + elif unsupported == 'encoder_decoder': + trainer.template.is_encoder_decoder = True + elif unsupported == 'unsloth': + trainer.args.tuner_backend = 'unsloth' + else: + trainer.args.use_liger_kernel = True + with pytest.raises(NotImplementedError, match='SP logits_to_keep'): + Seq2SeqTrainer.prepare_logits_to_keep(trainer, inputs) + assert 'logits_to_keep' not in inputs + + +def test_non_sp_logits_selection_unchanged(): + trainer = Seq2SeqTrainer.__new__(Seq2SeqTrainer) + trainer.template = SimpleNamespace(sequence_parallel_size=1) + inputs = {'labels': torch.tensor([[-100, -100, 2, 3]]), 'loss_scale': torch.tensor([[0., 0., 0.5, 2.]])} + Seq2SeqTrainer.prepare_logits_to_keep(trainer, inputs) + torch.testing.assert_close(inputs['labels'], torch.tensor([[-100, 2, 3]])) + torch.testing.assert_close(inputs['loss_scale'], torch.tensor([[0., 0.5, 2.]])) + torch.testing.assert_close(inputs['logits_to_keep'], torch.tensor([False, True, True, True])) + + +@pytest.mark.parametrize('training', [True, False]) +def test_shared_mixin_still_rejects_sp_logits_selection(training): + # DPO, KTO and GKD use this shared method, without SFT's loss restoration. + trainer = SimpleNamespace( + template=SimpleNamespace(sequence_parallel_size=2), model=SimpleNamespace(training=training)) + with pytest.raises(NotImplementedError): + SwiftMixin.prepare_logits_to_keep(trainer, {'labels': torch.tensor([[-100, 1]])}) + + +@pytest.mark.parametrize('key', [ + 'pixel_values', 'pixel_values_videos', 'image_grid_thw', 'video_grid_thw', 'mm_token_type_ids', 'inputs_embeds', + 'missing_input_ids' +]) +def test_qwen_moe_sp_selection_rejects_non_text_inputs(monkeypatch, key): + model = SimpleNamespace(training=True, model_meta=get_matched_model_meta('Qwen/Qwen3.6-35B-A3B')) + trainer = make_trainer(model) + monkeypatch.setattr(sequence_parallel, 'rp_world_size', 1) + inputs = {'input_ids': torch.tensor([[1, 2]]), 'labels': torch.tensor([[-100, 2]])} + if key == 'missing_input_ids': + inputs.pop('input_ids') + else: + inputs[key] = torch.ones(1) + with pytest.raises(NotImplementedError, match='text-only inputs'): + Seq2SeqTrainer.prepare_logits_to_keep(trainer, inputs) + assert 'logits_to_keep' not in inputs