[diffusion] quant: support fp8 mixed precision for cosmos3 (#36380)

Co-authored-by: Kedi Wu <kediw@nvidia.com>
This commit is contained in:
Kedi Wu
2026-09-04 09:23:38 +08:00
committed by GitHub
co-authored by Kedi Wu
parent 667bc043dc
commit 94eb15eb6c
5 changed files with 845 additions and 0 deletions
+20
View File
@@ -79,6 +79,9 @@ if TYPE_CHECKING:
SGLANG_DIFFUSION_FLASHINFER_FP4_GEMM_BACKEND: str | None = None
SGLANG_DIFFUSION_ENABLE_W8A8_FP8_GEMM: bool = False
SGLANG_DIFFUSION_FP8_WEIGHT_DEQUANT_CACHE: bool = True
SGLANG_DIFFUSION_ENABLE_COSMOS3_STEP_MIXED_PRECISION: bool = True
SGLANG_DIFFUSION_COSMOS3_STEP_MIXED_PRECISION_FIRST_STEPS: int = 3
SGLANG_DIFFUSION_COSMOS3_STEP_MIXED_PRECISION_LAST_STEPS: int = 3
SGLANG_DIFFUSION_VAE_CHANNELS_LAST_3D: str = "auto"
SGLANG_USE_ROCM_VAE: bool = False
SGLANG_USE_ROCM_CUDNN_BENCHMARK: bool = False
@@ -365,6 +368,23 @@ environment_variables: dict[str, Callable[[], Any]] = {
"SGLANG_DIFFUSION_FP8_WEIGHT_DEQUANT_CACHE": _lazy_bool(
"SGLANG_DIFFUSION_FP8_WEIGHT_DEQUANT_CACHE", "true"
),
# Run the first/last denoising steps of a ModelOpt FP8 (W8A8) Cosmos3 DiT
# as W8A16 when the checkpoint's diffusion_step_policy asks for it; the
# same FP8 weights are dequantized per call and fed to a 16-bit GEMM.
# Kill-switch: set 0 to run pure W8A8 regardless of the checkpoint.
"SGLANG_DIFFUSION_ENABLE_COSMOS3_STEP_MIXED_PRECISION": _lazy_bool(
"SGLANG_DIFFUSION_ENABLE_COSMOS3_STEP_MIXED_PRECISION", "true"
),
# Manual overrides for experiments: setting either explicitly overrides
# that field of the checkpoint policy, or force-enables mixed precision
# on a checkpoint without one (the other field then takes the default
# below). When neither is set, the checkpoint fully owns the behavior.
"SGLANG_DIFFUSION_COSMOS3_STEP_MIXED_PRECISION_FIRST_STEPS": _lazy_int(
"SGLANG_DIFFUSION_COSMOS3_STEP_MIXED_PRECISION_FIRST_STEPS", 3
),
"SGLANG_DIFFUSION_COSMOS3_STEP_MIXED_PRECISION_LAST_STEPS": _lazy_int(
"SGLANG_DIFFUSION_COSMOS3_STEP_MIXED_PRECISION_LAST_STEPS", 3
),
# ROCm: use AITer GroupNorm in VAE for improved performance
"SGLANG_USE_ROCM_VAE": _lazy_bool("SGLANG_USE_ROCM_VAE"),
# ROCm: enable cudnn.benchmark (MIOpen auto-tuning) for VAE conv layers
@@ -0,0 +1,376 @@
# SPDX-License-Identifier: Apache-2.0
"""Step-based mixed activation precision for ModelOpt FP8 DiT linears.
A ModelOpt FP8 checkpoint runs every quantized linear as W8A8 (FP8 weights,
statically quantized FP8 activations). The first and last denoising steps are
the most sensitive to activation quantization error, so this module lets a
transformer run those edge steps as W8A16 instead: the same resident FP8
weights are dequantized to the activation dtype per call and fed to a plain
16-bit GEMM, and ``input_scale`` is simply unused. Middle steps keep the
checkpoint's W8A8 scaled-mm path. No second checkpoint and no extra persistent
weight memory are needed.
The checkpoint owns the step schedule
(``quantization_config.runtime.diffusion_step_policy`` in the component's
``config.json``, schema shared with vLLM-Omni): no policy means no mixed
precision. Explicitly-set env vars act as a manual override for experiments.
The precision is selected once per denoising step (before any
transformer call for that step), so conditional and unconditional CFG branches
of the same step always share one selection. The reasoner (UND) path uses a
static per-request mode from the policy instead of the step schedule.
"""
from __future__ import annotations
import logging
import os
from collections.abc import Iterable, Mapping
import msgspec
import torch
import torch.nn as nn
import torch.nn.functional as F
import sglang.multimodal_gen.envs as envs
from sglang.multimodal_gen.runtime.layers.linear import LinearBase, LinearMethodBase
from sglang.multimodal_gen.runtime.layers.quantization.modelopt_fp8 import (
ModelOptFp8Config as FlatModelOptFp8Config,
)
from sglang.multimodal_gen.runtime.layers.quantization.modelopt_fp8 import (
ModelOptFp8LinearMethod as FlatModelOptFp8LinearMethod,
)
from sglang.multimodal_gen.runtime.layers.quantization.modelopt_quant import (
ModelOptFp8Config as HfModelOptFp8Config,
)
from sglang.multimodal_gen.runtime.layers.quantization.modelopt_quant import (
ModelOptFp8LinearMethod as HfModelOptFp8LinearMethod,
)
logger = logging.getLogger(__name__)
# Two ModelOpt FP8 static per-tensor implementations exist (flat
# `quant_method=modelopt` exports vs `modelopt_fp8` hf_quant_config ones).
# Both store the post-load weight as a column-major [in, out] FP8 view with a
# scalar-or-channelwise weight_scale, so one W8A16 path serves both.
MODELOPT_FP8_QUANT_CONFIGS = (
FlatModelOptFp8Config,
HfModelOptFp8Config,
)
MODELOPT_FP8_LINEAR_METHODS = (
FlatModelOptFp8LinearMethod,
HfModelOptFp8LinearMethod,
)
REASONER_PATH = "reasoner"
GENERATION_PATH = "generation"
# Checkpoint schema shared with vLLM-Omni; unknown or missing fields fail
# closed so a policy the runtime cannot honor never degrades silently.
_POLICY_FIELDS = frozenset(
{
"schema_version",
"type",
"index_space",
"scope",
"default_mode",
"first_steps",
"last_steps",
"overlap",
"reasoner",
}
)
_STEP_RANGE_FIELDS = frozenset({"count", "mode"})
_POLICY_COMPONENT = "transformer"
class StepPolicy(msgspec.Struct, frozen=True, kw_only=True):
first_steps: int
last_steps: int
reasoner_a16: bool = True
class StepMixedPrecisionController:
"""Holds the precision selected for the current denoising step."""
def __init__(
self, first_steps: int, last_steps: int, reasoner_a16: bool = True
) -> None:
if first_steps < 0 or last_steps < 0:
raise ValueError(
f"first_steps/last_steps must be non-negative, got "
f"{first_steps}/{last_steps}"
)
self.first_steps = first_steps
self.last_steps = last_steps
self.reasoner_a16 = reasoner_a16
self.high_precision = False
def use_high_precision(self, path: str) -> bool:
if path == REASONER_PATH:
return self.reasoner_a16
return self.high_precision
def set_step(self, step_index: int, num_steps: int) -> None:
if num_steps <= 0:
raise ValueError(f"num_steps must be positive, got {num_steps}")
if step_index < 0 or step_index >= num_steps:
raise IndexError(
f"step_index must be in [0, {num_steps}), got {step_index}"
)
# A one-step schedule is typically the engine warmup probe; keep it on
# the base W8A8 path rather than treating it as all-edge.
if num_steps == 1:
self.high_precision = False
return
self.high_precision = (
step_index < self.first_steps or step_index >= num_steps - self.last_steps
)
def reset(self) -> None:
self.high_precision = False
def read_checkpoint_step_policy(
quantization_config: Mapping | None,
) -> StepPolicy | None:
"""Parse ``runtime.diffusion_step_policy`` from a checkpoint quant config.
Missing metadata returns None (ordinary checkpoint behavior). Metadata
that is present but malformed or unsupported raises, matching vLLM-Omni's
fail-closed contract for this schema.
"""
if not isinstance(quantization_config, Mapping):
return None
if "runtime" not in quantization_config:
return None
runtime = quantization_config["runtime"]
if not isinstance(runtime, Mapping):
raise TypeError("quantization_config.runtime must be a mapping")
if "diffusion_step_policy" not in runtime:
return None
policy = runtime["diffusion_step_policy"]
if not isinstance(policy, Mapping):
raise TypeError(
"quantization_config.runtime.diffusion_step_policy must be a mapping"
)
return _parse_step_policy(policy)
def _parse_step_policy(policy: Mapping) -> StepPolicy | None:
unknown = set(policy) - _POLICY_FIELDS
if unknown:
raise ValueError(f"Unknown diffusion_step_policy fields: {sorted(unknown)}")
missing = _POLICY_FIELDS - set(policy)
if missing:
raise ValueError(f"Missing diffusion_step_policy fields: {sorted(missing)}")
schema_version = policy["schema_version"]
if (
not isinstance(schema_version, int)
or isinstance(schema_version, bool)
or schema_version != 1
):
raise ValueError("diffusion_step_policy.schema_version must be the integer 1")
if policy["type"] != "first_last_n":
raise ValueError("diffusion_step_policy.type must be 'first_last_n'")
if policy["index_space"] != "denoising_loop_iteration":
raise ValueError(
"diffusion_step_policy.index_space must be 'denoising_loop_iteration'"
)
if policy["default_mode"] != "native":
raise ValueError("diffusion_step_policy.default_mode must be 'native'")
if policy["overlap"] != "a16":
raise ValueError("diffusion_step_policy.overlap must be 'a16'")
scope = policy["scope"]
if (
not isinstance(scope, list)
or not scope
or not all(isinstance(item, str) for item in scope)
):
raise TypeError(
"diffusion_step_policy.scope must be a non-empty list of strings"
)
first_steps = _parse_step_range(policy["first_steps"], "first_steps")
last_steps = _parse_step_range(policy["last_steps"], "last_steps")
reasoner = policy["reasoner"]
if reasoner not in ("native", "a16"):
raise ValueError("diffusion_step_policy.reasoner must be 'native' or 'a16'")
if _POLICY_COMPONENT not in scope:
return None
return StepPolicy(
first_steps=first_steps,
last_steps=last_steps,
reasoner_a16=reasoner == "a16",
)
def _parse_step_range(value: object, name: str) -> int:
if not isinstance(value, Mapping):
raise TypeError(f"diffusion_step_policy.{name} must be a mapping")
unknown = set(value) - _STEP_RANGE_FIELDS
if unknown:
raise ValueError(
f"Unknown diffusion_step_policy.{name} fields: {sorted(unknown)}"
)
missing = _STEP_RANGE_FIELDS - set(value)
if missing:
raise ValueError(
f"Missing diffusion_step_policy.{name} fields: {sorted(missing)}"
)
if value["mode"] != "a16":
raise ValueError(f"diffusion_step_policy.{name}.mode must be 'a16'")
count = value["count"]
if not isinstance(count, int) or isinstance(count, bool) or count < 0:
raise ValueError(
f"diffusion_step_policy.{name}.count must be a non-negative integer"
)
return count
def resolve_step_policy(
quantization_config: Mapping | None,
) -> tuple[StepPolicy | None, str]:
"""Resolve the effective step policy and a human-readable source label.
The checkpoint owns the behavior: mixed precision runs only when the
checkpoint carries a diffusion_step_policy. The enable env var is a
kill-switch, and explicitly-set FIRST/LAST env vars are a manual
override — per field on top of a checkpoint policy, or standing alone
to force-enable on a checkpoint without one.
"""
if not envs.SGLANG_DIFFUSION_ENABLE_COSMOS3_STEP_MIXED_PRECISION:
return (
None,
"disabled by SGLANG_DIFFUSION_ENABLE_COSMOS3_STEP_MIXED_PRECISION=0",
)
checkpoint_policy = read_checkpoint_step_policy(quantization_config)
overridden = []
if "SGLANG_DIFFUSION_COSMOS3_STEP_MIXED_PRECISION_FIRST_STEPS" in os.environ:
overridden.append("first_steps")
if "SGLANG_DIFFUSION_COSMOS3_STEP_MIXED_PRECISION_LAST_STEPS" in os.environ:
overridden.append("last_steps")
if checkpoint_policy is None and not overridden:
return None, "checkpoint carries no diffusion_step_policy"
base = checkpoint_policy or StepPolicy(
first_steps=envs.SGLANG_DIFFUSION_COSMOS3_STEP_MIXED_PRECISION_FIRST_STEPS,
last_steps=envs.SGLANG_DIFFUSION_COSMOS3_STEP_MIXED_PRECISION_LAST_STEPS,
reasoner_a16=True,
)
if checkpoint_policy is None:
source = f"env vars ({', '.join(overridden)} set)"
elif overridden:
source = f"checkpoint with env override of {', '.join(overridden)}"
else:
source = "checkpoint"
first_steps = base.first_steps
last_steps = base.last_steps
if "first_steps" in overridden:
first_steps = envs.SGLANG_DIFFUSION_COSMOS3_STEP_MIXED_PRECISION_FIRST_STEPS
if "last_steps" in overridden:
last_steps = envs.SGLANG_DIFFUSION_COSMOS3_STEP_MIXED_PRECISION_LAST_STEPS
return (
StepPolicy(
first_steps=first_steps,
last_steps=last_steps,
reasoner_a16=base.reasoner_a16,
),
source,
)
def apply_fp8_w8a16_linear(
layer: nn.Module,
x: torch.Tensor,
bias: torch.Tensor | None,
) -> torch.Tensor:
"""16-bit GEMM against the layer's resident ModelOpt FP8 weight.
``ModelOptFp8LinearMethod.process_weights_after_loading`` stores the FP8
weight as a column-major ``[in, out]`` view; ``.t()`` recovers the
row-major ``[out, in]`` layout ``F.linear`` wants. ``weight_scale`` is
either the per-tensor scalar or its channelwise expansion (equal values),
so both broadcast correctly. ``input_scale`` is intentionally unused.
"""
weight = layer.weight.t()
scale = layer.weight_scale.to(x.dtype)
if scale.numel() > 1:
scale = scale.view(-1, 1)
return F.linear(x, weight.to(x.dtype) * scale, bias)
class StepMixedPrecisionFp8LinearMethod(LinearMethodBase):
"""Routes each call to W8A8 (base method) or W8A16 per the controller."""
def __init__(
self,
base_method: LinearMethodBase,
controller: StepMixedPrecisionController,
path: str = GENERATION_PATH,
) -> None:
self.base_method = base_method
self.controller = controller
self.path = path
def create_weights(self, layer: nn.Module, *args, **kwargs) -> None:
self.base_method.create_weights(layer, *args, **kwargs)
def process_weights_after_loading(self, layer: nn.Module) -> None:
self.base_method.process_weights_after_loading(layer)
def apply(
self,
layer: nn.Module,
x: torch.Tensor,
bias: torch.Tensor | None = None,
) -> torch.Tensor:
if self.controller.use_high_precision(self.path):
return apply_fp8_w8a16_linear(layer, x, bias)
return self.base_method.apply(layer, x, bias)
def install_step_mixed_precision(
reasoner_modules: Iterable[nn.Module],
generation_modules: Iterable[nn.Module],
controller: StepMixedPrecisionController,
) -> tuple[int, int]:
"""Wrap every ModelOpt FP8 linear for per-path precision dispatch.
Must run after the loader's ``process_weights_after_loading`` pass so the
wrapped method only ever dispatches ``apply``. Returns the wrapped counts
per path; (0, 0) means the model is not a ModelOpt FP8 checkpoint.
"""
return (
_wrap_path(reasoner_modules, controller, REASONER_PATH),
_wrap_path(generation_modules, controller, GENERATION_PATH),
)
def _wrap_path(
roots: Iterable[nn.Module],
controller: StepMixedPrecisionController,
path: str,
) -> int:
wrapped = 0
for root in roots:
for module in root.modules():
if not isinstance(module, LinearBase):
continue
if not isinstance(module.quant_method, MODELOPT_FP8_LINEAR_METHODS):
continue
module.quant_method = StepMixedPrecisionFp8LinearMethod(
base_method=module.quant_method,
controller=controller,
path=path,
)
wrapped += 1
return wrapped
@@ -43,6 +43,13 @@ from sglang.multimodal_gen.runtime.layers.linear import (
from sglang.multimodal_gen.runtime.layers.quantization.configs.base_config import (
QuantizationConfig,
)
from sglang.multimodal_gen.runtime.layers.quantization.modelopt_fp8_step_precision import (
MODELOPT_FP8_QUANT_CONFIGS,
StepMixedPrecisionController,
install_step_mixed_precision,
read_checkpoint_step_policy,
resolve_step_policy,
)
from sglang.multimodal_gen.runtime.layers.rotary_embedding import (
Qwen3VLTextRotaryEmbedding,
)
@@ -1281,6 +1288,12 @@ class Cosmos3OmniTransformer(CachableDiT, LayerwiseOffloadableModuleMixin):
self.cached_kv: dict[str, list[tuple[torch.Tensor, torch.Tensor]]] = {}
self.cached_gen_rope_inputs: dict[str, tuple[torch.Tensor, torch.Tensor]] = {}
# Installed in post_load_weights when step mixed precision is enabled.
self.step_precision_controller: StepMixedPrecisionController | None = None
self.modelopt_fp8_checkpoint = isinstance(
quant_config, MODELOPT_FP8_QUANT_CONFIGS
)
self.__post_init__()
self.layer_names = ["gen_layers", "language_model.layers"]
@@ -2003,5 +2016,76 @@ class Cosmos3OmniTransformer(CachableDiT, LayerwiseOffloadableModuleMixin):
if isinstance(module, RMSNorm):
module.to(target_dtype)
self._maybe_install_step_mixed_precision()
def _maybe_install_step_mixed_precision(self) -> None:
"""Wrap ModelOpt FP8 linears for per-denoising-step W8A16 dispatch.
The checkpoint owns the behavior: mixed precision runs only when the
checkpoint carries a diffusion_step_policy
(quantization_config.runtime in config.json). Explicitly-set env vars
act as a manual override, and
SGLANG_DIFFUSION_ENABLE_COSMOS3_STEP_MIXED_PRECISION=0 disables. Runs
at the end of post_load_weights so the base quant method has already
transposed weights and collapsed scales.
"""
checkpoint_quant_config = self.hf_config.get("quantization_config")
if not self.modelopt_fp8_checkpoint:
# A checkpoint that carries a step policy the runtime cannot honor
# must fail closed rather than silently run without it.
if read_checkpoint_step_policy(checkpoint_quant_config) is not None:
raise ValueError(
"Checkpoint carries a diffusion_step_policy but was not "
"loaded as a ModelOpt FP8 checkpoint; step mixed precision "
"supports only ModelOpt FP8 in sglang."
)
return
policy, source = resolve_step_policy(checkpoint_quant_config)
if policy is None:
logger.info(
"Step mixed precision off (%s); running W8A8 on every denoising step.",
source,
)
return
controller = StepMixedPrecisionController(
first_steps=policy.first_steps,
last_steps=policy.last_steps,
reasoner_a16=policy.reasoner_a16,
)
reasoner_wrapped, generation_wrapped = install_step_mixed_precision(
reasoner_modules=[self.language_model.layers],
generation_modules=[self.gen_layers],
controller=controller,
)
if reasoner_wrapped + generation_wrapped == 0:
logger.warning(
"ModelOpt FP8 quant config detected but no ModelOpt FP8 "
"linears were found; running without step mixed precision."
)
return
self.step_precision_controller = controller
logger.info(
"Step mixed precision enabled (policy source: %s): %d generation "
"FP8 linears run W8A16 on the first %d and last %d denoising "
"steps; %d reasoner FP8 linears run %s.",
source,
generation_wrapped,
controller.first_steps,
controller.last_steps,
reasoner_wrapped,
"W8A16" if controller.reasoner_a16 else "W8A8",
)
def set_denoising_step(self, step_index: int, num_steps: int) -> None:
"""Select this step's precision before any transformer call for it."""
if self.step_precision_controller is not None:
self.step_precision_controller.set_step(
step_index=step_index, num_steps=num_steps
)
def reset_denoising_step(self) -> None:
if self.step_precision_controller is not None:
self.step_precision_controller.reset()
EntryClass = Cosmos3OmniTransformer
@@ -1608,6 +1608,9 @@ class Cosmos3DenoisingStage(PipelineStage, RolloutDenoisingMixin):
)
for i, t in progress_bar:
# Precision is chosen once per step, before any transformer call,
# so all CFG branches of the step share the same selection.
self.transformer.set_denoising_step(step_index=i, num_steps=len(timesteps))
batch_dim = batch.latents.shape[0] if batch.latents is not None else 1
timestep = t.unsqueeze(0).expand(batch_dim) if t.dim() == 0 else t
# Outside the CFG window the effective scale collapses to 1.0,
@@ -1869,6 +1872,11 @@ class Cosmos3DenoisingStage(PipelineStage, RolloutDenoisingMixin):
if batch.profile and not batch.is_warmup:
self.step_profile()
# Hygiene only: the set_denoising_step at each loop head is what
# actually selects precision, so stale state cannot leak into the
# next request's steps.
self.transformer.reset_denoising_step()
if batch.rollout:
self._postprocess_rollout_outputs(
batch=batch,
@@ -0,0 +1,357 @@
# SPDX-License-Identifier: Apache-2.0
"""Unit tests for step-based W8A8/W8A16 mixed precision on ModelOpt FP8."""
import copy
import os
import unittest
from unittest import mock
import torch
from sglang.multimodal_gen.runtime.layers.linear import ReplicatedLinear
from sglang.multimodal_gen.runtime.layers.quantization.modelopt_fp8 import (
ModelOptFp8Config,
ModelOptFp8LinearMethod,
)
from sglang.multimodal_gen.runtime.layers.quantization.modelopt_fp8_step_precision import (
GENERATION_PATH,
REASONER_PATH,
StepMixedPrecisionController,
StepMixedPrecisionFp8LinearMethod,
StepPolicy,
install_step_mixed_precision,
read_checkpoint_step_policy,
resolve_step_policy,
)
from sglang.multimodal_gen.runtime.layers.quantization.modelopt_quant import (
ModelOptFp8Config as HfModelOptFp8Config,
)
# The published checkpoint schema (transformer/config.json,
# quantization_config.runtime.diffusion_step_policy), shared with vLLM-Omni.
CHECKPOINT_POLICY = {
"schema_version": 1,
"type": "first_last_n",
"index_space": "denoising_loop_iteration",
"scope": ["transformer"],
"default_mode": "native",
"first_steps": {"count": 3, "mode": "a16"},
"last_steps": {"count": 3, "mode": "a16"},
"overlap": "a16",
"reasoner": "a16",
}
def _quant_config_with_policy(policy: dict) -> dict:
return {"quant_algo": "FP8", "runtime": {"diffusion_step_policy": policy}}
def _make_loaded_fp8_linear(
in_features: int = 32, out_features: int = 16
) -> tuple[ReplicatedLinear, torch.Tensor, torch.Tensor]:
"""Build a ReplicatedLinear as the loader would leave it post-load.
Returns the layer plus the raw FP8 weight and per-tensor scale used to
fill it, for computing the expected W8A16 output.
"""
layer = ReplicatedLinear(
in_features,
out_features,
bias=False,
params_dtype=torch.bfloat16,
quant_config=ModelOptFp8Config(),
prefix="gen_layers.0.self_attn.to_qkv",
)
w16 = torch.randn(out_features, in_features, dtype=torch.float32) / 8
scale = (w16.abs().max() / 448.0).reshape(())
w_fp8 = (w16 / scale).to(torch.float8_e4m3fn)
layer.weight.data.copy_(w_fp8)
layer.weight_scale.data.fill_(scale.item())
layer.input_scale.data.fill_(1.0)
layer.quant_method.process_weights_after_loading(layer)
return layer, w_fp8, scale
class TestCheckpointPolicyParsing(unittest.TestCase):
def test_valid_policy_parses(self):
policy = read_checkpoint_step_policy(
_quant_config_with_policy(CHECKPOINT_POLICY)
)
self.assertEqual(
policy, StepPolicy(first_steps=3, last_steps=3, reasoner_a16=True)
)
def test_reasoner_native(self):
raw = copy.deepcopy(CHECKPOINT_POLICY)
raw["reasoner"] = "native"
raw["first_steps"]["count"] = 1
policy = read_checkpoint_step_policy(_quant_config_with_policy(raw))
self.assertEqual(
policy, StepPolicy(first_steps=1, last_steps=3, reasoner_a16=False)
)
def test_missing_metadata_returns_none(self):
self.assertIsNone(read_checkpoint_step_policy(None))
self.assertIsNone(read_checkpoint_step_policy({"quant_algo": "FP8"}))
self.assertIsNone(
read_checkpoint_step_policy({"quant_algo": "FP8", "runtime": {}})
)
def test_scope_without_transformer_returns_none(self):
raw = copy.deepcopy(CHECKPOINT_POLICY)
raw["scope"] = ["vae"]
self.assertIsNone(read_checkpoint_step_policy(_quant_config_with_policy(raw)))
def test_malformed_policy_fails_closed(self):
cases = [
("schema_version", 2),
("type", "sigmoid"),
("index_space", "sigma"),
("default_mode", "a16"),
("overlap", "native"),
("reasoner", "a8"),
("first_steps", {"count": -1, "mode": "a16"}),
("first_steps", {"count": 3, "mode": "a8"}),
("first_steps", {"count": 3}),
("scope", []),
]
for field, bad_value in cases:
raw = copy.deepcopy(CHECKPOINT_POLICY)
raw[field] = bad_value
with self.subTest(field=field):
with self.assertRaises((ValueError, TypeError)):
read_checkpoint_step_policy(_quant_config_with_policy(raw))
raw = copy.deepcopy(CHECKPOINT_POLICY)
raw["surprise"] = 1
with self.assertRaises(ValueError):
read_checkpoint_step_policy(_quant_config_with_policy(raw))
raw = copy.deepcopy(CHECKPOINT_POLICY)
del raw["reasoner"]
with self.assertRaises(ValueError):
read_checkpoint_step_policy(_quant_config_with_policy(raw))
class TestPolicyResolution(unittest.TestCase):
def _clean_environ(self):
patcher = mock.patch.dict(os.environ)
patcher.start()
self.addCleanup(patcher.stop)
for name in (
"SGLANG_DIFFUSION_ENABLE_COSMOS3_STEP_MIXED_PRECISION",
"SGLANG_DIFFUSION_COSMOS3_STEP_MIXED_PRECISION_FIRST_STEPS",
"SGLANG_DIFFUSION_COSMOS3_STEP_MIXED_PRECISION_LAST_STEPS",
):
os.environ.pop(name, None)
def test_checkpoint_policy_enables(self):
self._clean_environ()
raw = copy.deepcopy(CHECKPOINT_POLICY)
raw["first_steps"]["count"] = 1
raw["last_steps"]["count"] = 2
raw["reasoner"] = "native"
policy, source = resolve_step_policy(_quant_config_with_policy(raw))
self.assertEqual(
policy, StepPolicy(first_steps=1, last_steps=2, reasoner_a16=False)
)
self.assertEqual(source, "checkpoint")
def test_off_without_checkpoint_policy(self):
# The checkpoint owns the behavior: no diffusion_step_policy means
# mixed precision must not run.
self._clean_environ()
policy, source = resolve_step_policy({"quant_algo": "FP8"})
self.assertIsNone(policy)
self.assertEqual(source, "checkpoint carries no diffusion_step_policy")
def test_explicit_env_force_enables_without_checkpoint_policy(self):
self._clean_environ()
os.environ["SGLANG_DIFFUSION_COSMOS3_STEP_MIXED_PRECISION_FIRST_STEPS"] = "2"
policy, source = resolve_step_policy({"quant_algo": "FP8"})
self.assertEqual(
policy, StepPolicy(first_steps=2, last_steps=3, reasoner_a16=True)
)
self.assertIn("env vars", source)
def test_explicit_env_overrides_checkpoint_policy(self):
self._clean_environ()
os.environ["SGLANG_DIFFUSION_COSMOS3_STEP_MIXED_PRECISION_FIRST_STEPS"] = "5"
policy, source = resolve_step_policy(
_quant_config_with_policy(CHECKPOINT_POLICY)
)
self.assertEqual(
policy, StepPolicy(first_steps=5, last_steps=3, reasoner_a16=True)
)
self.assertIn("env override of first_steps", source)
def test_kill_switch_disables(self):
self._clean_environ()
os.environ["SGLANG_DIFFUSION_ENABLE_COSMOS3_STEP_MIXED_PRECISION"] = "0"
policy, source = resolve_step_policy(
_quant_config_with_policy(CHECKPOINT_POLICY)
)
self.assertIsNone(policy)
self.assertIn("disabled", source)
class TestStepPolicyDispatch(unittest.TestCase):
def test_edge_steps_select_high_precision(self):
controller = StepMixedPrecisionController(first_steps=3, last_steps=3)
selected = []
for step in range(10):
controller.set_step(step_index=step, num_steps=10)
selected.append(controller.high_precision)
self.assertEqual(
selected,
[True, True, True, False, False, False, False, True, True, True],
)
def test_single_step_schedule_stays_base_precision(self):
controller = StepMixedPrecisionController(first_steps=3, last_steps=3)
controller.set_step(step_index=0, num_steps=1)
self.assertFalse(controller.high_precision)
def test_reasoner_path_is_static(self):
controller = StepMixedPrecisionController(
first_steps=1, last_steps=1, reasoner_a16=False
)
controller.set_step(step_index=0, num_steps=10)
self.assertTrue(controller.use_high_precision(GENERATION_PATH))
self.assertFalse(controller.use_high_precision(REASONER_PATH))
controller = StepMixedPrecisionController(
first_steps=0, last_steps=0, reasoner_a16=True
)
controller.set_step(step_index=5, num_steps=10)
self.assertFalse(controller.use_high_precision(GENERATION_PATH))
self.assertTrue(controller.use_high_precision(REASONER_PATH))
def test_reset_returns_to_base_precision(self):
controller = StepMixedPrecisionController(first_steps=1, last_steps=0)
controller.set_step(step_index=0, num_steps=4)
self.assertTrue(controller.high_precision)
controller.reset()
self.assertFalse(controller.high_precision)
def test_invalid_inputs_raise(self):
with self.assertRaises(ValueError):
StepMixedPrecisionController(first_steps=-1, last_steps=0)
controller = StepMixedPrecisionController(first_steps=1, last_steps=1)
with self.assertRaises(ValueError):
controller.set_step(step_index=0, num_steps=0)
with self.assertRaises(IndexError):
controller.set_step(step_index=5, num_steps=5)
class TestInstallAndDispatch(unittest.TestCase):
def test_install_wraps_only_modelopt_fp8_linears(self):
fp8_layer, _, _ = _make_loaded_fp8_linear()
bf16_layer = ReplicatedLinear(
8, 8, bias=False, params_dtype=torch.bfloat16, prefix="norm_out"
)
root = torch.nn.ModuleList([fp8_layer, bf16_layer])
controller = StepMixedPrecisionController(first_steps=3, last_steps=3)
reasoner_wrapped, generation_wrapped = install_step_mixed_precision(
reasoner_modules=[], generation_modules=[root], controller=controller
)
self.assertEqual((reasoner_wrapped, generation_wrapped), (0, 1))
self.assertIsInstance(fp8_layer.quant_method, StepMixedPrecisionFp8LinearMethod)
self.assertEqual(fp8_layer.quant_method.path, GENERATION_PATH)
self.assertNotIsInstance(
bf16_layer.quant_method, StepMixedPrecisionFp8LinearMethod
)
def test_w8a16_matches_dequantized_reference(self):
layer, w_fp8, scale = _make_loaded_fp8_linear()
controller = StepMixedPrecisionController(first_steps=1, last_steps=0)
install_step_mixed_precision(
reasoner_modules=[], generation_modules=[layer], controller=controller
)
controller.set_step(step_index=0, num_steps=4)
self.assertTrue(controller.high_precision)
x = torch.randn(5, layer.input_size, dtype=torch.bfloat16)
out, _ = layer(x)
expected = torch.nn.functional.linear(
x, w_fp8.to(torch.bfloat16) * scale.to(torch.bfloat16)
)
torch.testing.assert_close(out, expected)
def test_install_wraps_hf_quant_config_variant(self):
# The `modelopt_fp8` hf_quant_config path uses a different
# ModelOptFp8LinearMethod class (modelopt_quant.py); the installer
# must wrap it too and the shared W8A16 dequant must hold.
layer = ReplicatedLinear(
32,
16,
bias=False,
params_dtype=torch.bfloat16,
quant_config=HfModelOptFp8Config(is_checkpoint_fp8_serialized=True),
prefix="gen_layers.0.self_attn.to_qkv",
)
w16 = torch.randn(16, 32, dtype=torch.float32) / 8
scale = (w16.abs().max() / 448.0).reshape(())
w_fp8 = (w16 / scale).to(torch.float8_e4m3fn)
layer.weight.data.copy_(w_fp8)
# Emulate this method's post-load state (its real pass needs CUDA
# quant kernels): transposed FP8 view plus collapsed scalar scales.
layer.weight.data = layer.weight.data.t()
layer.weight_scale.data = scale.clone()
layer.input_scale.data = torch.ones(())
controller = StepMixedPrecisionController(first_steps=1, last_steps=0)
_, generation_wrapped = install_step_mixed_precision(
reasoner_modules=[], generation_modules=[layer], controller=controller
)
self.assertEqual(generation_wrapped, 1)
controller.set_step(step_index=0, num_steps=4)
x = torch.randn(5, 32, dtype=torch.bfloat16)
out, _ = layer(x)
expected = torch.nn.functional.linear(
x, w_fp8.to(torch.bfloat16) * scale.to(torch.bfloat16)
)
torch.testing.assert_close(out, expected)
def test_reasoner_native_dispatches_to_w8a8_on_edge_steps(self):
layer, _, _ = _make_loaded_fp8_linear()
controller = StepMixedPrecisionController(
first_steps=1, last_steps=1, reasoner_a16=False
)
install_step_mixed_precision(
reasoner_modules=[layer], generation_modules=[], controller=controller
)
self.assertEqual(layer.quant_method.path, REASONER_PATH)
base = layer.quant_method.base_method
x = torch.randn(2, layer.input_size, dtype=torch.bfloat16)
with mock.patch.object(
base, "apply", return_value=torch.zeros(2, layer.output_size)
) as base_apply:
controller.set_step(step_index=0, num_steps=6)
layer(x)
base_apply.assert_called_once()
def test_base_steps_dispatch_to_w8a8_method(self):
layer, _, _ = _make_loaded_fp8_linear()
controller = StepMixedPrecisionController(first_steps=1, last_steps=1)
install_step_mixed_precision(
reasoner_modules=[], generation_modules=[layer], controller=controller
)
base = layer.quant_method.base_method
self.assertIsInstance(base, ModelOptFp8LinearMethod)
x = torch.randn(2, layer.input_size, dtype=torch.bfloat16)
with mock.patch.object(
base, "apply", return_value=torch.zeros(2, layer.output_size)
) as base_apply:
controller.set_step(step_index=2, num_steps=6)
layer(x)
base_apply.assert_called_once()
base_apply.reset_mock()
# Edge step: the W8A16 path runs and the base method is bypassed.
controller.set_step(step_index=5, num_steps=6)
layer(x)
base_apply.assert_not_called()
if __name__ == "__main__":
unittest.main()