Qwen-Image-Cyrillic-Config / controlnet_accuracy.py
777Radik's picture
Add files using upload-large-folder tool
4837bb3 verified
Raw
History Blame Contribute Delete
6.86 kB
"""OCR evaluation for outputs produced by controlnet_benchmark.py."""
from __future__ import annotations
import argparse
import json
import math
from pathlib import Path
from PIL import Image, ImageChops, ImageDraw, ImageFilter, ImageFont, ImageOps
from benchmarks.text_accuracy import (
centered_text_crop,
character_error_rate,
normalize_ocr,
run_ocr,
)
def _normalized_outline(
image: Image.Image,
size: tuple[int, int] = (384, 128),
*,
light_ink: bool = False,
) -> Image.Image:
grayscale = image.convert("L")
if light_ink:
ink = grayscale.point(lambda value: 255 if value >= 128 else 0)
else:
ink = grayscale.point(lambda value: 255 if value < 128 else 0)
bbox = ink.getbbox()
if bbox is None:
return Image.new("L", size, 0)
cropped = ImageOps.expand(ink.crop(bbox), border=4, fill=0)
normalized = cropped.resize(size, Image.Resampling.LANCZOS)
normalized = normalized.point(lambda value: 255 if value >= 128 else 0)
outer = normalized.filter(ImageFilter.MaxFilter(3))
inner = normalized.filter(ImageFilter.MinFilter(3))
return ImageChops.subtract(outer, inner)
def glyph_similarity(image: Image.Image, control: Image.Image) -> float:
generated = _normalized_outline(image.crop(centered_text_crop(image)))
expected = _normalized_outline(control, light_ink=True)
generated_dilated = generated.filter(ImageFilter.MaxFilter(7))
expected_dilated = expected.filter(ImageFilter.MaxFilter(7))
generated_pixels = sum(1 for value in generated.getdata() if value)
expected_pixels = sum(1 for value in expected.getdata() if value)
if not generated_pixels or not expected_pixels:
return 0.0
generated_hit = sum(
1
for edge, nearby in zip(generated.getdata(), expected_dilated.getdata())
if edge and nearby
)
expected_hit = sum(
1
for edge, nearby in zip(expected.getdata(), generated_dilated.getdata())
if edge and nearby
)
precision = generated_hit / generated_pixels
recall = expected_hit / expected_pixels
return round(2 * precision * recall / max(1e-9, precision + recall), 4)
def light_ink_margin(image: Image.Image) -> int:
ink = image.convert("L").point(lambda value: 255 if value >= 128 else 0)
bbox = ink.getbbox()
if bbox is None:
return -1
width, height = image.size
return min(bbox[0], bbox[1], width - bbox[2], height - bbox[3])
def write_contact_sheets(
results: list[dict[str, object]],
output_dir: Path,
*,
columns: int = 5,
rows: int = 4,
) -> list[str]:
page_size = columns * rows
thumb_size = 240
label_height = 44
font_path = Path("/usr/share/fonts/truetype/dejavu/DejaVuSans.ttf")
font = ImageFont.truetype(str(font_path), 22) if font_path.exists() else ImageFont.load_default()
names: list[str] = []
for page_index in range(math.ceil(len(results) / page_size)):
page_rows = results[page_index * page_size : (page_index + 1) * page_size]
sheet = Image.new(
"RGB",
(columns * thumb_size, rows * (thumb_size + label_height)),
"white",
)
draw = ImageDraw.Draw(sheet)
for cell_index, result in enumerate(page_rows):
image_path = output_dir / str(result["image"])
with Image.open(image_path) as source:
thumbnail = source.convert("RGB")
thumbnail.thumbnail((thumb_size, thumb_size), Image.Resampling.LANCZOS)
x = (cell_index % columns) * thumb_size
y = (cell_index // columns) * (thumb_size + label_height)
sheet.paste(
thumbnail,
(x + (thumb_size - thumbnail.width) // 2, y),
)
draw.text((x + 6, y + thumb_size + 4), str(result["text"]), fill="black", font=font)
name = f"contact-sheet-{page_index + 1:02d}.jpg"
sheet.save(output_dir / name, quality=90)
names.append(name)
return names
def summarize(results: list[dict[str, object]]) -> dict[str, dict[str, float | int]]:
strategies = sorted({str(result["strategy"]) for result in results})
summary: dict[str, dict[str, float | int]] = {}
for strategy in strategies:
rows = [result for result in results if result["strategy"] == strategy]
summary[strategy] = {
"samples": len(rows),
"exact": sum(bool(row["exact_match"]) for row in rows),
"mean_cer": round(sum(float(row["cer"]) for row in rows) / len(rows), 4),
"mean_glyph_similarity": round(
sum(float(row.get("glyph_similarity", 0.0)) for row in rows) / len(rows),
4,
),
"edge_touching_controls": sum(
int(row.get("control_margin_px", 1)) <= 0 for row in rows
),
}
return summary
def run(report_path: Path) -> dict[str, object]:
generation = json.loads(report_path.read_text(encoding="utf-8"))
output_dir = report_path.parent
results: list[dict[str, object]] = []
for row in generation["results"]:
image_path = output_dir / row["image"]
with Image.open(image_path) as image:
crop = centered_text_crop(image)
actual = run_ocr(image_path, "rus+eng", crop=crop)
gold = normalize_ocr(row["text"])
similarity = 0.0
control_margin_px: int | None = None
if row.get("control"):
with Image.open(image_path) as image, Image.open(output_dir / row["control"]) as control:
similarity = glyph_similarity(image, control)
control_margin_px = light_ink_margin(control)
results.append(
{
**row,
"gold": gold,
"ocr": actual,
"exact_match": actual == gold,
"cer": round(character_error_rate(gold, actual), 4),
"glyph_similarity": similarity,
"control_margin_px": control_margin_px,
}
)
contact_sheets = write_contact_sheets(results, output_dir)
report: dict[str, object] = {
"generation_report": report_path.name,
"summary": summarize(results),
"contact_sheets": contact_sheets,
"results": results,
}
(output_dir / "accuracy-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("generation_report", type=Path)
return parser.parse_args()
if __name__ == "__main__":
args = parse_args()
completed = run(args.generation_report.resolve())
print(json.dumps(completed, ensure_ascii=False))