From 1ac3e33622cd3dd19a068abb6677a944368519e7 Mon Sep 17 00:00:00 2001 From: Mick Date: Thu, 21 May 2026 13:20:24 +0800 Subject: [PATCH] [diffusion] optimize: enable inference mode in pipeline executor (#25891) --- .../runtime/layers/lora/linear.py | 32 +++- .../memory_managers/layerwise_offload.py | 19 +- .../executors/parallel_executor.py | 16 +- .../executors/pipeline_executor.py | 66 ++++++- .../pipelines_core/executors/sync_executor.py | 4 +- .../multimodal_gen/runtime/platforms/npu.py | 5 + .../test/unit/test_layerwise_offload.py | 24 +++ .../test/unit/test_lora_inference_mode.py | 38 ++++ .../test/unit/test_pipeline_executor.py | 171 ++++++++++++++++++ 9 files changed, 355 insertions(+), 20 deletions(-) create mode 100644 python/sglang/multimodal_gen/test/unit/test_lora_inference_mode.py create mode 100644 python/sglang/multimodal_gen/test/unit/test_pipeline_executor.py diff --git a/python/sglang/multimodal_gen/runtime/layers/lora/linear.py b/python/sglang/multimodal_gen/runtime/layers/lora/linear.py index 8ca749fd7..c0c553566 100644 --- a/python/sglang/multimodal_gen/runtime/layers/lora/linear.py +++ b/python/sglang/multimodal_gen/runtime/layers/lora/linear.py @@ -118,6 +118,14 @@ class BaseLayerWithLoRA(nn.Module): def slice_lora_b_weights(self, B: torch.Tensor) -> torch.Tensor: return B + @staticmethod + def _as_mutable_tensor(tensor: torch.Tensor) -> torch.Tensor: + # lora can be reconfigured after executor forwards create inference tensors + if tensor.is_inference(): + with torch.inference_mode(False): + return tensor.detach().clone() + return tensor + def set_lora_weights( self, A: torch.Tensor, @@ -291,6 +299,7 @@ class BaseLayerWithLoRA(nn.Module): data = self.base_layer.weight.data.to( get_local_torch_device() ).full_tensor() + data = self._as_mutable_tensor(data) target_dtype = data.dtype if ( merge_in_fp32 @@ -302,14 +311,17 @@ class BaseLayerWithLoRA(nn.Module): self._merge_lora_into_data(data, lora_list) unsharded_base_layer.weight = nn.Parameter( - data.to(current_device, dtype=target_dtype) + self._as_mutable_tensor(data.to(current_device, dtype=target_dtype)) ) if isinstance(getattr(self.base_layer, "bias", None), DTensor): - unsharded_base_layer.bias = nn.Parameter( + bias_data = ( self.base_layer.bias.to(get_local_torch_device(), non_blocking=True) .full_tensor() .to(current_device) ) + unsharded_base_layer.bias = nn.Parameter( + self._as_mutable_tensor(bias_data) + ) offload_policy = ( CPUOffloadPolicy() if "cpu" in str(current_device) else OffloadPolicy() @@ -325,6 +337,7 @@ class BaseLayerWithLoRA(nn.Module): else: current_device = self.base_layer.weight.data.device data = self.base_layer.weight.data.to(get_local_torch_device()) + data = self._as_mutable_tensor(data) target_dtype = data.dtype if ( merge_in_fp32 @@ -335,8 +348,8 @@ class BaseLayerWithLoRA(nn.Module): self._merge_lora_into_data(data, lora_list) - self.base_layer.weight.data = data.to( - current_device, dtype=target_dtype, non_blocking=True + self.base_layer.weight.data = self._as_mutable_tensor( + data.to(current_device, dtype=target_dtype, non_blocking=True) ) self.merged = True @@ -356,13 +369,20 @@ class BaseLayerWithLoRA(nn.Module): if isinstance(self.base_layer.weight, DTensor): device = self.base_layer.weight.data.device old_weight = self.base_layer.weight - new_weight_data = self.cpu_weight.to(device, non_blocking=True) + new_weight_data = self._as_mutable_tensor( + self.cpu_weight.to(device, non_blocking=True) + ) self.base_layer.weight = nn.Parameter(new_weight_data) del old_weight else: current_device = self.base_layer.weight.data.device cpu_weight_on_device = self.cpu_weight.to(current_device, non_blocking=True) - self.base_layer.weight.data.copy_(cpu_weight_on_device) + if self.base_layer.weight.data.is_inference(): + self.base_layer.weight.data = self._as_mutable_tensor( + cpu_weight_on_device + ) + else: + self.base_layer.weight.data.copy_(cpu_weight_on_device) if ( cpu_weight_on_device.data_ptr() != self.base_layer.weight.data.data_ptr() diff --git a/python/sglang/multimodal_gen/runtime/managers/memory_managers/layerwise_offload.py b/python/sglang/multimodal_gen/runtime/managers/memory_managers/layerwise_offload.py index 99db87030..95b66392c 100644 --- a/python/sglang/multimodal_gen/runtime/managers/memory_managers/layerwise_offload.py +++ b/python/sglang/multimodal_gen/runtime/managers/memory_managers/layerwise_offload.py @@ -284,7 +284,11 @@ class LayerwiseOffloadManager: # create gpu buffer and load from CPU buffer gpu_buffers: Dict[torch.dtype, torch.Tensor] = {} - with torch.get_device_module().stream(self.copy_stream): + with ( + torch.inference_mode(False), + torch.no_grad(), + torch.get_device_module().stream(self.copy_stream), + ): for dtype, cpu_buffer in self._consolidated_cpu_weights[layer_idx].items(): gpu_buffer = torch.empty( cpu_buffer.shape, dtype=dtype, device=self.device @@ -343,12 +347,13 @@ class LayerwiseOffloadManager: if layer_idx not in self._gpu_layers: return - for name, meta in self._weight_metadata.get(layer_idx, {}).items(): - target = self.get_target_with_name(name) - # Wraparound prefetch will reload the layer when it is needed again - target.data = self._get_shared_empty_tensor_for_target( - target, meta["dtype"] - ) + with torch.inference_mode(False), torch.no_grad(): + for name, meta in self._weight_metadata.get(layer_idx, {}).items(): + target = self.get_target_with_name(name) + # Wraparound prefetch will reload the layer when it is needed again + target.data = self._get_shared_empty_tensor_for_target( + target, meta["dtype"] + ) self._gpu_layers.discard(layer_idx) diff --git a/python/sglang/multimodal_gen/runtime/pipelines_core/executors/parallel_executor.py b/python/sglang/multimodal_gen/runtime/pipelines_core/executors/parallel_executor.py index 997981685..acaac0f68 100644 --- a/python/sglang/multimodal_gen/runtime/pipelines_core/executors/parallel_executor.py +++ b/python/sglang/multimodal_gen/runtime/pipelines_core/executors/parallel_executor.py @@ -78,7 +78,9 @@ class ParallelExecutor(PipelineExecutor): if rank == 0: # Only main rank executes, others just wait self.before_stage(stage, stage_index, batch, server_args) - batch = stage(batch, server_args) + batch = self.run_stage_with_context( + stage, batch, server_args, run_stage + ) self.after_stage(stage_index) torch.distributed.barrier() @@ -94,20 +96,26 @@ class ParallelExecutor(PipelineExecutor): if rank != 0: batch = broadcasted_list[0] self.before_stage(stage, stage_index, batch, server_args) - batch = stage(batch, server_args) + batch = self.run_stage_with_context( + stage, batch, server_args, run_stage + ) self.after_stage(stage_index) torch.distributed.barrier() elif paradigm == StageParallelismType.REPLICATED: self.before_stage(stage, stage_index, batch, server_args) - batch = stage(batch, server_args) + batch = self.run_stage_with_context( + stage, batch, server_args, run_stage + ) self.after_stage(stage_index) elif paradigm == StageParallelismType.MAIN_RANK_ONLY_AND_SEND_TO_OTHERS: if rank == 0: # Only main rank executes, others just wait self.before_stage(stage, stage_index, batch, server_args) - batch = stage(batch, server_args) + batch = self.run_stage_with_context( + stage, batch, server_args, run_stage + ) self.after_stage(stage_index) torch.distributed.barrier() diff --git a/python/sglang/multimodal_gen/runtime/pipelines_core/executors/pipeline_executor.py b/python/sglang/multimodal_gen/runtime/pipelines_core/executors/pipeline_executor.py index 2d0ee385d..e4960df97 100644 --- a/python/sglang/multimodal_gen/runtime/pipelines_core/executors/pipeline_executor.py +++ b/python/sglang/multimodal_gen/runtime/pipelines_core/executors/pipeline_executor.py @@ -9,8 +9,11 @@ import contextlib from abc import ABC, abstractmethod from typing import TYPE_CHECKING, List +import torch + from sglang.multimodal_gen.runtime.distributed import get_world_rank from sglang.multimodal_gen.runtime.pipelines_core.schedule_batch import OutputBatch, Req +from sglang.multimodal_gen.runtime.platforms import current_platform from sglang.multimodal_gen.runtime.server_args import ServerArgs from sglang.multimodal_gen.runtime.utils.logging_utils import init_logger from sglang.multimodal_gen.runtime.utils.perf_logger import StageProfiler @@ -81,7 +84,8 @@ class PipelineExecutor(ABC): ) -> OutputBatch: with self.profile_execution(batch, dump_rank=0): - batch = self.execute(stages, batch, server_args) + with current_platform.inference_mode(): + batch = self.execute(stages, batch, server_args) return batch @@ -93,9 +97,67 @@ class PipelineExecutor(ABC): ): """Execute a grouped request under the same profiler as a single request.""" with self.profile_execution(batches[0], dump_rank=0): - batches = self.execute_group(stages, batches, server_args) + with current_platform.inference_mode(): + batches = self.execute_group(stages, batches, server_args) return batches + @staticmethod + @contextlib.contextmanager + def _stage_execution_context(stage: "PipelineStage", server_args: ServerArgs): + if PipelineExecutor._stage_needs_version_counters(stage, server_args): + # fsdp and cpu-offload hooks need tensor version counters + with torch.inference_mode(False), torch.no_grad(): + yield + return + yield + + @staticmethod + def _stage_needs_version_counters( + stage: "PipelineStage", server_args: ServerArgs + ) -> bool: + if server_args.use_fsdp_inference: + return True + + stage_name = stage._active_component_stage_name() + for use in stage.component_uses(server_args, stage_name): + component_name = use.component_name + if server_args.dit_cpu_offload and component_name in ( + "transformer", + "transformer_2", + "video_dit", + "audio_dit", + ): + return True + if server_args.text_encoder_cpu_offload and component_name.startswith( + "text_encoder" + ): + return True + if server_args.image_encoder_cpu_offload and component_name in ( + "image_encoder", + "condition_image_encoder", + ): + return True + if server_args.vae_cpu_offload and component_name in ( + "vae", + "video_vae", + "audio_vae", + "vocoder", + "spatial_upsampler", + "condition_image_encoder", + ): + return True + return False + + def run_stage_with_context( + self, + stage: "PipelineStage", + payload, + server_args: ServerArgs, + run_stage, + ): + with self._stage_execution_context(stage, server_args): + return run_stage(stage, payload) + @abstractmethod def execute( self, diff --git a/python/sglang/multimodal_gen/runtime/pipelines_core/executors/sync_executor.py b/python/sglang/multimodal_gen/runtime/pipelines_core/executors/sync_executor.py index d1ed6fcf1..d19036080 100644 --- a/python/sglang/multimodal_gen/runtime/pipelines_core/executors/sync_executor.py +++ b/python/sglang/multimodal_gen/runtime/pipelines_core/executors/sync_executor.py @@ -33,7 +33,9 @@ class SyncExecutor(PipelineExecutor): try: for stage_index, stage in enumerate(stages): self.before_stage(stage, stage_index, payload, server_args) - payload = run_stage(stage, payload) + payload = self.run_stage_with_context( + stage, payload, server_args, run_stage + ) self.after_stage(stage_index) profiler = SGLDiffusionProfiler.get_instance() if profiler: diff --git a/python/sglang/multimodal_gen/runtime/platforms/npu.py b/python/sglang/multimodal_gen/runtime/platforms/npu.py index 710bfa651..322567aee 100644 --- a/python/sglang/multimodal_gen/runtime/platforms/npu.py +++ b/python/sglang/multimodal_gen/runtime/platforms/npu.py @@ -68,6 +68,11 @@ class NPUPlatformBase(Platform): return False return True + @classmethod + def inference_mode(cls): + # npu kernels in diffusion paths may need tensor version counters + return torch.no_grad() + @classmethod def is_full_nvlink(cls, physical_device_ids: list[int]) -> bool: logger.exception( diff --git a/python/sglang/multimodal_gen/test/unit/test_layerwise_offload.py b/python/sglang/multimodal_gen/test/unit/test_layerwise_offload.py index 6c11501af..2401c1006 100644 --- a/python/sglang/multimodal_gen/test/unit/test_layerwise_offload.py +++ b/python/sglang/multimodal_gen/test/unit/test_layerwise_offload.py @@ -204,6 +204,30 @@ def test_layerwise_offload_preserves_non_contiguous_stride(monkeypatch): assert torch.equal(reloaded_weight, original_weight) +def test_layerwise_offload_uses_normal_tensors_under_inference_mode(monkeypatch): + monkeypatch.setattr( + layerwise_offload_mod.torch, "get_device_module", lambda: _FakeDeviceModule + ) + monkeypatch.setattr(layerwise_offload_mod.current_platform, "device_type", "cpu") + + model = _DummyModel() + manager = LayerwiseOffloadManager( + model=model, + layers_attr_str="blocks", + num_layers=1, + enabled=True, + pin_cpu_memory=False, + prefetch_size=1, + ) + + with torch.inference_mode(): + manager.release_layer(0) + manager.prefetch_layer(0, non_blocking=False) + + assert model.blocks[0].weight._version >= 0 + assert model.blocks[0].bias._version >= 0 + + def test_layerwise_offload_keeps_shared_buffers_resident(monkeypatch): monkeypatch.setattr( layerwise_offload_mod.torch, "get_device_module", lambda: _FakeDeviceModule diff --git a/python/sglang/multimodal_gen/test/unit/test_lora_inference_mode.py b/python/sglang/multimodal_gen/test/unit/test_lora_inference_mode.py new file mode 100644 index 000000000..bac541163 --- /dev/null +++ b/python/sglang/multimodal_gen/test/unit/test_lora_inference_mode.py @@ -0,0 +1,38 @@ +import torch +from torch import nn + +from sglang.multimodal_gen.runtime.layers.lora.linear import LinearWithLoRA + + +def test_lora_merge_unmerge_handles_inference_base_weight(): + with torch.inference_mode(): + base_layer = nn.Linear(4, 3, bias=False) + + layer = LinearWithLoRA(base_layer, lora_rank=2, lora_alpha=2) + base_weight = layer.cpu_weight.clone() + + assert layer.base_layer.weight.is_inference() + assert not base_weight.is_inference() + + lora_a = torch.ones(2, 4) + lora_b = torch.full((3, 2), 0.5) + expected_merged = base_weight + lora_b @ lora_a + + with torch.inference_mode(False): + layer.set_lora_weights( + lora_a, + lora_b, + clear_existing=True, + merge_weights=True, + ) + + assert layer.merged + assert not layer.base_layer.weight.is_inference() + assert torch.allclose(layer.base_layer.weight, expected_merged) + + with torch.inference_mode(False): + layer.unmerge_lora_weights() + + assert not layer.merged + assert not layer.base_layer.weight.is_inference() + assert torch.allclose(layer.base_layer.weight, base_weight) diff --git a/python/sglang/multimodal_gen/test/unit/test_pipeline_executor.py b/python/sglang/multimodal_gen/test/unit/test_pipeline_executor.py new file mode 100644 index 000000000..c2d64f0bc --- /dev/null +++ b/python/sglang/multimodal_gen/test/unit/test_pipeline_executor.py @@ -0,0 +1,171 @@ +import contextlib +from types import SimpleNamespace + +import pytest +import torch + +from sglang.multimodal_gen.runtime.pipelines_core.executors import pipeline_executor +from sglang.multimodal_gen.runtime.pipelines_core.executors.pipeline_executor import ( + PipelineExecutor, +) +from sglang.multimodal_gen.runtime.platforms.npu import NPUPlatformBase + + +class _RecordingExecutor(PipelineExecutor): + def __init__(self): + super().__init__(server_args=SimpleNamespace()) + self.single_inference_mode = None + self.group_inference_mode = None + self.single_grad_enabled = None + self.group_grad_enabled = None + + def execute(self, stages, batch, server_args): + self.single_inference_mode = torch.is_inference_mode_enabled() + self.single_grad_enabled = torch.is_grad_enabled() + return batch + + def execute_group(self, stages, batches, server_args): + self.group_inference_mode = torch.is_inference_mode_enabled() + self.group_grad_enabled = torch.is_grad_enabled() + return batches + + +def _batch(): + return SimpleNamespace(profile=False, is_warmup=False) + + +def _server_args(**overrides): + values = { + "use_fsdp_inference": False, + "dit_cpu_offload": False, + "text_encoder_cpu_offload": False, + "image_encoder_cpu_offload": False, + "vae_cpu_offload": False, + "dit_layerwise_offload": False, + "layerwise_offload_components": (), + } + values.update(overrides) + return SimpleNamespace(**values) + + +class _NoGradPlatform: + @classmethod + @contextlib.contextmanager + def inference_mode(cls): + with torch.no_grad(): + yield + + +class _InferenceTensorPlatform: + @classmethod + def inference_mode(cls): + return torch.inference_mode(mode=True) + + +class _ComponentStage: + def __init__(self, *component_names): + self.component_names = component_names + + @staticmethod + def _active_component_stage_name(): + return "FakeStage" + + def component_uses(self, server_args, stage_name=None): + return [ + SimpleNamespace(component_name=component_name) + for component_name in self.component_names + ] + + +def test_execute_with_profiling_uses_inference_tensor_platform(monkeypatch): + monkeypatch.setattr(pipeline_executor, "current_platform", _InferenceTensorPlatform) + executor = _RecordingExecutor() + + with torch.inference_mode(False): + executor.execute_with_profiling([], _batch(), _server_args()) + + assert executor.single_inference_mode is True + assert executor.single_grad_enabled is False + + +def test_execute_with_profiling_uses_platform_inference_mode(monkeypatch): + monkeypatch.setattr(pipeline_executor, "current_platform", _NoGradPlatform) + executor = _RecordingExecutor() + + with torch.inference_mode(False): + executor.execute_with_profiling([], _batch(), _server_args()) + + assert executor.single_inference_mode is False + assert executor.single_grad_enabled is False + + +def test_execute_group_with_profiling_uses_platform_inference_mode(monkeypatch): + monkeypatch.setattr(pipeline_executor, "current_platform", _NoGradPlatform) + executor = _RecordingExecutor() + + with torch.inference_mode(False): + executor.execute_group_with_profiling([], [_batch(), _batch()], _server_args()) + + assert executor.group_inference_mode is False + assert executor.group_grad_enabled is False + + +@pytest.mark.parametrize( + ("server_args", "component_names"), + [ + (_server_args(use_fsdp_inference=True), ("transformer",)), + (_server_args(dit_cpu_offload=True), ("transformer",)), + (_server_args(text_encoder_cpu_offload=True), ("text_encoder",)), + (_server_args(image_encoder_cpu_offload=True), ("image_encoder",)), + (_server_args(vae_cpu_offload=True), ("vae",)), + ], +) +def test_stage_context_preserves_version_counters_when_needed( + server_args, component_names +): + stage = _ComponentStage(*component_names) + + with torch.inference_mode(): + with PipelineExecutor._stage_execution_context(stage, server_args): + tensor = torch.ones(1) + + assert torch.is_inference_mode_enabled() is False + assert torch.is_grad_enabled() is False + assert tensor._version == 0 + + +@pytest.mark.parametrize( + ("server_args", "component_names"), + [ + (_server_args(dit_layerwise_offload=True), ("transformer",)), + ( + _server_args( + text_encoder_cpu_offload=True, + layerwise_offload_components=("transformer",), + ), + ("transformer",), + ), + ( + _server_args(layerwise_offload_components=("text_encoder",)), + ("text_encoder",), + ), + ], +) +def test_stage_context_allows_layerwise_inference_tensor_mode( + server_args, component_names +): + stage = _ComponentStage(*component_names) + + with torch.inference_mode(): + with PipelineExecutor._stage_execution_context(stage, server_args): + assert torch.is_inference_mode_enabled() is True + + assert torch.is_inference_mode_enabled() is False + + +def test_npu_platform_inference_mode_preserves_version_counters(): + with torch.inference_mode(False), NPUPlatformBase.inference_mode(): + tensor = torch.ones(1) + + assert tensor._version == 0 + assert torch.is_inference_mode_enabled() is False