From ef2d4013d77cf1267e032d7a7a745dc3a6f74880 Mon Sep 17 00:00:00 2001 From: Prozac614 Date: Tue, 7 Apr 2026 09:50:23 +0800 Subject: [PATCH] [diffusion] CI: add consistency test (#15236) Co-authored-by: daiweitao --- .github/workflows/pr-test-amd-rocm720.yml | 2 + .github/workflows/pr-test-amd.yml | 2 + .../runtime/pipelines_core/lora_pipeline.py | 15 +- .../multimodal_gen/runtime/server_args.py | 7 + .../runtime/utils/hf_diffusers_utils.py | 23 +- .../server/ascend/testcase_configs_npu.py | 5 + .../test/server/consistency_threshold.json | 205 +++++++ .../test/server/test_server_common.py | 151 ++++++ .../test/server/testcase_configs.py | 8 + .../test/test_consistency_metrics.py | 115 ++++ .../sglang/multimodal_gen/test/test_utils.py | 502 +++++++++++++++++- 11 files changed, 1026 insertions(+), 9 deletions(-) create mode 100644 python/sglang/multimodal_gen/test/server/consistency_threshold.json create mode 100644 python/sglang/multimodal_gen/test/test_consistency_metrics.py diff --git a/.github/workflows/pr-test-amd-rocm720.yml b/.github/workflows/pr-test-amd-rocm720.yml index 24fb80ed3..ce6233110 100644 --- a/.github/workflows/pr-test-amd-rocm720.yml +++ b/.github/workflows/pr-test-amd-rocm720.yml @@ -633,6 +633,7 @@ jobs: -e SGLANG_NON_DENOISE_STAGE_TIME_TOLERANCE=0.6 \ -e SGLANG_DENOISE_STEP_TOLERANCE=0.6 \ -e SGLANG_DENOISE_AGG_TOLERANCE=0.3 \ + -e SGLANG_SKIP_CONSISTENCY=1 \ -e SGLANG_TEST_NUM_INFERENCE_STEPS=5 \ -e AITER_JIT_DIR=/sgl-data/aiter-kernels \ -e MIOPEN_USER_DB_PATH=/sgl-data/miopen-cache \ @@ -762,6 +763,7 @@ jobs: -e SGLANG_NON_DENOISE_STAGE_TIME_TOLERANCE=0.6 \ -e SGLANG_DENOISE_STEP_TOLERANCE=0.6 \ -e SGLANG_DENOISE_AGG_TOLERANCE=0.3 \ + -e SGLANG_SKIP_CONSISTENCY=1 \ -e SGLANG_TEST_NUM_INFERENCE_STEPS=5 \ -e AITER_JIT_DIR=/sgl-data/aiter-kernels \ -e MIOPEN_USER_DB_PATH=/sgl-data/miopen-cache \ diff --git a/.github/workflows/pr-test-amd.yml b/.github/workflows/pr-test-amd.yml index 2afa3cd37..67ab10853 100644 --- a/.github/workflows/pr-test-amd.yml +++ b/.github/workflows/pr-test-amd.yml @@ -638,6 +638,7 @@ jobs: -e SGLANG_NON_DENOISE_STAGE_TIME_TOLERANCE=0.6 \ -e SGLANG_DENOISE_STEP_TOLERANCE=0.6 \ -e SGLANG_DENOISE_AGG_TOLERANCE=0.3 \ + -e SGLANG_SKIP_CONSISTENCY=1 \ -e SGLANG_TEST_NUM_INFERENCE_STEPS=5 \ -e AITER_JIT_DIR=/sgl-data/aiter-kernels \ -e MIOPEN_USER_DB_PATH=/sgl-data/miopen-cache \ @@ -765,6 +766,7 @@ jobs: -e SGLANG_NON_DENOISE_STAGE_TIME_TOLERANCE=0.6 \ -e SGLANG_DENOISE_STEP_TOLERANCE=0.6 \ -e SGLANG_DENOISE_AGG_TOLERANCE=0.3 \ + -e SGLANG_SKIP_CONSISTENCY=1 \ -e SGLANG_TEST_NUM_INFERENCE_STEPS=5 \ -e AITER_JIT_DIR=/sgl-data/aiter-kernels \ -e MIOPEN_USER_DB_PATH=/sgl-data/miopen-cache \ diff --git a/python/sglang/multimodal_gen/runtime/pipelines_core/lora_pipeline.py b/python/sglang/multimodal_gen/runtime/pipelines_core/lora_pipeline.py index f5d3fb24e..a4777fbfe 100644 --- a/python/sglang/multimodal_gen/runtime/pipelines_core/lora_pipeline.py +++ b/python/sglang/multimodal_gen/runtime/pipelines_core/lora_pipeline.py @@ -506,16 +506,25 @@ class LoRAPipeline(ComposedPipelineBase): return bool(self.cur_adapter_name) return target in self.cur_adapter_name - def load_lora_adapter(self, lora_path: str, lora_nickname: str, rank: int): + def load_lora_adapter( + self, + lora_path: str, + lora_nickname: str, + rank: int, + weight_name: str | None = None, + ): """ Load the LoRA, and setup the lora_adapters for later weight replacement """ assert lora_path is not None + if weight_name is None and lora_path == self.server_args.lora_path: + weight_name = self.server_args.lora_weight_name + # Only rank 0 downloads to avoid race conditions where other ranks # try to load incomplete downloads if rank == 0: - lora_local_path = maybe_download_lora(lora_path) + lora_local_path = maybe_download_lora(lora_path, weight_name=weight_name) else: lora_local_path = None @@ -525,7 +534,7 @@ class LoRAPipeline(ComposedPipelineBase): # Non-rank-0 workers now download (will hit cache since rank 0 completed) if rank != 0: - lora_local_path = maybe_download_lora(lora_path) + lora_local_path = maybe_download_lora(lora_path, weight_name=weight_name) raw_state_dict = load_file(lora_local_path) lora_state_dict = normalize_lora_state_dict(raw_state_dict, logger=logger) diff --git a/python/sglang/multimodal_gen/runtime/server_args.py b/python/sglang/multimodal_gen/runtime/server_args.py index 0d0c82cfb..4402c6eb8 100644 --- a/python/sglang/multimodal_gen/runtime/server_args.py +++ b/python/sglang/multimodal_gen/runtime/server_args.py @@ -149,6 +149,7 @@ class ServerArgs: lora_path: str | None = None lora_nickname: str = "default" # for swapping adapters in the pipeline lora_scale: float = 1.0 # LoRA scale for merging (e.g., 0.125 for Hyper-SD) + lora_weight_name: str | None = None # Component path overrides (key = model_index.json component name, value = path) component_paths: dict[str, str] = field(default_factory=dict) @@ -849,6 +850,12 @@ class ServerArgs: default=ServerArgs.lora_scale, help="LoRA scale for merging (e.g., 0.125 for Hyper-SD). Same as lora_scale in Diffusers", ) + parser.add_argument( + "--lora-weight-name", + type=str, + default=ServerArgs.lora_weight_name, + help="Specific safetensors filename to load from a multi-file LoRA repo", + ) # Add pipeline configuration arguments PipelineConfig.add_cli_args(parser) diff --git a/python/sglang/multimodal_gen/runtime/utils/hf_diffusers_utils.py b/python/sglang/multimodal_gen/runtime/utils/hf_diffusers_utils.py index 25af21d61..c7cd5a1fb 100644 --- a/python/sglang/multimodal_gen/runtime/utils/hf_diffusers_utils.py +++ b/python/sglang/multimodal_gen/runtime/utils/hf_diffusers_utils.py @@ -378,7 +378,10 @@ def check_gguf_file(model: str | os.PathLike) -> bool: def maybe_download_lora( - model_name_or_path: str, local_dir: str | None = None, download: bool = True + model_name_or_path: str, + local_dir: str | None = None, + download: bool = True, + weight_name: str | None = None, ) -> str: """ Check if the model path is a Hugging Face Hub model ID and download it if needed. @@ -386,6 +389,8 @@ def maybe_download_lora( model_name_or_path: Local path or Hugging Face Hub model ID local_dir: Local directory to save the model download: Whether to download the model from Hugging Face Hub + weight_name: Specific safetensors filename to load (pins deterministic selection + for repos with multiple weight files) Returns: Local path to the model @@ -403,14 +408,22 @@ def maybe_download_lora( if os.path.isfile(local_path): return local_path - weight_name = _best_guess_weight_name(local_path, file_extension=".safetensors") + if weight_name is not None: + target = os.path.join(local_path, weight_name) + if not os.path.isfile(target): + raise FileNotFoundError( + f"Specified lora_weight_name '{weight_name}' not found in {local_path}" + ) + return target + + guessed = _best_guess_weight_name(local_path, file_extension=".safetensors") # AMD workaround: PR 15813 changed from model_name_or_path to local_path, # which can return None. Fall back to original behavior on ROCm. - if weight_name is None and current_platform.is_rocm(): - weight_name = _best_guess_weight_name( + if guessed is None and current_platform.is_rocm(): + guessed = _best_guess_weight_name( model_name_or_path, file_extension=".safetensors" ) - return os.path.join(local_path, weight_name) + return os.path.join(local_path, guessed) def verify_model_config_and_directory(model_path: str) -> dict[str, Any]: diff --git a/python/sglang/multimodal_gen/test/server/ascend/testcase_configs_npu.py b/python/sglang/multimodal_gen/test/server/ascend/testcase_configs_npu.py index 0a78eb35f..4086dcd74 100644 --- a/python/sglang/multimodal_gen/test/server/ascend/testcase_configs_npu.py +++ b/python/sglang/multimodal_gen/test/server/ascend/testcase_configs_npu.py @@ -15,6 +15,7 @@ ONE_NPU_CASES: list[DiffusionTestCase] = [ modality="image", ), T2I_sampling_params, + run_consistency_check=False, ), # === Text to Video (T2V) === DiffusionTestCase( @@ -27,6 +28,7 @@ ONE_NPU_CASES: list[DiffusionTestCase] = [ DiffusionSamplingParams( prompt=T2V_PROMPT, ), + run_consistency_check=False, ), ] @@ -41,6 +43,7 @@ TWO_NPU_CASES: list[DiffusionTestCase] = [ tp_size=2, ), T2I_sampling_params, + run_consistency_check=False, ), DiffusionTestCase( "qwen_image_t2i_2npu", @@ -53,6 +56,7 @@ TWO_NPU_CASES: list[DiffusionTestCase] = [ ring_degree=2, ), T2I_sampling_params, + run_consistency_check=False, ), ] @@ -70,5 +74,6 @@ EIGHT_NPU_CASES: list[DiffusionTestCase] = [ DiffusionSamplingParams( prompt=T2V_PROMPT, ), + run_consistency_check=False, ), ] diff --git a/python/sglang/multimodal_gen/test/server/consistency_threshold.json b/python/sglang/multimodal_gen/test/server/consistency_threshold.json new file mode 100644 index 000000000..3795a9f6e --- /dev/null +++ b/python/sglang/multimodal_gen/test/server/consistency_threshold.json @@ -0,0 +1,205 @@ +{ + "_comment": "Some cases use lower thresholds; raise them if quality/perf improves later.", + "cases": { + "qwen_image_t2i": { + "clip_threshold": 0.97, + "ssim_threshold": 0.84, + "psnr_threshold": 16.0, + "mean_abs_diff_threshold": 11.5 + }, + "flux_image_t2i": { + "clip_threshold": 0.92, + "ssim_threshold": 0.95, + "psnr_threshold": 24.0, + "mean_abs_diff_threshold": 8.0 + }, + "flux_2_klein_image_t2i": { + "clip_threshold": 0.86, + "ssim_threshold": 0.60, + "psnr_threshold": 10.0, + "mean_abs_diff_threshold": 56.0 + }, + "zimage_image_t2i": { + "clip_threshold": 0.92, + "ssim_threshold": 0.90, + "psnr_threshold": 22.0, + "mean_abs_diff_threshold": 8.0 + }, + "zimage_image_t2i_multi_lora": { + "clip_threshold": 0.92, + "ssim_threshold": 0.92, + "psnr_threshold": 22.0, + "mean_abs_diff_threshold": 8.0 + }, + "qwen_image_t2i_cache_dit_enabled": { + "clip_threshold": 0.92, + "ssim_threshold": 0.86, + "psnr_threshold": 17.0, + "mean_abs_diff_threshold": 10.0 + }, + "layerwise_offload": { + "clip_threshold": 0.92, + "ssim_threshold": 0.91, + "psnr_threshold": 21.5, + "mean_abs_diff_threshold": 8.0 + }, + "zimage_image_t2i_fp8": { + "clip_threshold": 0.92, + "ssim_threshold": 0.84, + "psnr_threshold": 19.0, + "mean_abs_diff_threshold": 10.0 + }, + "sana_image_t2i": { + "clip_threshold": 0.91, + "ssim_threshold": 0.88, + "psnr_threshold": 21.0, + "mean_abs_diff_threshold": 8.4 + }, + "qwen_image_edit_2509_ti2i": { + "clip_threshold": 0.92, + "ssim_threshold": 0.65, + "psnr_threshold": 13.0, + "mean_abs_diff_threshold": 26.0 + }, + "qwen_image_layered_i2i": { + "clip_threshold": 0.92, + "ssim_threshold": 0.94, + "psnr_threshold": 28.0, + "mean_abs_diff_threshold": 8.0 + }, + "mova_360p_1gpu": { + "clip_threshold": 0.90, + "ssim_threshold": 0.87, + "psnr_threshold": 24.0, + "mean_abs_diff_threshold": 10.0 + }, + "wan2_1_t2v_1_3b_lora_1gpu": { + "clip_threshold": 0.54, + "ssim_threshold": 0.40, + "psnr_threshold": 13.2, + "mean_abs_diff_threshold": 32.0 + }, + "wan2_2_ti2v_5b": { + "clip_threshold": 0.90, + "ssim_threshold": 0.81, + "psnr_threshold": 20.4, + "mean_abs_diff_threshold": 10.0 + }, + "fastwan2_2_ti2v_5b": { + "clip_threshold": 0.90, + "ssim_threshold": 0.88, + "psnr_threshold": 24.0, + "mean_abs_diff_threshold": 10.0 + }, + "turbo_wan2_1_t2v_1.3b": { + "clip_threshold": 0.90, + "ssim_threshold": 0.52, + "psnr_threshold": 9.5, + "mean_abs_diff_threshold": 46.0 + }, + "zimage_image_t2i_multi_lora": { + "clip_threshold": 0.90, + "ssim_threshold": 0.52, + "psnr_threshold": 9.5, + "mean_abs_diff_threshold": 46.0 + }, + "fsdp-inference": { + "clip_threshold": 0.92, + "ssim_threshold": 0.90, + "psnr_threshold": 21.5, + "mean_abs_diff_threshold": 8.0 + }, + "zimage_image_t2i_2_gpus_non_square": { + "clip_threshold": 0.92, + "ssim_threshold": 0.76, + "psnr_threshold": 14.8, + "mean_abs_diff_threshold": 17.5 + }, + "flux_2_image_t2i_2_gpus": { + "clip_threshold": 0.54, + "ssim_threshold": 0.9, + "psnr_threshold": 19, + "mean_abs_diff_threshold": 8.0 + }, + "zimage_image_t2i_2_gpus": { + "clip_threshold": 0.92, + "ssim_threshold": 0.90, + "psnr_threshold": 21.5, + "mean_abs_diff_threshold": 8.0 + }, + "flux_image_t2i_2_gpus": { + "clip_threshold": 0.92, + "ssim_threshold": 0.90, + "psnr_threshold": 18.7, + "mean_abs_diff_threshold": 8.0 + }, + "flux_2_klein_ti2i_2_gpus": { + "clip_threshold": 0.92, + "ssim_threshold": 0.77, + "psnr_threshold": 18.4, + "mean_abs_diff_threshold": 18.0 + }, + "wan2_2_t2v_a14b_teacache_2gpu": { + "clip_threshold": 0.90, + "ssim_threshold": 0.72, + "psnr_threshold": 17.8, + "mean_abs_diff_threshold": 16.0 + }, + "wan2_1_t2v_14b_2gpu": { + "clip_threshold": 0.90, + "ssim_threshold": 0.84, + "psnr_threshold": 24.0, + "mean_abs_diff_threshold": 10.0 + }, + "mova_360p_ring1_uly2": { + "clip_threshold": 0.90, + "ssim_threshold": 0.91, + "psnr_threshold": 24.0, + "mean_abs_diff_threshold": 10.0 + }, + "wan2_1_i2v_14b_lora_2gpu": { + "clip_threshold": 0.90, + "ssim_threshold": 0.81, + "psnr_threshold": 24.0, + "mean_abs_diff_threshold": 10.0 + }, + "wan2_2_t2v_a14b_2gpu": { + "clip_threshold": 0.90, + "ssim_threshold": 0.72, + "psnr_threshold": 17.8, + "mean_abs_diff_threshold": 16.0 + }, + "wan2_2_t2v_a14b_lora_2gpu": { + "clip_threshold": 0.90, + "ssim_threshold": 0.81, + "psnr_threshold": 22.2, + "mean_abs_diff_threshold": 10.0 + }, + "mova_360p_ring2_uly1": { + "clip_threshold": 0.90, + "ssim_threshold": 0.91, + "psnr_threshold": 24.0, + "mean_abs_diff_threshold": 10.0 + }, + "wan2_1_i2v_14b_480P_2gpu": { + "clip_threshold": 0.76, + "ssim_threshold": 0.51, + "psnr_threshold": 14.8, + "mean_abs_diff_threshold": 23.0 + }, + "wan2_1_i2v_14b_720P_2gpu": { + "clip_threshold": 0.90, + "ssim_threshold": 0.89, + "psnr_threshold": 24.0, + "mean_abs_diff_threshold": 10.0 + } + }, + "default_clip_threshold_image": 0.92, + "default_clip_threshold_video": 0.90, + "default_ssim_threshold_image": 0.95, + "default_psnr_threshold_image": 28.0, + "default_mean_abs_diff_threshold_image": 8.0, + "default_ssim_threshold_video": 0.92, + "default_psnr_threshold_video": 24.0, + "default_mean_abs_diff_threshold_video": 10.0 +} diff --git a/python/sglang/multimodal_gen/test/server/test_server_common.py b/python/sglang/multimodal_gen/test/server/test_server_common.py index 29d985d21..dd48e7e0c 100644 --- a/python/sglang/multimodal_gen/test/server/test_server_common.py +++ b/python/sglang/multimodal_gen/test/server/test_server_common.py @@ -35,8 +35,16 @@ from sglang.multimodal_gen.test.server.testcase_configs import ( ) from sglang.multimodal_gen.test.test_utils import ( _consistency_gt_filenames, + _get_consistency_gt_dir, + compare_with_gt, extract_key_frames_from_video, + get_consistency_gt_candidates, + get_consistency_gt_remote_files, + get_consistency_thresholds, get_dynamic_server_port, + gt_exists, + image_bytes_to_numpy, + load_consistency_gt, wait_for_req_perf_record, ) @@ -445,6 +453,146 @@ Consider updating perf_baselines.json with the snippets below: """ logger.error(output) + def _validate_consistency( + self, + case: DiffusionTestCase, + content: bytes, + ) -> None: + """Validate output consistency against ground truth using CLIP similarity.""" + if os.environ.get("SGLANG_SKIP_CONSISTENCY", "0") == "1": + logger.info( + f"[Consistency] Skipping consistency check for {case.id} (SGLANG_SKIP_CONSISTENCY=1)" + ) + return + + if not content: + logger.warning( + f"[Consistency] Skipping consistency check for {case.id}: " + "content is empty (generation may have timed out)" + ) + return + + num_gpus = case.server_args.num_gpus + is_video = case.server_args.modality == "video" + output_format = case.sampling_params.output_format + + if not gt_exists( + case.id, num_gpus, is_video=is_video, output_format=output_format + ): + if _get_consistency_gt_dir() is not None: + names = ", ".join( + get_consistency_gt_candidates( + case.id, num_gpus, is_video, output_format + ) + ) + else: + names = ", ".join( + _consistency_gt_filenames( + case.id, num_gpus, is_video, output_format + ) + ) + logger.error(f""" +--- MISSING GROUND TRUTH DETECTED --- +GT image(s) not found for '{case.id}'. + +Add the expected file(s) to sglang-ci-data in diffusion-ci/consistency_gt/ with naming (n=num_gpus). + Image: {case.id}_{{n}}gpu. (ext from output_format: png, jpg, webp) + Video: {case.id}_{{n}}gpu_frame_0.png, {case.id}_{{n}}gpu_frame_mid.png, {case.id}_{{n}}gpu_frame_last.png + +For this case, expected file(s): {names} + +Repository: https://github.com/sglang-bot/sglang-ci-data (path: diffusion-ci/consistency_gt/) + +(Optional) Per-case override in consistency_threshold.json: + "cases": {{ + "{case.id}": {{ + "clip_threshold": 0.92, + "ssim_threshold": 0.95, + "psnr_threshold": 28.0, + "mean_abs_diff_threshold": 8.0 + }} + }} +""") + pytest.fail( + f"GT not found for {case.id}. See logs for instructions to add GT." + ) + + gt_data = load_consistency_gt( + case.id, num_gpus, is_video=is_video, output_format=output_format + ) + thresholds = get_consistency_thresholds(case.id, is_video=is_video) + + if is_video: + output_frames = extract_key_frames_from_video(content) + else: + output_frames = [image_bytes_to_numpy(content)] + + result = compare_with_gt( + output_frames=output_frames, + gt_data=gt_data, + thresholds=thresholds, + case_id=case.id, + ) + + if not result.passed: + failed_frames = [] + video_gt_info = "" + if is_video: + gt_remote_files = get_consistency_gt_remote_files( + case.id, + num_gpus, + is_video=True, + output_format=output_format, + ) + video_gt_info = "\n".join( + f" - {filename}: {url}" for filename, url in gt_remote_files + ) + for metric in result.frame_metrics: + failed_metrics = [] + if not metric.clip_passed: + failed_metrics.append("clip") + if not metric.ssim_passed: + failed_metrics.append("ssim") + if not metric.psnr_passed: + failed_metrics.append("psnr") + if not metric.mean_abs_diff_passed: + failed_metrics.append("mean_abs_diff") + if failed_metrics: + failed_frames.append( + f" - f{metric.frame_index} " + f"[{', '.join(failed_metrics)}] " + f"clip={metric.clip_similarity:.4f} " + f"ssim={metric.ssim:.4f} " + f"psnr={metric.psnr:.4f} " + f"mean_abs_diff={metric.mean_abs_diff:.4f}" + ) + pytest.fail( + f"Consistency check failed for {case.id}:\n" + f" Metrics: sim={result.min_similarity:.4f}, " + f"ssim={result.min_ssim:.4f}, " + f"psnr={result.min_psnr:.4f}, " + f"mean_abs_diff={result.max_mean_abs_diff:.4f}\n" + f" Thresholds: clip>={result.thresholds.clip_threshold}, " + f"ssim>={result.thresholds.ssim_threshold}, " + f"psnr>={result.thresholds.psnr_threshold}, " + f"mean_abs_diff<={result.thresholds.mean_abs_diff_threshold}\n" + f" Failed frames:\n" + + "\n".join(failed_frames) + + ( + f"\n Compared GT frame files and links:\n{video_gt_info}" + if video_gt_info + else "" + ) + ) + + logger.info( + f"[Consistency] {case.id}: PASSED " + f"(min_similarity={result.min_similarity:.4f}, " + f"min_ssim={result.min_ssim:.4f}, " + f"min_psnr={result.min_psnr:.4f}, " + f"max_mean_abs_diff={result.max_mean_abs_diff:.4f})" + ) + def _save_gt_output( self, case: DiffusionTestCase, @@ -903,6 +1051,9 @@ Consider updating perf_baselines.json with the snippets below: if case.run_t2v_input_reference_check: self._test_t2v_rejects_input_reference(diffusion_server, case) + if case.run_consistency_check: + self._validate_consistency(case, content) + # LoRA API functionality test with E2E validation (only for LoRA-enabled cases) if case.run_lora_basic_api_check: self._test_lora_api_functionality(diffusion_server, case, generate_fn) diff --git a/python/sglang/multimodal_gen/test/server/testcase_configs.py b/python/sglang/multimodal_gen/test/server/testcase_configs.py index 07e2e0f21..d879adce6 100644 --- a/python/sglang/multimodal_gen/test/server/testcase_configs.py +++ b/python/sglang/multimodal_gen/test/server/testcase_configs.py @@ -248,6 +248,7 @@ class DiffusionTestCase: server_args: DiffusionServerArgs sampling_params: DiffusionSamplingParams run_perf_check: bool = True + run_consistency_check: bool = True run_models_api_check: bool = True run_t2v_input_reference_check: bool = True run_lora_basic_api_check: bool = False @@ -777,6 +778,7 @@ if not current_platform.is_hip(): enable_warmup=False, ), HUNYUAN3D_SHAPE_sampling_params, + run_consistency_check=False, ), ) # Skip turbowan on AMD: Triton requires 81920 shared memory, but AMD only has 65536. @@ -802,6 +804,7 @@ ONE_GPU_CASES_C = [ modality="image", ), T2I_sampling_params, + run_consistency_check=False, ) ] @@ -851,6 +854,10 @@ TWO_GPU_CASES_A = [ custom_validator="video", num_gpus=2, lora_path="Cseti/wan2.2-14B-Arcane_Jinx-lora-v1", + extras=[ + "--lora-weight-name", + "985347-wan22_14B-low-Nfj1nx-e65.safetensors", + ], ), DiffusionSamplingParams( prompt="Nfj1nx with blue hair, a woman walking in a cyberpunk city at night", @@ -938,6 +945,7 @@ TWO_GPU_CASES_A = [ extras=["--pipeline-class-name LTX2TwoStagePipeline"], ), T2V_sampling_params, + run_consistency_check=False, ), ] diff --git a/python/sglang/multimodal_gen/test/test_consistency_metrics.py b/python/sglang/multimodal_gen/test/test_consistency_metrics.py new file mode 100644 index 000000000..78906a19c --- /dev/null +++ b/python/sglang/multimodal_gen/test/test_consistency_metrics.py @@ -0,0 +1,115 @@ +import math + +import numpy as np + +from sglang.multimodal_gen.test import test_utils +from sglang.multimodal_gen.test.test_utils import ( + ConsistencyThresholds, + LoadedConsistencyGT, + compare_with_gt, + compute_mean_abs_diff, + compute_psnr, + compute_ssim, +) + + +def _solid_image(value: int, size: int = 32) -> np.ndarray: + return np.full((size, size, 3), value, dtype=np.uint8) + + +def test_pixel_metrics_identical_image(): + image = _solid_image(128) + + ssim = compute_ssim(image, image) + psnr = compute_psnr(image, image) + mean_abs_diff = compute_mean_abs_diff(image, image) + + assert ssim == 1.0 + assert math.isinf(psnr) + assert mean_abs_diff == 0.0 + + +def test_pixel_metrics_detect_different_image(): + image = _solid_image(128) + other = _solid_image(0) + + ssim = compute_ssim(image, other) + psnr = compute_psnr(image, other) + mean_abs_diff = compute_mean_abs_diff(image, other) + + assert ssim < 0.95 + assert psnr < 28.0 + assert mean_abs_diff > 8.0 + + +def test_compare_with_gt_passes_for_identical_image(monkeypatch): + gt_image = _solid_image(128) + + monkeypatch.setattr( + test_utils, + "compute_clip_embedding", + lambda image: np.array([1.0, 0.0], dtype=np.float32), + ) + + result = compare_with_gt( + output_frames=[gt_image.copy()], + gt_data=LoadedConsistencyGT( + images=[gt_image.copy()], + embeddings=[np.array([1.0, 0.0], dtype=np.float32)], + ), + thresholds=ConsistencyThresholds( + clip_threshold=0.92, + ssim_threshold=0.95, + psnr_threshold=28.0, + mean_abs_diff_threshold=8.0, + ), + case_id="unit_image_pass", + ) + + assert result.passed is True + assert result.min_similarity == 1.0 + assert result.min_ssim == 1.0 + assert math.isinf(result.min_psnr) + assert result.max_mean_abs_diff == 0.0 + + +def test_compare_with_gt_uses_worst_frame_for_video(monkeypatch): + gt_frame_0 = _solid_image(128) + gt_frame_1 = _solid_image(128) + bad_frame = _solid_image(0) + + monkeypatch.setattr( + test_utils, + "compute_clip_embedding", + lambda image: np.array([1.0, 0.0], dtype=np.float32), + ) + + result = compare_with_gt( + output_frames=[gt_frame_0.copy(), bad_frame], + gt_data=LoadedConsistencyGT( + images=[gt_frame_0.copy(), gt_frame_1.copy()], + embeddings=[ + np.array([1.0, 0.0], dtype=np.float32), + np.array([1.0, 0.0], dtype=np.float32), + ], + ), + thresholds=ConsistencyThresholds( + clip_threshold=0.92, + ssim_threshold=0.95, + psnr_threshold=28.0, + mean_abs_diff_threshold=8.0, + ), + case_id="unit_video_fail", + ) + + assert result.passed is False + assert result.min_similarity == 1.0 + assert result.min_ssim < 0.95 + assert result.min_psnr < 28.0 + assert result.max_mean_abs_diff > 8.0 + assert any( + not metric.ssim_passed + or not metric.psnr_passed + or not metric.mean_abs_diff_passed + for metric in result.frame_metrics + ) diff --git a/python/sglang/multimodal_gen/test/test_utils.py b/python/sglang/multimodal_gen/test/test_utils.py index 2b2f6cb40..d25d6b18a 100644 --- a/python/sglang/multimodal_gen/test/test_utils.py +++ b/python/sglang/multimodal_gen/test/test_utils.py @@ -7,12 +7,15 @@ import socket import subprocess import tempfile import time +from dataclasses import dataclass from pathlib import Path +from typing import TYPE_CHECKING, Any from urllib.parse import urljoin import cv2 import httpx import numpy as np +import requests from PIL import Image from sglang.multimodal_gen.runtime.utils.common import get_bool_env_var @@ -22,8 +25,27 @@ from sglang.multimodal_gen.runtime.utils.perf_logger import ( get_diffusion_perf_log_dir, ) +if TYPE_CHECKING: + from sglang.multimodal_gen.test.server.testcase_configs import DiffusionTestCase + logger = init_logger(__name__) +SGL_TEST_FILES_CONSISTENCY_GT_BASE = "https://raw.githubusercontent.com/sglang-bot/sglang-ci-data/main/diffusion-ci/consistency_gt" +CONSISTENCY_THRESHOLD_JSON_PATH = ( + Path(__file__).resolve().parent / "server" / "consistency_threshold.json" +) +CLIP_MODEL_NAME = "openai/clip-vit-large-patch14" +DEFAULT_CLIP_THRESHOLD_IMAGE = 0.92 +DEFAULT_CLIP_THRESHOLD_VIDEO = 0.90 +DEFAULT_SSIM_THRESHOLD_IMAGE = 0.95 +DEFAULT_PSNR_THRESHOLD_IMAGE = 28.0 +DEFAULT_MEAN_ABS_DIFF_THRESHOLD_IMAGE = 8.0 +DEFAULT_SSIM_THRESHOLD_VIDEO = 0.92 +DEFAULT_PSNR_THRESHOLD_VIDEO = 24.0 +DEFAULT_MEAN_ABS_DIFF_THRESHOLD_VIDEO = 10.0 +_clip_model_cache: dict[str, Any] = {} +_consistency_gt_cache: dict[str, Any] = {} + # --------------------------------------------------------------------------- # Common model IDs for diffusion tests # @@ -560,10 +582,232 @@ def validate_video_file( ), f"Video height mismatch: expected {expected_height}, got {actual_height}" +def _load_threshold_json() -> dict[str, Any]: + """Load consistency_threshold.json; returns {} if missing.""" + if not CONSISTENCY_THRESHOLD_JSON_PATH.exists(): + return {} + with CONSISTENCY_THRESHOLD_JSON_PATH.open("r", encoding="utf-8") as f: + return json.load(f) + + +@dataclass +class ConsistencyThresholds: + clip_threshold: float + ssim_threshold: float + psnr_threshold: float + mean_abs_diff_threshold: float + + +def get_consistency_thresholds( + case_id: str, + is_video: bool, + metadata: dict[str, Any] | None = None, +) -> ConsistencyThresholds: + """Get all consistency thresholds for a case.""" + if metadata is None: + metadata = _load_threshold_json() + + case_meta = metadata.get("cases", {}).get(case_id, {}) + suffix = "video" if is_video else "image" + + defaults = { + "clip_threshold": metadata.get( + f"default_clip_threshold_{suffix}", + DEFAULT_CLIP_THRESHOLD_VIDEO if is_video else DEFAULT_CLIP_THRESHOLD_IMAGE, + ), + "ssim_threshold": metadata.get( + f"default_ssim_threshold_{suffix}", + DEFAULT_SSIM_THRESHOLD_VIDEO if is_video else DEFAULT_SSIM_THRESHOLD_IMAGE, + ), + "psnr_threshold": metadata.get( + f"default_psnr_threshold_{suffix}", + DEFAULT_PSNR_THRESHOLD_VIDEO if is_video else DEFAULT_PSNR_THRESHOLD_IMAGE, + ), + "mean_abs_diff_threshold": metadata.get( + f"default_mean_abs_diff_threshold_{suffix}", + ( + DEFAULT_MEAN_ABS_DIFF_THRESHOLD_VIDEO + if is_video + else DEFAULT_MEAN_ABS_DIFF_THRESHOLD_IMAGE + ), + ), + } + + return ConsistencyThresholds( + clip_threshold=float( + case_meta.get("clip_threshold", defaults["clip_threshold"]) + ), + ssim_threshold=float( + case_meta.get("ssim_threshold", defaults["ssim_threshold"]) + ), + psnr_threshold=float( + case_meta.get("psnr_threshold", defaults["psnr_threshold"]) + ), + mean_abs_diff_threshold=float( + case_meta.get( + "mean_abs_diff_threshold", defaults["mean_abs_diff_threshold"] + ) + ), + ) + + +def get_clip_threshold( + case: "DiffusionTestCase", + metadata: dict[str, Any] | None = None, +) -> float: + """Get CLIP similarity threshold for a consistency test case.""" + return get_consistency_thresholds( + case_id=case.id, + is_video=case.server_args.modality == "video", + metadata=metadata, + ).clip_threshold + + +@dataclass +class FrameConsistencyMetrics: + frame_index: int + clip_similarity: float + ssim: float + psnr: float + mean_abs_diff: float + clip_passed: bool + ssim_passed: bool + psnr_passed: bool + mean_abs_diff_passed: bool + + +@dataclass +class ConsistencyResult: + """Result of a consistency comparison.""" + + case_id: str + passed: bool + similarity_scores: list[float] + min_similarity: float + threshold: float + min_ssim: float + min_psnr: float + max_mean_abs_diff: float + thresholds: ConsistencyThresholds + frame_metrics: list[FrameConsistencyMetrics] + + +@dataclass +class LoadedConsistencyGT: + images: list[np.ndarray] + embeddings: list[np.ndarray] + + +def get_clip_model() -> tuple[Any, Any]: + """Get CLIP model and processor.""" + global _clip_model_cache + + if "model" not in _clip_model_cache: + try: + import torch + from transformers import CLIPModel, CLIPProcessor + except ImportError as exc: + raise ImportError( + "transformers and torch are required for CLIP consistency check." + ) from exc + + logger.info(f"Loading CLIP model: {CLIP_MODEL_NAME}") + processor = CLIPProcessor.from_pretrained(CLIP_MODEL_NAME) + model = CLIPModel.from_pretrained(CLIP_MODEL_NAME) + + device = "cuda" if torch.cuda.is_available() else "cpu" + model = model.to(device) + model.eval() + + _clip_model_cache["model"] = model + _clip_model_cache["processor"] = processor + _clip_model_cache["device"] = device + logger.info(f"CLIP model loaded on {device}") + + return _clip_model_cache["model"], _clip_model_cache["processor"] + + +def compute_clip_embedding(image: np.ndarray) -> np.ndarray: + """Compute a normalized CLIP image embedding.""" + try: + import torch + except ImportError as exc: + raise ImportError("torch is required for CLIP consistency check.") from exc + + model, processor = get_clip_model() + device = _clip_model_cache["device"] + + pil_image = Image.fromarray(image) + inputs = processor(images=pil_image, return_tensors="pt") + inputs = {k: v.to(device) for k, v in inputs.items()} + + with torch.no_grad(): + image_features = model.get_image_features(**inputs) + if hasattr(image_features, "image_embeds"): + image_features = image_features.image_embeds + elif hasattr(image_features, "pooler_output"): + image_features = image_features.pooler_output + image_features = image_features / image_features.norm(dim=-1, keepdim=True) + + return image_features.cpu().numpy().flatten() + + +def compute_clip_similarity(emb1: np.ndarray, emb2: np.ndarray) -> float: + """Compute cosine similarity between two CLIP embeddings.""" + return float(np.dot(emb1, emb2)) + + +def _ensure_rgb_uint8_image(image: np.ndarray) -> np.ndarray: + """Normalize image input for pixel-wise consistency metrics.""" + if image.ndim != 3 or image.shape[2] != 3: + raise ValueError(f"Expected RGB HWC image, got shape={image.shape}") + if image.dtype == np.uint8: + return image + image = np.clip(image, 0, 255) + return image.astype(np.uint8) + + +def compute_ssim(image: np.ndarray, gt_image: np.ndarray) -> float: + """Compute SSIM between two RGB images.""" + from skimage.metrics import structural_similarity + + image = _ensure_rgb_uint8_image(image) + gt_image = _ensure_rgb_uint8_image(gt_image) + if image.shape != gt_image.shape: + raise ValueError( + f"Image shape mismatch for SSIM: output={image.shape}, gt={gt_image.shape}" + ) + return float(structural_similarity(image, gt_image, channel_axis=2, data_range=255)) + + +def compute_psnr(image: np.ndarray, gt_image: np.ndarray) -> float: + """Compute PSNR between two RGB images.""" + from skimage.metrics import peak_signal_noise_ratio + + image = _ensure_rgb_uint8_image(image) + gt_image = _ensure_rgb_uint8_image(gt_image) + if image.shape != gt_image.shape: + raise ValueError( + f"Image shape mismatch for PSNR: output={image.shape}, gt={gt_image.shape}" + ) + return float(peak_signal_noise_ratio(gt_image, image, data_range=255)) + + +def compute_mean_abs_diff(image: np.ndarray, gt_image: np.ndarray) -> float: + """Compute mean absolute pixel difference between two RGB images.""" + image = _ensure_rgb_uint8_image(image) + gt_image = _ensure_rgb_uint8_image(gt_image) + if image.shape != gt_image.shape: + raise ValueError( + f"Image shape mismatch for mean_abs_diff: output={image.shape}, gt={gt_image.shape}" + ) + return float(np.abs(image.astype(np.float32) - gt_image.astype(np.float32)).mean()) + + def output_format_to_ext(output_format: str | None) -> str: """Map output_format to file extension. Used by GT naming and consistency check.""" if not output_format: - return "png" + return "jpg" of = output_format.lower() if of == "jpeg": return "jpg" @@ -587,6 +831,151 @@ def _consistency_gt_filenames( return [f"{case_id}_{n}gpu.{ext}"] +def get_consistency_gt_candidates( + case_id: str, num_gpus: int, is_video: bool, output_format: str | None = None +) -> list[str]: + """Return candidate GT filenames for local consistency data.""" + n = num_gpus + if is_video: + return [ + f"{case_id}_{n}gpu_frame_0.png", + f"{case_id}_{n}gpu_frame_mid.png", + f"{case_id}_{n}gpu_frame_last.png", + ] + base = f"{case_id}_{n}gpu" + preferred = output_format_to_ext(output_format) + exts = [preferred] + [e for e in ("png", "jpg", "webp") if e != preferred] + return [f"{base}.{e}" for e in exts] + + +def get_consistency_gt_remote_files( + case_id: str, num_gpus: int, is_video: bool, output_format: str | None = None +) -> list[tuple[str, str]]: + """Return GT filenames with their remote raw URLs.""" + filenames = _consistency_gt_filenames(case_id, num_gpus, is_video, output_format) + return [ + (filename, f"{SGL_TEST_FILES_CONSISTENCY_GT_BASE}/{filename}") + for filename in filenames + ] + + +def _get_consistency_gt_dir() -> Path | None: + """Return the local GT directory when configured.""" + d = os.environ.get("SGLANG_CONSISTENCY_GT_DIR") + if not d: + return None + return Path(d).resolve() + + +def _get_consistency_gt_cache_key( + case_id: str, + num_gpus: int, + is_video: bool, + output_format: str | None, +) -> str: + gt_dir = _get_consistency_gt_dir() + source = str(gt_dir) if gt_dir is not None else "remote" + return f"{case_id}:{num_gpus}:{is_video}:{output_format or ''}:{source}" + + +def load_consistency_gt( + case_id: str, + num_gpus: int, + is_video: bool = False, + output_format: str | None = None, +) -> LoadedConsistencyGT: + """Load GT images and CLIP embeddings for consistency checks.""" + cache_key = _get_consistency_gt_cache_key( + case_id, num_gpus, is_video, output_format + ) + cached = _consistency_gt_cache.get(cache_key) + if cached is not None: + return cached + + filenames = _consistency_gt_filenames(case_id, num_gpus, is_video, output_format) + images: list[np.ndarray] = [] + + gt_dir = _get_consistency_gt_dir() + if gt_dir is not None: + candidates = get_consistency_gt_candidates( + case_id, num_gpus, is_video, output_format + ) + if is_video: + for fn in candidates: + path = gt_dir / fn + if not path.exists(): + raise FileNotFoundError(f"GT image not found: {path}") + arr = np.array(Image.open(path).convert("RGB")) + images.append(arr) + else: + path = None + for fn in candidates: + candidate = gt_dir / fn + if candidate.exists(): + path = candidate + break + if path is None: + raise FileNotFoundError( + f"GT image not found in {gt_dir}. Tried: {', '.join(candidates)}" + ) + images.append(np.array(Image.open(path).convert("RGB"))) + logger.info(f"Loaded {len(images)} GT images for {case_id} from {gt_dir}") + else: + for fn in filenames: + url = f"{SGL_TEST_FILES_CONSISTENCY_GT_BASE}/{fn}" + resp = requests.get(url, timeout=30) + if resp.status_code != 200: + raise FileNotFoundError(f"GT image not found: {url}") + images.append(np.array(Image.open(io.BytesIO(resp.content)).convert("RGB"))) + logger.info(f"Loaded {len(images)} GT images for {case_id} from sglang-ci-data") + + embeddings = [compute_clip_embedding(arr) for arr in images] + loaded_gt = LoadedConsistencyGT(images=images, embeddings=embeddings) + _consistency_gt_cache[cache_key] = loaded_gt + return loaded_gt + + +def load_gt_embeddings( + case_id: str, + num_gpus: int, + is_video: bool = False, + output_format: str | None = None, +) -> list[np.ndarray]: + """Load GT images and convert them into CLIP embeddings.""" + return load_consistency_gt( + case_id=case_id, + num_gpus=num_gpus, + is_video=is_video, + output_format=output_format, + ).embeddings + + +def gt_exists( + case_id: str, + num_gpus: int, + is_video: bool = False, + output_format: str | None = None, +) -> bool: + """Check whether GT image(s) exist.""" + gt_dir = _get_consistency_gt_dir() + if gt_dir is not None: + candidates = get_consistency_gt_candidates( + case_id, num_gpus, is_video, output_format + ) + if is_video: + return all((gt_dir / c).exists() for c in candidates) + return any((gt_dir / c).exists() for c in candidates) + + filenames = _consistency_gt_filenames(case_id, num_gpus, is_video, output_format) + fn = filenames[0] + url = f"{SGL_TEST_FILES_CONSISTENCY_GT_BASE}/{fn}" + try: + r = requests.head(url, timeout=10) + return r.status_code == 200 + except Exception: + return False + + def extract_key_frames_from_video( video_bytes: bytes, num_frames: int | None = None, @@ -643,3 +1032,114 @@ def image_bytes_to_numpy(image_bytes: bytes) -> np.ndarray: """Convert image bytes to numpy array.""" img = Image.open(io.BytesIO(image_bytes)).convert("RGB") return np.array(img) + + +def compare_with_gt( + output_frames: list[np.ndarray], + gt_data: LoadedConsistencyGT, + thresholds: ConsistencyThresholds, + case_id: str, +) -> ConsistencyResult: + """Compare output frames with GT using CLIP and pixel-level metrics.""" + if len(output_frames) != len(gt_data.embeddings): + raise ValueError( + f"Frame count mismatch: output={len(output_frames)}, gt={len(gt_data.embeddings)}" + ) + + similarity_scores = [] + frame_metrics: list[FrameConsistencyMetrics] = [] + + for i, (out_frame, gt_frame, gt_emb) in enumerate( + zip(output_frames, gt_data.images, gt_data.embeddings) + ): + out_frame = _ensure_rgb_uint8_image(out_frame) + gt_frame = _ensure_rgb_uint8_image(gt_frame) + if out_frame.shape != gt_frame.shape: + raise ValueError( + f"Frame shape mismatch for case {case_id}, frame {i}: " + f"output={out_frame.shape}, gt={gt_frame.shape}" + ) + out_emb = compute_clip_embedding(out_frame) + clip_similarity = compute_clip_similarity(out_emb, gt_emb) + ssim = compute_ssim(out_frame, gt_frame) + psnr = compute_psnr(out_frame, gt_frame) + mean_abs_diff = compute_mean_abs_diff(out_frame, gt_frame) + similarity_scores.append(clip_similarity) + frame_metrics.append( + FrameConsistencyMetrics( + frame_index=i, + clip_similarity=clip_similarity, + ssim=ssim, + psnr=psnr, + mean_abs_diff=mean_abs_diff, + clip_passed=clip_similarity >= thresholds.clip_threshold, + ssim_passed=ssim >= thresholds.ssim_threshold, + psnr_passed=psnr >= thresholds.psnr_threshold, + mean_abs_diff_passed=( + mean_abs_diff <= thresholds.mean_abs_diff_threshold + ), + ) + ) + + min_similarity = min(similarity_scores) + min_ssim = min(metric.ssim for metric in frame_metrics) + min_psnr = min(metric.psnr for metric in frame_metrics) + max_mean_abs_diff = max(metric.mean_abs_diff for metric in frame_metrics) + passed = all( + metric.clip_passed + and metric.ssim_passed + and metric.psnr_passed + and metric.mean_abs_diff_passed + for metric in frame_metrics + ) + + result = ConsistencyResult( + case_id=case_id, + passed=passed, + similarity_scores=similarity_scores, + min_similarity=min_similarity, + threshold=thresholds.clip_threshold, + min_ssim=min_ssim, + min_psnr=min_psnr, + max_mean_abs_diff=max_mean_abs_diff, + thresholds=thresholds, + frame_metrics=frame_metrics, + ) + + status = "PASSED" if passed else "FAILED" + print(f"\n{'=' * 60}") + print(f"[CLIP Consistency] {case_id}: {status}") + print( + " Thresholds: " + f"clip>={thresholds.clip_threshold}, " + f"ssim>={thresholds.ssim_threshold}, " + f"psnr>={thresholds.psnr_threshold}, " + f"mean_abs_diff<={thresholds.mean_abs_diff_threshold}" + ) + print(f" Min similarity: {min_similarity:.4f}") + print(f" Min SSIM: {min_ssim:.4f}") + print(f" Min PSNR: {min_psnr:.4f}") + print(f" Max mean_abs_diff: {max_mean_abs_diff:.4f}") + print(" Frame details:") + for metric in frame_metrics: + frame_status = ( + "PASS" + if ( + metric.clip_passed + and metric.ssim_passed + and metric.psnr_passed + and metric.mean_abs_diff_passed + ) + else "FAIL" + ) + print( + f" Frame {metric.frame_index}: " + f"clip={metric.clip_similarity:.4f} " + f"ssim={metric.ssim:.4f} " + f"psnr={metric.psnr:.4f} " + f"mean_abs_diff={metric.mean_abs_diff:.4f} " + f"{frame_status}" + ) + print(f"{'=' * 60}\n") + + return result