From 4b316f038f09c5398401004574da1e57f45dd461 Mon Sep 17 00:00:00 2001 From: wassname <1103714+wassname@users.noreply.github.com> Date: Wed, 14 Jan 2026 18:26:23 +0800 Subject: [PATCH] computes focus in raw space (not fisher which leads to invalid ratios) --- antipasto/train/inner_contrastive_loss.py | 6 ++++-- 1 file changed, 4 insertions(+), 2 deletions(-) diff --git a/antipasto/train/inner_contrastive_loss.py b/antipasto/train/inner_contrastive_loss.py index b685744..5cd500b 100644 --- a/antipasto/train/inner_contrastive_loss.py +++ b/antipasto/train/inner_contrastive_loss.py @@ -371,6 +371,8 @@ def contrastive_steering_loss_with_ref( cos_neg_ref = F.cosine_similarity(v_neg, v_ref, dim=-1) # [b] # Concentration-aware: weight by subspace focus (||proj|| / ||full||) + # focus ∈ [0, 1] = fraction of delta energy in loss subspace + # Clamped to [0, 1] because numerically can exceed 1 due to Fisher weighting mismatch cos_pos_ref_used = cos_pos_ref cos_neg_ref_used = cos_neg_ref focus_pos = None @@ -378,8 +380,8 @@ def contrastive_steering_loss_with_ref( if delta_pos_norm_full is not None and delta_neg_norm_full is not None: proj_norm_pos = antisym_pos_agg.norm(dim=-1) # [b] proj_norm_neg = antisym_neg_agg.norm(dim=-1) # [b] - focus_pos = proj_norm_pos / delta_pos_norm_full.clamp(min=eps) - focus_neg = proj_norm_neg / delta_neg_norm_full.clamp(min=eps) + focus_pos = (proj_norm_pos / delta_pos_norm_full.clamp(min=eps)).clamp(max=1.0) + focus_neg = (proj_norm_neg / delta_neg_norm_full.clamp(min=eps)).clamp(max=1.0) if focus_softness > 0: focus_pos = focus_pos.pow(1.0 - focus_softness) focus_neg = focus_neg.pow(1.0 - focus_softness)