mirror of
https://github.com/wassname/AntiPaSTO.git
synced 2026-09-09 11:12:52 +08:00
restore some alt subspaces
This commit is contained in:
@@ -1,15 +1,18 @@
|
||||
#!/usr/bin/env python3
|
||||
"""Centralized layer selection logic for AntiPaSTO training.
|
||||
|
||||
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.
|
||||
|
||||
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
|
||||
- Computes subspaces: write, read, taskdiff, and combinations
|
||||
|
||||
Supported loss_subspace types:
|
||||
- taskdiff_x_suppressed_x_write (default): empirical stenographic + writable
|
||||
- taskdiff_x_logits_read: task signal that affects lm_head output
|
||||
- taskdiff_x_write_not_read: task signal in static write-not-read space
|
||||
- taskdiff_x_write_x_notlogits: task ∩ write ∩ (lm_head^⊥)
|
||||
- write, taskdiff: basic building blocks
|
||||
|
||||
Key functions:
|
||||
- compute_simple_layer_selection(): Main entry point for layer/subspace selection
|
||||
@@ -36,10 +39,15 @@ from transformers import DataCollatorWithPadding
|
||||
import gc
|
||||
from antipasto.peft_utils.subspaces import (
|
||||
compute_lm_head_svd,
|
||||
compute_lm_head_subspace,
|
||||
compute_suppressed_from_hidden_states,
|
||||
compute_task_diff_from_hidden_states,
|
||||
compute_module_subspace_from_svds,
|
||||
compute_stenographic_subspace,
|
||||
compute_write_not_read_subspace,
|
||||
compute_write_x_notlogits_subspace,
|
||||
compute_task_lm_head_subspace,
|
||||
compute_task_wnr_subspace,
|
||||
find_write_modules,
|
||||
find_read_modules,
|
||||
approx_intersection,
|
||||
@@ -965,7 +973,55 @@ def compute_simple_layer_selection(
|
||||
steno_intersect_write = approx_intersection(steno_subspace, write_subspace, top_k=INTERMEDIATE_SUBSPACE_RANK)
|
||||
subspaces.set('taskdiff_x_suppressed_x_write', steno_intersect_write)
|
||||
|
||||
logger.info(f"Activation-based subspaces computed: taskdiff, suppressed, taskdiff_x_suppressed, taskdiff_x_suppressed_x_write")
|
||||
# =====================================================================
|
||||
# STATIC (model-intrinsic) subspaces - alternative to empirical suppressed
|
||||
# =====================================================================
|
||||
|
||||
# taskdiff_x_logits_read = task signal that AFFECTS output (opposite of suppressed)
|
||||
taskdiff_logits_read = 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', taskdiff_logits_read)
|
||||
|
||||
# Static write-not-read and its task intersection
|
||||
if write_subspace is not None and read_subspace is not None:
|
||||
# write_not_read = static: write^⊥_read^⊥_lmhead
|
||||
write_not_read = compute_write_not_read_subspace(
|
||||
write_subspace=write_subspace,
|
||||
read_subspace=read_subspace,
|
||||
lm_head_subspace=lm_head_sub,
|
||||
top_k=INTERMEDIATE_SUBSPACE_RANK,
|
||||
)
|
||||
subspaces.set('write_not_read', write_not_read)
|
||||
|
||||
# taskdiff_x_write_not_read = task in static hidden channels
|
||||
taskdiff_wnr = compute_task_wnr_subspace(
|
||||
task_diff_subspace=task_diff_subspace,
|
||||
write_not_read_subspace=write_not_read,
|
||||
top_k=INTERMEDIATE_SUBSPACE_RANK,
|
||||
)
|
||||
subspaces.set('taskdiff_x_write_not_read', taskdiff_wnr)
|
||||
|
||||
# write_x_notlogits = static: write ∩ (lm_head^⊥)
|
||||
if write_subspace is not None:
|
||||
write_x_notlogits = compute_write_x_notlogits_subspace(
|
||||
write_subspace=write_subspace,
|
||||
lm_head_subspace=lm_head_sub,
|
||||
top_k=INTERMEDIATE_SUBSPACE_RANK,
|
||||
)
|
||||
subspaces.set('write_x_notlogits', write_x_notlogits)
|
||||
|
||||
# taskdiff_x_write_x_notlogits = task ∩ write ∩ (lm_head^⊥)
|
||||
taskdiff_write_notlogits = approx_intersection(
|
||||
task_diff_subspace, write_x_notlogits, top_k=INTERMEDIATE_SUBSPACE_RANK
|
||||
)
|
||||
taskdiff_write_notlogits.name = 'taskdiff_x_write_x_notlogits'
|
||||
subspaces.set('taskdiff_x_write_x_notlogits', taskdiff_write_notlogits)
|
||||
|
||||
computed_subs = list(subspaces._subspaces.keys())
|
||||
logger.info(f"Subspaces computed: {computed_subs}")
|
||||
|
||||
# Cleanup hidden states if collected
|
||||
if hs_stacked is not None:
|
||||
|
||||
@@ -19,18 +19,15 @@ For geometric intuition and taxonomy of named subspaces, see docs/steering_metho
|
||||
|
||||
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
|
||||
Supported subspace types:
|
||||
- taskdiff_x_suppressed_x_write (default) - empirical stenographic + writable
|
||||
- taskdiff_x_logits_read - task signal that affects lm_head output
|
||||
- taskdiff_x_write_not_read - task signal in static write-not-read space
|
||||
- write_x_notlogits - static: write ∩ (lm_head^⊥)
|
||||
- write_not_read - static: write^⊥_read (optionally also ^⊥_lmhead)
|
||||
- write, taskdiff - basic building blocks
|
||||
|
||||
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
|
||||
See docs/steering_methods.qmd for geometric intuition.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
@@ -676,6 +673,164 @@ def compute_lm_head_svd(model: nn.Module) -> tuple[Tensor, Tensor]:
|
||||
return S, Vh
|
||||
|
||||
|
||||
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 (logits_read)
|
||||
"""
|
||||
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_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).
|
||||
|
||||
This is STATIC (model-intrinsic, not task-specific). Intersect with taskdiff
|
||||
to get task-relevant write-not-read directions.
|
||||
|
||||
Args:
|
||||
write_subspace: Subspace of write directions
|
||||
read_subspace: Subspace of read directions
|
||||
lm_head_subspace: Optional subspace readable by lm_head
|
||||
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_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).
|
||||
|
||||
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.
|
||||
|
||||
This is STATIC (model-intrinsic). Intersect with taskdiff for task relevance.
|
||||
|
||||
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_task_lm_head_subspace(
|
||||
task_diff_subspace: Subspace,
|
||||
lm_head_subspace: Subspace,
|
||||
top_k: int = 256,
|
||||
) -> Subspace:
|
||||
"""Compute taskdiff_x_logits_read: 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 WILL affect output logits. Opposite of suppressed/hidden.
|
||||
|
||||
Args:
|
||||
task_diff_subspace: Subspace of task differences
|
||||
lm_head_subspace: Subspace readable by lm_head
|
||||
top_k: Number of components
|
||||
|
||||
Returns:
|
||||
Subspace of task signal that affects output
|
||||
"""
|
||||
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 taskdiff_x_write_not_read: task signal in static WnR space.
|
||||
|
||||
taskdiff_x_write_not_read = taskdiff ∩ write_not_read
|
||||
|
||||
These are task-discriminative directions that are written to residual
|
||||
but not read by later layers or lm_head. Uses STATIC write_not_read
|
||||
(model-intrinsic) rather than empirical suppressed.
|
||||
|
||||
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 in static hidden channels
|
||||
"""
|
||||
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_stenographic_subspace(
|
||||
task_diff_subspace: Subspace,
|
||||
suppressed_subspace: Subspace,
|
||||
|
||||
Reference in New Issue
Block a user