Spaces:
Build error
Build error
Download src/editable_image/bitmap_extract.py from tonigi/make_editable_image: direct link, hf CLI and curl.
- Browser
- Download file 7.92 kB
-
https://huggingface.co/spaces/tonigi/make_editable_image/resolve/main/src/editable_image/bitmap_extract.py
- Command line
-
hf download hf://spaces/tonigi/make_editable_image/src/editable_image/bitmap_extract.py
-
curl -L -o bitmap_extract.py https://huggingface.co/spaces/tonigi/make_editable_image/resolve/main/src/editable_image/bitmap_extract.py
7.92 kB
| 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) | |