[diffusion] UX: clean up startup and offload logs (#36034)
This commit is contained in:
@@ -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"
|
||||||
)
|
)
|
||||||
|
|||||||
+3
-1
@@ -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)
|
||||||
|
|||||||
+6
-4
@@ -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)"
|
||||||
)
|
)
|
||||||
|
|||||||
+48
-17
@@ -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()
|
||||||
|
|||||||
Reference in New Issue
Block a user