hallucination / extra_materials /graph /Feature_Graph_Trans.py
ToiTenBao's picture
Upload hallucination folder
a2ffd07 verified
Raw
History Blame Contribute Delete
112 kB
from tracemalloc import start
from typing import Literal
from .Feature_Graph import *
from transformer_lens.utils import get_act_name
def add_cache(all_cache: Dict[str, float | Tensor], cache: Dict[str, Tensor]):
for key, value in cache.items():
if key not in all_cache:
all_cache[key] = 0
all_cache[key] += value
def _hook_name(sae_name: str, name: str | None):
return sae_name + "." + str(name)
@t.no_grad()
def get_ln_constant(
cache: Dict[str, Tensor] | ActivationCache,
vector: Tensor, # (1, 1, d_model) | (b, 1, d_model) | (b, seq, d_model)
layer: int,
pos: int | None,
ln: Literal['ln1', 'ln2'] = 'ln2',
eps: float = 1e-6,
) -> Tensor:
assert ln in ['ln1', 'ln2'], f"ln must be either 'ln1' or 'ln2', got {ln}"
x_act_name = get_act_name('resid_mid', layer) if ln == "ln2" else get_act_name('resid_pre', layer)
y_act_name = get_act_name('normalized', layer, ln)
if pos is not None:
x = cache[x_act_name][:, pos].unsqueeze(1) # (b, 1, d_model)
y = cache[y_act_name][:, pos].unsqueeze(1) # (b, 1, d_model)
if vector.squeeze(1).ndim == 3:
vector = vector[:, pos].unsqueeze(1) # (b, 1, d_model)
denom = t.sum(x * vector, dim=-1, keepdim=True) # (b, 1, 1)
numer = t.sum(y * vector, dim=-1, keepdim=True) # (b, 1, 1)
else:
x = cache[x_act_name] # (b, seq, d_model)
y = cache[y_act_name] # (b, seq, d_model)
denom = t.sum(x * vector, dim=-1, keepdim=True) # (b, seq, 1)
numer = t.sum(y * vector, dim=-1, keepdim=True) # (b, seq, 1)
# Handle zero division
mask = t.abs(denom) < eps
safe_denom = t.where(mask, t.ones_like(denom), denom)
result = numer / safe_denom
# Set result to 0 where denominator is zero
return t.where(mask, t.zeros_like(result), result) * vector # (b, 1 or seq, d_model)
@t.no_grad()
def get_attn_head_contribs(
vector: Tensor, # (1, 1, d_model) | (b, 1, d_model) | (b, seq, d_model)
model: HookedSAETransformer,
cache: Dict[str, Tensor] | ActivationCache,
layer: int,
dst_pos: int | None,
sum_head: bool = True,
):
attn_pattern = cache[get_act_name('pattern', layer)] # (b, h, dst, src)
batch_size = attn_pattern.shape[0]
if vector.shape[0] != batch_size: # upsample to batchsize
vector = vector.repeat((batch_size, 1, 1)) # (b, 1 or seq, d_model)
if dst_pos is not None:
attn_pattern = attn_pattern[:, :, dst_pos] # (b, h, src)
if vector.squeeze(1).ndim == 3:
vector = vector[:, dst_pos].unsqueeze(1) # (b, 1, d_model2)
grad_outs = einops.einsum(
model.W_O[layer], vector.squeeze(1),
'h d_head d_model2, b d_model2 -> b h d_head',
)
grad_vals = einops.einsum(
attn_pattern, grad_outs,
'b h src, b h d_head -> b h src d_head',
)
else:
grad_outs = einops.einsum(
model.W_O[layer], vector,
'h d_head d_model2, b dst d_model2 -> b h dst d_head',
)
grad_vals = einops.einsum(
attn_pattern, grad_outs,
'b h dst src, b h dst d_head -> b h src d_head',
)
grad = einops.einsum(
model.W_V[layer], grad_vals,
'h d_model1 d_head, b h src d_head -> b h src d_model1'
)
if sum_head:
grad = t.sum(grad, dim=1) # (b, seq, d_model)
return grad # (b, seq, d_model) | (b, h. seq, d_model)
def _check_shape(vec: Tensor, pos: int | None):
if vec.squeeze(1).ndim != 3:
assert pos is not None
else:
assert pos is None
@t.no_grad()
def gradient_ln_only(
end_feature_vec: Tensor, # (1, 1, d_model) | (b, 1, d_model) | (b, seq, d_model)
start_feature_vec: Tensor | None, # (d_sae, d_model)
layer_end: int,
pos_end: int | None,
seq_length: int,
batch_size: int,
use_error_term: bool,
cache: Dict[str, Tensor] | ActivationCache,
device: str | t.device,
ln: Literal["ln1", "ln2"] | None = None,
) -> SparseAct:
d_model = end_feature_vec.shape[-1]
_check_shape(end_feature_vec, pos_end)
if ln is not None:
feature_scale = get_ln_constant( # (b, 1 or seq, d_model)
cache,
end_feature_vec,
layer_end,
pos_end,
ln=ln,
)
else:
feature_scale = end_feature_vec
if pos_end is not None:
grad = t.zeros((batch_size, seq_length, d_model), device=device) # (b, seq, d_model)
grad[:, pos_end] = feature_scale.squeeze(1) # (b, d_model)
else:
grad = feature_scale # (b, seq, d_model)
if start_feature_vec is not None:
return SparseAct(
act=grad @ start_feature_vec.T, # (b, seq, d_sae)
res=grad if use_error_term else None, # (b, seq, d_model)
resc=None,
)
else:
return SparseAct(
act=grad, # (b, seq, d_model)
res=grad if use_error_term else None, # (b, seq, d_model)
resc=None,
)
@t.no_grad()
def gradient_with_mlp(
end_feature_vec: Tensor, # (1, 1, d_model) | (b, 1, d_model) | (b, seq, d_model)
start_feature_vec: Tensor | None, # (d_sae, d_model)
transcoder_enc: Tensor, # (d_model, d_sae)
transcoder_dec: Tensor, # (d_sae, d_model)
transcoder_act: Tensor, # (b, seq, d_sae)
layer_end: int,
pos_end: int | None,
seq_length: int,
batch_size: int,
use_error_term: bool,
cache: Dict[str, Tensor] | ActivationCache,
device: str | t.device,
) -> SparseAct:
d_model = end_feature_vec.shape[-1]
_check_shape(end_feature_vec, pos_end)
if pos_end is not None:
# [(1 or b, 1, d_model) @ (d_model, d_sae)] * (b, 1, d_sae) -> (b, 1, d_sae)
grad_post = (end_feature_vec @ transcoder_dec.T) * transcoder_act[:, pos_end].unsqueeze(1)
grad_enc = grad_post @ transcoder_enc.T # (b, 1, d_model)
feature_scale = get_ln_constant( # (b, 1, d_model)
cache,
grad_enc,
layer_end,
pos_end,
ln="ln2",
)
grad = t.zeros((batch_size, seq_length, d_model), device=device) # (b, seq, d_model)
grad[:, pos_end] = feature_scale.squeeze(1) # (b, d_model)
else:
# [(b, seq, d_model) @ (d_model, d_sae)] * (b, seq, d_sae) -> (b, seq, d_sae)
grad_post = (end_feature_vec @ transcoder_dec.T) * transcoder_act
grad_enc = grad_post @ transcoder_enc.T # (b, seq, d_model)
grad = get_ln_constant( # (b, seq, d_model)
cache,
grad_enc,
layer_end,
pos_end,
ln="ln2",
)
if start_feature_vec is not None:
return SparseAct(
act=grad @ start_feature_vec.T, # (b, seq, d_sae)
res=grad if use_error_term else None, # (b, seq, d_model)
resc=None,
)
else:
return SparseAct(
act=grad, # (b, seq, d_model)
res=grad if use_error_term else None, # (b, seq, d_model)
resc=None,
)
@t.no_grad()
def gradient_with_attn(
model: HookedSAETransformer,
end_feature_vec: Tensor, # (1, 1, d_model) | (b, 1, d_model) | (b, seq, d_model)
start_feature_vec: Tensor | None, # (d_sae, d_model)
layer_end: int,
pos_end: int | None,
use_error_term: bool,
cache: Dict[str, Tensor] | ActivationCache,
device: str | t.device,
) -> SparseAct:
_check_shape(end_feature_vec, pos_end)
grad_through_attn = get_attn_head_contribs(end_feature_vec, model, cache, layer_end, pos_end) # (b, seq, d_model)
grad = get_ln_constant( # (b, seq, d_model)
cache,
grad_through_attn,
layer_end,
None,
ln="ln1",
)
if start_feature_vec is not None:
return SparseAct(
act=grad @ start_feature_vec.T, # (b, seq, d_sae)
res=grad if use_error_term else None, # (b, seq, d_model)
resc=None,
)
else:
return SparseAct(
act=grad, # (b, seq, d_model)
res=grad if use_error_term else None, # (b, seq, d_model)
resc=None,
)
class Feature_Graph_Trans(Feature_Graph):
def __init__(
self,
model: HookedSAETransformer,
saes: Dict[int, List[Tuple[str, Any]]], # {layer: list[{hook_position: HookedSAE}]}, can define granularity here
use_error_term: bool = False,
):
super().__init__(model, saes, use_error_term)
self.process_transcoder()
def process_transcoder(self):
self.input_hooks = []
self.output_hooks = []
self.transcoders = []
self.non_transcoders = []
check_mlp_out = False
self.check_attn_out = False
self.check_resid_pre = False
for sae in self.dict_saes.values():
sae = sae.to(self.device)
if getattr(sae, 'input_hook', False) and getattr(sae, 'output_hook', False):
self.input_hooks.append(sae.input_hook)
self.output_hooks.append(sae.output_hook)
self.transcoders.append(sae)
if "mlp_out" in sae.output_hook:
check_mlp_out = True
else:
self.non_transcoders.append(sae)
if "attn_out" in sae.cfg.hook_name:
self.check_attn_out = True
if "resid_pre" in sae.cfg.hook_name:
self.check_resid_pre = True
assert check_mlp_out, "Transcoder needs to be provided at the mlp_out hook."
def forward(
self,
clean_token: Tensor,
corrupt_cache: ActivationCache | Dict[str, Tensor] | None,
patch_deleted_comp: bool = False,
**kwargs,
) -> Tuple[Tensor, Dict[str, SparseAct]]:
'''
Forward pass of the graph with clean tokens, if the edge exists, replace the activation with corrupted activation
'''
self._check_graph()
self.model.reset_hooks()
self.model_setup()
fwd_cache = {}
with t.no_grad():
with self._setup_forward_model_hook(transfer_grad=False):
with self._setup_fwd_sae_hook(
fwd_cache=fwd_cache, corrupt_cache=corrupt_cache, patch_deleted_comp=patch_deleted_comp
):
logits = self.model(clean_token)
cache = {}
for sae_name in self.dict_saes.keys():
cache[sae_name] = cache_to_sparseact(
fwd_cache,
sae_hook_name(sae_name),
error_term_name(sae_name) if self.use_error_term else None,
)
for sae in self.dict_saes.values():
sae.reset_hooks()
self.model.reset_hooks()
return logits, cache
def forward_backward_gradient(
self,
clean_token: Tensor,
corrupt_cache: ActivationCache | Dict[str, Tensor],
metric: Callable[[Tensor], Tensor],
retain_graph: bool = False,
mode: str = 'node',
gradient_mode: str = 'standard',
pass_through_grad: bool = False,
verbose: bool = False,
**kwargs,
) -> Tuple[
Dict[Tuple[Node, Index], SparseAct], # node effects
Dict[Tuple[Node, Index, Node, Index], Tensor], # edge effects
]:
if mode == 'node':
if verbose:
print("Calculating node gradients...")
if gradient_mode == "standard":
node_grads, clean_cache = self._gradient_wrt_nodes(
clean_token, metric, retain_graph, pass_through_grad, **kwargs
)
elif gradient_mode == "ig":
node_grads, clean_cache = self._gradient_wrt_nodes_ig(
clean_token, corrupt_cache, metric, retain_graph, verbose, **kwargs
)
elif gradient_mode == "virtual_weight":
node_grads, clean_cache = self._gradient_wrt_nodes_vw(
clean_token, metric, retain_graph, pass_through_grad, **kwargs
)
elif gradient_mode == "virtual_weight_ig":
node_grads, clean_cache = self._gradient_wrt_nodes_vw_ig(
clean_token, corrupt_cache, metric, retain_graph, **kwargs
)
else:
raise NotImplementedError(f"gradient_mode {gradient_mode} is not supported")
node_effect = self._attrib_effect_node(node_grads, corrupt_cache, clean_cache)
return node_effect, {}
elif mode == 'edge':
if kwargs.get('node_grads', None) is None or kwargs.get('node_effect', None) is None:
# run the _gradient_wrt_nodes to get node_grads, and use node_effect to prune out unimportant nodes
if verbose:
print("Calculating node gradients...")
if gradient_mode == "standard":
node_grads, clean_cache = self._gradient_wrt_nodes(
clean_token, metric, retain_graph, pass_through_grad, **kwargs
)
elif gradient_mode == "ig":
node_grads, clean_cache = self._gradient_wrt_nodes_ig(
clean_token, corrupt_cache, metric, retain_graph, verbose, **kwargs
)
elif gradient_mode == "virtual_weight":
node_grads, clean_cache = self._gradient_wrt_nodes_vw(
clean_token, metric, retain_graph, pass_through_grad, **kwargs
)
elif gradient_mode == "virtual_weight_ig":
node_grads, clean_cache = self._gradient_wrt_nodes_vw_ig(
clean_token, corrupt_cache, metric, retain_graph, **kwargs
)
else:
raise NotImplementedError(f"gradient_mode {gradient_mode} is not supported")
node_effect = self._attrib_effect_node(node_grads, corrupt_cache, clean_cache)
else:
node_grads: Dict[Tuple[Node, Index], SparseAct] = kwargs.get('node_grads') # type: ignore
node_effect: Dict[Tuple[Node, Index], SparseAct] = kwargs.get('node_effect') # type: ignore
# Pruning
if kwargs.get('prune', False):
if verbose:
print("Pruning nodes...")
self._prune_nodes(node_effect, verbose, **kwargs)
if kwargs.get('gradient_only', False):
if verbose:
print("Returning edge gradients only...")
for name, sparse_act in node_grads.items():
node_grads[name] = sparse_act.to_sparse_like_self(t.ones_like(sparse_act.to_tensor()))
del sparse_act
edge_grads, _ = self._gradient_wrt_edges(
clean_token, corrupt_cache, node_grads, verbose, **kwargs
)
return node_effect, edge_grads
else:
raise NotImplementedError(f"mode {mode} is not supported")
def _gradient_wrt_nodes(
self,
clean_token: Tensor,
metric: Callable[[Tensor], Tensor],
retain_graph: bool = False,
pass_through_grad: bool = False,
verbose: bool = False,
**kwargs,
) -> Tuple[
Dict[Tuple[Node, Index], SparseAct], # node effects
Dict[str, SparseAct]
]:
self._check_graph()
self.model_setup()
self.model.reset_hooks()
for _, sae in self.dict_saes.items():
sae.reset_hooks()
fwd_cache = {}
bwd_cache = {}
with t.set_grad_enabled(True):
with self._detach_error_term(True):
with self._setup_forward_model_hook(transfer_grad=kwargs.get("transfer_grad", True)):
with self._setup_fwd_bwd_grad_sae_hook(
fwd_cache=fwd_cache,
bwd_cache=bwd_cache,
pass_through_grad=pass_through_grad,
):
metric(self.model(clean_token)).backward(retain_graph=retain_graph)
node_grads = {}
for node, index in self.nodes.keys():
node_grads[(node, index)] = cache_to_sparseact(
bwd_cache,
sae_hook_name(node.name),
error_term_name(node.name) if self.use_error_term else None,
)
cache = {}
for sae_name in self.dict_saes.keys():
cache[sae_name] = cache_to_sparseact(
fwd_cache,
sae_hook_name(sae_name),
error_term_name(sae_name) if self.use_error_term else None,
)
self.model.reset_hooks()
for sae in self.dict_saes.values():
sae.reset_hooks()
return node_grads, cache
def _gradient_wrt_nodes_ig(
self,
clean_token: Tensor,
corrupt_cache: ActivationCache | Dict[str, Tensor],
metric: Callable[[Tensor], Tensor],
retain_graph: bool = False,
verbose: bool = False,
**kwargs,
) -> Tuple[
Dict[Tuple[Node, Index], SparseAct], # node effects
Dict[str, SparseAct]
]:
steps = kwargs.get("steps", 10)
self._check_graph()
self.model_setup()
self.model.reset_hooks()
for _, sae in self.dict_saes.items():
sae.reset_hooks()
fwd_cache = {}
bwd_cache = {}
with t.set_grad_enabled(True):
with self._detach_error_term(True):
with self._setup_forward_model_hook(transfer_grad=kwargs.get("transfer_grad", True)):
for target_name in self.dict_saes.keys():
for step in range(steps):
frac = step / steps
with self._setup_fwd_bwd_grad_sae_hook_ig(
target_name=target_name,
frac=frac,
fwd_cache=fwd_cache,
bwd_cache=bwd_cache,
corrupt_cache=corrupt_cache,
):
metric(self.model(clean_token)).backward(retain_graph=retain_graph)
# average the gradients
for key in bwd_cache.keys():
bwd_cache[key] /= steps
node_grads = {}
for node, index in self.nodes.keys():
node_grads[(node, index)] = cache_to_sparseact(
bwd_cache,
sae_hook_name(node.name),
error_term_name(node.name) if self.use_error_term else None,
)
cache = {}
for sae_name in self.dict_saes.keys():
cache[sae_name] = cache_to_sparseact(
fwd_cache,
sae_hook_name(sae_name),
error_term_name(sae_name) if self.use_error_term else None,
)
self.model.reset_hooks()
for sae in self.dict_saes.values():
sae.reset_hooks()
return node_grads, cache
def _gradient_wrt_nodes_vw(
self,
clean_token: Tensor,
metric: Callable[[Tensor], Tensor],
retain_graph: bool = False,
pass_through_grad: bool = False,
verbose: bool = False,
**kwargs,
) -> Tuple[
Dict[Tuple[Node, Index], SparseAct], # node effects
Dict[str, SparseAct]
]:
'''
Using virtual weight to compute node grad
'''
self._check_graph()
self.model_setup()
self.model.reset_hooks()
# unpatched cache on clean token
_, unpatch_clean_cache = self.model.run_with_cache(clean_token)
sink_hook_name = get_act_name("resid_post", self.n_layers-1) # last layer resid post hook
sink_node_cache: Dict[str, Tensor] = {}
def hook_bwd(tens: Tensor, hook: HookPoint):
sink_node_cache[hook.name] = tens.detach() # type: ignore
fwd_cache = {}
with t.set_grad_enabled(True):
with self._detach_error_term(True):
with self.model.hooks(
bwd_hooks=[(sink_hook_name, hook_bwd)]
):
with self._setup_forward_model_hook(transfer_grad=False): # save memory, we don't need transfer_grad here
with self._setup_fwd_bwd_grad_sae_hook(
fwd_cache=fwd_cache,
bwd_cache={},
pass_through_grad=pass_through_grad,
):
metric(self.model(clean_token)).backward(retain_graph=retain_graph)
current_grad = sink_node_cache[sink_hook_name] # (b, seq, d_model)
bwd_cache = self._TE_using_virtual_weight(
current_grad,
fwd_cache=fwd_cache,
unpatch_clean_cache=unpatch_clean_cache,
)
node_grads = {}
for node, index in self.nodes.keys():
node_grads[(node, index)] = bwd_cache[node.name]
cache = {}
for sae_name in self.dict_saes.keys():
cache[sae_name] = cache_to_sparseact(
fwd_cache,
sae_hook_name(sae_name),
error_term_name(sae_name) if self.use_error_term else None,
)
self.model.reset_hooks()
for sae in self.dict_saes.values():
sae.reset_hooks()
return node_grads, cache
def _gradient_wrt_nodes_vw_ig(
self,
clean_token: Tensor,
corrupt_cache: ActivationCache | Dict[str, Tensor],
metric: Callable[[Tensor], Tensor],
retain_graph: bool = False,
verbose: bool = False,
**kwargs,
) -> Tuple[
Dict[Tuple[Node, Index], SparseAct], # node effects
Dict[str, SparseAct]
]:
'''
Using virtual weight to compute node grad
'''
steps = kwargs.get("steps", 10)
self._check_graph()
self.model_setup()
self.model.reset_hooks()
sink_hook_name = get_act_name("resid_post", self.n_layers-1) # last layer resid post hook
sink_node_cache: Dict[str, Tensor] = {}
def hook_bwd(tens: Tensor, hook: HookPoint):
sink_node_cache[hook.name] = tens.detach() # type: ignore
all_fwd_cache = {}
all_bwd_cache = {}
for target_name in self.dict_saes.keys():
for step in range(steps):
frac = step / steps
fwd_cache = {}
with t.set_grad_enabled(True):
with self._detach_error_term(True):
with self.model.hooks(
bwd_hooks=[(sink_hook_name, hook_bwd)]
):
with self._setup_forward_model_hook(transfer_grad=False): # save memory, we don't need transfer_grad here
with self._setup_fwd_bwd_grad_sae_hook_ig(
target_name=target_name,
frac=frac,
fwd_cache=fwd_cache,
bwd_cache={},
corrupt_cache=corrupt_cache,
):
metric(self.model(clean_token)).backward(retain_graph=retain_graph)
# run the model on interpolate intervention
with self._setup_forward_model_hook(transfer_grad=False): # save memory, we don't need transfer_grad here
with self._setup_virtual_weight_sae_hook_ig(
target_name=target_name,
frac=frac,
corrupt_cache=corrupt_cache,
):
_, unpatch_clean_cache = self.model.run_with_cache(clean_token)
current_grad = sink_node_cache[sink_hook_name] # (b, seq, d_model)
bwd_cache = self._TE_using_virtual_weight(
current_grad,
fwd_cache=fwd_cache,
unpatch_clean_cache=unpatch_clean_cache,
)
if step == 0:
add_cache(all_fwd_cache, fwd_cache)
add_cache(all_bwd_cache, bwd_cache)
# average the gradients
for key in all_bwd_cache.keys():
all_bwd_cache[key] /= steps
node_grads = {}
for node, index in self.nodes.keys():
node_grads[(node, index)] = all_bwd_cache[node.name]
cache = {}
for sae_name in self.dict_saes.keys():
cache[sae_name] = cache_to_sparseact(
all_fwd_cache,
sae_hook_name(sae_name),
error_term_name(sae_name) if self.use_error_term else None,
)
self.model.reset_hooks()
for sae in self.dict_saes.values():
sae.reset_hooks()
return node_grads, cache
def _TE_using_virtual_weight(
self,
current_grad: Tensor,
fwd_cache: ActivationCache | Dict[str, Tensor],
unpatch_clean_cache: ActivationCache | Dict[str, Tensor],
) -> Dict[str, Tensor]:
bwd_cache = {}
for layer in reversed(range(self.n_layers)):
mlp_name = get_act_name("mlp_out", layer)
attn_name = get_act_name("attn_out", layer)
resid_pre_name = get_act_name("resid_pre", layer)
# (b, seq, d_model) @ (d_model, d_sae) -> (b, seq, d_sae)
bwd_cache[mlp_name] = SparseAct(
act=current_grad @ self.dict_saes[mlp_name].W_dec.T,
res = current_grad if self.use_error_term else None,
)
grad_through_mlp = gradient_with_mlp(
end_feature_vec=current_grad, # (b, seq, d_model)
start_feature_vec=None,
transcoder_enc=self.dict_saes[mlp_name].W_enc, # (d_model, d_sae)
transcoder_dec=self.dict_saes[mlp_name].W_dec, # (d_sae, d_model)
transcoder_act=fwd_cache[sae_hook_name(mlp_name)], # (b, seq, d_sae)
layer_end=layer,
pos_end=None,
seq_length=self.seq_length, # type: ignore
batch_size=current_grad.shape[0],
use_error_term=self.use_error_term,
cache=unpatch_clean_cache,
device=self.device # type: ignore
).act # (b, seq, d_model)
current_grad = current_grad + grad_through_mlp # resid_mid_grad
if self.check_attn_out:
# (b, seq, d_model) @ (d_model, d_sae) -> (b, seq, d_sae)
bwd_cache[attn_name] = SparseAct(
act=current_grad @ self.dict_saes[attn_name].W_dec.T,
res = current_grad if self.use_error_term else None,
)
grad_through_attn = gradient_with_attn(
model=self.model,
end_feature_vec=current_grad, # (b, seq, d_model)
start_feature_vec=None,
layer_end=layer,
pos_end=None,
use_error_term=self.use_error_term,
cache=unpatch_clean_cache,
device=self.device # type: ignore
).act # (b, seq, d_model)
current_grad = current_grad + grad_through_attn # resid_pre_grad
if self.check_resid_pre:
# (b, seq, d_model) @ (d_model, d_sae) -> (b, seq, d_sae)
bwd_cache[resid_pre_name] = SparseAct(
act=current_grad @ self.dict_saes[resid_pre_name].W_dec.T,
res = current_grad if self.use_error_term else None,
)
return bwd_cache
def _gradient_wrt_edges(
self,
clean_token: Tensor,
corrupt_cache: ActivationCache | Dict[str, Tensor],
node_grads: Dict[Tuple[Node, Index], SparseAct],
verbose: bool = False,
**kwargs,
) -> Tuple[
Dict[Tuple[Node, Index, Node, Index], Tensor], # edge effects
Dict[str, SparseAct]
]:
self._check_graph()
self.model_setup()
self.model.reset_hooks()
self.model.reset_saes()
for _, sae in self.dict_saes.items():
sae.reset_hooks()
gradient_mode = kwargs.get('edge_gradient_mode', 'virtual_weight')
# activation cache on clean token
_, unpatch_clean_cache = self.model.run_with_cache(clean_token)
# SAE act cache on clean token
_, clean_cache = self.forward(clean_token, corrupt_cache=None)
edge_grads: Dict[Tuple[Node, Index, Node, Index], Tensor] = {}
for layer, connection in tqdm(self.connection.items(), disable=not verbose):
if verbose:
print(f"Layer {layer}:")
for hook_position_end, list_hook_positions_start in tqdm(connection.items(), disable=not verbose):
assert hook_position_end in node_grads, f"Node gradient of {hook_position_end} is not provided."
for hook_position_start in list_hook_positions_start:
corrupt_sparse_act = cache_to_sparseact(
corrupt_cache,
sae_hook_name(hook_position_start[0].name),
error_term_name(hook_position_start[0].name) if self.use_error_term else None,
)
right_vec = corrupt_sparse_act - clean_cache[hook_position_start[0].name]
if gradient_mode == 'virtual_weight':
edge_grads[hook_position_start + hook_position_end] = self._edge_attribution_trans(
unpatch_clean_cache,
hook_position_end,
hook_position_start,
node_grads[hook_position_end],
right_vec,
**kwargs,
)
elif gradient_mode == 'gradient':
edge_grads[hook_position_start + hook_position_end] = self._edge_attribution(
clean_token,
hook_position_end,
hook_position_start,
node_grads[hook_position_end],
right_vec,
layer,
**kwargs,
)
else:
raise NotImplementedError(f"gradient_mode {gradient_mode} is not supported")
return edge_grads, clean_cache
def _edge_attribution_trans(
self,
unpatched_clean_cache: ActivationCache | Dict[str, Tensor],
hook_position_end: Tuple[Node, Index],
hook_position_start: Tuple[Node, Index],
leftvec: SparseAct,
rightvec: SparseAct,
**kwargs,
) -> Tensor:
d_sae_end = self.dict_saes[hook_position_end[0].name].cfg.d_sae
d_sae_start = self.dict_saes[hook_position_start[0].name].cfg.d_sae
aggregate_dim = [0] if self.token_wise else [0, 1]
edge_effect = {}
all_error = []
for end_node, end_index in self.active_nodes(*hook_position_end):
if isinstance(end_index, ErrorIndex):
all_error.append((end_node, end_index))
elif isinstance(end_index, FeatureIndex):
feat_id = end_index.idx[-1]
pos_end = end_index.idx[-2]
index = t.tensor(list(end_index.idx), device=self.device)
end_node_grad = leftvec.act[:, pos_end, feat_id].unsqueeze(-1).unsqueeze(-1) # (b, 1, 1)
end_feature_vec = self.dict_saes[end_node.name].W_enc[:, feat_id].unsqueeze(0).unsqueeze(0) # (1, 1, d_model)
grad_dot_leftvec_tensor = self._DE_using_virtual_weight( # (b, seq, d_model)
grad=end_feature_vec,
pos_end=pos_end,
batch_size=leftvec.act.shape[0],
unpatched_clean_cache=unpatched_clean_cache,
hook_position_end=hook_position_end,
hook_position_start=hook_position_start,
) * end_node_grad
grad_dot_leftvec = SparseAct(
# (b, seq, d_model) @ (d_model, d_sae)
act=grad_dot_leftvec_tensor @ self.dict_saes[hook_position_start[0].name].W_dec.T,
res=grad_dot_leftvec_tensor if self.use_error_term else None
)
'''
edge_effect shape (seq, d_sae+1, seq, d_sae+1) or (d_sae+1, d_sae+1) in sparse_coo tensor
the sparse_coo will have the shape:
--> indices of shape (2, num_active) or (1, num_active)
--> values of shape (num_active, seq, d_sae+1) or (num_active, d_sae+1)
'''
edge_effect[index] = ( # (seq, d_sae+1) || (d_sae+1)
grad_dot_leftvec @ rightvec
).sum(aggregate_dim).to_tensor()
else:
raise ValueError(f"end_index of type {type(end_index)} is not supported.")
'''
The gradient of error node to upstream node is 1 - gradient of sum end_feature_node
We multiply error grad so that we only have to backward once (Jacobian vector product)
The "end_feature_dependent" sums all of the gradient of end_feature_node.
'''
if self.use_error_term:
all_end_node_grad: Tensor = leftvec.res # type: ignore
feature_coef_to_cal_error_edge = einops.einsum(
all_end_node_grad, self.dict_saes[hook_position_end[0].name].W_dec,
"b seq d_model, d_sae_end d_model -> b seq d_sae_end",
)
end_feature_dependent = einops.einsum(
feature_coef_to_cal_error_edge, self.dict_saes[hook_position_end[0].name].W_enc,
"b seq d_sae_end, d_model d_sae_end -> b seq d_model"
)
for end_error_node, end_error_index in all_error:
pos_end = end_error_index.idx[0]
index = t.tensor(list(end_error_index.idx + (d_sae_end,)), device=self.device)
end_node_grad = leftvec.res[:, pos_end].unsqueeze(1) # (b, 1, d_model) # type: ignore
end_error_grad_tensor = self._DE_using_virtual_weight( # (b, seq, d_model)
grad=end_node_grad - end_feature_dependent[:, pos_end].unsqueeze(1), # (b, 1, d_model)
pos_end=pos_end,
batch_size=leftvec.act.shape[0],
unpatched_clean_cache=unpatched_clean_cache,
hook_position_end=hook_position_end,
hook_position_start=hook_position_start,
)
end_error_grad = SparseAct(
# (b, seq, d_model) @ (d_model, d_sae)
act=end_error_grad_tensor @ self.dict_saes[hook_position_start[0].name].W_dec.T,
res=end_error_grad_tensor if self.use_error_term else None
)
edge_effect[index] = ( # (seq, d_sae_start+1) | (d_sae_start+1)
end_error_grad @ rightvec
).sum(aggregate_dim).to_tensor()
seq = int(self.seq_length) # type: ignore
num_end = d_sae_end
num_start = d_sae_start
if self.use_error_term:
num_end += 1
num_start += 1
if len(edge_effect.keys()) != 0:
indices = t.stack(list(edge_effect.keys()), dim=0).T # shape (2, num_active) or (1, num_active)
values = t.stack([value for value in edge_effect.values()], dim=0) # shape (num_active, seq, d_sae+1) or (num_active, d_sae+1)
# if no active nodes, return empty tensor
else:
indices = t.empty((2, 0) if self.token_wise else (1, 0), dtype=t.long).to(self.device)
values = t.empty((0, seq, num_start) if self.token_wise else (0, num_start), dtype=t.float).to(self.device)
if self.token_wise:
return t.sparse_coo_tensor(indices, values, size=(seq, num_end, seq, num_start)).coalesce()
else:
return t.sparse_coo_tensor(indices, values, size=(num_end, num_start)).coalesce()
def _edge_attribution(
self,
clean_token: Tensor,
hook_position_end: Tuple[Node, Index],
hook_position_start: Tuple[Node, Index],
leftvec: SparseAct,
rightvec: SparseAct,
layer: int,
**kwargs,
) -> Tensor:
d_sae_end = self.dict_saes[hook_position_end[0].name].cfg.d_sae
d_sae_start = self.dict_saes[hook_position_start[0].name].cfg.d_sae
to_bwd_cache = {}
bwd_cache = {}
edge_effect = {}
with t.set_grad_enabled(True):
with self._detach_error_term(False, hook_position_end[0].name):
with self._setup_forward_model_hook(transfer_grad=False): # must be false, no gradient from intermediate component
with self._setup_fwd_bwd_edge_grad_sae_hook(
bwd_cache=bwd_cache,
to_bwd_cache=to_bwd_cache,
hook_position_start=hook_position_start,
hook_position_end=hook_position_end,
):
self.model.forward(clean_token, return_type=None, stop_at_layer=layer+1)
aggregate_dim = [0] if self.token_wise else [0, 1]
to_bwd = ( # (seq, d_sae+1) || (d_sae+1)
cache_to_sparseact(
to_bwd_cache,
sae_hook_name(hook_position_end[0].name),
error_term_name(hook_position_end[0].name) if self.use_error_term else None,
) @ leftvec.detach()
).sum(aggregate_dim).to_tensor()
del to_bwd_cache
for end_node, end_index in self.active_nodes(*hook_position_end):
if isinstance(end_index, ErrorIndex):
# the last index is error: shape (d_sae+1) so last index is d_sae
to_bwd[end_index.idx + (d_sae_end,)].backward(retain_graph=True)
index = t.tensor(list(end_index.idx + (d_sae_end,)), device=self.device)
elif isinstance(end_index, FeatureIndex):
to_bwd[end_index.idx].backward(retain_graph=True)
index = t.tensor(list(end_index.idx), device=self.device)
else:
raise ValueError(f"end_index of type {type(end_index)} is not supported.")
'''
edge_effect shape (seq, d_sae+1, seq, d_sae+1) or (d_sae+1, d_sae+1) in sparse_coo tensor
the sparse_coo will have the shape:
--> indices of shape (2, num_active) or (1, num_active)
--> values of shape (num_active, seq, d_sae+1) or (num_active, d_sae+1)
'''
edge_effect[index] = ( # (seq, d_sae+1) || (d_sae+1)
cache_to_sparseact(
bwd_cache,
sae_hook_name(hook_position_start[0].name),
error_term_name(hook_position_start[0].name) if self.use_error_term else None,
) @ rightvec
).sum(aggregate_dim).to_tensor()
del bwd_cache
seq = int(self.seq_length) # type: ignore
num_end = d_sae_end
num_start = d_sae_start
if self.use_error_term:
num_end += 1
num_start += 1
if len(edge_effect.keys()) != 0:
indices = t.stack(list(edge_effect.keys()), dim=0).T # shape (2, num_active) or (1, num_active)
values = t.stack([value for value in edge_effect.values()], dim=0) # shape (num_active, seq, d_sae+1) or (num_active, d_sae+1)
# if no active nodes, return empty tensor
else:
indices = t.empty((2, 0) if self.token_wise else (1, 0), dtype=t.long).to(self.device)
values = t.empty((0, seq, num_start) if self.token_wise else (0, num_start), dtype=t.float).to(self.device)
if self.token_wise:
return t.sparse_coo_tensor(indices, values, size=(seq, num_end, seq, num_start)).coalesce()
else:
return t.sparse_coo_tensor(indices, values, size=(num_end, num_start)).coalesce()
def _DE_using_virtual_weight(
self,
grad: Tensor, # (1, 1, d_model) or (b, 1, d_model)
pos_end: int,
batch_size: int,
unpatched_clean_cache: ActivationCache | Dict[str, Tensor],
hook_position_end: Tuple[Node, Index],
hook_position_start: Tuple[Node, Index],
):
start_layer = int(hook_position_start[0].name.split(".")[1])
end_layer = int(hook_position_end[0].name.split(".")[1])
path = [] # path attention for gradient, if we have attention circuit then the path is empty
if not self.check_attn_out:
if not "resid_pre" in hook_position_end[0].name:
path.append(("attn", end_layer)) # attn or mlp at end hook -> goes through attn at end_layer
for layer in reversed(range(start_layer+1, end_layer)):
path.append(("attn", layer))
if "resid_pre" in hook_position_start[0].name:
path.append(("attn", start_layer)) # resid_pre at start hook -> goes through attn at start_layer
if "mlp_out" in hook_position_end[0].name:
ln = "ln2"
elif "attn_out" in hook_position_end[0].name:
ln = "ln1"
else:
ln = None
if self.check_attn_out and "attn_out" in hook_position_end[0].name:
current_grad = gradient_with_attn(
model=self.model,
end_feature_vec=grad, # (b, seq, d_model)
start_feature_vec=None,
layer_end=end_layer,
pos_end=pos_end,
use_error_term=self.use_error_term,
cache=unpatched_clean_cache,
device=self.device # type: ignore
).act # (b, seq, d_model)
else:
current_grad = gradient_ln_only( # (b, seq, d_model)
end_feature_vec=grad,
start_feature_vec=None,
layer_end=end_layer,
pos_end=pos_end,
seq_length=self.seq_length, # type: ignore
batch_size=batch_size,
use_error_term=self.use_error_term,
cache=unpatched_clean_cache,
device=self.device, # type: ignore
ln=ln,
).act
for _, layer in path:
grad_through_attn = gradient_with_attn(
model=self.model,
end_feature_vec=current_grad, # (b, seq, d_model)
start_feature_vec=None,
layer_end=layer,
pos_end=None,
use_error_term=self.use_error_term,
cache=unpatched_clean_cache,
device=self.device # type: ignore
).act # (b, seq, d_model)
current_grad = current_grad + grad_through_attn
return current_grad
def run_model(
self,
toks: Tensor,
use_error_term: bool | None = None,
) -> Tuple[Tensor, ActivationCache]:
"""
Runs an MLP transcoder(s) on a batch of tokens.
"""
fwd_cache = {}
with self._setup_forward_model_hook(use_error_term=use_error_term, transfer_grad=False):
with self._setup_fwd_sae_hook(
fwd_cache=fwd_cache, corrupt_cache=None, patch_deleted_comp=False, use_error_term=use_error_term
):
logits = self.model(toks)
return logits, ActivationCache(cache_dict=fwd_cache, model=self.model) # type: ignore
@contextmanager
def _setup_fwd_sae_hook(
self,
fwd_cache: Dict[str, Tensor],
corrupt_cache: ActivationCache | Dict[str, Tensor] | None,
patch_deleted_comp: bool = False,
use_error_term: bool | None = None,
):
use_error_term = use_error_term if use_error_term is not None else self.use_error_term
def hook_sae_fwd(act: Tensor, hook: HookPoint, sae_name: str) -> Tensor:
if hook.name == "hook_sae_acts_post" or hook.name == sae_hook_name(sae_name):
if patch_deleted_comp and corrupt_cache is not None:
act_mask = self.nodes[(sae_name, (None,))].act == 0 # type: ignore
act[:, act_mask] = corrupt_cache[sae_hook_name(sae_name)][:, act_mask]
fwd_cache[sae_hook_name(sae_name)] = act.detach()
elif (hook.name == "hook_sae_error" or hook.name == error_term_name(sae_name)) and use_error_term:
if patch_deleted_comp and corrupt_cache is not None:
resc_mask = self.nodes[(sae_name, (None,))].resc == 0 # type: ignore
act[:, resc_mask] = corrupt_cache[error_term_name(sae_name)][:, resc_mask]
fwd_cache[error_term_name(sae_name)] = act.detach()
return act
try:
with self._setup_error_term(use_error_term):
for sae_name, sae in self.dict_saes.items():
sae.add_hook(
lambda name: True,
partial(hook_sae_fwd, sae_name=sae_name),
dir="fwd",
)
yield
finally:
for sae in self.dict_saes.values():
sae.reset_hooks()
@contextmanager
def _setup_fwd_bwd_grad_sae_hook(
self,
fwd_cache: Dict[str, Tensor],
bwd_cache: Dict[str, Tensor],
pass_through_grad: bool,
):
def hook_sae_fwd(act: Tensor, hook: HookPoint, sae_name: str) -> Tensor:
if hook.name == "hook_sae_acts_post" or hook.name == sae_hook_name(sae_name):
fwd_cache[sae_hook_name(sae_name)] = act.detach()
elif (hook.name == "hook_sae_error" or error_term_name(sae_name) == hook.name) and self.use_error_term:
fwd_cache[error_term_name(sae_name)] = act.detach()
return act
pass_through_cache = {}
def hook_sae_bwd(grad: Tensor, hook: HookPoint, sae_name: str) -> None:
if hook.name == "hook_sae_acts_post" or hook.name == sae_hook_name(sae_name):
bwd_cache[sae_hook_name(sae_name)] = grad.detach()
elif hook.name == "hook_sae_output" or hook.name == output_hook_name(sae_name):
if self.use_error_term:
# we have: output = recon + stop_grad(error_term)
# so, the TRUE error_grad (if not stop_grad) is output_grad
# due to: output = recon + error_term
bwd_cache[error_term_name(sae_name)] = grad.detach()
# avoid pass through grad for transcoder
if pass_through_grad and sae_name not in self.output_hooks:
pass_through_cache[output_hook_name(sae_name)] = grad.detach()
elif hook.name == "hook_sae_input" or hook.name == input_hook_name(sae_name):
if pass_through_grad:
# we have to modify inplace instead of grad = ... and then return grad
# because, returning a tensor in bwd pass hook is buggy somehow
if sae_name not in self.output_hooks:
# no pass through grad for transcoder
grad.copy_(pass_through_cache[output_hook_name(sae_name)])
else:
# avoid gradient flow to upstream from the transcoder
grad.zero_()
try:
with self._setup_error_term(self.use_error_term):
for sae_name, sae in self.dict_saes.items():
sae.add_hook(
lambda name: True,
partial(hook_sae_fwd, sae_name=sae_name),
dir="fwd",
)
sae.add_hook(
lambda name: True,
partial(hook_sae_bwd, sae_name=sae_name),
dir="bwd",
)
yield
finally:
for sae in self.dict_saes.values():
sae.reset_hooks()
@contextmanager
def _setup_fwd_bwd_grad_sae_hook_ig(
self,
target_name: str,
frac: float,
fwd_cache: Dict[str, Tensor],
bwd_cache: Dict[str, Tensor],
corrupt_cache: ActivationCache | Dict[str, Tensor],
):
def hook_sae_fwd(act: Tensor, hook: HookPoint, sae_name: str, target_name: str, frac: float) -> Tensor:
if hook.name == "hook_sae_acts_post" or hook.name == sae_hook_name(sae_name):
# interpolate for integrated gradients
if hook.name == sae_hook_name(target_name) or _hook_name(sae_name, hook.name) == sae_hook_name(target_name):
act = interpolate(
corrupt_cache[sae_hook_name(sae_name)],
act,
frac,
)
fwd_cache[sae_hook_name(sae_name)] = act.detach()
elif (hook.name == "hook_sae_error" or error_term_name(sae_name) == hook.name) and self.use_error_term:
# interpolate for integrated gradients
if hook.name == error_term_name(target_name) or _hook_name(sae_name, hook.name) == error_term_name(target_name):
act = interpolate(
corrupt_cache[error_term_name(sae_name)],
act,
frac,
)
fwd_cache[error_term_name(sae_name)] = act.detach()
return act
def hook_sae_bwd(grad: Tensor, hook: HookPoint, sae_name: str) -> None:
if hook.name == "hook_sae_acts_post" or hook.name == sae_hook_name(sae_name):
create_list(bwd_cache, sae_hook_name(sae_name))
bwd_cache[sae_hook_name(sae_name)] += grad.detach()
elif hook.name == "hook_sae_output" or hook.name == output_hook_name(sae_name):
if self.use_error_term:
# we have: output = recon + stop_grad(error_term)
# so, the TRUE error_grad (if not stop_grad) is output_grad
# due to: output = recon + error_term
create_list(bwd_cache, error_term_name(sae_name))
bwd_cache[error_term_name(sae_name)] += grad.detach()
try:
with self._setup_error_term(self.use_error_term):
for sae_name, sae in self.dict_saes.items():
sae.add_hook(
lambda name: True,
partial(hook_sae_fwd, sae_name=sae_name, target_name=target_name, frac=frac),
dir="fwd",
)
sae.add_hook(
lambda name: True,
partial(hook_sae_bwd, sae_name=sae_name),
dir="bwd",
)
yield
finally:
for sae in self.dict_saes.values():
sae.reset_hooks()
@contextmanager
def _setup_virtual_weight_sae_hook_ig(
self,
target_name: str,
frac: float,
corrupt_cache: ActivationCache | Dict[str, Tensor],
):
def hook_sae_fwd(act: Tensor, hook: HookPoint, sae_name: str, target_name: str, frac: float) -> Tensor:
if hook.name == "hook_sae_acts_post" or hook.name == sae_hook_name(sae_name):
# interpolate for integrated gradients
if hook.name == sae_hook_name(target_name) or _hook_name(sae_name, hook.name) == sae_hook_name(target_name):
act = interpolate(
corrupt_cache[sae_hook_name(sae_name)],
act,
frac,
)
elif (hook.name == "hook_sae_error" or error_term_name(sae_name) == hook.name) and self.use_error_term:
# interpolate for integrated gradients
if hook.name == error_term_name(target_name) or _hook_name(sae_name, hook.name) == error_term_name(target_name):
act = interpolate(
corrupt_cache[error_term_name(sae_name)],
act,
frac,
)
return act
try:
with self._setup_error_term(self.use_error_term):
for sae_name, sae in self.dict_saes.items():
sae.add_hook(
lambda name: True,
partial(hook_sae_fwd, sae_name=sae_name, target_name=target_name, frac=frac),
dir="fwd",
)
yield
finally:
for sae in self.dict_saes.values():
sae.reset_hooks()
@contextmanager
def _setup_fwd_bwd_edge_grad_sae_hook(
self,
bwd_cache: Dict[str, Tensor],
to_bwd_cache: Dict[str, Tensor],
hook_position_start: Tuple[Node, Index],
hook_position_end: Tuple[Node, Index],
):
# print(hook_position_end[0].name, hook_position_start[0].name)
def hook_sae_bwd(grad: Tensor, hook: HookPoint, sae_name: str, bwd_cache: Dict) -> None:
# print(sae_name, hook.name, t.sum(grad != 0).item())
if hook.name == "hook_sae_acts_post" or hook.name == sae_hook_name(sae_name):
if hook.name == sae_hook_name(hook_position_start[0].name) or _hook_name(sae_name, hook.name) == sae_hook_name(hook_position_start[0].name):
# only store the gradient at the start position
bwd_cache[sae_hook_name(sae_name)] = grad.detach()
elif hook.name == "hook_sae_output" or hook.name == output_hook_name(sae_name):
if hook.name == output_hook_name(hook_position_start[0].name) or _hook_name(sae_name, hook.name) == output_hook_name(hook_position_start[0].name):
# we have: output = recon + stop_grad(error_term)
# so, the TRUE error_grad (if not stop_grad) is output_grad
# due to: output = recon + error_term
if self.use_error_term:
# only store the gradient at the start position
bwd_cache[error_term_name(sae_name)] = grad.detach()
# IMPORTANT NOTE for reproducibility:
# we zero grad of intermediate components
# but zero grad the output hook of SAE will ALSO ZERO GRAD the resid mid / post --> no downstream grads
# so, we instead zero grad of the input hook of SAE
elif hook.name == "hook_sae_input" or hook.name == input_hook_name(sae_name):
if hook.name == input_hook_name(hook_position_end[0].name) or _hook_name(sae_name, hook.name) == input_hook_name(hook_position_end[0].name):
pass # no zero grad for end position
else:
grad.zero_()
def hook_sae_fwd(act: Tensor, hook: HookPoint, sae_name: str, to_bwd_cache: Dict) -> Tensor:
if hook.name == "hook_sae_acts_post" or hook.name == sae_hook_name(sae_name):
if hook.name == sae_hook_name(hook_position_end[0].name) or _hook_name(sae_name, hook.name) == sae_hook_name(hook_position_end[0].name):
# store activation for backward later
to_bwd_cache[sae_hook_name(sae_name)] = act
elif (hook.name == "hook_sae_error" or error_term_name(sae_name) == hook.name) and self.use_error_term:
if hook.name == error_term_name(hook_position_end[0].name) or _hook_name(sae_name, hook.name) == error_term_name(hook_position_end[0].name):
# store activation for backward later
to_bwd_cache[error_term_name(sae_name)] = act
return act
try:
with self._setup_error_term(self.use_error_term):
for sae_name, sae in self.dict_saes.items():
sae.add_hook(
lambda name: True,
partial(hook_sae_fwd, sae_name=sae_name, to_bwd_cache=to_bwd_cache),
dir="fwd",
)
sae.add_hook(
lambda name: True,
partial(hook_sae_bwd, sae_name=sae_name, bwd_cache=bwd_cache),
dir="bwd",
)
yield
finally:
for sae in self.dict_saes.values():
sae.reset_hooks()
@contextmanager
def _setup_forward_model_hook(self, use_error_term: bool | None = None, transfer_grad: bool = True):
trans_cache = {}
# Hook function at transcoder input: caches the activations before transcoder
def hook_transcoder_input(activations: Tensor, hook: HookPoint, transcoder_idx: int):
trans_cache[transcoder_idx] = activations.clone()
# Hook function at transcoder output: replaces activations with transcoder output
def hook_transcoder_output(activations: Tensor, hook: HookPoint, transcoder_idx: int):
trans_input = trans_cache[transcoder_idx]
if transfer_grad:
# allow gradient to pass through the transcoder output and original activations
return self.transcoders[transcoder_idx]((trans_input, activations)) + (activations - activations.detach())
else:
# no original activations gradient, but save memory
return self.transcoders[transcoder_idx]((trans_input, activations))
fwd_hooks = []
for i in range(len(self.transcoders)):
fwd_hooks.append((self.input_hooks[i], partial(hook_transcoder_input, transcoder_idx=i)))
fwd_hooks.append((self.output_hooks[i], partial(hook_transcoder_output, transcoder_idx=i)))
use_error_term = use_error_term if use_error_term is not None else self.use_error_term
try:
for sae in self._saes_to_list():
self.model.add_sae(sae, use_error_term)
for hook, func in fwd_hooks:
self.model.add_hook(hook, func, dir="fwd")
yield
finally:
self.model.reset_saes()
self.model.reset_hooks()
@contextmanager
def _setup_error_term(self, use_error_term: bool | None = None):
if use_error_term is None:
use_error_term = self.use_error_term
orig_use_error_term = {}
try:
for sae_name, sae in self.dict_saes.items():
sae.use_error_term = use_error_term
orig_use_error_term[sae_name] = sae.use_error_term
yield
finally:
for sae_name, sae in self.dict_saes.items():
sae.use_error_term = orig_use_error_term[sae_name]
def _saes_to_list(self) -> List[Any]:
return self.non_transcoders
class ESAE_FG_Trans(
ESAE_FG,
Feature_Graph_Trans,
):
def __init__(
self,
model: HookedSAETransformer,
saes: Dict[int, List[Tuple[str, Any]]],
esaes: Dict[int, List[Tuple[str, Any]]],
use_esae_error_term: bool = False,
) -> None:
super().__init__(model, saes, esaes, use_esae_error_term)
self.process_transcoder()
def forward(
self,
clean_token: Tensor,
corrupt_cache: ActivationCache | Dict[str, Tensor] | None,
patch_deleted_comp: bool = False,
**kwargs,
) -> Tuple[Tensor, Dict[str, SparseAct]]:
'''
Forward pass of the graph with clean tokens, if the edge exists, replace the activation with corrupted activation
'''
self._check_graph()
self.model.reset_hooks()
self.model_setup()
fwd_cache = {}
with t.no_grad():
with self._setup_forward_model_hook(transfer_grad=False):
with self._setup_fwd_sae_hook(
fwd_cache=fwd_cache, corrupt_cache=corrupt_cache, patch_deleted_comp=patch_deleted_comp
):
logits = self.model(clean_token)
cache = {}
for sae_name in self.dict_saes.keys():
cache[sae_name] = cache_to_sparseact(
fwd_cache,
sae_hook_name(sae_name),
sae_hook_name(error_term_name(sae_name)) if self.use_error_term else None,
error_term_name(error_term_name(sae_name)) if self.use_esae_error_term else None,
)
for sae in self.dict_saes.values():
sae.reset_hooks()
self.model.reset_hooks()
return logits, cache
def forward_backward_gradient(
self,
clean_token: Tensor,
corrupt_cache: ActivationCache | Dict[str, Tensor],
metric: Callable[[Tensor], Tensor],
retain_graph: bool = False,
mode: str = 'node',
gradient_mode: str = 'standard',
pass_through_grad: bool = False,
verbose: bool = False,
**kwargs,
) -> Tuple[
Dict[Tuple[Node, Index], SparseAct], # node effects
Dict[Tuple[Node, Index, Node, Index], Tensor], # edge effects
]:
return Feature_Graph_Trans.forward_backward_gradient(
self,
clean_token=clean_token,
corrupt_cache=corrupt_cache,
metric=metric,
retain_graph=retain_graph,
mode=mode,
gradient_mode=gradient_mode,
pass_through_grad=pass_through_grad,
verbose=verbose,
**kwargs,
)
def _gradient_wrt_nodes(
self,
clean_token: Tensor,
metric: Callable[[Tensor], Tensor],
retain_graph: bool = False,
pass_through_grad: bool = False,
verbose: bool = False,
**kwargs,
) -> Tuple[
Dict[Tuple[Node, Index], SparseAct], # node effects
Dict[str, SparseAct]
]:
self._check_graph()
self.model_setup()
self.model.reset_hooks()
for _, sae in self.dict_saes.items():
sae.reset_hooks()
fwd_cache = {}
bwd_cache = {}
with t.set_grad_enabled(True):
with self._detach_error_term(True):
with self._setup_forward_model_hook(transfer_grad=kwargs.get("transfer_grad", True)):
with self._setup_fwd_bwd_grad_sae_hook(
fwd_cache=fwd_cache,
bwd_cache=bwd_cache,
pass_through_grad=pass_through_grad,
):
metric(self.model(clean_token)).backward(retain_graph=retain_graph)
node_grads = {}
for node, index in self.nodes.keys():
node_grads[(node, index)] = cache_to_sparseact(
bwd_cache,
sae_hook_name(node.name),
sae_hook_name(error_term_name(node.name)) if self.use_error_term else None,
error_term_name(error_term_name(node.name)) if self.use_esae_error_term else None,
)
cache = {}
for sae_name in self.dict_saes.keys():
cache[sae_name] = cache_to_sparseact(
fwd_cache,
sae_hook_name(sae_name),
sae_hook_name(error_term_name(sae_name)) if self.use_error_term else None,
error_term_name(error_term_name(sae_name)) if self.use_esae_error_term else None,
)
self.model.reset_hooks()
for sae in self.dict_saes.values():
sae.reset_hooks()
return node_grads, cache
def _gradient_wrt_nodes_ig(
self,
clean_token: Tensor,
corrupt_cache: ActivationCache | Dict[str, Tensor],
metric: Callable[[Tensor], Tensor],
retain_graph: bool = False,
verbose: bool = False,
**kwargs,
) -> Tuple[
Dict[Tuple[Node, Index], SparseAct], # node effects
Dict[str, SparseAct]
]:
steps = kwargs.get("steps", 10)
self._check_graph()
self.model_setup()
self.model.reset_hooks()
for _, sae in self.dict_saes.items():
sae.reset_hooks()
fwd_cache = {}
bwd_cache = {}
with t.set_grad_enabled(True):
with self._detach_error_term(True):
with self._setup_forward_model_hook(transfer_grad=kwargs.get("transfer_grad", True)):
for target_name in self.dict_saes.keys():
for step in range(steps):
frac = step / steps
with self._setup_fwd_bwd_grad_sae_hook_ig(
target_name=target_name,
frac=frac,
fwd_cache=fwd_cache,
bwd_cache=bwd_cache,
corrupt_cache=corrupt_cache,
):
metric(self.model(clean_token)).backward(retain_graph=retain_graph)
# average the gradients
for key in bwd_cache.keys():
bwd_cache[key] /= steps
node_grads = {}
for node, index in self.nodes.keys():
node_grads[(node, index)] = cache_to_sparseact(
bwd_cache,
sae_hook_name(node.name),
sae_hook_name(error_term_name(node.name)) if self.use_error_term else None,
error_term_name(error_term_name(node.name)) if self.use_esae_error_term else None,
)
cache = {}
for sae_name in self.dict_saes.keys():
cache[sae_name] = cache_to_sparseact(
fwd_cache,
sae_hook_name(sae_name),
sae_hook_name(error_term_name(sae_name)) if self.use_error_term else None,
error_term_name(error_term_name(sae_name)) if self.use_esae_error_term else None,
)
self.model.reset_hooks()
for sae in self.dict_saes.values():
sae.reset_hooks()
return node_grads, cache
def _gradient_wrt_nodes_vw(
self,
clean_token: Tensor,
metric: Callable[[Tensor], Tensor],
retain_graph: bool = False,
pass_through_grad: bool = False,
verbose: bool = False,
**kwargs,
) -> Tuple[
Dict[Tuple[Node, Index], SparseAct], # node effects
Dict[str, SparseAct]
]:
'''
Using virtual weight to compute node grad
'''
self._check_graph()
self.model_setup()
self.model.reset_hooks()
# unpatched cache on clean token
_, unpatch_clean_cache = self.model.run_with_cache(clean_token)
sink_hook_name = get_act_name("resid_post", self.n_layers-1) # last layer resid post hook
sink_node_cache: Dict[str, Tensor] = {}
def hook_bwd(tens: Tensor, hook: HookPoint):
sink_node_cache[hook.name] = tens.detach() # type: ignore
fwd_cache = {}
with t.set_grad_enabled(True):
with self._detach_error_term(True):
with self.model.hooks(
bwd_hooks=[(sink_hook_name, hook_bwd)]
):
with self._setup_forward_model_hook(transfer_grad=False):
with self._setup_fwd_bwd_grad_sae_hook(
fwd_cache=fwd_cache,
bwd_cache={},
pass_through_grad=pass_through_grad,
):
metric(self.model(clean_token)).backward(retain_graph=retain_graph)
current_grad = sink_node_cache[sink_hook_name] # (b, seq, d_model)
bwd_cache = {}
for layer in reversed(range(self.n_layers)):
mlp_name = get_act_name("mlp_out", layer)
attn_name = get_act_name("attn_out", layer)
resid_pre_name = get_act_name("resid_pre", layer)
# (b, seq, d_model) @ (d_model, d_sae) -> (b, seq, d_sae)
bwd_cache[mlp_name] = SparseAct(
act=current_grad @ self.dict_saes[mlp_name].W_dec.T,
res = current_grad @ self.dict_esaes[mlp_name].W_dec.T if self.use_error_term else None,
resc = current_grad if self.use_esae_error_term else None,
)
grad_through_mlp = gradient_with_mlp(
end_feature_vec=current_grad, # (b, seq, d_model)
start_feature_vec=None,
transcoder_enc=self.dict_saes[mlp_name].W_enc, # (d_model, d_sae)
transcoder_dec=self.dict_saes[mlp_name].W_dec, # (d_sae, d_model)
transcoder_act=fwd_cache[sae_hook_name(mlp_name)], # (b, seq, d_sae)
layer_end=layer,
pos_end=None,
seq_length=self.seq_length, # type: ignore
batch_size=current_grad.shape[0],
use_error_term=self.use_error_term,
cache=unpatch_clean_cache,
device=self.device # type: ignore
).act # (b, seq, d_model)
current_grad = current_grad + grad_through_mlp # resid_mid_grad
if self.check_attn_out:
# (b, seq, d_model) @ (d_model, d_sae) -> (b, seq, d_sae)
bwd_cache[attn_name] = SparseAct(
act=current_grad @ self.dict_saes[attn_name].W_dec.T,
res = current_grad @ self.dict_esaes[attn_name].W_dec.T if self.use_error_term else None,
resc = current_grad if self.use_esae_error_term else None,
)
grad_through_attn = gradient_with_attn(
model=self.model,
end_feature_vec=current_grad, # (b, seq, d_model)
start_feature_vec=None,
layer_end=layer,
pos_end=None,
use_error_term=self.use_error_term,
cache=unpatch_clean_cache,
device=self.device # type: ignore
).act # (b, seq, d_model)
current_grad = current_grad + grad_through_attn # resid_pre_grad
if self.check_resid_pre:
# (b, seq, d_model) @ (d_model, d_sae) -> (b, seq, d_sae)
bwd_cache[resid_pre_name] = SparseAct(
act=current_grad @ self.dict_saes[resid_pre_name].W_dec.T,
res = current_grad @ self.dict_esaes[resid_pre_name].W_dec.T if self.use_error_term else None,
resc = current_grad if self.use_esae_error_term else None,
)
node_grads = {}
for node, index in self.nodes.keys():
node_grads[(node, index)] = bwd_cache[node.name]
cache = {}
for sae_name in self.dict_saes.keys():
cache[sae_name] = cache_to_sparseact(
fwd_cache,
sae_hook_name(sae_name),
sae_hook_name(error_term_name(sae_name)) if self.use_error_term else None,
error_term_name(error_term_name(sae_name)) if self.use_esae_error_term else None,
)
self.model.reset_hooks()
for sae in self.dict_saes.values():
sae.reset_hooks()
return node_grads, cache
def _gradient_wrt_edges(
self,
clean_token: Tensor,
corrupt_cache: ActivationCache | Dict[str, Tensor],
node_grads: Dict[Tuple[Node, Index], SparseAct],
verbose: bool = False,
**kwargs,
) -> Tuple[
Dict[Tuple[Node, Index, Node, Index], Tensor], # edge effects
Dict[str, SparseAct]
]:
self._check_graph()
self.model_setup()
self.model.reset_hooks()
self.model.reset_saes()
for _, sae in self.dict_saes.items():
sae.reset_hooks()
gradient_mode = kwargs.get('edge_gradient_mode', 'virtual_weight')
# unpatched cache on clean token
_, unpatch_clean_cache = self.model.run_with_cache(clean_token)
# potentially patched cache on clean token
_, clean_cache = self.forward(clean_token, corrupt_cache=None)
edge_grads: Dict[Tuple[Node, Index, Node, Index], Tensor] = {}
for layer, connection in tqdm(self.connection.items(), disable=not verbose):
if verbose:
print(f"Layer {layer}:")
for hook_position_end, list_hook_positions_start in tqdm(connection.items(), disable=not verbose):
assert hook_position_end in node_grads, f"Node gradient of {hook_position_end} is not provided."
for hook_position_start in list_hook_positions_start:
corrupt_sparse_act = cache_to_sparseact(
corrupt_cache,
sae_hook_name(hook_position_start[0].name),
sae_hook_name(error_term_name(hook_position_start[0].name)) if self.use_error_term else None,
error_term_name(error_term_name(hook_position_start[0].name)) if self.use_esae_error_term else None,
)
right_vec = corrupt_sparse_act - clean_cache[hook_position_start[0].name]
if gradient_mode == "virtual_weight":
edge_grads[hook_position_start + hook_position_end] = self._edge_attribution_trans(
unpatch_clean_cache,
hook_position_end,
hook_position_start,
node_grads[hook_position_end],
right_vec,
**kwargs,
)
elif gradient_mode == 'gradient':
edge_grads[hook_position_start + hook_position_end] = self._edge_attribution(
clean_token,
hook_position_end,
hook_position_start,
node_grads[hook_position_end],
right_vec,
layer,
**kwargs,
)
else:
raise NotImplementedError(f"gradient_mode {gradient_mode} is not supported")
return edge_grads, clean_cache
def _edge_attribution_trans(
self,
unpatched_clean_cache: ActivationCache | Dict[str, Tensor],
hook_position_end: Tuple[Node, Index],
hook_position_start: Tuple[Node, Index],
leftvec: SparseAct,
rightvec: SparseAct,
**kwargs,
) -> Tensor:
d_sae_end = self.dict_saes[hook_position_end[0].name].cfg.d_sae
d_sae_start = self.dict_saes[hook_position_start[0].name].cfg.d_sae
d_esae_end = self.dict_esaes[hook_position_end[0].name].cfg.d_sae
d_esae_start = self.dict_esaes[hook_position_start[0].name].cfg.d_sae
aggregate_dim = [0] if self.token_wise else [0, 1]
edge_effect = {}
all_error = []
all_feature_error = []
for end_node, end_index in self.active_nodes(*hook_position_end):
if isinstance(end_index, ErrorIndex):
all_error.append((end_node, end_index))
elif isinstance(end_index, FeatureErrorIndex):
all_feature_error.append((end_node, end_index))
elif isinstance(end_index, FeatureIndex):
feat_id = end_index.idx[-1]
pos_end = end_index.idx[-2]
index = t.tensor(list(end_index.idx), device=self.device)
end_node_grad = leftvec.act[:, pos_end, feat_id].unsqueeze(-1).unsqueeze(-1) # (b, 1, 1)
end_feature_vec = self.dict_saes[end_node.name].W_enc[:, feat_id].unsqueeze(0).unsqueeze(0) # (1, 1, d_model)
grad_dot_leftvec_tensor = self._DE_using_virtual_weight( # (b, seq, d_model)
grad=end_feature_vec,
pos_end=pos_end,
batch_size=leftvec.act.shape[0],
unpatched_clean_cache=unpatched_clean_cache,
hook_position_end=hook_position_end,
hook_position_start=hook_position_start,
) * end_node_grad
grad_dot_leftvec = SparseAct(
# (b, seq, d_model) @ (d_model, d_sae)
act=grad_dot_leftvec_tensor @ self.dict_saes[hook_position_start[0].name].W_dec.T,
res=grad_dot_leftvec_tensor @ self.dict_esaes[hook_position_start[0].name].W_dec.T if self.use_error_term else None,
resc=grad_dot_leftvec_tensor if self.use_esae_error_term else None,
)
'''
edge_effect shape (seq, d_sae+1, seq, d_sae+1) or (d_sae+1, d_sae+1) in sparse_coo tensor
the sparse_coo will have the shape:
--> indices of shape (2, num_active) or (1, num_active)
--> values of shape (num_active, seq, d_sae+1) or (num_active, d_sae+1)
'''
effect = (
grad_dot_leftvec * rightvec
).sum(aggregate_dim)
if self.use_esae_error_term:
effect.contract()
edge_effect[index] = effect.to_tensor() # (seq, d_sae_start+d_esae_start+1) | (d_sae_start+d_esae_start+1)
else:
raise ValueError(f"end_index of type {type(end_index)} is not supported.")
'''
The gradient of feature error node to upstream node is f_esae_enc - sum gradient of end_feature_node
The "grad_through_end_feat_error" computes the jacobian of f_esae_enc going through f_sae_dec and f_sae_enc
'''
if self.use_error_term:
grad_through_end_feat_error = einops.einsum(
self.dict_saes[hook_position_end[0].name].W_dec.T, self.dict_saes[hook_position_end[0].name].W_enc.T,
"d_model1 d_sae_end, d_sae_end d_model2 -> d_model1 d_model2"
)
grad_through_end_feat_error = einops.einsum( # gradient through sae_feature
self.dict_esaes[hook_position_end[0].name].W_enc.T, grad_through_end_feat_error,
"d_esae_end d_model1, d_model1 d_model2 -> d_esae_end d_model2"
)
for end_feature_error_node, end_feature_error_index in all_feature_error:
pos_end = end_feature_error_index.idx[0]
feat_id = end_feature_error_index.idx[-1]
revised_index = list(end_feature_error_index.idx)
revised_index[-1] += d_sae_end
index = t.tensor(revised_index, device=self.device)
end_node_grad = leftvec.res[:, pos_end, feat_id].unsqueeze(-1).unsqueeze(-1) # (b, 1, 1) # type: ignore
end_feature_vec = self.dict_esaes[end_feature_error_node.name].W_enc[:, feat_id].unsqueeze(0).unsqueeze(0) # (1, 1, d_model)
end_feature_error_grad_tensor = self._DE_using_virtual_weight( # (b, seq, d_model)
grad=end_feature_vec - grad_through_end_feat_error[feat_id].unsqueeze(0).unsqueeze(0), # (1, 1, d_model)
pos_end=pos_end,
batch_size=leftvec.act.shape[0],
unpatched_clean_cache=unpatched_clean_cache,
hook_position_end=hook_position_end,
hook_position_start=hook_position_start,
)
end_feature_error_grad = SparseAct(
# (b, seq, d_model) @ (d_model, d_sae)
act=end_feature_error_grad_tensor @ self.dict_saes[hook_position_start[0].name].W_dec.T,
res=end_feature_error_grad_tensor @ self.dict_esaes[hook_position_start[0].name].W_dec.T,
resc=end_feature_error_grad_tensor if self.use_esae_error_term else None,
)
effect = (
end_feature_error_grad * rightvec
).sum(aggregate_dim)
if self.use_esae_error_term:
effect.contract()
edge_effect[index] = effect.to_tensor() # (seq, d_sae_start+d_esae_start+1) | (d_sae_start+d_esae_start+1)
'''
The gradient of feature error node to upstream node is:
f_esae_enc - sum gradient of end_feature_node - sum gradient of end_feaeture_error_node
The gradient of feature error node to upstream node is f_esae_enc - sum gradient of end_feature_node (see above)
The "grad_through_end_feat_error" computes the jacobian of f_esae_enc going through f_sae_dec and f_sae_enc
The "feature_error_coef_to_cal_error_edge" calculates the leftvec (metric gradient) at the esae_error | esae_feature
We can then have the contribution of feature_error_node by:
The contribution of feature_end_node is the gradient of f_esae_enc - sum grad_through_end_feat_error * feature_error_coef_to_cal_error_edge
The "end_feature_dependent" sums all of the gradient of end_feature_node.
'''
if self.use_esae_error_term and self.use_error_term:
all_end_node_grad: Tensor = leftvec.resc # type: ignore
'''
End feauture contribution
'''
feature_coef_to_cal_error_edge = einops.einsum(
all_end_node_grad, self.dict_saes[hook_position_end[0].name].W_dec,
"b seq d_model, d_sae_end d_model -> b seq d_sae_end",
)
end_feature_dependent = einops.einsum(
feature_coef_to_cal_error_edge, self.dict_saes[hook_position_end[0].name].W_enc,
"b seq d_sae_end, d_model d_sae_end -> b seq d_model"
)
'''
End feature error contribution
'''
grad_through_end_feat_error = einops.einsum(
self.dict_saes[hook_position_end[0].name].W_dec.T, self.dict_saes[hook_position_end[0].name].W_enc.T,
"d_model1 d_sae_end, d_sae_end d_model2 -> d_model1 d_model2"
)
grad_through_end_feat_error = einops.einsum( # gradient through sae_feature
self.dict_esaes[hook_position_end[0].name].W_enc.T, grad_through_end_feat_error,
"d_esae_end d_model1, d_model1 d_model2 -> d_esae_end d_model2"
)
feature_error_coef_to_cal_error_edge = einops.einsum(
all_end_node_grad, self.dict_esaes[hook_position_end[0].name].W_dec,
"b seq d_model, d_esae_end d_model -> b seq d_esae_end",
)
feature_error_dependent = einops.einsum(
self.dict_esaes[hook_position_end[0].name].W_enc.T - grad_through_end_feat_error,
feature_error_coef_to_cal_error_edge,
"d_esae_end d_model, b seq d_esae_end -> b seq d_model",
)
for end_error_node, end_error_index in all_error:
pos_end = end_error_index.idx[0]
index = t.tensor(list(end_error_index.idx + (d_sae_end+d_esae_end,)), device=self.device)
end_node_grad = leftvec.resc[:, pos_end].unsqueeze(1) # (b, 1, d_model) # type: ignore
end_error_grad_tensor = self._DE_using_virtual_weight( # (b, seq, d_model)
# (b, 1, d_model)
grad= end_node_grad - end_feature_dependent[:, pos_end].unsqueeze(1) - feature_error_dependent[:, pos_end].unsqueeze(1),
pos_end=pos_end,
batch_size=leftvec.act.shape[0],
unpatched_clean_cache=unpatched_clean_cache,
hook_position_end=hook_position_end,
hook_position_start=hook_position_start,
)
end_error_grad = SparseAct(
# (b, seq, d_model) @ (d_model, d_sae)
act=end_error_grad_tensor @ self.dict_saes[hook_position_start[0].name].W_dec.T,
res=end_error_grad_tensor @ self.dict_esaes[hook_position_start[0].name].W_dec.T,
resc=end_error_grad_tensor,
)
effect = (
end_error_grad * rightvec
).sum(aggregate_dim)
if self.use_esae_error_term:
effect.contract()
edge_effect[index] = effect.to_tensor() # (seq, d_sae_start+d_esae_start+1) | (d_sae_start+d_esae_start+1)
seq = int(self.seq_length) # type: ignore
num_end = d_sae_end
num_start = d_sae_start
if self.use_error_term:
num_end += d_esae_end
num_start += d_esae_start
if self.use_esae_error_term:
num_end += 1
num_start += 1
if len(edge_effect.keys()) != 0:
indices = t.stack(list(edge_effect.keys()), dim=0).T # shape (2, num_active) or (1, num_active)
values = t.stack([value for value in edge_effect.values()], dim=0) # shape (num_active, seq, d_sae+d_esae+1) or (num_active, d_sae+d_esae+1)
# if no active nodes, return empty tensor
else:
indices = t.empty((2, 0) if self.token_wise else (1, 0), dtype=t.long).to(self.device)
values = t.empty((0, seq, num_start) if self.token_wise else (0, num_start), dtype=t.float).to(self.device)
if self.token_wise:
return t.sparse_coo_tensor(indices, values, size=(seq, num_end, seq, num_start)).coalesce()
else:
return t.sparse_coo_tensor(indices, values, size=(num_end, num_start)).coalesce()
def _edge_attribution(
self,
clean_token: Tensor,
hook_position_end: Tuple[Node, Index],
hook_position_start: Tuple[Node, Index],
leftvec: SparseAct,
rightvec: SparseAct,
layer: int,
**kwargs,
) -> Tensor:
d_sae_end = self.dict_saes[hook_position_end[0].name].cfg.d_sae
d_esae_end = self.dict_esaes[hook_position_end[0].name].cfg.d_sae
d_sae_start = self.dict_saes[hook_position_start[0].name].cfg.d_sae
d_esae_start = self.dict_esaes[hook_position_start[0].name].cfg.d_sae
to_bwd_cache = {}
bwd_cache = {}
edge_effect = {}
with t.set_grad_enabled(True):
with self._detach_error_term(False, hook_position_end[0].name):
with self._setup_forward_model_hook(transfer_grad=False): # must be false, no gradient from intermediate component
with self._setup_fwd_bwd_edge_grad_sae_hook(
bwd_cache=bwd_cache,
to_bwd_cache=to_bwd_cache,
hook_position_start=hook_position_start,
hook_position_end=hook_position_end,
):
self.model.forward(clean_token, return_type=None, stop_at_layer=layer+1)
aggregate_dim = [0] if self.token_wise else [0, 1]
to_bwd = ( # (seq, d_sae+1) || (d_sae+1)
cache_to_sparseact(
to_bwd_cache,
sae_hook_name(hook_position_end[0].name),
sae_hook_name(error_term_name(hook_position_end[0].name)) if self.use_error_term else None,
error_term_name(error_term_name(hook_position_end[0].name)) if self.use_esae_error_term else None,
) * leftvec.detach()
).sum(aggregate_dim)
if self.use_esae_error_term:
to_bwd = to_bwd.contract()
to_bwd = to_bwd.to_tensor() # (seq, d_sae+d_esae+1) || (d_sae+d_esae+1)
del to_bwd_cache
for end_node, end_index in self.active_nodes(*hook_position_end):
if isinstance(end_index, ErrorIndex):
# the last index is error: shape (d_sae+d_esae+1) so last index is d_sae+d_esae
to_bwd[end_index.idx + (d_sae_end+d_esae_end,)].backward(retain_graph=True)
index = t.tensor(list(end_index.idx + (d_sae_end+d_esae_end,)), device=self.device)
elif isinstance(end_index, FeatureErrorIndex):
revised_index = list(end_index.idx)
revised_index[-1] += d_sae_end
to_bwd[tuple(revised_index)].backward(retain_graph=True)
index = t.tensor(revised_index, device=self.device)
elif isinstance(end_index, FeatureIndex):
to_bwd[end_index.idx].backward(retain_graph=True)
index = t.tensor(list(end_index.idx), device=self.device)
else:
raise ValueError(f"end_index of type {type(end_index)} is not supported.")
'''
edge_effect shape (seq, d_sae+d_esae+1, seq, d_sae+d_esae+1) or (d_sae+d_esae+1, d_sae+d_esae+1) in sparse_coo tensor
the sparse_coo will have the shape:
--> indices of shape (2, num_active) or (1, num_active)
--> values of shape (num_active, seq, d_sae+d_esae+1) or (num_active, d_sae+d_esae+1)
'''
effect = (
cache_to_sparseact(
bwd_cache,
sae_hook_name(hook_position_start[0].name),
sae_hook_name(error_term_name(hook_position_start[0].name)) if self.use_error_term else None,
error_term_name(error_term_name(hook_position_start[0].name)) if self.use_esae_error_term else None,
) * rightvec
).sum(aggregate_dim)
if self.use_esae_error_term:
effect.contract()
edge_effect[index] = effect.to_tensor() # (seq, d_sae+d_esae+1) || (d_sae+d_esae+1)
del bwd_cache
seq = int(self.seq_length) # type: ignore
num_end = d_sae_end
num_start = d_sae_start
if self.use_error_term:
num_end += d_esae_end
num_start += d_esae_start
if self.use_esae_error_term:
num_end += 1
num_start += 1
if len(edge_effect.keys()) != 0:
indices = t.stack(list(edge_effect.keys()), dim=0).T # shape (2, num_active) or (1, num_active)
values = t.stack([value for value in edge_effect.values()], dim=0) # shape (num_active, seq, d_sae+d_esae+1) or (num_active, d_sae+d_esae+1)
# if no active nodes, return empty tensor
else:
indices = t.empty((2, 0) if self.token_wise else (1, 0), dtype=t.long).to(self.device)
values = t.empty((0, seq, num_start) if self.token_wise else (0, num_start), dtype=t.float).to(self.device)
if self.token_wise:
return t.sparse_coo_tensor(indices, values, size=(seq, num_end, seq, num_start)).coalesce()
else:
return t.sparse_coo_tensor(indices, values, size=(num_end, num_start)).coalesce()
@contextmanager
def _setup_fwd_sae_hook(
self,
fwd_cache: Dict[str, Tensor],
corrupt_cache: ActivationCache | Dict[str, Tensor] | None,
patch_deleted_comp: bool = False,
use_error_term: bool | None = None,
):
use_esae_error_term = use_error_term if use_error_term is not None else self.use_esae_error_term
def hook_sae_fwd(act: Tensor, hook: HookPoint, sae_name: str) -> Tensor:
if hook.name == "hook_sae_acts_post" or hook.name == sae_hook_name(sae_name):
if patch_deleted_comp and corrupt_cache is not None:
act_mask = self.nodes[(sae_name, (None,))].act == 0 # type: ignore
act[:, act_mask] = corrupt_cache[sae_hook_name(sae_name)][:, act_mask]
fwd_cache[sae_hook_name(sae_name)] = act.detach()
elif (
hook.name == "hook_sae_error.hook_sae_acts_post" or hook.name == sae_hook_name(error_term_name(sae_name))
):
if patch_deleted_comp and corrupt_cache is not None:
act_mask = self.nodes[(sae_name, (None,))].res == 0 # type: ignore
act[:, act_mask] = corrupt_cache[sae_hook_name(error_term_name(sae_name))][:, act_mask]
fwd_cache[sae_hook_name(error_term_name(sae_name))] = act.detach()
elif use_esae_error_term and (
hook.name == "hook_sae_error.hook_sae_error" or hook.name == error_term_name(error_term_name(sae_name))
):
if patch_deleted_comp and corrupt_cache is not None:
resc_mask = self.nodes[(sae_name, (None,))].resc == 0 # type: ignore
act[:, resc_mask] = corrupt_cache[error_term_name(error_term_name(sae_name))][:, resc_mask]
fwd_cache[error_term_name(error_term_name(sae_name))] = act.detach()
return act
try:
with self._hook_esaes_to_saes(use_esae_error_term=use_esae_error_term):
with self._setup_error_term(True):
for sae_name, sae in self.dict_saes.items():
sae.add_hook(
lambda name: True,
partial(hook_sae_fwd, sae_name=sae_name),
dir="fwd",
)
yield
finally:
for sae in self.dict_saes.values():
sae.reset_hooks()
@contextmanager
def _setup_fwd_bwd_grad_sae_hook(
self,
fwd_cache: Dict[str, Tensor],
bwd_cache: Dict[str, Tensor],
pass_through_grad: bool,
):
def hook_sae_fwd(act: Tensor, hook: HookPoint, sae_name: str) -> Tensor:
if hook.name == "hook_sae_acts_post" or hook.name == sae_hook_name(sae_name):
fwd_cache[sae_hook_name(sae_name)] = act.detach()
elif self.use_error_term and (
hook.name == "hook_sae_error.hook_sae_acts_post" or hook.name == sae_hook_name(error_term_name(sae_name))
):
fwd_cache[sae_hook_name(error_term_name(sae_name))] = act.detach()
elif self.use_esae_error_term and (
hook.name == "hook_sae_error.hook_sae_error" or hook.name == error_term_name(error_term_name(sae_name))
):
fwd_cache[error_term_name(error_term_name(sae_name))] = act.detach()
return act
pass_through_cache = {}
def hook_sae_bwd(grad: Tensor, hook: HookPoint, sae_name: str) -> None:
if hook.name == "hook_sae_acts_post" or hook.name == sae_hook_name(sae_name):
bwd_cache[sae_hook_name(sae_name)] = grad.detach()
elif hook.name == "hook_sae_output" or hook.name == output_hook_name(sae_name):
# avoid pass through grad for transcoder
if pass_through_grad and sae_name not in self.output_hooks:
pass_through_cache[output_hook_name(sae_name)] = grad.detach()
elif self.use_error_term and (
hook.name == "hook_sae_error.hook_sae_acts_post" or hook.name == sae_hook_name(error_term_name(sae_name))
):
bwd_cache[sae_hook_name(error_term_name(sae_name))] = grad.detach()
elif self.use_esae_error_term and (
hook.name == "hook_sae_error.hook_sae_error" or hook.name == error_term_name(error_term_name(sae_name))
):
# the error_grad is the output_grad, due to: output = recon + error_term
bwd_cache[error_term_name(error_term_name(sae_name))] = grad.detach()
elif hook.name == "hook_sae_input" or hook.name == input_hook_name(sae_name):
if pass_through_grad:
# we have to modify inplace instead of grad = ... and then return grad
# because, returning a tensor in bwd pass hook is buggy somehow
if sae_name not in self.output_hooks:
# no pass through grad for transcoder
grad.copy_(pass_through_cache[output_hook_name(sae_name)])
else:
# avoid gradient flow to upstream from the transcoder
grad.zero_()
try:
with self._hook_esaes_to_saes(use_esae_error_term=self.use_esae_error_term):
with self._setup_error_term(self.use_error_term):
for sae_name, sae in self.dict_saes.items():
sae.add_hook(
lambda name: True,
partial(hook_sae_fwd, sae_name=sae_name),
dir="fwd",
)
sae.add_hook(
lambda name: True,
partial(hook_sae_bwd, sae_name=sae_name),
dir="bwd",
)
yield
finally:
for sae in self.dict_saes.values():
sae.reset_hooks()
@contextmanager
def _setup_fwd_bwd_edge_grad_sae_hook(
self,
bwd_cache: Dict[str, Tensor],
to_bwd_cache: Dict[str, Tensor],
hook_position_start: Tuple[Node, Index],
hook_position_end: Tuple[Node, Index],
):
# print(hook_position_end[0].name, hook_position_start[0].name)
def hook_sae_bwd(grad: Tensor, hook: HookPoint, sae_name: str, bwd_cache: Dict) -> None:
# print(sae_name, hook.name, t.sum(grad != 0).item())
if hook.name == "hook_sae_acts_post" or hook.name == sae_hook_name(sae_name):
if hook.name == sae_hook_name(hook_position_start[0].name) or _hook_name(sae_name, hook.name) == sae_hook_name(hook_position_start[0].name):
# only store the gradient at the start position
bwd_cache[sae_hook_name(sae_name)] = grad.detach()
elif self.use_error_term and (
hook.name == "hook_sae_error.hook_sae_acts_post" or hook.name == sae_hook_name(error_term_name(sae_name))
):
if hook.name == sae_hook_name(error_term_name(hook_position_start[0].name)) or _hook_name(sae_name, hook.name) == sae_hook_name(error_term_name(hook_position_start[0].name)):
# only store the gradient at the start position
bwd_cache[sae_hook_name(error_term_name(sae_name))] = grad.detach()
elif self.use_esae_error_term and (
hook.name == "hook_sae_error.hook_sae_output" or hook.name == output_hook_name(error_term_name(sae_name))
):
if hook.name == output_hook_name(error_term_name(hook_position_start[0].name)) or _hook_name(sae_name, hook.name) == output_hook_name(error_term_name(hook_position_start[0].name)):
# we have: output = recon + stop_grad(error_term)
# so, the TRUE error_grad (if not stop_grad) is output_grad
# due to: output = recon + error_term
# only store the gradient at the start position
bwd_cache[error_term_name(error_term_name(sae_name))] = grad.detach()
# IMPORTANT NOTE for reproducibility:
# we zero grad of intermediate components
# but zero grad the output hook of SAE will ALSO ZERO GRAD the resid mid / post --> no downstream grads
# so, we instead zero grad of the input hook of SAE
elif hook.name == "hook_sae_input" or hook.name == input_hook_name(sae_name):
if hook.name == input_hook_name(hook_position_end[0].name) or _hook_name(sae_name, hook.name) == input_hook_name(hook_position_end[0].name):
pass # no zero grad for end position
else:
grad.zero_()
def hook_sae_fwd(act: Tensor, hook: HookPoint, sae_name: str, to_bwd_cache: Dict) -> Tensor:
if hook.name == "hook_sae_acts_post" or hook.name == sae_hook_name(sae_name):
if hook.name == sae_hook_name(hook_position_end[0].name) or _hook_name(sae_name, hook.name) == sae_hook_name(hook_position_end[0].name):
# store activation for backward later
to_bwd_cache[sae_hook_name(sae_name)] = act
elif self.use_error_term and (
hook.name == "hook_sae_error.hook_sae_acts_post" or sae_hook_name(error_term_name(sae_name)) == hook.name
):
# if not at the end position, we intervene as usual
if hook.name == sae_hook_name(error_term_name(hook_position_end[0].name)) or _hook_name(sae_name, hook.name) == sae_hook_name(error_term_name(hook_position_end[0].name)):
to_bwd_cache[sae_hook_name(error_term_name(sae_name))] = act
elif self.use_esae_error_term and (
hook.name == "hook_sae_error.hook_sae_error" or error_term_name(error_term_name(sae_name)) == hook.name
):
# if not at the end position, we intervene as usual
if hook.name == error_term_name(error_term_name(hook_position_end[0].name)) or _hook_name(sae_name, hook.name) == error_term_name(error_term_name(hook_position_end[0].name)):
to_bwd_cache[error_term_name(error_term_name(sae_name))] = act
return act
try:
with self._hook_esaes_to_saes(use_esae_error_term=self.use_esae_error_term):
with self._setup_error_term(self.use_error_term):
for sae_name, sae in self.dict_saes.items():
sae.add_hook(
lambda name: True,
partial(hook_sae_fwd, sae_name=sae_name, to_bwd_cache=to_bwd_cache),
dir="fwd",
)
sae.add_hook(
lambda name: True,
partial(hook_sae_bwd, sae_name=sae_name, bwd_cache=bwd_cache),
dir="bwd",
)
yield
finally:
for sae in self.dict_saes.values():
sae.reset_hooks()
@contextmanager
def _setup_fwd_bwd_grad_sae_hook_ig(
self,
target_name: str,
frac: float,
fwd_cache: Dict[str, Tensor],
bwd_cache: Dict[str, Tensor],
corrupt_cache: ActivationCache | Dict[str, Tensor],
):
def hook_sae_fwd(act: Tensor, hook: HookPoint, sae_name: str, target_name: str, frac: float) -> Tensor:
if hook.name == "hook_sae_acts_post" or hook.name == sae_hook_name(sae_name):
# interpolate for integrated gradients
if hook.name == sae_hook_name(target_name) or _hook_name(sae_name, hook.name) == sae_hook_name(target_name):
act = interpolate(
corrupt_cache[sae_hook_name(sae_name)],
act,
frac,
)
fwd_cache[sae_hook_name(sae_name)] = act.detach()
elif self.use_error_term and (
hook.name == "hook_sae_error.hook_sae_acts_post" or hook.name == sae_hook_name(error_term_name(sae_name))
):
# interpolate for integrated gradients
if (
hook.name == sae_hook_name(error_term_name(target_name)) or
_hook_name(sae_name, hook.name) == sae_hook_name(error_term_name(target_name))
):
act = interpolate(
corrupt_cache[sae_hook_name(error_term_name(sae_name))],
act,
frac,
)
fwd_cache[sae_hook_name(error_term_name(sae_name))] = act.detach()
elif self.use_esae_error_term and (
hook.name == "hook_sae_error.hook_sae_error" or hook.name == error_term_name(error_term_name(sae_name))
):
if (
hook.name == error_term_name(error_term_name(target_name)) or
_hook_name(sae_name, hook.name) == error_term_name(error_term_name(target_name))
):
act = interpolate(
corrupt_cache[error_term_name(error_term_name(sae_name))],
act,
frac,
)
fwd_cache[error_term_name(error_term_name(sae_name))] = act.detach()
return act
def hook_sae_bwd(grad: Tensor, hook: HookPoint, sae_name: str) -> None:
if hook.name == "hook_sae_acts_post" or hook.name == sae_hook_name(sae_name):
create_list(bwd_cache, sae_hook_name(sae_name))
bwd_cache[sae_hook_name(sae_name)] += grad.detach()
elif self.use_error_term and (
hook.name == "hook_sae_error.hook_sae_acts_post" or hook.name == sae_hook_name(error_term_name(sae_name))
):
create_list(bwd_cache, sae_hook_name(error_term_name(sae_name)))
bwd_cache[sae_hook_name(error_term_name(sae_name))] += grad.detach()
elif self.use_esae_error_term and (
hook.name == "hook_sae_error.hook_sae_error" or hook.name == error_term_name(error_term_name(sae_name))
):
create_list(bwd_cache, error_term_name(error_term_name(sae_name)))
bwd_cache[error_term_name(error_term_name(sae_name))] += grad.detach()
try:
with self._hook_esaes_to_saes(use_esae_error_term=self.use_esae_error_term):
with self._setup_error_term(self.use_error_term):
for sae_name, sae in self.dict_saes.items():
sae.add_hook(
lambda name: True,
partial(hook_sae_fwd, sae_name=sae_name, target_name=target_name, frac=frac),
dir="fwd",
)
sae.add_hook(
lambda name: True,
partial(hook_sae_bwd, sae_name=sae_name),
dir="bwd",
)
yield
finally:
for sae in self.dict_saes.values():
sae.reset_hooks()
def run_model(
self,
toks: Tensor,
use_error_term: bool | None = None,
) -> Tuple[Tensor, ActivationCache]:
"""
Runs an MLP transcoder(s) on a batch of tokens.
"""
return Feature_Graph_Trans.run_model(
self,
toks=toks,
use_error_term=use_error_term
)
def _saes_to_list(self) -> List[Any]:
return Feature_Graph_Trans._saes_to_list(self)