diff --git a/examples/v1/train_full/train_full_chunk_loss.yaml b/examples/v1/train_full/train_full_chunk_loss.yaml new file mode 100644 index 000000000..7b001fe66 --- /dev/null +++ b/examples/v1/train_full/train_full_chunk_loss.yaml @@ -0,0 +1,27 @@ +model: Qwen/Qwen3-0.6B +trust_remote_code: true +model_class: llm + +kernel_config: + name: auto + +# FSDP Config +dist_config: + name: fsdp2 + dcp_path: null + +### data +train_dataset: data/v1_sft_demo.yaml + +### training +output_dir: outputs/test_chunk_loss +micro_batch_size: 1 +cutoff_len: 2048 +# Maximum flattened token rows per logits/CE chunk; this is not the sequence length. +chunk_loss_size: 256 +learning_rate: 1.0e-4 +max_steps: 10 + +### sample +sample_backend: hf +max_new_tokens: 128 diff --git a/examples/v1/train_full/train_full_multimodal_ulysses_cp.yaml b/examples/v1/train_full/train_full_multimodal_ulysses_cp.yaml new file mode 100644 index 000000000..98d5b015b --- /dev/null +++ b/examples/v1/train_full/train_full_multimodal_ulysses_cp.yaml @@ -0,0 +1,25 @@ +model: Qwen/Qwen3.5-0.8B +trust_remote_code: true +model_class: llm + +flash_attn: flash_attention_2 + +# FSDP Config +dist_config: + name: fsdp2 + dcp_path: null + +cp_mode: ulysses +cp_size: 2 + +### data +train_dataset: data/v1_multimodal_demo.yaml + +### training +output_dir: outputs/test_multimodal_ulysses_cp +micro_batch_size: 1 +batching_strategy: normal +cutoff_len: 2048 +learning_rate: 1.0e-4 +bf16: false +max_steps: 10 diff --git a/src/llamafactory/v1/config/training_args.py b/src/llamafactory/v1/config/training_args.py index f86f80a6d..5c01d93e0 100644 --- a/src/llamafactory/v1/config/training_args.py +++ b/src/llamafactory/v1/config/training_args.py @@ -144,6 +144,10 @@ class TrainingArguments: default=1, metadata={"help": "Log metrics every N optimizer steps."}, ) + chunk_loss_size: int | None = field( + default=None, + metadata={"help": "Maximum flattened token rows per Chunk Loss chunk. None disables Chunk Loss."}, + ) pref_loss: Literal["sigmoid", "orpo", "simpo"] = field( default="sigmoid", metadata={"help": "The type of DPO loss to use."}, @@ -173,6 +177,8 @@ class TrainingArguments: self.dist_config = get_plugin_config(self.dist_config) self.optim_config = get_plugin_config(self.optim_config) self.lr_scheduler_config = get_plugin_config(self.lr_scheduler_config) + if self.chunk_loss_size is not None and self.chunk_loss_size <= 0: + raise ValueError("`chunk_loss_size` must be positive.") try: from ..plugins.model_plugins.deepspeed_utils import register_deepspeed_dist_config diff --git a/src/llamafactory/v1/core/base_trainer.py b/src/llamafactory/v1/core/base_trainer.py index 06aa55670..e04e8a9b2 100644 --- a/src/llamafactory/v1/core/base_trainer.py +++ b/src/llamafactory/v1/core/base_trainer.py @@ -239,7 +239,12 @@ class BaseTrainer: @abstractmethod def compute_loss(self, batch: BatchInput) -> Tensor: - """Compute the scalar loss.""" + """Compute the scalar loss. + + Subclasses must handle sequence-parallel layout and loss aggregation when + `self.cp_size > 1`, or reject context parallelism during initialization. + The shared training loop does not dispatch sequence-parallel loss. + """ ... def fit(self) -> None: @@ -265,14 +270,7 @@ class BaseTrainer: step_valid_tokens = DistributedInterface().all_reduce(step_valid_tokens, op=ReduceOp.SUM) num_micro = len(micro_batches) for i, micro_batch in enumerate(micro_batches): - if self.args.cp_size > 1: - from ..plugins.model_plugins.parallelization.sequence_parallel import ( - SequenceParallelLossPlugin, - ) - - loss = SequenceParallelLossPlugin("sequence_parallel_loss")(self.model, micro_batch) - else: - loss = self.compute_loss(micro_batch) + loss = self.compute_loss(micro_batch) mini_step_valid_tokens = compute_valid_tokens([micro_batch]) # fsdp uses mean reduction so we need to scale the loss by dp_size loss = loss * mini_step_valid_tokens * self.dp_size / (step_valid_tokens + 1e-6) diff --git a/src/llamafactory/v1/plugins/model_plugins/chunk_loss.py b/src/llamafactory/v1/plugins/model_plugins/chunk_loss.py new file mode 100644 index 000000000..8a4dd01aa --- /dev/null +++ b/src/llamafactory/v1/plugins/model_plugins/chunk_loss.py @@ -0,0 +1,218 @@ +# Copyright 2026 the LlamaFactory team. +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. + +"""Chunked linear cross-entropy for SFT.""" + +from __future__ import annotations + +from dataclasses import dataclass + +import torch +import torch.nn as nn +import torch.nn.functional as F +from torch import Tensor + +from ...utils.constants import IGNORE_INDEX +from ...utils.plugin import BasePlugin +from ...utils.types import BatchInput, HFModel, ModelOutput + + +class LossPlugin(BasePlugin): + def __call__(self, model: HFModel, chunk_size: int) -> ChunkLoss: + return super().__call__(model, chunk_size) + + +class _ChunkedLinearCrossEntropy(torch.autograd.Function): + @staticmethod + def forward( + ctx, + hidden_states: Tensor, + head_weight: Tensor, + head_bias: Tensor | None, + labels: Tensor, + loss_weights: Tensor, + chunk_size: int, + ) -> Tensor: + needs_hidden_grad, needs_weight_grad, needs_bias_grad = ctx.needs_input_grad[:3] + hidden_states_flat = hidden_states.reshape(-1, hidden_states.size(-1)) + labels_flat = labels.reshape(-1) + loss_weights_flat = loss_weights.reshape(-1) + + loss = torch.zeros((), device=hidden_states.device, dtype=torch.float32) + grad_hidden = torch.empty_like(hidden_states_flat) if needs_hidden_grad else None + # Avoid repeated BF16 rounding when summing head gradients across chunks. + grad_weight = torch.zeros_like(head_weight, dtype=torch.float32) if needs_weight_grad else None + grad_bias = ( + torch.zeros_like(head_bias, dtype=torch.float32) if head_bias is not None and needs_bias_grad else None + ) + + for start in range(0, hidden_states_flat.size(0), chunk_size): + end = start + chunk_size + with torch.enable_grad(): + hidden_arg = hidden_states_flat[start:end].detach().requires_grad_(needs_hidden_grad) + weight_arg = head_weight.detach().requires_grad_(needs_weight_grad) + bias_arg = head_bias.detach().requires_grad_(needs_bias_grad) if head_bias is not None else None + logits = F.linear(hidden_arg, weight_arg, bias_arg).float() + token_loss = F.cross_entropy( + logits, + labels_flat[start:end], + reduction="none", + ignore_index=IGNORE_INDEX, + ) + chunk_loss = (token_loss * loss_weights_flat[start:end]).sum() + grad_targets = [ + tensor + for tensor, needed in ( + (hidden_arg, needs_hidden_grad), + (weight_arg, needs_weight_grad), + (bias_arg, needs_bias_grad), + ) + if tensor is not None and needed + ] + chunk_grads = torch.autograd.grad(chunk_loss, grad_targets) if grad_targets else () + + loss.add_(chunk_loss.detach()) + grad_index = 0 + if grad_hidden is not None: + grad_hidden[start:end].copy_(chunk_grads[grad_index]) + grad_index += 1 + if grad_weight is not None: + grad_weight.add_(chunk_grads[grad_index]) + grad_index += 1 + if grad_bias is not None: + grad_bias.add_(chunk_grads[grad_index]) + + ctx.save_for_backward( + grad_hidden.reshape_as(hidden_states) if grad_hidden is not None else None, + grad_weight.to(head_weight.dtype) if grad_weight is not None else None, + grad_bias.to(head_bias.dtype) if grad_bias is not None else None, + ) + return loss + + @staticmethod + def backward(ctx, grad_output: Tensor): + grad_hidden, grad_weight, grad_bias = ctx.saved_tensors + return ( + grad_hidden * grad_output if grad_hidden is not None else None, + grad_weight * grad_output if grad_weight is not None else None, + grad_bias * grad_output if grad_bias is not None else None, + None, + None, + None, + ) + + +@dataclass +class _ChunkLossState: + labels: Tensor + loss_weights: Tensor + loss: Tensor | None = None + loss_version: int = 0 + + +@LossPlugin("chunk_loss").register() +class ChunkLoss: + """Install Chunk Loss before distributed wrapping. + + The model must call a plain Linear output head once and return its output + directly as logits, without subsequent transformations. Other forwards keep + their normal logits, including forwards of independently installed models. + """ + + def __init__(self, model: HFModel, chunk_size: int) -> None: + self._output_head = model.get_output_embeddings() + if type(self._output_head) is not nn.Linear: + raise TypeError("Chunk Loss requires `get_output_embeddings()` to return a plain `torch.nn.Linear`.") + + self.chunk_size = chunk_size + self._active_state: _ChunkLossState | None = None + self._original_forward = self._output_head.forward + self._output_head.forward = self._head_forward + model.register_forward_hook(self._check_model_output) + + def __call__( + self, + model: HFModel, + model_inputs: dict[str, Tensor], + labels: Tensor, + loss_weights: Tensor, + ) -> Tensor: + """Return a local weighted loss sum for already shifted targets.""" + state = _ChunkLossState(labels=labels, loss_weights=loss_weights) + self._active_state = state + try: + outputs: ModelOutput = model(**model_inputs) + finally: + self._active_state = None + + if state.loss is None: + raise RuntimeError("Chunk Loss did not reach the model output head.") + # Use the outer output so distributed wrappers retain their backward hooks. + return outputs.logits + + def compute_loss( + self, + model: HFModel, + batch: BatchInput, + *, + device: torch.device, + uses_mrope: bool, + ) -> Tensor: + """Prepare an unsharded SFT batch and compute its weighted mean Chunk Loss.""" + model_inputs = { + key: value.to(device, non_blocking=True) for key, value in batch.items() if isinstance(value, torch.Tensor) + } + labels = model_inputs.pop("labels") + loss_weights = model_inputs.pop("loss_weights") + if uses_mrope: + model_inputs.pop("position_ids", None) + + # Align each hidden state with its next-token target, as in the CP batch preparation. + labels = F.pad(labels[..., 1:].contiguous(), (0, 1), value=IGNORE_INDEX) + loss_weights = F.pad(loss_weights[..., 1:], (0, 1), value=0.0) + numerator = self(model, model_inputs, labels, loss_weights) + return numerator / (loss_weights.sum() + 1e-6) + + def _head_forward(self, hidden_states: Tensor) -> Tensor: + state = self._active_state + if state is None: + return self._original_forward(hidden_states) + if state.loss is not None: + raise RuntimeError("Chunk Loss expects one output-head call per model forward.") + if hidden_states.shape[:-1] != state.labels.shape or state.labels.shape != state.loss_weights.shape: + raise ValueError( + "Chunk Loss hidden states, labels, and loss weights must share the same token layout: " + f"hidden_states={tuple(hidden_states.shape)}, labels={tuple(state.labels.shape)}, " + f"loss_weights={tuple(state.loss_weights.shape)}." + ) + + state.loss = _ChunkedLinearCrossEntropy.apply( + hidden_states, + self._output_head.weight, + self._output_head.bias, + state.labels, + state.loss_weights, + self.chunk_size, + ) + state.loss_version = state.loss._version + return state.loss + + def _check_model_output(self, _model, _args, outputs: ModelOutput) -> None: + state = self._active_state + if state is None or state.loss is None: + return + # Validate before DDP and other wrappers can replace the output tensor. + # The version counter also catches in-place logits transformations. + if outputs.logits is not state.loss or state.loss._version != state.loss_version: + raise NotImplementedError("Chunk Loss does not support transformations after the model output head.") diff --git a/src/llamafactory/v1/plugins/model_plugins/parallelization/batch.py b/src/llamafactory/v1/plugins/model_plugins/parallelization/batch.py new file mode 100644 index 000000000..745d2bffe --- /dev/null +++ b/src/llamafactory/v1/plugins/model_plugins/parallelization/batch.py @@ -0,0 +1,108 @@ +# Copyright 2026 the LlamaFactory team. +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. + +"""Explicit batch ownership and sequence layout helpers for sequence parallelism.""" + +from dataclasses import dataclass + +import torch +import torch.nn.functional as F + +from ....utils.constants import IGNORE_INDEX +from ....utils.types import BatchInput, Tensor + + +# These tensors belong to a multimodal encoder tower and remain replicated +# until the outer model has fused them into the global language sequence. +MULTIMODAL_ENCODER_INPUT_KEYS = frozenset( + { + "pixel_values", + "image_grid_thw", + "pixel_values_videos", + "video_grid_thw", + "second_per_grid_ts", + "video_second_per_grid", + "input_features", + "feature_attention_mask", + } +) + +# Only tensors in this list are padded along the global language sequence. +# They remain full while the outer model performs multimodal fusion and mRoPE setup. +SEQUENCE_PARALLEL_INPUT_KEYS = frozenset( + { + "input_ids", + "attention_mask", + "position_ids", + "mm_token_type_ids", + } +) + + +@dataclass(frozen=True) +class PreparedSequenceParallelBatch: + """Full model inputs plus CP-local next-token targets.""" + + model_inputs: dict[str, Tensor] + local_shift_labels: Tensor + local_shift_loss_weights: Tensor + global_loss_weight_sum: Tensor + + +def split_sequence_tensor(tensor: Tensor, device_mesh, dim: int = -1) -> Tensor: + """Take the contiguous sequence shard owned by the local CP rank.""" + cp_mesh = device_mesh["cp"] + cp_size = cp_mesh.size() + sequence_length = tensor.shape[dim] + if sequence_length == 0 or sequence_length % cp_size != 0: + raise ValueError(f"Sequence length {sequence_length} must be positive and divisible by CP size {cp_size}.") + + cp_rank = cp_mesh.get_local_rank() + return torch.chunk(tensor, chunks=cp_size, dim=dim)[cp_rank].contiguous() + + +def prepare_sequence_parallel_batch( + batch: BatchInput, + *, + device: torch.device, + device_mesh, + uses_mrope: bool = False, +) -> PreparedSequenceParallelBatch: + """Pad the global language sequence while preserving encoder-owned layouts.""" + model_inputs = { + key: value.to(device, non_blocking=True) for key, value in batch.items() if isinstance(value, torch.Tensor) + } + labels = model_inputs.pop("labels") + loss_weights = model_inputs.pop("loss_weights") + + sequence_length = model_inputs["input_ids"].shape[-1] + cp_size = device_mesh["cp"].size() + pad_size = -sequence_length % cp_size + has_multimodal_inputs = bool(MULTIMODAL_ENCODER_INPUT_KEYS.intersection(model_inputs)) + if uses_mrope and has_multimodal_inputs: + model_inputs.pop("position_ids", None) + + # Only language tensors are padded; encoder tensors retain their original layouts. + for key in SEQUENCE_PARALLEL_INPUT_KEYS.intersection(model_inputs): + model_inputs[key] = F.pad(model_inputs[key], (0, pad_size), value=0) + + shift_labels = F.pad(labels[..., 1:], (0, pad_size + 1), value=IGNORE_INDEX) + shift_loss_weights = F.pad(loss_weights[..., 1:], (0, pad_size + 1), value=0.0) + + return PreparedSequenceParallelBatch( + model_inputs=model_inputs, + local_shift_labels=split_sequence_tensor(shift_labels, device_mesh), + local_shift_loss_weights=split_sequence_tensor(shift_loss_weights, device_mesh), + global_loss_weight_sum=shift_loss_weights.sum(), + ) diff --git a/src/llamafactory/v1/plugins/model_plugins/parallelization/gdn_attention.py b/src/llamafactory/v1/plugins/model_plugins/parallelization/gdn_attention.py index 9f28249f5..4edb4056d 100644 --- a/src/llamafactory/v1/plugins/model_plugins/parallelization/gdn_attention.py +++ b/src/llamafactory/v1/plugins/model_plugins/parallelization/gdn_attention.py @@ -17,6 +17,7 @@ import torch import torch.distributed as dist import torch.nn.functional as F +from ....utils import logging from .seq_comm import SeqAllToAll4D from .ulysses import ( get_ulysses_sequence_parallel_group, @@ -24,6 +25,9 @@ from .ulysses import ( ) +logger = logging.get_logger(__name__) + + def is_gdn_layer(layer) -> bool: """Return True if the module is a GDN (linear attention) layer or a DecoderLayer containing one.""" if hasattr(layer, "layer_type") and layer.layer_type == "linear_attention": @@ -107,6 +111,7 @@ def gdn_forward_with_cp(self, hidden_states, attention_mask=None, **kwargs): global_position_ids = torch.cat(global_position_ids, dim=-1).contiguous() try: from transformers.modeling_flash_attention_utils import prepare_fa_kwargs_from_position_ids + cu_seqlens = prepare_fa_kwargs_from_position_ids(global_position_ids)[0][0] except ImportError: cu_seqlens = None @@ -248,3 +253,19 @@ def gdn_forward_with_cp(self, hidden_states, attention_mask=None, **kwargs): # Output projection in CP layout output = self.out_proj(norm_out) return output + + +def apply_gdn_attention(model, cp_size: int) -> None: + """Install the sequence-parallel GDN forward on each unique linear-attention module.""" + if cp_size > 1: + replaced_modules = set() + for name, module in model.named_modules(): + if is_gdn_layer(module): + gdn_module = _get_gdn_module(module) + if id(gdn_module) in replaced_modules: + continue + replaced_modules.add(id(gdn_module)) + gdn_module.original_forward = gdn_module.forward + gdn_module.forward = gdn_forward_with_cp.__get__(gdn_module, type(gdn_module)) + gdn_name = name if gdn_module is module else f"{name}.linear_attn" + logger.info_rank0(f"Replaced GDN forward in {gdn_name} with gdn_forward_with_cp for context parallel.") diff --git a/src/llamafactory/v1/plugins/model_plugins/parallelization/hook.py b/src/llamafactory/v1/plugins/model_plugins/parallelization/hook.py new file mode 100644 index 000000000..fff388dfa --- /dev/null +++ b/src/llamafactory/v1/plugins/model_plugins/parallelization/hook.py @@ -0,0 +1,119 @@ +# Copyright 2026 the LlamaFactory team. +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. + +"""Common language-boundary hook for text-only and multimodal CP.""" + +from __future__ import annotations + +from typing import TYPE_CHECKING, Any + +from torch import nn + +from ....accelerator.interface import Dim, DistributedInterface +from ....utils import logging +from .batch import MULTIMODAL_ENCODER_INPUT_KEYS, split_sequence_tensor + + +if TYPE_CHECKING: + from ....utils.types import HFModel + + +logger = logging.get_logger(__name__) + +_MULTIMODAL_ENCODER_NAMES = ("visual", "audio_tower") + + +def _resolve_multimodal_boundary(model: HFModel) -> nn.Module | None: + cores = (model, getattr(model, "model", None)) + for boundary_name in ("language_model", "model"): + for core in cores: + has_encoder = any(isinstance(getattr(core, name, None), nn.Module) for name in _MULTIMODAL_ENCODER_NAMES) + boundary = getattr(core, boundary_name, None) + if has_encoder and isinstance(boundary, nn.Module): + return boundary + + return None + + +def _split_deepstack_inputs(kwargs: dict[str, Any], device_mesh) -> None: + """Align optional DeepStack visual rows with the local language shard.""" + visual_pos_masks = kwargs.get("visual_pos_masks") + if visual_pos_masks is None: + return + + deepstack_visual_embeds = kwargs["deepstack_visual_embeds"] + visual_ordinals = visual_pos_masks.reshape(-1).long().cumsum(dim=0) - 1 + visual_ordinals = visual_ordinals.view_as(visual_pos_masks) + local_visual_pos_masks = split_sequence_tensor(visual_pos_masks, device_mesh) + local_visual_ordinals = split_sequence_tensor(visual_ordinals, device_mesh)[local_visual_pos_masks] + + kwargs["visual_pos_masks"] = local_visual_pos_masks + kwargs["deepstack_visual_embeds"] = [ + visual_embeds.index_select(0, local_visual_ordinals.to(visual_embeds.device)) + for visual_embeds in deepstack_visual_embeds + ] + + +def install_sequence_parallel_hook(model: HFModel) -> None: + """Install the common CP split at the model's language boundary.""" + get_base_model = getattr(model, "get_base_model", None) + if callable(get_base_model): + model = get_base_model() + + boundary = _resolve_multimodal_boundary(model) + requires_fused_inputs = boundary is not None + if boundary is None: + boundary = model.base_model + + device_mesh = DistributedInterface().get_device_mesh(Dim.CP) + + def sequence_parallel_pre_hook(_module, args, kwargs): + encoder_inputs = {key for key in MULTIMODAL_ENCODER_INPUT_KEYS if kwargs.get(key) is not None} + if not requires_fused_inputs and encoder_inputs: + raise ValueError( + "Sequence parallelism reached a text language boundary with multimodal encoder inputs " + f"{sorted(encoder_inputs)}; this model structure is not supported." + ) + + input_ids = kwargs.get("input_ids") + inputs_embeds = kwargs.get("inputs_embeds") + if requires_fused_inputs and input_ids is not None: + raise ValueError( + "Multimodal sequence parallelism must enter the language boundary through fused " + "`inputs_embeds`; received non-null `input_ids`." + ) + + uses_inputs_embeds = inputs_embeds is not None + sequence_tensor = inputs_embeds if uses_inputs_embeds else input_ids + sequence_length = sequence_tensor.shape[1] + + attention_mask = kwargs.get("attention_mask") + + position_ids = kwargs["position_ids"] + if position_ids.shape[-1] != sequence_length: + raise ValueError("position_ids must match the global sequence length before CP splitting.") + + _split_deepstack_inputs(kwargs, device_mesh) + + sequence_name = "inputs_embeds" if uses_inputs_embeds else "input_ids" + kwargs[sequence_name] = split_sequence_tensor(sequence_tensor, device_mesh, dim=1) + if attention_mask is not None: + kwargs["attention_mask"] = split_sequence_tensor(attention_mask, device_mesh) + kwargs["position_ids"] = split_sequence_tensor(position_ids, device_mesh) + + kwargs["use_cache"] = False + return args, kwargs + + boundary.register_forward_pre_hook(sequence_parallel_pre_hook, with_kwargs=True) + logger.info_rank0("Installed sequence-parallel pre-hook at the language boundary.") diff --git a/src/llamafactory/v1/plugins/model_plugins/parallelization/sequence_parallel.py b/src/llamafactory/v1/plugins/model_plugins/parallelization/sequence_parallel.py index 517d9a7a4..fd2f53ef6 100644 --- a/src/llamafactory/v1/plugins/model_plugins/parallelization/sequence_parallel.py +++ b/src/llamafactory/v1/plugins/model_plugins/parallelization/sequence_parallel.py @@ -12,29 +12,16 @@ # See the License for the specific language governing permissions and # limitations under the License. -import sys -from functools import partial - import torch import torch.distributed as dist import torch.nn.functional as F -import transformers from ....accelerator.interface import Dim, DistributedInterface -from ....utils import logging +from ....utils.constants import IGNORE_INDEX from ....utils.plugin import BasePlugin -from ....utils.types import ModelOutput -from .gdn_attention import _get_gdn_module, gdn_forward_with_cp, is_gdn_layer -from .ulysses import ( - UlyssesAttention, - get_ulysses_sequence_parallel_group, - get_ulysses_sequence_parallel_rank, - get_ulysses_sequence_parallel_world_size, - set_ulysses_sequence_parallel_group, -) - - -logger = logging.get_logger(__name__) +from .batch import prepare_sequence_parallel_batch +from .gdn_attention import apply_gdn_attention +from .ulysses import apply_ulysses_attention class SequenceParallelModelPlugin(BasePlugin): @@ -47,171 +34,48 @@ class SequenceParallelLossPlugin(BasePlugin): return super().__call__(model, inputs, *args, **kwargs) -def new_flash_attn_forward( - query_states, - key_states, - value_states, - attention_mask, - sequence_parallel_size=1, - dropout=0, - deterministic=False, - is_causal=True, - group=None, - mode="ulysses", - attn_fn=None, - target_dtype=None, - **kwargs, -): - if mode == "ulysses": - dist_attn = UlyssesAttention(sequence_process_group=group, attn_fn=attn_fn) - attn_output = dist_attn( - query_states, - key_states, - value_states, - attention_mask, - query_length=query_states.shape[1] * sequence_parallel_size, - deterministic=deterministic, - dropout_p=dropout, - causal=is_causal, - position_ids=kwargs.get("position_ids", None), - target_dtype=target_dtype, - ) - else: - raise NotImplementedError("Other sequence parallel modes are to be implemented.") - - return attn_output - - @SequenceParallelModelPlugin("ulysses").register() def apply_sequence_parallel(model, cp_size: int): - # Replace _flash_attention_forward with new_flash_attn_forward - module = sys.modules[model.__module__] + from .hook import install_sequence_parallel_hook - set_ulysses_sequence_parallel_group(DistributedInterface().get_group(Dim.CP)) - - try: - num_attention_heads, num_key_value_heads = ( - model.config.num_attention_heads, - model.config.num_key_value_heads, - ) - except AttributeError: - num_attention_heads, num_key_value_heads = ( - model.config.text_config.num_attention_heads, - model.config.text_config.num_key_value_heads, - ) - - assert num_attention_heads % cp_size == 0, "num_attention_heads must be divisible by cp_size" - assert num_key_value_heads % cp_size == 0 or cp_size % num_key_value_heads == 0, ( - "num_key_value_heads must be divisible by cp_size" - ) - - origin_attn = transformers.modeling_flash_attention_utils._flash_attention_forward - new_flash_attention_forward = partial( - new_flash_attn_forward, - group=get_ulysses_sequence_parallel_group(), - mode="ulysses", - attn_fn=origin_attn, - sequence_parallel_size=cp_size, - ) - - for module_name, module in list(sys.modules.items()): - try: - if ( - hasattr(module, "__file__") - and "transformers" in module.__file__ - and getattr(module._flash_attention_forward, "__name__", "") == "_flash_attention_forward" - ): - module._flash_attention_forward = new_flash_attention_forward - logger.info_rank0( - f"Replaced _flash_attention_forward in module {module_name} with new_flash_attn_forward for sequence parallel." - ) - except (AttributeError, TypeError): - continue - - # Register GDN forward for CP support - if cp_size > 1: - replaced_modules = set() - for name, module in model.named_modules(): - if is_gdn_layer(module): - gdn_module = _get_gdn_module(module) - if id(gdn_module) in replaced_modules: - continue - replaced_modules.add(id(gdn_module)) - gdn_module.original_forward = gdn_module.forward - gdn_module.forward = gdn_forward_with_cp.__get__(gdn_module, type(gdn_module)) - gdn_name = name if gdn_module is module else f"{name}.linear_attn" - logger.info_rank0(f"Replaced GDN forward in {gdn_name} with gdn_forward_with_cp for context parallel.") - - -def padding_and_split_data(data, device_mesh=None): - if device_mesh is not None: - cp_size = device_mesh["cp"].size() - cp_rank = device_mesh["cp"].get_local_rank() - cp_group = device_mesh["cp"].get_group() - for k, v in data.items(): - if isinstance(v, torch.Tensor) and v.ndim > 1: - data_len = torch.tensor(v.shape[-1], device=v.device, dtype=torch.int64) - global_data_len = [torch.empty_like(data_len) for _ in range(cp_size)] - dist.all_gather(global_data_len, data_len, group=cp_group) - max_data_len = max(global_data_len) - pad_size = max_data_len - v.shape[-1] + (cp_size - max_data_len % cp_size) % cp_size - if k == "labels": - pad_value = -100 - elif k == "loss_weights": - pad_value = 0.0 - else: - pad_value = 0 - pad_data = F.pad(v, (0, pad_size), value=pad_value) - data[k] = torch.chunk(pad_data, chunks=cp_size, dim=-1)[cp_rank].contiguous() - return data + install_sequence_parallel_hook(model) + group = DistributedInterface().get_group(Dim.CP) + apply_ulysses_attention(model, cp_size, group) + apply_gdn_attention(model, cp_size) @SequenceParallelLossPlugin("sequence_parallel_loss").register() -def sequence_parallel_loss(model, model_inputs): +def sequence_parallel_loss(model, model_inputs, loss_fn=None, *, uses_mrope: bool = False): + """Prepare CP targets and aggregate weighted CE, optionally using a custom loss function. + + ``loss_fn`` receives ``(model, model_inputs, labels, loss_weights)``. Labels + and weights are already shifted globally and sharded for the local CP rank. + It must return a differentiable FP32 scalar weighted loss sum, without + shifting targets again, normalizing, or performing CP collectives. + """ device_mesh = DistributedInterface().get_device_mesh(Dim.CP) - model_inputs = { - k: v.to(dist.get_rank(), non_blocking=True) for k, v in model_inputs.items() if isinstance(v, torch.Tensor) - } + prepared = prepare_sequence_parallel_batch( + model_inputs, + device=DistributedInterface().current_device, + device_mesh=device_mesh, + uses_mrope=uses_mrope, + ) + labels = prepared.local_shift_labels + loss_weights = prepared.local_shift_loss_weights + if loss_fn is None: + logits = model(**prepared.model_inputs).logits.float() + token_loss = F.cross_entropy( + logits.reshape(-1, logits.size(-1)), labels.reshape(-1), reduction="none", ignore_index=IGNORE_INDEX + ) + local_numerator = (token_loss * loss_weights.reshape(-1)).sum() + else: + local_numerator = loss_fn(model, prepared.model_inputs, labels, loss_weights) + cp_group = device_mesh["cp"].get_group() - model_inputs = padding_and_split_data(model_inputs, device_mesh) - - batch_size, _ = model_inputs["labels"].shape - - outputs: ModelOutput = model(**model_inputs) - - logits = outputs.logits.float() - - labels = model_inputs["labels"] - - cp_group = get_ulysses_sequence_parallel_group() - cp_world_size = get_ulysses_sequence_parallel_world_size(cp_group) - cp_rank = get_ulysses_sequence_parallel_rank(cp_group) - - # use all_gather to collect labels from all sequence parallel processes - global_labels = [torch.empty_like(labels) for _ in range(cp_world_size)] - dist.all_gather(global_labels, labels, group=cp_group) - labels = torch.cat(global_labels, dim=1).contiguous() - shift_labels = labels[..., 1:].contiguous() - shift_labels = F.pad(shift_labels, (0, 1), value=-100) - shift_labels = torch.chunk(shift_labels, chunks=cp_world_size, dim=1)[cp_rank].contiguous() - - # use all_gather to collect loss_weights from all sequence parallel processes - loss_weights = model_inputs["loss_weights"] - global_loss_weights = [torch.empty_like(loss_weights) for _ in range(cp_world_size)] - dist.all_gather(global_loss_weights, loss_weights, group=cp_group) - shift_loss_weights = torch.cat(global_loss_weights, dim=1).contiguous() - shift_loss_weights = shift_loss_weights[..., 1:].contiguous() - - shift_logits = logits.view(-1, logits.size(-1)).contiguous() - shift_labels = shift_labels.view(-1).contiguous() - - # use all_gather to collect log_probs from all sequence parallel processes - log_probs = -F.cross_entropy(shift_logits, shift_labels, reduction="none").view(batch_size, -1) - global_log_probs = dist.nn.all_gather(log_probs, group=cp_group) - global_log_probs = torch.cat(global_log_probs, dim=1).contiguous() - log_probs = global_log_probs[..., :-1].contiguous() - - loss = (-log_probs * shift_loss_weights).sum() / (shift_loss_weights.sum() + 1e-6) - - return loss + # Do not average local mean losses: CP shards can own different supervised-token weights. + # Gather the differentiable weighted numerators instead, reducing communication from + # [batch, local_sequence] log probabilities to one scalar per CP rank. + global_loss_numerators = dist.nn.all_gather(local_numerator.reshape(1), group=cp_group) + global_loss_numerator = torch.cat(global_loss_numerators).sum() + return global_loss_numerator / (prepared.global_loss_weight_sum + 1e-6) diff --git a/src/llamafactory/v1/plugins/model_plugins/parallelization/ulysses.py b/src/llamafactory/v1/plugins/model_plugins/parallelization/ulysses.py index 6ba744da9..3c8bfe9b2 100644 --- a/src/llamafactory/v1/plugins/model_plugins/parallelization/ulysses.py +++ b/src/llamafactory/v1/plugins/model_plugins/parallelization/ulysses.py @@ -15,16 +15,22 @@ # See the License for the specific language governing permissions and # limitations under the License. +import sys +from functools import partial from typing import Any, Optional import torch import torch.distributed as dist +import transformers from torch import Tensor from torch.distributed import ProcessGroup +from ....utils import logging from .seq_comm import SeqAllToAll4D +logger = logging.get_logger(__name__) + _ULYSSES_SEQUENCE_PARALLEL_GROUP = None @@ -187,3 +193,95 @@ class UlyssesAttention(torch.nn.Module): # out e.g., [s/p::h] return output + + +def new_flash_attn_forward( + query_states, + key_states, + value_states, + attention_mask, + sequence_parallel_size=1, + dropout=0, + deterministic=False, + is_causal=True, + group=None, + mode="ulysses", + attn_fn=None, + target_dtype=None, + **kwargs, +): + """Route causal language attention through Ulysses and leave replicated encoders native.""" + if mode == "ulysses": + if not is_causal: + return attn_fn( + query_states, + key_states, + value_states, + attention_mask, + is_causal=False, + dropout=dropout, + deterministic=deterministic, + target_dtype=target_dtype, + **kwargs, + ) + + dist_attn = UlyssesAttention(sequence_process_group=group, attn_fn=attn_fn) + attn_output = dist_attn( + query_states, + key_states, + value_states, + attention_mask, + query_length=query_states.shape[1] * sequence_parallel_size, + deterministic=deterministic, + dropout_p=dropout, + causal=is_causal, + position_ids=kwargs.get("position_ids", None), + target_dtype=target_dtype, + ) + else: + raise NotImplementedError("Other sequence parallel modes are to be implemented.") + + return attn_output + + +def apply_ulysses_attention(model, cp_size: int, group: dist.ProcessGroup) -> None: + """Validate and install the Ulysses FlashAttention bridge for one process group.""" + # Replace _flash_attention_forward with new_flash_attn_forward + set_ulysses_sequence_parallel_group(group) + + try: + num_attention_heads, num_key_value_heads = ( + model.config.num_attention_heads, + model.config.num_key_value_heads, + ) + except AttributeError: + num_attention_heads, num_key_value_heads = ( + model.config.text_config.num_attention_heads, + model.config.text_config.num_key_value_heads, + ) + + assert num_attention_heads % cp_size == 0, "num_attention_heads must be divisible by cp_size" + assert num_key_value_heads % cp_size == 0, "num_key_value_heads must be divisible by cp_size" + + origin_attn = transformers.modeling_flash_attention_utils._flash_attention_forward + new_flash_attention_forward = partial( + new_flash_attn_forward, + group=get_ulysses_sequence_parallel_group(), + mode="ulysses", + attn_fn=origin_attn, + sequence_parallel_size=cp_size, + ) + + for module_name, module in list(sys.modules.items()): + try: + if ( + hasattr(module, "__file__") + and "transformers" in module.__file__ + and getattr(module._flash_attention_forward, "__name__", "") == "_flash_attention_forward" + ): + module._flash_attention_forward = new_flash_attention_forward + logger.info_rank0( + f"Replaced _flash_attention_forward in module {module_name} with new_flash_attn_forward for sequence parallel." + ) + except (AttributeError, TypeError): + continue diff --git a/src/llamafactory/v1/trainers/dpo_trainer.py b/src/llamafactory/v1/trainers/dpo_trainer.py index fb2dfb59d..b22512373 100644 --- a/src/llamafactory/v1/trainers/dpo_trainer.py +++ b/src/llamafactory/v1/trainers/dpo_trainer.py @@ -96,6 +96,8 @@ class DPOTrainer(BaseTrainer): ) -> None: if args.cp_size > 1: raise NotImplementedError("DPO trainer currently only supports cp_size == 1.") + if args.chunk_loss_size is not None: + raise NotImplementedError("Chunk Loss currently only supports SFT training.") self.pref_loss = args.pref_loss self.pref_beta = args.pref_beta diff --git a/src/llamafactory/v1/trainers/rm_trainer.py b/src/llamafactory/v1/trainers/rm_trainer.py index 78694d47b..2a32ae2c6 100644 --- a/src/llamafactory/v1/trainers/rm_trainer.py +++ b/src/llamafactory/v1/trainers/rm_trainer.py @@ -78,6 +78,8 @@ class RMTrainer(BaseTrainer): ) -> None: if args.cp_size > 1: raise NotImplementedError("RM trainer currently only supports cp_size == 1.") + if args.chunk_loss_size is not None: + raise NotImplementedError("Chunk Loss currently only supports SFT training.") super().__init__(args, model, renderer, train_dataset, callbacks) diff --git a/src/llamafactory/v1/trainers/sft_trainer.py b/src/llamafactory/v1/trainers/sft_trainer.py index d84263283..cc92bac04 100644 --- a/src/llamafactory/v1/trainers/sft_trainer.py +++ b/src/llamafactory/v1/trainers/sft_trainer.py @@ -14,15 +14,42 @@ from ..accelerator.interface import DistributedInterface -from ..config import InputArgument, get_args +from ..config import InputArgument, TrainingArguments, get_args from ..core.base_trainer import BaseTrainer from ..core.data_engine import DataEngine from ..core.model_engine import ModelEngine -from ..utils.types import BatchInput, Tensor +from ..core.rendering import Renderer +from ..utils.callbacks import TrainerCallback +from ..utils.types import BatchInput, HFModel, Tensor, TorchDataset class SFTTrainer(BaseTrainer): + def __init__( + self, + args: TrainingArguments, + model: HFModel, + renderer: Renderer, + train_dataset: TorchDataset, + callbacks: list[TrainerCallback] | None = None, + ) -> None: + self._chunk_loss = None + if args.chunk_loss_size is not None: + from ..plugins.model_plugins.chunk_loss import LossPlugin + + self._chunk_loss = LossPlugin("chunk_loss")(model, args.chunk_loss_size) + + super().__init__(args, model, renderer, train_dataset, callbacks) + def compute_loss(self, batch: BatchInput) -> Tensor: + if self.cp_size > 1: + from ..plugins.model_plugins.parallelization.sequence_parallel import SequenceParallelLossPlugin + + return SequenceParallelLossPlugin("sequence_parallel_loss")( + self.model, batch, loss_fn=self._chunk_loss, uses_mrope=self._uses_mrope + ) + if self._chunk_loss is not None: + return self._chunk_loss.compute_loss(self.model, batch, device=self.device, uses_mrope=self._uses_mrope) + shift_loss_weights = batch["loss_weights"].to(self.device, non_blocking=True)[..., 1:] log_probs = self.compute_log_probs(self.model, batch) loss = (-log_probs * shift_loss_weights).sum() / (shift_loss_weights.sum() + 1e-6) diff --git a/tests_v1/plugins/model_plugins/test_chunk_loss.py b/tests_v1/plugins/model_plugins/test_chunk_loss.py new file mode 100644 index 000000000..b9a5736ea --- /dev/null +++ b/tests_v1/plugins/model_plugins/test_chunk_loss.py @@ -0,0 +1,154 @@ +# Copyright 2026 the LlamaFactory team. +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. + +from copy import deepcopy + +import pytest +import torch +import torch.distributed as dist +import torch.nn.functional as F +from torch import nn +from torch.nn.parallel import DistributedDataParallel as DDP +from transformers.modeling_outputs import CausalLMOutput + +from llamafactory.v1.plugins.model_plugins.chunk_loss import LossPlugin, _ChunkedLinearCrossEntropy +from llamafactory.v1.trainers.sft_trainer import SFTTrainer +from llamafactory.v1.utils.constants import IGNORE_INDEX +from llamafactory.v1.utils.env import find_available_port +from llamafactory.v1.utils.pytest import dist_env + + +class _TinyCausalLM(nn.Module): + def __init__(self): + super().__init__() + self.embed_tokens = nn.Embedding(31, 16) + self.lm_head = nn.Linear(16, 31, bias=False) + + def get_output_embeddings(self): + return self.lm_head + + def forward(self, input_ids, **_): + return CausalLMOutput(logits=self.lm_head(self.embed_tokens(input_ids))) + + +def _make_model(): + return _TinyCausalLM() + + +def _assert_gradients_close(actual: torch.Tensor, expected: torch.Tensor) -> None: + error = torch.linalg.vector_norm(actual.float() - expected.float()) + reference = torch.linalg.vector_norm(expected.float()) + assert error <= 2 * torch.finfo(expected.dtype).eps * reference + + +def _weighted_cross_entropy(logits, labels, loss_weights): + losses = F.cross_entropy(logits.flatten(0, -2).float(), labels.flatten(), reduction="none") + return (losses * loss_weights.flatten()).sum() + + +@pytest.mark.parametrize("frozen_head", [False, True]) +def test_chunk_loss_matches_eager_loss_and_gradients(frozen_head): + torch.manual_seed(0) + eager_head = nn.Linear(4, 7).to(torch.bfloat16) + eager_head.requires_grad_(not frozen_head) + chunk_head = deepcopy(eager_head) + eager_hidden = torch.randn(2, 5, 4, dtype=torch.bfloat16, requires_grad=True) + chunk_hidden = eager_hidden.detach().clone().requires_grad_() + labels = torch.tensor([[0, 1, IGNORE_INDEX, 3, 4], [5, 6, 0, 1, 2]]) + loss_weights = torch.tensor([[0.0, 0.25, 1.0, 0.75, 1.0], [1.0, 0.5, 0.0, 0.25, 1.0]]) + + eager_loss = _weighted_cross_entropy(eager_head(eager_hidden), labels, loss_weights) + chunk_loss = _ChunkedLinearCrossEntropy.apply( + chunk_hidden, chunk_head.weight, chunk_head.bias, labels, loss_weights, 3 + ) + scale = 0.07 / (loss_weights.sum() + 1e-6) + (eager_loss * scale).backward() + (chunk_loss * scale).backward() + + torch.testing.assert_close(chunk_loss, eager_loss) + _assert_gradients_close(chunk_hidden.grad, eager_hidden.grad) + for actual, expected in zip(chunk_head.parameters(), eager_head.parameters()): + if frozen_head: + assert actual.grad is None + else: + _assert_gradients_close(actual.grad, expected.grad) + + +@pytest.mark.parametrize("zero_supervision", [False, True]) +def test_chunk_sft_loss_matches_reference(zero_supervision): + model = _make_model() + input_ids = torch.tensor([[1, 2, 3, 4, 5], [6, 7, 8, 9, 10]]) + labels = input_ids.clone() + labels[0, 2] = IGNORE_INDEX + batch = { + "input_ids": input_ids, + "attention_mask": torch.ones_like(input_ids), + "position_ids": torch.arange(5).expand(2, -1), + "labels": labels, + "loss_weights": torch.tensor([[0.0, 0.25, 0.0, 1.0, 0.5], [0.0, 0.0, 0.75, 0.0, 1.0]]), + } + if zero_supervision: + batch["labels"].fill_(IGNORE_INDEX) + batch["loss_weights"].zero_() + original_batch = {key: value.clone() for key, value in batch.items()} + + trainer = object.__new__(SFTTrainer) + trainer.model = model + trainer.device = torch.device("cpu") + trainer.cp_size = 1 + trainer._uses_mrope = False + trainer._chunk_loss = None + eager_loss = trainer.compute_loss(batch) + + chunk_model = deepcopy(model) + trainer.model = chunk_model + trainer._chunk_loss = LossPlugin("chunk_loss")(chunk_model, chunk_size=3) + chunk_loss = trainer.compute_loss(batch) + + torch.testing.assert_close(chunk_loss, eager_loss) + for key in batch: + torch.testing.assert_close(batch[key], original_batch[key]) + torch.testing.assert_close(chunk_model(input_ids=input_ids).logits, model(input_ids=input_ids).logits) + + +@pytest.mark.skipif(not dist.is_available() or not dist.is_gloo_available(), reason="Requires the CPU Gloo backend.") +def test_chunk_loss_preserves_ddp_output_backward_hooks(): + torch.manual_seed(7) + eager_model = _make_model() + chunk_model = deepcopy(eager_model) + loss_fn = LossPlugin("chunk_loss")(chunk_model, chunk_size=2) + input_ids = torch.tensor([[1, 2, 3, 4, 5]]) + model_inputs = {"input_ids": input_ids, "use_cache": False} + labels = torch.tensor([[2, 3, 4, 5, IGNORE_INDEX]]) + loss_weights = torch.tensor([[0.0, 0.25, 1.0, 0.5, 0.0]]) + outer_scale = 0.3 / loss_weights.sum() + outer_outputs = [] + + def capture_outer_output(_model, _args, output): + outer_outputs.append(output.logits) + + with dist_env(master_port=find_available_port()): + dist.init_process_group("gloo") + wrapped_model = DDP(chunk_model, find_unused_parameters=True) + wrapped_model.register_forward_hook(capture_outer_output) + eager_loss = _weighted_cross_entropy(eager_model(**model_inputs).logits, labels, loss_weights) + chunk_loss = loss_fn(wrapped_model, model_inputs, labels, loss_weights) + + assert loss_fn._active_state is None + assert chunk_loss is outer_outputs.pop() + torch.testing.assert_close(chunk_loss, eager_loss) + (eager_loss * outer_scale).backward() + (chunk_loss * outer_scale).backward() + for expected, actual in zip(eager_model.parameters(), chunk_model.parameters(), strict=True): + torch.testing.assert_close(actual.grad, expected.grad) diff --git a/tests_v1/plugins/model_plugins/test_ulysses_cp.py b/tests_v1/plugins/model_plugins/test_ulysses_cp.py index e16286266..4d4fcb79f 100644 --- a/tests_v1/plugins/model_plugins/test_ulysses_cp.py +++ b/tests_v1/plugins/model_plugins/test_ulysses_cp.py @@ -12,19 +12,25 @@ # See the License for the specific language governing permissions and # limitations under the License. +from types import SimpleNamespace + import pytest import torch import torch.multiprocessing as mp +from torch import nn +import llamafactory.v1.plugins.model_plugins.parallelization.hook as hook_module from llamafactory.v1.accelerator.interface import DistributedInterface from llamafactory.v1.config.model_args import ModelArguments from llamafactory.v1.config.training_args import TrainingArguments from llamafactory.v1.core.model_engine import ModelEngine from llamafactory.v1.plugins.model_plugins.parallelization import ulysses +from llamafactory.v1.plugins.model_plugins.parallelization.batch import prepare_sequence_parallel_batch from llamafactory.v1.plugins.model_plugins.parallelization.sequence_parallel import ( SequenceParallelModelPlugin, sequence_parallel_loss, ) +from llamafactory.v1.utils.constants import IGNORE_INDEX from llamafactory.v1.utils.env import find_available_port from llamafactory.v1.utils.pytest import dist_env @@ -99,3 +105,102 @@ def test_sequence_parallel_loss(cp_size, dp_size, batch_size): mp.spawn( _test_sequence_parallel_loss, args=(world_size, master_port, cp_size, dp_size, batch_size), nprocs=world_size ) + + +def test_non_causal_multimodal_encoder_attention_bypasses_ulysses(): + captured_is_causal = None + + def fake_native_attention(query, _key, _value, _attention_mask, **kwargs): + nonlocal captured_is_causal + captured_is_causal = kwargs["is_causal"] + return query + 1 + + query = torch.zeros(1, 4, 2, 8) + output = ulysses.new_flash_attn_forward(query, query, query, None, is_causal=False, attn_fn=fake_native_attention) + + torch.testing.assert_close(output, query + 1) + assert captured_is_causal is False + + +def _device_mesh(rank=0, size=2): + return {"cp": SimpleNamespace(size=lambda: size, get_local_rank=lambda: rank)} + + +class _RecordingLanguageModel(nn.Module): + def forward(self, **kwargs): + return kwargs + + +def test_multimodal_sequence_parallel_hook(monkeypatch): + # One shard gets non-contiguous visual rows while the next shard is empty. + cp_rank = [1] + device_mesh = {"cp": SimpleNamespace(size=lambda: 3, get_local_rank=lambda: cp_rank[0])} + distributed = SimpleNamespace(get_device_mesh=lambda _dim: device_mesh) + monkeypatch.setattr(hook_module, "DistributedInterface", lambda: distributed) + + model = nn.Module() + model.model = core = nn.Module() + boundary = _RecordingLanguageModel() + core.visual = nn.Identity() + core.language_model = boundary + hook_module.install_sequence_parallel_hook(SimpleNamespace(get_base_model=lambda: model)) + + fused_inputs = torch.arange(24, dtype=torch.float32).view(2, 6, 2) + attention_mask = torch.tensor([[1, 1, 1, 1, 1, 0], [1, 1, 1, 1, 0, 0]]) + position_ids = torch.arange(36).view(3, 2, 6) + visual_mask = torch.tensor([[True, False, False, True, False, False], [True, True, True, False, False, False]]) + visual_embeds = torch.arange(10, dtype=torch.float32).view(5, 2).requires_grad_() + model_inputs = { + "input_ids": None, + "inputs_embeds": fused_inputs, + "attention_mask": attention_mask, + "position_ids": position_ids, + "visual_pos_masks": visual_mask, + "deepstack_visual_embeds": [visual_embeds], + } + + outputs = boundary(**model_inputs) + torch.testing.assert_close(outputs["inputs_embeds"], fused_inputs[:, 2:4]) + torch.testing.assert_close(outputs["attention_mask"], attention_mask[:, 2:4]) + torch.testing.assert_close(outputs["position_ids"], position_ids[..., 2:4]) + torch.testing.assert_close(outputs["visual_pos_masks"], visual_mask[:, 2:4]) + torch.testing.assert_close(outputs["deepstack_visual_embeds"][0], visual_embeds[[1, 4]]) + assert outputs["use_cache"] is False + + outputs["deepstack_visual_embeds"][0].sum().backward() + expected_grad = torch.zeros_like(visual_embeds) + expected_grad[[1, 4]] = 1 + torch.testing.assert_close(visual_embeds.grad, expected_grad) + + cp_rank[0] = 2 + visual_embeds.grad = None + empty_visual_embeds = boundary(**model_inputs)["deepstack_visual_embeds"][0] + assert empty_visual_embeds.shape == (0, 2) + empty_visual_embeds.sum().backward() + torch.testing.assert_close(visual_embeds.grad, torch.zeros_like(visual_embeds)) + + +def test_prepare_multimodal_sequence_parallel_batch_preserves_encoder_inputs_and_shifts_targets(): + pixel_values = torch.arange(12, dtype=torch.float32).view(3, 4) + batch = { + "input_ids": torch.tensor([[1, 2, 3]]), + "attention_mask": torch.ones(1, 3, dtype=torch.long), + "position_ids": torch.tensor([[0, 1, 2]]), + "mm_token_type_ids": torch.tensor([[0, 1, 1]]), + "labels": torch.tensor([[1, 2, 3]]), + "loss_weights": torch.tensor([[9.0, 0.5, 2.0]]), + "pixel_values": pixel_values, + } + for rank in range(2): + prepared = prepare_sequence_parallel_batch( + batch, device=torch.device("cpu"), device_mesh=_device_mesh(rank), uses_mrope=True + ) + assert prepared.model_inputs["input_ids"].tolist() == [[1, 2, 3, 0]] + assert prepared.model_inputs["mm_token_type_ids"].tolist() == [[0, 1, 1, 0]] + assert "labels" not in prepared.model_inputs and "loss_weights" not in prepared.model_inputs + assert "position_ids" not in prepared.model_inputs + assert prepared.model_inputs["attention_mask"].tolist() == [[1, 1, 1, 0]] + torch.testing.assert_close(prepared.model_inputs["pixel_values"], pixel_values) + assert prepared.local_shift_labels.tolist() == ([[2, 3]] if rank == 0 else [[IGNORE_INDEX, IGNORE_INDEX]]) + assert prepared.local_shift_loss_weights.tolist() == ([[0.5, 2.0]] if rank == 0 else [[0.0, 0.0]]) + assert prepared.global_loss_weight_sum.item() == 2.5