make_editable_image / src /editable_image /bitmap_extract.py
tonigi's picture
Add complete Gradio interface
3c2cd23 verified
Raw History Blame Contribute Delete
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)