diff --git a/README.md b/README.md index c2c6b54..bdcd1f5 100644 --- a/README.md +++ b/README.md @@ -17,7 +17,7 @@ Gradient-based steering in SVD transformation space, trained on internal represe ```sh uv sync --all-groups uv run python nbs/train.py tiny --quick # al dente check -# Training complete. Final loss: -2.9062 +# Training complete. Final loss: -6.1250 uv run python nbs/train.py # full course (Gemma-3-1B) ``` diff --git a/antipasto/peft_utils/layer_selection.py b/antipasto/peft_utils/layer_selection.py index ea43bb8..f971b39 100644 --- a/antipasto/peft_utils/layer_selection.py +++ b/antipasto/peft_utils/layer_selection.py @@ -33,7 +33,6 @@ import torch import torch.nn as nn import torch.nn.functional as F # no not remove from baukit.nethook import TraceDict -from antipasto.train.inner_contrastive_loss import compute_fisher_t from torch.utils.data import DataLoader, Subset from transformers import DataCollatorWithPadding import gc diff --git a/antipasto/train/inner_contrastive_loss.py b/antipasto/train/inner_contrastive_loss.py index 61ca898..b685744 100644 --- a/antipasto/train/inner_contrastive_loss.py +++ b/antipasto/train/inner_contrastive_loss.py @@ -42,156 +42,41 @@ def mask_agg_tokens_dim( return weighted / count -def symlog(x: torch.Tensor) -> torch.Tensor: - """Symmetric log: sign(x) * log(1 + |x|). - - Compresses large values to log-scale while preserving sign and smoothness at zero. - Commonly used for signed values that span many orders of magnitude. - """ - return torch.sign(x) * torch.log1p(x.abs()) - - -def compute_fisher_t( - diff: Float[Tensor, "b r"], - eps: float = 1e-6, - var_floor_frac: float = 0.1, - abs_std_floor: float = 0.05, - detach_std: bool = False, -) -> tuple[Float[Tensor, "r"], dict]: - """ - Compute signed t-statistic per dimension: mu / sqrt(var). - - High |t| = large, consistent separation in that dimension. - Sign indicates direction of separation (cho > rej or cho < rej). - - This is the core of Fisher-based loss: dimensions with high variance - (inconsistent across samples) get downweighted automatically. - - The variance floor prevents t-explosion when variance collapses: - - var_floor = var_floor_frac * mean(var) ensures relative scaling - - abs_std_floor provides absolute minimum (for few samples where variance is noisy) - - Together these cap max |t| to prevent gradient explosion - - Args: - diff: cho-rej difference in projeciton-space [b, r] - eps: numerical stability for variance - var_floor_frac: variance floor as fraction of median std (0.1 = 10%) - abs_std_floor: absolute minimum std (prevents t-explosion with <10 samples) - detach_std: if True, detach std to prevent zero-variance hacking (legacy) - - Returns: - t: signed t-statistic per dimension [r] - info: dict with floor_activation_rate (fraction of dims hitting floor) - """ - # Check for NaNs in input immediately - fail fast to find root cause - if not torch.isfinite(diff).all(): - n_nan = torch.isnan(diff).sum() - n_inf = torch.isinf(diff).sum() - raise ValueError(f"compute_fisher_t received non-finite inputs: {n_nan} NaNs, {n_inf} Infs. " - f"Range: [{diff.min():.2e}, {diff.max():.2e}]. " - "Likely causes: learning rate too high (exploding grads), or SVD projection issues.") - - # Clamp input only to prevent float32 overflow during squaring, not to hide NaNs - diff = diff.clamp(-1e4, 1e4) - - mu = reduce(diff, 'b r -> r', 'mean') - - # Compute standard deviation: std = sqrt(var + eps) - # CRITICAL: eps INSIDE sqrt to bound gradient at 0. d/dx sqrt(x) = 1/(2*sqrt(x)) → ∞ as x→0 - # sqrt(x).clamp() still has infinite gradient at 0; (x + eps).sqrt() doesn't - var_raw = reduce((diff - mu.unsqueeze(0)).pow(2), 'b r -> r', 'mean') - std_raw = (var_raw + eps).sqrt() # eps inside sqrt, not clamp after - - # Std floor: fraction of median std across dims - # This prevents division by tiny numbers in dimensions that haven't learned anything yet - std_median = std_raw.median() - std_floor = max(var_floor_frac * std_median + eps, abs_std_floor) - std = std_raw.clamp(min=std_floor) - - # Track how many dims are hitting the floor (diagnostic for tuning floor params) - floor_activation_rate = (std_raw < std_floor).float().mean().item() - - # Optionally detach std to prevent zero-variance hacking (legacy behavior) - # With floors in place, detach is less necessary but still an option - if detach_std: - std = std.detach() - - t = mu / std # [r] - - info = { - "floor_activation_rate": floor_activation_rate, - "std_floor": std_floor, - "std_min": std_raw.min().item(), - "std_median": std_median.item(), - } - return t, info - - -def compute_fisher_scale( +def compute_fisher_std( diff: Float[Tensor, "b r"], eps: float = 1e-6, var_floor_frac: float = 0.1, abs_std_floor: float = 0.05, - std: Tensor | None = None, - detach_std: bool = False, - return_scaled: bool = False, ) -> tuple[Float[Tensor, "r"], Float[Tensor, "r"], dict]: """Compute (mu, std) over batch with the same flooring rules as compute_fisher_t. This is used when we want a *shared* per-dimension scale (e.g., std from ref) but still want gradients through the *means* of other tensors. - If std is provided, we use it as the denominator (no recomputation/flooring here) - and return either (mu, std, info) or (mu/std, std, info) depending on return_scaled. - Returns: mu: mean over batch per dimension [r] std: floored std per dimension [r] info: diagnostics dict """ if not torch.isfinite(diff).all(): - n_nan = torch.isnan(diff).sum() - n_inf = torch.isinf(diff).sum() - raise ValueError( - "compute_fisher_scale received non-finite inputs: " - f"{n_nan} NaNs, {n_inf} Infs. Range: [{diff.min():.2e}, {diff.max():.2e}]." - ) + raise ValueError(f"compute_fisher_std received non-finite inputs") diff = diff.clamp(-1e4, 1e4) mu = reduce(diff, "b r -> r", "mean") + var_raw = reduce((diff - mu.unsqueeze(0)).pow(2), "b r -> r", "mean") + std_raw = (var_raw + eps).sqrt() - if std is None: - var_raw = reduce((diff - mu.unsqueeze(0)).pow(2), "b r -> r", "mean") - std_raw = (var_raw + eps).sqrt() + std_median = std_raw.median() + std_floor = max(var_floor_frac * std_median + eps, abs_std_floor) + std = std_raw.clamp(min=std_floor) - std_median = std_raw.median() - std_floor = max(var_floor_frac * std_median + eps, abs_std_floor) - std = std_raw.clamp(min=std_floor) - - floor_activation_rate = (std_raw < std_floor).float().mean().item() - - if std_raw.numel() > 0: - std_min = std_raw.min().item() - std_median_val = std_median.item() - else: - std_min = 0.0 - std_median_val = 0.0 - - info = { - "floor_activation_rate": floor_activation_rate, - "std_floor": std_floor, - "std_min": std_min, - "std_median": std_median_val, - } - else: - info = {} - - if detach_std: - std = std.detach() - - if return_scaled: - return mu / std, std, info - return mu, std, info + floor_activation_rate = (std_raw < std_floor).float().mean().item() + info = { + "fisher_floor_rate": floor_activation_rate, + "fisher_std_floor": std_floor, + "fisher_std_min": std_raw.min().item() if std_raw.numel() > 0 else 0.0, + } + return std, info # ============================================================================= # COHERENCE LOSS COMPONENTS @@ -390,9 +275,6 @@ def contrastive_steering_loss_with_ref( fisher_var_floor_frac: float = 0.1, fisher_abs_std_floor: float = 0.05, fisher_detach_std: bool = False, - fisher_stats: dict | None = None, - fisher_stats_key: str | None = None, - fisher_std_ema_beta: float = 0.1, ): """ Bidirectional antisymmetric separation loss for reversible SVD steering adapters. @@ -452,103 +334,55 @@ def contrastive_steering_loss_with_ref( antisym_pos_agg = mask_agg_tokens_dim(delta_pos, hs_mask) # [b, r] antisym_neg_agg = mask_agg_tokens_dim(delta_neg, hs_mask) # [b, r] - # === Fisher t-space: normalize by std to focus on reliable dimensions === - # Use std computed from *reference* for pos/neg/ref, so we live in one geometry. - # EMA on std_ref reduces noise when batch is small. - fisher_info = {} - b = diff_pos_agg.shape[0] - - if fisher_stats is not None and fisher_stats_key is None: - raise ValueError("fisher_stats_key must be provided when fisher_stats is not None") - - fisher_scale_kwargs = dict( + # Fisher normalization: shared std from ref, applied per-sample + std_ref, fisher_info = compute_fisher_std( + diff_ref_agg, var_floor_frac=fisher_var_floor_frac, abs_std_floor=fisher_abs_std_floor, ) - # Compute batch std from ref (with floors), then optionally EMA it. - _mu_ref, std_ref_batch, info_ref = compute_fisher_scale(diff_ref_agg, **fisher_scale_kwargs) - std_ref = std_ref_batch - - if fisher_stats is not None: - ema_key = f"fisher_std_ema/{fisher_stats_key}" - std_ref_detached = std_ref.detach() - if ema_key in fisher_stats: - fisher_stats[ema_key] = (1 - fisher_std_ema_beta) * fisher_stats[ema_key] + fisher_std_ema_beta * std_ref_detached - else: - fisher_stats[ema_key] = std_ref_detached - std_ref = fisher_stats[ema_key].to(device=std_ref.device, dtype=std_ref.dtype) - - # Build Fisher-like per-sample vectors using a shared std_ref. - # std_ref is estimated over the batch (and optionally EMA'd), then applied to each sample. + # Normalize by shared std_ref (optionally detached) std_for_div = std_ref.detach() if fisher_detach_std else std_ref v_ref = diff_ref_agg / std_for_div.unsqueeze(0) # [b, r] v_pos = antisym_pos_agg / std_for_div.unsqueeze(0) # [b, r] v_neg = antisym_neg_agg / std_for_div.unsqueeze(0) # [b, r] - fisher_info = { - "fisher_floor_rate": info_ref["floor_activation_rate"], - "fisher_std_floor": info_ref["std_floor"], - "fisher_std_min": info_ref["std_min"], - } - - # Compute per-sample dot product and cosine of delta vectors - dot_delta = (v_pos * v_neg).sum(dim=-1) # [b] δ+ · δ-, want negative (antisymmetric) - dot_ref = (v_ref * v_ref).sum(dim=-1) # [b] + # Compute per-sample cosine of delta vectors (for diagnostics) cos_delta = F.cosine_similarity(v_pos, v_neg, dim=-1) # [b] cos(δ+, δ-), want -1 - mag_pos = v_pos.norm(p=2, dim=-1) # [b] - mag_neg = v_neg.norm(p=2, dim=-1) # [b] - separation_norm = v_pos.norm(p=2, dim=-1).mean() # scalar for logging + mag_pos = v_pos.norm(p=2, dim=-1) # [b] + mag_neg = v_neg.norm(p=2, dim=-1) # [b] # Orthogonal penalty: penalize energy not in shared antiparallel axis - # Uses v_pos/v_neg (already in Fisher t-space) - # Normalized by dot_ref to be dimensionless and scale-invariant with rank r. if orth_weight > 0: + dot_delta = (v_pos * v_neg).sum(dim=-1) # [b] + dot_ref = (v_ref * v_ref).sum(dim=-1) # [b] mag_sq_pos = mag_pos * mag_pos mag_sq_neg = mag_neg * mag_neg - orth_waste_sq = ((mag_sq_pos + mag_sq_neg) - 2 * dot_delta.abs()).clamp(min=0) - - # Normalize by dot_ref to make dimensionless (comparable to symlog proj_diff) - # dot_ref = ||t_ref||² which scales with rank, so this removes rank dependence orth_ratio = orth_waste_sq / dot_ref.clamp(min=1.0) - - # sqrt(ratio) gives scale-free penalty; eps inside sqrt for gradient stability at 0 loss_orth = (orth_ratio + 1e-6).sqrt() * orth_weight else: - loss_orth = torch.zeros_like(dot_delta) - orth_waste_sq = None + loss_orth = torch.zeros(mag_pos.shape[0], device=mag_pos.device, dtype=mag_pos.dtype) # === Antisymmetry formulation: ALIGN mode === # cos(delta_pos, ref) × cos(delta_neg, ref) < 0 means one aligns, one anti-aligns # with the reference direction. This constrains steering to the ref axis. - - # Compute vector-level cosines using Fisher-weighted vectors (t-statistics) - # This normalizes by std_ref, making dimensions with high variance less influential - cos_pos_ref = F.cosine_similarity(v_pos, v_ref, dim=-1) # [b] - Fisher-weighted - cos_neg_ref = F.cosine_similarity(v_neg, v_ref, dim=-1) # [b] - Fisher-weighted + cos_pos_ref = F.cosine_similarity(v_pos, v_ref, dim=-1) # [b] + cos_neg_ref = F.cosine_similarity(v_neg, v_ref, dim=-1) # [b] - # Make alignment concentration-aware: weight each cosine by how much of the - # full-space delta energy lies in the loss subspace. - # This yields: (axis alignment) × (subspace concentration) + # Concentration-aware: weight by subspace focus (||proj|| / ||full||) cos_pos_ref_used = cos_pos_ref cos_neg_ref_used = cos_neg_ref focus_pos = None focus_neg = None - focus_pos_raw = None - focus_neg_raw = None 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_raw = proj_norm_pos / delta_pos_norm_full.clamp(min=eps) - focus_neg_raw = proj_norm_neg / delta_neg_norm_full.clamp(min=eps) - # Soften: focus^(1-softness). softness=0→raw, 0.5→sqrt, 1→ignore. + focus_pos = proj_norm_pos / delta_pos_norm_full.clamp(min=eps) + focus_neg = proj_norm_neg / delta_neg_norm_full.clamp(min=eps) if focus_softness > 0: - focus_pos = focus_pos_raw.pow(1.0 - focus_softness) - focus_neg = focus_neg_raw.pow(1.0 - focus_softness) - else: - focus_pos = focus_pos_raw - focus_neg = focus_neg_raw + focus_pos = focus_pos.pow(1.0 - focus_softness) + focus_neg = focus_neg.pow(1.0 - focus_softness) cos_pos_ref_used = cos_pos_ref * focus_pos cos_neg_ref_used = cos_neg_ref * focus_neg @@ -559,86 +393,42 @@ def contrastive_steering_loss_with_ref( # Scale to bounded range [-30, 30] regardless of rank # cos_product_used ∈ [-1, 1], scaled gives consistent gradient magnitude PROJ_SCALE = 30.0 - r = antisym_pos_agg.shape[-1] # Keep for diagnostics shifted = cos_product_used * PROJ_SCALE + antisym_margin # [b], ∈ [-30, 30] # Linear + quadratic: linear keeps pushing, quadratic penalizes positive (bad) - # symlog compresses to prevent runaway: proj_raw ∈ [-30, ~90] → loss ∈ [-3.4, 4.5] proj_raw = shifted + F.relu(shifted).pow(2) # [b] - loss_proj = (proj_raw) + loss_orth # [b] + loss_proj = proj_raw + loss_orth # [b] - # For diagnostics: fake per-dim tensor to keep logging API consistent - per_dim_antisym = cos_product_used.unsqueeze(-1).expand(-1, r) # [b, r] - assert torch.isfinite(loss_proj).all(), f"Non-finite projection loss {loss_proj}" result = { "loss_proj": loss_proj, - "dot_delta": dot_delta.mean(), # δ+ · δ-, want large negative - "dot_ref": dot_ref.mean(), "cos_delta": cos_delta.mean(), # cos(δ+, δ-), want -1 - # separation_norm should respect the same token masking as the loss. - # We report the norm of the aggregated separation vector. - "separation_norm": separation_norm, - "mag_plus": mag_pos.mean(), # Magnitude at α=+1 - "mag_minus": mag_neg.mean(), # Magnitude at α=-1 - "mag_ratio": (torch.minimum(mag_pos, mag_neg) / (torch.maximum(mag_pos, mag_neg) + eps)).mean(), # min/max, want close to 1 + "mag_ratio": (torch.minimum(mag_pos, mag_neg) / (torch.maximum(mag_pos, mag_neg) + eps)).mean(), + "mag_plus": mag_pos, # [b], for train_adapter logging + "mag_minus": mag_neg, # [b], for train_adapter logging } - # Alignment diagnostics + # Alignment diagnostics (core geometry) result["cos_pos_ref_mean"] = cos_pos_ref.mean() result["cos_neg_ref_mean"] = cos_neg_ref.mean() result["cos_product_mean"] = cos_product.mean() - # Subspace focus weighting diagnostics (how much delta energy is in loss subspace) + # Subspace focus weighting (how much delta energy is in loss subspace) if delta_pos_norm_full is not None and delta_neg_norm_full is not None: - assert focus_pos is not None and focus_neg is not None and focus_pos_raw is not None - result["focus_pos_mean"] = focus_pos.mean() # Softened if focus_softness > 0 + assert focus_pos is not None and focus_neg is not None + result["focus_pos_mean"] = focus_pos.mean() result["focus_neg_mean"] = focus_neg.mean() - if focus_softness > 0: - result["focus_pos_raw_mean"] = focus_pos_raw.mean() - result["focus_neg_raw_mean"] = focus_neg_raw.mean() - result["cos_pos_ref_used_mean"] = cos_pos_ref_used.mean() - result["cos_neg_ref_used_mean"] = cos_neg_ref_used.mean() result["cos_product_used_mean"] = cos_product_used.mean() if orth_weight > 0: result["loss_orth"] = loss_orth.mean() - result["orth_waste_sq"] = orth_waste_sq.mean() - result["orth_ratio"] = orth_ratio.mean() # Normalized metric for comparison - result["antisym_separation_ratio"] = (-dot_delta / dot_ref.abs().clamp(min=0.1)).mean() result.update(fisher_info) # Add floor diagnostics - # Loss component metrics (shifted is now [b], shifted_diag is [b, r] for legacy diagnostics) - past_margin = (shifted < 0).float().mean() # Fraction of batch past margin (good) - quad_penalty = F.relu(shifted).pow(2) # [b] - quadratic penalty on bad samples - result["straddle_frac"] = past_margin.item() # Want high (all samples past margin) - result["antisym_mean"] = per_dim_antisym.mean().item() # Avg antisymmetry (want << 0) - result["shifted_mean"] = shifted.mean().item() # Mean shifted value (want negative) - result["proj_raw"] = proj_raw.mean().item() # Pre-symlog loss (want negative) - result["quad_penalty"] = quad_penalty.mean().item() # Quadratic penalty on bad dims (want ~0) - result["antisym_margin"] = antisym_margin # The margin used - result["ref_norm_sq_mean"] = diff_ref_agg.pow(2).sum(dim=-1).mean().item() # Mean ||ref||² (for diagnostics) - - # Subspace concentration diagnostics - if delta_pos_norm_full is not None and delta_neg_norm_full is not None: - # Ratio of projected energy to full-space energy - proj_norm_pos = antisym_pos_agg.norm(dim=-1) # [b] - proj_norm_neg = antisym_neg_agg.norm(dim=-1) # [b] - subspace_ratio_pos = (proj_norm_pos / delta_pos_norm_full.clamp(min=eps)).mean() - subspace_ratio_neg = (proj_norm_neg / delta_neg_norm_full.clamp(min=eps)).mean() - result["subspace_ratio_pos"] = subspace_ratio_pos.item() # Want close to 1 - result["subspace_ratio_neg"] = subspace_ratio_neg.item() # Want close to 1 + # Loss component metrics + result["quad_penalty"] = F.relu(shifted).pow(2).mean().item() - # If this is non-zero, we're living in the eps clamp regime and gradients can get sharp. - norm_prod = delta_pos_norm_full * delta_neg_norm_full # [b] - result["delta_full_norm_prod_min"] = norm_prod.min().item() - result["delta_full_norm_prod_mean"] = norm_prod.mean().item() - result["delta_full_norm_prod_clamp_frac"] = (norm_prod < eps).float().mean().item() - result["delta_pos_norm_full_min"] = delta_pos_norm_full.min().item() - result["delta_neg_norm_full_min"] = delta_neg_norm_full.min().item() - return result diff --git a/antipasto/train/train_adapter.py b/antipasto/train/train_adapter.py index fd08d3a..a8cb3c5 100644 --- a/antipasto/train/train_adapter.py +++ b/antipasto/train/train_adapter.py @@ -258,8 +258,6 @@ def compute_batch_loss( fisher_var_floor_frac=config.fisher_var_floor_frac, fisher_abs_std_floor=config.fisher_abs_std_floor, fisher_detach_std=config.fisher_detach_std, - fisher_stats=flip_stats, - fisher_stats_key=basis_module_path, ) proj_losses = {basis_module_path: loss_dict["loss_proj"]}