from __future__ import annotations import cv2 import numpy as np from .models import ( BitmapExtractor, BitmapRegion, Box, DocumentSpec, LinearGradientPaint, ShapeRegion, SolidPaint, ) from .shape_recognition import geometry_mask, shape_source_mask def extract_bitmap_layers( rgba: np.ndarray, document: DocumentSpec, text_mask: np.ndarray, extractor: BitmapExtractor, device: str, ) -> tuple[list[BitmapRegion], dict[str, np.ndarray]]: height, width = rgba.shape[:2] rgb = rgba[:, :, :3] text_exclusion = cv2.dilate(text_mask, np.ones((5, 5), np.uint8)) for region in document.regions: polygon = np.array( [(point.x, point.y) for point in region.source_quad], dtype=np.int32, ) cv2.fillPoly(text_exclusion, [polygon], 255) vector_foreground = np.zeros((height, width), dtype=np.uint8) for shape in document.shapes: if not shape.is_container and shape.include_overlay: vector_foreground = cv2.bitwise_or( vector_foreground, geometry_mask(shape.geometry, width, height), ) regions: list[BitmapRegion] = [] assets: dict[str, np.ndarray] = {} for shape in document.shapes: if not shape.is_container or not shape.include_overlay: continue container_mask = geometry_mask(shape.geometry, width, height) predicted = render_fill(shape, width, height) source_lab = cv2.cvtColor(rgb, cv2.COLOR_RGB2LAB).astype(np.float32) predicted_lab = cv2.cvtColor(predicted, cv2.COLOR_RGB2LAB).astype(np.float32) delta = np.linalg.norm(source_lab - predicted_lab, axis=2) border_exclusion = cv2.subtract( container_mask, cv2.erode(container_mask, np.ones((7, 7), np.uint8)), ) eligible = ( (container_mask > 0) & (text_exclusion == 0) & (vector_foreground == 0) & (border_exclusion == 0) ) values = delta[eligible] if values.size < 20: continue baseline = float(np.percentile(values, 35)) threshold = max(9.0, baseline + 6.0) initial = np.zeros((height, width), dtype=np.uint8) initial[eligible & (delta >= threshold * 0.65)] = cv2.GC_PR_FGD initial[eligible & (delta >= threshold * 1.35)] = cv2.GC_FGD initial[(container_mask > 0) & (initial == 0)] = cv2.GC_BGD if np.count_nonzero(initial == cv2.GC_FGD) < 8: continue mask = refine_opencv(rgb, initial, container_mask) if extractor == BitmapExtractor.SAM2: from .sam2 import refine_with_sam2 mask = refine_with_sam2(rgb, mask, container_mask, device) mask[text_exclusion > 0] = 0 mask[vector_foreground > 0] = 0 mask = filter_components(mask, max(8, int(width * height * 0.00001))) if cv2.countNonZero(mask) < 12: continue alpha = soft_alpha(delta, mask, threshold) x, y, w, h = cv2.boundingRect((alpha > 0).astype(np.uint8)) if w <= 0 or h <= 0: continue padding = 2 x0, y0 = max(0, x - padding), max(0, y - padding) x1, y1 = min(width, x + w + padding), min(height, y + h + padding) crop = np.dstack([rgb[y0:y1, x0:x1], alpha[y0:y1, x0:x1]]) asset_name = f"bitmap-{len(regions) + 1:04d}.png" region = BitmapRegion( id=f"bitmap-{len(regions) + 1:04d}", asset_name=asset_name, box=Box(x=x0, y=y0, width=x1 - x0, height=y1 - y0), z_index=500, ) regions.append(region) assets[region.id] = crop return regions, assets def refine_opencv(rgb: np.ndarray, initial: np.ndarray, container_mask: np.ndarray) -> np.ndarray: mask = initial.copy() try: cv2.grabCut( cv2.cvtColor(rgb, cv2.COLOR_RGB2BGR), mask, None, np.zeros((1, 65), np.float64), np.zeros((1, 65), np.float64), 2, cv2.GC_INIT_WITH_MASK, ) output = np.isin(mask, [cv2.GC_FGD, cv2.GC_PR_FGD]).astype(np.uint8) * 255 except cv2.error: output = (initial >= cv2.GC_PR_FGD).astype(np.uint8) * 255 output[container_mask == 0] = 0 return output def filter_components(mask: np.ndarray, minimum_area: int) -> np.ndarray: count, labels, stats, _ = cv2.connectedComponentsWithStats((mask > 0).astype(np.uint8), 8) output = np.zeros_like(mask, dtype=np.uint8) for index in range(1, count): if stats[index, cv2.CC_STAT_AREA] >= minimum_area: output[labels == index] = 255 return output def soft_alpha(delta: np.ndarray, mask: np.ndarray, threshold: float) -> np.ndarray: strength = np.clip((delta - threshold * 0.45) / max(threshold, 1), 0, 1) alpha = np.rint(strength * 255).astype(np.uint8) alpha[mask == 0] = 0 alpha = cv2.GaussianBlur(alpha, (3, 3), 0.65) alpha[cv2.dilate(mask, np.ones((3, 3), np.uint8)) == 0] = 0 return alpha def render_fill(shape: ShapeRegion, width: int, height: int) -> np.ndarray: paint = shape.style.fill if isinstance(paint, SolidPaint): color = parse_hex(paint.color) return np.broadcast_to(color, (height, width, 3)).copy() assert isinstance(paint, LinearGradientPaint) first, last = paint.stops[0], paint.stops[-1] color0 = parse_hex(first.color).astype(np.float32) color1 = parse_hex(last.color).astype(np.float32) radians = np.deg2rad(paint.angle) direction = np.array([np.cos(radians), np.sin(radians)], dtype=np.float32) yy, xx = np.mgrid[0:height, 0:width] projection = xx * direction[0] + yy * direction[1] minimum, maximum = float(projection.min()), float(projection.max()) t = (projection - minimum) / max(maximum - minimum, 1) result = color0[None, None, :] * (1 - t[:, :, None]) + color1[None, None, :] * t[:, :, None] return np.clip(np.rint(result), 0, 255).astype(np.uint8) def reconstruct_background( rgba: np.ndarray, document: DocumentSpec, combined_mask: np.ndarray, inpaint_callback, ) -> np.ndarray: height, width = rgba.shape[:2] working = rgba.copy() container_union = np.zeros((height, width), dtype=np.uint8) for shape in sorted(document.shapes, key=lambda item: geometry_area(item), reverse=True): if not shape.is_container or not shape.remove_source: continue mask = shape_source_mask(shape, width, height) if not cv2.countNonZero(mask): continue ring = cv2.subtract( cv2.dilate(mask, np.ones((15, 15), np.uint8)), cv2.dilate(mask, np.ones((3, 3), np.uint8)), ) samples = working[:, :, :3][ring > 0] fill = np.median(samples, axis=0) if samples.size else np.array([255, 255, 255]) working[:, :, :3][mask > 0] = np.rint(fill).astype(np.uint8) container_union = cv2.bitwise_or(container_union, mask) remaining = combined_mask.copy() remaining[container_union > 0] = 0 if cv2.countNonZero(remaining): working = inpaint_callback(working, remaining) return working def geometry_area(shape: ShapeRegion) -> float: geometry = shape.geometry if hasattr(geometry, "box"): return geometry.box.width * geometry.box.height if hasattr(geometry, "rx"): return float(np.pi * geometry.rx * geometry.ry) points = np.array([(point.x, point.y) for point in geometry.points], dtype=np.float32) return abs(float(cv2.contourArea(points))) if len(points) >= 3 else 0 def parse_hex(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.uint8)