[diffusion] perf: support zero-cost weight offload and overlap with compute for wan-series (#15511)

This commit is contained in:
Xiaoyu Zhang
2025-12-20 22:52:40 +08:00
committed by GitHub
parent dce2ed4467
commit 8999ce754f
6 changed files with 322 additions and 6 deletions
@@ -32,6 +32,7 @@ If you only need to use the distributed environment without model parallelism,
"""
import contextlib
import os
import time
import weakref
from collections import namedtuple
from collections.abc import Callable
@@ -67,8 +68,6 @@ _DP: Optional[GroupCoordinator] = None
_DIT: Optional[GroupCoordinator] = None
_VAE: Optional[GroupCoordinator] = None
logger = init_logger(__name__)
TensorMetadata = namedtuple("TensorMetadata", ["device", "dtype", "size"])
@@ -433,6 +432,9 @@ def initialize_model_parallel(
ring_group=PROCESS_GROUP.RING_PG,
)
if ulysses_degree > 1:
_warmup_ulysses_communication()
global _TP
assert _TP is None, "Tensor parallel group is already initialized"
_TP = init_parallel_group_coordinator(
@@ -945,6 +947,49 @@ def get_ring_parallel_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
def get_pp_group() -> PipelineGroupCoordinator:
assert _PP is not None, "pipeline model parallel group is not initialized"