feat(ara): add methods for obtaining module I/O

This commit is contained in:
Philipp Emanuel Weidmann
2026-02-25 17:44:46 +05:30
parent 27097bfe8e
commit ea7c59a55a
2 changed files with 139 additions and 1 deletions
+13
View File
@@ -411,6 +411,19 @@ def run():
evaluator.get_score() evaluator.get_score()
return return
def tensor_shape_repr(self: torch.Tensor):
return f"tensor(shape={tuple(self.shape)}, dtype={self.dtype}, device={self.device})"
torch.Tensor.__repr__ = tensor_shape_repr # ty:ignore[invalid-assignment]
print()
print("Obtaining module I/O for good prompts...")
good_module_io = model.get_module_io_batched(good_prompts)
print("Obtaining module I/O for bad prompts...")
bad_module_io = model.get_module_io_batched(bad_prompts)
print(good_module_io)
print() print()
print("Calculating per-layer refusal directions...") print("Calculating per-layer refusal directions...")
print("* Obtaining residuals for good prompts...") print("* Obtaining residuals for good prompts...")
+126 -1
View File
@@ -4,7 +4,7 @@
import math import math
from contextlib import suppress from contextlib import suppress
from dataclasses import dataclass from dataclasses import dataclass
from typing import Any, Type, cast from typing import Any, Callable, Type, cast
import bitsandbytes as bnb import bitsandbytes as bnb
import torch import torch
@@ -14,6 +14,7 @@ from peft import LoraConfig, PeftModel, get_peft_model
from peft.tuners.lora.layer import Linear from peft.tuners.lora.layer import Linear
from torch import FloatTensor, LongTensor, Tensor from torch import FloatTensor, LongTensor, Tensor
from torch.nn import Module, ModuleList from torch.nn import Module, ModuleList
from torch.utils.hooks import RemovableHandle
from transformers import ( from transformers import (
AutoModelForCausalLM, AutoModelForCausalLM,
AutoModelForImageTextToText, AutoModelForImageTextToText,
@@ -653,6 +654,130 @@ class Model:
return torch.cat(residuals, dim=0) return torch.cat(residuals, dim=0)
def get_module_io(
self,
prompts: list[Prompt],
) -> list[dict[str, dict[int, tuple[Tensor, Tensor]]]]:
# The list contains one element per layer.
# Each element maps from the component name to a (possibly sparse) mapping
# from the module index to an (input, output) tuple containing the I/O
# tensors of shape (prompt, component).
module_io: list[dict[str, dict[int, tuple[Tensor, Tensor]]]] = []
def get_hook(
layer_index: int,
component: str,
module_index: int,
) -> Callable[[Module, tuple[Tensor, ...], Tensor], None]:
def hook(
module: Module,
inputs: tuple[Tensor, ...],
outputs: Tensor,
) -> None:
if len(module_io) == layer_index:
# First invocation of the hook for this layer.
module_io.append({})
# Layers are invoked in order during inference,
# so this should always hold.
assert len(module_io) == layer_index + 1
if component not in module_io[layer_index]:
module_io[layer_index][component] = {}
# Each module should be invoked at most once per inference step.
assert module_index not in module_io[layer_index][component]
# inputs[0] and outputs have shape (prompt, position, component),
# so this extracts the input/output at the end of each prompt.
input = inputs[0][:, -1, :].detach()
output = outputs[:, -1, :].detach()
# The modules associated with a component (e.g. expert MLPs)
# are not necessarily invoked in order, nor are all of them
# necessarily invoked in each inference step, so we cannot
# use a list here.
module_io[layer_index][component][module_index] = (input, output)
return hook
hook_handles: list[RemovableHandle] = []
for layer_index in range(len(self.get_layers())):
for component, modules in self.get_layer_modules(layer_index).items():
for module_index, module in enumerate(modules):
hook_handles.append(
module.register_forward_hook(
get_hook(layer_index, component, module_index)
)
)
self.generate(prompts, max_new_tokens=1)
for hook_handle in hook_handles:
hook_handle.remove()
return module_io
def get_module_io_batched(
self,
prompts: list[Prompt],
) -> list[dict[str, dict[int, tuple[Tensor, Tensor]]]]:
# Aggregating batch results is more complicated for module I/O
# than for other get_*_batched methods, because the structure of the results
# might differ between batches, as whether individual modules activate
# can depend on the prompt (in particular for MoE models).
# In practice, inhomogeneous results should be very rare, but to be fully
# generic, this logic is required.
module_io_batches = [
self.get_module_io(batch)
for batch in batchify(prompts, self.settings.batch_size)
]
module_io: list[dict[str, dict[int, tuple[Tensor, Tensor]]]] = []
for layer_index in range(len(self.get_layers())):
module_io.append({})
for module_io_batch in module_io_batches:
for component, io_map in module_io_batch[layer_index].items():
if component not in module_io[layer_index]:
module_io[layer_index][component] = {}
for module_index in io_map:
if module_index not in module_io[layer_index][component]:
# This is a placeholder; the actual aggregation happens below.
# We need to iterate over the batches twice because we don't
# know in advance which components and module indices are present.
module_io[layer_index][component][module_index] = (
torch.empty(0),
torch.empty(0),
)
for component, io_map in module_io[layer_index].items():
for module_index in io_map:
inputs_outputs = [
module_io_batch[layer_index][component][module_index]
for module_io_batch in module_io_batches
if component in module_io_batch[layer_index]
and module_index in module_io_batch[layer_index][component]
]
input = torch.cat(
[input_output[0] for input_output in inputs_outputs],
dim=0,
)
output = torch.cat(
[input_output[1] for input_output in inputs_outputs],
dim=0,
)
# The key already exists, and replacing existing values
# in a dictionary while iterating over the same dictionary
# is safe in Python.
module_io[layer_index][component][module_index] = (input, output)
return module_io
# We work with logprobs rather than probabilities for numerical stability # We work with logprobs rather than probabilities for numerical stability
# when computing the KL divergence. # when computing the KL divergence.
def get_logprobs(self, prompts: list[Prompt]) -> Tensor: def get_logprobs(self, prompts: list[Prompt]) -> Tensor: