[diffusion] CI: add consistency test (#15236)
Co-authored-by: daiweitao <dwti614707404@163.com>
This commit is contained in:
@@ -633,6 +633,7 @@ jobs:
|
|||||||
-e SGLANG_NON_DENOISE_STAGE_TIME_TOLERANCE=0.6 \
|
-e SGLANG_NON_DENOISE_STAGE_TIME_TOLERANCE=0.6 \
|
||||||
-e SGLANG_DENOISE_STEP_TOLERANCE=0.6 \
|
-e SGLANG_DENOISE_STEP_TOLERANCE=0.6 \
|
||||||
-e SGLANG_DENOISE_AGG_TOLERANCE=0.3 \
|
-e SGLANG_DENOISE_AGG_TOLERANCE=0.3 \
|
||||||
|
-e SGLANG_SKIP_CONSISTENCY=1 \
|
||||||
-e SGLANG_TEST_NUM_INFERENCE_STEPS=5 \
|
-e SGLANG_TEST_NUM_INFERENCE_STEPS=5 \
|
||||||
-e AITER_JIT_DIR=/sgl-data/aiter-kernels \
|
-e AITER_JIT_DIR=/sgl-data/aiter-kernels \
|
||||||
-e MIOPEN_USER_DB_PATH=/sgl-data/miopen-cache \
|
-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_NON_DENOISE_STAGE_TIME_TOLERANCE=0.6 \
|
||||||
-e SGLANG_DENOISE_STEP_TOLERANCE=0.6 \
|
-e SGLANG_DENOISE_STEP_TOLERANCE=0.6 \
|
||||||
-e SGLANG_DENOISE_AGG_TOLERANCE=0.3 \
|
-e SGLANG_DENOISE_AGG_TOLERANCE=0.3 \
|
||||||
|
-e SGLANG_SKIP_CONSISTENCY=1 \
|
||||||
-e SGLANG_TEST_NUM_INFERENCE_STEPS=5 \
|
-e SGLANG_TEST_NUM_INFERENCE_STEPS=5 \
|
||||||
-e AITER_JIT_DIR=/sgl-data/aiter-kernels \
|
-e AITER_JIT_DIR=/sgl-data/aiter-kernels \
|
||||||
-e MIOPEN_USER_DB_PATH=/sgl-data/miopen-cache \
|
-e MIOPEN_USER_DB_PATH=/sgl-data/miopen-cache \
|
||||||
|
|||||||
@@ -638,6 +638,7 @@ jobs:
|
|||||||
-e SGLANG_NON_DENOISE_STAGE_TIME_TOLERANCE=0.6 \
|
-e SGLANG_NON_DENOISE_STAGE_TIME_TOLERANCE=0.6 \
|
||||||
-e SGLANG_DENOISE_STEP_TOLERANCE=0.6 \
|
-e SGLANG_DENOISE_STEP_TOLERANCE=0.6 \
|
||||||
-e SGLANG_DENOISE_AGG_TOLERANCE=0.3 \
|
-e SGLANG_DENOISE_AGG_TOLERANCE=0.3 \
|
||||||
|
-e SGLANG_SKIP_CONSISTENCY=1 \
|
||||||
-e SGLANG_TEST_NUM_INFERENCE_STEPS=5 \
|
-e SGLANG_TEST_NUM_INFERENCE_STEPS=5 \
|
||||||
-e AITER_JIT_DIR=/sgl-data/aiter-kernels \
|
-e AITER_JIT_DIR=/sgl-data/aiter-kernels \
|
||||||
-e MIOPEN_USER_DB_PATH=/sgl-data/miopen-cache \
|
-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_NON_DENOISE_STAGE_TIME_TOLERANCE=0.6 \
|
||||||
-e SGLANG_DENOISE_STEP_TOLERANCE=0.6 \
|
-e SGLANG_DENOISE_STEP_TOLERANCE=0.6 \
|
||||||
-e SGLANG_DENOISE_AGG_TOLERANCE=0.3 \
|
-e SGLANG_DENOISE_AGG_TOLERANCE=0.3 \
|
||||||
|
-e SGLANG_SKIP_CONSISTENCY=1 \
|
||||||
-e SGLANG_TEST_NUM_INFERENCE_STEPS=5 \
|
-e SGLANG_TEST_NUM_INFERENCE_STEPS=5 \
|
||||||
-e AITER_JIT_DIR=/sgl-data/aiter-kernels \
|
-e AITER_JIT_DIR=/sgl-data/aiter-kernels \
|
||||||
-e MIOPEN_USER_DB_PATH=/sgl-data/miopen-cache \
|
-e MIOPEN_USER_DB_PATH=/sgl-data/miopen-cache \
|
||||||
|
|||||||
@@ -506,16 +506,25 @@ class LoRAPipeline(ComposedPipelineBase):
|
|||||||
return bool(self.cur_adapter_name)
|
return bool(self.cur_adapter_name)
|
||||||
return target in 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
|
Load the LoRA, and setup the lora_adapters for later weight replacement
|
||||||
"""
|
"""
|
||||||
assert lora_path is not None
|
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
|
# Only rank 0 downloads to avoid race conditions where other ranks
|
||||||
# try to load incomplete downloads
|
# try to load incomplete downloads
|
||||||
if rank == 0:
|
if rank == 0:
|
||||||
lora_local_path = maybe_download_lora(lora_path)
|
lora_local_path = maybe_download_lora(lora_path, weight_name=weight_name)
|
||||||
else:
|
else:
|
||||||
lora_local_path = None
|
lora_local_path = None
|
||||||
|
|
||||||
@@ -525,7 +534,7 @@ class LoRAPipeline(ComposedPipelineBase):
|
|||||||
|
|
||||||
# Non-rank-0 workers now download (will hit cache since rank 0 completed)
|
# Non-rank-0 workers now download (will hit cache since rank 0 completed)
|
||||||
if rank != 0:
|
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)
|
raw_state_dict = load_file(lora_local_path)
|
||||||
lora_state_dict = normalize_lora_state_dict(raw_state_dict, logger=logger)
|
lora_state_dict = normalize_lora_state_dict(raw_state_dict, logger=logger)
|
||||||
|
|||||||
@@ -149,6 +149,7 @@ class ServerArgs:
|
|||||||
lora_path: str | None = None
|
lora_path: str | None = None
|
||||||
lora_nickname: str = "default" # for swapping adapters in the pipeline
|
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_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 path overrides (key = model_index.json component name, value = path)
|
||||||
component_paths: dict[str, str] = field(default_factory=dict)
|
component_paths: dict[str, str] = field(default_factory=dict)
|
||||||
@@ -849,6 +850,12 @@ class ServerArgs:
|
|||||||
default=ServerArgs.lora_scale,
|
default=ServerArgs.lora_scale,
|
||||||
help="LoRA scale for merging (e.g., 0.125 for Hyper-SD). Same as lora_scale in Diffusers",
|
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
|
# Add pipeline configuration arguments
|
||||||
PipelineConfig.add_cli_args(parser)
|
PipelineConfig.add_cli_args(parser)
|
||||||
|
|
||||||
|
|||||||
@@ -378,7 +378,10 @@ def check_gguf_file(model: str | os.PathLike) -> bool:
|
|||||||
|
|
||||||
|
|
||||||
def maybe_download_lora(
|
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:
|
) -> str:
|
||||||
"""
|
"""
|
||||||
Check if the model path is a Hugging Face Hub model ID and download it if needed.
|
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
|
model_name_or_path: Local path or Hugging Face Hub model ID
|
||||||
local_dir: Local directory to save the model
|
local_dir: Local directory to save the model
|
||||||
download: Whether to download the model from Hugging Face Hub
|
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:
|
Returns:
|
||||||
Local path to the model
|
Local path to the model
|
||||||
@@ -403,14 +408,22 @@ def maybe_download_lora(
|
|||||||
if os.path.isfile(local_path):
|
if os.path.isfile(local_path):
|
||||||
return 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,
|
# AMD workaround: PR 15813 changed from model_name_or_path to local_path,
|
||||||
# which can return None. Fall back to original behavior on ROCm.
|
# which can return None. Fall back to original behavior on ROCm.
|
||||||
if weight_name is None and current_platform.is_rocm():
|
if guessed is None and current_platform.is_rocm():
|
||||||
weight_name = _best_guess_weight_name(
|
guessed = _best_guess_weight_name(
|
||||||
model_name_or_path, file_extension=".safetensors"
|
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]:
|
def verify_model_config_and_directory(model_path: str) -> dict[str, Any]:
|
||||||
|
|||||||
@@ -15,6 +15,7 @@ ONE_NPU_CASES: list[DiffusionTestCase] = [
|
|||||||
modality="image",
|
modality="image",
|
||||||
),
|
),
|
||||||
T2I_sampling_params,
|
T2I_sampling_params,
|
||||||
|
run_consistency_check=False,
|
||||||
),
|
),
|
||||||
# === Text to Video (T2V) ===
|
# === Text to Video (T2V) ===
|
||||||
DiffusionTestCase(
|
DiffusionTestCase(
|
||||||
@@ -27,6 +28,7 @@ ONE_NPU_CASES: list[DiffusionTestCase] = [
|
|||||||
DiffusionSamplingParams(
|
DiffusionSamplingParams(
|
||||||
prompt=T2V_PROMPT,
|
prompt=T2V_PROMPT,
|
||||||
),
|
),
|
||||||
|
run_consistency_check=False,
|
||||||
),
|
),
|
||||||
]
|
]
|
||||||
|
|
||||||
@@ -41,6 +43,7 @@ TWO_NPU_CASES: list[DiffusionTestCase] = [
|
|||||||
tp_size=2,
|
tp_size=2,
|
||||||
),
|
),
|
||||||
T2I_sampling_params,
|
T2I_sampling_params,
|
||||||
|
run_consistency_check=False,
|
||||||
),
|
),
|
||||||
DiffusionTestCase(
|
DiffusionTestCase(
|
||||||
"qwen_image_t2i_2npu",
|
"qwen_image_t2i_2npu",
|
||||||
@@ -53,6 +56,7 @@ TWO_NPU_CASES: list[DiffusionTestCase] = [
|
|||||||
ring_degree=2,
|
ring_degree=2,
|
||||||
),
|
),
|
||||||
T2I_sampling_params,
|
T2I_sampling_params,
|
||||||
|
run_consistency_check=False,
|
||||||
),
|
),
|
||||||
]
|
]
|
||||||
|
|
||||||
@@ -70,5 +74,6 @@ EIGHT_NPU_CASES: list[DiffusionTestCase] = [
|
|||||||
DiffusionSamplingParams(
|
DiffusionSamplingParams(
|
||||||
prompt=T2V_PROMPT,
|
prompt=T2V_PROMPT,
|
||||||
),
|
),
|
||||||
|
run_consistency_check=False,
|
||||||
),
|
),
|
||||||
]
|
]
|
||||||
|
|||||||
@@ -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
|
||||||
|
}
|
||||||
@@ -35,8 +35,16 @@ from sglang.multimodal_gen.test.server.testcase_configs import (
|
|||||||
)
|
)
|
||||||
from sglang.multimodal_gen.test.test_utils import (
|
from sglang.multimodal_gen.test.test_utils import (
|
||||||
_consistency_gt_filenames,
|
_consistency_gt_filenames,
|
||||||
|
_get_consistency_gt_dir,
|
||||||
|
compare_with_gt,
|
||||||
extract_key_frames_from_video,
|
extract_key_frames_from_video,
|
||||||
|
get_consistency_gt_candidates,
|
||||||
|
get_consistency_gt_remote_files,
|
||||||
|
get_consistency_thresholds,
|
||||||
get_dynamic_server_port,
|
get_dynamic_server_port,
|
||||||
|
gt_exists,
|
||||||
|
image_bytes_to_numpy,
|
||||||
|
load_consistency_gt,
|
||||||
wait_for_req_perf_record,
|
wait_for_req_perf_record,
|
||||||
)
|
)
|
||||||
|
|
||||||
@@ -445,6 +453,146 @@ Consider updating perf_baselines.json with the snippets below:
|
|||||||
"""
|
"""
|
||||||
logger.error(output)
|
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> (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(
|
def _save_gt_output(
|
||||||
self,
|
self,
|
||||||
case: DiffusionTestCase,
|
case: DiffusionTestCase,
|
||||||
@@ -903,6 +1051,9 @@ Consider updating perf_baselines.json with the snippets below:
|
|||||||
if case.run_t2v_input_reference_check:
|
if case.run_t2v_input_reference_check:
|
||||||
self._test_t2v_rejects_input_reference(diffusion_server, case)
|
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)
|
# LoRA API functionality test with E2E validation (only for LoRA-enabled cases)
|
||||||
if case.run_lora_basic_api_check:
|
if case.run_lora_basic_api_check:
|
||||||
self._test_lora_api_functionality(diffusion_server, case, generate_fn)
|
self._test_lora_api_functionality(diffusion_server, case, generate_fn)
|
||||||
|
|||||||
@@ -248,6 +248,7 @@ class DiffusionTestCase:
|
|||||||
server_args: DiffusionServerArgs
|
server_args: DiffusionServerArgs
|
||||||
sampling_params: DiffusionSamplingParams
|
sampling_params: DiffusionSamplingParams
|
||||||
run_perf_check: bool = True
|
run_perf_check: bool = True
|
||||||
|
run_consistency_check: bool = True
|
||||||
run_models_api_check: bool = True
|
run_models_api_check: bool = True
|
||||||
run_t2v_input_reference_check: bool = True
|
run_t2v_input_reference_check: bool = True
|
||||||
run_lora_basic_api_check: bool = False
|
run_lora_basic_api_check: bool = False
|
||||||
@@ -777,6 +778,7 @@ if not current_platform.is_hip():
|
|||||||
enable_warmup=False,
|
enable_warmup=False,
|
||||||
),
|
),
|
||||||
HUNYUAN3D_SHAPE_sampling_params,
|
HUNYUAN3D_SHAPE_sampling_params,
|
||||||
|
run_consistency_check=False,
|
||||||
),
|
),
|
||||||
)
|
)
|
||||||
# Skip turbowan on AMD: Triton requires 81920 shared memory, but AMD only has 65536.
|
# Skip turbowan on AMD: Triton requires 81920 shared memory, but AMD only has 65536.
|
||||||
@@ -802,6 +804,7 @@ ONE_GPU_CASES_C = [
|
|||||||
modality="image",
|
modality="image",
|
||||||
),
|
),
|
||||||
T2I_sampling_params,
|
T2I_sampling_params,
|
||||||
|
run_consistency_check=False,
|
||||||
)
|
)
|
||||||
]
|
]
|
||||||
|
|
||||||
@@ -851,6 +854,10 @@ TWO_GPU_CASES_A = [
|
|||||||
custom_validator="video",
|
custom_validator="video",
|
||||||
num_gpus=2,
|
num_gpus=2,
|
||||||
lora_path="Cseti/wan2.2-14B-Arcane_Jinx-lora-v1",
|
lora_path="Cseti/wan2.2-14B-Arcane_Jinx-lora-v1",
|
||||||
|
extras=[
|
||||||
|
"--lora-weight-name",
|
||||||
|
"985347-wan22_14B-low-Nfj1nx-e65.safetensors",
|
||||||
|
],
|
||||||
),
|
),
|
||||||
DiffusionSamplingParams(
|
DiffusionSamplingParams(
|
||||||
prompt="Nfj1nx with blue hair, a woman walking in a cyberpunk city at night",
|
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"],
|
extras=["--pipeline-class-name LTX2TwoStagePipeline"],
|
||||||
),
|
),
|
||||||
T2V_sampling_params,
|
T2V_sampling_params,
|
||||||
|
run_consistency_check=False,
|
||||||
),
|
),
|
||||||
]
|
]
|
||||||
|
|
||||||
|
|||||||
@@ -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
|
||||||
|
)
|
||||||
@@ -7,12 +7,15 @@ import socket
|
|||||||
import subprocess
|
import subprocess
|
||||||
import tempfile
|
import tempfile
|
||||||
import time
|
import time
|
||||||
|
from dataclasses import dataclass
|
||||||
from pathlib import Path
|
from pathlib import Path
|
||||||
|
from typing import TYPE_CHECKING, Any
|
||||||
from urllib.parse import urljoin
|
from urllib.parse import urljoin
|
||||||
|
|
||||||
import cv2
|
import cv2
|
||||||
import httpx
|
import httpx
|
||||||
import numpy as np
|
import numpy as np
|
||||||
|
import requests
|
||||||
from PIL import Image
|
from PIL import Image
|
||||||
|
|
||||||
from sglang.multimodal_gen.runtime.utils.common import get_bool_env_var
|
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,
|
get_diffusion_perf_log_dir,
|
||||||
)
|
)
|
||||||
|
|
||||||
|
if TYPE_CHECKING:
|
||||||
|
from sglang.multimodal_gen.test.server.testcase_configs import DiffusionTestCase
|
||||||
|
|
||||||
logger = init_logger(__name__)
|
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
|
# Common model IDs for diffusion tests
|
||||||
#
|
#
|
||||||
@@ -560,10 +582,232 @@ def validate_video_file(
|
|||||||
), f"Video height mismatch: expected {expected_height}, got {actual_height}"
|
), 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:
|
def output_format_to_ext(output_format: str | None) -> str:
|
||||||
"""Map output_format to file extension. Used by GT naming and consistency check."""
|
"""Map output_format to file extension. Used by GT naming and consistency check."""
|
||||||
if not output_format:
|
if not output_format:
|
||||||
return "png"
|
return "jpg"
|
||||||
of = output_format.lower()
|
of = output_format.lower()
|
||||||
if of == "jpeg":
|
if of == "jpeg":
|
||||||
return "jpg"
|
return "jpg"
|
||||||
@@ -587,6 +831,151 @@ def _consistency_gt_filenames(
|
|||||||
return [f"{case_id}_{n}gpu.{ext}"]
|
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(
|
def extract_key_frames_from_video(
|
||||||
video_bytes: bytes,
|
video_bytes: bytes,
|
||||||
num_frames: int | None = None,
|
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."""
|
"""Convert image bytes to numpy array."""
|
||||||
img = Image.open(io.BytesIO(image_bytes)).convert("RGB")
|
img = Image.open(io.BytesIO(image_bytes)).convert("RGB")
|
||||||
return np.array(img)
|
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
|
||||||
|
|||||||
Reference in New Issue
Block a user