mirror of
https://github.com/wassname/AntiPaSTO.git
synced 2026-09-10 11:40:28 +08:00
simplify
This commit is contained in:
@@ -1,28 +1,22 @@
|
||||
#!/usr/bin/env python3
|
||||
"""Centralized layer selection logic for AntiPaSTO training.
|
||||
|
||||
FLOW OVERVIEW (v2.2+)
|
||||
=====================
|
||||
1. compute_simple_layer_selection() [RECOMMENDED]:
|
||||
- Computes SVD for all linear layers (needed for adapter init anyway)
|
||||
- Selects layers uniformly across valid depth range
|
||||
- Uses top-r singular values for dimension selection (no gradient)
|
||||
- Computes weight-only subspaces (write, write_x_notlogits)
|
||||
- No backward pass = no OOM on large models (12B+)
|
||||
SIMPLIFIED FOR PUBLICATION (2026-01-13):
|
||||
Only 3 loss_subspace types are supported: taskdiff_x_suppressed_x_write (default), write, taskdiff.
|
||||
Deprecated gradient-based layer selection has been removed.
|
||||
|
||||
2. compute_gradient_layer_selection() [DEPRECATED]:
|
||||
- Expensive backward pass for gradient-based ranking
|
||||
- Ablations show it doesn't improve over simple selection
|
||||
- Kept for research/debugging purposes
|
||||
FLOW:
|
||||
1. compute_simple_layer_selection():
|
||||
- Computes SVD for all linear layers (needed for adapter init)
|
||||
- Selects layers uniformly across valid depth range
|
||||
- Computes subspaces: write, taskdiff, taskdiff_x_suppressed_x_write
|
||||
|
||||
Key functions:
|
||||
- compute_simple_layer_selection(): Uniform layer selection, top-S dims, weight-only subspaces
|
||||
- compute_gradient_layer_selection(): Gradient-based selection (deprecated, OOMs on 12B+)
|
||||
- compute_simple_layer_selection(): Main entry point for layer/subspace selection
|
||||
- find_linear_layers(): Discover all linear modules in model
|
||||
- resolve_target_modules(): Expand "residual-writers" etc. to concrete module lists
|
||||
|
||||
Subspace operations (compute_write_subspace, compute_write_x_notlogits, etc.)
|
||||
are in antipasto/peft_utils/subspaces.py
|
||||
Subspace operations are in antipasto/peft_utils/subspaces.py
|
||||
"""
|
||||
import re
|
||||
import pandas as pd
|
||||
@@ -41,23 +35,11 @@ from torch.utils.data import DataLoader, Subset
|
||||
from transformers import DataCollatorWithPadding
|
||||
import gc
|
||||
from antipasto.peft_utils.subspaces import (
|
||||
compute_lm_head_subspace,
|
||||
compute_lm_head_svd,
|
||||
compute_suppressed_from_hidden_states,
|
||||
compute_churn_from_hidden_states,
|
||||
compute_churn_constructive_from_hidden_states,
|
||||
compute_churn_suppressive_from_hidden_states,
|
||||
compute_task_diff_from_hidden_states,
|
||||
compute_task_diff_constructive_from_hidden_states,
|
||||
compute_task_read_subspace,
|
||||
compute_task_lm_head_subspace,
|
||||
compute_task_wnr_subspace,
|
||||
compute_module_subspace_from_svds,
|
||||
compute_write_not_read_subspace,
|
||||
compute_stenographic_subspace,
|
||||
compute_write_x_notlogits_subspace,
|
||||
compute_logits_tail_subspace,
|
||||
compute_taskdiff_x_write_x_notlogits_subspace,
|
||||
find_write_modules,
|
||||
find_read_modules,
|
||||
approx_intersection,
|
||||
@@ -944,304 +926,46 @@ def compute_simple_layer_selection(
|
||||
S=lm_head_S_full[:INTERMEDIATE_SUBSPACE_RANK].to(device=device, dtype=dtype).detach(),
|
||||
)
|
||||
if lm_head_sub is not None:
|
||||
subspaces.set('logits_read', lm_head_sub) # Full Subspace with S
|
||||
|
||||
# write_x_notlogits = write projected into notlogits (compound: needs full geometry)
|
||||
if write_subspace is not None and lm_head_sub is not None:
|
||||
hfl_sub = compute_write_x_notlogits_subspace(
|
||||
write_subspace=write_subspace,
|
||||
lm_head_subspace=lm_head_sub,
|
||||
top_k=INTERMEDIATE_SUBSPACE_RANK, # Full geometry for subtraction
|
||||
)
|
||||
subspaces.set('write_x_notlogits', hfl_sub) # Full Subspace
|
||||
logger.info(f"write_x_notlogits subspace: rank={hfl_sub.V.shape[1]}")
|
||||
|
||||
# Read subspace from read modules (q_proj, k_proj, etc.) row spaces
|
||||
read_modules = find_read_modules(model)
|
||||
read_subspace = compute_module_subspace_from_svds(
|
||||
layer_svds=layer_svd_cpu,
|
||||
layer_info=layer_info_full,
|
||||
module_filter=read_modules,
|
||||
use_column_space=False, # Row space = read directions
|
||||
top_k=INTERMEDIATE_SUBSPACE_RANK,
|
||||
device=device,
|
||||
dtype=dtype,
|
||||
name="read",
|
||||
)
|
||||
if read_subspace is not None:
|
||||
subspaces.set('read', read_subspace) # Full Subspace
|
||||
logger.info(f"read subspace: rank={read_subspace.V.shape[1]}")
|
||||
|
||||
# NEW: Specific read subspaces (Query, Key, Value)
|
||||
for name, suffix in [('query_read', 'q_proj'), ('key_read', 'k_proj'), ('value_read', 'v_proj')]:
|
||||
modules = [m for m in read_modules if suffix in m]
|
||||
if modules:
|
||||
sub = compute_module_subspace_from_svds(
|
||||
layer_svds=layer_svd_cpu,
|
||||
layer_info=layer_info_full,
|
||||
module_filter=modules,
|
||||
use_column_space=False, # Row space
|
||||
top_k=INTERMEDIATE_SUBSPACE_RANK,
|
||||
device=device,
|
||||
dtype=dtype,
|
||||
name=name,
|
||||
)
|
||||
if sub is not None:
|
||||
subspaces.set(name, sub)
|
||||
|
||||
# Attention Read (Union of Q, K, V)
|
||||
attn_read_modules = [m for m in read_modules if any(s in m for s in ['q_proj', 'k_proj', 'v_proj'])]
|
||||
if attn_read_modules:
|
||||
attn_read_subspace = compute_module_subspace_from_svds(
|
||||
layer_svds=layer_svd_cpu,
|
||||
layer_info=layer_info_full,
|
||||
module_filter=attn_read_modules,
|
||||
use_column_space=False, # Row space = read directions
|
||||
top_k=INTERMEDIATE_SUBSPACE_RANK,
|
||||
device=device,
|
||||
dtype=dtype,
|
||||
name="attn_read",
|
||||
)
|
||||
|
||||
# NEW: Attention Sink = Write - Attention_Read
|
||||
if write_subspace is not None and attn_read_subspace is not None:
|
||||
attn_sink = compute_write_not_read_subspace(
|
||||
write_subspace=write_subspace,
|
||||
read_subspace=attn_read_subspace,
|
||||
lm_head_subspace=None, # Don't subtract lm_head for pure attention sink
|
||||
top_k=INTERMEDIATE_SUBSPACE_RANK,
|
||||
)
|
||||
subspaces.set('attention_sink', attn_sink)
|
||||
logger.info(f"attention_sink subspace: intermediate={attn_sink.V.shape[1]}, stored={top_k}")
|
||||
|
||||
# NEW: Communication Channel = Write & Read
|
||||
if write_subspace is not None and read_subspace is not None:
|
||||
comm_channel = approx_intersection(write_subspace, read_subspace, top_k=INTERMEDIATE_SUBSPACE_RANK)
|
||||
subspaces.set('communication_channel', comm_channel)
|
||||
logger.info(f"communication_channel subspace: intermediate={comm_channel.V.shape[1]}, stored={top_k}")
|
||||
|
||||
# write_not_read = write - read - lm_head (what's written but ignored)
|
||||
if write_subspace is not None and read_subspace is not None:
|
||||
wnr_sub = compute_write_not_read_subspace(
|
||||
write_subspace=write_subspace,
|
||||
read_subspace=read_subspace,
|
||||
lm_head_subspace=lm_head_sub, # Also subtract lm_head if available
|
||||
top_k=INTERMEDIATE_SUBSPACE_RANK,
|
||||
)
|
||||
subspaces.set('write_not_read', wnr_sub)
|
||||
logger.info(f"write_not_read subspace: intermediate={wnr_sub.V.shape[1]}, stored={top_k}")
|
||||
subspaces.set('logits_read', lm_head_sub)
|
||||
|
||||
# =========================================================================
|
||||
# WANDA_X_LOGITS_NULL SUBSPACE
|
||||
# =========================================================================
|
||||
if loss_subspace == 'wanda_x_notlogits':
|
||||
if hs_stacked is None:
|
||||
raise ValueError("wanda_x_notlogits requires hidden states")
|
||||
|
||||
logger.info("Computing wanda_x_notlogits subspace...")
|
||||
# Compute lm_head SVD to get S and Vh
|
||||
W = model.lm_head.weight.data.float().cpu()
|
||||
_, S_lm_head, Vh_lm_head = torch.linalg.svd(W, full_matrices=False)
|
||||
|
||||
# logits_tail uses its own internal PCA - pass intermediate rank
|
||||
# for full geometry, then truncate at storage
|
||||
active_null_sub = compute_logits_tail_subspace(
|
||||
hidden_states=hs_stacked,
|
||||
lm_head_S=S_lm_head,
|
||||
lm_head_Vh=Vh_lm_head,
|
||||
top_k=INTERMEDIATE_SUBSPACE_RANK
|
||||
)
|
||||
subspaces.set('wanda_x_notlogits', active_null_sub)
|
||||
logger.info(f"wanda_x_notlogits subspace: intermediate={active_null_sub.V.shape[1]}, stored={top_k}")
|
||||
|
||||
# Random subspace (sanity baseline)
|
||||
if loss_subspace == 'random':
|
||||
if write_subspace is not None:
|
||||
d_model = write_subspace.V.shape[0]
|
||||
elif lm_head_sub is not None:
|
||||
d_model = lm_head_sub.V.shape[0]
|
||||
else:
|
||||
# Fall back to any SVD to infer d_model (input dim for residual-connected linears)
|
||||
any_path = next(iter(layer_svd_cpu.keys()))
|
||||
_, _, any_Vh = layer_svd_cpu[any_path]
|
||||
d_model = any_Vh.shape[1]
|
||||
gen_random = torch.Generator(device=device)
|
||||
gen_random.manual_seed(_stable_u32(f"{seed}:loss_subspace:random:{d_model}:{top_k}"))
|
||||
random_basis = torch.randn(d_model, top_k, device=device, dtype=dtype, generator=gen_random)
|
||||
random_basis = torch.linalg.qr(random_basis.float())[0].to(dtype) # Orthonormalize
|
||||
subspaces.set('random', random_basis)
|
||||
logger.info(f"random subspace: shape={random_basis.shape}")
|
||||
|
||||
# =========================================================================
|
||||
# ACTIVATION-BASED SUBSPACES (uses hidden states collected earlier)
|
||||
# Always computed when hidden states are available (dataset_pt was provided)
|
||||
# ACTIVATION-BASED SUBSPACES (simplified for publication)
|
||||
# Only computes subspaces needed for 3 supported loss_subspace types:
|
||||
# - taskdiff, taskdiff_x_suppressed_x_write (default)
|
||||
# =========================================================================
|
||||
if hs_stacked is not None:
|
||||
logger.info(f"Computing activation-based subspaces (taskdiff, churn, suppressed, etc.)...")
|
||||
logger.info(f"Computing activation-based subspaces (taskdiff, suppressed)...")
|
||||
|
||||
# Compute taskdiff subspace - use INTERMEDIATE rank for full geometry
|
||||
# Compute taskdiff subspace - PCA on cho-rej difference
|
||||
task_diff_subspace = compute_task_diff_from_hidden_states(
|
||||
hidden_states=hs_stacked,
|
||||
top_k=INTERMEDIATE_SUBSPACE_RANK,
|
||||
layer_frac=loss_hs_frac_for_task,
|
||||
)
|
||||
subspaces.set('taskdiff', task_diff_subspace) # Full Subspace with S for energy thresholding
|
||||
|
||||
# taskdiff_write: per-layer contributions that differ between cho/rej
|
||||
taskdiff_write_subspace = compute_task_diff_from_hidden_states(
|
||||
hidden_states=hs_stacked,
|
||||
top_k=INTERMEDIATE_SUBSPACE_RANK,
|
||||
layer_frac=loss_hs_frac_for_task,
|
||||
use_layer_diffs=True,
|
||||
)
|
||||
subspaces.set('taskdiff_write', taskdiff_write_subspace)
|
||||
subspaces.set('taskdiff', task_diff_subspace)
|
||||
|
||||
# Compute suppressed subspace (from layer diffs) - use INTERMEDIATE rank
|
||||
# Compute suppressed subspace (written but erased by later layers)
|
||||
suppressed_subspace = compute_suppressed_from_hidden_states(
|
||||
hidden_states=hs_stacked,
|
||||
lm_head_subspace=lm_head_sub,
|
||||
top_k=INTERMEDIATE_SUBSPACE_RANK,
|
||||
)
|
||||
subspaces.set('suppressed', suppressed_subspace) # Full Subspace with S for energy thresholding
|
||||
subspaces.set('suppressed', suppressed_subspace)
|
||||
|
||||
# Compute churn subspace - use INTERMEDIATE rank
|
||||
churn_subspace = compute_churn_from_hidden_states(
|
||||
hidden_states=hs_stacked,
|
||||
top_k=INTERMEDIATE_SUBSPACE_RANK,
|
||||
)
|
||||
subspaces.set('churn', churn_subspace) # Full Subspace with S for energy thresholding
|
||||
|
||||
# taskdiff_x_suppressed = taskdiff ∩ suppressed (compound: needs full geometry)
|
||||
# taskdiff_x_suppressed = taskdiff ∩ suppressed (stenographic signal)
|
||||
steno_subspace = compute_stenographic_subspace(
|
||||
task_diff_subspace=task_diff_subspace,
|
||||
suppressed_subspace=suppressed_subspace,
|
||||
top_k=INTERMEDIATE_SUBSPACE_RANK, # Full geometry for intersection
|
||||
)
|
||||
subspaces.set('taskdiff_x_suppressed', steno_subspace) # Store full Subspace for energy thresholding
|
||||
|
||||
# NEW: Prediction Suppression = Suppressed & Read
|
||||
if suppressed_subspace is not None and read_subspace is not None:
|
||||
pred_supp = approx_intersection(suppressed_subspace, read_subspace, top_k=INTERMEDIATE_SUBSPACE_RANK)
|
||||
subspaces.set('prediction_suppression', pred_supp)
|
||||
logger.info(f"prediction_suppression subspace: intermediate={pred_supp.V.shape[1]}, stored={top_k}")
|
||||
|
||||
# Compound subspaces: taskdiff ∩ X (all need full geometry for intersection)
|
||||
if hfl_sub is not None:
|
||||
task_intersect_hfl = approx_intersection(task_diff_subspace, hfl_sub, top_k=INTERMEDIATE_SUBSPACE_RANK)
|
||||
subspaces.set('taskdiff_x_write_x_notlogits', task_intersect_hfl)
|
||||
|
||||
if write_subspace is not None:
|
||||
task_intersect_write = approx_intersection(task_diff_subspace, write_subspace, top_k=INTERMEDIATE_SUBSPACE_RANK)
|
||||
subspaces.set('taskdiff_x_write', task_intersect_write)
|
||||
|
||||
task_intersect_churn = approx_intersection(task_diff_subspace, churn_subspace, top_k=INTERMEDIATE_SUBSPACE_RANK)
|
||||
subspaces.set('taskdiff_x_churn', task_intersect_churn)
|
||||
|
||||
task_intersect_steno = approx_intersection(task_diff_subspace, steno_subspace, top_k=INTERMEDIATE_SUBSPACE_RANK)
|
||||
subspaces.set('taskdiff_x_taskdiff_x_suppressed', task_intersect_steno)
|
||||
|
||||
# Churn variants (constructive = magnitude increase, suppressive = magnitude decrease)
|
||||
churn_constructive = compute_churn_constructive_from_hidden_states(
|
||||
hidden_states=hs_stacked,
|
||||
top_k=INTERMEDIATE_SUBSPACE_RANK,
|
||||
)
|
||||
subspaces.set('churn_constructive', churn_constructive)
|
||||
|
||||
churn_suppressive = compute_churn_suppressive_from_hidden_states(
|
||||
hidden_states=hs_stacked,
|
||||
top_k=INTERMEDIATE_SUBSPACE_RANK,
|
||||
)
|
||||
subspaces.set('churn_suppressive', churn_suppressive)
|
||||
|
||||
# task_diff_constructive = directions where task magnitude INCREASES
|
||||
task_diff_constructive = compute_task_diff_constructive_from_hidden_states(
|
||||
hidden_states=hs_stacked,
|
||||
top_k=INTERMEDIATE_SUBSPACE_RANK,
|
||||
layer_range=(min_adapter_layer_frac, loss_layer_frac),
|
||||
)
|
||||
subspaces.set('taskdiff_constructive', task_diff_constructive)
|
||||
|
||||
# Task compound subspaces with specific weight subspaces
|
||||
if write_subspace is not None:
|
||||
task_write = approx_intersection(task_diff_subspace, write_subspace, top_k=INTERMEDIATE_SUBSPACE_RANK)
|
||||
subspaces.set('taskdiff_x_write', task_write)
|
||||
|
||||
if read_subspace is not None:
|
||||
# task ∩ read: task-discriminative directions that are read by attention/MLP inputs
|
||||
task_read = approx_intersection(task_diff_subspace, read_subspace, top_k=INTERMEDIATE_SUBSPACE_RANK)
|
||||
subspaces.set('taskdiff_read', task_read)
|
||||
|
||||
# task ∩ lm_head: task directions that affect output logits
|
||||
if lm_head_sub is not None:
|
||||
task_lm_head = compute_task_lm_head_subspace(
|
||||
task_diff_subspace=task_diff_subspace,
|
||||
lm_head_subspace=lm_head_sub,
|
||||
top_k=INTERMEDIATE_SUBSPACE_RANK,
|
||||
)
|
||||
subspaces.set('taskdiff_x_logits_read', task_lm_head)
|
||||
|
||||
if wnr_sub is not None:
|
||||
task_wnr = compute_task_wnr_subspace(
|
||||
task_diff_subspace=task_diff_subspace,
|
||||
write_not_read_subspace=wnr_sub,
|
||||
top_k=INTERMEDIATE_SUBSPACE_RANK,
|
||||
)
|
||||
subspaces.set('taskdiff_x_write_not_read', task_wnr)
|
||||
|
||||
# Additional compound intersections for sweeps
|
||||
task_intersect_churn_constructive = approx_intersection(task_diff_subspace, churn_constructive, top_k=INTERMEDIATE_SUBSPACE_RANK)
|
||||
subspaces.set('taskdiff_x_churn_constructive', task_intersect_churn_constructive)
|
||||
|
||||
if wnr_sub is not None and write_subspace is not None:
|
||||
taskdiff_write_intersect_wnr = approx_intersection(taskdiff_write_subspace, wnr_sub, top_k=INTERMEDIATE_SUBSPACE_RANK)
|
||||
subspaces.set('taskdiff_write_x_write_not_read', taskdiff_write_intersect_wnr)
|
||||
|
||||
taskdiff_write_intersect_suppressed = approx_intersection(taskdiff_write_subspace, suppressed_subspace, top_k=INTERMEDIATE_SUBSPACE_RANK)
|
||||
subspaces.set('taskdiff_write_x_suppressed', taskdiff_write_intersect_suppressed)
|
||||
|
||||
taskdiff_write_intersect_churn = approx_intersection(taskdiff_write_subspace, churn_subspace, top_k=INTERMEDIATE_SUBSPACE_RANK)
|
||||
subspaces.set('taskdiff_write_x_churn', taskdiff_write_intersect_churn)
|
||||
|
||||
# task_constructive_intersect_* (task_diff_constructive ∩ X)
|
||||
if hfl_sub is not None:
|
||||
task_constructive_intersect_hfl = approx_intersection(task_diff_constructive, hfl_sub, top_k=INTERMEDIATE_SUBSPACE_RANK)
|
||||
subspaces.set('taskdiff_constructive_x_write_x_notlogits', task_constructive_intersect_hfl)
|
||||
|
||||
task_constructive_intersect_suppressed = approx_intersection(task_diff_constructive, suppressed_subspace, top_k=INTERMEDIATE_SUBSPACE_RANK)
|
||||
subspaces.set('taskdiff_constructive_x_suppressed', task_constructive_intersect_suppressed)
|
||||
|
||||
task_constructive_intersect_churn = approx_intersection(task_diff_constructive, churn_subspace, top_k=INTERMEDIATE_SUBSPACE_RANK)
|
||||
subspaces.set('taskdiff_constructive_x_churn', task_constructive_intersect_churn)
|
||||
|
||||
task_constructive_intersect_churn_constructive = approx_intersection(task_diff_constructive, churn_constructive, top_k=INTERMEDIATE_SUBSPACE_RANK)
|
||||
subspaces.set('taskdiff_constructive_x_churn_constructive', task_constructive_intersect_churn_constructive)
|
||||
|
||||
# taskdiff_x_suppressed_x_* (taskdiff_x_suppressed ∩ X)
|
||||
steno_intersect_churn = approx_intersection(steno_subspace, churn_subspace, top_k=INTERMEDIATE_SUBSPACE_RANK)
|
||||
subspaces.set('taskdiff_x_suppressed_x_churn', steno_intersect_churn)
|
||||
subspaces.set('taskdiff_x_suppressed', steno_subspace)
|
||||
|
||||
# taskdiff_x_suppressed_x_write = steno ∩ write (default loss subspace)
|
||||
if write_subspace is not None:
|
||||
steno_intersect_write = approx_intersection(steno_subspace, write_subspace, top_k=INTERMEDIATE_SUBSPACE_RANK)
|
||||
subspaces.set('taskdiff_x_suppressed_x_write', steno_intersect_write)
|
||||
|
||||
if hfl_sub is not None:
|
||||
steno_intersect_hfl = approx_intersection(steno_subspace, hfl_sub, top_k=INTERMEDIATE_SUBSPACE_RANK)
|
||||
subspaces.set('taskdiff_x_suppressed_x_write_x_notlogits', steno_intersect_hfl)
|
||||
|
||||
# taskdiff_write_x_notlogits: task-discriminative directions in write ∩ notlogits
|
||||
# (like write_x_notlogits but weighted by cho-rej difference)
|
||||
if write_subspace is not None and lm_head_sub is not None:
|
||||
taskdiff_x_write_x_notlogits = compute_taskdiff_x_write_x_notlogits_subspace(
|
||||
hidden_states=hs_stacked,
|
||||
write_subspace=write_subspace,
|
||||
lm_head_S=lm_head_S_full,
|
||||
lm_head_Vh=lm_head_Vh_full,
|
||||
top_k=INTERMEDIATE_SUBSPACE_RANK,
|
||||
layer_frac=loss_layer_frac,
|
||||
)
|
||||
subspaces.set('taskdiff_write_x_notlogits', taskdiff_x_write_x_notlogits)
|
||||
|
||||
logger.info(f"Activation-based subspaces computed (intermediate={INTERMEDIATE_SUBSPACE_RANK}, stored={top_k}): {[k for k in subspaces.keys() if 'task' in k or 'steno' in k or 'churn' in k]}")
|
||||
logger.info(f"Activation-based subspaces computed: taskdiff, suppressed, taskdiff_x_suppressed, taskdiff_x_suppressed_x_write")
|
||||
|
||||
# Cleanup hidden states if collected
|
||||
if hs_stacked is not None:
|
||||
|
||||
@@ -18,6 +18,19 @@ Naming conventions for subspace functions:
|
||||
For geometric intuition and taxonomy of named subspaces, see docs/steering_methods.qmd.
|
||||
|
||||
All bases are detached (frozen) to prevent gradient hacking.
|
||||
|
||||
SIMPLIFICATION NOTE (2026-01-13):
|
||||
This file was simplified for publication. Many experimental subspace types were removed.
|
||||
Only the following are supported:
|
||||
- taskdiff_x_suppressed_x_write (default)
|
||||
- write
|
||||
- taskdiff
|
||||
|
||||
See git history (pre-2026-01-13) for removed experimental types:
|
||||
- write_not_read, stenographic, write_x_notlogits, logits_tail
|
||||
- taskdiff_x_write_x_notlogits, task_read, task_lm_head, task_wnr
|
||||
- churn variants (constructive, suppressive), taskdiff_constructive
|
||||
- lm_head, embed subspaces
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
@@ -647,37 +660,6 @@ def compute_suppressed_from_hidden_states(
|
||||
return Subspace(V_supp, name="suppressed", S=S_supp)
|
||||
|
||||
|
||||
# ============================================================================
|
||||
# Legacy subspace computation (kept for compatibility)
|
||||
# ============================================================================
|
||||
|
||||
def compute_lm_head_subspace(model: nn.Module, top_k: int = 256) -> Subspace:
|
||||
"""Compute subspace read by lm_head (directions that affect output logits).
|
||||
|
||||
Uses right singular vectors (V) of lm_head.weight since it reads from residual.
|
||||
lm_head computes logits = h @ W.T, so it reads directions in row-space of W.
|
||||
|
||||
Args:
|
||||
model: Model with lm_head
|
||||
top_k: Number of components
|
||||
|
||||
Returns:
|
||||
Subspace of directions lm_head reads
|
||||
"""
|
||||
device = next(model.parameters()).device
|
||||
dtype = next(model.parameters()).dtype
|
||||
|
||||
# SVD: W = U @ S @ Vh, row-space = span of Vh rows = right singular vectors
|
||||
S, Vh = compute_lm_head_svd(model)
|
||||
V_read: Float[Tensor, "d_model top_k"] = Vh[:top_k, :].T # transpose: [d_model, top_k]
|
||||
S_read = S[:top_k].to(dtype).to(device).detach()
|
||||
|
||||
V_read = V_read.to(dtype).to(device).detach()
|
||||
logger.debug(f"logits_read subspace: rank={V_read.shape[1]}")
|
||||
|
||||
return Subspace(V_read, name="logits_read", S=S_read)
|
||||
|
||||
|
||||
def compute_lm_head_svd(model: nn.Module) -> tuple[Tensor, Tensor]:
|
||||
"""Return (S, Vh) for lm_head.weight SVD.
|
||||
|
||||
@@ -694,68 +676,6 @@ def compute_lm_head_svd(model: nn.Module) -> tuple[Tensor, Tensor]:
|
||||
return S, Vh
|
||||
|
||||
|
||||
def compute_embed_subspace(model: nn.Module, top_k: int = 256) -> Subspace:
|
||||
"""Compute subspace written by embedding layer.
|
||||
|
||||
Uses left singular vectors (U) of embed_tokens.weight since it writes to residual.
|
||||
|
||||
Args:
|
||||
model: Model with embed_tokens
|
||||
top_k: Number of components
|
||||
|
||||
Returns:
|
||||
Subspace of directions embedding writes
|
||||
"""
|
||||
device = next(model.parameters()).device
|
||||
dtype = next(model.parameters()).dtype
|
||||
|
||||
# embed_tokens.weight: [vocab_size, d_model], output is row-indexed
|
||||
# Column space = write directions
|
||||
W = model.model.embed_tokens.weight.data # [vocab, d_model]
|
||||
|
||||
# Column space via transpose
|
||||
U, S, _ = torch.linalg.svd(W.T.float().cpu(), full_matrices=False)
|
||||
V_write = U[:, :top_k] # [d_model, top_k]
|
||||
|
||||
V_write = V_write.to(dtype).to(device).detach()
|
||||
logger.info(f"Embed write subspace: rank={V_write.shape[1]}")
|
||||
|
||||
return Subspace(V_write, name="embed_write")
|
||||
|
||||
|
||||
def compute_write_not_read_subspace(
|
||||
write_subspace: Subspace,
|
||||
read_subspace: Subspace,
|
||||
lm_head_subspace: Optional[Subspace] = None,
|
||||
top_k: int =256,
|
||||
) -> Subspace:
|
||||
"""Compute Write-Not-Read subspace: directions written but not read.
|
||||
|
||||
Notation: WnR = Write_perp_Read = Π_{Read^⊥}(Write).
|
||||
|
||||
If `lm_head_subspace` is provided, also subtract directions readable by
|
||||
the lm_head (since those are "read" at the output interface).
|
||||
|
||||
Args:
|
||||
write_subspace: Subspace of write directions
|
||||
read_subspace: Subspace of read directions
|
||||
top_k: Number of components
|
||||
|
||||
Returns:
|
||||
Subspace of directions written but ignored by reading layers
|
||||
"""
|
||||
wnr = project_subspace_into_perp(write_subspace, read_subspace)
|
||||
if lm_head_subspace is not None:
|
||||
wnr = project_subspace_into_perp(wnr, lm_head_subspace)
|
||||
|
||||
if wnr.rank > top_k:
|
||||
wnr = Subspace(wnr.V[:, :top_k], name="write_not_read")
|
||||
else:
|
||||
wnr.name = "write_not_read"
|
||||
|
||||
return wnr
|
||||
|
||||
|
||||
def compute_stenographic_subspace(
|
||||
task_diff_subspace: Subspace,
|
||||
suppressed_subspace: Subspace,
|
||||
@@ -778,504 +698,6 @@ def compute_stenographic_subspace(
|
||||
return steno
|
||||
|
||||
|
||||
def compute_write_x_notlogits_subspace(
|
||||
write_subspace: Subspace,
|
||||
lm_head_subspace: Subspace,
|
||||
top_k: int =256,
|
||||
) -> Subspace:
|
||||
"""Compute write_x_notlogits: write projected into (logits_read)^perp.
|
||||
|
||||
Notation: write_x_notlogits = Write_perp_logits_read = Π_{(logits_read)^⊥}(Write).
|
||||
|
||||
In code this uses project_subspace_into_perp(write, logits), which performs an
|
||||
orthogonal-complement projection (see project_bases_into_perp docstring), not a set
|
||||
difference.
|
||||
|
||||
These directions are written to residual by model layers but don't affect
|
||||
output logits (lm_head can't read them). Simpler than write_not_read since
|
||||
it ignores layer-to-layer reads.
|
||||
|
||||
Note it includes write to avoid token embeddings that prepopulate the residual stream
|
||||
|
||||
Args:
|
||||
write_subspace: Subspace of write directions
|
||||
lm_head_subspace: Subspace readable by lm_head
|
||||
top_k: Number of components
|
||||
|
||||
Returns:
|
||||
Subspace of directions hidden from final output
|
||||
"""
|
||||
hfl = project_subspace_into_perp(write_subspace, lm_head_subspace)
|
||||
|
||||
if hfl.rank > top_k:
|
||||
hfl = Subspace(hfl.V[:, :top_k], name="write_x_notlogits")
|
||||
else:
|
||||
hfl.name = "write_x_notlogits"
|
||||
|
||||
return hfl
|
||||
|
||||
|
||||
def compute_logits_tail_subspace(
|
||||
hidden_states: Float[Tensor, "batch n_layers_plus1 d_model"],
|
||||
lm_head_S: Float[Tensor, "rank"],
|
||||
lm_head_Vh: Float[Tensor, "rank d_model"],
|
||||
top_k: int = 64,
|
||||
layer_range: Optional[tuple] = None,
|
||||
null_frac: float = 0.5,
|
||||
) -> Subspace:
|
||||
"""Compute wanda_x_notlogits subspace: tail lm_head singular dirs weighted by activation.
|
||||
|
||||
Like write_x_notlogits but empirical: uses actual activations to weight directions.
|
||||
|
||||
Method (WANDA-inspired):
|
||||
1. Take bottom `null_frac` of lm_head singular directions (low S = low output gain)
|
||||
2. Project hidden states into this tail subspace
|
||||
3. Weight each direction by activation magnitude (WANDA: ||X||_2 per direction)
|
||||
4. PCA on weighted projections to find most-used directions within tail space
|
||||
|
||||
This differs from write_x_notlogits (static weight subtraction) by incorporating
|
||||
which directions are actually used, not just which could theoretically be hidden.
|
||||
|
||||
Args:
|
||||
hidden_states: [batch, n_layers+1, d_model] from model output
|
||||
lm_head_S: [rank] singular values of lm_head (descending order from SVD)
|
||||
lm_head_Vh: [rank, d_model] right singular vectors (rows are basis vectors)
|
||||
top_k: Number of components to return
|
||||
layer_range: (start_frac, end_frac) for which layers to use (default 0.3-0.8)
|
||||
null_frac: Fraction of bottom singular directions to use (default 0.5)
|
||||
|
||||
Returns:
|
||||
Subspace of actively-used low-gain directions
|
||||
"""
|
||||
device = hidden_states.device
|
||||
dtype = hidden_states.dtype
|
||||
d_model = hidden_states.shape[-1]
|
||||
n_layers_plus1 = hidden_states.shape[1]
|
||||
n_layers = n_layers_plus1 - 1
|
||||
|
||||
if layer_range is None:
|
||||
layer_range = (0.3, 0.8)
|
||||
|
||||
start_idx = max(1, int(layer_range[0] * n_layers))
|
||||
end_idx = min(n_layers, int(layer_range[1] * n_layers))
|
||||
|
||||
# Get relevant hidden states [batch, selected_layers, d_model]
|
||||
hs_selected = hidden_states[:, start_idx:end_idx, :]
|
||||
|
||||
# Take bottom null_frac of singular directions (low S = lm_head ignores)
|
||||
rank = lm_head_Vh.shape[0]
|
||||
null_start = int((1 - null_frac) * rank)
|
||||
null_rank = rank - null_start
|
||||
|
||||
if null_rank < top_k:
|
||||
logger.warning(f"null_frac={null_frac} gives {null_rank} dims < top_k={top_k}. Expanding.")
|
||||
null_start = max(0, rank - top_k * 2)
|
||||
null_rank = rank - null_start
|
||||
|
||||
# V_tail: [d_model, null_rank] - bottom singular vectors
|
||||
V_tail = lm_head_Vh[null_start:, :].T.to(device).to(dtype) # [d_model, null_rank]
|
||||
S_tail = lm_head_S[null_start:].to(device).float() # [null_rank]
|
||||
|
||||
# Project hidden states into tail subspace
|
||||
hs_flat = hs_selected.reshape(-1, d_model).float()
|
||||
z = hs_flat @ V_tail.float() # [n, null_rank]
|
||||
|
||||
# WANDA-style: weight by activation magnitude (L2 norm per direction)
|
||||
# ||X_j||_2 = sqrt(sum_i x_ij^2), captures total energy in each direction
|
||||
activation_norm = z.norm(dim=0) # [null_rank] - L2 norm across samples
|
||||
|
||||
# Weight projections by activation norm
|
||||
z_weighted = z * activation_norm # [n, null_rank]
|
||||
|
||||
# PCA on weighted projections
|
||||
z_centered = z_weighted - z_weighted.mean(dim=0, keepdim=True)
|
||||
_, S_pca, Vh_pca = torch.linalg.svd(z_centered, full_matrices=False)
|
||||
|
||||
# Top-k directions in tail basis
|
||||
k = min(top_k, Vh_pca.shape[0])
|
||||
U_top = Vh_pca[:k, :] # [k, null_rank]
|
||||
|
||||
# Map back to residual basis: [k, null_rank] @ [null_rank, d_model] -> [k, d_model]
|
||||
V_result = (U_top @ V_tail.T.float()).T # [d_model, k]
|
||||
|
||||
# Orthonormalize
|
||||
V_result, _ = torch.linalg.qr(V_result)
|
||||
V_result = V_result[:, :k].to(dtype).to(device).detach()
|
||||
|
||||
explained_var = (S_pca[:k] ** 2).sum() / ((S_pca ** 2).sum() + 1e-8)
|
||||
act_range = f"{activation_norm.min():.2f}-{activation_norm.max():.2f}"
|
||||
s_range = f"{S_tail.min():.2e}-{S_tail.max():.2e}"
|
||||
logger.info(f"wanda_x_notlogits subspace: rank={V_result.shape[1]}, null_dims={null_rank}, "
|
||||
f"explained_var={explained_var:.1%}, activation_range={act_range}, S_range={s_range}")
|
||||
log_topk_explained_variance(S_pca ** 2, "wanda_x_notlogits") # squared because we used variance formula
|
||||
|
||||
return Subspace(V_result, name="wanda_x_notlogits")
|
||||
|
||||
|
||||
def compute_taskdiff_x_write_x_notlogits_subspace(
|
||||
hidden_states: Float[Tensor, "batch n_layers_plus1 d_model"],
|
||||
write_subspace: "Subspace",
|
||||
lm_head_S: Float[Tensor, "rank"],
|
||||
lm_head_Vh: Float[Tensor, "rank d_model"],
|
||||
top_k: int = 64,
|
||||
layer_frac: float = 0.7,
|
||||
null_frac: float = 0.5,
|
||||
) -> "Subspace":
|
||||
"""Task-discriminative directions in write ∩ lm_head_null.
|
||||
|
||||
Finds directions that are:
|
||||
1. Writable (in column space of o_proj/down_proj)
|
||||
2. Hidden from lm_head (in bottom singular vectors of lm_head)
|
||||
3. Task-discriminative (high cho-rej difference magnitude)
|
||||
|
||||
Unlike `write_x_notlogits` (weight-only), this uses cho-rej activations
|
||||
to find WHICH hidden directions carry task-relevant signal.
|
||||
|
||||
Unlike `logits_tail` (sample-specific), this weights by cho-rej
|
||||
DIFFERENCE, not total activation magnitude.
|
||||
|
||||
Args:
|
||||
hidden_states: [batch, n_layers+1, d_model] from contrastive pairs
|
||||
Assumes batch dimension alternates cho/rej: [cho_0, rej_0, cho_1, rej_1, ...]
|
||||
write_subspace: Subspace of write directions (from compute_write_subspace)
|
||||
lm_head_S: [rank] singular values of lm_head (descending)
|
||||
lm_head_Vh: [rank, d_model] right singular vectors
|
||||
top_k: Number of components to return
|
||||
layer_frac: Which layer to use (fraction of total layers)
|
||||
null_frac: Fraction of bottom singular vectors to use as "null" (default 0.5)
|
||||
|
||||
Returns:
|
||||
Subspace of task-discriminative write-lm_null directions
|
||||
"""
|
||||
device = hidden_states.device
|
||||
dtype = hidden_states.dtype
|
||||
batch, n_layers_plus1, d_model = hidden_states.shape
|
||||
n_layers = n_layers_plus1 - 1
|
||||
|
||||
# Get layer hidden states
|
||||
layer_idx = int(layer_frac * n_layers)
|
||||
hs = hidden_states[:, layer_idx, :].float() # [batch, d]
|
||||
|
||||
# Split cho/rej (assumes alternating)
|
||||
hs_cho = hs[0::2] # [n_pairs, d]
|
||||
hs_rej = hs[1::2] # [n_pairs, d]
|
||||
diff = hs_cho - hs_rej # [n_pairs, d]
|
||||
|
||||
# Step 1: Get lm_head null space (bottom singular vectors = low output gain)
|
||||
rank = lm_head_Vh.shape[0]
|
||||
null_start = int((1 - null_frac) * rank)
|
||||
V_lm_null = lm_head_Vh[null_start:, :].T.to(device).float() # [d, null_rank]
|
||||
|
||||
# Step 2: Intersect with write space
|
||||
V_write = write_subspace.V.to(device).float() # [d, write_rank]
|
||||
V_write_lmnull, _ = approx_intersection_bases(V_write, V_lm_null, top_k=256) # [d, intersect_rank]
|
||||
|
||||
if V_write_lmnull.shape[1] < 2:
|
||||
logger.warning(f"taskdiff_x_write_x_notlogits: write ∩ lm_null intersection too small ({V_write_lmnull.shape[1]}), using write only")
|
||||
V_write_lmnull = V_write
|
||||
|
||||
# Step 3: Project differences into write ∩ lm_null
|
||||
z_diff = diff @ V_write_lmnull # [n_pairs, intersect_rank]
|
||||
|
||||
# Step 4: Weight by task-discriminative magnitude (mean absolute difference)
|
||||
task_weight = z_diff.abs().mean(dim=0) # [intersect_rank]
|
||||
|
||||
# Step 5: PCA on weighted projections to find most task-discriminative directions
|
||||
z_weighted = z_diff * task_weight
|
||||
z_centered = z_weighted - z_weighted.mean(dim=0, keepdim=True)
|
||||
_, S_pca, Vh_pca = torch.linalg.svd(z_centered, full_matrices=False)
|
||||
|
||||
# Top-k directions in intersection basis
|
||||
k = min(top_k, Vh_pca.shape[0])
|
||||
U_top = Vh_pca[:k, :] # [k, intersect_rank]
|
||||
|
||||
# Map back to residual basis: [k, intersect_rank] @ [intersect_rank, d] -> [k, d]
|
||||
V_result = (U_top @ V_write_lmnull.T).T # [d, k]
|
||||
|
||||
# Orthonormalize
|
||||
V_result = orthonormalize(V_result).to(dtype).to(device).detach()
|
||||
|
||||
explained_var = (S_pca[:k] ** 2).sum() / ((S_pca ** 2).sum() + 1e-8)
|
||||
intersect_rank = V_write_lmnull.shape[1]
|
||||
weight_range = f"{task_weight.min():.2f}-{task_weight.max():.2f}"
|
||||
logger.info(f"taskdiff_write_x_notlogits subspace: rank={V_result.shape[1]}, intersect_rank={intersect_rank}, "
|
||||
f"explained_var={explained_var:.1%}, task_weight_range={weight_range}")
|
||||
log_topk_explained_variance(S_pca ** 2, "taskdiff_write_x_notlogits")
|
||||
|
||||
return Subspace(V_result, name="taskdiff_write_x_notlogits")
|
||||
|
||||
|
||||
# ============================================================================
|
||||
# Subspace computation from precomputed SVDs (used by layer_selection.py)
|
||||
# ============================================================================
|
||||
|
||||
def compute_churn_from_hidden_states(
|
||||
hidden_states: Float[Tensor, "batch n_layers_plus1 d_model"],
|
||||
top_k: int =256,
|
||||
) -> Subspace:
|
||||
"""Compute churn subspace: PCA of layer-to-layer changes.
|
||||
|
||||
Churn captures "active computation lanes" - directions where layers
|
||||
add and remove energy during processing.
|
||||
|
||||
Args:
|
||||
hidden_states: [batch, n_layers+1, d_model] - all layer outputs
|
||||
top_k: Number of components to keep
|
||||
|
||||
Returns:
|
||||
Subspace of high-churn directions
|
||||
"""
|
||||
device = hidden_states.device
|
||||
dtype = hidden_states.dtype
|
||||
d_model = hidden_states.shape[-1]
|
||||
|
||||
# Layer diffs: [batch, n_layers, d_model]
|
||||
layer_diffs: Float[Tensor, "batch n_layers d"] = hidden_states[:, 1:, :] - hidden_states[:, :-1, :]
|
||||
layer_diffs_flat: Float[Tensor, "n d"] = layer_diffs.reshape(-1, d_model).float()
|
||||
|
||||
# PCA of layer diffs. normalize_samples=False: layer diffs are already comparable
|
||||
# (same scale within a model), and we want magnitude-weighted to capture where
|
||||
# most computation happens.
|
||||
sub = pca_subspace(
|
||||
layer_diffs_flat,
|
||||
top_k=top_k,
|
||||
normalize_samples=False,
|
||||
name="churn",
|
||||
device=device,
|
||||
dtype=dtype,
|
||||
)
|
||||
return sub
|
||||
|
||||
|
||||
def compute_churn_constructive_from_hidden_states(
|
||||
hidden_states: Float[Tensor, "batch n_layers_plus1 d_model"],
|
||||
top_k: int =256,
|
||||
layer_range: Optional[tuple] = None,
|
||||
) -> Subspace:
|
||||
"""Compute constructive churn: directions where magnitude INCREASES across layers.
|
||||
|
||||
Standard churn is unsigned (PCA of layer diffs). This variant filters to directions
|
||||
where the residual stream is actively BUILDING signal (amplifying), not erasing it.
|
||||
|
||||
Method: For each churn PC, compute whether ||h @ v||^2 increases from early to late layers.
|
||||
Keep only PCs where slope > 0 (magnitude growing).
|
||||
|
||||
Args:
|
||||
hidden_states: [batch, n_layers+1, d_model] - all layer outputs
|
||||
top_k: Number of components to keep
|
||||
layer_range: Optional (start_frac, end_frac) for slope computation (default 0.2-0.8)
|
||||
|
||||
Returns:
|
||||
Subspace of constructive (amplifying) churn directions
|
||||
"""
|
||||
device = hidden_states.device
|
||||
dtype = hidden_states.dtype
|
||||
d_model = hidden_states.shape[-1]
|
||||
n_layers_plus1 = hidden_states.shape[1]
|
||||
n_layers = n_layers_plus1 - 1
|
||||
|
||||
if layer_range is None:
|
||||
layer_range = (0.2, 0.8)
|
||||
|
||||
start_idx = max(1, int(layer_range[0] * n_layers))
|
||||
end_idx = min(n_layers, int(layer_range[1] * n_layers))
|
||||
|
||||
# First compute regular churn PCs
|
||||
layer_diffs: Float[Tensor, "batch n_layers d"] = hidden_states[:, 1:, :] - hidden_states[:, :-1, :]
|
||||
layer_diffs_flat: Float[Tensor, "n d"] = layer_diffs.reshape(-1, d_model).float()
|
||||
layer_diffs_centered = layer_diffs_flat - layer_diffs_flat.mean(dim=0, keepdim=True)
|
||||
_, S, Vh = torch.linalg.svd(layer_diffs_centered, full_matrices=False)
|
||||
|
||||
# Get more PCs than we need to filter
|
||||
n_candidates = min(top_k * 3, Vh.shape[0])
|
||||
V_candidates: Float[Tensor, "d k"] = Vh[:n_candidates, :].T # [d_model, n_candidates]
|
||||
|
||||
# For each PC, compute magnitude trend across layers
|
||||
# Project hidden states onto each PC: [batch, n_layers+1, n_candidates]
|
||||
proj_mag_sq = (hidden_states.float() @ V_candidates) ** 2 # [batch, n_layers+1, n_candidates]
|
||||
|
||||
# Compute slope via early vs late layer magnitude.
|
||||
# We average over a 3-layer window at each endpoint for noise reduction.
|
||||
# The "constructive" signal is (late_mag - early_mag) > 0, meaning
|
||||
# magnitude in this PC direction is INCREASING through the network.
|
||||
# Window size 3 is a tradeoff: smaller = more sensitive but noisier.
|
||||
early_mag = proj_mag_sq[:, start_idx:start_idx+3, :].mean(dim=(0, 1)) # [n_candidates]
|
||||
late_mag = proj_mag_sq[:, end_idx-3:end_idx, :].mean(dim=(0, 1)) # [n_candidates]
|
||||
|
||||
# Constructive = late > early (magnitude increasing)
|
||||
mag_slope = late_mag - early_mag # positive = constructive
|
||||
|
||||
# Select top-k by constructiveness (positive slope), sorted by magnitude
|
||||
constructive_mask = mag_slope > 0
|
||||
if constructive_mask.sum() < top_k:
|
||||
# Fallback: take all with positive slope, fill with least negative
|
||||
logger.warning(f"Only {constructive_mask.sum()} constructive PCs found, taking {top_k} least suppressive")
|
||||
sorted_indices = torch.argsort(mag_slope, descending=True)[:top_k]
|
||||
else:
|
||||
# Among constructive, sort by explained variance (S) and take top-k
|
||||
constructive_indices = torch.where(constructive_mask)[0]
|
||||
# Weight by both constructiveness and variance explained
|
||||
scores = mag_slope[constructive_indices] * S[constructive_indices]
|
||||
sorted_by_score = torch.argsort(scores, descending=True)[:top_k]
|
||||
sorted_indices = constructive_indices[sorted_by_score]
|
||||
|
||||
V_constructive: Float[Tensor, "d k"] = V_candidates[:, sorted_indices].to(dtype).to(device).detach()
|
||||
|
||||
n_positive = (mag_slope[sorted_indices] > 0).sum().item()
|
||||
logger.info(f"Churn_constructive subspace: rank={V_constructive.shape[1]}, {n_positive}/{top_k} strictly constructive")
|
||||
|
||||
return Subspace(V_constructive, name="churn_constructive")
|
||||
|
||||
|
||||
def compute_churn_suppressive_from_hidden_states(
|
||||
hidden_states: Float[Tensor, "batch n_layers_plus1 d_model"],
|
||||
top_k: int =256,
|
||||
layer_range: Optional[tuple] = None,
|
||||
) -> Subspace:
|
||||
"""Compute suppressive churn: directions where magnitude DECREASES across layers.
|
||||
|
||||
Complement to constructive churn. These are directions the model is actively
|
||||
ERASING or damping during processing. Steering these could fight the model's flow.
|
||||
|
||||
Args:
|
||||
hidden_states: [batch, n_layers+1, d_model] - all layer outputs
|
||||
top_k: Number of components to keep
|
||||
layer_range: Optional (start_frac, end_frac) for slope computation (default 0.2-0.8)
|
||||
|
||||
Returns:
|
||||
Subspace of suppressive (erasing) churn directions
|
||||
"""
|
||||
device = hidden_states.device
|
||||
dtype = hidden_states.dtype
|
||||
d_model = hidden_states.shape[-1]
|
||||
n_layers_plus1 = hidden_states.shape[1]
|
||||
n_layers = n_layers_plus1 - 1
|
||||
|
||||
if layer_range is None:
|
||||
layer_range = (0.2, 0.8)
|
||||
|
||||
start_idx = max(1, int(layer_range[0] * n_layers))
|
||||
end_idx = min(n_layers, int(layer_range[1] * n_layers))
|
||||
|
||||
# First compute regular churn PCs
|
||||
layer_diffs: Float[Tensor, "batch n_layers d"] = hidden_states[:, 1:, :] - hidden_states[:, :-1, :]
|
||||
layer_diffs_flat: Float[Tensor, "n d"] = layer_diffs.reshape(-1, d_model).float()
|
||||
layer_diffs_centered = layer_diffs_flat - layer_diffs_flat.mean(dim=0, keepdim=True)
|
||||
_, S, Vh = torch.linalg.svd(layer_diffs_centered, full_matrices=False)
|
||||
|
||||
n_candidates = min(top_k * 3, Vh.shape[0])
|
||||
V_candidates: Float[Tensor, "d k"] = Vh[:n_candidates, :].T
|
||||
|
||||
proj_mag_sq = (hidden_states.float() @ V_candidates) ** 2
|
||||
early_mag = proj_mag_sq[:, start_idx:start_idx+3, :].mean(dim=(0, 1))
|
||||
late_mag = proj_mag_sq[:, end_idx-3:end_idx, :].mean(dim=(0, 1))
|
||||
mag_slope = late_mag - early_mag # negative = suppressive
|
||||
|
||||
# Select top-k by suppressiveness (negative slope)
|
||||
suppressive_mask = mag_slope < 0
|
||||
if suppressive_mask.sum() < top_k:
|
||||
logger.warning(f"Only {suppressive_mask.sum()} suppressive PCs found, taking {top_k} most suppressive")
|
||||
sorted_indices = torch.argsort(mag_slope, descending=False)[:top_k] # Most negative first
|
||||
else:
|
||||
suppressive_indices = torch.where(suppressive_mask)[0]
|
||||
scores = -mag_slope[suppressive_indices] * S[suppressive_indices] # Higher = more suppressive
|
||||
sorted_by_score = torch.argsort(scores, descending=True)[:top_k]
|
||||
sorted_indices = suppressive_indices[sorted_by_score]
|
||||
|
||||
V_suppressive: Float[Tensor, "d k"] = V_candidates[:, sorted_indices].to(dtype).to(device).detach()
|
||||
|
||||
n_negative = (mag_slope[sorted_indices] < 0).sum().item()
|
||||
logger.info(f"Churn_suppressive subspace: rank={V_suppressive.shape[1]}, {n_negative}/{top_k} strictly suppressive")
|
||||
|
||||
return Subspace(V_suppressive, name="churn_suppressive")
|
||||
|
||||
|
||||
def compute_task_diff_constructive_from_hidden_states(
|
||||
hidden_states: Float[Tensor, "batch n_layers_plus1 d_model"],
|
||||
top_k: int =256,
|
||||
layer_range: Optional[tuple] = None,
|
||||
) -> Subspace:
|
||||
"""Compute constructive task_diff: task-discriminative directions being AMPLIFIED.
|
||||
|
||||
Standard task_diff is unsigned PCA of (h_cho - h_rej). This variant filters to
|
||||
directions where the cho/rej separation is INCREASING across layers - i.e., the
|
||||
model is actively building this distinction, not inheriting it from embeddings.
|
||||
|
||||
Method: For each task_diff PC, compute slope of |h_cho @ v| - |h_rej @ v| across layers.
|
||||
Keep only PCs where separation is growing (constructive discrimination).
|
||||
|
||||
Args:
|
||||
hidden_states: [batch, n_layers+1, d_model] with interleaved cho/rej pairs
|
||||
top_k: Number of components
|
||||
layer_range: Optional (start_frac, end_frac) for slope (default 0.3-0.8)
|
||||
|
||||
Returns:
|
||||
Subspace of constructively-discriminating task directions
|
||||
"""
|
||||
device = hidden_states.device
|
||||
dtype = hidden_states.dtype
|
||||
d_model = hidden_states.shape[-1]
|
||||
n_layers_plus1 = hidden_states.shape[1]
|
||||
n_layers = n_layers_plus1 - 1
|
||||
|
||||
if layer_range is None:
|
||||
layer_range = (0.3, 0.8)
|
||||
|
||||
start_idx = max(1, int(layer_range[0] * n_layers))
|
||||
end_idx = min(n_layers, int(layer_range[1] * n_layers))
|
||||
|
||||
# Extract cho and rej (interleaved)
|
||||
hs_cho: Float[Tensor, "n_pairs layers d"] = hidden_states[::2]
|
||||
hs_rej: Float[Tensor, "n_pairs layers d"] = hidden_states[1::2]
|
||||
|
||||
# First compute regular task_diff PCs (on mean diff across layers)
|
||||
task_diffs: Float[Tensor, "n_pairs d"] = (
|
||||
hs_cho[:, start_idx:end_idx+1, :] - hs_rej[:, start_idx:end_idx+1, :]
|
||||
).mean(dim=1).float()
|
||||
|
||||
# Per-sample normalize: each pair votes equally regardless of cho-rej magnitude.
|
||||
# Without this, pairs with large ||cho - rej|| dominate PCA.
|
||||
task_diffs_norm = normalize_rows(task_diffs)
|
||||
task_diffs_centered = task_diffs_norm - task_diffs_norm.mean(dim=0, keepdim=True)
|
||||
_, S, Vh = torch.linalg.svd(task_diffs_centered, full_matrices=False)
|
||||
|
||||
n_candidates = min(top_k * 3, Vh.shape[0])
|
||||
V_candidates: Float[Tensor, "d k"] = Vh[:n_candidates, :].T # [d_model, n_candidates]
|
||||
|
||||
# For each PC, compute magnitude separation trend across layers
|
||||
# |h_cho @ v| - |h_rej @ v| should increase for constructive directions
|
||||
proj_cho = (hs_cho.float() @ V_candidates).abs() # [n_pairs, n_layers+1, n_candidates]
|
||||
proj_rej = (hs_rej.float() @ V_candidates).abs()
|
||||
separation = proj_cho - proj_rej # positive = cho more aligned
|
||||
|
||||
# Compute slope: early vs late separation
|
||||
early_sep = separation[:, start_idx:start_idx+3, :].mean(dim=(0, 1)) # [n_candidates]
|
||||
late_sep = separation[:, end_idx-3:end_idx, :].mean(dim=(0, 1))
|
||||
sep_slope = late_sep - early_sep # positive = constructive (separation growing)
|
||||
|
||||
# Also check that the direction is actually discriminative (|late_sep| > threshold)
|
||||
discriminative = late_sep.abs() > 0.01 # Nonzero separation
|
||||
|
||||
# Select: constructive AND discriminative
|
||||
valid_mask = (sep_slope > 0) & discriminative
|
||||
if valid_mask.sum() < top_k:
|
||||
logger.warning(f"Only {valid_mask.sum()} constructive+discriminative PCs, taking {top_k} best")
|
||||
scores = sep_slope * late_sep.abs() # Favor growing + large separation
|
||||
sorted_indices = torch.argsort(scores, descending=True)[:top_k]
|
||||
else:
|
||||
valid_indices = torch.where(valid_mask)[0]
|
||||
scores = sep_slope[valid_indices] * S[valid_indices]
|
||||
sorted_by_score = torch.argsort(scores, descending=True)[:top_k]
|
||||
sorted_indices = valid_indices[sorted_by_score]
|
||||
|
||||
V_constructive: Float[Tensor, "d k"] = V_candidates[:, sorted_indices].to(dtype).to(device).detach()
|
||||
|
||||
n_valid = ((sep_slope[sorted_indices] > 0) & (late_sep[sorted_indices].abs() > 0.01)).sum().item()
|
||||
logger.info(f"Task_diff_constructive subspace: rank={V_constructive.shape[1]}, {n_valid}/{top_k} constructive+discriminative")
|
||||
|
||||
return Subspace(V_constructive, name="taskdiff_constructive")
|
||||
|
||||
|
||||
def compute_task_diff_from_hidden_states(
|
||||
hidden_states: Float[Tensor, "batch n_layers_plus1 d_model"],
|
||||
top_k: int =256,
|
||||
@@ -1373,88 +795,6 @@ def compute_task_diff_from_hidden_states(
|
||||
return Subspace(sub.V, name="taskdiff", S=sub.S)
|
||||
|
||||
|
||||
def compute_task_read_subspace(
|
||||
task_diff_subspace: Subspace,
|
||||
read_subspace: Subspace,
|
||||
top_k: int = 256,
|
||||
) -> Subspace:
|
||||
"""Compute task_read subspace: task signal readable by transformer blocks.
|
||||
|
||||
task_read = task_diff ∩ read
|
||||
|
||||
These are task-discriminative directions that are read by attention/MLP inputs
|
||||
(q/k/v projections, up/gate projections, etc.).
|
||||
|
||||
Args:
|
||||
task_diff_subspace: Subspace of task differences
|
||||
read_subspace: Subspace readable by residual readers
|
||||
top_k: Number of components
|
||||
|
||||
Returns:
|
||||
Subspace of task signal that read-modules can read
|
||||
"""
|
||||
task_read = approx_intersection(task_diff_subspace, read_subspace)
|
||||
|
||||
if task_read.rank > top_k:
|
||||
task_read = Subspace(task_read.V[:, :top_k], name="taskdiff_read")
|
||||
else:
|
||||
task_read.name = "taskdiff_read"
|
||||
|
||||
return task_read
|
||||
|
||||
|
||||
def compute_task_lm_head_subspace(
|
||||
task_diff_subspace: Subspace,
|
||||
lm_head_subspace: Subspace,
|
||||
top_k: int = 256,
|
||||
) -> Subspace:
|
||||
"""Compute taskdiff_x_logits_read subspace: task signal readable by lm_head.
|
||||
|
||||
taskdiff_x_logits_read = taskdiff ∩ logits_read
|
||||
|
||||
These are task-discriminative directions that lm_head can read,
|
||||
i.e. they affect output logits.
|
||||
"""
|
||||
taskdiff_logits_read = approx_intersection(task_diff_subspace, lm_head_subspace)
|
||||
|
||||
if taskdiff_logits_read.rank > top_k:
|
||||
taskdiff_logits_read = Subspace(taskdiff_logits_read.V[:, :top_k], name="taskdiff_x_logits_read")
|
||||
else:
|
||||
taskdiff_logits_read.name = "taskdiff_x_logits_read"
|
||||
|
||||
return taskdiff_logits_read
|
||||
|
||||
|
||||
def compute_task_wnr_subspace(
|
||||
task_diff_subspace: Subspace,
|
||||
write_not_read_subspace: Subspace,
|
||||
top_k: int =256,
|
||||
) -> Subspace:
|
||||
"""Compute task_wnr subspace: task signal written but not read.
|
||||
|
||||
task_wnr = task_diff ∩ write_not_read
|
||||
|
||||
These are task-discriminative directions that are written to residual
|
||||
but not read by later layers or lm_head.
|
||||
|
||||
Args:
|
||||
task_diff_subspace: Subspace of task differences
|
||||
write_not_read_subspace: Subspace of write-not-read directions
|
||||
top_k: Number of components
|
||||
|
||||
Returns:
|
||||
Subspace of task signal that's written but ignored
|
||||
"""
|
||||
taskdiff_write_not_read = approx_intersection(task_diff_subspace, write_not_read_subspace)
|
||||
|
||||
if taskdiff_write_not_read.rank > top_k:
|
||||
taskdiff_write_not_read = Subspace(taskdiff_write_not_read.V[:, :top_k], name="taskdiff_x_write_not_read")
|
||||
else:
|
||||
taskdiff_write_not_read.name = "taskdiff_x_write_not_read"
|
||||
|
||||
return taskdiff_write_not_read
|
||||
|
||||
|
||||
def compute_module_subspace_from_svds(
|
||||
layer_svds: Dict[str, tuple],
|
||||
layer_info: Dict[str, dict],
|
||||
|
||||
Reference in New Issue
Block a user