Add recolor pass to clean residual contamination colour

Add bgfilter/recolor.py: after despill, propagate clean neighbour chroma into
pixels whose colour is still off (residual green, or the magenta a green unmix
overshoots into), keeping luminance partway so hair texture survives. It installs
a real observed colour rather than subtracting, so it cannot overshoot. Despill
already neutralises the matte edge (off ~ 0 there), so recolor only touches the
contaminated interior and the wisps the dropout-fill rescued.

Runs as a final colour pass on the despilled foreground; pymatting + despill keep
their stronger edge handling. On the samples the magenta fringe drops sharply
(TestImage vis-magenta 0.065->0.014, TestImage2 0.085->0.031) with edge green
unchanged (vis-green ~0.02). Debug output color_mask.png replaces the despill /
foreground-background / foreground-correction maps; smoke updated.

Co-Authored-By: Claude Opus 4.8 <noreply@anthropic.com>
This commit is contained in:
2026-06-30 22:19:15 +08:00
parent 42d65b0d25
commit b0a634dfd4
6 changed files with 119 additions and 9 deletions
+5 -1
View File
@@ -11,6 +11,7 @@ from .settings import (
ForegroundSettings,
ModelSettings,
PipelineSettings,
RecolorSettings,
SegmentationSettings,
TrimapSettings,
)
@@ -43,7 +44,7 @@ def _update_dataclass(instance: T, values: dict[str, Any] | None) -> T:
def settings_from_dict(data: dict[str, Any]) -> PipelineSettings:
allowed_sections = {"chroma", "trimap", "alpha_post", "foreground", "despill", "model", "segmentation"}
allowed_sections = {"chroma", "trimap", "alpha_post", "foreground", "despill", "recolor", "model", "segmentation"}
unknown_sections = sorted(set(data) - allowed_sections)
if unknown_sections:
raise ValueError(f"Unknown config section(s): {', '.join(unknown_sections)}")
@@ -54,6 +55,7 @@ def settings_from_dict(data: dict[str, Any]) -> PipelineSettings:
alpha_post=_update_dataclass(AlphaPostSettings(), data.get("alpha_post")),
foreground=_update_dataclass(ForegroundSettings(), data.get("foreground")),
despill=_update_dataclass(DespillSettings(), data.get("despill")),
recolor=_update_dataclass(RecolorSettings(), data.get("recolor")),
model=_update_dataclass(ModelSettings(), data.get("model")),
segmentation=_update_dataclass(SegmentationSettings(), data.get("segmentation")),
)
@@ -71,6 +73,7 @@ def override_settings(settings: PipelineSettings, **overrides: Any) -> PipelineS
alpha_post = settings.alpha_post
foreground = settings.foreground
despill = settings.despill
recolor = settings.recolor
model = settings.model
segmentation = settings.segmentation
@@ -102,6 +105,7 @@ def override_settings(settings: PipelineSettings, **overrides: Any) -> PipelineS
alpha_post=alpha_post,
foreground=foreground,
despill=_update_dataclass(despill, despill_updates),
recolor=recolor,
model=_update_dataclass(model, model_updates),
segmentation=_update_dataclass(segmentation, seg_updates),
)
+7 -5
View File
@@ -15,6 +15,7 @@ from .alpha_post import (
from .chroma import compute_bg_confidence
from .despill import despill_green
from .foreground import estimate_foreground_rgb
from .recolor import recolor_foreground
from .io import load_rgb, save_gray, save_rgb, save_rgba, write_text
from .qa import BACKGROUND_COLORS, composite, make_qa_grid, save_previews
from .settings import PipelineSettings
@@ -96,9 +97,12 @@ def _run_image(
foreground = estimate_foreground_rgb(
rgb, alpha, bg_confidence, model, settings.foreground
)
corrected_rgb, despill_mask = despill_green(
corrected_rgb, color_mask = despill_green(
foreground.rgb, alpha, model, settings.despill
)
if settings.recolor.enabled:
recolored = recolor_foreground(corrected_rgb, alpha, seg_mask, settings.recolor)
corrected_rgb, color_mask = recolored.rgb, recolored.mask
save_rgba(output_path, corrected_rgb, alpha)
result = {
@@ -118,10 +122,8 @@ def _run_image(
if seg_mask is not None:
save_gray(debug / "seg_mask.png", seg_mask)
save_gray(debug / "alpha.png", alpha)
save_rgb(debug / "foreground_rgb.png", foreground.rgb)
save_rgb(debug / "foreground_background.png", foreground.background)
save_gray(debug / "foreground_correction.png", foreground.correction)
save_gray(debug / "despill_mask.png", despill_mask)
save_rgb(debug / "foreground_rgb.png", corrected_rgb)
save_gray(debug / "color_mask.png", color_mask)
save_previews(debug, corrected_rgb, alpha)
tiles = {
"input": rgb,
+81
View File
@@ -0,0 +1,81 @@
from __future__ import annotations
from dataclasses import dataclass
import numpy as np
from .deps import require_cv2
from .settings import RecolorSettings
@dataclass(frozen=True)
class RecolorResult:
rgb: np.ndarray # recoloured foreground colour (uint8)
mask: np.ndarray # per-pixel recolour weight applied (float32, 0..1)
def _smoothstep(x: np.ndarray, edge0: float, edge1: float) -> np.ndarray:
t = np.clip((x - edge0) / max(edge1 - edge0, 1e-6), 0.0, 1.0)
return t * t * (3.0 - 2.0 * t)
def recolor_foreground(
rgb: np.ndarray,
alpha: np.ndarray,
seg_mask: np.ndarray | None,
settings: RecolorSettings,
) -> RecolorResult:
"""Repair AI green-contamination by propagating clean neighbour colour.
Green that the generator bled onto the subject is recoloured by pulling the
chroma from nearby *clean* (non-green) subject pixels, while keeping each
pixel's luminance so hair texture and shading survive. Because it installs a
real observed colour rather than subtracting green, it does not overshoot into
magenta the way a despill/unmix does. No assumption that the contaminated
region is hair -- the target is inferred locally per pixel.
"""
rgb_f = rgb.astype(np.float32) / 255.0
if not settings.enabled:
return RecolorResult(rgb=rgb.copy(), mask=np.zeros(alpha.shape, dtype=np.float32))
cv2 = require_cv2()
r, g, b = rgb_f[..., 0], rgb_f[..., 1], rgb_f[..., 2]
# Residual off-colour after despill: leftover green or the magenta that an
# unmix can overshoot into. Despill already neutralises the matte edge, so the
# clean edge has off ~ 0 and is left alone; this targets contaminated interior.
off = np.maximum(g - np.maximum(r, b), np.minimum(r, b) - g)
subject = np.clip(seg_mask if seg_mask is not None else alpha, 0.0, 1.0)
coverage = _smoothstep(np.clip(alpha, 0.0, 1.0), 0.0, settings.coverage_soft)
contam = _smoothstep(off, settings.contam_low, settings.contam_high) * coverage
clean = (subject >= settings.person_threshold).astype(np.float32) * (
off < settings.clean_dom_max
).astype(np.float32)
lab = cv2.cvtColor(rgb_f, cv2.COLOR_RGB2LAB)
L, A, B = lab[..., 0], lab[..., 1], lab[..., 2]
radius = max(3, int(round(max(alpha.shape) * settings.propagate_radius_ratio)))
k = (radius * 2 + 1, radius * 2 + 1)
den = cv2.GaussianBlur(clean, k, 0)
inv_den = 1.0 / np.maximum(den, 1e-6)
target_a = cv2.GaussianBlur(A * clean, k, 0) * inv_den
target_b = cv2.GaussianBlur(B * clean, k, 0) * inv_den
target_l = cv2.GaussianBlur(L * clean, k, 0) * inv_den
# Fade the recolour out where there is too little clean reference nearby
# (a fully-contaminated pocket has no trustworthy colour to borrow).
avail = _smoothstep(den, settings.min_reference_weight * 0.5, settings.min_reference_weight)
w = contam * avail
out_a = (1.0 - w) * A + w * target_a
out_b = (1.0 - w) * B + w * target_b
# Green is bright, so pull luminance partway to the clean neighbour to avoid a
# bright fringe, but only partway so hair texture is not flattened.
wl = w * settings.luminance_strength
out_l = (1.0 - wl) * L + wl * target_l
lab_out = np.stack([out_l, out_a, out_b], axis=-1)
out = cv2.cvtColor(lab_out, cv2.COLOR_LAB2RGB)
return RecolorResult(
rgb=np.clip(out * 255.0, 0, 255).astype(np.uint8),
mask=w.astype(np.float32),
)
+14
View File
@@ -79,6 +79,19 @@ class DespillSettings:
alpha_weight_floor: float = 0.35
@dataclass(frozen=True)
class RecolorSettings:
enabled: bool = True
person_threshold: float = 0.5 # clean reference = subject above this seg/alpha
coverage_soft: float = 0.08 # recolour ramps in over alpha 0..this
contam_low: float = 0.02 # residual off-colour where recolour starts
contam_high: float = 0.12 # residual off-colour for full recolour
clean_dom_max: float = 0.03 # reference = subject pixels below this off-colour
propagate_radius_ratio: float = 0.02
min_reference_weight: float = 0.05
luminance_strength: float = 0.5 # how far to pull luminance to the clean neighbour
@dataclass(frozen=True)
class ModelSettings:
model_name: str = "hustvl/vitmatte-small-composition-1k"
@@ -102,5 +115,6 @@ class PipelineSettings:
alpha_post: AlphaPostSettings = AlphaPostSettings()
foreground: ForegroundSettings = ForegroundSettings()
despill: DespillSettings = DespillSettings()
recolor: RecolorSettings = RecolorSettings()
model: ModelSettings = ModelSettings()
segmentation: SegmentationSettings = SegmentationSettings()
+11
View File
@@ -52,3 +52,14 @@ despill:
green_excess_margin: 0.015
edge_expand_radius: 2
alpha_weight_floor: 0.35
recolor:
enabled: true
person_threshold: 0.5
coverage_soft: 0.08
contam_low: 0.02
contam_high: 0.12
clean_dom_max: 0.03
propagate_radius_ratio: 0.02
min_reference_weight: 0.05
luminance_strength: 0.5
+1 -3
View File
@@ -18,10 +18,8 @@ REQUIRED_DEBUG_FILES = {
"bg_confidence.png",
"trimap.png",
"alpha.png",
"despill_mask.png",
"foreground_rgb.png",
"foreground_background.png",
"foreground_correction.png",
"color_mask.png",
"preview_black.png",
"preview_white.png",
"preview_gray.png",