[diffusion] UX: clean up startup and offload logs (#36034)

This commit is contained in:
Mick
2026-08-23 16:50:25 +08:00
committed by GitHub
parent 340391a297
commit e3a008a9db
10 changed files with 113 additions and 42 deletions
-5
View File
@@ -272,11 +272,6 @@ def _discover_and_register_pipelines():
cls.pipeline_config_cls, cls.pipeline_config_cls,
cls.sampling_params_cls, cls.sampling_params_cls,
) )
logger.debug(
f"Auto-registered config classes for pipeline '{cls.pipeline_name}': "
f"PipelineConfig={cls.pipeline_config_cls.__name__}, "
f"SamplingParams={cls.sampling_params_cls.__name__}"
)
logger.debug( logger.debug(
f"Registering pipelines complete, {len(_PIPELINE_REGISTRY)} pipelines registered" f"Registering pipelines complete, {len(_PIPELINE_REGISTRY)} pipelines registered"
) )
@@ -11,6 +11,8 @@ import torch.distributed as dist
from torch import Tensor from torch import Tensor
from torch.distributed import ProcessGroup, ReduceOp from torch.distributed import ProcessGroup, ReduceOp
from sglang.multimodal_gen.runtime.distributed.utils import all_gather_single
def _ipc_all_to_all_4d(group, input_, scatter_dim): def _ipc_all_to_all_4d(group, input_, scatter_dim):
"""2-rank IPC path for AllToAll4D; None when the transport is unavailable.""" """2-rank IPC path for AllToAll4D; None when the transport is unavailable."""
@@ -103,7 +105,7 @@ class DistributedAutograd:
output_size, dtype=input_.dtype, device=input_.device output_size, dtype=input_.dtype, device=input_.device
) )
dist.all_gather_into_tensor(output_tensor, input_, group=group) all_gather_single(output_tensor, input_, group=group)
output_tensor = output_tensor.reshape((world_size,) + input_size) output_tensor = output_tensor.reshape((world_size,) + input_size)
output_tensor = output_tensor.movedim(0, dim) output_tensor = output_tensor.movedim(0, dim)
@@ -9,6 +9,8 @@ import os
import torch import torch
from torch.distributed import ProcessGroup from torch.distributed import ProcessGroup
from sglang.multimodal_gen.runtime.distributed.utils import all_gather_single
from .base_device_communicator import DeviceCommunicatorBase from .base_device_communicator import DeviceCommunicatorBase
@@ -26,6 +28,7 @@ class CpuCommunicator(DeviceCommunicatorBase):
super().__init__(cpu_group, device, device_group, unique_name) super().__init__(cpu_group, device, device_group, unique_name)
self.dist_module = torch.distributed self.dist_module = torch.distributed
self._all_gather_single = all_gather_single
if ( if (
(current_platform.get_cpu_architecture() == CpuArchEnum.X86) (current_platform.get_cpu_architecture() == CpuArchEnum.X86)
@@ -33,6 +36,7 @@ class CpuCommunicator(DeviceCommunicatorBase):
and unique_name.startswith("tp") and unique_name.startswith("tp")
): ):
self.dist_module = _CPUSHMDistributed(self) self.dist_module = _CPUSHMDistributed(self)
self._all_gather_single = self.dist_module.all_gather_single
def all_reduce( def all_reduce(
self, self,
@@ -89,9 +93,7 @@ class CpuCommunicator(DeviceCommunicatorBase):
output_size, dtype=input_.dtype, device=input_.device output_size, dtype=input_.dtype, device=input_.device
) )
# All-gather. # All-gather.
self.dist_module.all_gather_into_tensor( self._all_gather_single(output_tensor, input_, group=self.device_group)
output_tensor, input_, group=self.device_group
)
# Reshape # Reshape
output_tensor = output_tensor.reshape((self.world_size,) + input_size) output_tensor = output_tensor.reshape((self.world_size,) + input_size)
@@ -153,7 +155,7 @@ class _CPUSHMDistributed:
torch.distributed.get_group_rank(group, dst), torch.distributed.get_group_rank(group, dst),
) )
def all_gather_into_tensor( def all_gather_single(
self, self,
output: torch.Tensor, output: torch.Tensor,
input: torch.Tensor, input: torch.Tensor,
@@ -24,6 +24,7 @@ from sglang.multimodal_gen.runtime.distributed.device_communicators.base_device_
from sglang.multimodal_gen.runtime.distributed.device_communicators.cpu_communicator import ( from sglang.multimodal_gen.runtime.distributed.device_communicators.cpu_communicator import (
CpuCommunicator, CpuCommunicator,
) )
from sglang.multimodal_gen.runtime.distributed.utils import all_gather_single
from sglang.multimodal_gen.runtime.platforms import current_platform from sglang.multimodal_gen.runtime.platforms import current_platform
from sglang.multimodal_gen.runtime.utils.logging_utils import ( from sglang.multimodal_gen.runtime.utils.logging_utils import (
init_logger, init_logger,
@@ -417,9 +418,7 @@ class GroupCoordinator:
): ):
return torch.ops.sgl_kernel.shm_allgather(input_, dim) return torch.ops.sgl_kernel.shm_allgather(input_, dim)
else: else:
torch.distributed.all_gather_into_tensor( all_gather_single(output_tensor, input_, group=self.device_group)
output_tensor, input_, group=self.device_group
)
if dim != 0: if dim != 0:
input_size[0] //= world_size input_size[0] //= world_size
@@ -18,8 +18,15 @@ from typing import Any
import torch import torch
from torch.distributed import TCPStore from torch.distributed import TCPStore
try:
from torch.distributed import all_gather_single as _all_gather_single
except ImportError:
from torch.distributed import all_gather_into_tensor as _all_gather_single
from sglang.multimodal_gen.runtime.utils.logging_utils import init_logger from sglang.multimodal_gen.runtime.utils.logging_utils import init_logger
all_gather_single = _all_gather_single
logger = init_logger(__name__) logger = init_logger(__name__)
@@ -343,16 +343,16 @@ class TransformerLoader(ComponentLoader):
) )
checkpoint_key_filter = _minimax_h3_adaln_cache_key_filter checkpoint_key_filter = _minimax_h3_adaln_cache_key_filter
if ( runtime_quant_config = init_params["quant_config"]
init_params["quant_config"] is None if runtime_quant_config is not None:
and component_server_args.transformer_weights_path is not None logger.debug(
): "Runtime quantization: %s", type(runtime_quant_config).__name__
)
elif component_server_args.transformer_weights_path is not None:
logger.info( logger.info(
"Using an unquantized transformer weight override from %s", "Using an unquantized transformer weight override from %s",
component_server_args.transformer_weights_path, component_server_args.transformer_weights_path,
) )
else:
logger.debug("quantization config: %s", init_params["quant_config"])
local_torch_device = get_local_torch_device() local_torch_device = get_local_torch_device()
checkpoint_load_device = ( checkpoint_load_device = (
@@ -257,9 +257,9 @@ def describe_host_memory() -> str:
capped = cgroup_memory_limit_bytes() capped = cgroup_memory_limit_bytes()
available = host_memory_available_bytes() available = host_memory_available_bytes()
if capped is None: if capped is None:
return f"host memory available: {available / GIB_BYTES:.1f} GiB (no cgroup cap)" return f"{available / GIB_BYTES:.1f} GiB available (no cgroup cap)"
limit, usage = capped limit, usage = capped
return ( return (
f"host memory available: {available / GIB_BYTES:.1f} GiB " f"{available / GIB_BYTES:.1f} GiB available "
f"(cgroup cap {limit / GIB_BYTES:.1f} GiB, in use {usage / GIB_BYTES:.1f} GiB)" f"(cgroup cap {limit / GIB_BYTES:.1f} GiB, in use {usage / GIB_BYTES:.1f} GiB)"
) )
@@ -4,6 +4,7 @@ import re
import threading import threading
from collections.abc import Mapping, Sequence from collections.abc import Mapping, Sequence
from contextlib import nullcontext from contextlib import nullcontext
from time import perf_counter
from typing import Any, Dict, List, Optional, Set, Tuple from typing import Any, Dict, List, Optional, Set, Tuple
import torch import torch
@@ -811,9 +812,6 @@ class LayerwiseOffloadManager:
self.register_forward_hooks() self.register_forward_hooks()
self._configured = True self._configured = True
logger.debug(
f"LayerwiseOffloadManager initialized with num prefetched layer: {self.prefetch_size}, num resident layers: {self.resident_layers}, total num layers: {self.num_layers}, residency policy: {self.residency_policy}"
)
if self.residency_policy == RESIDENCY_POLICY_STRIDED and self._streamed_order: if self.residency_policy == RESIDENCY_POLICY_STRIDED and self._streamed_order:
# Printed because the layout is the whole point of the policy, and # Printed because the layout is the whole point of the policy, and
# "did it actually stride?" is otherwise only answerable from a # "did it actually stride?" is otherwise only answerable from a
@@ -1657,7 +1655,7 @@ class LayerwiseOffloadableModuleMixin:
) )
self.layerwise_offload_managers = [] self.layerwise_offload_managers = []
named_modules = dict(self.named_modules()) named_modules = dict(self.named_modules())
configured_layer_names = [] layer_specs = []
# `--dit-*` is the group default these fall back to, not a scope. # `--dit-*` is the group default these fall back to, not a scope.
prefetch_value, resident_value, residency_policy = ( prefetch_value, resident_value, residency_policy = (
server_args.layerwise_tuning_for( server_args.layerwise_tuning_for(
@@ -1687,6 +1685,32 @@ class LayerwiseOffloadableModuleMixin:
else: else:
resident_layers = min(num_layers, int(resident_value)) resident_layers = min(num_layers, int(resident_value))
layer_specs.append((layer_name, num_layers, prefetch_size, resident_layers))
if not layer_specs:
logger.debug(
"No layerwise-offloadable ModuleList found for %s. Candidates: %s",
self.__class__.__name__,
self.layer_names,
)
return
component_label = (
f"{component_name} ({self.__class__.__name__})"
if component_name is not None
else self.__class__.__name__
)
logger.info(
"Configuring layerwise offload for %s: %s",
component_label,
", ".join(
f"{layer_name} ({num_layers} layers)"
for layer_name, num_layers, _, _ in layer_specs
),
)
started_at = perf_counter()
for layer_name, num_layers, prefetch_size, resident_layers in layer_specs:
# Pinning these weights is what lets the copy stream run ahead of # Pinning these weights is what lets the copy stream run ahead of
# compute, but pinned pages are the ones the kernel cannot reclaim, # compute, but pinned pages are the ones the kernel cannot reclaim,
# so they are handed out only while the budget lasts. The budget goes # so they are handed out only while the budget lasts. The budget goes
@@ -1709,7 +1733,6 @@ class LayerwiseOffloadableModuleMixin:
residency_policy=residency_policy, residency_policy=residency_policy,
) )
self.layerwise_offload_managers.append(manager) self.layerwise_offload_managers.append(manager)
configured_layer_names.append(layer_name)
if current_platform.is_mps(): if current_platform.is_mps():
for manager in self.layerwise_offload_managers: for manager in self.layerwise_offload_managers:
@@ -1746,17 +1769,25 @@ class LayerwiseOffloadableModuleMixin:
for manager in enabled_managers: for manager in enabled_managers:
manager._finalize_initialization() manager._finalize_initialization()
if configured_layer_names: managers = self.layerwise_offload_managers
logger.debug( prefetch_sizes = ", ".join(
"Enabled layerwise offload for %s on modules: %s", str(value)
self.__class__.__name__, for value in sorted({manager.prefetch_size for manager in managers})
configured_layer_names,
) )
else: policies = ", ".join(sorted({manager.residency_policy for manager in managers}))
logger.debug( total_layers = sum(manager.num_layers for manager in managers)
"No layerwise-offloadable ModuleList found for %s. Candidates: %s", resident_layers = sum(manager.resident_layers for manager in managers)
self.__class__.__name__, logger.info(
self.layer_names, "Layerwise offload ready for %s in %.2fs: groups=%d, layers=%d, "
"prefetch/group=%s, resident=%d/%d, policy=%s",
component_label,
perf_counter() - started_at,
len(managers),
total_layers,
prefetch_sizes,
resident_layers,
total_layers,
policies,
) )
def prepare_for_next_req(self): def prepare_for_next_req(self):
@@ -2073,7 +2104,7 @@ def configure_layerwise_offload_modules(
reverse=True, reverse=True,
) )
pin_budget = HostPinBudget() pin_budget = HostPinBudget()
logger.info("Layerwise offload: %s", describe_host_memory()) logger.info("Layerwise offload host memory: %s", describe_host_memory())
for component_name in selected_pipeline_component_names: for component_name in selected_pipeline_component_names:
module = modules[component_name] module = modules[component_name]
@@ -2108,7 +2139,7 @@ def configure_layerwise_offload_modules(
) )
logger.info( logger.info(
"Enabled layerwise offload for pipeline components: %s", "Layerwise offload summary: %s",
", ".join( ", ".join(
f"{name} ({format_component_residency(modules[name])})" f"{name} ({format_component_residency(modules[name])})"
for name in configured_component_names for name in configured_component_names
@@ -25,6 +25,7 @@ from sglang.multimodal_gen.runtime.distributed import (
get_sp_world_size, get_sp_world_size,
model_parallel_is_initialized, model_parallel_is_initialized,
) )
from sglang.multimodal_gen.runtime.distributed.utils import all_gather_single
from sglang.multimodal_gen.runtime.managers.memory_managers.layerwise_offload import ( from sglang.multimodal_gen.runtime.managers.memory_managers.layerwise_offload import (
LayerwiseOffloadableModuleMixin, LayerwiseOffloadableModuleMixin,
) )
@@ -369,7 +370,7 @@ class ParallelTiledVAE(ABC, nn.Module, LayerwiseOffloadableModuleMixin):
.repeat(world_size, *[1] * len(padded_results.shape)) .repeat(world_size, *[1] * len(padded_results.shape))
.contiguous() .contiguous()
) )
dist.all_gather_into_tensor(gathered_results, padded_results) all_gather_single(gathered_results, padded_results)
dist.all_gather_object(gathered_dim_metadata, local_dim_metadata) dist.all_gather_object(gathered_dim_metadata, local_dim_metadata)
gathered_dim_metadata = cast(list[list[torch.Size]], gathered_dim_metadata) gathered_dim_metadata = cast(list[list[torch.Size]], gathered_dim_metadata)
@@ -510,7 +511,7 @@ class ParallelTiledVAE(ABC, nn.Module, LayerwiseOffloadableModuleMixin):
device=padded_results.device, device=padded_results.device,
dtype=padded_results.dtype, dtype=padded_results.dtype,
) )
dist.all_gather_into_tensor(gathered_results, padded_results) all_gather_single(gathered_results, padded_results)
dec = z.new_empty( dec = z.new_empty(
( (
@@ -994,6 +994,40 @@ def test_configure_resolves_residency_policy(monkeypatch):
) )
def test_configure_logs_component_start_and_completion(monkeypatch):
_patch_fake_device(monkeypatch)
logs = []
timestamps = iter((10.0, 12.345))
monkeypatch.setattr(
layerwise_offload_mod.logger,
"info",
lambda message, *args: logs.append(message % args),
)
monkeypatch.setattr(
layerwise_offload_mod,
"perf_counter",
lambda: next(timestamps),
)
comp = _ResidentComponent(8)
comp.configure_layerwise_offload(
_server_args(
dit_offload_prefetch_size=2,
dit_layerwise_resident_layers=3,
),
component_name="transformer",
)
assert logs[0] == (
"Configuring layerwise offload for transformer (_ResidentComponent): "
"blocks (8 layers)"
)
assert logs[-1] == (
"Layerwise offload ready for transformer (_ResidentComponent) in 2.35s: "
"groups=1, layers=8, prefetch/group=2, resident=3/8, policy=leading"
)
def test_configure_offloads_all_layer_groups_before_moving_non_layers(monkeypatch): def test_configure_offloads_all_layer_groups_before_moving_non_layers(monkeypatch):
_patch_fake_device(monkeypatch) _patch_fake_device(monkeypatch)
model = _MultiGroupComponent() model = _MultiGroupComponent()