mirror of
https://github.com/wassname/AntiPaSTO.git
synced 2026-09-11 11:52:42 +08:00
restore alt subspaces
This commit is contained in:
+11
-2
@@ -149,7 +149,11 @@ class TrainingConfig:
|
||||
|
||||
loss_subspace: Literal[
|
||||
# Recommended (default)
|
||||
"taskdiff_x_suppressed_x_write", # Task-discriminative ∩ suppressed ∩ write
|
||||
"taskdiff_x_suppressed_x_write", # Task-discriminative ∩ suppressed ∩ write (empirical)
|
||||
# Static alternatives (model-intrinsic, not task-specific base)
|
||||
"taskdiff_x_logits_read", # Task signal that AFFECTS output (opposite of suppressed)
|
||||
"taskdiff_x_write_not_read", # Task signal in static write-not-read space
|
||||
"taskdiff_x_write_x_notlogits", # Task ∩ write ∩ (lm_head^⊥)
|
||||
# Simpler alternatives
|
||||
"write", # Write space only (o_proj, down_proj column space)
|
||||
"taskdiff", # Task-discriminative PCA only
|
||||
@@ -158,8 +162,13 @@ class TrainingConfig:
|
||||
|
||||
Default: taskdiff_x_suppressed_x_write = intersection of:
|
||||
- taskdiff: PCA on cho-rej difference (task-discriminative directions)
|
||||
- suppressed: Written to residual but erased by later layers
|
||||
- suppressed: Written to residual but erased by later layers (EMPIRICAL)
|
||||
- write: Column space of o_proj and down_proj (writable by model)
|
||||
|
||||
Alternative static (model-intrinsic) options:
|
||||
- taskdiff_x_logits_read: Task signal readable by lm_head (affects output)
|
||||
- taskdiff_x_write_not_read: Task signal in static write-not-read space
|
||||
- taskdiff_x_write_x_notlogits: Task ∩ write ∩ (lm_head^⊥)
|
||||
"""
|
||||
|
||||
loss_subspace_rank: Optional[int] = 8
|
||||
|
||||
@@ -339,11 +339,6 @@ class AntiPaSTOLayer(BaseTunerLayer):
|
||||
"""
|
||||
# Soft clamp rotation angle: small θ ensures R(θ)@S ≈ -R(-θ)@S (first-order approx)
|
||||
# This gives additive output symmetry: Δy(+1) ≈ -Δy(-1) around base model
|
||||
|
||||
# if max_angle is not None and max_angle < float('inf'):
|
||||
# A_clamped = max_angle * torch.tanh(A / max_angle)
|
||||
# else:
|
||||
# A_clamped = A
|
||||
|
||||
if max_angle is not None and max_angle < (torch.pi - 1e-6):
|
||||
# Convert desired max rotation angle to A-space limit
|
||||
@@ -446,7 +441,7 @@ class AntiPaSTOLayer(BaseTunerLayer):
|
||||
if not torch.isfinite(U_rot).all():
|
||||
raise ValueError(f"NaNs in U_rot for adapter {adapter}. alpha={alpha}, max_angle={max_angle}")
|
||||
if not torch.isfinite(S_scaled).all():
|
||||
raise ValueError(f"NaNs in S_scaled for adapter {adapter}. scale_mode={scale_mode}")
|
||||
raise ValueError(f"NaNs in S_scaled for adapter {adapter}.")
|
||||
|
||||
# Efficient forward: x @ V_rot @ diag(S_scaled) @ U_rot^T
|
||||
x_projected = x @ V_rot # [..., r]
|
||||
|
||||
@@ -145,19 +145,24 @@ def compute_loss_subspace_basis(
|
||||
Returns:
|
||||
L_subspace: [d_model, top_k] tensor
|
||||
|
||||
Note: Many loss_subspace options were removed in Jan 2026 cleanup.
|
||||
Only taskdiff_x_suppressed_x_write, write, taskdiff are supported.
|
||||
See git history for steer*, null, notlogits, weight_svd implementations.
|
||||
Supported loss_subspace options:
|
||||
- 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
|
||||
"""
|
||||
top_k = config.loss_subspace_rank # None means auto by energy
|
||||
|
||||
# Get from cache - SubspaceCache now stores Subspace objects
|
||||
L_cached_sub = subspaces.get(config.loss_subspace)
|
||||
if L_cached_sub is None:
|
||||
available = list(subspaces._subspaces.keys())
|
||||
raise ValueError(
|
||||
f"Subspace '{config.loss_subspace}' not found in cache. "
|
||||
f"Available: {list(subspaces._subspaces.keys())}. "
|
||||
f"Valid options: taskdiff_x_suppressed_x_write, write, taskdiff"
|
||||
f"Available: {available}. "
|
||||
f"Valid options: taskdiff_x_suppressed_x_write, taskdiff_x_logits_read, "
|
||||
f"taskdiff_x_write_not_read, taskdiff_x_write_x_notlogits, write, taskdiff"
|
||||
)
|
||||
|
||||
L_cached = L_cached_sub.V
|
||||
|
||||
Reference in New Issue
Block a user