s0sp
cpu intensive without tweak best result
a9e6328
Raw
History Blame Contribute Delete
7.68 kB
"""
Multi-model line art extraction with intelligent fusion.
"""
import cv2
import numpy as np
import onnxruntime as ort
from app.models.registry import ModelRegistry
from app.config import LineArtStyle
from app.processing.utils import (
multi_scale_close,
directional_close,
bridge_endpoints,
hysteresis_threshold,
declutter_lines,
)
class MultiLineArtExtractor:
"""Extracts line art using one or more AI models."""
def __init__(self, registry: ModelRegistry):
self.registry = registry
def _run_model(
self,
session: ort.InferenceSession,
image_bgr: np.ndarray,
target_size: int = 512,
) -> np.ndarray:
"""Run a single ONNX line art model."""
original_h, original_w = image_bgr.shape[:2]
rgb = cv2.cvtColor(image_bgr, cv2.COLOR_BGR2RGB)
resized = cv2.resize(rgb, (target_size, target_size))
blob = resized.astype(np.float32) / 255.0
blob = np.transpose(blob, (2, 0, 1))
blob = np.expand_dims(blob, axis=0)
input_name = session.get_inputs()[0].name
outputs = session.run(None, {input_name: blob})
raw = np.squeeze(outputs[0])
raw = np.clip(raw, 0.0, 1.0)
return cv2.resize(raw, (original_w, original_h))
def _preprocess(self, image: np.ndarray) -> np.ndarray:
"""Bilateral filter + CLAHE for cleaner model input."""
filtered = cv2.bilateralFilter(image, d=9, sigmaColor=75, sigmaSpace=75)
lab = cv2.cvtColor(filtered, cv2.COLOR_BGR2LAB)
clahe = cv2.createCLAHE(clipLimit=2.0, tileGridSize=(8, 8))
lab[:, :, 0] = clahe.apply(lab[:, :, 0])
return cv2.cvtColor(lab, cv2.COLOR_LAB2BGR)
def extract(
self,
image: np.ndarray,
style: LineArtStyle = LineArtStyle.FUSED,
density: str = "normal",
bridge_gaps: bool = True,
) -> np.ndarray:
"""
Extract line art using specified style.
"""
preprocessed = self._preprocess(image)
outputs: list[np.ndarray] = []
# ── Decide which models to run ──
if style == LineArtStyle.AUTO:
style = self._detect_best_style(image)
if style == LineArtStyle.FUSED:
models_to_run = ["informative_drawings", "anyline"]
elif style == LineArtStyle.MANGA:
models_to_run = ["manga_line"]
elif style == LineArtStyle.ANYLINE:
models_to_run = ["anyline"]
else:
models_to_run = ["informative_drawings"]
print(f" [LineArt] Styles to run: {models_to_run}")
# ── Run each available model ──
for model_name in models_to_run:
session = self.registry.get(model_name)
if session is None:
continue
sizes = self._get_inference_sizes(model_name, density)
scale_outputs = []
for size in sizes:
out = self._run_model(session, preprocessed, target_size=size)
scale_outputs.append(out)
merged = scale_outputs[0]
for o in scale_outputs[1:]:
merged = np.maximum(merged, o)
outputs.append(merged)
# ── Fallback: Canny if no models available ──
if not outputs:
gray = cv2.cvtColor(image, cv2.COLOR_BGR2GRAY)
edges = cv2.Canny(gray, 50, 150)
outputs.append(edges.astype(np.float32) / 255.0)
# ── Fuse multiple model outputs ──
if len(outputs) == 1:
fused = outputs[0]
else:
print(f" [LineArt] Fusing {len(outputs)} model outputs...")
fused = self._fuse_outputs(outputs)
# ── To uint8 ──
gray = (fused * 255).astype(np.uint8)
# ── Hysteresis threshold ──
thresholds = {
"sparse": (100, 200),
"normal": (90, 180),
"dense": (80, 160),
}
lo, hi = thresholds.get(density, (90, 180))
binary = hysteresis_threshold(gray, low=lo, high=hi)
binary = cv2.bitwise_not(binary) # lines = 0, bg = 255
# ── Aggressive closing pipeline ──
print(" [LineArt] Running closing pipeline...")
binary = self._close_pipeline(binary, density, bridge_gaps=bridge_gaps)
return binary
def _fuse_outputs(self, outputs: list[np.ndarray]) -> np.ndarray:
"""Intelligent fusion of multiple model outputs."""
if len(outputs) == 2:
a, b = outputs[0], outputs[1]
agreement = np.minimum(a, b)
union = np.maximum(a, b)
fused = 0.7 * agreement + 0.3 * union
both_strong = (a > 0.3) & (b > 0.3)
fused[both_strong] = np.maximum(fused[both_strong], 0.8)
return np.clip(fused, 0.0, 1.0)
else:
stacked = np.stack(outputs)
avg = np.mean(stacked, axis=0)
agree_count = np.sum(stacked > 0.3, axis=0).astype(np.float32)
boost = agree_count / len(outputs)
fused = avg * (0.5 + 0.5 * boost)
return np.clip(fused, 0.0, 1.0)
def _close_pipeline(self, binary: np.ndarray, density: str, bridge_gaps: bool = True) -> np.ndarray:
"""Full morphological closing pipeline."""
h, w = binary.shape
scale_factor = min(w, h) / 1024 # Normalize based on 1024px baseline
configs = {
"sparse": {"kernels": [3], "dir_len": 5, "iters": 1, "bridge": 5},
"normal": {"kernels": [3, 5], "dir_len": 7, "iters": 1, "bridge": 10},
"dense": {"kernels": [3, 5, 7], "dir_len": 9, "iters": 2, "bridge": 15},
}
cfg = configs.get(density, configs["normal"])
# Scale bridge radius
bridge_dist = max(1, int(cfg["bridge"] * scale_factor))
if scale_factor < 0.5: # Small image
bridge_dist = min(bridge_dist, 8)
print(" - Multi-scale close...")
result = multi_scale_close(binary, cfg["kernels"], cfg["iters"])
print(" - Directional close...")
result = directional_close(result, cfg["dir_len"], cfg["iters"])
smooth = cv2.getStructuringElement(cv2.MORPH_ELLIPSE, (3, 3))
result = cv2.morphologyEx(result, cv2.MORPH_CLOSE, smooth, iterations=1)
if bridge_gaps:
print(f" - Bridging endpoints (radius {bridge_dist})...")
result = bridge_endpoints(result, search_radius=bridge_dist)
# ── Declutter: Merge close lines and thin out ──
print(" - Decluttering lines (merge & thin)...")
result = declutter_lines(result, merge_radius=2, thin_out=True)
return result
def _detect_best_style(self, image: np.ndarray) -> LineArtStyle:
"""Auto-detect content type to pick the best model."""
hsv = cv2.cvtColor(image, cv2.COLOR_BGR2HSV)
saturation = hsv[:, :, 1].mean()
if saturation > 120:
if self.registry.available("manga_line"):
return LineArtStyle.MANGA
return LineArtStyle.FUSED
def _get_inference_sizes(self, model_name: str, density: str) -> list[int]:
"""Get multi-scale inference sizes per model and density."""
base = {
"informative_drawings": 512,
"anyline": 768,
"manga_line": 512,
}.get(model_name, 512)
if density == "sparse":
return [base]
elif density == "dense":
return [max(256, base - 128), base, base + 256]
else:
return [base, base + 256]