Minimal, efficient covariance-based engram extraction for editing HuggingFace LLMs
Project description
ai-engram
Closed-form, covariance-based engram extraction for editing HuggingFace LLMs — forward-only, no gradient descent.
An engram is the slice of a layer's weights attributable to a target set of inputs. ai-engram isolates it analytically:
W_engram = W · Σ_target · pinv(Σ_total)
Σ_target and Σ_total are input covariances over the forget set and the reference set. Subtracting it — W ← W − α·W_engram — removes that knowledge while keeping the rest: fast, training-free unlearning / model editing.
- Closed-form — one pseudo-inverse per layer; no optimization loop, no labels.
- Forward-only — covariances via forward pre-hooks; no backprop.
- HF-native — Llama, Mistral, Qwen, Gemma, Phi … and GPT-2 (
Conv1D) out of the box. - Affine-correct — automatic bias absorption for bias-bearing layers.
Milestone 1 (this release): statistics collection + engram extraction. Applying the edit, a one-call
edit_llmhelper, adaptive scaling, registries, and metrics come in later milestones — and the extraction already reproduces TOFU unlearning (see Validation).
Install
pip install ai-engram
Pulls torch, tqdm, and transformers — HF LLMs and GPT-2 work out of the box. Distribution name ai-engram; import name engram.
📖 Documentation: https://jeakwon.github.io/ai-engram/
Quickstart
Any nn.Linear (or GPT-2 Conv1D) model:
import torch
from engram import EngramEditor, EditorConfig
editor = EngramEditor(model, EditorConfig())
target_cov = editor.collect_statistics(forget_loader) # Σ over data to isolate
total_cov = editor.collect_statistics(total_loader) # Σ over the reference set
weight_engrams, bias_engrams = editor.compute_engram_weights(target_cov, total_cov)
# weight_engrams[name] matches the layer's .weight; bias_engrams[name] its .bias
HuggingFace LLM (answer-token masked)
from engram import EngramEditor, EditorConfig
editor = EngramEditor(model, EditorConfig())
batch_fn = lambda b: {"input_ids": b["input_ids"], "attention_mask": b["attention_mask"]}
mask_fn = lambda b: b["labels"] != -100 # covariance over answer tokens only
g_forget = editor.collect_statistics(forget_loader, batch_fn=batch_fn, mask_fn=mask_fn)
g_total = editor.collect_statistics(total_loader, batch_fn=batch_fn, mask_fn=mask_fn)
weight_engrams, _ = editor.compute_engram_weights(g_forget, g_total)
# apply — Milestone 2 will expose this as editor.edit(...)
import copy
edited = copy.deepcopy(model)
mods = dict(edited.named_modules())
for name, w in weight_engrams.items():
mods[name].weight.data -= (0.6 * w).to(mods[name].weight.dtype)
Restrict the edit to specific modules with target_modules — the same convention
as LoRA/PEFT (["down_proj"] by name suffix, or a regex string), plus
layers_to_transform for decoder-layer indices. See the
Quickstart guide for details.
Mixture-of-experts. Answer-token masking reaches the experts automatically on
transformers <5; on transformers ≥5 (fused experts) opt in to the
detachable engram.moe adapter — EngramEditor(model, adapters=[FusedExpertAdapter()]) —
covering ~35 fused MoE architectures (Mixtral, Qwen2/3/3.5-MoE, DeepSeek-V3,
GLM4-MoE, MiniMax, Mistral4, OLMoE, Phi-MoE, …).
How it works
| step | what | cost |
|---|---|---|
| 1. collect | forward pre-hooks accumulate Σ = Σ xᵀx per layer |
one forward pass, no backward |
| 2. compute | W_engram = W · Σ_target · pinv(Σ_total) |
one pseudo-inverse per layer |
| 3. apply (M2) | W ← W − α·W_engram |
a single subtraction |
Efficient by construction — forward-only hooks, in-place accumulation, CPU/GPU split (covariances on storage_device), a float32 solve cast back to the model dtype, and answer-token masking. Handles nn.Linear, GPT-2 Conv1D (a transposed linear), and masked variants; full details in the Guide.
Configuration (EditorConfig)
| field | default | purpose |
|---|---|---|
storage_device |
model's device | where covariances are held; set "cpu" if the D×D matrices don't fit in VRAM (large models) |
absorb_bias |
True |
absorb bias into the edit for bias-bearing layers |
Validation
On TOFU forget10 with tofu_Llama-3.2-1B-Instruct, the engram extraction reproduces the paper's official 14-metric Overall within ~0.01:
| condition | ai-engram | paper |
|---|---|---|
| gold (retain90) | 0.998 | 0.998 |
| plain (α=0.6) | 0.706 | 0.698 |
| adaptive-norm (α=1.0, p=1) | 0.817 | 0.818 |
Answer-token NLL confirms strong, selective forgetting — the forget set's NLL jumps ~16× while retain is preserved, and adaptive-norm beats plain on both axes. Runnable end-to-end in
tests/ and
examples/; see the TOFU page.
API
collect_statistics(loader, target_modules=None, batch_fn=None, mask_fn=None, layers_to_transform=None) -> {name: Σ}compute_engram_weights(target_cov, total_cov) -> (weight_engrams, bias_engrams)merge_statistics(*stats)·save_statistics(stats, path)·load_statistics(path)
Full reference (auto-generated from docstrings): API docs.
License
MIT © Jeakwon Kim
Project details
Download files
Download the file for your platform. If you're not sure which to choose, learn more about installing packages.
Source Distribution
Built Distribution
Filter files by name, interpreter, ABI, and platform.
If you're not sure about the file name format, learn more about wheel file names.
Copy a direct link to the current filters
File details
Details for the file ai_engram-0.5.0.tar.gz.
File metadata
- Download URL: ai_engram-0.5.0.tar.gz
- Upload date:
- Size: 34.5 kB
- Tags: Source
- Uploaded using Trusted Publishing? No
- Uploaded via: twine/6.2.0 CPython/3.10.19
File hashes
| Algorithm | Hash digest | |
|---|---|---|
| SHA256 |
2cccc212af661559c053e7df6fa8d331bddd6fac978e9bf77d8a95bc95200e29
|
|
| MD5 |
cf9e417eda72d8ccf4e2193e0845e8b4
|
|
| BLAKE2b-256 |
1f333344c2778155e4d2e68c56e26cc612e5df684587b413f06d1efe1f0dd7e7
|
File details
Details for the file ai_engram-0.5.0-py3-none-any.whl.
File metadata
- Download URL: ai_engram-0.5.0-py3-none-any.whl
- Upload date:
- Size: 19.3 kB
- Tags: Python 3
- Uploaded using Trusted Publishing? No
- Uploaded via: twine/6.2.0 CPython/3.10.19
File hashes
| Algorithm | Hash digest | |
|---|---|---|
| SHA256 |
29cddef44efa42b4b2904c4806021cbca1ab2953c7bb51e67f21f7fad38f07de
|
|
| MD5 |
215d6451bda38f6ff7621c281db881e8
|
|
| BLAKE2b-256 |
6badd18a6c87ee2a7883acdc1e34d15f2a83469a9b7e24396bf1068242d6d827
|