-
Notifications
You must be signed in to change notification settings - Fork 3
Expand file tree
/
Copy pathsteering_method.py
More file actions
101 lines (78 loc) · 3.82 KB
/
Copy pathsteering_method.py
File metadata and controls
101 lines (78 loc) · 3.82 KB
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
"""
Abstract base class for steering methods in the Cross-Value Transfer experiment.
Each concrete subclass encapsulates:
- How pre-computed steering vectors are loaded from disk.
- How a forward hook is installed / removed for a given steering vector.
This interface is intentionally minimal so that methods such as plain CAA,
SAE-based steering (QwenScopeCAA), spherical steering, etc. can all be
plugged in by implementing three abstract methods.
Note: This ABC is experiment-scoped and is separate from the internal
CAA/Geometry/steering/base.py which is tightly coupled to the
activation-extraction pipeline.
"""
from __future__ import annotations
from abc import ABC, abstractmethod
from typing import Any, Dict
import torch
class SteeringMethod(ABC):
"""Abstract steering method for the cross-value transfer experiment.
Subclass contract
-----------------
1. ``name`` — unique identifier used for output directory naming
(e.g. ``"caa"``, ``"qwenscope_caa"``).
2. ``layer`` — the model layer index at which this method applies steering.
Implementations may resolve this lazily (e.g. by reading
``selected_layer.json`` from the run directory).
3. ``load_vectors()`` — load and return {value_name: steering_tensor} for
all 20 Schwartz values. Vectors should be returned on CPU; the runner
will move them to the appropriate device before calling ``apply_hook``.
Implementations should normalise vectors to unit length here.
4. ``apply_hook(model_info, vector, alpha)`` — register a forward hook that
injects ``alpha * vector`` (or the method's equivalent) into the
residual stream at the layer returned by ``self.layer``. Returns an
opaque handle that will be passed to ``remove_hook``.
5. ``remove_hook(handle)`` — remove the hook registered by ``apply_hook``.
"""
# ── identity ─────────────────────────────────────────────────────────────
@property
@abstractmethod
def name(self) -> str:
"""Unique string identifier for this method (used in output paths)."""
@property
@abstractmethod
def layer(self) -> int:
"""Model layer index at which steering is applied."""
# ── vectors ──────────────────────────────────────────────────────────────
@abstractmethod
def load_vectors(self) -> Dict[str, torch.Tensor]:
"""Load per-value steering vectors from disk.
Returns
-------
dict mapping each of the 20 Schwartz value names to a 1-D CPU tensor
of shape ``(hidden_dim,)``. Vectors should be unit-normalised.
"""
# ── hooks ────────────────────────────────────────────────────────────────
@abstractmethod
def apply_hook(
self,
model_info: Any,
vector: torch.Tensor,
alpha: float,
) -> Any:
"""Register a steering hook on the model.
Parameters
----------
model_info:
``ModelInfo`` object from ``CAA/Geometry/model_loader.py``.
vector:
Unit-normalised steering vector (CPU tensor; implementations
should move to ``model_info.device`` internally).
alpha:
Steering strength multiplier.
Returns
-------
An opaque handle (or list of handles) to be passed to ``remove_hook``.
"""
@abstractmethod
def remove_hook(self, handle: Any) -> None:
"""Remove the hook(s) returned by ``apply_hook``."""