Skip to content
Merged
Show file tree
Hide file tree
Changes from 25 commits
Commits
Show all changes
35 commits
Select commit Hold shift + click to select a range
043a941
[diffusion][CI]: Add individual component accuracy CI for diffusion m…
Ratish1 Feb 4, 2026
5f93ead
remove fallback
Ratish1 Feb 5, 2026
d18275e
Merge remote-tracking branch 'upstream/main' into feat/accuracy-test
Ratish1 Feb 9, 2026
3df80fe
upd
Ratish1 Feb 12, 2026
668e52e
Merge remote-tracking branch 'upstream/main' into feat/accuracy-test
Ratish1 Feb 23, 2026
8d4061f
test: migrate diffusion accuracy to native hook architecture
Ratish1 Feb 23, 2026
6958597
upd
Ratish1 Feb 26, 2026
61fce33
upd
Ratish1 Feb 26, 2026
d594a1c
upd
Ratish1 Mar 9, 2026
3754993
fix conflict
Ratish1 Mar 19, 2026
0dbabd7
Merge remote-tracking branch 'upstream/main' into feat/accuracy-test
Ratish1 Mar 25, 2026
61916ea
test(multimodal-gen): stabilize 2-GPU component accuracy harness
Ratish1 Mar 27, 2026
d804385
upd
Ratish1 Mar 27, 2026
5f623f1
move shard context helper into accuracy utils
Ratish1 Mar 27, 2026
d48a1ee
move accuracy runtime helpers into accuracy utils
Ratish1 Mar 27, 2026
ea39b1a
deduplicate text encoder module resolution
Ratish1 Mar 27, 2026
265f6e8
inline native forward output capture
Ratish1 Mar 27, 2026
fa816f4
deduplicate native output normalization
Ratish1 Mar 28, 2026
34c17e6
remove generic native profile registry
Ratish1 Mar 28, 2026
b7e6a68
rename accuracy hook helpers for clarity
Ratish1 Mar 28, 2026
ed16d98
make native profile and text encoder helpers more explicit
Ratish1 Mar 28, 2026
0300ad9
reduce unnecessary memory cleanup work between accuracy stages
Ratish1 Mar 28, 2026
0e5f665
remove unused accuracy helpers
Ratish1 Mar 28, 2026
9edb668
[diffusion] test: trim accuracy harness comments
Ratish1 Mar 28, 2026
4d8ccb0
Merge remote-tracking branch 'upstream/main' into feat/accuracy-test
Ratish1 Mar 28, 2026
280fd24
[diffusion] test: reduce duplicate accuracy coverage and stabilize 2-…
Ratish1 Mar 28, 2026
eb82f80
Merge remote-tracking branch 'upstream/main' into feat/accuracy-test
Ratish1 Mar 28, 2026
6534ea3
Merge branch 'main' into feat/accuracy-test
BBuf Mar 30, 2026
93307ad
fix ci names
Ratish1 Mar 30, 2026
c5f54e9
fix ci path
Ratish1 Mar 30, 2026
fe41a77
upd
Ratish1 Mar 30, 2026
9d121fc
fix OOM on 2 gpu cases
Ratish1 Mar 30, 2026
6d4245b
add CI tests to diffusion workflow
Ratish1 Mar 30, 2026
c1edd23
fix
Ratish1 Mar 30, 2026
96c5fdb
fix
Ratish1 Mar 30, 2026
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
Original file line number Diff line number Diff line change
Expand Up @@ -156,7 +156,9 @@ def _prepare_weights(
return hf_folder, hf_weights_files, use_safetensors

def _get_weights_iterator(
self, source: "Source", to_cpu: bool
self,
source: "Source",
to_cpu: bool,
) -> Generator[tuple[str, torch.Tensor], None, None]:
"""get an iterator for the model weights based on the load format."""
hf_folder, hf_weights_files, use_safetensors = self._prepare_weights(
Expand All @@ -166,7 +168,8 @@ def _get_weights_iterator(
)
if use_safetensors:
weights_iterator = safetensors_weights_iterator(
hf_weights_files, to_cpu=to_cpu
hf_weights_files,
to_cpu=to_cpu,
)
else:
weights_iterator = pt_weights_iterator(hf_weights_files, to_cpu=to_cpu)
Expand All @@ -186,17 +189,27 @@ def _get_all_weights(
fall_back_to_pt=getattr(model, "fall_back_to_pt_during_load", True),
allow_patterns_overrides=getattr(model, "allow_patterns_overrides", None),
)
yield from self._get_weights_iterator(primary_weights, to_cpu)
yield from self._get_weights_iterator(
primary_weights,
to_cpu,
)

secondary_weights = cast(
Iterable[TextEncoderLoader.Source],
getattr(model, "secondary_weights", ()),
)
for source in secondary_weights:
yield from self._get_weights_iterator(source, to_cpu)
yield from self._get_weights_iterator(
source,
to_cpu,
)

def load_customized(
self, component_model_path: str, server_args: ServerArgs, component_name: str
self,
component_model_path: str,
server_args: ServerArgs,
component_name: str,
cpu_offload_flag: bool | None = None,
):
"""Load the text encoders based on the model path, and inference args."""
diffusers_pretrained_config = get_config(
Expand Down Expand Up @@ -227,6 +240,7 @@ def is_not_first_encoder(module_name):
encoder_config,
server_args,
encoder_dtype,
cpu_offload_flag=cpu_offload_flag,
)

def load_model(
Expand All @@ -240,7 +254,10 @@ def load_model(
# Determine CPU offload behavior and target device

local_torch_device = get_local_torch_device()
should_offload = self.should_offload(server_args, model_config)
fsdp_cpu_offload = self.should_offload(server_args, model_config)
should_offload = (
cpu_offload_flag if cpu_offload_flag is not None else fsdp_cpu_offload
)

if should_offload and not current_platform.is_mps():
model_device = torch.device("cpu")
Expand All @@ -263,21 +280,21 @@ def load_model(

weights_to_load = {name for name, _ in model.named_parameters()}
loaded_weights = model.load_weights(
self._get_all_weights(model, model_path, to_cpu=should_offload)
self._get_all_weights(
model,
model_path,
to_cpu=should_offload,
)
)

# Explicitly move model to target device after loading weights
if not should_offload:
model = model.to(local_torch_device)

if should_offload:
# Disable FSDP for MPS as it's not compatible
if current_platform.is_mps():
logger.info(
"Disabling FSDP sharding for MPS platform as it's not compatible"
)
model = model.to(local_torch_device)
else:
elif fsdp_cpu_offload:
mesh = init_device_mesh(
current_platform.device_type,
mesh_shape=(1, dist.get_world_size()),
Expand All @@ -292,6 +309,8 @@ def load_model(
or getattr(model, "_fsdp_shard_conditions", None),
pin_cpu_memory=server_args.pin_cpu_memory,
)
else:
model = model.to("cpu")
else:
model = model.to(local_torch_device)
# We only enable strict check for non-quantized models
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -24,6 +24,7 @@
except ImportError:
HAS_RUNAI_MODEL_STREAMER = False

from sglang.multimodal_gen import envs
from sglang.multimodal_gen.runtime.distributed import get_local_torch_device
from sglang.multimodal_gen.runtime.utils.logging_utils import init_logger

Expand Down Expand Up @@ -135,13 +136,17 @@ def _validate_safetensors_file(file_path: str) -> bool:
def safetensors_weights_iterator(
hf_weights_files: list[str],
to_cpu: bool = True,
use_runai_model_streamer: bool = HAS_RUNAI_MODEL_STREAMER,
use_runai_model_streamer: bool | None = None,
) -> Generator[tuple[str, torch.Tensor], None, None]:
"""Iterate over the weights in the model safetensor files."""
enable_tqdm = (
not torch.distributed.is_initialized() or torch.distributed.get_rank() == 0
)
device = "cpu" if to_cpu else str(get_local_torch_device())
if use_runai_model_streamer is None:
use_runai_model_streamer = (
HAS_RUNAI_MODEL_STREAMER and envs.SGLANG_USE_RUNAI_MODEL_STREAMER
)

# Validate files before loading
corrupted_files = [
Expand Down
214 changes: 214 additions & 0 deletions python/sglang/multimodal_gen/test/server/accuracy_config.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,214 @@
from __future__ import annotations

from dataclasses import dataclass
from enum import Enum
from typing import Dict, Optional

from sglang.multimodal_gen.test.server.testcase_configs import DiffusionTestCase


class ComponentType(str, Enum):
VAE = "vae"
TRANSFORMER = "transformer"
TEXT_ENCODER = "text_encoder"


@dataclass(frozen=True)
class ComponentSkip:
reason: str


DEFAULT_TIMESTEP = 500.0
TIMESTEP_NORMALIZATION_FACTOR = 1000.0
I2V_IMAGE_DIM = 1280
I2V_TEXT_ENCODER_DIM = 5120

DEFAULT_TEXT_ENCODER_VOCAB_SIZE = 32000
TEXT_ENCODER_INPUT_SEED = 42
TEXT_ENCODER_TOKEN_MIN = 100
TEXT_ENCODER_TOKEN_MAX = 30000
TEXT_ENCODER_TOKEN_LENGTH = 32

# Default thresholds by component. Override per component/case if needed.
DEFAULT_THRESHOLDS = {
ComponentType.VAE: 0.999,
ComponentType.TRANSFORMER: 0.995,
ComponentType.TEXT_ENCODER: 0.98,
}

# Optional per-case overrides: {case_id: {ComponentType: threshold}}
CASE_THRESHOLDS: Dict[str, Dict[ComponentType, float]] = {
# Add overrides here when a specific model/component needs a different threshold.
"flux_2_image_t2i": {ComponentType.TRANSFORMER: 0.99},
"flux_2_image_t2i_layerwise_offload": {ComponentType.TRANSFORMER: 0.99},
"flux_2_image_t2i_2_gpus": {ComponentType.TRANSFORMER: 0.99},
"flux_2_klein_ti2i_2_gpus": {ComponentType.TRANSFORMER: 0.975},
"flux_2_ti2i": {ComponentType.TRANSFORMER: 0.99},
"flux_2_t2i_customized_vae_path": {ComponentType.TRANSFORMER: 0.99},
"fast_hunyuan_video": {ComponentType.TRANSFORMER: 0.99},
"fsdp-inference": {ComponentType.TRANSFORMER: 0.9935},
"wan2_2_i2v_a14b_2gpu": {ComponentType.TRANSFORMER: 0.99},
"wan2_2_t2v_a14b_2gpu": {ComponentType.TRANSFORMER: 0.99},
"wan2_2_t2v_a14b_teacache_2gpu": {ComponentType.TRANSFORMER: 0.99},
"wan2_2_t2v_a14b_lora_2gpu": {ComponentType.TRANSFORMER: 0.99},
"zimage_image_t2i_2_gpus": {ComponentType.TRANSFORMER: 0.9935},
"zimage_image_t2i_2_gpus_non_square": {ComponentType.TRANSFORMER: 0.9935},
}

# Active skip policy. Keep this limited to cases with current, concrete evidence
# of real divergence or unsupported reference loading in the harness.
SKIP_COMPONENTS: Dict[str, Dict[ComponentType, ComponentSkip]] = {
"flux_image_t2i": {
ComponentType.TEXT_ENCODER: ComponentSkip(
"Text encoder diverges from HF baseline despite 100% matched weights (CosSim ~0.47)"
)
},
"sana_image_t2i": {
ComponentType.VAE: ComponentSkip(
"HF AutoencoderDC checkpoint leaves required to_qkv_multiscale weights missing, so VAE transfer would compare against partially initialized reference weights"
)
},
"mova_360p_1gpu": {
ComponentType.TRANSFORMER: ComponentSkip(
"HF reference transformer cannot be materialized from the video_dit repo layout"
)
},
"flux_2_t2i_customized_vae_path": {
ComponentType.VAE: ComponentSkip(
"Customized VAE override points to FLUX.2 Tiny AutoEncoder, but the HF reference loader does not yet materialize a trustworthy matching VAE baseline"
)
},
"wan2_2_ti2v_5b": {
ComponentType.TRANSFORMER: ComponentSkip(
"SGLang transformer loader rejects new parameters in HF checkpoint"
)
},
"fastwan2_2_ti2v_5b": {
ComponentType.TRANSFORMER: ComponentSkip(
"SGLang transformer loader rejects new parameters in HF checkpoint"
)
},
"turbo_wan2_1_t2v_1.3b": {
ComponentType.TRANSFORMER: ComponentSkip(
"Weight transfer match ratio too low for reliable comparison"
)
},
"wan2_1_i2v_14b_480P_2gpu": {
ComponentType.TRANSFORMER: ComponentSkip(
"Transformer diverges from Diffusers baseline in 2-GPU accuracy run (CosSim ~0.71) after full weight transfer and matching output shape"
),
ComponentType.TEXT_ENCODER: ComponentSkip(
"Text encoder diverges from HF baseline in 2-GPU SP-folded accuracy run (CosSim ~0.31) after 100% matched weight transfer"
),
},
"wan2_1_i2v_14b_lora_2gpu": {
ComponentType.TRANSFORMER: ComponentSkip(
"Transformer diverges from Diffusers baseline in 2-GPU accuracy run (CosSim ~0.68) after full weight transfer and matching output shape"
),
ComponentType.TEXT_ENCODER: ComponentSkip(
"Text encoder diverges from HF baseline in 2-GPU SP-folded accuracy run (CosSim ~0.31) after 100% matched weight transfer"
),
},
"wan2_1_i2v_14b_720P_2gpu": {
ComponentType.TRANSFORMER: ComponentSkip(
"Transformer diverges from Diffusers baseline in 2-GPU accuracy run (CosSim ~0.68) after full weight transfer and matching output shape"
),
ComponentType.TEXT_ENCODER: ComponentSkip(
"Text encoder diverges from HF baseline in 2-GPU SP-folded accuracy run (CosSim ~0.31) after 100% matched weight transfer"
),
},
"wan2_2_i2v_a14b_2gpu": {
ComponentType.TEXT_ENCODER: ComponentSkip(
"Text encoder diverges from HF baseline in 2-GPU SP-folded accuracy run (CosSim ~0.31) after 100% matched weight transfer"
)
},
"wan2_2_t2v_a14b_2gpu": {
ComponentType.TEXT_ENCODER: ComponentSkip(
"Text encoder diverges from HF baseline in 2-GPU SP-folded accuracy run (CosSim ~0.31) after 100% matched weight transfer"
)
},
"wan2_2_t2v_a14b_teacache_2gpu": {
ComponentType.TEXT_ENCODER: ComponentSkip(
"Text encoder diverges from HF baseline in 2-GPU SP-folded accuracy run (CosSim ~0.31) after 100% matched weight transfer"
)
},
"wan2_2_t2v_a14b_lora_2gpu": {
ComponentType.TEXT_ENCODER: ComponentSkip(
"Text encoder diverges from HF baseline in 2-GPU SP-folded accuracy run (CosSim ~0.31) after 100% matched weight transfer"
)
},
"wan2_1_t2v_14b_2gpu": {
ComponentType.TEXT_ENCODER: ComponentSkip(
"Text encoder diverges from HF baseline in 2-GPU SP-folded accuracy run (CosSim ~0.31) after 100% matched weight transfer"
)
},
"wan2_1_t2v_1.3b_cfg_parallel": {
ComponentType.TEXT_ENCODER: ComponentSkip(
"Text encoder diverges from HF baseline in 2-GPU accuracy run (CosSim ~0.31) after 100% matched weight transfer"
)
},
"mova_360p_tp2": {
ComponentType.TRANSFORMER: ComponentSkip(
"HF reference transformer cannot be materialized from the MOVA video_dit repo layout"
),
ComponentType.TEXT_ENCODER: ComponentSkip(
"Text encoder diverges from HF baseline in 2-GPU accuracy run (CosSim ~0.31) after 100% matched weight transfer"
),
},
"mova_360p_ring1_uly2": {
ComponentType.TRANSFORMER: ComponentSkip(
"HF reference transformer cannot be materialized from the MOVA video_dit repo layout"
),
ComponentType.TEXT_ENCODER: ComponentSkip(
"Text encoder diverges from HF baseline in 2-GPU accuracy run (CosSim ~0.31) after 100% matched weight transfer"
),
},
"mova_360p_ring2_uly1": {
ComponentType.TRANSFORMER: ComponentSkip(
"HF reference transformer cannot be materialized from the MOVA video_dit repo layout"
),
ComponentType.TEXT_ENCODER: ComponentSkip(
"Text encoder diverges from HF baseline in 2-GPU accuracy run (CosSim ~0.31) after 100% matched weight transfer"
),
},
"flux_image_t2i_2_gpus": {
ComponentType.TEXT_ENCODER: ComponentSkip(
"Text encoder diverges from HF baseline in 2-GPU accuracy run (CosSim ~0.47) after 100% matched weight transfer"
)
},
"flux_2_image_t2i_2_gpus": {
ComponentType.TRANSFORMER: ComponentSkip(
"2-GPU FLUX.2 transformer diverges strongly from Diffusers baseline (CosSim ~0.54) despite full weight transfer"
)
},
"hunyuan3d_shape_gen": {
ComponentType.VAE: ComponentSkip(
"HF config cannot be parsed as valid JSON for component reference loading"
),
ComponentType.TRANSFORMER: ComponentSkip(
"HF config cannot be parsed as valid JSON for component reference loading"
),
ComponentType.TEXT_ENCODER: ComponentSkip(
"HF config cannot be parsed as valid JSON for component reference loading"
),
},
}

# TODO: If a model needs extra compatibility logic, prefer adding a skip or an
# explicit override here instead of adding more ad-hoc hacks in the engine.


def get_threshold(case_id: str, component: ComponentType) -> float:
overrides = CASE_THRESHOLDS.get(case_id, {})
return overrides.get(component, DEFAULT_THRESHOLDS[component])


def get_skip_reason(case: DiffusionTestCase, component: ComponentType) -> Optional[str]:
skip_entry = SKIP_COMPONENTS.get(case.id, {}).get(component)
if skip_entry is None:
return None
return skip_entry.reason


def should_skip_component(case: DiffusionTestCase, component: ComponentType) -> bool:
return get_skip_reason(case, component) is not None
Loading
Loading