multimodalart's picture
multimodalart HF Staff
Bernini-Diffusers-v2 r2v demo
fed6c68 verified
Raw
History Blame Contribute Delete
4.58 kB
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