[diffusion][CI]: route multimodal component accuracy through run_suite (#21960)
This commit is contained in:
@@ -56,6 +56,14 @@ SUITES = {
|
||||
"test_server_2_gpu_b.py",
|
||||
# add new 2-gpu test files here
|
||||
],
|
||||
"component-accuracy-1-gpu": [
|
||||
"test_accuracy_1_gpu_a.py",
|
||||
"test_accuracy_1_gpu_b.py",
|
||||
],
|
||||
"component-accuracy-2-gpu": [
|
||||
"test_accuracy_2_gpu_a.py",
|
||||
"test_accuracy_2_gpu_b.py",
|
||||
],
|
||||
"1-gpu-b200": [
|
||||
"test_server_c.py",
|
||||
],
|
||||
@@ -78,6 +86,10 @@ suites_ascend = {
|
||||
|
||||
SUITES.update(suites_ascend)
|
||||
STRICT_SUITES = {"unit"}
|
||||
COMPONENT_ACCURACY_SUITES = {
|
||||
"component-accuracy-1-gpu",
|
||||
"component-accuracy-2-gpu",
|
||||
}
|
||||
|
||||
|
||||
def parse_args():
|
||||
@@ -261,6 +273,52 @@ def run_pytest(files, filter_expr=None, exitfirst=False):
|
||||
return returncode
|
||||
|
||||
|
||||
def partition_test_files(files, partition_id, total_partitions):
|
||||
return [
|
||||
file_path
|
||||
for i, file_path in enumerate(files)
|
||||
if i % total_partitions == partition_id
|
||||
]
|
||||
|
||||
|
||||
def run_component_accuracy_files(
|
||||
files, suite: str, filter_expr=None, continue_on_error=False
|
||||
):
|
||||
exit_code = 0
|
||||
for file_path in files:
|
||||
if suite == "component-accuracy-2-gpu":
|
||||
cmd = [
|
||||
sys.executable,
|
||||
"-m",
|
||||
"torch.distributed.run",
|
||||
"--nproc_per_node=2",
|
||||
"-m",
|
||||
"pytest",
|
||||
"-s",
|
||||
"-v",
|
||||
]
|
||||
else:
|
||||
cmd = [sys.executable, "-m", "pytest", "-s", "-v"]
|
||||
|
||||
if filter_expr:
|
||||
cmd.extend(["-k", filter_expr])
|
||||
cmd.append(file_path)
|
||||
|
||||
print(f"Running command: {' '.join(cmd)}")
|
||||
file_exit_code = subprocess.call(cmd)
|
||||
if file_exit_code == 5:
|
||||
print(
|
||||
"No tests collected (exit code 5). This is expected when filters "
|
||||
"deselect all tests in a file. Treating as success."
|
||||
)
|
||||
file_exit_code = 0
|
||||
if file_exit_code != 0 and exit_code == 0:
|
||||
exit_code = file_exit_code
|
||||
if file_exit_code != 0 and not continue_on_error:
|
||||
return file_exit_code
|
||||
return exit_code
|
||||
|
||||
|
||||
def _is_in_ci() -> bool:
|
||||
return os.environ.get("SGLANG_IS_IN_CI", "").lower() in ("1", "true", "yes", "on")
|
||||
|
||||
@@ -314,6 +372,49 @@ def main():
|
||||
print(f"No valid test files found for suite '{args.suite}'.")
|
||||
sys.exit(1 if args.suite in STRICT_SUITES else 0)
|
||||
|
||||
if args.suite in COMPONENT_ACCURACY_SUITES:
|
||||
my_files = partition_test_files(
|
||||
suite_files_abs, args.partition_id, args.total_partitions
|
||||
)
|
||||
partition_info = (
|
||||
f"{args.partition_id + 1}/{args.total_partitions} "
|
||||
f"(0-based id={args.partition_id})"
|
||||
)
|
||||
headers = ["Suite", "Partition"]
|
||||
rows = [[args.suite, partition_info]]
|
||||
msg = tabulate.tabulate(rows, headers=headers, tablefmt="psql") + "\n"
|
||||
msg += f"✅ Enabled {len(my_files)} file(s):\n"
|
||||
for file_path in my_files:
|
||||
msg += f" - {file_path}\n"
|
||||
print(msg, flush=True)
|
||||
print(
|
||||
f"Suite: {args.suite} | Partition: {args.partition_id}/{args.total_partitions}"
|
||||
)
|
||||
print(f"Selected {len(suite_files_abs)} files:")
|
||||
for f in suite_files_abs:
|
||||
print(f" - {os.path.basename(f)}")
|
||||
|
||||
if not my_files:
|
||||
print("No files assigned to this partition. Exiting success.")
|
||||
sys.exit(0)
|
||||
|
||||
print(f"Running {len(my_files)} files in this shard: {', '.join(my_files)}")
|
||||
|
||||
exit_code = run_component_accuracy_files(
|
||||
my_files,
|
||||
suite=args.suite,
|
||||
filter_expr=args.filter,
|
||||
continue_on_error=args.continue_on_error,
|
||||
)
|
||||
|
||||
msg = "\n" + tabulate.tabulate(rows, headers=headers, tablefmt="psql") + "\n"
|
||||
msg += f"✅ Executed {len(my_files)} file(s):\n"
|
||||
for file_path in my_files:
|
||||
msg += f" - {file_path}\n"
|
||||
print(msg, flush=True)
|
||||
|
||||
sys.exit(exit_code)
|
||||
|
||||
# 3. collect all test items and partition by items (not files)
|
||||
all_test_items = collect_test_items(suite_files_abs, filter_expr=args.filter)
|
||||
|
||||
@@ -352,7 +453,11 @@ def main():
|
||||
|
||||
# 4. execute with the specific test items
|
||||
# Fast-fail: stop on first failure unless --continue-on-error is set
|
||||
exit_code = run_pytest(my_items, exitfirst=not args.continue_on_error)
|
||||
exit_code = run_pytest(
|
||||
my_items,
|
||||
filter_expr=args.filter,
|
||||
exitfirst=not args.continue_on_error,
|
||||
)
|
||||
|
||||
# Print tests again at the end for visibility
|
||||
msg = "\n" + tabulate.tabulate(rows, headers=headers, tablefmt="psql") + "\n"
|
||||
|
||||
@@ -73,6 +73,14 @@ SKIP_COMPONENTS: Dict[str, Dict[ComponentType, ComponentSkip]] = {
|
||||
"HF reference transformer cannot be materialized from the video_dit repo layout"
|
||||
)
|
||||
},
|
||||
"ltx_2.3_one_stage_ti2v": {
|
||||
ComponentType.VAE: ComponentSkip(
|
||||
"LTX-2.3 VAE component diverges from the HF reference after local overlay materialization; weight transfer matched 96/176 (54.55%), below the minimum threshold for trustworthy comparison"
|
||||
),
|
||||
ComponentType.TRANSFORMER: ComponentSkip(
|
||||
"LTX-2.3 transformer component does not match the HF reference architecture after local overlay materialization; scale_shift_table parameters load as [9, ...] in the checkpoint but [6, ...] in the reference model"
|
||||
),
|
||||
},
|
||||
"qwen_image_t2i_cache_dit_enabled": {
|
||||
ComponentType.VAE: ComponentSkip(
|
||||
"Representative VAE accuracy is already covered by qwen_image_t2i for the same source component and topology"
|
||||
@@ -353,6 +361,11 @@ SKIP_COMPONENTS: Dict[str, Dict[ComponentType, ComponentSkip]] = {
|
||||
"2-GPU FLUX.2 transformer diverges strongly from Diffusers baseline (CosSim ~0.54) despite full weight transfer"
|
||||
)
|
||||
},
|
||||
"ltx_2_two_stage_t2v": {
|
||||
ComponentType.TRANSFORMER: ComponentSkip(
|
||||
"Transformer output shape mismatch after 100% matched weight transfer: SGL [1, 128, 4, 16, 16] vs Diffusers [1, 1024, 128]"
|
||||
)
|
||||
},
|
||||
"hunyuan3d_shape_gen": {
|
||||
ComponentType.VAE: ComponentSkip(
|
||||
"HF config cannot be parsed as valid JSON for component reference loading"
|
||||
|
||||
@@ -30,6 +30,7 @@ DEFAULT_TEXT_SEQ_LEN = 64
|
||||
DEFAULT_TOKEN_LAYOUT_SIZE = 32
|
||||
REDUCED_TOKEN_LAYOUT_SIZE = 16
|
||||
DEFAULT_VIDEO_FRAME_COUNT = 4
|
||||
DEFAULT_AUDIO_FRAME_COUNT = 16
|
||||
DEFAULT_IMAGE_TOKEN_COUNT = 257
|
||||
ALIAS_ROTARY_TEXT_PAD_MULTIPLE = 32
|
||||
DEFAULT_TRANSFORMER_IN_CHANNELS = 16
|
||||
@@ -210,6 +211,13 @@ def _build_transformer_hook_inputs(
|
||||
|
||||
rng = _DeterministicRNG()
|
||||
layout = _infer_transformer_layout(param_names)
|
||||
requires_audio_stream_inputs = (
|
||||
"audio_hidden_states" in param_names
|
||||
and "audio_encoder_hidden_states" in param_names
|
||||
)
|
||||
requires_audio_video_shape_inputs = requires_audio_stream_inputs and all(
|
||||
key in param_names for key in ("num_frames", "height", "width")
|
||||
)
|
||||
in_channels = _read_config_value(
|
||||
model,
|
||||
[
|
||||
@@ -238,6 +246,16 @@ def _build_transformer_hook_inputs(
|
||||
],
|
||||
default=DEFAULT_TRANSFORMER_TEXT_CHANNELS,
|
||||
)
|
||||
audio_in_channels = _read_config_value(
|
||||
model,
|
||||
[
|
||||
"arch_config.audio_in_channels",
|
||||
"audio_in_channels",
|
||||
"arch_config.audio_out_channels",
|
||||
"audio_out_channels",
|
||||
],
|
||||
default=in_channels,
|
||||
)
|
||||
pooled_channels = _read_config_value(
|
||||
model,
|
||||
[
|
||||
@@ -256,7 +274,21 @@ def _build_transformer_hook_inputs(
|
||||
default=I2V_IMAGE_DIM,
|
||||
)
|
||||
|
||||
if layout == "token_shapes":
|
||||
if requires_audio_video_shape_inputs:
|
||||
patch_size = getattr(model, "patch_size", None)
|
||||
if not (
|
||||
isinstance(patch_size, tuple)
|
||||
and len(patch_size) == 3
|
||||
and all(isinstance(dim, int) and dim > 0 for dim in patch_size)
|
||||
):
|
||||
patch_size = (1, 2, 2)
|
||||
patch_t, patch_h, patch_w = patch_size
|
||||
num_frames = DEFAULT_VIDEO_FRAME_COUNT * patch_t
|
||||
height = REDUCED_TOKEN_LAYOUT_SIZE * patch_h
|
||||
width = REDUCED_TOKEN_LAYOUT_SIZE * patch_w
|
||||
seq_len = (num_frames // patch_t) * (height // patch_h) * (width // patch_w)
|
||||
hidden_states = rng.randn((1, seq_len, in_channels), device, torch.bfloat16)
|
||||
elif layout == "token_shapes":
|
||||
height, width = DEFAULT_TOKEN_LAYOUT_SIZE, DEFAULT_TOKEN_LAYOUT_SIZE
|
||||
seq_len = (height // 2) * (width // 2)
|
||||
hidden_states = rng.randn((1, seq_len, in_channels), device, torch.bfloat16)
|
||||
@@ -307,6 +339,24 @@ def _build_transformer_hook_inputs(
|
||||
"guidance": torch.tensor([1.0], device=device, dtype=torch.bfloat16),
|
||||
}
|
||||
|
||||
if requires_audio_stream_inputs:
|
||||
inputs["audio_hidden_states"] = rng.randn(
|
||||
(1, DEFAULT_AUDIO_FRAME_COUNT, audio_in_channels),
|
||||
device,
|
||||
torch.bfloat16,
|
||||
)
|
||||
inputs["audio_encoder_hidden_states"] = rng.randn(
|
||||
(1, DEFAULT_TEXT_SEQ_LEN, text_channels),
|
||||
device,
|
||||
torch.bfloat16,
|
||||
)
|
||||
inputs["audio_timestep"] = inputs["timestep"].clone()
|
||||
inputs["audio_num_frames"] = DEFAULT_AUDIO_FRAME_COUNT
|
||||
if requires_audio_video_shape_inputs:
|
||||
inputs["num_frames"] = num_frames
|
||||
inputs["height"] = height
|
||||
inputs["width"] = width
|
||||
|
||||
if "pooled_projections" in param_names:
|
||||
inputs["pooled_projections"] = rng.randn(
|
||||
(1, pooled_channels), device, torch.bfloat16
|
||||
@@ -320,6 +370,10 @@ def _build_transformer_hook_inputs(
|
||||
)
|
||||
inputs["encoder_attention_mask"] = attention_mask
|
||||
inputs["encoder_hidden_states_mask"] = attention_mask
|
||||
if "audio_encoder_attention_mask" in param_names:
|
||||
inputs["audio_encoder_attention_mask"] = torch.ones(
|
||||
1, DEFAULT_TEXT_SEQ_LEN, device=device, dtype=torch.bool
|
||||
)
|
||||
if "encoder_hidden_states_image" in param_names and _supports_image_conditioning(
|
||||
model
|
||||
):
|
||||
@@ -471,8 +525,16 @@ def _prepare_transformer_hook_call(
|
||||
"txt_seq_lens",
|
||||
"freqs_cis",
|
||||
"additional_t_cond",
|
||||
"audio_hidden_states",
|
||||
"audio_encoder_hidden_states",
|
||||
"audio_timestep",
|
||||
"encoder_attention_mask",
|
||||
"encoder_hidden_states_mask",
|
||||
"audio_encoder_attention_mask",
|
||||
"num_frames",
|
||||
"height",
|
||||
"width",
|
||||
"audio_num_frames",
|
||||
):
|
||||
if key in param_names and key in inputs:
|
||||
kwargs[key] = inputs[key]
|
||||
@@ -491,6 +553,17 @@ def _prepare_transformer_reference_call(module: nn.Module, inputs: Inputs) -> Ho
|
||||
return _prepare_transformer_hook_call(module, inputs, side="reference")
|
||||
|
||||
|
||||
def _normalize_transformer_reference_output(output: Any) -> torch.Tensor:
|
||||
sample = getattr(output, "sample", None)
|
||||
if (
|
||||
isinstance(sample, (list, tuple))
|
||||
and sample
|
||||
and all(isinstance(item, torch.Tensor) for item in sample)
|
||||
):
|
||||
return torch.stack(list(sample), dim=0)
|
||||
return extract_output_tensor(output)
|
||||
|
||||
|
||||
class _VAEDecodeModule(nn.Module):
|
||||
def __init__(self, vae: nn.Module):
|
||||
super().__init__()
|
||||
@@ -562,6 +635,7 @@ TRANSFORMER_NATIVE_PROFILE = NativeHookProfile(
|
||||
build_inputs=_build_transformer_hook_inputs,
|
||||
prepare_sglang_call=_prepare_transformer_sglang_call,
|
||||
prepare_reference_call=_prepare_transformer_reference_call,
|
||||
normalize_reference_output=_normalize_transformer_reference_output,
|
||||
)
|
||||
|
||||
VAE_NATIVE_PROFILE = NativeHookProfile(
|
||||
|
||||
@@ -2,6 +2,7 @@ from __future__ import annotations
|
||||
|
||||
import json
|
||||
import os
|
||||
from contextlib import nullcontext
|
||||
from dataclasses import dataclass
|
||||
from typing import Any, Dict, List, Optional, Tuple
|
||||
|
||||
@@ -19,8 +20,12 @@ from sglang.multimodal_gen.runtime.distributed.parallel_state import (
|
||||
model_parallel_is_initialized,
|
||||
)
|
||||
from sglang.multimodal_gen.runtime.layers.utils import get_group_rank, get_group_size
|
||||
from sglang.multimodal_gen.runtime.server_args import ServerArgs
|
||||
from sglang.multimodal_gen.runtime.server_args import ServerArgs, get_global_server_args
|
||||
from sglang.multimodal_gen.runtime.utils.hf_diffusers_utils import maybe_download_model
|
||||
from sglang.multimodal_gen.runtime.utils.model_overlay import (
|
||||
load_overlay_manifest_if_present,
|
||||
resolve_model_overlay_target,
|
||||
)
|
||||
from sglang.multimodal_gen.test.server.accuracy_config import (
|
||||
DEFAULT_TEXT_ENCODER_VOCAB_SIZE,
|
||||
I2V_TEXT_ENCODER_DIM,
|
||||
@@ -32,27 +37,6 @@ from sglang.multimodal_gen.test.server.accuracy_config import (
|
||||
get_threshold,
|
||||
)
|
||||
|
||||
STAGED_1GPU_NATIVE_CASE_IDS = {
|
||||
"flux_2_image_t2i",
|
||||
"qwen_image_layered_i2i",
|
||||
"flux_2_image_t2i_upscaling_4x",
|
||||
"flux_2_ti2i",
|
||||
"flux_2_t2i_customized_vae_path",
|
||||
"flux_2_ti2i_multi_image_cache_dit",
|
||||
}
|
||||
|
||||
# These case allowlists are accuracy-runner policy. They select the few 1-GPU
|
||||
# cases that need sequential SGLang/reference execution to stay within memory
|
||||
# limits during CI and local correctness runs.
|
||||
STAGED_1GPU_TEXT_ENCODER_CASE_IDS = {
|
||||
"flux_2_image_t2i",
|
||||
"flux_2_image_t2i_upscaling_4x",
|
||||
"mova_360p_1gpu",
|
||||
"flux_2_ti2i",
|
||||
"flux_2_t2i_customized_vae_path",
|
||||
"flux_2_ti2i_multi_image_cache_dit",
|
||||
}
|
||||
|
||||
SOURCE_PREFIXES = (
|
||||
"module.",
|
||||
"model.",
|
||||
@@ -81,9 +65,7 @@ class ComponentSelection:
|
||||
base_model_id: str
|
||||
base_model_root: str
|
||||
component_paths: Dict[str, str]
|
||||
source_root: str
|
||||
source_path: str
|
||||
source_subfolder: str
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
@@ -153,7 +135,7 @@ def _resolve_component_subfolder(
|
||||
|
||||
def resolve_component_path(
|
||||
local_root: str, component: ComponentType, model_index_keys: Tuple[str, ...]
|
||||
) -> Tuple[str, str]:
|
||||
) -> str:
|
||||
model_index_path = os.path.join(local_root, "model_index.json")
|
||||
model_index = read_json_file(model_index_path)
|
||||
|
||||
@@ -169,13 +151,13 @@ def resolve_component_path(
|
||||
candidate
|
||||
):
|
||||
continue
|
||||
return candidate, subfolder
|
||||
return candidate
|
||||
|
||||
if has_component_files(local_root):
|
||||
if component != ComponentType.TEXT_ENCODER or is_text_encoder_config(
|
||||
local_root
|
||||
):
|
||||
return local_root, ""
|
||||
return local_root
|
||||
|
||||
raise FileNotFoundError(
|
||||
f"Could not resolve {component.value} from model_index.json under {local_root}"
|
||||
@@ -230,37 +212,48 @@ def select_component_source(
|
||||
model_index_keys: Tuple[str, ...],
|
||||
) -> ComponentSelection:
|
||||
component_paths = extract_component_path_overrides(extra_args)
|
||||
base_model_root = maybe_download_model(model_id)
|
||||
force_diffusers_model = resolve_model_overlay_target(model_id) is not None or (
|
||||
os.path.exists(model_id)
|
||||
and load_overlay_manifest_if_present(model_id) is not None
|
||||
)
|
||||
base_model_root = maybe_download_model(
|
||||
model_id, force_diffusers_model=force_diffusers_model
|
||||
)
|
||||
search_keys = [component.value]
|
||||
for key in model_index_keys:
|
||||
if key not in search_keys:
|
||||
search_keys.append(key)
|
||||
|
||||
source_root = base_model_root
|
||||
component_key = component.value
|
||||
for key in search_keys:
|
||||
override_path = component_paths.get(key)
|
||||
if override_path:
|
||||
source_root = maybe_download_model(override_path)
|
||||
component_key = key
|
||||
break
|
||||
if override_path is None:
|
||||
continue
|
||||
assert has_component_files(override_path), (
|
||||
f"Component override for {component.value} must point directly to a "
|
||||
f"component directory: {override_path}"
|
||||
)
|
||||
if component == ComponentType.TEXT_ENCODER:
|
||||
assert is_text_encoder_config(override_path), (
|
||||
f"Text encoder override must point to a text encoder directory: "
|
||||
f"{override_path}"
|
||||
)
|
||||
return ComponentSelection(
|
||||
base_model_id=model_id,
|
||||
base_model_root=base_model_root,
|
||||
component_paths=component_paths,
|
||||
source_path=override_path,
|
||||
)
|
||||
|
||||
ordered_keys = [component_key]
|
||||
for key in search_keys:
|
||||
if key not in ordered_keys:
|
||||
ordered_keys.append(key)
|
||||
source_path, source_subfolder = resolve_component_path(
|
||||
source_root,
|
||||
source_path = resolve_component_path(
|
||||
base_model_root,
|
||||
component,
|
||||
tuple(ordered_keys),
|
||||
tuple(search_keys),
|
||||
)
|
||||
return ComponentSelection(
|
||||
base_model_id=model_id,
|
||||
base_model_root=base_model_root,
|
||||
component_paths=component_paths,
|
||||
source_root=source_root,
|
||||
source_path=source_path,
|
||||
source_subfolder=source_subfolder,
|
||||
)
|
||||
|
||||
|
||||
@@ -680,16 +673,6 @@ def run_text_encoder_accuracy_pair(
|
||||
)
|
||||
|
||||
|
||||
def _should_stage_case(case: Any, component: ComponentType, num_gpus: int) -> bool:
|
||||
if num_gpus == 2:
|
||||
return True
|
||||
if num_gpus != 1:
|
||||
return False
|
||||
if component == ComponentType.TEXT_ENCODER:
|
||||
return case.id in STAGED_1GPU_TEXT_ENCODER_CASE_IDS
|
||||
return case.id in STAGED_1GPU_NATIVE_CASE_IDS
|
||||
|
||||
|
||||
def _run_single_text_encoder_forward(
|
||||
model: nn.Module, input_ids: torch.Tensor, attention_mask: torch.Tensor
|
||||
) -> torch.Tensor:
|
||||
@@ -724,26 +707,55 @@ def _run_staged_native_component_accuracy_case(
|
||||
component,
|
||||
library,
|
||||
num_gpus,
|
||||
materialize_sgl_on_device=(component != ComponentType.TRANSFORMER),
|
||||
materialize_ref_on_device=False,
|
||||
)
|
||||
if component == ComponentType.TRANSFORMER:
|
||||
sgl = sgl.to(device=device, dtype=torch.bfloat16).eval()
|
||||
profile = resolve_component_native_profile(component)
|
||||
inputs = profile.build_inputs(case, sgl, device, ref)
|
||||
runtime_server_args = get_global_server_args()
|
||||
use_transformer_autocast = (
|
||||
component == ComponentType.TRANSFORMER
|
||||
and not runtime_server_args.disable_autocast
|
||||
and torch.device(device).type != "cpu"
|
||||
)
|
||||
|
||||
sgl_call = profile.prepare_sglang_call(sgl, inputs)
|
||||
with torch.no_grad():
|
||||
sgl_autocast = (
|
||||
torch.autocast(
|
||||
device_type=torch.device(device).type,
|
||||
dtype=torch.bfloat16,
|
||||
enabled=True,
|
||||
)
|
||||
if use_transformer_autocast
|
||||
else nullcontext()
|
||||
)
|
||||
with torch.no_grad(), sgl_autocast:
|
||||
sgl_raw = engine_cls._execute_with_native_hook(sgl_call)
|
||||
sgl_out = profile.normalize_sglang_output(sgl_raw)
|
||||
sgl_out = engine_cls._apply_output_transforms(sgl_out, sgl_call).detach().cpu()
|
||||
|
||||
del sgl_call
|
||||
del sgl_raw
|
||||
if component == ComponentType.TRANSFORMER and num_gpus == 1:
|
||||
engine_cls.prepare_component_for_release(sgl)
|
||||
del sgl
|
||||
sgl = None
|
||||
engine_cls.clear_memory()
|
||||
|
||||
ref = ref.to(device=device, dtype=torch.bfloat16).eval()
|
||||
ref_call = profile.prepare_reference_call(ref, inputs)
|
||||
with torch.no_grad():
|
||||
ref_autocast = (
|
||||
torch.autocast(
|
||||
device_type=torch.device(device).type,
|
||||
dtype=torch.bfloat16,
|
||||
enabled=True,
|
||||
)
|
||||
if use_transformer_autocast
|
||||
else nullcontext()
|
||||
)
|
||||
with torch.no_grad(), ref_autocast:
|
||||
ref_raw = engine_cls._execute_with_native_hook(ref_call)
|
||||
ref_out = profile.normalize_reference_output(ref_raw)
|
||||
ref_out = engine_cls._apply_output_transforms(ref_out, ref_call).detach().cpu()
|
||||
@@ -758,6 +770,8 @@ def _run_staged_native_component_accuracy_case(
|
||||
)
|
||||
finally:
|
||||
if sgl is not None:
|
||||
if component == ComponentType.TRANSFORMER and num_gpus == 1:
|
||||
engine_cls.prepare_component_for_release(sgl)
|
||||
del sgl
|
||||
if ref is not None:
|
||||
del ref
|
||||
@@ -823,58 +837,10 @@ def run_native_component_accuracy_case(
|
||||
library: str,
|
||||
num_gpus: int,
|
||||
) -> None:
|
||||
if _should_stage_case(case, component, num_gpus):
|
||||
_run_staged_native_component_accuracy_case(
|
||||
engine_cls, case, component, library, num_gpus
|
||||
)
|
||||
return
|
||||
engine_cls.clear_memory()
|
||||
sgl = None
|
||||
ref = None
|
||||
try:
|
||||
sgl, ref, device = engine_cls.load_component_pair(
|
||||
case, component, library, num_gpus
|
||||
)
|
||||
sgl_out, ref_out = engine_cls.run_component_pair_native(
|
||||
case, component, sgl, ref, device
|
||||
)
|
||||
engine_cls.check_accuracy(
|
||||
sgl_out,
|
||||
ref_out,
|
||||
f"{case.id}_{component.value}",
|
||||
get_threshold(case.id, component),
|
||||
)
|
||||
finally:
|
||||
if sgl is not None:
|
||||
del sgl
|
||||
if ref is not None:
|
||||
del ref
|
||||
engine_cls.reset_parallel_runtime()
|
||||
engine_cls.clear_memory()
|
||||
_run_staged_native_component_accuracy_case(
|
||||
engine_cls, case, component, library, num_gpus
|
||||
)
|
||||
|
||||
|
||||
def run_text_encoder_accuracy_case(engine_cls: Any, case: Any, num_gpus: int) -> None:
|
||||
if _should_stage_case(case, ComponentType.TEXT_ENCODER, num_gpus):
|
||||
_run_staged_text_encoder_accuracy_case(engine_cls, case, num_gpus)
|
||||
return
|
||||
engine_cls.clear_memory()
|
||||
sgl = None
|
||||
ref = None
|
||||
try:
|
||||
sgl, ref, _device = engine_cls.load_component_pair(
|
||||
case, ComponentType.TEXT_ENCODER, "transformers", num_gpus
|
||||
)
|
||||
sgl_out, ref_out = run_text_encoder_accuracy_pair(sgl, ref)
|
||||
engine_cls.check_accuracy(
|
||||
sgl_out,
|
||||
ref_out,
|
||||
f"{case.id}_encoder",
|
||||
get_threshold(case.id, ComponentType.TEXT_ENCODER),
|
||||
)
|
||||
finally:
|
||||
if sgl is not None:
|
||||
del sgl
|
||||
if ref is not None:
|
||||
del ref
|
||||
engine_cls.reset_parallel_runtime()
|
||||
engine_cls.clear_memory()
|
||||
_run_staged_text_encoder_accuracy_case(engine_cls, case, num_gpus)
|
||||
|
||||
@@ -26,6 +26,7 @@ except ImportError:
|
||||
|
||||
import sglang.multimodal_gen.runtime.managers.forward_context as fc_mod
|
||||
from sglang.multimodal_gen.runtime.distributed.parallel_state import (
|
||||
cleanup_dist_env_and_memory,
|
||||
destroy_model_parallel,
|
||||
get_local_torch_device,
|
||||
get_tensor_model_parallel_rank,
|
||||
@@ -149,13 +150,34 @@ def _load_wan_reference_vae(comp_path: str, pipeline_config) -> nn.Module:
|
||||
return vae
|
||||
|
||||
|
||||
def _load_reference_component_from_local_safetensors(
|
||||
component_cls: type[nn.Module],
|
||||
comp_path: str,
|
||||
component_name: str,
|
||||
) -> nn.Module:
|
||||
config = component_cls.load_config(comp_path)
|
||||
component = component_cls.from_config(config)
|
||||
missing_keys, unexpected_keys = load_checkpoint_weights(component, comp_path)
|
||||
if missing_keys:
|
||||
logger.warning(
|
||||
"Reference %s missing keys from local safetensors: %s",
|
||||
component_name,
|
||||
missing_keys,
|
||||
)
|
||||
if unexpected_keys:
|
||||
logger.warning(
|
||||
"Reference %s unexpected keys from local safetensors: %s",
|
||||
component_name,
|
||||
unexpected_keys,
|
||||
)
|
||||
return component
|
||||
|
||||
|
||||
def _load_reference_component(
|
||||
comp_path: str,
|
||||
source_root: str,
|
||||
component: ComponentType,
|
||||
hub_id: str,
|
||||
pipeline_config,
|
||||
subfolder: str,
|
||||
) -> nn.Module:
|
||||
# WAN VAE does not have a clean generic diffusers auto-load path here, and we
|
||||
# explicitly need checkpoint-loaded weights for reference-side transfer/parity.
|
||||
@@ -168,9 +190,14 @@ def _load_reference_component(
|
||||
cls = getattr(diffusers, str(class_name), None) if class_name else None
|
||||
if cls is None:
|
||||
cls = diffusers.AutoencoderKL
|
||||
if cls is not diffusers.AutoencoderKL and os.path.exists(
|
||||
os.path.join(comp_path, "model.safetensors")
|
||||
):
|
||||
return _load_reference_component_from_local_safetensors(
|
||||
cls, comp_path, component.value
|
||||
)
|
||||
return cls.from_pretrained(
|
||||
source_root,
|
||||
subfolder=subfolder,
|
||||
comp_path,
|
||||
torch_dtype=torch.bfloat16,
|
||||
trust_remote_code=True,
|
||||
)
|
||||
@@ -182,6 +209,14 @@ def _load_reference_component(
|
||||
"torch_dtype": torch.bfloat16,
|
||||
"trust_remote_code": True,
|
||||
}
|
||||
if class_name:
|
||||
maybe_cls = getattr(diffusers, str(class_name), None)
|
||||
if maybe_cls is not None and os.path.exists(
|
||||
os.path.join(comp_path, "model.safetensors")
|
||||
):
|
||||
return _load_reference_component_from_local_safetensors(
|
||||
maybe_cls, comp_path, component.value
|
||||
)
|
||||
if cfg:
|
||||
for k, out_k in [
|
||||
("in_dim", "in_channels"),
|
||||
@@ -192,10 +227,8 @@ def _load_reference_component(
|
||||
if k in cfg:
|
||||
load_kwargs[out_k] = cfg[k]
|
||||
candidates = [diffusers.AutoModel]
|
||||
if class_name:
|
||||
maybe_cls = getattr(diffusers, str(class_name), None)
|
||||
if maybe_cls is not None:
|
||||
candidates.insert(0, maybe_cls)
|
||||
if class_name and maybe_cls is not None:
|
||||
candidates.insert(0, maybe_cls)
|
||||
last_error: Optional[Exception] = None
|
||||
for cls in candidates:
|
||||
try:
|
||||
@@ -241,11 +274,30 @@ def _load_reference_component(
|
||||
|
||||
# Public accuracy engine
|
||||
class AccuracyEngine:
|
||||
@staticmethod
|
||||
def prepare_component_for_release(module: nn.Module) -> None:
|
||||
for submodule in module.modules():
|
||||
reset_teacache_state = getattr(submodule, "reset_teacache_state", None)
|
||||
if callable(reset_teacache_state):
|
||||
reset_teacache_state()
|
||||
|
||||
seen_names: set[str] = set()
|
||||
for cls in type(submodule).__mro__:
|
||||
for name, attr in cls.__dict__.items():
|
||||
if name in seen_names:
|
||||
continue
|
||||
seen_names.add(name)
|
||||
|
||||
cache_clear = getattr(attr, "cache_clear", None)
|
||||
if callable(cache_clear):
|
||||
cache_clear()
|
||||
|
||||
@staticmethod
|
||||
def reset_parallel_runtime() -> None:
|
||||
if torch.distributed.is_initialized():
|
||||
torch.distributed.barrier()
|
||||
if model_parallel_is_initialized():
|
||||
if torch.distributed.get_world_size() == 1:
|
||||
cleanup_dist_env_and_memory()
|
||||
elif model_parallel_is_initialized():
|
||||
destroy_model_parallel()
|
||||
gc.collect()
|
||||
if torch.cuda.is_available():
|
||||
@@ -476,6 +528,8 @@ class AccuracyEngine:
|
||||
num_gpus,
|
||||
component_selection.component_paths,
|
||||
)
|
||||
if component == ComponentType.TRANSFORMER and not materialize_sgl_on_device:
|
||||
sgl_args.dit_cpu_offload = True
|
||||
initialize_parallel_runtime(sgl_args)
|
||||
set_global_server_args(sgl_args)
|
||||
|
||||
@@ -497,11 +551,9 @@ class AccuracyEngine:
|
||||
|
||||
ref_component = _load_reference_component(
|
||||
component_selection.source_path,
|
||||
component_selection.source_root,
|
||||
component,
|
||||
hub_id,
|
||||
sgl_args.pipeline_config,
|
||||
component_selection.source_subfolder,
|
||||
)
|
||||
if materialize_ref_on_device:
|
||||
ref_component = ref_component.to(device=device, dtype=torch.bfloat16)
|
||||
|
||||
@@ -179,6 +179,9 @@ class ServerContext:
|
||||
# Clean up downloaded models if HF cache is not persistent
|
||||
# This prevents disk exhaustion in CI when cache is not mounted
|
||||
self._cleanup_hf_cache_if_not_persistent()
|
||||
else:
|
||||
# Give the runtime a brief cooldown after server shutdown.
|
||||
time.sleep(2)
|
||||
|
||||
def _cleanup_hf_cache_if_not_persistent(self) -> None:
|
||||
"""Clean up HF cache if it's not on a persistent volume.
|
||||
|
||||
Reference in New Issue
Block a user