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
This commit is contained in:
+5
-3
@@ -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]
|
||||
|
||||
+4
-3
@@ -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))
|
||||
|
||||
|
||||
+6
-25
@@ -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.
|
||||
|
||||
+1
-6
@@ -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 ---
|
||||
|
||||
Reference in New Issue
Block a user