[diffusion] chore: remove redundant ulysses nccl warmup (#16301)

This commit is contained in:
Xiaoyu Zhang
2026-01-03 00:19:35 +08:00
committed by GitHub
parent 0d244116d2
commit 1cfd2b2ded
@@ -32,7 +32,6 @@ If you only need to use the distributed environment without model parallelism,
""" """
import contextlib import contextlib
import os import os
import time
import weakref import weakref
from collections import namedtuple from collections import namedtuple
from collections.abc import Callable from collections.abc import Callable
@@ -434,9 +433,6 @@ def initialize_model_parallel(
ring_group=PROCESS_GROUP.RING_PG, ring_group=PROCESS_GROUP.RING_PG,
) )
if ulysses_degree > 1:
_warmup_ulysses_communication()
global _TP global _TP
assert _TP is None, "Tensor parallel group is already initialized" assert _TP is None, "Tensor parallel group is already initialized"
_TP = init_parallel_group_coordinator( _TP = init_parallel_group_coordinator(
@@ -949,49 +945,6 @@ def get_ring_parallel_rank():
return get_sp_group().ring_rank return get_sp_group().ring_rank
def _warmup_ulysses_communication():
"""
Warmup NCCL communication for Ulysses all-to-all to avoid first-step latency.
This function performs a dummy all-to-all operation to initialize NCCL communication
channels, which can take several seconds on the first call.
"""
logger.info("Warming up Ulysses all-to-all communication...")
try:
import torch.distributed._functional_collectives as ft_c
ulysses_pg = get_sp_group().ulysses_group
if ulysses_pg is None:
logger.warning("Ulysses group not initialized, skipping warmup")
return
warmup_start = time.time()
device = torch.device(f"cuda:{get_world_group().local_rank}")
dummy_tensor = torch.zeros(1024, device=device, dtype=torch.float32)
output = ft_c.all_to_all_single(
dummy_tensor,
output_split_sizes=None,
input_split_sizes=None,
group=ulysses_pg,
)
if isinstance(output, ft_c.AsyncCollectiveTensor):
output = output.wait()
torch.cuda.synchronize()
warmup_time = (time.time() - warmup_start) * 1000
logger.info(f"Ulysses communication warmup completed in {warmup_time:.2f}ms")
except Exception as e:
logger.warning(
f"Ulysses communication warmup failed: {e}. Continuing without warmup."
)
# PP # PP
def get_pp_group() -> PipelineGroupCoordinator: def get_pp_group() -> PipelineGroupCoordinator:
assert _PP is not None, "pipeline model parallel group is not initialized" assert _PP is not None, "pipeline model parallel group is not initialized"