make_editable_image / src /editable_image /shape_recognition.py
tonigi's picture
Add complete Gradio interface
3c2cd23 verified
Raw History Blame Contribute Delete
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,
)
@dataclass(slots=True)
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))