| |
| """Template matching, handwriting isolation, and lossless template-text merging.""" |
|
|
| from __future__ import annotations |
|
|
| import json |
| import re |
| from dataclasses import dataclass |
| from pathlib import Path |
| from typing import Any |
|
|
| import cv2 |
| import numpy as np |
| from PIL import Image |
|
|
|
|
| PLACEHOLDER = "{{HANDWRITING}}" |
| VARIANTS = { |
| "adaptive65": {"sigma": 2.5, "floor": 8.0, "percentile": 65.0}, |
| "adaptive75": {"sigma": 4.0, "floor": 10.0, "percentile": 75.0}, |
| "adaptive85": {"sigma": 5.5, "floor": 12.0, "percentile": 85.0}, |
| "fixed20": {"fixed": 20.0}, |
| "residual2": {"mode": "residual", "floor": 2.0, "gain": 2.0}, |
| "residual3": {"mode": "residual", "floor": 3.0, "gain": 3.0}, |
| "residual4": {"mode": "residual", "floor": 4.0, "gain": 4.0}, |
| "residual6": {"mode": "residual", "floor": 6.0, "gain": 4.0}, |
| "highpass08": { |
| "mode": "highpass", |
| "sigma": 0.8, |
| "floor": 1.5, |
| "gain": 4.0, |
| }, |
| "highpass12": { |
| "mode": "highpass", |
| "sigma": 1.2, |
| "floor": 2.0, |
| "gain": 4.0, |
| }, |
| "highpass20": { |
| "mode": "highpass", |
| "sigma": 2.0, |
| "floor": 2.0, |
| "gain": 3.0, |
| }, |
| "highpass30": { |
| "mode": "highpass", |
| "sigma": 3.0, |
| "floor": 2.5, |
| "gain": 3.0, |
| }, |
| } |
|
|
|
|
| @dataclass(frozen=True) |
| class Template: |
| id: str |
| file: str |
| image: np.ndarray |
| text_region: tuple[int, int, int, int] |
| template_text: str |
| blocks: tuple[dict[str, Any], ...] |
|
|
|
|
| @dataclass(frozen=True) |
| class TemplateMatch: |
| template: Template |
| median_abs_delta: float |
| p90_abs_delta: float |
| candidate_count: int |
|
|
|
|
| def normalize_text(value: str) -> str: |
| return " ".join(str(value).split()) |
|
|
|
|
| def merge_template_text(template_text: str, handwriting: str) -> tuple[str, int]: |
| """Insert handwriting once and preserve every Gemini template token.""" |
| template_text = str(template_text) |
| handwriting = normalize_text(handwriting) |
| count = template_text.count(PLACEHOLDER) |
| if count: |
| merged = template_text.replace(PLACEHOLDER, handwriting, 1) |
| merged = merged.replace(PLACEHOLDER, "") |
| elif normalize_text(template_text): |
| merged = f"{template_text} {handwriting}" |
| else: |
| merged = handwriting |
| return normalize_text(merged), count |
|
|
|
|
| class TemplateCatalog: |
| def __init__(self, background_root: Path, ocr_root: Path): |
| background_root = Path(background_root) |
| ocr_root = Path(ocr_root) |
| manifest = json.loads( |
| (background_root / "manifest.json").read_text(encoding="utf-8") |
| ) |
| images: dict[str, np.ndarray] = {} |
| templates: list[Template] = [] |
| for item in manifest: |
| filename = str(item["file"]) |
| if filename not in images: |
| with Image.open(background_root / filename) as opened: |
| images[filename] = np.asarray(opened.convert("RGB")) |
| result: dict[str, Any] = {} |
| ocr_path = ocr_root / f"{item['id']}.json" |
| if ocr_path.is_file(): |
| result = json.loads( |
| ocr_path.read_text(encoding="utf-8") |
| ).get("result", {}) |
| templates.append( |
| Template( |
| id=str(item["id"]), |
| file=filename, |
| image=images[filename], |
| text_region=tuple(int(value) for value in item["text_region"]), |
| template_text=str(result.get("template_text", PLACEHOLDER)), |
| blocks=tuple(result.get("blocks", [])), |
| ) |
| ) |
| self.templates = tuple(templates) |
| self.by_id = {template.id: template for template in templates} |
|
|
| @staticmethod |
| def difference(image: np.ndarray, template: np.ndarray) -> np.ndarray: |
| if image.shape != template.shape: |
| raise ValueError( |
| f"image/template shape mismatch: {image.shape} != {template.shape}" |
| ) |
| return np.max( |
| np.abs(image.astype(np.int16) - template.astype(np.int16)), |
| axis=2, |
| ).astype(np.float32) |
|
|
| def match(self, image: Image.Image | np.ndarray) -> TemplateMatch: |
| actual = ( |
| np.asarray(image.convert("RGB")) |
| if isinstance(image, Image.Image) |
| else np.asarray(image) |
| ) |
| height, width = actual.shape[:2] |
| candidates = [ |
| template |
| for template in self.templates |
| if template.image.shape[:2] == (height, width) |
| ] |
| if not candidates: |
| raise ValueError(f"no template has native size {width}x{height}") |
|
|
| by_file: dict[str, tuple[float, float, np.ndarray]] = {} |
| for template in candidates: |
| if template.file in by_file: |
| continue |
| delta = self.difference(actual, template.image) |
| by_file[template.file] = ( |
| float(np.median(delta)), |
| float(np.percentile(delta, 90)), |
| delta, |
| ) |
| best_file = min( |
| by_file, |
| key=lambda name: ( |
| by_file[name][0], |
| by_file[name][1], |
| name, |
| ), |
| ) |
| same_file = [ |
| template for template in candidates if template.file == best_file |
| ] |
| median, p90, delta = by_file[best_file] |
| if len(same_file) == 1: |
| selected = same_file[0] |
| else: |
| center = float(np.median(delta)) |
| spread = float(np.median(np.abs(delta - center))) |
| threshold = max(8.0, center + 3.0 * max(spread, 1.0)) |
|
|
| def region_energy(template: Template) -> float: |
| x0, y0, x1, y1 = template.text_region |
| x0 = max(0, min(width, x0)) |
| x1 = max(0, min(width, x1)) |
| y0 = max(0, min(height, y0)) |
| y1 = max(0, min(height, y1)) |
| region = delta[y0:y1, x0:x1] |
| if not region.size: |
| return -1.0 |
| return float(np.maximum(region - threshold, 0.0).sum()) / region.size |
|
|
| selected = max(same_file, key=lambda template: (region_energy(template), template.id)) |
| return TemplateMatch( |
| template=selected, |
| median_abs_delta=median, |
| p90_abs_delta=p90, |
| candidate_count=len(candidates), |
| ) |
|
|
| @staticmethod |
| def _mask(delta: np.ndarray, variant: str) -> tuple[np.ndarray, float]: |
| if variant not in VARIANTS: |
| raise ValueError(f"unknown isolation variant: {variant}") |
| options = VARIANTS[variant] |
| if "fixed" in options: |
| threshold = float(options["fixed"]) |
| else: |
| center = float(np.median(delta)) |
| spread = float(np.median(np.abs(delta - center))) |
| threshold = max( |
| float(options["floor"]), |
| center + float(options["sigma"]) * max(spread, 1.0), |
| float(np.percentile(delta, float(options["percentile"]))), |
| ) |
| mask = np.uint8(delta >= threshold) |
| count, labels, stats, _ = cv2.connectedComponentsWithStats(mask, 8) |
| filtered = np.zeros_like(mask) |
| for component in range(1, count): |
| if int(stats[component, cv2.CC_STAT_AREA]) >= 3: |
| filtered[labels == component] = 1 |
| filtered = cv2.dilate( |
| filtered, |
| cv2.getStructuringElement(cv2.MORPH_ELLIPSE, (3, 3)), |
| iterations=1, |
| ) |
| return filtered.astype(bool), threshold |
|
|
| @staticmethod |
| def _residual( |
| actual: np.ndarray, |
| template: np.ndarray, |
| variant: str, |
| ) -> tuple[np.ndarray, dict[str, float]]: |
| options = VARIANTS[variant] |
| channel_delta = template.astype(np.float32) - actual.astype(np.float32) |
| offsets = np.median(channel_delta.reshape(-1, 3), axis=0) |
| corrected = channel_delta - offsets.reshape(1, 1, 3) |
| signal = np.maximum(corrected.max(axis=2) - float(options["floor"]), 0.0) |
| signal = np.clip(signal * float(options["gain"]), 0.0, 245.0) |
| raw_mask = np.uint8(signal >= 3.0) |
| count, labels, stats, _ = cv2.connectedComponentsWithStats(raw_mask, 8) |
| mask = np.zeros_like(raw_mask) |
| for component in range(1, count): |
| if int(stats[component, cv2.CC_STAT_AREA]) >= 2: |
| mask[labels == component] = 1 |
| darkness = np.zeros_like(signal) |
| darkness[mask.astype(bool)] = signal[mask.astype(bool)] |
| darkness = cv2.dilate( |
| darkness, |
| cv2.getStructuringElement(cv2.MORPH_ELLIPSE, (2, 2)), |
| iterations=1, |
| ) |
| output = np.full_like(actual, 255) |
| value = np.uint8(np.clip(255.0 - darkness, 0.0, 255.0)) |
| retained = darkness >= 3.0 |
| output[retained] = np.repeat(value[:, :, None], 3, axis=2)[retained] |
| return output, { |
| "threshold": float(options["floor"]), |
| "gain": float(options["gain"]), |
| "kept_fraction": float(retained.mean()), |
| "channel_offset_r": float(offsets[0]), |
| "channel_offset_g": float(offsets[1]), |
| "channel_offset_b": float(offsets[2]), |
| } |
|
|
| @staticmethod |
| def _highpass( |
| actual: np.ndarray, |
| template: np.ndarray, |
| variant: str, |
| ) -> tuple[np.ndarray, dict[str, float]]: |
| options = VARIANTS[variant] |
| sigma = float(options["sigma"]) |
| actual_gray = cv2.cvtColor(actual, cv2.COLOR_RGB2GRAY).astype(np.float32) |
| template_gray = cv2.cvtColor(template, cv2.COLOR_RGB2GRAY).astype(np.float32) |
| actual_darkness = cv2.GaussianBlur(actual_gray, (0, 0), sigma) - actual_gray |
| template_darkness = ( |
| cv2.GaussianBlur(template_gray, (0, 0), sigma) - template_gray |
| ) |
| signal = np.maximum( |
| actual_darkness - template_darkness - float(options["floor"]), |
| 0.0, |
| ) |
| signal = np.clip(signal * float(options["gain"]), 0.0, 245.0) |
| raw_mask = np.uint8(signal >= 3.0) |
| count, labels, stats, _ = cv2.connectedComponentsWithStats(raw_mask, 8) |
| mask = np.zeros_like(raw_mask) |
| for component in range(1, count): |
| if int(stats[component, cv2.CC_STAT_AREA]) >= 2: |
| mask[labels == component] = 1 |
| darkness = np.zeros_like(signal) |
| darkness[mask.astype(bool)] = signal[mask.astype(bool)] |
| darkness = cv2.dilate( |
| darkness, |
| cv2.getStructuringElement(cv2.MORPH_ELLIPSE, (2, 2)), |
| iterations=1, |
| ) |
| retained = darkness >= 3.0 |
| value = np.uint8(np.clip(255.0 - darkness, 0.0, 255.0)) |
| output = np.repeat(value[:, :, None], 3, axis=2) |
| return output, { |
| "threshold": float(options["floor"]), |
| "gain": float(options["gain"]), |
| "sigma": sigma, |
| "kept_fraction": float(retained.mean()), |
| } |
|
|
| @classmethod |
| def _render( |
| cls, |
| actual: np.ndarray, |
| template: np.ndarray, |
| variant: str, |
| ) -> tuple[np.ndarray, dict[str, float]]: |
| options = VARIANTS.get(variant) |
| if options is None: |
| raise ValueError(f"unknown isolation variant: {variant}") |
| delta = cls.difference(actual, template) |
| if options.get("mode") == "residual": |
| output, metrics = cls._residual(actual, template, variant) |
| elif options.get("mode") == "highpass": |
| output, metrics = cls._highpass(actual, template, variant) |
| else: |
| mask, threshold = cls._mask(delta, variant) |
| output = np.full_like(actual, 255) |
| output[mask] = actual[mask] |
| metrics = { |
| "threshold": threshold, |
| "kept_fraction": float(mask.mean()), |
| } |
| metrics.update( |
| { |
| "median_abs_delta": float(np.median(delta)), |
| "p90_abs_delta": float(np.percentile(delta, 90)), |
| } |
| ) |
| return output, metrics |
|
|
| def isolate_with_template( |
| self, |
| image: Image.Image | np.ndarray, |
| template: Template, |
| variant: str, |
| ) -> tuple[Image.Image, dict[str, float]]: |
| actual = ( |
| np.asarray(image.convert("RGB")) |
| if isinstance(image, Image.Image) |
| else np.asarray(image) |
| ) |
| output, metrics = self._render(actual, template.image, variant) |
| return Image.fromarray(output, mode="RGB"), metrics |
|
|
| def isolate( |
| self, |
| image: Image.Image | np.ndarray, |
| variant: str, |
| ) -> tuple[Image.Image, TemplateMatch, dict[str, float]]: |
| match = self.match(image) |
| isolated, metrics = self.isolate_with_template( |
| image, |
| match.template, |
| variant, |
| ) |
| return isolated, match, metrics |
|
|
| def isolate_crop( |
| self, |
| image: Image.Image | np.ndarray, |
| template_id: str, |
| crop_bbox: list[int] | tuple[int, int, int, int], |
| variant: str, |
| ) -> tuple[Image.Image, dict[str, float]]: |
| actual = ( |
| np.asarray(image.convert("RGB")) |
| if isinstance(image, Image.Image) |
| else np.asarray(image) |
| ) |
| template = self.by_id[template_id] |
| x0, y0, x1, y1 = (int(value) for value in crop_bbox) |
| background = template.image[y0:y1, x0:x1] |
| if background.shape != actual.shape: |
| background = cv2.resize( |
| background, |
| (actual.shape[1], actual.shape[0]), |
| interpolation=cv2.INTER_LANCZOS4, |
| ) |
| output, metrics = self._render(actual, background, variant) |
| return Image.fromarray(output, mode="RGB"), metrics |
|
|
|
|
| def index_from_id(value: str) -> int: |
| match = re.fullmatch(r"handwriting:(\d+)", str(value)) |
| if not match: |
| raise ValueError(f"invalid handwriting row id: {value}") |
| return int(match.group(1)) |
|
|