Spaces:
Build error
Build error
Download src/editable_image/shape_recognition.py from tonigi/make_editable_image: direct link, hf CLI and curl.
- Browser
- Download file 31 kB
-
https://huggingface.co/spaces/tonigi/make_editable_image/resolve/main/src/editable_image/shape_recognition.py
- Command line
-
hf download hf://spaces/tonigi/make_editable_image/src/editable_image/shape_recognition.py
-
curl -L -o shape_recognition.py https://huggingface.co/spaces/tonigi/make_editable_image/resolve/main/src/editable_image/shape_recognition.py
31 kB
| from __future__ import annotations | |
| from dataclasses import dataclass | |
| from math import atan2, degrees | |
| import cv2 | |
| import numpy as np | |
| from .models import ( | |
| Box, | |
| EllipseGeometry, | |
| GradientStop, | |
| LinearGradientPaint, | |
| PathGeometry, | |
| Point, | |
| PolygonGeometry, | |
| RectGeometry, | |
| ShapeRegion, | |
| ShapeStyle, | |
| SolidPaint, | |
| ) | |
| class Candidate: | |
| shape: ShapeRegion | |
| mask: np.ndarray | |
| score: float | |
| def detect_shapes(rgba: np.ndarray, protected_mask: np.ndarray) -> list[ShapeRegion]: | |
| """Recognize low-complexity diagram primitives without tracing illustrations.""" | |
| rgb = rgba[:, :, :3] | |
| height, width = rgb.shape[:2] | |
| clusters = color_clusters(rgb, protected_mask) | |
| candidates: list[Candidate] = smoothed_container_candidates(rgb, protected_mask) | |
| candidates.extend( | |
| arrow_candidates( | |
| rgb, | |
| protected_mask, | |
| [*clusters, *arrow_color_masks(rgb, protected_mask)], | |
| ) | |
| ) | |
| for cluster_index, cluster_mask in enumerate(clusters): | |
| cleaned = cv2.morphologyEx(cluster_mask, cv2.MORPH_OPEN, np.ones((2, 2), np.uint8)) | |
| cleaned = cv2.morphologyEx(cleaned, cv2.MORPH_CLOSE, np.ones((5, 5), np.uint8)) | |
| contours, _ = cv2.findContours(cleaned, cv2.RETR_EXTERNAL, cv2.CHAIN_APPROX_SIMPLE) | |
| for contour in contours: | |
| candidate = fit_cluster_contour( | |
| rgb, | |
| cleaned, | |
| contour, | |
| protected_mask, | |
| width, | |
| height, | |
| cluster_index, | |
| ) | |
| if candidate is not None: | |
| candidates.append(candidate) | |
| candidates.extend(edge_rectangle_candidates(rgb, protected_mask)) | |
| accepted = deduplicate(candidates) | |
| containers = [candidate for candidate in accepted if candidate.shape.is_container] | |
| accepted = [ | |
| candidate | |
| for candidate in accepted | |
| if not ( | |
| isinstance(candidate.shape.geometry, EllipseGeometry) | |
| and any( | |
| cv2.countNonZero(cv2.bitwise_and(candidate.mask, container.mask)) | |
| / max(cv2.countNonZero(candidate.mask), 1) | |
| >= 0.9 | |
| for container in containers | |
| ) | |
| and cv2.countNonZero(cv2.bitwise_and(candidate.mask, protected_mask)) < 3 | |
| ) | |
| ] | |
| accepted = filter_arrow_context(accepted) | |
| accepted.sort( | |
| key=lambda item: (-item.shape.is_container, -int(item.mask.sum()), item.shape.id) | |
| ) | |
| for index, candidate in enumerate(accepted, 1): | |
| candidate.shape.id = f"shape-{index:04d}" | |
| candidate.shape.z_index = 100 if candidate.shape.is_container else 700 | |
| return [candidate.shape for candidate in accepted] | |
| def color_clusters(rgb: np.ndarray, protected_mask: np.ndarray, count: int = 10) -> list[np.ndarray]: | |
| lab = cv2.cvtColor(rgb, cv2.COLOR_RGB2LAB).astype(np.float32) | |
| pixels = lab.reshape(-1, 3) | |
| allowed = protected_mask.reshape(-1) == 0 | |
| indices = np.flatnonzero(allowed) | |
| if indices.size == 0: | |
| return [] | |
| stride = max(1, indices.size // 50_000) | |
| sample = pixels[indices[::stride]].copy() | |
| count = min(count, max(2, sample.shape[0] // 100)) | |
| cv2.setRNGSeed(1729) | |
| _, _, centers = cv2.kmeans( | |
| sample, | |
| count, | |
| None, | |
| (cv2.TERM_CRITERIA_EPS + cv2.TERM_CRITERIA_MAX_ITER, 30, 0.4), | |
| 3, | |
| cv2.KMEANS_PP_CENTERS, | |
| ) | |
| distances = np.linalg.norm(pixels[:, None, :] - centers[None, :, :], axis=2) | |
| labels = np.argmin(distances, axis=1).reshape(lab.shape[:2]) | |
| output: list[np.ndarray] = [] | |
| for index in range(count): | |
| mask = ((labels == index) & (protected_mask == 0)).astype(np.uint8) * 255 | |
| output.append(mask) | |
| return output | |
| def arrow_candidates( | |
| rgb: np.ndarray, | |
| protected_mask: np.ndarray, | |
| clusters: list[np.ndarray], | |
| ) -> list[Candidate]: | |
| """Find connected, concave silhouettes with a supported arrow tip and tail.""" | |
| height, width = rgb.shape[:2] | |
| image_area = width * height | |
| output: list[Candidate] = [] | |
| for cluster_index, cluster_mask in enumerate(clusters): | |
| cleaned = cv2.morphologyEx(cluster_mask, cv2.MORPH_OPEN, np.ones((2, 2), np.uint8)) | |
| cleaned = cv2.morphologyEx(cleaned, cv2.MORPH_CLOSE, np.ones((3, 3), np.uint8)) | |
| contours, _ = cv2.findContours(cleaned, cv2.RETR_EXTERNAL, cv2.CHAIN_APPROX_SIMPLE) | |
| for contour in contours: | |
| area = float(cv2.contourArea(contour)) | |
| x, y, w, h = cv2.boundingRect(contour) | |
| aspect = max(w / max(h, 1), h / max(w, 1)) | |
| if ( | |
| area < max(45, image_area * 0.000025) | |
| or w * h < 600 | |
| or max(w, h) < 24 | |
| # Compact arrowheads and short curved arrows can be almost square. | |
| or aspect < 1.03 | |
| ): | |
| continue | |
| perimeter = cv2.arcLength(contour, True) | |
| if perimeter <= 0: | |
| continue | |
| approx = cv2.approxPolyDP(contour, max(1.2, perimeter * 0.008), True) | |
| if not 5 <= len(approx) <= 48: | |
| continue | |
| hull = cv2.convexHull(contour) | |
| hull_area = float(cv2.contourArea(hull)) | |
| solidity = area / max(hull_area, 1) | |
| fill_ratio = area / max(w * h, 1) | |
| if ( | |
| solidity >= 0.93 | |
| or fill_ratio > 0.55 | |
| or (len(approx) > 14 and fill_ratio >= 0.28) | |
| or (area > image_area * 0.003 and aspect < 4) | |
| ): | |
| continue | |
| tip = supported_arrow_tip(approx) | |
| if tip is None: | |
| continue | |
| tip_index, tip_score = tip | |
| points = [Point(x=float(item[0][0]), y=float(item[0][1])) for item in approx] | |
| geometry = PathGeometry(points=points, closed=True) | |
| fitted_mask = np.zeros((height, width), dtype=np.uint8) | |
| cv2.drawContours(fitted_mask, [contour], -1, 255, -1) | |
| sample_mask = cv2.bitwise_and(cleaned, fitted_mask) | |
| sample_mask[protected_mask > 0] = 0 | |
| style = estimate_style(rgb, sample_mask, geometry) | |
| confidence = min(0.99, 0.72 + tip_score * 0.2 + (1 - solidity) * 0.1) | |
| shape = ShapeRegion( | |
| id=f"arrow-{cluster_index}-{tip_index}", | |
| confidence=confidence, | |
| geometry=geometry, | |
| source_contours=[contour_points(contour)], | |
| style=style, | |
| semantic_role="arrow", | |
| ) | |
| output.append(Candidate(shape=shape, mask=fitted_mask, score=confidence)) | |
| return output | |
| def arrow_color_masks(rgb: np.ndarray, protected_mask: np.ndarray) -> list[np.ndarray]: | |
| hsv = cv2.cvtColor(rgb, cv2.COLOR_RGB2HSV) | |
| masks: list[np.ndarray] = [] | |
| gray = ( | |
| (hsv[:, :, 1] < 55) | |
| & (hsv[:, :, 2] >= 45) | |
| & (hsv[:, :, 2] <= 220) | |
| & (protected_mask == 0) | |
| ) | |
| masks.append(gray.astype(np.uint8) * 255) | |
| # Overlapping bins keep anti-aliased and shaded arrow pixels connected while | |
| # still separating unrelated diagram colors. | |
| for start in range(0, 165, 15): | |
| colored = ( | |
| (hsv[:, :, 0] >= start) | |
| & (hsv[:, :, 0] < start + 30) | |
| & (hsv[:, :, 1] >= 55) | |
| & (hsv[:, :, 2] >= 35) | |
| & (hsv[:, :, 2] <= 235) | |
| & (protected_mask == 0) | |
| ) | |
| if np.count_nonzero(colored) >= 40: | |
| masks.append(colored.astype(np.uint8) * 255) | |
| red = ( | |
| ((hsv[:, :, 0] < 15) | (hsv[:, :, 0] >= 165)) | |
| & (hsv[:, :, 1] >= 55) | |
| & (hsv[:, :, 2] >= 35) | |
| & (hsv[:, :, 2] <= 235) | |
| & (protected_mask == 0) | |
| ) | |
| if np.count_nonzero(red) >= 40: | |
| masks.append(red.astype(np.uint8) * 255) | |
| return masks | |
| def filter_arrow_context(candidates: list[Candidate]) -> list[Candidate]: | |
| containers = [candidate for candidate in candidates if candidate.shape.is_container] | |
| arrow_groups: dict[int, list[Candidate]] = {index: [] for index in range(len(containers))} | |
| outside: set[int] = set() | |
| for candidate in candidates: | |
| if candidate.shape.semantic_role != "arrow": | |
| continue | |
| area = max(cv2.countNonZero(candidate.mask), 1) | |
| containing = [ | |
| index | |
| for index, container in enumerate(containers) | |
| if cv2.countNonZero(cv2.bitwise_and(candidate.mask, container.mask)) / area >= 0.9 | |
| ] | |
| if not containing: | |
| outside.add(id(candidate)) | |
| else: | |
| arrow_groups[containing[0]].append(candidate) | |
| accepted_inside: set[int] = set() | |
| for container_index, group in arrow_groups.items(): | |
| container = containers[container_index] | |
| eligible = [ | |
| candidate | |
| for candidate in group | |
| if color_distance(candidate.shape, container.shape) >= 25 | |
| ] | |
| peer_counts = { | |
| id(candidate): sum( | |
| color_distance(candidate.shape, other.shape) <= 75 | |
| and similar_arrow_scale(candidate.mask, other.mask) | |
| for other in eligible | |
| ) | |
| for candidate in eligible | |
| } | |
| stable = [candidate for candidate in eligible if peer_counts[id(candidate)] >= 4] | |
| for candidate in stable: | |
| stable_peers = sum( | |
| color_distance(candidate.shape, other.shape) <= 75 | |
| and similar_arrow_scale(candidate.mask, other.mask) | |
| for other in stable | |
| ) | |
| # A coherent repeated set is strong evidence inside illustration | |
| # panels; isolated arrow-like glyphs are usually raster detail. | |
| if stable_peers >= 4: | |
| accepted_inside.add(id(candidate)) | |
| return [ | |
| candidate | |
| for candidate in candidates | |
| if candidate.shape.semantic_role != "arrow" | |
| or id(candidate) in outside | |
| or id(candidate) in accepted_inside | |
| ] | |
| def color_distance(first: ShapeRegion, second: ShapeRegion) -> float: | |
| if not isinstance(first.style.fill, SolidPaint) or not isinstance(second.style.fill, SolidPaint): | |
| return float("inf") | |
| return float(np.linalg.norm(hex_rgb(first.style.fill.color) - hex_rgb(second.style.fill.color))) | |
| def similar_arrow_scale(first: np.ndarray, second: np.ndarray) -> bool: | |
| first_area = max(cv2.countNonZero(first), 1) | |
| second_area = max(cv2.countNonZero(second), 1) | |
| ratio = max(first_area, second_area) / min(first_area, second_area) | |
| return ratio <= 3.5 | |
| def hex_rgb(value: str) -> np.ndarray: | |
| value = value.lstrip("#") | |
| if len(value) == 3: | |
| value = "".join(character * 2 for character in value) | |
| return np.array([int(value[index:index + 2], 16) for index in (0, 2, 4)], dtype=np.float32) | |
| def supported_arrow_tip(approx: np.ndarray) -> tuple[int, float] | None: | |
| points = approx[:, 0, :].astype(np.float32) | |
| centroid = np.mean(points, axis=0) | |
| best: tuple[int, float] | None = None | |
| for index, tip in enumerate(points): | |
| previous = points[(index - 1) % len(points)] | |
| following = points[(index + 1) % len(points)] | |
| side_a = previous - tip | |
| side_b = following - tip | |
| norm_a, norm_b = np.linalg.norm(side_a), np.linalg.norm(side_b) | |
| if min(norm_a, norm_b) < 2: | |
| continue | |
| cosine = float(np.clip(np.dot(side_a, side_b) / (norm_a * norm_b), -1, 1)) | |
| angle = float(np.degrees(np.arccos(cosine))) | |
| if angle > 95: | |
| continue | |
| base_midpoint = (previous + following) / 2 | |
| direction = tip - base_midpoint | |
| direction_norm = np.linalg.norm(direction) | |
| if direction_norm < 2: | |
| continue | |
| direction /= direction_norm | |
| projections = (points - centroid) @ direction | |
| tip_projection = float((tip - centroid) @ direction) | |
| projection_span = float(projections.max() - projections.min()) | |
| if projection_span < 8 or tip_projection < float(projections.max()) - projection_span * 0.08: | |
| continue | |
| head_width = float(np.linalg.norm(previous - following)) | |
| if head_width < 4 or projection_span < head_width * 0.65: | |
| continue | |
| base_projection = float((base_midpoint - centroid) @ direction) | |
| tail_points = int(np.count_nonzero(projections < base_projection - max(2, projection_span * 0.08))) | |
| if tail_points < 2: | |
| continue | |
| symmetry = 1 - min(1.0, abs(norm_a - norm_b) / max(norm_a, norm_b)) | |
| sharpness = max(0.0, 1 - angle / 100) | |
| extremity = min(1.0, (tip_projection - float(np.median(projections))) / max(projection_span * 0.45, 1)) | |
| score = sharpness * 0.4 + symmetry * 0.25 + extremity * 0.35 | |
| if score >= 0.54 and (best is None or score > best[1]): | |
| best = (index, score) | |
| return best | |
| def smoothed_container_candidates( | |
| rgb: np.ndarray, | |
| protected_mask: np.ndarray, | |
| ) -> list[Candidate]: | |
| """Find large low-frequency panels after suppressing their foreground details.""" | |
| height, width = rgb.shape[:2] | |
| scale = min(1.0, 500 / max(width, height)) | |
| small_width = max(32, int(round(width * scale))) | |
| small_height = max(32, int(round(height * scale))) | |
| small = cv2.resize(rgb, (small_width, small_height), interpolation=cv2.INTER_AREA) | |
| kernel = min(21, min(small_width, small_height) // 2 * 2 - 1) | |
| kernel = max(5, kernel) | |
| smooth = cv2.medianBlur(small, kernel) | |
| lab = cv2.cvtColor(smooth, cv2.COLOR_RGB2LAB).astype(np.float32) | |
| pixels = lab.reshape(-1, 3) | |
| cv2.setRNGSeed(2718) | |
| _, labels, centers = cv2.kmeans( | |
| pixels, | |
| 12, | |
| None, | |
| (cv2.TERM_CRITERIA_EPS + cv2.TERM_CRITERIA_MAX_ITER, 30, 0.35), | |
| 4, | |
| cv2.KMEANS_PP_CENTERS, | |
| ) | |
| labels = labels.reshape(small_height, small_width) | |
| output: list[Candidate] = [] | |
| for index, center in enumerate(centers): | |
| mask = (labels == index).astype(np.uint8) * 255 | |
| mask = cv2.morphologyEx(mask, cv2.MORPH_CLOSE, np.ones((5, 5), np.uint8)) | |
| contours, _ = cv2.findContours(mask, cv2.RETR_EXTERNAL, cv2.CHAIN_APPROX_SIMPLE) | |
| for contour in contours: | |
| x, y, w, h = cv2.boundingRect(contour) | |
| area = cv2.contourArea(contour) | |
| if not ( | |
| width * 0.12 <= w / scale <= width * 0.28 | |
| and height * 0.45 <= h / scale <= height * 0.78 | |
| and area / max(w * h, 1) >= 0.55 | |
| ): | |
| continue | |
| box = Box( | |
| x=x / scale, | |
| y=y / scale, | |
| width=min(width - x / scale, w / scale), | |
| height=min(height - y / scale, h / scale), | |
| ) | |
| box = refine_container_box(rgb, box) | |
| geometry = RectGeometry( | |
| box=box, | |
| rx=min(box.width, box.height) * 0.035, | |
| ry=min(box.width, box.height) * 0.035, | |
| ) | |
| fitted = geometry_mask(geometry, width, height) | |
| full_lab = cv2.cvtColor(rgb, cv2.COLOR_RGB2LAB).astype(np.float32) | |
| close = np.linalg.norm(full_lab - center[None, None, :], axis=2) < 12 | |
| sample = ((fitted > 0) & close & (protected_mask == 0)).astype(np.uint8) * 255 | |
| style = estimate_style(rgb, sample, geometry) | |
| corners = [ | |
| Point(x=box.x, y=box.y), | |
| Point(x=box.x + box.width, y=box.y), | |
| Point(x=box.x + box.width, y=box.y + box.height), | |
| Point(x=box.x, y=box.y + box.height), | |
| ] | |
| confidence = min(0.98, 0.82 + area / max(w * h, 1) * 0.16) | |
| shape = ShapeRegion( | |
| id=f"smooth-{index}", | |
| confidence=confidence, | |
| geometry=geometry, | |
| source_contours=[corners], | |
| style=style, | |
| is_container=True, | |
| ) | |
| output.append(Candidate(shape=shape, mask=fitted, score=confidence)) | |
| return output | |
| def refine_container_box(rgb: np.ndarray, box: Box) -> Box: | |
| height, width = rgb.shape[:2] | |
| lab = cv2.cvtColor(rgb, cv2.COLOR_RGB2LAB).astype(np.float32) | |
| border = np.concatenate( | |
| [lab[:8].reshape(-1, 3), lab[-8:].reshape(-1, 3), lab[:, :8].reshape(-1, 3), lab[:, -8:].reshape(-1, 3)] | |
| ) | |
| delta = np.linalg.norm(lab - np.median(border, axis=0), axis=2) | |
| x0, y0 = max(0, int(box.x)), max(0, int(box.y)) | |
| x1 = min(width, int(np.ceil(box.x + box.width))) | |
| y1 = min(height, int(np.ceil(box.y + box.height))) | |
| row_signal = np.median(delta[y0:y1, x0:x1], axis=1) > 2.5 | |
| row_run = longest_run(row_signal) | |
| if row_run and row_run[1] - row_run[0] >= box.height * 0.5: | |
| y0, y1 = y0 + row_run[0], y0 + row_run[1] | |
| column_signal = np.median(delta[y0:y1, x0:x1], axis=0) > 2.5 | |
| column_run = longest_run(column_signal) | |
| if column_run and column_run[1] - column_run[0] >= box.width * 0.65: | |
| x0, x1 = x0 + column_run[0], x0 + column_run[1] | |
| return Box(x=x0, y=y0, width=max(1, x1 - x0), height=max(1, y1 - y0)) | |
| def longest_run(values: np.ndarray) -> tuple[int, int] | None: | |
| mask = (values.astype(np.uint8) * 255)[:, None] | |
| mask = cv2.morphologyEx(mask, cv2.MORPH_CLOSE, np.ones((7, 1), np.uint8))[:, 0] > 0 | |
| padded = np.pad(mask.astype(np.int8), (1, 1)) | |
| changes = np.diff(padded) | |
| starts = np.flatnonzero(changes == 1) | |
| ends = np.flatnonzero(changes == -1) | |
| if not len(starts): | |
| return None | |
| index = int(np.argmax(ends - starts)) | |
| return int(starts[index]), int(ends[index]) | |
| def fit_cluster_contour( | |
| rgb: np.ndarray, | |
| cluster_mask: np.ndarray, | |
| contour: np.ndarray, | |
| protected_mask: np.ndarray, | |
| width: int, | |
| height: int, | |
| cluster_index: int, | |
| ) -> Candidate | None: | |
| area = float(cv2.contourArea(contour)) | |
| image_area = width * height | |
| if area < max(45, image_area * 0.00003): | |
| return None | |
| perimeter = cv2.arcLength(contour, True) | |
| if perimeter <= 0: | |
| return None | |
| approx = cv2.approxPolyDP(contour, max(1.5, perimeter * 0.012), True) | |
| x, y, w, h = cv2.boundingRect(contour) | |
| if w < 5 or h < 5: | |
| return None | |
| contour_mask = np.zeros((height, width), dtype=np.uint8) | |
| cv2.drawContours(contour_mask, [contour], -1, 255, -1) | |
| circularity = 4 * np.pi * area / max(perimeter * perimeter, 1) | |
| geometry: RectGeometry | EllipseGeometry | PolygonGeometry | PathGeometry | |
| fitted_mask = np.zeros_like(contour_mask) | |
| rectangularity = area / max(w * h, 1) | |
| if len(approx) == 4 and rectangularity >= 0.58: | |
| rotated = cv2.minAreaRect(contour) | |
| (cx, cy), (rw, rh), angle = rotated | |
| if rw < rh: | |
| rw, rh = rh, rw | |
| angle += 90 | |
| angle = normalize_angle(angle) | |
| if cx - rw / 2 < 0 or cy - rh / 2 < 0: | |
| return None | |
| radius = min(rw, rh) * (0.045 if rectangularity < 0.97 else 0) | |
| geometry = RectGeometry( | |
| box=Box(x=cx - rw / 2, y=cy - rh / 2, width=rw, height=rh), | |
| rx=radius, | |
| ry=radius, | |
| rotation=angle, | |
| ) | |
| fitted_mask = geometry_mask(geometry, width, height) | |
| elif circularity >= 0.72 and len(contour) >= 5: | |
| (cx, cy), (ew, eh), angle = cv2.fitEllipse(contour) | |
| if ew < 5 or eh < 5: | |
| return None | |
| geometry = EllipseGeometry(cx=cx, cy=cy, rx=ew / 2, ry=eh / 2, rotation=angle) | |
| fitted_mask = geometry_mask(geometry, width, height) | |
| elif 3 <= len(approx) <= 10: | |
| points = [Point(x=float(item[0][0]), y=float(item[0][1])) for item in approx] | |
| geometry = PolygonGeometry(points=points) | |
| fitted_mask = geometry_mask(geometry, width, height) | |
| elif len(approx) <= 20 and (w / h >= 3.5 or h / w >= 3.5): | |
| points = [Point(x=float(item[0][0]), y=float(item[0][1])) for item in approx] | |
| geometry = PathGeometry(points=points, closed=True) | |
| fitted_mask = geometry_mask(geometry, width, height) | |
| else: | |
| return None | |
| large_container = area >= image_area * 0.012 and isinstance(geometry, RectGeometry) | |
| if not large_container: | |
| if isinstance(geometry, RectGeometry) and area < image_area * 0.00045: | |
| return None | |
| if isinstance(geometry, EllipseGeometry) and min(geometry.rx, geometry.ry) < 12: | |
| return None | |
| if isinstance(geometry, PolygonGeometry): | |
| aspect = max(w / h, h / w) | |
| if area < image_area * 0.00018 or not ( | |
| len(geometry.points) == 3 or aspect >= 1.8 or area >= image_area * 0.0015 | |
| ): | |
| return None | |
| fitted_area = cv2.countNonZero(fitted_mask) | |
| if fitted_area == 0: | |
| return None | |
| support = cv2.countNonZero(cv2.bitwise_and(cluster_mask, fitted_mask)) / fitted_area | |
| overlap = cv2.countNonZero(cv2.bitwise_and(contour_mask, fitted_mask)) | |
| union = cv2.countNonZero(cv2.bitwise_or(contour_mask, fitted_mask)) | |
| iou = overlap / max(union, 1) | |
| minimum_support = 0.38 if large_container else 0.64 | |
| if support < minimum_support or iou < (0.58 if large_container else 0.68): | |
| return None | |
| sample_mask = cv2.bitwise_and(cluster_mask, fitted_mask) | |
| sample_mask[protected_mask > 0] = 0 | |
| style = estimate_style(rgb, sample_mask, geometry) | |
| source = contour_points(contour) | |
| confidence = min(0.99, 0.45 * support + 0.45 * iou + 0.1) | |
| shape = ShapeRegion( | |
| id=f"cluster-{cluster_index}", | |
| confidence=confidence, | |
| geometry=geometry, | |
| source_contours=[source], | |
| style=style, | |
| is_container=large_container, | |
| ) | |
| return Candidate(shape=shape, mask=fitted_mask, score=confidence) | |
| def edge_rectangle_candidates(rgb: np.ndarray, protected_mask: np.ndarray) -> list[Candidate]: | |
| height, width = rgb.shape[:2] | |
| gray = cv2.cvtColor(rgb, cv2.COLOR_RGB2GRAY) | |
| edges = cv2.Canny(gray, 40, 120) | |
| edges[cv2.dilate(protected_mask, np.ones((3, 3), np.uint8)) > 0] = 0 | |
| edges = cv2.morphologyEx(edges, cv2.MORPH_CLOSE, np.ones((5, 5), np.uint8)) | |
| contours, _ = cv2.findContours(edges, cv2.RETR_EXTERNAL, cv2.CHAIN_APPROX_SIMPLE) | |
| output: list[Candidate] = [] | |
| image_area = width * height | |
| for contour in contours: | |
| area = float(cv2.contourArea(contour)) | |
| if area < image_area * 0.001: | |
| continue | |
| perimeter = cv2.arcLength(contour, True) | |
| approx = cv2.approxPolyDP(contour, max(2, perimeter * 0.02), True) | |
| x, y, w, h = cv2.boundingRect(contour) | |
| if not 4 <= len(approx) <= 8 or w < 20 or h < 20: | |
| continue | |
| rectangularity = area / max(w * h, 1) | |
| circularity = 4 * np.pi * area / max(perimeter * perimeter, 1) | |
| if 0.8 <= w / h <= 1.2 and circularity >= 0.68: | |
| continue | |
| if rectangularity < 0.72: | |
| continue | |
| geometry = RectGeometry( | |
| box=Box(x=x, y=y, width=w, height=h), | |
| rx=min(w, h) * (0.045 if rectangularity < 0.97 else 0), | |
| ry=min(w, h) * (0.045 if rectangularity < 0.97 else 0), | |
| ) | |
| mask = geometry_mask(geometry, width, height) | |
| interior = cv2.erode(mask, np.ones((5, 5), np.uint8)) | |
| interior[protected_mask > 0] = 0 | |
| style = estimate_style(rgb, interior, geometry) | |
| ring = cv2.subtract(mask, cv2.erode(mask, np.ones((5, 5), np.uint8))) | |
| ring_pixels = rgb[ring > 0] | |
| if ring_pixels.size: | |
| style.stroke = rgb_hex(np.median(ring_pixels, axis=0)) | |
| style.stroke_width = 2 | |
| confidence = min(0.94, 0.55 + rectangularity * 0.4) | |
| shape = ShapeRegion( | |
| id="edge-rect", | |
| confidence=confidence, | |
| geometry=geometry, | |
| source_contours=[contour_points(contour)], | |
| style=style, | |
| is_container=area >= image_area * 0.012, | |
| ) | |
| output.append(Candidate(shape=shape, mask=mask, score=confidence)) | |
| return output | |
| def estimate_style( | |
| rgb: np.ndarray, | |
| sample_mask: np.ndarray, | |
| geometry: RectGeometry | EllipseGeometry | PolygonGeometry | PathGeometry, | |
| ) -> ShapeStyle: | |
| ys, xs = np.nonzero(sample_mask) | |
| if xs.size == 0: | |
| return ShapeStyle(fill=SolidPaint(color="#000000")) | |
| values = rgb[ys, xs].astype(np.float32) | |
| median = np.median(values, axis=0) | |
| solid_error = float(np.mean((values - median) ** 2)) | |
| if xs.size >= 100 and solid_error > 5: | |
| x_span = max(float(xs.max() - xs.min()), 1) | |
| y_span = max(float(ys.max() - ys.min()), 1) | |
| tx = (xs - xs.min()) / x_span | |
| ty = (ys - ys.min()) / y_span | |
| fit_x, error_x = linear_color_fit(values, tx) | |
| fit_y, error_y = linear_color_fit(values, ty) | |
| fit, error, angle = (fit_x, error_x, 0.0) if error_x <= error_y else (fit_y, error_y, 90.0) | |
| if error < solid_error * 0.65 and np.linalg.norm(fit[1] - fit[0]) >= 5: | |
| return ShapeStyle( | |
| fill=LinearGradientPaint( | |
| angle=angle, | |
| stops=[ | |
| GradientStop(offset=0, color=rgb_hex(fit[0])), | |
| GradientStop(offset=1, color=rgb_hex(fit[1])), | |
| ], | |
| ) | |
| ) | |
| return ShapeStyle(fill=SolidPaint(color=rgb_hex(median))) | |
| def linear_color_fit(values: np.ndarray, t: np.ndarray) -> tuple[np.ndarray, float]: | |
| design = np.column_stack([np.ones_like(t), t]) | |
| coefficients, _, _, _ = np.linalg.lstsq(design, values, rcond=None) | |
| predicted = design @ coefficients | |
| endpoints = np.clip(np.vstack([coefficients[0], coefficients[0] + coefficients[1]]), 0, 255) | |
| return endpoints, float(np.mean((values - predicted) ** 2)) | |
| def deduplicate(candidates: list[Candidate]) -> list[Candidate]: | |
| accepted: list[Candidate] = [] | |
| # Preserve semantic interpretations when their silhouette overlaps a more | |
| # generic polygon/path fit. Containers still establish the layout first. | |
| for candidate in sorted( | |
| candidates, | |
| key=lambda item: ( | |
| item.shape.is_container, | |
| item.shape.semantic_role == "arrow", | |
| item.score, | |
| ), | |
| reverse=True, | |
| ): | |
| if ( | |
| isinstance(candidate.shape.geometry, RectGeometry) | |
| and not candidate.shape.is_container | |
| and any( | |
| existing.shape.is_container | |
| and cv2.countNonZero(cv2.bitwise_and(candidate.mask, existing.mask)) | |
| / max(cv2.countNonZero(candidate.mask), 1) | |
| >= 0.9 | |
| for existing in accepted | |
| ) | |
| ): | |
| continue | |
| duplicate = False | |
| for existing in accepted: | |
| intersection = cv2.countNonZero(cv2.bitwise_and(candidate.mask, existing.mask)) | |
| union = cv2.countNonZero(cv2.bitwise_or(candidate.mask, existing.mask)) | |
| smaller = min(cv2.countNonZero(candidate.mask), cv2.countNonZero(existing.mask)) | |
| container_containment = ( | |
| candidate.shape.is_container | |
| and existing.shape.is_container | |
| and smaller > 0 | |
| and intersection / smaller >= 0.78 | |
| ) | |
| arrow_containment = ( | |
| candidate.shape.semantic_role == "arrow" | |
| and existing.shape.semantic_role == "arrow" | |
| and smaller > 0 | |
| and intersection / smaller >= 0.7 | |
| ) | |
| if union and ( | |
| intersection / union >= 0.78 | |
| or container_containment | |
| or arrow_containment | |
| ): | |
| duplicate = True | |
| break | |
| if not duplicate: | |
| accepted.append(candidate) | |
| return accepted | |
| def geometry_mask( | |
| geometry: RectGeometry | EllipseGeometry | PolygonGeometry | PathGeometry, | |
| width: int, | |
| height: int, | |
| ) -> np.ndarray: | |
| mask = np.zeros((height, width), dtype=np.uint8) | |
| if isinstance(geometry, RectGeometry): | |
| box = geometry.box | |
| center = (box.x + box.width / 2, box.y + box.height / 2) | |
| if abs(geometry.rotation) > 0.01: | |
| points = cv2.boxPoints((center, (box.width, box.height), geometry.rotation)).astype(np.int32) | |
| cv2.fillPoly(mask, [points], 255) | |
| else: | |
| p1 = (int(round(box.x)), int(round(box.y))) | |
| p2 = (int(round(box.x + box.width)), int(round(box.y + box.height))) | |
| radius = int(round(max(geometry.rx, geometry.ry))) | |
| if radius > 0: | |
| cv2.rectangle(mask, (p1[0] + radius, p1[1]), (p2[0] - radius, p2[1]), 255, -1) | |
| cv2.rectangle(mask, (p1[0], p1[1] + radius), (p2[0], p2[1] - radius), 255, -1) | |
| for cx, cy in ( | |
| (p1[0] + radius, p1[1] + radius), | |
| (p2[0] - radius, p1[1] + radius), | |
| (p1[0] + radius, p2[1] - radius), | |
| (p2[0] - radius, p2[1] - radius), | |
| ): | |
| cv2.circle(mask, (cx, cy), radius, 255, -1) | |
| else: | |
| cv2.rectangle(mask, p1, p2, 255, -1) | |
| elif isinstance(geometry, EllipseGeometry): | |
| cv2.ellipse( | |
| mask, | |
| (int(round(geometry.cx)), int(round(geometry.cy))), | |
| (int(round(geometry.rx)), int(round(geometry.ry))), | |
| geometry.rotation, | |
| 0, | |
| 360, | |
| 255, | |
| -1, | |
| ) | |
| else: | |
| points = np.array([(point.x, point.y) for point in geometry.points], dtype=np.int32) | |
| if isinstance(geometry, PolygonGeometry) or geometry.closed: | |
| cv2.fillPoly(mask, [points], 255) | |
| else: | |
| cv2.polylines(mask, [points], False, 255, 2, cv2.LINE_AA) | |
| return mask | |
| def shape_source_mask(shape: ShapeRegion, width: int, height: int) -> np.ndarray: | |
| if not shape.source_contours: | |
| return geometry_mask(shape.geometry, width, height) | |
| mask = np.zeros((height, width), dtype=np.uint8) | |
| contours = [ | |
| np.array([(point.x, point.y) for point in contour], dtype=np.int32) | |
| for contour in shape.source_contours | |
| if len(contour) >= 3 | |
| ] | |
| if contours: | |
| cv2.fillPoly(mask, contours, 255) | |
| return mask | |
| def contour_points(contour: np.ndarray) -> list[Point]: | |
| simplified = cv2.approxPolyDP(contour, max(1, cv2.arcLength(contour, True) * 0.004), True) | |
| return [Point(x=float(item[0][0]), y=float(item[0][1])) for item in simplified] | |
| def rgb_hex(value: np.ndarray) -> str: | |
| red, green, blue = np.clip(np.rint(value), 0, 255).astype(np.uint8) | |
| return f"#{red:02x}{green:02x}{blue:02x}" | |
| def normalize_angle(value: float) -> float: | |
| while value > 90: | |
| value -= 180 | |
| while value < -90: | |
| value += 180 | |
| return float(value) | |
| def path_angle(points: list[Point]) -> float: | |
| start, end = points[0], points[-1] | |
| return degrees(atan2(end.y - start.y, end.x - start.x)) | |