AQIT 0.1.0
Loading...
Searching...
No Matches
aquin.compute.layer_analysis Namespace Reference

Functions

tuple[list[str], bool] _resolve_stability_prompts (list[str]|None prompts)
tuple[float, float] _pca_variance_ratios (torch.Tensor acts, int top_k)
dict[int, torch.Tensor] _collect_layer_activations (HookedTransformer model, list[str] prompts)
dict run_activation_stability (HookedTransformer model, list[str]|None prompts=None, int top_k=10)
float _mean_pairwise_cos (torch.Tensor vecs)
float _mean_cross_cos (torch.Tensor a, torch.Tensor b)
float _rbf_mmd (torch.Tensor x, torch.Tensor y)
dict run_ood_similarity (HookedTransformer model, list[str]|None in_domain_prompts=None, list[str]|None ood_prompts=None)
dict run_layer_analysis (HookedTransformer model, dict args)

Variables

list DEFAULT_STABILITY_PROMPTS
list DEFAULT_IN_DOMAIN_PROMPTS
list DEFAULT_OOD_PROMPTS
float COLLAPSE_THRESHOLD = 0.85
float DEAD_THRESHOLD = 0.01
int MIN_STABILITY_PROMPTS = 2

Function Documentation

◆ _collect_layer_activations()

dict[int, torch.Tensor] _collect_layer_activations ( HookedTransformer model,
list[str] prompts )
protected
Last-token resid_post per layer. layer -> (n_prompts, d_model).

Definition at line 78 of file layer_analysis.py.

Referenced by run_activation_stability(), and run_ood_similarity().

◆ _mean_cross_cos()

float _mean_cross_cos ( torch.Tensor a,
torch.Tensor b )
protected

Definition at line 160 of file layer_analysis.py.

Referenced by run_ood_similarity().

◆ _mean_pairwise_cos()

float _mean_pairwise_cos ( torch.Tensor vecs)
protected

Definition at line 151 of file layer_analysis.py.

Referenced by run_ood_similarity().

◆ _pca_variance_ratios()

tuple[float, float] _pca_variance_ratios ( torch.Tensor acts,
int top_k )
protected
acts: (n_samples, d_model). Returns (top1_ratio, topk_ratio).

Definition at line 59 of file layer_analysis.py.

Referenced by run_activation_stability().

◆ _rbf_mmd()

float _rbf_mmd ( torch.Tensor x,
torch.Tensor y )
protected

Definition at line 168 of file layer_analysis.py.

Referenced by run_ood_similarity().

◆ _resolve_stability_prompts()

tuple[list[str], bool] _resolve_stability_prompts ( list[str] | None prompts)
protected
PCA needs ≥2 samples. Pad with defaults when the user passes fewer.

Definition at line 44 of file layer_analysis.py.

Referenced by run_activation_stability().

◆ run_activation_stability()

dict run_activation_stability ( HookedTransformer model,
list[str] | None prompts = None,
int top_k = 10 )

◆ run_layer_analysis()

dict run_layer_analysis ( HookedTransformer model,
dict args )

Definition at line 234 of file layer_analysis.py.

References run_activation_stability(), and run_ood_similarity().

◆ run_ood_similarity()

dict run_ood_similarity ( HookedTransformer model,
list[str] | None in_domain_prompts = None,
list[str] | None ood_prompts = None )

Variable Documentation

◆ COLLAPSE_THRESHOLD

float aquin.compute.layer_analysis.COLLAPSE_THRESHOLD = 0.85

Definition at line 39 of file layer_analysis.py.

◆ DEAD_THRESHOLD

float aquin.compute.layer_analysis.DEAD_THRESHOLD = 0.01

Definition at line 40 of file layer_analysis.py.

◆ DEFAULT_IN_DOMAIN_PROMPTS

list aquin.compute.layer_analysis.DEFAULT_IN_DOMAIN_PROMPTS
Initial value:
= [
"The capital of France is Paris.",
"Water boils at 100 degrees Celsius.",
"Python is a programming language.",
"The Earth orbits the Sun.",
]

Definition at line 25 of file layer_analysis.py.

◆ DEFAULT_OOD_PROMPTS

list aquin.compute.layer_analysis.DEFAULT_OOD_PROMPTS
Initial value:
= [
"asdf jkl qwerty zxcv",
"!!! ??? ### @@@",
"florp sniggle wumpus",
"9918273645 192837465",
]

Definition at line 32 of file layer_analysis.py.

◆ DEFAULT_STABILITY_PROMPTS

list aquin.compute.layer_analysis.DEFAULT_STABILITY_PROMPTS
Initial value:
= [
"The capital of France is",
"Write a function that sorts a list",
"Explain photosynthesis in simple terms",
"Hello, how are you today?",
"The quick brown fox jumps over the lazy dog",
]

Definition at line 17 of file layer_analysis.py.

◆ MIN_STABILITY_PROMPTS

int aquin.compute.layer_analysis.MIN_STABILITY_PROMPTS = 2

Definition at line 41 of file layer_analysis.py.