"""Zero-shot Cyrillic benchmark for Qwen-Image Blockwise ControlNet Canny.""" from __future__ import annotations import argparse import csv import json import re from datetime import UTC, datetime from pathlib import Path import torch from diffsynth.pipelines.qwen_image import ControlNetInput, ModelConfig, QwenImagePipeline from PIL import Image, ImageChops, ImageFilter from lora_server import model_config DEFAULT_TEXTS = ("ЁЖИК", "ПОДЪЁМ", "СЪЕЗД", "ЩЁТКА") CAPTION_TEXT = re.compile(r'\btext "([^"]+)" on\b') def glyph_edge(mask: Image.Image, size: int) -> Image.Image: """Convert a filled glyph raster into a Canny-like white outline.""" binary = mask.convert("L").resize((size, size), Image.Resampling.LANCZOS) binary = binary.point(lambda value: 255 if value >= 128 else 0) inner = binary.filter(ImageFilter.MinFilter(3)) edge = ImageChops.subtract(binary, inner).filter(ImageFilter.MaxFilter(3)) return edge.convert("RGB") def glyph_control(mask: Image.Image, size: int, mode: str) -> Image.Image: if mode == "edge": return glyph_edge(mask, size) if mode == "filled": filled = mask.convert("L").resize((size, size), Image.Resampling.LANCZOS) return filled.point(lambda value: 255 if value >= 128 else 0).convert("RGB") raise ValueError(f"unsupported control mode: {mode}") def fit_glyph_mask(mask: Image.Image, size: int) -> Image.Image: binary = mask.convert("L").point(lambda value: 255 if value >= 128 else 0) bbox = binary.getbbox() if bbox is None: raise ValueError("glyph mask is empty") cropped = binary.crop(bbox) scale = min(size * 0.8 / cropped.width, size * 0.4 / cropped.height) resized = cropped.resize( (max(1, round(cropped.width * scale)), max(1, round(cropped.height * scale))), Image.Resampling.LANCZOS, ) canvas = Image.new("L", (size, size), 0) canvas.paste(resized, ((size - resized.width) // 2, (size - resized.height) // 2)) return canvas def heldout_glyphs(dataset_dir: Path, texts: tuple[str, ...]) -> dict[str, Path]: with (dataset_dir / "heldout.csv").open(encoding="utf-8", newline="") as handle: rows = list(csv.DictReader(handle)) found: dict[str, Path] = {} for row in rows: match = CAPTION_TEXT.search(row.get("prompt", "")) if match and match.group(1) in texts: found[match.group(1)] = dataset_dir / row["glyph"] missing = [text for text in texts if text not in found] if missing: raise ValueError(f"held-out glyphs missing: {missing}") return found def heldout_texts(dataset_dir: Path) -> tuple[str, ...]: with (dataset_dir / "heldout.csv").open(encoding="utf-8", newline="") as handle: rows = list(csv.DictReader(handle)) texts: list[str] = [] for row in rows: match = CAPTION_TEXT.search(row.get("prompt", "")) if match is None: raise ValueError(f"held-out prompt has no exact text: {row.get('prompt')!r}") texts.append(match.group(1)) if not texts or len(set(texts)) != len(texts): raise ValueError("held-out texts must be non-empty and unique") return tuple(texts) def low_vram_model_config(path: str | list[str]) -> ModelConfig: """Official DiffSynth disk-offload profile with FP8 weight onload.""" return ModelConfig( path=path, offload_dtype="disk", offload_device="disk", onload_dtype=torch.float8_e4m3fn, onload_device="cpu", preparing_dtype=torch.float8_e4m3fn, preparing_device="cuda", computation_dtype=torch.bfloat16, computation_device="cuda", ) def resolve_controlnet_path(path: Path) -> Path: checkpoint = path / "model.safetensors" if path.is_dir() else path if not checkpoint.is_file(): raise FileNotFoundError(f"ControlNet checkpoint missing: {checkpoint}") return checkpoint def resolve_transformer_files(base_dir: Path, transformer_dir: Path | None = None) -> list[str]: directory = transformer_dir if transformer_dir is not None else base_dir / "transformer" files = sorted(str(path) for path in directory.glob("*.safetensors")) if not files: raise FileNotFoundError(f"transformer checkpoints missing: {directory}") return files def load_pipeline( base_dir: Path, controlnet_path: Path, vram_limit_gib: float | None = None, transformer_dir: Path | None = None, ) -> QwenImagePipeline: transformer = resolve_transformer_files(base_dir, transformer_dir) text_encoder = sorted(str(path) for path in (base_dir / "text_encoder").glob("*.safetensors")) controlnet = resolve_controlnet_path(controlnet_path) required = [ *map(Path, transformer), *map(Path, text_encoder), base_dir / "vae" / "diffusion_pytorch_model.safetensors", base_dir / "tokenizer", controlnet, ] missing = [str(path) for path in required if not path.exists()] if missing: raise FileNotFoundError(f"model components missing: {missing}") paths: list[str | list[str]] = [ transformer, text_encoder, str(base_dir / "vae" / "diffusion_pytorch_model.safetensors"), str(controlnet), ] configs = ( [low_vram_model_config(path) for path in paths] if vram_limit_gib is not None else [ model_config(transformer, "fp8"), model_config(text_encoder, "fp8"), model_config(paths[2], "fp8"), model_config(paths[3], "bf16"), ] ) return QwenImagePipeline.from_pretrained( torch_dtype=torch.bfloat16, device="cuda", model_configs=configs, tokenizer_config=ModelConfig(path=str(base_dir / "tokenizer")), vram_limit=vram_limit_gib, ) def control_variants(scales: tuple[float, ...]) -> tuple[tuple[str, float], ...]: if not scales: raise ValueError("at least one control scale is required") if any(scale <= 0 for scale in scales): raise ValueError("control scales must be positive") if len(set(scales)) != len(scales): raise ValueError("control scales must be unique") if len(scales) == 1: return (("controlnet", scales[0]),) return tuple((f"controlnet-{scale:g}", scale) for scale in scales) def run(args: argparse.Namespace) -> dict[str, object]: output_dir = args.output_dir.resolve() output_dir.mkdir(parents=True, exist_ok=True) glyphs = heldout_glyphs(args.dataset_dir, args.texts) pipe = load_pipeline( args.base_dir, args.controlnet_path, args.vram_limit_gib, args.transformer_dir, ) results: list[dict[str, object]] = [] for index, text in enumerate(args.texts): seed = args.seed + index source = Image.open(glyphs[text]) if args.fit_control: source = fit_glyph_mask(source, args.size) edge = glyph_control(source, args.size, args.control_mode) edge_path = output_dir / f"{index + 1:02d}-control.png" edge.save(edge_path) prompt = ( "A clean professional typographic poster on a plain neutral background. " f'Display exactly the single centered Russian word "{text}". ' "No other letters, words, logos, or decorations." ) common = { "prompt": prompt, "negative_prompt": "misspelled text, extra letters, duplicated glyphs, watermark", "height": args.size, "width": args.size, "seed": seed, "num_inference_steps": args.steps, } baseline_variants: tuple[tuple[str, float | None], ...] = ( () if args.skip_baseline else (("baseline", None),) ) variants = (*baseline_variants, *control_variants(args.control_scales)) for strategy, scale in variants: controls = ( None if scale is None else [ControlNetInput(image=edge, scale=scale)] ) image = pipe(**common, blockwise_controlnet_inputs=controls) image_path = output_dir / f"{strategy}-{index + 1:02d}.png" image.save(image_path) results.append( { "strategy": strategy, "text": text, "seed": seed, "prompt": prompt, "image": image_path.name, "control": edge_path.name if controls else None, "control_scale": scale, } ) report: dict[str, object] = { "generated_at": datetime.now(UTC).isoformat(), "base_dir": str(args.base_dir), "transformer_dir": str(args.transformer_dir) if args.transformer_dir else None, "controlnet_path": str(args.controlnet_path), "size": args.size, "steps": args.steps, "seed": args.seed, "control_scales": args.control_scales, "control_mode": args.control_mode, "fit_control": args.fit_control, "vram_limit_gib": args.vram_limit_gib, "skip_baseline": args.skip_baseline, "results": results, } (output_dir / "generation-report.json").write_text( json.dumps(report, ensure_ascii=False, indent=2), encoding="utf-8", ) return report def parse_args() -> argparse.Namespace: parser = argparse.ArgumentParser(description=__doc__) parser.add_argument("--base-dir", type=Path, default=Path("/models/Qwen-Image-2512")) parser.add_argument( "--transformer-dir", type=Path, help="Override transformer shards while reusing the other --base-dir components.", ) parser.add_argument( "--controlnet-dir", "--controlnet-path", dest="controlnet_path", type=Path, default=Path("/models/Qwen-Image-Blockwise-ControlNet-Canny"), ) parser.add_argument("--dataset-dir", type=Path, default=Path("/workspace/dataset")) parser.add_argument("--output-dir", type=Path, default=Path("/workspace/controlnet-benchmark")) text_group = parser.add_mutually_exclusive_group() text_group.add_argument("--text", action="append", dest="texts") text_group.add_argument("--all-heldout", action="store_true") parser.add_argument("--size", type=int, default=512) parser.add_argument("--control-mode", choices=("edge", "filled"), default="edge") parser.add_argument("--fit-control", action="store_true") parser.add_argument("--steps", type=int, default=20) parser.add_argument("--seed", type=int, default=2512) parser.add_argument("--skip-baseline", action="store_true") parser.add_argument( "--vram-limit-gib", type=float, help="enable official disk-offload mode and cap managed model VRAM", ) parser.add_argument( "--control-scale", type=float, action="append", dest="control_scales", help="repeat to benchmark multiple scales with one model load", ) args = parser.parse_args() if args.all_heldout: args.texts = heldout_texts(args.dataset_dir) else: args.texts = tuple(args.texts or DEFAULT_TEXTS) args.control_scales = tuple(args.control_scales or (1.0,)) return args if __name__ == "__main__": completed = run(parse_args()) print(json.dumps(completed, ensure_ascii=False))