from typing import Optional import torch import torch.distributed as dist from torch import Tensor from torch.distributed import ProcessGroup from ...distributed.parallel_state import get_parallel_state from .comm import get_ulysses_sequence_parallel_group, get_unified_sequence_parallel_group from .ulysses import _Gather, _Slice from .utils import pad_tensor, unpadding_tensor_for_seqeunce_parallel def sp_pad_and_slice( tensor: torch.Tensor, dim: int = -1, pad_value: int = 0, pad_scale: int = 1, ) -> torch.Tensor: """ Pads and slices a tensor for sequence parallelism (SP) distribution. This function ensures the tensor can be evenly distributed across SP ranks by: 1. Padding the tensor to make its length divisible by (sp_size * pad_scale) 2. Slicing the padded tensor to extract the chunk for the current SP rank Args: tensor: Input tensor to pad and slice dim: Dimension along which to pad and slice (default: -1) pad_value: Value to use for padding (default: 0) pad_scale: Scaling factor for SP size during padding (default: 1). This is needed for some VLMs that perform token merging to ensure padding is handled correctly before the merge operation Returns: The sliced tensor chunk for the current SP rank """ # Get sequence parallelism configuration sp_size = get_parallel_state().sp_size sp_rank = get_parallel_state().sp_rank # Phase 1: Pad the tensor to align with (sp_size * pad_scale) # This ensures the tensor can be evenly split across all SP ranks seq_length = tensor.size(dim) scale_sp_size = sp_size * pad_scale # Calculate the chunk size after scaling, rounding up to ensure full coverage sp_chunk_size = (seq_length + scale_sp_size - 1) // scale_sp_size # Calculate how much padding is needed to reach the target length pad_size = sp_chunk_size * scale_sp_size - seq_length if pad_size != 0: # Create padding tensor with the same shape except for the target dimension pad_shape = list(tensor.shape) pad_shape[dim] = pad_size pad = torch.full(pad_shape, fill_value=pad_value, dtype=tensor.dtype, device=tensor.device) # Concatenate padding to the end of the tensor tensor = torch.cat((tensor, pad), dim=dim) # Phase 2: Slice the padded tensor for the current SP rank # After padding, recalculate the chunk size based on the actual sp_size seq_length = tensor.size(dim) sp_chunk_size = (seq_length + sp_size - 1) // sp_size # Extract the chunk for this rank: each rank gets a contiguous slice # narrow(dim, start, length) extracts tensor[start:start+length] along dim return tensor.narrow(dim, sp_rank * sp_chunk_size, sp_chunk_size) def slice_input_tensor( x: Tensor, dim: int, padding: bool = True, padding_value: int = 0, group: ProcessGroup = None, ) -> Tensor: """ A func to slice the input sequence in sequence parallel """ group = get_unified_sequence_parallel_group() if group is None else group if not group: return x sp_rank = dist.get_rank(group) sp_world = dist.get_world_size(group) dim_size = x.shape[dim] unit = (dim_size + sp_world - 1) // sp_world if padding and dim_size % sp_world: padding_size = sp_world - (dim_size % sp_world) x = pad_tensor(x, dim, padding_size, padding_value) slc = [slice(None)] * len(x.shape) slc[dim] = slice(unit * sp_rank, unit * (sp_rank + 1)) return x[tuple(slc)].contiguous() def slice_input_tensor_scale_grad( x: Tensor, dim: int, group: ProcessGroup = None, scale_grad=True, ): """ A func to gather the outputs for the model result in sequence parallel """ group = get_ulysses_sequence_parallel_group() if group is None else group if not group: return x x = _Slice.apply(group, x, dim, scale_grad) return x def gather_outputs( x: Tensor, gather_dim: int, padding_dim: Optional[int] = None, unpad_dim_size: Optional[int] = None, scale_grad=False, group: ProcessGroup = None, ): """ A func to gather the outputs for the model result in sequence parallel """ group = get_unified_sequence_parallel_group() if group is None else group if not group: return x x = _Gather.apply(group, x, gather_dim, scale_grad) if padding_dim is not None: x = unpadding_tensor_for_seqeunce_parallel(x, padding_dim, unpad_dim_size, group) return x