Qwen-Image-Cyrillic-Config / controlnet_benchmark.py
777Radik's picture
Add files using upload-large-folder tool
4837bb3 verified
Raw
History Blame Contribute Delete
11.5 kB
"""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))