Skip to content

distributed

Classes

fastvideo.distributed.StatelessProcessGroup dataclass

StatelessProcessGroup(rank: int, world_size: int, store: Store, data_expiration_seconds: int = 3600, send_dst_counter: dict[int, int] = dict(), recv_src_counter: dict[int, int] = dict(), broadcast_send_counter: int = 0, broadcast_recv_src_counter: dict[int, int] = dict(), entries: deque[tuple[str, float]] = deque())

A dataclass to hold a metadata store, and the rank, world_size of the group. Only use it to communicate metadata between processes. For data-plane communication, create NCCL-related objects.

Methods:

fastvideo.distributed.StatelessProcessGroup.all_gather_obj
all_gather_obj(obj: Any) -> list[Any]

All gather an object from all ranks.

Source code in fastvideo/distributed/utils.py
def all_gather_obj(self, obj: Any) -> list[Any]:
    """All gather an object from all ranks."""
    gathered_objs = []
    for i in range(self.world_size):
        if i == self.rank:
            gathered_objs.append(obj)
            self.broadcast_obj(obj, src=self.rank)
        else:
            recv_obj = self.broadcast_obj(None, src=i)
            gathered_objs.append(recv_obj)
    return gathered_objs
fastvideo.distributed.StatelessProcessGroup.barrier
barrier()

A barrier to synchronize all ranks.

Source code in fastvideo/distributed/utils.py
def barrier(self):
    """A barrier to synchronize all ranks."""
    for i in range(self.world_size):
        if i == self.rank:
            self.broadcast_obj(None, src=self.rank)
        else:
            self.broadcast_obj(None, src=i)
fastvideo.distributed.StatelessProcessGroup.broadcast_obj
broadcast_obj(obj: Any | None, src: int) -> Any

Broadcast an object from a source rank to all other ranks. It does not clean up after all ranks have received the object. Use it for limited times, e.g., for initialization.

Source code in fastvideo/distributed/utils.py
def broadcast_obj(self, obj: Any | None, src: int) -> Any:
    """Broadcast an object from a source rank to all other ranks.
    It does not clean up after all ranks have received the object.
    Use it for limited times, e.g., for initialization.
    """
    if self.rank == src:
        self.expire_data()
        key = (f"broadcast_from/{src}/"
               f"{self.broadcast_send_counter}")
        self.store.set(key, pickle.dumps(obj))
        self.broadcast_send_counter += 1
        self.entries.append((key, time.perf_counter()))
        return obj
    else:
        key = (f"broadcast_from/{src}/"
               f"{self.broadcast_recv_src_counter[src]}")
        recv_obj = pickle.loads(self.store.get(key))
        self.broadcast_recv_src_counter[src] += 1
        return recv_obj
fastvideo.distributed.StatelessProcessGroup.create staticmethod
create(host: str, port: int, rank: int, world_size: int, data_expiration_seconds: int = 3600) -> StatelessProcessGroup

A replacement for torch.distributed.init_process_group that does not pollute the global state.

If we have process A and process B called torch.distributed.init_process_group to form a group, and then we want to form another group with process A, B, C, D, it is not possible in PyTorch, because process A and process B have already formed a group, and process C and process D cannot join that group. This function is a workaround for this issue.

torch.distributed.init_process_group is a global call, while this function is a stateless call. It will return a StatelessProcessGroup object that can be used for exchanging metadata. With this function, process A and process B can call StatelessProcessGroup.create to form a group, and then process A, B, C, and D can call StatelessProcessGroup.create to form another group.

Source code in fastvideo/distributed/utils.py
@staticmethod
def create(
    host: str,
    port: int,
    rank: int,
    world_size: int,
    data_expiration_seconds: int = 3600,
) -> "StatelessProcessGroup":
    """A replacement for `torch.distributed.init_process_group` that does not
    pollute the global state.

    If we have process A and process B called `torch.distributed.init_process_group`
    to form a group, and then we want to form another group with process A, B, C,
    D, it is not possible in PyTorch, because process A and process B have already
    formed a group, and process C and process D cannot join that group. This
    function is a workaround for this issue.

    `torch.distributed.init_process_group` is a global call, while this function
    is a stateless call. It will return a `StatelessProcessGroup` object that can be
    used for exchanging metadata. With this function, process A and process B
    can call `StatelessProcessGroup.create` to form a group, and then process A, B,
    C, and D can call `StatelessProcessGroup.create` to form another group.
    """ # noqa
    store = TCPStore(
        host_name=host,
        port=port,
        world_size=world_size,
        is_master=(rank == 0),
    )

    return StatelessProcessGroup(rank=rank,
                                 world_size=world_size,
                                 store=store,
                                 data_expiration_seconds=data_expiration_seconds)
fastvideo.distributed.StatelessProcessGroup.expire_data
expire_data() -> None

Expire data that is older than data_expiration_seconds seconds.

Source code in fastvideo/distributed/utils.py
def expire_data(self) -> None:
    """Expire data that is older than `data_expiration_seconds` seconds."""
    while self.entries:
        # check the oldest entry
        key, timestamp = self.entries[0]
        if time.perf_counter() - timestamp > self.data_expiration_seconds:
            self.store.delete_key(key)
            self.entries.popleft()
        else:
            break
fastvideo.distributed.StatelessProcessGroup.recv_obj
recv_obj(src: int) -> Any

Receive an object from a source rank.

Source code in fastvideo/distributed/utils.py
def recv_obj(self, src: int) -> Any:
    """Receive an object from a source rank."""
    obj = pickle.loads(self.store.get(f"send_to/{self.rank}/{self.recv_src_counter[src]}"))
    self.recv_src_counter[src] += 1
    return obj
fastvideo.distributed.StatelessProcessGroup.send_obj
send_obj(obj: Any, dst: int)

Send an object to a destination rank.

Source code in fastvideo/distributed/utils.py
def send_obj(self, obj: Any, dst: int):
    """Send an object to a destination rank."""
    self.expire_data()
    key = f"send_to/{dst}/{self.send_dst_counter[dst]}"
    self.store.set(key, pickle.dumps(obj))
    self.send_dst_counter[dst] += 1
    self.entries.append((key, time.perf_counter()))

Functions:

fastvideo.distributed.compute_padding_for_sp

compute_padding_for_sp(seq_len: int, sp_world_size: int) -> tuple[int, int]

Compute padding needed for sequence parallel.

Parameters:

Name Type Description Default
seq_len int

Original sequence length

required
sp_world_size int

Sequence parallel world size

required

Returns:

Name Type Description
tuple tuple[int, int]

(padded_seq_len, padding_amount)

Source code in fastvideo/distributed/utils.py
def compute_padding_for_sp(seq_len: int, sp_world_size: int) -> tuple[int, int]:
    """
    Compute padding needed for sequence parallel.

    Args:
        seq_len: Original sequence length
        sp_world_size: Sequence parallel world size

    Returns:
        tuple: (padded_seq_len, padding_amount)
    """
    if seq_len % sp_world_size == 0:
        return seq_len, 0

    padding_amount = sp_world_size - (seq_len % sp_world_size)
    padded_seq_len = seq_len + padding_amount

    return padded_seq_len, padding_amount

fastvideo.distributed.divide

divide(numerator: int, denominator: int) -> int

Ensure that numerator is divisible by the denominator and return the division value.

Source code in fastvideo/distributed/utils.py
def divide(numerator: int, denominator: int) -> int:
    """Ensure that numerator is divisible by the denominator and return
    the division value."""
    ensure_divisibility(numerator, denominator)
    return numerator // denominator

fastvideo.distributed.ensure_divisibility

ensure_divisibility(numerator, denominator) -> None

Ensure that numerator is divisible by the denominator.

Source code in fastvideo/distributed/utils.py
def ensure_divisibility(numerator, denominator) -> None:
    """Ensure that numerator is divisible by the denominator."""
    assert numerator % denominator == 0, "{} is not divisible by {}".format(numerator, denominator)

fastvideo.distributed.get_dp_rank

get_dp_rank() -> int

Return my rank for the data parallel group.

Source code in fastvideo/distributed/parallel_state.py
def get_dp_rank() -> int:
    """Return my rank for the data parallel group."""
    return get_dp_group().rank_in_group

fastvideo.distributed.get_dp_world_size

get_dp_world_size() -> int

Return world size for the data parallel group.

Source code in fastvideo/distributed/parallel_state.py
def get_dp_world_size() -> int:
    """Return world size for the data parallel group."""
    return get_dp_group().world_size

fastvideo.distributed.get_local_torch_device

get_local_torch_device() -> device

Return the torch device for the current rank.

Source code in fastvideo/distributed/parallel_state.py
def get_local_torch_device() -> torch.device:
    """Return the torch device for the current rank."""
    from fastvideo.platforms import current_platform
    if current_platform.is_npu():
        device = torch.device(f"npu:{envs.LOCAL_RANK}")
    elif current_platform.is_cuda_alike() or current_platform.is_cuda():
        device = torch.device(f"cuda:{envs.LOCAL_RANK}")
    else:
        device = torch.device("mps")
    return device

fastvideo.distributed.get_sp_parallel_rank

get_sp_parallel_rank() -> int

Return my rank for the sequence model parallel group.

Source code in fastvideo/distributed/parallel_state.py
def get_sp_parallel_rank() -> int:
    """Return my rank for the sequence model parallel group."""
    return get_sp_group().rank_in_group

fastvideo.distributed.get_sp_world_size

get_sp_world_size() -> int

Return world size for the sequence model parallel group.

Source code in fastvideo/distributed/parallel_state.py
def get_sp_world_size() -> int:
    """Return world size for the sequence model parallel group."""
    return get_sp_group().world_size

fastvideo.distributed.get_tp_rank

get_tp_rank() -> int

Return my rank for the tensor model parallel group.

Source code in fastvideo/distributed/parallel_state.py
def get_tp_rank() -> int:
    """Return my rank for the tensor model parallel group."""
    return get_tp_group().rank_in_group

fastvideo.distributed.get_tp_world_size

get_tp_world_size() -> int

Return world size for the tensor model parallel group.

Source code in fastvideo/distributed/parallel_state.py
def get_tp_world_size() -> int:
    """Return world size for the tensor model parallel group."""
    return get_tp_group().world_size

fastvideo.distributed.get_world_rank

get_world_rank() -> int

Return my rank for the world group.

Source code in fastvideo/distributed/parallel_state.py
def get_world_rank() -> int:
    """Return my rank for the world group."""
    return get_world_group().rank

fastvideo.distributed.get_world_size

get_world_size() -> int

Return world size for the world group.

Source code in fastvideo/distributed/parallel_state.py
def get_world_size() -> int:
    """Return world size for the world group."""
    return get_world_group().world_size

fastvideo.distributed.init_logger

init_logger(name: str) -> _FastvideoLogger

The main purpose of this function is to ensure that loggers are retrieved in such a way that we can be sure the root fastvideo logger has already been configured.

Source code in fastvideo/logger.py
def init_logger(name: str) -> _FastvideoLogger:
    """The main purpose of this function is to ensure that loggers are
    retrieved in such a way that we can be sure the root fastvideo logger has
    already been configured."""

    logger = logging.getLogger(name)

    methods_to_patch = {
        "info_once": _print_info_once,
        "warning_once": _print_warning_once,
        "info": _info,
    }

    for method_name, method in methods_to_patch.items():
        setattr(logger, method_name, MethodType(method, logger))  # type: ignore[arg-type]

    return cast(_FastvideoLogger, logger)

fastvideo.distributed.initialize_model_parallel

initialize_model_parallel(tensor_model_parallel_size: int = 1, sequence_model_parallel_size: int = 1, data_parallel_size: int = 1, backend: str | None = None) -> None

Initialize model parallel groups.

Parameters:

Name Type Description Default
tensor_model_parallel_size int

number of GPUs used for tensor model parallelism (used for language encoder).

1
sequence_model_parallel_size int

number of GPUs used for sequence model parallelism (used for DiT).

1
Source code in fastvideo/distributed/parallel_state.py
def initialize_model_parallel(
    tensor_model_parallel_size: int = 1,
    sequence_model_parallel_size: int = 1,
    data_parallel_size: int = 1,
    backend: str | None = None,
) -> None:
    """
    Initialize model parallel groups.

    Arguments:
        tensor_model_parallel_size: number of GPUs used for tensor model
            parallelism (used for language encoder).
        sequence_model_parallel_size: number of GPUs used for sequence model
            parallelism (used for DiT).
    """
    # Get world size and rank. Ensure some consistencies.
    assert _WORLD is not None, "world group is not initialized, please call init_distributed_environment first"
    world_size: int = get_world_size()
    backend = backend or torch.distributed.get_backend(get_world_group().device_group)
    assert world_size >= tensor_model_parallel_size, f"world_size({world_size}) must be greater than or equal to tensor_model_parallel_size({tensor_model_parallel_size})"
    num_tensor_model_parallel_groups: int = (world_size // tensor_model_parallel_size)
    global _TP
    assert _TP is None, ("tensor model parallel group is already initialized")
    group_ranks = []
    for i in range(num_tensor_model_parallel_groups):
        ranks = list(range(i * tensor_model_parallel_size, (i + 1) * tensor_model_parallel_size))
        group_ranks.append(ranks)

    # message queue broadcaster is only used in tensor model parallel group
    _TP = init_model_parallel_group(group_ranks,
                                    get_world_group().local_rank,
                                    backend,
                                    use_message_queue_broadcaster=True,
                                    group_name="tp")

    # Build the sequence model-parallel groups.
    num_sequence_model_parallel_groups: int = (world_size // sequence_model_parallel_size)
    global _SP
    assert _SP is None, ("sequence model parallel group is already initialized")
    group_ranks = []

    # Since SP is incompatible with TP and PP, we can use a simpler group creation logic
    for i in range(num_sequence_model_parallel_groups):
        # Create groups of consecutive ranks
        ranks = list(range(i * sequence_model_parallel_size, (i + 1) * sequence_model_parallel_size))
        group_ranks.append(ranks)

    _SP = init_model_parallel_group(group_ranks, get_world_group().local_rank, backend, group_name="sp")

    # Build the data parallel groups.
    num_data_parallel_groups: int = sequence_model_parallel_size
    global _DP
    assert _DP is None, ("data parallel group is already initialized")
    group_ranks = []

    for i in range(num_data_parallel_groups):
        ranks = list(range(i, world_size, num_data_parallel_groups))
        group_ranks.append(ranks)

    _DP = init_model_parallel_group(group_ranks, get_world_group().local_rank, backend, group_name="dp")

fastvideo.distributed.model_parallel_is_initialized

model_parallel_is_initialized() -> bool

Check if tensor, sequence parallel groups are initialized.

Source code in fastvideo/distributed/parallel_state.py
def model_parallel_is_initialized() -> bool:
    """Check if tensor, sequence parallel groups are initialized."""
    return _TP is not None and _SP is not None and _DP is not None

fastvideo.distributed.pad_sequence_tensor

pad_sequence_tensor(tensor: Tensor, target_seq_len: int, seq_dim: int = 1, pad_value: float = 0.0) -> Tensor

Pad a tensor along the sequence dimension.

Parameters:

Name Type Description Default
tensor Tensor

Input tensor to pad

required
target_seq_len int

Target sequence length after padding

required
seq_dim int

Dimension to pad along (default: 1)

1
pad_value float

Value to use for padding (default: 0.0)

0.0

Returns:

Name Type Description
Tensor Tensor

Padded tensor

Source code in fastvideo/distributed/utils.py
def pad_sequence_tensor(
    tensor: torch.Tensor,
    target_seq_len: int,
    seq_dim: int = 1,
    pad_value: float = 0.0,
) -> torch.Tensor:
    """
    Pad a tensor along the sequence dimension.

    Args:
        tensor: Input tensor to pad
        target_seq_len: Target sequence length after padding
        seq_dim: Dimension to pad along (default: 1)
        pad_value: Value to use for padding (default: 0.0)

    Returns:
        Tensor: Padded tensor
    """
    current_seq_len = tensor.shape[seq_dim]

    if current_seq_len >= target_seq_len:
        return tensor

    padding_amount = target_seq_len - current_seq_len

    # Create padding shape
    pad_shape = list(tensor.shape)
    pad_shape[seq_dim] = padding_amount

    # Create padding tensor
    padding = torch.full(
        pad_shape,
        pad_value,
        dtype=tensor.dtype,
        device=tensor.device,
    )

    # Concatenate along sequence dimension
    padded_tensor = torch.cat([tensor, padding], dim=seq_dim)

    return padded_tensor

fastvideo.distributed.sequence_model_parallel_all_gather

sequence_model_parallel_all_gather(input_: Tensor, dim: int = -1) -> Tensor

All-gather the input tensor across model parallel group.

Source code in fastvideo/distributed/communication_op.py
def sequence_model_parallel_all_gather(input_: torch.Tensor, dim: int = -1) -> torch.Tensor:
    """All-gather the input tensor across model parallel group."""
    return get_sp_group().all_gather(input_, dim)

fastvideo.distributed.sequence_model_parallel_all_gather_with_unpad

sequence_model_parallel_all_gather_with_unpad(input_: Tensor, original_seq_len: int, dim: int = -1) -> Tensor

All-gather the input tensor and remove padding.

Parameters:

Name Type Description Default
input_ Tensor

Sharded (and possibly padded) tensor to gather

required
original_seq_len int

Original sequence length before padding

required
dim int

Dimension to gather along (default: -1)

-1

Returns:

Name Type Description
Tensor Tensor

Gathered and unpadded tensor

Source code in fastvideo/distributed/communication_op.py
def sequence_model_parallel_all_gather_with_unpad(input_: torch.Tensor,
                                                  original_seq_len: int,
                                                  dim: int = -1) -> torch.Tensor:
    """All-gather the input tensor and remove padding.

    Args:
        input_: Sharded (and possibly padded) tensor to gather
        original_seq_len: Original sequence length before padding
        dim: Dimension to gather along (default: -1)

    Returns:
        Tensor: Gathered and unpadded tensor
    """

    # First gather across all ranks
    gathered = get_sp_group().all_gather(input_, dim)

    current_seq_len = gathered.shape[dim]
    if current_seq_len > original_seq_len:
        gathered = unpad_sequence_tensor(gathered, original_seq_len, seq_dim=dim)

    return gathered

fastvideo.distributed.sequence_model_parallel_all_to_all_4D

sequence_model_parallel_all_to_all_4D(input_: Tensor, scatter_dim: int = 2, gather_dim: int = 1) -> Tensor

All-to-all communication of 4D tensors (e.g. QKV matrices) across sequence parallel group.

Source code in fastvideo/distributed/communication_op.py
def sequence_model_parallel_all_to_all_4D(input_: torch.Tensor,
                                          scatter_dim: int = 2,
                                          gather_dim: int = 1) -> torch.Tensor:
    """All-to-all communication of 4D tensors (e.g. QKV matrices) across sequence parallel group."""
    return get_sp_group().all_to_all_4D(input_, scatter_dim, gather_dim)

fastvideo.distributed.sequence_model_parallel_shard

sequence_model_parallel_shard(input_: Tensor, dim: int = 1) -> tuple[Tensor, int]

Shard the input tensor across model parallel group with optional padding.

Parameters:

Name Type Description Default
input_ Tensor

Input tensor to shard

required
dim int

Dimension to shard along (default: 1)

1

Returns:

Name Type Description
tuple tuple[Tensor, int]

(sharded_tensor, original_seq_len) - sharded_tensor: The sharded (and possibly padded) tensor - original_seq_len: Original sequence length before padding

Source code in fastvideo/distributed/communication_op.py
def sequence_model_parallel_shard(input_: torch.Tensor, dim: int = 1) -> tuple[torch.Tensor, int]:
    """Shard the input tensor across model parallel group with optional padding.

    Args:
        input_: Input tensor to shard
        dim: Dimension to shard along (default: 1)

    Returns:
        tuple: (sharded_tensor, original_seq_len)
            - sharded_tensor: The sharded (and possibly padded) tensor
            - original_seq_len: Original sequence length before padding
    """

    sp_world_size = get_sp_world_size()

    original_seq_len = input_.shape[dim]

    # Compute padding if needed
    padded_seq_len, padding_amount = compute_padding_for_sp(original_seq_len, sp_world_size)

    # Pad if necessary
    if padding_amount > 0:
        input_ = pad_sequence_tensor(input_, padded_seq_len, seq_dim=dim)

    # Sharding with autograd-aware backward (all-gather in backward).
    input_ = get_sp_group().shard(input_, dim=dim, scale_grad=True)

    return input_, original_seq_len

fastvideo.distributed.split_tensor_along_last_dim

split_tensor_along_last_dim(tensor: Tensor, num_partitions: int, contiguous_split_chunks: bool = False) -> Sequence[Tensor]

Split a tensor along its last dimension.

Parameters:

Name Type Description Default
tensor Tensor

input tensor.

required
num_partitions int

number of partitions to split the tensor

required
contiguous_split_chunks bool

If True, make each chunk contiguous in memory.

False

Returns:

Type Description
Sequence[Tensor]

A list of Tensors

Source code in fastvideo/distributed/utils.py
def split_tensor_along_last_dim(
    tensor: torch.Tensor,
    num_partitions: int,
    contiguous_split_chunks: bool = False,
) -> Sequence[torch.Tensor]:
    """ Split a tensor along its last dimension.

        Arguments:
            tensor: input tensor.
            num_partitions: number of partitions to split the tensor
            contiguous_split_chunks: If True, make each chunk contiguous
                                     in memory.

        Returns:
            A list of Tensors
    """
    # Get the size and dimension.
    last_dim = tensor.dim() - 1
    last_dim_size = divide(tensor.size()[last_dim], num_partitions)
    # Split.
    tensor_list = torch.split(tensor, last_dim_size, dim=last_dim)
    # NOTE: torch.split does not create contiguous tensors by default.
    if contiguous_split_chunks:
        return tuple(chunk.contiguous() for chunk in tensor_list)

    return tuple(tensor_list)

fastvideo.distributed.tensor_model_parallel_all_gather

tensor_model_parallel_all_gather(input_: Tensor, dim: int = -1) -> Tensor

All-gather the input tensor across model parallel group.

Source code in fastvideo/distributed/communication_op.py
def tensor_model_parallel_all_gather(input_: torch.Tensor, dim: int = -1) -> torch.Tensor:
    """All-gather the input tensor across model parallel group."""
    return get_tp_group().all_gather(input_, dim)

fastvideo.distributed.tensor_model_parallel_all_reduce

tensor_model_parallel_all_reduce(input_: Tensor) -> Tensor

All-reduce the input tensor across model parallel group.

Source code in fastvideo/distributed/communication_op.py
def tensor_model_parallel_all_reduce(input_: torch.Tensor) -> torch.Tensor:
    """All-reduce the input tensor across model parallel group."""
    return get_tp_group().all_reduce(input_)

fastvideo.distributed.unpad_sequence_tensor

unpad_sequence_tensor(tensor: Tensor, original_seq_len: int, seq_dim: int = 1) -> Tensor

Remove padding from a tensor along the sequence dimension.

Parameters:

Name Type Description Default
tensor Tensor

Padded tensor

required
original_seq_len int

Original sequence length (before padding)

required
seq_dim int

Dimension to unpad along (default: 1)

1

Returns:

Name Type Description
Tensor Tensor

Unpadded tensor

Source code in fastvideo/distributed/utils.py
def unpad_sequence_tensor(
    tensor: torch.Tensor,
    original_seq_len: int,
    seq_dim: int = 1,
) -> torch.Tensor:
    """
    Remove padding from a tensor along the sequence dimension.

    Args:
        tensor: Padded tensor
        original_seq_len: Original sequence length (before padding)
        seq_dim: Dimension to unpad along (default: 1)

    Returns:
        Tensor: Unpadded tensor
    """
    # Use slice to remove padding
    indices = [slice(None)] * tensor.dim()
    indices[seq_dim] = slice(0, original_seq_len)

    return tensor[tuple(indices)]

fastvideo.distributed.warmup_sequence_parallel_communication

warmup_sequence_parallel_communication(device: device | None = None) -> None

Warmup NCCL communicators for sequence parallel all-to-all operations.

The first NCCL collective operation is slow due to lazy communicator initialization. This function runs dummy all-to-all operations to trigger the initialization upfront, before the first real forward pass.

Parameters:

Name Type Description Default
device device | None

Device to use for warmup tensors. If None, uses CUDA device 0.

None
Source code in fastvideo/distributed/communication_op.py
def warmup_sequence_parallel_communication(device: torch.device | None = None) -> None:
    """Warmup NCCL communicators for sequence parallel all-to-all operations.

    The first NCCL collective operation is slow due to lazy communicator
    initialization. This function runs dummy all-to-all operations to
    trigger the initialization upfront, before the first real forward pass.

    Args:
        device: Device to use for warmup tensors. If None, uses CUDA device 0.
    """
    global _sp_warmup_done

    if _sp_warmup_done:
        return

    if not model_parallel_is_initialized():
        return

    sp_world_size = get_sp_world_size()
    if sp_world_size <= 1:
        _sp_warmup_done = True
        return

    if device is None:
        device = torch.device("cuda")

    logger.info("Warming up sequence parallel communication (SP=%d)...", sp_world_size)

    # Use small but representative tensor shapes for warmup
    # Shape: [batch, seq_len, num_heads, head_dim]
    # The all-to-all patterns used in attention:
    #   1. scatter_dim=2 (heads), gather_dim=1 (seq) - before attention
    #   2. scatter_dim=1 (seq), gather_dim=2 (heads) - after attention
    batch_size = 1
    seq_len_per_rank = 16  # Small sequence per rank
    num_heads = sp_world_size * 4  # Must be divisible by sp_world_size
    head_dim = 64

    # Create dummy tensor for warmup
    dummy = torch.zeros(batch_size, seq_len_per_rank, num_heads, head_dim, device=device, dtype=torch.bfloat16)

    # Warmup pattern 1: scatter heads, gather sequence (before attention)
    _ = sequence_model_parallel_all_to_all_4D(dummy, scatter_dim=2, gather_dim=1)

    # Warmup pattern 2: scatter sequence, gather heads (after attention)
    dummy2 = torch.zeros(batch_size,
                         seq_len_per_rank * sp_world_size,
                         num_heads // sp_world_size,
                         head_dim,
                         device=device,
                         dtype=torch.bfloat16)
    _ = sequence_model_parallel_all_to_all_4D(dummy2, scatter_dim=1, gather_dim=2)

    # Warmup all-gather (used for replicated tokens)
    dummy3 = torch.zeros(batch_size, 8, num_heads // sp_world_size, head_dim, device=device, dtype=torch.bfloat16)
    _ = sequence_model_parallel_all_gather(dummy3, dim=2)

    # Synchronize to ensure warmup completes
    torch.cuda.synchronize(device)

    # Clean up
    del dummy, dummy2, dummy3

    _sp_warmup_done = True
    logger.info("Sequence parallel communication warmup complete.")