From e9b7d7481da9a25812032600bdfd9cbbd5ac8ed3 Mon Sep 17 00:00:00 2001 From: facok <128763816+facok@users.noreply.github.com> Date: Mon, 23 Mar 2026 16:07:06 +0800 Subject: [PATCH] Add missing core modules required by LCSColorAnchor core/bilateral.py and core/relationships.py were referenced by nodes/anchor.py but never committed, causing ModuleNotFoundError on fresh clones. --- core/bilateral.py | 77 ++++++++++++++++++++++++ core/relationships.py | 136 ++++++++++++++++++++++++++++++++++++++++++ 2 files changed, 213 insertions(+) create mode 100644 core/bilateral.py create mode 100644 core/relationships.py diff --git a/core/bilateral.py b/core/bilateral.py new file mode 100644 index 0000000..5c1941d --- /dev/null +++ b/core/bilateral.py @@ -0,0 +1,77 @@ +"""Bilateral filter in LCS space for smooth color anchoring.""" + +import torch +import torch.nn.functional as F + + +def estimate_bilateral_params(c, h_len, w_len): + """Estimate bilateral filter parameters from local color statistics. + + Computes per-channel spatial std of c across the grid, takes the median + to derive sigma_color. sigma_spatial is fixed at 1.5 (5x5 kernel is small). + + c: [B, L, 3] LCS coordinates + Returns: (sigma_spatial, sigma_color) floats + """ + B = c.shape[0] + grid = c.reshape(B, h_len, w_len, 3) # [B, H, W, 3] + # Per-channel std across spatial dims → [B, 3] + channel_std = grid.reshape(B, -1, 3).std(dim=1) # [B, 3] + # Median across batch and channels + median_std = float(channel_std.median()) + sigma_color = max(0.05, min(3.0, 0.75 * median_std)) + sigma_spatial = 1.5 + return sigma_spatial, sigma_color + + +def bilateral_filter_lcs(c, h_len, w_len, sigma_spatial, sigma_color, kernel_radius=2): + """Bilateral filter on [B, L, 3] LCS coordinates arranged on h_len x w_len grid. + + Uses spatial distance + LCS color distance as joint weights. + kernel_radius=2 -> 5x5 neighborhood (25 lookups per patch). + Returns [B, L, 3] filtered coordinates. + """ + B = c.shape[0] + # Reshape to spatial grid + grid = c.reshape(B, h_len, w_len, 3) # [B, H, W, 3] + + # Pad by kernel_radius (replicate) — pad last two spatial dims + # F.pad on [B, H, W, 3]: need to pad dims -3 and -2 (H and W) + # Permute to [B, 3, H, W] for F.pad, then back + grid_chw = grid.permute(0, 3, 1, 2) # [B, 3, H, W] + r = kernel_radius + padded = F.pad(grid_chw, (r, r, r, r), mode="replicate") # [B, 3, H+2r, W+2r] + + # Precompute spatial Gaussian weights for each offset in kernel + inv_2ss = -0.5 / (sigma_spatial * sigma_spatial) + inv_2sc = -0.5 / (sigma_color * sigma_color) + + # Accumulate weighted sum + weight_sum = torch.zeros(B, 1, h_len, w_len, device=c.device, dtype=c.dtype) + value_sum = torch.zeros(B, 3, h_len, w_len, device=c.device, dtype=c.dtype) + + for dy in range(-r, r + 1): + for dx in range(-r, r + 1): + # Spatial weight (constant per offset) + spatial_dist_sq = float(dy * dy + dx * dx) + w_spatial = torch.exp(torch.tensor(spatial_dist_sq * inv_2ss, device=c.device, dtype=c.dtype)) + + # Extract neighbor values from padded grid + y_start = r + dy + x_start = r + dx + neighbor = padded[:, :, y_start:y_start + h_len, x_start:x_start + w_len] # [B, 3, H, W] + + # Color distance weight (per-pixel) + diff = neighbor - grid_chw # [B, 3, H, W] + color_dist_sq = (diff * diff).sum(dim=1, keepdim=True) # [B, 1, H, W] + w_color = torch.exp(color_dist_sq * inv_2sc) # [B, 1, H, W] + + w = w_spatial * w_color + weight_sum = weight_sum + w + value_sum = value_sum + w * neighbor + + # Normalize + result = value_sum / weight_sum.clamp(min=1e-8) # [B, 3, H, W] + + # Back to [B, L, 3] + return result.permute(0, 2, 3, 1).reshape(B, -1, 3) diff --git a/core/relationships.py b/core/relationships.py new file mode 100644 index 0000000..71b46bf --- /dev/null +++ b/core/relationships.py @@ -0,0 +1,136 @@ +"""Local color relationship analysis for drift detection and correction.""" + +import torch +import torch.nn.functional as F + + +def compute_local_relationships(c, h_len, w_len, kernel_radius=2): + """Compute per-patch relationship vector from 5x5 neighborhood. + + For each patch, cosine similarity with each of up to 24 neighbors. + Returns [B, L, N_neighbors] relationship vectors where N_neighbors = (2*r+1)^2 - 1. + """ + B = c.shape[0] + r = kernel_radius + k_size = 2 * r + 1 + n_neighbors = k_size * k_size - 1 # 24 for r=2 + + # Reshape to spatial grid + grid = c.reshape(B, h_len, w_len, 3) # [B, H, W, 3] + + # Permute to [B, 3, H, W] for padding + grid_chw = grid.permute(0, 3, 1, 2) # [B, 3, H, W] + padded = F.pad(grid_chw, (r, r, r, r), mode="replicate") # [B, 3, H+2r, W+2r] + + # Center values — normalize for cosine similarity + center = grid_chw # [B, 3, H, W] + center_norm = center / center.norm(dim=1, keepdim=True).clamp(min=1e-8) + + # Collect cosine similarities with each neighbor + similarities = [] + for dy in range(-r, r + 1): + for dx in range(-r, r + 1): + if dy == 0 and dx == 0: + continue + y_start = r + dy + x_start = r + dx + neighbor = padded[:, :, y_start:y_start + h_len, x_start:x_start + w_len] + neighbor_norm = neighbor / neighbor.norm(dim=1, keepdim=True).clamp(min=1e-8) + # Cosine similarity per pixel + sim = (center_norm * neighbor_norm).sum(dim=1) # [B, H, W] + similarities.append(sim) + + # Stack to [B, H, W, N_neighbors] -> [B, L, N_neighbors] + rel = torch.stack(similarities, dim=-1) # [B, H, W, N_neighbors] + return rel.reshape(B, -1, n_neighbors) + + +def detect_anomalies(r_current, r_reference, threshold=0.3): + """Compare current vs reference relationships, return per-patch anomaly. + + Returns anomaly_magnitude [B, L, 1] -- 0.0 where relationships match, + >0 where disrupted, scaled by deviation magnitude. + """ + # Mean absolute difference across neighbor relationships + diff = (r_current - r_reference).abs().mean(dim=-1, keepdim=True) # [B, L, 1] + + # Soft threshold: below threshold -> 0, above -> linear ramp + anomaly = (diff - threshold).clamp(min=0.0) + + # Normalize so max anomaly ~ 1.0 (diff ranges from 0 to ~2 for cosine) + # Max possible diff for cosine sims is 2.0, minus threshold + max_range = 2.0 - threshold + anomaly = anomaly / max(max_range, 1e-8) + + return anomaly + + +def detect_anomalies_adaptive(r_current, r_reference): + """Compare current vs reference relationships with adaptive threshold. + + Uses per-batch robust outlier detection: threshold = median + 3.0 * 1.4826 * MAD. + Returns anomaly_magnitude [B, L, 1] in [0, 1]. + """ + # Mean absolute difference across neighbor relationships + diff = (r_current - r_reference).abs().mean(dim=-1) # [B, L] + + # Per-batch robust statistics + median = diff.median(dim=-1, keepdim=True).values # [B, 1] + mad = (diff - median).abs().median(dim=-1, keepdim=True).values # [B, 1] + threshold = median + 3.0 * 1.4826 * mad # [B, 1] + + # Soft ramp above threshold, normalized to [0, 1] + anomaly = (diff - threshold).clamp(min=0.0) # [B, L] + # Normalize per-batch: max anomaly → 1.0 + amax = anomaly.amax(dim=-1, keepdim=True).clamp(min=1e-8) # [B, 1] + anomaly = anomaly / amax + + return anomaly.unsqueeze(-1) # [B, L, 1] + + +def infer_color_from_neighbors(c, r_ref, anomaly_mag, h_len, w_len, kernel_radius=2): + """For anomalous patches, infer correct color from non-anomalous neighbors. + + Uses inverse-anomaly weighting: patches with low anomaly contribute more. + Returns [B, L, 3] corrected colors (blended: anomalous patches get + neighbor-inferred values, non-anomalous patches keep their original). + """ + B = c.shape[0] + r = kernel_radius + + # Reshape to spatial grid + grid = c.reshape(B, h_len, w_len, 3) + anom_grid = anomaly_mag.reshape(B, h_len, w_len, 1) + + # Pad both grid and anomaly + grid_chw = grid.permute(0, 3, 1, 2) # [B, 3, H, W] + anom_chw = anom_grid.permute(0, 3, 1, 2) # [B, 1, H, W] + padded_c = F.pad(grid_chw, (r, r, r, r), mode="replicate") + padded_a = F.pad(anom_chw, (r, r, r, r), mode="replicate") + + # Weight neighbors by how non-anomalous they are + weight_sum = torch.zeros(B, 1, h_len, w_len, device=c.device, dtype=c.dtype) + value_sum = torch.zeros(B, 3, h_len, w_len, device=c.device, dtype=c.dtype) + + for dy in range(-r, r + 1): + for dx in range(-r, r + 1): + if dy == 0 and dx == 0: + continue + y_start = r + dy + x_start = r + dx + neighbor_c = padded_c[:, :, y_start:y_start + h_len, x_start:x_start + w_len] + neighbor_a = padded_a[:, :, y_start:y_start + h_len, x_start:x_start + w_len] + + # Weight: 1 - anomaly (non-anomalous neighbors get high weight) + w = (1.0 - neighbor_a).clamp(min=0.01) # [B, 1, H, W] + weight_sum = weight_sum + w + value_sum = value_sum + w * neighbor_c + + # Inferred color from neighbors + inferred = value_sum / weight_sum.clamp(min=1e-8) # [B, 3, H, W] + inferred = inferred.permute(0, 2, 3, 1).reshape(B, -1, 3) # [B, L, 3] + + # Blend: anomalous patches use inferred, non-anomalous keep original + # anomaly_mag is [B, L, 1], range [0, ~1] + blend = anomaly_mag.clamp(0, 1) + return c * (1.0 - blend) + inferred * blend