From d6bb658244b58c92dfcc66e1f2b276e801023834 Mon Sep 17 00:00:00 2001 From: Codex Date: Tue, 30 Jun 2026 16:25:56 +0800 Subject: [PATCH] Add foreground color estimation pass --- README.md | 13 +++-- bgfilter/config.py | 6 +- bgfilter/despill.py | 43 +------------- bgfilter/foreground.py | 81 +++++++++++++++++++++++++++ bgfilter/pipeline.py | 12 +++- bgfilter/settings.py | 22 ++++++-- configs/default.yaml | 19 +++++-- docs/green_screen_matting_workflow.md | 26 ++++++--- scripts/smoke_samples.py | 6 +- 9 files changed, 158 insertions(+), 70 deletions(-) create mode 100644 bgfilter/foreground.py diff --git a/README.md b/README.md index 871a03a..fd2a20e 100644 --- a/README.md +++ b/README.md @@ -84,6 +84,10 @@ For each processed image, the CLI writes an RGBA PNG and optional debug files: bg_confidence.png trimap.png alpha.png +foreground_mask.png +foreground_unmix_mask.png +foreground_local_mask.png +foreground_rgb.png despill_mask.png preview_black.png preview_white.png @@ -102,7 +106,7 @@ edge pixels. ```powershell D:\MiniConda\envs\lightML\python.exe -m bgfilter.quality_cli ` Outputs\TestImage_rgba.png ` - --max-edge-green-excess-p95 0.50 + --max-edge-green-excess-p95 0.30 ``` Run the bundled sample smoke check: @@ -114,13 +118,12 @@ D:\MiniConda\envs\lightML\python.exe scripts\smoke_samples.py ` --config configs\default.yaml ` --device cpu ` --fallback-to-chroma-alpha ` - --max-edge-green-excess-p95 0.50 + --max-edge-green-excess-p95 0.30 ``` ## Current Notes - `Samples/` and `Outputs/` are ignored by Git. - `ViTMatte` model weights are loaded from Hugging Face on first use. -- Current despill is edge-local and includes a lightweight foreground unmixing pass. - Strong green spill around hair may still need more advanced foreground color - estimation. +- Foreground color estimation runs before final de-spill and writes debug masks + for unmixing and local foreground propagation. diff --git a/bgfilter/config.py b/bgfilter/config.py index 1b99194..377c7f3 100644 --- a/bgfilter/config.py +++ b/bgfilter/config.py @@ -8,6 +8,7 @@ from .settings import ( AlphaPostSettings, ChromaSettings, DespillSettings, + ForegroundSettings, ModelSettings, PipelineSettings, TrimapSettings, @@ -41,7 +42,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", "despill", "model"} + allowed_sections = {"chroma", "trimap", "alpha_post", "foreground", "despill", "model"} unknown_sections = sorted(set(data) - allowed_sections) if unknown_sections: raise ValueError(f"Unknown config section(s): {', '.join(unknown_sections)}") @@ -50,6 +51,7 @@ def settings_from_dict(data: dict[str, Any]) -> PipelineSettings: chroma=_update_dataclass(ChromaSettings(), data.get("chroma")), trimap=_update_dataclass(TrimapSettings(), data.get("trimap")), alpha_post=_update_dataclass(AlphaPostSettings(), data.get("alpha_post")), + foreground=_update_dataclass(ForegroundSettings(), data.get("foreground")), despill=_update_dataclass(DespillSettings(), data.get("despill")), model=_update_dataclass(ModelSettings(), data.get("model")), ) @@ -65,6 +67,7 @@ def override_settings(settings: PipelineSettings, **overrides: Any) -> PipelineS chroma = settings.chroma trimap = settings.trimap alpha_post = settings.alpha_post + foreground = settings.foreground despill = settings.despill model = settings.model @@ -91,6 +94,7 @@ def override_settings(settings: PipelineSettings, **overrides: Any) -> PipelineS chroma=chroma, trimap=_update_dataclass(trimap, trimap_updates), alpha_post=alpha_post, + foreground=foreground, despill=_update_dataclass(despill, despill_updates), model=_update_dataclass(model, model_updates), ) diff --git a/bgfilter/despill.py b/bgfilter/despill.py index 3a032f5..4796bbe 100644 --- a/bgfilter/despill.py +++ b/bgfilter/despill.py @@ -3,22 +3,9 @@ from __future__ import annotations import numpy as np from .chroma import BackgroundModel -from .deps import require_cv2 from .settings import DespillSettings -def _weighted_blur(values: np.ndarray, weights: np.ndarray, radius: int) -> np.ndarray: - if radius <= 0: - return values - cv2 = require_cv2() - kernel_size = radius * 2 + 1 - numerator = cv2.GaussianBlur( - values * weights[..., None], (kernel_size, kernel_size), 0 - ) - denominator = cv2.GaussianBlur(weights, (kernel_size, kernel_size), 0) - return numerator / np.maximum(denominator[..., None], 1e-6) - - def despill_green( rgb: np.ndarray, alpha: np.ndarray, @@ -30,9 +17,11 @@ def despill_green( if not settings.enabled: return rgb.copy(), np.zeros(alpha.shape, dtype=np.float32) - cv2 = require_cv2() edge = (alpha > settings.edge_low) & (alpha < settings.edge_high) if settings.edge_expand_radius > 0: + from .deps import require_cv2 + + cv2 = require_cv2() radius = settings.edge_expand_radius kernel = cv2.getStructuringElement( cv2.MORPH_ELLIPSE, (radius * 2 + 1, radius * 2 + 1) @@ -69,30 +58,4 @@ def despill_green( out[..., 0] = np.clip(out[..., 0] + compensation * 0.5, 0.0, 1.0) out[..., 2] = np.clip(out[..., 2] + compensation * 0.5, 0.0, 1.0) - # Approximate foreground recovery for semi-transparent contaminated pixels. - # C = alpha * F + (1 - alpha) * B, so F = (C - (1 - alpha) * B) / alpha. - # Blend it only where green excess is present; this is intentionally local - # because incorrect unmixing on opaque clothing would damage normal colors. - bg = np.asarray(model.rgb_center, dtype=np.float32) - a = np.clip(alpha[..., None], settings.unmix_alpha_min, 1.0) - recovered = np.clip((rgb_f - (1.0 - alpha[..., None]) * bg) / a, 0.0, 1.0) - unmix_mask = ( - mask - * settings.unmix_strength - * np.clip((1.0 - alpha) / max(1.0 - settings.unmix_alpha_min, 1e-6), 0.0, 1.0) - ) - out = out * (1.0 - unmix_mask[..., None]) + recovered * unmix_mask[..., None] - - local_fg_weights = ( - (alpha >= settings.local_fg_alpha_threshold) - & (bg_confidence <= settings.local_fg_bg_confidence_max) - ).astype(np.float32) - local_fg = _weighted_blur(out, local_fg_weights, settings.local_fg_blur_radius) - local_mask = ( - mask - * settings.local_fg_strength - * np.clip((1.0 - alpha) / max(1.0 - settings.unmix_alpha_min, 1e-6), 0.0, 1.0) - ) - out = out * (1.0 - local_mask[..., None]) + local_fg * local_mask[..., None] - return np.clip(out * 255.0, 0, 255).astype(np.uint8), mask.astype(np.float32) diff --git a/bgfilter/foreground.py b/bgfilter/foreground.py new file mode 100644 index 0000000..bca3f30 --- /dev/null +++ b/bgfilter/foreground.py @@ -0,0 +1,81 @@ +from __future__ import annotations + +from dataclasses import dataclass + +import numpy as np + +from .chroma import BackgroundModel +from .deps import require_cv2 +from .settings import ForegroundSettings + + +@dataclass(frozen=True) +class ForegroundEstimate: + rgb: np.ndarray + mask: np.ndarray + unmix_mask: np.ndarray + local_mask: np.ndarray + + +def _weighted_blur(values: np.ndarray, weights: np.ndarray, radius: int) -> np.ndarray: + if radius <= 0: + return values + cv2 = require_cv2() + kernel_size = radius * 2 + 1 + numerator = cv2.GaussianBlur( + values * weights[..., None], (kernel_size, kernel_size), 0 + ) + denominator = cv2.GaussianBlur(weights, (kernel_size, kernel_size), 0) + return numerator / np.maximum(denominator[..., None], 1e-6) + + +def estimate_foreground_rgb( + rgb: np.ndarray, + alpha: np.ndarray, + bg_confidence: np.ndarray, + model: BackgroundModel, + settings: ForegroundSettings, +) -> ForegroundEstimate: + rgb_f = rgb.astype(np.float32) / 255.0 + zero = np.zeros(alpha.shape, dtype=np.float32) + if not settings.enabled: + return ForegroundEstimate(rgb=rgb.copy(), mask=zero, unmix_mask=zero, local_mask=zero) + + r = rgb_f[..., 0] + g = rgb_f[..., 1] + b = rgb_f[..., 2] + green_excess = np.maximum(g - (np.maximum(r, b) + settings.green_excess_margin), 0.0) + edge = (alpha > settings.edge_low) & (alpha < settings.edge_high) + green_weight = np.clip( + green_excess / max(settings.green_excess_margin * 8.0, 1e-6), 0.0, 1.0 + ) + bg_weight = ( + settings.bg_confidence_weight * np.clip(bg_confidence, 0.0, 1.0) + + (1.0 - settings.bg_confidence_weight) + ) + base_mask = edge.astype(np.float32) * green_weight * bg_weight + + bg = np.asarray(model.rgb_center, dtype=np.float32) + a = np.clip(alpha[..., None], settings.min_unmix_alpha, 1.0) + unmixed = np.clip((rgb_f - (1.0 - alpha[..., None]) * bg) / a, 0.0, 1.0) + low_alpha_weight = np.clip( + (1.0 - alpha) / max(1.0 - settings.min_unmix_alpha, 1e-6), 0.0, 1.0 + ) + unmix_mask = base_mask * settings.unmix_strength * low_alpha_weight + out = rgb_f * (1.0 - unmix_mask[..., None]) + unmixed * unmix_mask[..., None] + + local_fg_weights = ( + (alpha >= settings.local_alpha_threshold) + & (bg_confidence <= settings.local_bg_confidence_max) + ).astype(np.float32) + local_fg = _weighted_blur(out, local_fg_weights, settings.local_blur_radius) + local_mask = base_mask * settings.local_strength * low_alpha_weight + out = out * (1.0 - local_mask[..., None]) + local_fg * local_mask[..., None] + + combined_mask = np.clip(unmix_mask + local_mask, 0.0, 1.0) + return ForegroundEstimate( + rgb=np.clip(out * 255.0, 0, 255).astype(np.uint8), + mask=combined_mask.astype(np.float32), + unmix_mask=unmix_mask.astype(np.float32), + local_mask=local_mask.astype(np.float32), + ) diff --git a/bgfilter/pipeline.py b/bgfilter/pipeline.py index 89ea2c7..b6d6529 100644 --- a/bgfilter/pipeline.py +++ b/bgfilter/pipeline.py @@ -9,7 +9,8 @@ import numpy as np from .alpha_post import clean_alpha, enforce_trimap from .chroma import compute_bg_confidence from .despill import despill_green -from .io import load_rgb, save_gray, save_rgba, write_text +from .foreground import estimate_foreground_rgb +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 from .trimap import generate_trimap, trimap_to_alpha_seed @@ -71,8 +72,11 @@ def _run_image( alpha = enforce_trimap(alpha, trimap) alpha = clean_alpha(alpha, trimap, settings.alpha_post) + foreground = estimate_foreground_rgb( + rgb, alpha, bg_confidence, model, settings.foreground + ) corrected_rgb, despill_mask = despill_green( - rgb, alpha, bg_confidence, model, settings.despill + foreground.rgb, alpha, bg_confidence, model, settings.despill ) save_rgba(output_path, corrected_rgb, alpha) @@ -91,6 +95,10 @@ def _run_image( save_gray(debug / "bg_confidence.png", bg_confidence) save_gray(debug / "trimap.png", trimap) save_gray(debug / "alpha.png", alpha) + save_gray(debug / "foreground_mask.png", foreground.mask) + save_gray(debug / "foreground_unmix_mask.png", foreground.unmix_mask) + save_gray(debug / "foreground_local_mask.png", foreground.local_mask) + save_rgb(debug / "foreground_rgb.png", foreground.rgb) save_gray(debug / "despill_mask.png", despill_mask) save_previews(debug, corrected_rgb, alpha) tiles = { diff --git a/bgfilter/settings.py b/bgfilter/settings.py index 3cd3d86..770d39f 100644 --- a/bgfilter/settings.py +++ b/bgfilter/settings.py @@ -35,6 +35,21 @@ class AlphaPostSettings: alpha_ceil: float = 0.998 +@dataclass(frozen=True) +class ForegroundSettings: + enabled: bool = True + edge_low: float = 0.005 + edge_high: float = 0.995 + min_unmix_alpha: float = 0.08 + unmix_strength: float = 0.70 + local_strength: float = 0.45 + local_blur_radius: int = 11 + local_alpha_threshold: float = 0.92 + local_bg_confidence_max: float = 0.20 + green_excess_margin: float = 0.015 + bg_confidence_weight: float = 0.70 + + @dataclass(frozen=True) class DespillSettings: enabled: bool = True @@ -45,12 +60,6 @@ class DespillSettings: edge_expand_radius: int = 2 bg_confidence_weight: float = 0.65 alpha_weight_floor: float = 0.35 - unmix_strength: float = 0.55 - unmix_alpha_min: float = 0.08 - local_fg_strength: float = 0.35 - local_fg_blur_radius: int = 9 - local_fg_alpha_threshold: float = 0.92 - local_fg_bg_confidence_max: float = 0.20 @dataclass(frozen=True) @@ -66,5 +75,6 @@ class PipelineSettings: chroma: ChromaSettings = ChromaSettings() trimap: TrimapSettings = TrimapSettings() alpha_post: AlphaPostSettings = AlphaPostSettings() + foreground: ForegroundSettings = ForegroundSettings() despill: DespillSettings = DespillSettings() model: ModelSettings = ModelSettings() diff --git a/configs/default.yaml b/configs/default.yaml index 6e9b125..90cf10f 100644 --- a/configs/default.yaml +++ b/configs/default.yaml @@ -10,6 +10,19 @@ trimap: unknown_radius_ratio: 0.012 fg_safe_radius_ratio: 0.006 +foreground: + enabled: true + edge_low: 0.005 + edge_high: 0.995 + min_unmix_alpha: 0.08 + unmix_strength: 0.70 + local_strength: 0.45 + local_blur_radius: 11 + local_alpha_threshold: 0.92 + local_bg_confidence_max: 0.20 + green_excess_margin: 0.015 + bg_confidence_weight: 0.70 + despill: enabled: true edge_low: 0.005 @@ -19,9 +32,3 @@ despill: edge_expand_radius: 2 bg_confidence_weight: 0.65 alpha_weight_floor: 0.35 - unmix_strength: 0.55 - unmix_alpha_min: 0.08 - local_fg_strength: 0.35 - local_fg_blur_radius: 9 - local_fg_alpha_threshold: 0.92 - local_fg_bg_confidence_max: 0.20 diff --git a/docs/green_screen_matting_workflow.md b/docs/green_screen_matting_workflow.md index 4c10610..6b4958f 100644 --- a/docs/green_screen_matting_workflow.md +++ b/docs/green_screen_matting_workflow.md @@ -71,6 +71,8 @@ trimap 约束 alpha ↓ alpha 轻量清理 ↓ +foreground color estimation + ↓ de-spill 去绿边 ↓ RGBA PNG 导出 @@ -370,6 +372,10 @@ result_rgba.png debug/bg_confidence.png debug/trimap.png debug/alpha.png +debug/foreground_mask.png +debug/foreground_unmix_mask.png +debug/foreground_local_mask.png +debug/foreground_rgb.png debug/despill_mask.png debug/preview_black.png debug/preview_white.png @@ -419,12 +425,13 @@ debug/qa_grid.png ```text 1. 能生成 RGBA PNG。 2. 能生成 bg_confidence / trimap / alpha debug 图。 -3. 能生成黑白灰红蓝多背景 QA 图。 -4. 外部绿幕基本无残留。 -5. 发丝区域不是硬切边。 -6. 镂空洞露出绿幕时可透明。 -7. 边缘没有明显绿色污染。 -8. `scripts/smoke_samples.py` 能处理样例图并验证 RGBA、debug 产物和边缘绿色污染指标。 +3. 能生成 foreground color estimation debug 图。 +4. 能生成黑白灰红蓝多背景 QA 图。 +5. 外部绿幕基本无残留。 +6. 发丝区域不是硬切边。 +7. 镂空洞露出绿幕时可透明。 +8. 边缘没有明显绿色污染。 +9. `scripts/smoke_samples.py` 能处理样例图并验证 RGBA、debug 产物和边缘绿色污染指标。 ``` ## 19. 推荐落地顺序 @@ -438,7 +445,8 @@ debug/qa_grid.png 6. 接入 ViTMatte GPU 推理。 7. 加 trimap alpha 强约束。 8. 实现轻量 alpha 清理。 -9. 实现 de-spill。 -10. 实现 QA grid。 -11. 用样例图调默认参数。 +9. 实现 foreground color estimation。 +10. 实现 de-spill。 +11. 实现 QA grid。 +12. 用样例图调默认参数。 ``` diff --git a/scripts/smoke_samples.py b/scripts/smoke_samples.py index 3d77342..96ca103 100644 --- a/scripts/smoke_samples.py +++ b/scripts/smoke_samples.py @@ -19,6 +19,10 @@ REQUIRED_DEBUG_FILES = { "trimap.png", "alpha.png", "despill_mask.png", + "foreground_mask.png", + "foreground_unmix_mask.png", + "foreground_local_mask.png", + "foreground_rgb.png", "preview_black.png", "preview_white.png", "preview_gray.png", @@ -62,7 +66,7 @@ def main() -> int: default=None, help="Override config and allow chroma alpha when ViTMatte fails.", ) - parser.add_argument("--max-edge-green-excess-p95", type=float, default=0.50) + parser.add_argument("--max-edge-green-excess-p95", type=float, default=0.30) args = parser.parse_args() images = sorted(args.samples_dir.glob("*.png"))