From cf5ad12612b01a88ed14e15191433232fdc60cfc Mon Sep 17 00:00:00 2001 From: Ratish P <114130421+Ratish1@users.noreply.github.com> Date: Fri, 10 Apr 2026 20:36:03 +0530 Subject: [PATCH] [diffusion][CI]: route multimodal component accuracy through run_suite (#21960) --- .github/workflows/pr-test-multimodal-gen.yml | 110 +++++++++++ .github/workflows/pr-test.yml | 2 + .../sglang/multimodal_gen/test/run_suite.py | 107 ++++++++++- .../test/server/accuracy_config.py | 13 ++ .../test/server/accuracy_hooks.py | 76 +++++++- .../test/server/accuracy_utils.py | 180 +++++++----------- .../test/server/component_accuracy.py | 76 ++++++-- .../test/server/test_server_utils.py | 3 + scripts/ci/utils/slash_command_handler.py | 2 + 9 files changed, 448 insertions(+), 121 deletions(-) diff --git a/.github/workflows/pr-test-multimodal-gen.yml b/.github/workflows/pr-test-multimodal-gen.yml index 1fd8ed24e..d92b3a141 100644 --- a/.github/workflows/pr-test-multimodal-gen.yml +++ b/.github/workflows/pr-test-multimodal-gen.yml @@ -159,6 +159,116 @@ jobs: with: artifact-suffix: ${{ matrix.part }} + multimodal-gen-component-accuracy-1-gpu: + if: | + (inputs.target_stage == 'multimodal-gen-component-accuracy-1-gpu') || + ( + !inputs.target_stage && + ((github.event_name == 'schedule' || inputs.test_parallel_dispatch == 'true') || (inputs.caller_needs_failure != 'true' && !cancelled())) && + inputs.multimodal_gen == 'true' + ) + runs-on: 1-gpu-h100 + timeout-minutes: 240 + strategy: + fail-fast: false + matrix: + part: [0, 1] + steps: + - name: Checkout code + uses: actions/checkout@v4 + with: + ref: ${{ inputs.pr_head_sha || inputs.git_ref || github.sha }} + + - uses: ./.github/actions/check-stage-health + + - uses: ./.github/actions/check-maintenance + + - name: Download artifacts + if: inputs.sgl_kernel == 'true' + uses: actions/download-artifact@v4 + with: + path: sgl-kernel/dist/ + merge-multiple: true + pattern: wheel-python3.10-cuda12.9 + + - name: Install dependencies + timeout-minutes: 20 + run: | + CUSTOM_BUILD_SGL_KERNEL=${{inputs.sgl_kernel}} bash scripts/ci/cuda/ci_install_dependency.sh diffusion + + - name: Run diffusion component accuracy tests (1-GPU) + timeout-minutes: 240 + env: + RUNAI_STREAMER_MEMORY_LIMIT: 0 + CONTINUE_ON_ERROR_FLAG: ${{ inputs.continue_on_error == 'true' && '--continue-on-error' || '' }} + run: | + cd python + python3 sglang/multimodal_gen/test/run_suite.py \ + --suite component-accuracy-1-gpu \ + --partition-id ${{ matrix.part }} \ + --total-partitions 2 \ + $CONTINUE_ON_ERROR_FLAG + + - uses: ./.github/actions/upload-cuda-coredumps + if: always() + with: + artifact-suffix: ${{ matrix.part }} + + multimodal-gen-component-accuracy-2-gpu: + if: | + (inputs.target_stage == 'multimodal-gen-component-accuracy-2-gpu') || + ( + !inputs.target_stage && + ((github.event_name == 'schedule' || inputs.test_parallel_dispatch == 'true') || (inputs.caller_needs_failure != 'true' && !cancelled())) && + inputs.multimodal_gen == 'true' + ) + runs-on: 2-gpu-h100 + timeout-minutes: 240 + strategy: + fail-fast: false + matrix: + part: [0, 1] + steps: + - name: Checkout code + uses: actions/checkout@v4 + with: + ref: ${{ inputs.pr_head_sha || inputs.git_ref || github.sha }} + + - uses: ./.github/actions/check-stage-health + + - uses: ./.github/actions/check-maintenance + + - name: Download artifacts + if: inputs.sgl_kernel == 'true' + uses: actions/download-artifact@v4 + with: + path: sgl-kernel/dist/ + merge-multiple: true + pattern: wheel-python3.10-cuda12.9 + + - name: Install dependencies + timeout-minutes: 20 + run: | + CUSTOM_BUILD_SGL_KERNEL=${{inputs.sgl_kernel}} bash scripts/ci/cuda/ci_install_dependency.sh diffusion + + - name: Run diffusion component accuracy tests (2-GPU) + timeout-minutes: 240 + env: + RUNAI_STREAMER_MEMORY_LIMIT: 0 + CONTINUE_ON_ERROR_FLAG: ${{ inputs.continue_on_error == 'true' && '--continue-on-error' || '' }} + run: | + cd python + python3 sglang/multimodal_gen/test/run_suite.py \ + --suite component-accuracy-2-gpu \ + --partition-id ${{ matrix.part }} \ + --total-partitions 2 \ + $CONTINUE_ON_ERROR_FLAG + + - uses: ./.github/actions/upload-cuda-coredumps + if: always() + with: + artifact-suffix: ${{ matrix.part }} + multimodal-gen-test-1-b200: if: | (inputs.target_stage == 'multimodal-gen-test-1-b200') || diff --git a/.github/workflows/pr-test.yml b/.github/workflows/pr-test.yml index b4cc681a3..8feb364bd 100644 --- a/.github/workflows/pr-test.yml +++ b/.github/workflows/pr-test.yml @@ -891,6 +891,8 @@ jobs: ( inputs.target_stage == 'multimodal-gen-test-1-gpu' || inputs.target_stage == 'multimodal-gen-test-2-gpu' || + inputs.target_stage == 'multimodal-gen-component-accuracy-1-gpu' || + inputs.target_stage == 'multimodal-gen-component-accuracy-2-gpu' || inputs.target_stage == 'multimodal-gen-test-1-b200' || inputs.target_stage == 'multimodal-gen-unit-test' || ( diff --git a/python/sglang/multimodal_gen/test/run_suite.py b/python/sglang/multimodal_gen/test/run_suite.py index a6fef42e8..553224def 100644 --- a/python/sglang/multimodal_gen/test/run_suite.py +++ b/python/sglang/multimodal_gen/test/run_suite.py @@ -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" diff --git a/python/sglang/multimodal_gen/test/server/accuracy_config.py b/python/sglang/multimodal_gen/test/server/accuracy_config.py index d125e289d..1ec0b487b 100644 --- a/python/sglang/multimodal_gen/test/server/accuracy_config.py +++ b/python/sglang/multimodal_gen/test/server/accuracy_config.py @@ -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" diff --git a/python/sglang/multimodal_gen/test/server/accuracy_hooks.py b/python/sglang/multimodal_gen/test/server/accuracy_hooks.py index 0ba46437d..a495904d8 100644 --- a/python/sglang/multimodal_gen/test/server/accuracy_hooks.py +++ b/python/sglang/multimodal_gen/test/server/accuracy_hooks.py @@ -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( diff --git a/python/sglang/multimodal_gen/test/server/accuracy_utils.py b/python/sglang/multimodal_gen/test/server/accuracy_utils.py index 48ecace6f..943fb4e34 100644 --- a/python/sglang/multimodal_gen/test/server/accuracy_utils.py +++ b/python/sglang/multimodal_gen/test/server/accuracy_utils.py @@ -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) diff --git a/python/sglang/multimodal_gen/test/server/component_accuracy.py b/python/sglang/multimodal_gen/test/server/component_accuracy.py index ba59063c0..1e97beb76 100644 --- a/python/sglang/multimodal_gen/test/server/component_accuracy.py +++ b/python/sglang/multimodal_gen/test/server/component_accuracy.py @@ -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) diff --git a/python/sglang/multimodal_gen/test/server/test_server_utils.py b/python/sglang/multimodal_gen/test/server/test_server_utils.py index 7525a6a38..8b28b9cf0 100644 --- a/python/sglang/multimodal_gen/test/server/test_server_utils.py +++ b/python/sglang/multimodal_gen/test/server/test_server_utils.py @@ -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. diff --git a/scripts/ci/utils/slash_command_handler.py b/scripts/ci/utils/slash_command_handler.py index 818b001d9..b075aac21 100644 --- a/scripts/ci/utils/slash_command_handler.py +++ b/scripts/ci/utils/slash_command_handler.py @@ -271,6 +271,8 @@ def handle_rerun_stage( "stage-c-test-deepep-8-gpu-h200", "multimodal-gen-test-1-gpu", "multimodal-gen-test-2-gpu", + "multimodal-gen-component-accuracy-1-gpu", + "multimodal-gen-component-accuracy-2-gpu", "multimodal-gen-test-1-b200", ]