diff --git a/docs/diffusion/quantization.md b/docs/diffusion/quantization.md index 96c872d4f..5a33a3544 100644 --- a/docs/diffusion/quantization.md +++ b/docs/diffusion/quantization.md @@ -43,7 +43,7 @@ backend. | quant_family | checkpoint form | canonical CLI | supported models | extra dependency | platform / notes | |------------------|--------------------------------------------------------------------------------------------|------------------------------------------------------|--------------------------------------------------------------|---------------------------------------|-----------------------------------------------------------------------------------------------------------------------| | `fp8` | Quantized transformer component folder, or safetensors with `quantization_config` metadata | `--transformer-path` or `--transformer-weights-path` | ALL | None | Component-folder and single-file flows are both supported | -| `modelopt-fp8` | Converted ModelOpt FP8 transformer directory or repo with `config.json` | `--transformer-weights-path` | FLUX.2, Wan2.2 | None | Override config is read from the quantized transformer repo; `dit_cpu_offload` and `dit_layerwise_offload` are auto-disabled | +| `modelopt-fp8` | Converted ModelOpt FP8 transformer directory or repo with `config.json` | `--transformer-weights-path` | FLUX.2, Wan2.2 | None | Override config is read from the quantized transformer repo; `dit_layerwise_offload` is supported and `dit_cpu_offload` stays disabled | | `nvfp4-modelopt` | NVFP4 safetensors file, sharded directory, or repo providing transformer weights | `--transformer-weights-path` | FLUX.2 | `comfy-kitchen` optional on Blackwell | Blackwell can use a best-performance kit when available; otherwise SGLang falls back to the generic ModelOpt FP4 path | | `nunchaku-svdq` | Pre-quantized Nunchaku transformer weights, usually named `svdq-{int4\|fp4}_r{rank}-...` | `--transformer-weights-path` | Model-specific support such as Qwen-Image, FLUX, and Z-Image | `nunchaku` | SGLang can infer precision and rank from the filename and supports both `int4` and `nvfp4` | | `msmodelslim` | Pre-quantized msmodelslim transformer weights | `--model-path` | Wan2.2 family | None | Currently only compatible with the Ascend NPU family and supports both `w8a8` and `w4a4` | @@ -98,9 +98,10 @@ sglang generate \ - If the override repo or local directory contains its own `config.json`, SGLang reads the quantization config from that override instead of relying on the base model config. -- `dit_cpu_offload` and `dit_layerwise_offload` are automatically disabled for - ModelOpt FP8 checkpoints because the runtime expects the transformed FP8 - weights to remain GPU-resident in their column-major layout. +- `dit_layerwise_offload` is supported for ModelOpt FP8 checkpoints. +- `dit_cpu_offload` still stays disabled for ModelOpt FP8 checkpoints. +- The layerwise offload path now preserves the non-contiguous FP8 weight stride + expected by the runtime FP8 GEMM path. - To build the converted checkpoint yourself from a ModelOpt diffusers export, use `python -m sglang.multimodal_gen.tools.build_modelopt_fp8_transformer`. diff --git a/python/sglang/multimodal_gen/.claude/skills/sglang-diffusion-modelopt-quant/SKILL.md b/python/sglang/multimodal_gen/.claude/skills/sglang-diffusion-modelopt-quant/SKILL.md index 46dd8137c..4e81ef4eb 100644 --- a/python/sglang/multimodal_gen/.claude/skills/sglang-diffusion-modelopt-quant/SKILL.md +++ b/python/sglang/multimodal_gen/.claude/skills/sglang-diffusion-modelopt-quant/SKILL.md @@ -21,7 +21,7 @@ This skill owns the ModelOpt-to-SGLang bridge. It is not a generic kernel-tuning - Use ModelOpt's official `quantize.py` as the PTQ source of truth. - Keep the workflow generic. Put model-specific fallback logic in small isolated branches, not in the main conversion path. - Benchmark only when BF16 and quantized commands are identical except for the checkpoint override being tested. -- For diffusion FP8, pin `dit_cpu_offload=false` and `dit_layerwise_offload=false`. +- For diffusion FP8, keep `dit_cpu_offload=false`. `dit_layerwise_offload=true` is valid on the fixed path when you want lower DiT residency. - For multi-transformer pipelines, use per-component overrides when different components need different checkpoints. - When a branch is missing the validated helper tools, refresh `python/sglang/multimodal_gen/tools/build_modelopt_fp8_transformer.py`, `python/sglang/multimodal_gen/tools/build_modelopt_nvfp4_transformer.py`, and `python/sglang/multimodal_gen/tools/compare_diffusion_trajectory_similarity.py` instead of inventing one-off scripts elsewhere. - After validating a new ModelOpt quant path, update the ModelOpt support matrix in `docs/diffusion/quantization.md` before closing the task. @@ -52,7 +52,7 @@ This repo now contains: - diffusion-side ModelOpt FP8 linear loading - diffusion-side NVFP4 loading from ModelOpt exports - FLUX.2 packed-QKV detection that distinguishes packed NVFP4 checkpoints from standard diffusers exports -- automatic protection against incompatible FP8 offload modes +- automatic protection against incompatible FP8 CPU offload while keeping layerwise DiT offload available - FP8 transformer build: [`python/sglang/multimodal_gen/tools/build_modelopt_fp8_transformer.py`](../../../tools/build_modelopt_fp8_transformer.py) - trajectory similarity validation: @@ -279,18 +279,18 @@ Do not turn one validated model quirk into a generic rule unless another family Current diffusion ModelOpt FP8 support requires: - `dit_cpu_offload=false` -- `dit_layerwise_offload=false` +- `dit_layerwise_offload` may be enabled when you want lower DiT residency Reason: - the FP8 linear path depends on a CUTLASS-compatible weight layout after loading -- the offload and restore path does not preserve that layout -- in particular, layerwise offload can flatten and rebuild FP8 weights in a way that breaks the column-major requirement used by the FP8 GEMM path +- `dit_cpu_offload` is still treated conservatively +- the fixed layerwise offload path now preserves non-contiguous tensor strides instead of flattening and rebuilding FP8 weights into a contiguous layout Runtime behavior: -- SGLang currently force-disables these two flags when it detects `modelopt_fp8` -- benchmark commands should still pin them explicitly so the command line itself makes the comparison rule obvious +- SGLang still force-disables `dit_cpu_offload` when it detects `modelopt_fp8` +- benchmark commands should still pin offload flags explicitly so the command line itself makes the comparison rule obvious ## Claim Discipline diff --git a/python/sglang/multimodal_gen/runtime/loader/transformer_load_utils.py b/python/sglang/multimodal_gen/runtime/loader/transformer_load_utils.py index 5ecffdfd0..5885fa03b 100644 --- a/python/sglang/multimodal_gen/runtime/loader/transformer_load_utils.py +++ b/python/sglang/multimodal_gen/runtime/loader/transformer_load_utils.py @@ -186,26 +186,13 @@ class _ModelOptFp8OffloadAdapter(_TransformerQuantAdapter): if quant_name != "modelopt_fp8": return - disabled_args: list[str] = [] - if server_args.dit_cpu_offload: server_args.dit_cpu_offload = False - disabled_args.append("dit_cpu_offload") - - if server_args.dit_layerwise_offload: - server_args.dit_layerwise_offload = False - disabled_args.append("dit_layerwise_offload") - - if not disabled_args: - return - - logger.warning( - "ModelOpt FP8 diffusion checkpoints currently require the transformer " - "FP8 weights to stay GPU-resident in their column-major layout; " - "disabling %s for this run. Text encoder / VAE offload settings are " - "left unchanged.", - ", ".join(disabled_args), - ) + logger.warning( + "ModelOpt FP8 diffusion checkpoints currently keep dit_cpu_offload " + "disabled. Layerwise DiT offload stays enabled because the runtime " + "now preserves the restored FP8 tensor strides.", + ) def prepare(self) -> None: _ModelOptFp8OffloadAdapter._maybe_disable_incompatible_dit_offload_modes( diff --git a/python/sglang/multimodal_gen/runtime/utils/layerwise_offload.py b/python/sglang/multimodal_gen/runtime/utils/layerwise_offload.py index 05a3d676d..29d601845 100644 --- a/python/sglang/multimodal_gen/runtime/utils/layerwise_offload.py +++ b/python/sglang/multimodal_gen/runtime/utils/layerwise_offload.py @@ -56,6 +56,9 @@ class LayerwiseOffloadManager: # layer_idx -> {dtype: consolidated_pinned_cpu_tensor} # stores the consolidated weight from a same layer, of same dtype self._consolidated_cpu_weights: Dict[int, Dict[torch.dtype, torch.Tensor]] = {} + # layer_idx -> {name: pinned_cpu_tensor_with_original_stride} + # stores tensors whose original non-contiguous stride/layout must be preserved + self._strided_cpu_weights: Dict[int, Dict[str, torch.Tensor]] = {} # layer_idx -> {name: {dtype, offset, numel, shape}} # stores the offset and numel of each weight from a same layer, of same dtype self._weight_metadata: Dict[int, Dict[str, Dict[str, Any]]] = {} @@ -88,6 +91,21 @@ class LayerwiseOffloadManager: self._offload_placeholders[dtype] = placeholder return placeholder + @staticmethod + def _get_alignment_numel(dtype: torch.dtype, alignment_bytes: int = 32) -> int: + element_size = torch.empty((), dtype=dtype).element_size() + return max(1, alignment_bytes // element_size) + + @classmethod + def _align_numel_offset( + cls, offset: int, dtype: torch.dtype, alignment_bytes: int = 32 + ) -> int: + alignment_numel = cls._get_alignment_numel(dtype, alignment_bytes) + remainder = offset % alignment_numel + if remainder == 0: + return offset + return offset + alignment_numel - remainder + @torch.compiler.disable def _initialize(self) -> None: if not self.enabled: @@ -110,10 +128,49 @@ class LayerwiseOffloadManager: # 2. concat and offload (in pinned memory) for layer_idx, dtype_to_params in layer_groups.items(): self._consolidated_cpu_weights[layer_idx] = {} + self._strided_cpu_weights[layer_idx] = {} self._weight_metadata[layer_idx] = {} for dtype, weights in dtype_to_params.items(): - total_numel = sum(t.numel() for _, t in weights) + contiguous_weights: List[Tuple[str, torch.Tensor]] = [] + for name, weight in weights: + if weight.is_contiguous(): + contiguous_weights.append((name, weight)) + continue + + # Preserve non-contiguous layouts such as the transposed FP8 + # weight views expected by CUTLASS kernels. + cpu_tensor = torch.empty_strided( + size=weight.shape, + stride=weight.stride(), + dtype=dtype, + pin_memory=self.pin_cpu_memory, + ) + cpu_tensor.copy_(weight) + self._strided_cpu_weights[layer_idx][name] = cpu_tensor + self._weight_metadata[layer_idx][name] = { + "dtype": dtype, + "shape": weight.shape, + "stride": weight.stride(), + "preserve_strides": True, + } + weight.data = self._get_shared_empty_tensor(dtype) + + if not contiguous_weights: + continue + + current_offset = 0 + aligned_offsets: Dict[str, int] = {} + for name, weight in contiguous_weights: + # Some fused diffusion kernels require tensor base pointers to + # satisfy a 32-byte alignment contract. Reusing one flat buffer + # is still fine, but each logical tensor slice must start on an + # aligned offset inside that buffer. + current_offset = self._align_numel_offset(current_offset, dtype) + aligned_offsets[name] = current_offset + current_offset += weight.numel() + + total_numel = current_offset # create concatenated CPU buffer (in pinned memory) cpu_buffer = torch.empty( @@ -121,8 +178,8 @@ class LayerwiseOffloadManager: ) # offload weights to the buffer - current_offset = 0 - for name, weight in weights: + for name, weight in contiguous_weights: + current_offset = aligned_offsets[name] numel = weight.numel() cpu_buffer[current_offset : current_offset + numel].copy_( weight.flatten() @@ -132,6 +189,8 @@ class LayerwiseOffloadManager: "offset": current_offset, "numel": numel, "shape": weight.shape, + "stride": weight.stride(), + "preserve_strides": False, } weight.data = self._get_shared_empty_tensor(dtype) @@ -190,22 +249,39 @@ class LayerwiseOffloadManager: gpu_buffer.copy_(cpu_buffer, non_blocking=non_blocking) gpu_buffers[dtype] = gpu_buffer - # record the prefetch event of this layer + # restore model's weights by their metadata using the same copy stream + # so the recorded event covers both flat-buffer and stride-preserving copies. + for name, meta in self._weight_metadata[layer_idx].items(): + target = self.get_target_with_name(name) + if meta.get("preserve_strides", False): + # Recreate the original view layout instead of flatten+view. + # ModelOpt FP8 relies on a transposed runtime weight layout, + # so preserving stride is part of correctness, not just an + # optimization detail. + cpu_tensor = self._strided_cpu_weights[layer_idx][name] + gpu_tensor = torch.empty_strided( + size=meta["shape"], + stride=meta["stride"], + dtype=meta["dtype"], + device=self.device, + ) + gpu_tensor.copy_(cpu_tensor, non_blocking=non_blocking) + target.data = gpu_tensor + continue + + dtype = meta["dtype"] + gpu_buffer = gpu_buffers[dtype] + + # map the parameter's data to the correct slice of the GPU buffer + target.data = gpu_buffer[ + meta["offset"] : meta["offset"] + meta["numel"] + ].view(meta["shape"]) + + # record the prefetch event of this layer after all copies are enqueued event = torch.get_device_module().Event() event.record(self.copy_stream) self._prefetch_events[layer_idx] = event - # restore model's weights by their metadata using gpu buffer - for name, meta in self._weight_metadata[layer_idx].items(): - dtype = meta["dtype"] - gpu_buffer = gpu_buffers[dtype] - - # map the parameter's data to the correct slice of the GPU buffer - target = self.get_target_with_name(name) - target.data = gpu_buffer[ - meta["offset"] : meta["offset"] + meta["numel"] - ].view(meta["shape"]) - self._gpu_layers.add(layer_idx) @torch.compiler.disable @@ -266,6 +342,10 @@ class LayerwiseOffloadManager: # Collect current GPU weights and write back to CPU buffer for name, meta in self._weight_metadata.get(layer_idx, {}).items(): target = self.get_target_with_name(name) + if meta.get("preserve_strides", False): + self._strided_cpu_weights[layer_idx][name].copy_(target.data.cpu()) + continue + gpu_weight = target.data.flatten().cpu() dtype = meta["dtype"] @@ -331,12 +411,17 @@ class LayerwiseOffloadManager: ) dtype = meta["dtype"] - offset = meta["offset"] - numel = meta["numel"] - cpu_buffer = self._consolidated_cpu_weights[layer_idx][dtype] - cpu_buffer[offset : offset + numel].copy_( - loaded_weight.to(dtype=dtype).flatten() - ) + if meta.get("preserve_strides", False): + self._strided_cpu_weights[layer_idx][name].copy_( + loaded_weight.to(dtype=dtype) + ) + else: + offset = meta["offset"] + numel = meta["numel"] + cpu_buffer = self._consolidated_cpu_weights[layer_idx][dtype] + cpu_buffer[offset : offset + numel].copy_( + loaded_weight.to(dtype=dtype).flatten() + ) # If this layer is currently on GPU, update the live parameter. if layer_idx in self._gpu_layers: @@ -358,6 +443,14 @@ class LayerwiseOffloadManager: """ for layer_idx in sorted(self._weight_metadata): for name, meta in self._weight_metadata[layer_idx].items(): + if meta.get("preserve_strides", False): + # Some quantized weights rely on a non-contiguous layout. + # Yield the strided tensor directly instead of rebuilding it + # from the flat buffer, which would silently lose the + # original stride information. + yield name, self._strided_cpu_weights[layer_idx][name] + continue + dtype = meta["dtype"] offset = meta["offset"] numel = meta["numel"] diff --git a/python/sglang/multimodal_gen/test/unit/test_layerwise_offload.py b/python/sglang/multimodal_gen/test/unit/test_layerwise_offload.py new file mode 100644 index 000000000..73b256683 --- /dev/null +++ b/python/sglang/multimodal_gen/test/unit/test_layerwise_offload.py @@ -0,0 +1,166 @@ +from contextlib import nullcontext +from types import SimpleNamespace + +import torch + +from sglang.multimodal_gen.runtime.layers.quantization.modelopt_quant import ( + ModelOptFp8Config, +) +from sglang.multimodal_gen.runtime.loader.transformer_load_utils import ( + _ModelOptFp8OffloadAdapter, +) +from sglang.multimodal_gen.runtime.utils import ( + layerwise_offload as layerwise_offload_mod, +) +from sglang.multimodal_gen.runtime.utils.layerwise_offload import ( + LayerwiseOffloadManager, +) + + +class _FakeStream: + def wait_stream(self, _stream) -> None: + return None + + def wait_event(self, _event) -> None: + return None + + +class _FakeEvent: + def record(self, _stream) -> None: + return None + + +class _FakeDeviceModule: + Stream = _FakeStream + Event = _FakeEvent + + @staticmethod + def is_available() -> bool: + return True + + @staticmethod + def current_device() -> int: + return 0 + + @staticmethod + def current_stream() -> _FakeStream: + return _FakeStream() + + @staticmethod + def stream(_stream): + return nullcontext() + + +class _DummyBlock(torch.nn.Module): + def __init__(self) -> None: + super().__init__() + base = torch.arange(12, dtype=torch.float32).reshape(3, 4) + self.weight = torch.nn.Parameter(base.t()) + self.bias = torch.nn.Parameter(torch.arange(3, dtype=torch.float32)) + + +class _DummyModel(torch.nn.Module): + def __init__(self) -> None: + super().__init__() + self.blocks = torch.nn.ModuleList([_DummyBlock()]) + + +def test_layerwise_offload_preserves_non_contiguous_stride(monkeypatch): + monkeypatch.setattr( + layerwise_offload_mod.torch, "get_device_module", lambda: _FakeDeviceModule + ) + monkeypatch.setattr(layerwise_offload_mod.current_platform, "device_type", "cpu") + + model = _DummyModel() + original_weight = model.blocks[0].weight.detach().clone() + original_stride = model.blocks[0].weight.stride() + assert not model.blocks[0].weight.is_contiguous() + + manager = LayerwiseOffloadManager( + model=model, + layers_attr_str="blocks", + num_layers=1, + enabled=True, + pin_cpu_memory=False, + prefetch_size=1, + ) + + meta = manager._weight_metadata[0]["blocks.0.weight"] + assert meta["preserve_strides"] is True + + restored_weight = model.blocks[0].weight.data + assert restored_weight.shape == original_weight.shape + assert restored_weight.stride() == original_stride + assert not restored_weight.is_contiguous() + assert torch.equal(restored_weight, original_weight) + + manager.release_layer(0) + manager.prefetch_layer(0, non_blocking=False) + + reloaded_weight = model.blocks[0].weight.data + assert reloaded_weight.stride() == original_stride + assert not reloaded_weight.is_contiguous() + assert torch.equal(reloaded_weight, original_weight) + + +def test_modelopt_fp8_adapter_keeps_layerwise_offload_enabled(): + server_args = SimpleNamespace( + dit_cpu_offload=True, + dit_layerwise_offload=True, + ) + quant_config = ModelOptFp8Config(is_checkpoint_fp8_serialized=True) + + _ModelOptFp8OffloadAdapter._maybe_disable_incompatible_dit_offload_modes( + server_args=server_args, + quant_config=quant_config, + ) + + assert server_args.dit_cpu_offload is False + assert server_args.dit_layerwise_offload is True + + +def test_layerwise_offload_aligns_contiguous_tensor_offsets(monkeypatch): + monkeypatch.setattr( + layerwise_offload_mod.torch, "get_device_module", lambda: _FakeDeviceModule + ) + monkeypatch.setattr(layerwise_offload_mod.current_platform, "device_type", "cpu") + + class _AlignedDummyBlock(torch.nn.Module): + def __init__(self) -> None: + super().__init__() + self.weight = torch.nn.Parameter( + torch.arange(9, dtype=torch.float32).reshape(3, 3) + ) + self.bias = torch.nn.Parameter(torch.arange(3, dtype=torch.float32)) + + class _AlignedDummyModel(torch.nn.Module): + def __init__(self) -> None: + super().__init__() + self.blocks = torch.nn.ModuleList([_AlignedDummyBlock()]) + + model = _AlignedDummyModel() + original_weight = model.blocks[0].weight.detach().clone() + original_bias = model.blocks[0].bias.detach().clone() + + manager = LayerwiseOffloadManager( + model=model, + layers_attr_str="blocks", + num_layers=1, + enabled=True, + pin_cpu_memory=False, + prefetch_size=1, + ) + + weight_meta = manager._weight_metadata[0]["blocks.0.weight"] + bias_meta = manager._weight_metadata[0]["blocks.0.bias"] + assert weight_meta["preserve_strides"] is False + assert bias_meta["preserve_strides"] is False + assert weight_meta["offset"] == 0 + assert bias_meta["offset"] % 8 == 0 + + restored_weight = model.blocks[0].weight.data + restored_bias = model.blocks[0].bias.data + assert restored_weight.data_ptr() % 32 == 0 + assert restored_bias.data_ptr() % 32 == 0 + assert torch.equal(restored_weight, original_weight) + assert torch.equal(restored_bias, original_bias)