From f5ed2687ec6a748dd9922e14b88c97bac8567612 Mon Sep 17 00:00:00 2001 From: Mick Date: Fri, 22 May 2026 19:07:20 +0800 Subject: [PATCH] [diffusion] CI: guard gt publishing from corrupt output (#26044) --- .github/workflows/diffusion-ci-gt-gen.yml | 36 +++ .../utils/diffusion/publish_diffusion_gt.py | 296 +++++++++++++++++- 2 files changed, 328 insertions(+), 4 deletions(-) diff --git a/.github/workflows/diffusion-ci-gt-gen.yml b/.github/workflows/diffusion-ci-gt-gen.yml index 39b267a0e..2c935f0a1 100644 --- a/.github/workflows/diffusion-ci-gt-gen.yml +++ b/.github/workflows/diffusion-ci-gt-gen.yml @@ -294,6 +294,15 @@ jobs: python/${{ env.OUTPUT_NAME }}/official_ltx23_manifest.json retention-days: 7 + - name: Validate generated GT images + env: + GITHUB_TOKEN: ${{ secrets.GH_PAT_FOR_NIGHTLY_CI_DATA }} + run: | + python scripts/ci/utils/diffusion/publish_diffusion_gt.py \ + --source-dir python/${{ env.OUTPUT_NAME }} \ + --target-dir "${{ env.PUBLISH_TARGET_DIR }}" \ + --check-only + - name: Publish official GT images to sgl-project/ci-data env: GITHUB_TOKEN: ${{ secrets.GH_PAT_FOR_NIGHTLY_CI_DATA }} @@ -394,6 +403,15 @@ jobs: path: python/${{ env.OUTPUT_NAME }} retention-days: 7 + - name: Validate generated GT images + env: + GITHUB_TOKEN: ${{ secrets.GH_PAT_FOR_NIGHTLY_CI_DATA }} + run: | + python scripts/ci/utils/diffusion/publish_diffusion_gt.py \ + --source-dir python/${{ env.OUTPUT_NAME }} \ + --target-dir "${{ env.PUBLISH_TARGET_DIR }}" \ + --check-only + - name: Publish GT images to sgl-project/ci-data env: GITHUB_TOKEN: ${{ secrets.GH_PAT_FOR_NIGHTLY_CI_DATA }} @@ -460,6 +478,15 @@ jobs: path: python/${{ env.OUTPUT_NAME }} retention-days: 7 + - name: Validate generated GT images + env: + GITHUB_TOKEN: ${{ secrets.GH_PAT_FOR_NIGHTLY_CI_DATA }} + run: | + python scripts/ci/utils/diffusion/publish_diffusion_gt.py \ + --source-dir python/${{ env.OUTPUT_NAME }} \ + --target-dir "${{ env.PUBLISH_TARGET_DIR }}" \ + --check-only + - name: Publish GT images to sgl-project/ci-data env: GITHUB_TOKEN: ${{ secrets.GH_PAT_FOR_NIGHTLY_CI_DATA }} @@ -526,6 +553,15 @@ jobs: path: python/${{ env.OUTPUT_NAME }} retention-days: 7 + - name: Validate generated GT images + env: + GITHUB_TOKEN: ${{ secrets.GH_PAT_FOR_NIGHTLY_CI_DATA }} + run: | + python scripts/ci/utils/diffusion/publish_diffusion_gt.py \ + --source-dir python/${{ env.OUTPUT_NAME }} \ + --target-dir "${{ env.PUBLISH_TARGET_DIR }}" \ + --check-only + - name: Publish GT images to sgl-project/ci-data env: GITHUB_TOKEN: ${{ secrets.GH_PAT_FOR_NIGHTLY_CI_DATA }} diff --git a/scripts/ci/utils/diffusion/publish_diffusion_gt.py b/scripts/ci/utils/diffusion/publish_diffusion_gt.py index 9912a2c19..b2b3745b1 100644 --- a/scripts/ci/utils/diffusion/publish_diffusion_gt.py +++ b/scripts/ci/utils/diffusion/publish_diffusion_gt.py @@ -4,13 +4,19 @@ via the GitHub API (same pattern as publish_traces.py). """ import argparse +import base64 import hashlib +import io import json import os import sys +from dataclasses import dataclass from pathlib import Path from urllib.error import HTTPError +import numpy as np +from PIL import Image, ImageFilter + # Reuse GitHub API helpers from publish_traces. # Support both direct script execution and package-style imports. if __package__: @@ -47,6 +53,32 @@ BRANCH = "main" DEFAULT_TARGET_DIR = "diffusion-ci/consistency_gt/sglang_generated" IMAGE_EXTENSIONS = {".png", ".jpg", ".jpeg", ".webp"} +QUALITY_MAX_SIDE = 256 +LOW_DETAIL_STD_THRESHOLD = 0.075 +LOW_DETAIL_ENTROPY_THRESHOLD = 0.55 +LOW_DETAIL_BLUR_RESIDUAL_THRESHOLD = 0.035 +LOW_DETAIL_GRADIENT_P95_THRESHOLD = 0.045 +RANDOM_NOISE_CORRELATION_THRESHOLD = 0.55 +RANDOM_NOISE_LOW_FREQUENCY_THRESHOLD = 0.20 +RANDOM_NOISE_BLUR_RESIDUAL_THRESHOLD = 0.045 +OLD_NEW_MIN_SSIM = 0.20 +OLD_NEW_MAX_MEAN_ABS_DIFF = 45.0 + + +@dataclass(frozen=True) +class ImageQualityMetrics: + luminance_std: float + entropy: float + blur_residual: float + gradient_p95: float + neighbor_correlation: float + low_frequency_ratio: float + + +@dataclass(frozen=True) +class OldNewMetrics: + ssim: float + mean_abs_diff: float def collect_images(source_dir, target_dir): @@ -72,6 +104,15 @@ def git_blob_sha(content): def get_remote_blob_shas(repo_owner, repo_name, target_dir, token): + return { + path: item["sha"] + for path, item in get_remote_image_entries( + repo_owner, repo_name, target_dir, token + ).items() + } + + +def get_remote_image_entries(repo_owner, repo_name, target_dir, token): url = ( f"https://api.github.com/repos/{repo_owner}/{repo_name}/contents/" f"{target_dir}?ref={BRANCH}" @@ -84,9 +125,11 @@ def get_remote_blob_shas(repo_owner, repo_name, target_dir, token): raise entries = json.loads(response) return { - item["path"]: item["sha"] + item["path"]: item for item in entries - if item.get("type") == "file" and "sha" in item + if item.get("type") == "file" + and "sha" in item + and os.path.splitext(item["path"])[1].lower() in IMAGE_EXTENSIONS } @@ -98,6 +141,237 @@ def filter_changed_files(files, remote_blob_shas): ] +def get_remote_blob_content(repo_owner, repo_name, blob_sha, token): + url = f"https://api.github.com/repos/{repo_owner}/{repo_name}/git/blobs/{blob_sha}" + response = make_github_request(url, token) + blob = json.loads(response) + if blob.get("encoding") != "base64": + raise ValueError( + f"Unexpected blob encoding for {blob_sha}: {blob.get('encoding')}" + ) + return base64.b64decode(blob["content"]) + + +def _load_quality_image(content): + with Image.open(io.BytesIO(content)) as image: + image = image.convert("RGB") + image.thumbnail((QUALITY_MAX_SIDE, QUALITY_MAX_SIDE), Image.Resampling.BICUBIC) + return image.copy() + + +def _image_to_rgb_array(image): + return np.asarray(image, dtype=np.float32) + + +def _luminance(rgb): + return 0.299 * rgb[..., 0] + 0.587 * rgb[..., 1] + 0.114 * rgb[..., 2] + + +def _neighbor_correlation(luma): + def corr(a, b): + a = a.ravel() + b = b.ravel() + if a.std() < 1e-6 or b.std() < 1e-6: + return 1.0 + return float(np.corrcoef(a, b)[0, 1]) + + return (corr(luma[:, 1:], luma[:, :-1]) + corr(luma[1:, :], luma[:-1, :])) / 2 + + +def _low_frequency_ratio(luma): + centered = luma - luma.mean() + power = np.abs(np.fft.fftshift(np.fft.fft2(centered))) ** 2 + total_power = power.sum() + if total_power < 1e-12: + return 0.0 + + height, width = luma.shape + y, x = np.ogrid[:height, :width] + center_y = height // 2 + center_x = width // 2 + radius = np.sqrt((y - center_y) ** 2 + (x - center_x) ** 2) + low_frequency_radius = min(height, width) * 0.08 + return float(power[radius <= low_frequency_radius].sum() / total_power) + + +def compute_image_quality_metrics(content): + image = _load_quality_image(content) + rgb = _image_to_rgb_array(image) + luma = _luminance(rgb) / 255.0 + + gradients = np.concatenate( + [ + np.abs(np.diff(luma, axis=1)).ravel(), + np.abs(np.diff(luma, axis=0)).ravel(), + ] + ) + histogram, _ = np.histogram(luma, bins=32, range=(0, 1)) + probabilities = histogram / histogram.sum() + nonzero_probabilities = probabilities[probabilities > 0] + entropy = float( + -(nonzero_probabilities * np.log2(nonzero_probabilities)).sum() / 5.0 + ) + blurred = _image_to_rgb_array(image.filter(ImageFilter.GaussianBlur(radius=3))) + + return ImageQualityMetrics( + luminance_std=float(luma.std()), + entropy=entropy, + blur_residual=float(np.mean(np.abs(rgb - blurred)) / 255.0), + gradient_p95=float(np.percentile(gradients, 95)), + neighbor_correlation=_neighbor_correlation(luma), + low_frequency_ratio=_low_frequency_ratio(luma), + ) + + +def get_quality_failure_reasons(metrics): + reasons = [] + low_detail_static = ( + metrics.luminance_std < LOW_DETAIL_STD_THRESHOLD + and metrics.entropy < LOW_DETAIL_ENTROPY_THRESHOLD + and ( + metrics.blur_residual < LOW_DETAIL_BLUR_RESIDUAL_THRESHOLD + or metrics.gradient_p95 < LOW_DETAIL_GRADIENT_P95_THRESHOLD + ) + ) + high_frequency_noise = ( + metrics.neighbor_correlation < RANDOM_NOISE_CORRELATION_THRESHOLD + and metrics.low_frequency_ratio < RANDOM_NOISE_LOW_FREQUENCY_THRESHOLD + and metrics.blur_residual > RANDOM_NOISE_BLUR_RESIDUAL_THRESHOLD + ) + if low_detail_static: + reasons.append("low-contrast low-detail output") + if high_frequency_noise: + reasons.append("high-frequency random noise") + return reasons + + +def _resize_for_old_new_compare(content, size=None): + with Image.open(io.BytesIO(content)) as image: + image = image.convert("RGB") + if size is None: + image.thumbnail( + (QUALITY_MAX_SIDE, QUALITY_MAX_SIDE), Image.Resampling.BICUBIC + ) + else: + image = image.resize(size, Image.Resampling.BICUBIC) + return _image_to_rgb_array(image) + + +def compute_old_new_metrics(old_content, new_content): + old_rgb = _resize_for_old_new_compare(old_content) + new_rgb = _resize_for_old_new_compare( + new_content, size=(old_rgb.shape[1], old_rgb.shape[0]) + ) + old_luma = _luminance(old_rgb) / 255.0 + new_luma = _luminance(new_rgb) / 255.0 + + old_mean = old_luma.mean() + new_mean = new_luma.mean() + old_variance = old_luma.var() + new_variance = new_luma.var() + covariance = ((old_luma - old_mean) * (new_luma - new_mean)).mean() + c1 = 0.01**2 + c2 = 0.03**2 + ssim = ( + (2 * old_mean * new_mean + c1) + * (2 * covariance + c2) + / ((old_mean**2 + new_mean**2 + c1) * (old_variance + new_variance + c2)) + ) + + return OldNewMetrics( + ssim=float(ssim), + mean_abs_diff=float(np.abs(old_rgb - new_rgb).mean()), + ) + + +def _format_quality_metrics(metrics): + return ( + f"std={metrics.luminance_std:.4f}, entropy={metrics.entropy:.4f}, " + f"blur_residual={metrics.blur_residual:.4f}, " + f"gradient_p95={metrics.gradient_p95:.4f}, " + f"neighbor_corr={metrics.neighbor_correlation:.4f}, " + f"low_freq={metrics.low_frequency_ratio:.4f}" + ) + + +def _format_old_new_metrics(metrics): + return f"ssim={metrics.ssim:.4f}, mean_abs_diff={metrics.mean_abs_diff:.2f}" + + +def validate_gt_files(files_to_upload, changed_files, remote_image_entries, token): + failures = [] + for path, content in files_to_upload: + quality_metrics = compute_image_quality_metrics(content) + quality_reasons = get_quality_failure_reasons(quality_metrics) + if quality_reasons: + failures.append( + f"{path}: {', '.join(quality_reasons)} " + f"({_format_quality_metrics(quality_metrics)})" + ) + + for path, content in changed_files: + remote_entry = remote_image_entries.get(path) + if not remote_entry: + continue + + old_content = get_remote_blob_content( + REPO_OWNER, REPO_NAME, remote_entry["sha"], token + ) + old_quality_metrics = compute_image_quality_metrics(old_content) + old_quality_reasons = get_quality_failure_reasons(old_quality_metrics) + if old_quality_reasons: + print( + f"Skipping old/new drift check for {path} because existing GT is " + f"already suspicious: {', '.join(old_quality_reasons)} " + f"({_format_quality_metrics(old_quality_metrics)})" + ) + continue + + old_new_metrics = compute_old_new_metrics(old_content, content) + if ( + old_new_metrics.ssim < OLD_NEW_MIN_SSIM + and old_new_metrics.mean_abs_diff > OLD_NEW_MAX_MEAN_ABS_DIFF + ): + failures.append( + f"{path}: changed too far from existing GT " + f"({_format_old_new_metrics(old_new_metrics)})" + ) + + if not failures: + print( + f"GT quality gate passed for {len(files_to_upload)} generated image(s) " + f"and {len(changed_files)} changed image(s)." + ) + return + + print("GT quality gate failed; refusing to publish suspicious image updates:") + for failure in failures: + print(f" - {failure}") + sys.exit(1) + + +def check_quality(source_dir, target_dir=None): + target_dir = target_dir or DEFAULT_TARGET_DIR + token = os.getenv("GITHUB_TOKEN") + if not token: + print("Error: GITHUB_TOKEN environment variable not set") + sys.exit(1) + + files_to_upload = collect_images(source_dir, target_dir) + if not files_to_upload: + print(f"No image files found in {source_dir}") + return + + remote_image_entries = get_remote_image_entries( + REPO_OWNER, REPO_NAME, target_dir, token + ) + remote_blob_shas = { + path: item["sha"] for path, item in remote_image_entries.items() + } + changed_files = filter_changed_files(files_to_upload, remote_blob_shas) + validate_gt_files(files_to_upload, changed_files, remote_image_entries, token) + + def publish(source_dir, target_dir=None): target_dir = target_dir or DEFAULT_TARGET_DIR token = os.getenv("GITHUB_TOKEN") @@ -129,10 +403,16 @@ def publish(source_dir, target_dir=None): try: branch_sha = get_branch_sha(REPO_OWNER, REPO_NAME, BRANCH, token) tree_sha = get_tree_sha(REPO_OWNER, REPO_NAME, branch_sha, token) - remote_blob_shas = get_remote_blob_shas( + remote_image_entries = get_remote_image_entries( REPO_OWNER, REPO_NAME, target_dir, token ) + remote_blob_shas = { + path: item["sha"] for path, item in remote_image_entries.items() + } changed_files = filter_changed_files(files_to_upload, remote_blob_shas) + validate_gt_files( + files_to_upload, changed_files, remote_image_entries, token + ) if not changed_files: print("No image changes to publish.") return @@ -211,8 +491,16 @@ def main(): default=None, help=f"Target directory in the remote repo (default: {DEFAULT_TARGET_DIR})", ) + parser.add_argument( + "--check-only", + action="store_true", + help="Validate generated GT images without publishing them", + ) args = parser.parse_args() - publish(args.source_dir, args.target_dir) + if args.check_only: + check_quality(args.source_dir, args.target_dir) + else: + publish(args.source_dir, args.target_dir) if __name__ == "__main__":