| """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)) |
|
|