From 44812b79a632a995be5915c6f0c8f4fcb2b81e1b Mon Sep 17 00:00:00 2001 From: facok <128763816+facok@users.noreply.github.com> Date: Mon, 23 Mar 2026 16:33:06 +0800 Subject: [PATCH] Clean up anchor and calibration code after review - Remove dead detect_anomalies() (superseded by adaptive variant) - Remove unused r_ref param from infer_color_from_neighbors() - Remove write-only c_ema state from anchor hook - Use math.exp() instead of torch.tensor()+torch.exp() in bilateral loop - Use in-place .add_() for accumulation in bilateral filter - Pre-normalize padded tensor once in compute_local_relationships() - Fix video VAE fallback producing duplicate vectors in calibration --- core/bilateral.py | 8 +++++--- core/calibration.py | 7 ++++--- core/relationships.py | 31 ++++++------------------------- nodes/anchor.py | 7 +------ 4 files changed, 16 insertions(+), 37 deletions(-) diff --git a/core/bilateral.py b/core/bilateral.py index 5c1941d..27c6815 100644 --- a/core/bilateral.py +++ b/core/bilateral.py @@ -1,5 +1,7 @@ """Bilateral filter in LCS space for smooth color anchoring.""" +import math + import torch import torch.nn.functional as F @@ -54,7 +56,7 @@ def bilateral_filter_lcs(c, h_len, w_len, sigma_spatial, sigma_color, kernel_rad 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)) + w_spatial = math.exp(spatial_dist_sq * inv_2ss) # Extract neighbor values from padded grid y_start = r + dy @@ -67,8 +69,8 @@ def bilateral_filter_lcs(c, h_len, w_len, sigma_spatial, sigma_color, kernel_rad 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 + weight_sum.add_(w) + value_sum.add_(w * neighbor) # Normalize result = value_sum / weight_sum.clamp(min=1e-8) # [B, 3, H, W] diff --git a/core/calibration.py b/core/calibration.py index b22f853..a9fbc80 100644 --- a/core/calibration.py +++ b/core/calibration.py @@ -103,11 +103,12 @@ def calibrate(vae, num_colors=512, image_size=512, batch_size=8): # Normal VAE: batch encode worked vectors.extend(avg.unbind(0)) else: - # Video VAE: batch not supported, encode one by one - vectors.extend(avg.unbind(0)) - for k in range(1, actual_batch): + # Video VAE or unexpected batch collapse — encode one by one + for k in range(actual_batch): single = imgs[k:k+1, :, :, :3] lat = vae.encode(single) + if lat.ndim == 5: + lat = lat[:, :, 0, :, :] p, _, _, _ = patchify(lat) vectors.append(p.mean(dim=1).cpu().squeeze(0)) diff --git a/core/relationships.py b/core/relationships.py index 71b46bf..4706595 100644 --- a/core/relationships.py +++ b/core/relationships.py @@ -23,8 +23,10 @@ def compute_local_relationships(c, h_len, w_len, kernel_radius=2): 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) + center_norm = grid_chw / grid_chw.norm(dim=1, keepdim=True).clamp(min=1e-8) + + # Pre-normalize padded tensor once (avoids per-neighbor normalization in loop) + padded_norm = padded / padded.norm(dim=1, keepdim=True).clamp(min=1e-8) # Collect cosine similarities with each neighbor similarities = [] @@ -34,8 +36,7 @@ def compute_local_relationships(c, h_len, w_len, kernel_radius=2): 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) + neighbor_norm = padded_norm[:, :, y_start:y_start + h_len, x_start:x_start + w_len] # Cosine similarity per pixel sim = (center_norm * neighbor_norm).sum(dim=1) # [B, H, W] similarities.append(sim) @@ -45,26 +46,6 @@ def compute_local_relationships(c, h_len, w_len, kernel_radius=2): 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. @@ -88,7 +69,7 @@ def detect_anomalies_adaptive(r_current, r_reference): return anomaly.unsqueeze(-1) # [B, L, 1] -def infer_color_from_neighbors(c, r_ref, anomaly_mag, h_len, w_len, kernel_radius=2): +def infer_color_from_neighbors(c, 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. diff --git a/nodes/anchor.py b/nodes/anchor.py index 55ed799..d347ff7 100644 --- a/nodes/anchor.py +++ b/nodes/anchor.py @@ -74,7 +74,6 @@ def _build_adaptive_anchor_fn(lcs_data, mode, intensity, mask, "envelope": None, "correction_index": 0, "r_ema": None, - "c_ema": None, "prev_c_mean": None, "drift_sum": 0.0, "drift_count": 0, @@ -140,10 +139,8 @@ def _build_adaptive_anchor_fn(lcs_data, mode, intensity, mask, decay = 0.8 if state["r_ema"] is None: state["r_ema"] = r_current.detach().clone() - state["c_ema"] = c_norm.detach().clone() else: state["r_ema"] = decay * state["r_ema"] + (1 - decay) * r_current.detach() - state["c_ema"] = decay * state["c_ema"] + (1 - decay) * c_norm.detach() # Collect step-to-step drift for auto_intensity (self_anchor) c_mean_now = c_norm.detach().mean(dim=1, keepdim=True) @@ -225,18 +222,16 @@ def _build_adaptive_anchor_fn(lcs_data, mode, intensity, mask, if state["r_ema"] is None: # Seed EMA — first step, no correction yet (anomalies will be zero) state["r_ema"] = r_current.detach().clone() - state["c_ema"] = c_norm.detach().clone() anomaly_mag = detect_anomalies_adaptive(r_current, state["r_ema"]) c_corrected = infer_color_from_neighbors( - c_norm, state["r_ema"], anomaly_mag, h_len, w_len + c_norm, anomaly_mag, h_len, w_len ) new_c_norm = c_norm + step_strength * (c_corrected - c_norm) # Update EMA (slow decay during correction) decay = 0.95 state["r_ema"] = decay * state["r_ema"] + (1 - decay) * r_current.detach() - state["c_ema"] = decay * state["c_ema"] + (1 - decay) * c_norm.detach() state["prev_c_mean"] = c_norm.detach().mean(dim=1, keepdim=True) # --- Apply mask ---