Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
Show all changes
24 commits
Select commit Hold shift + click to select a range
355dd1a
minor
thiswillbeyourgithub Sep 3, 2025
131ad80
minor: perf
thiswillbeyourgithub Sep 3, 2025
1f0b5bb
feat: support for chat templates
thiswillbeyourgithub Sep 3, 2025
f9b509b
use loguru instead of warnings
thiswillbeyourgithub Sep 3, 2025
c0cc68f
feat: support for layer zones in addition to layer ids
thiswillbeyourgithub Sep 3, 2025
07418cc
import make_dataset from utils instead of defining it in tests.py
thiswillbeyourgithub Sep 3, 2025
b0f8bd3
test: add a test for make_dataset
thiswillbeyourgithub Sep 3, 2025
e2802a6
test: update test values not passing
thiswillbeyourgithub Sep 3, 2025
de746ba
doc: mention how to use chat templates
thiswillbeyourgithub Sep 3, 2025
95c89fe
doc: add a link related to OOM in transformers related to gguf
thiswillbeyourgithub Sep 3, 2025
200bb43
Merge branch 'main' into chat-templates-and-layer-zones
thiswillbeyourgithub Sep 3, 2025
29963f0
fix: typo in loguru dep
thiswillbeyourgithub Sep 5, 2025
6e6139c
fix: control layer zone
thiswillbeyourgithub Sep 5, 2025
0e2ea13
new: add more tqdm descriptions
thiswillbeyourgithub Sep 5, 2025
b03d055
actually in the readme we should omit autocorrect as it now is a fall…
thiswillbeyourgithub Sep 6, 2025
a1df843
fix: mapping layers had a wrong edgecase for zone starting at 0
thiswillbeyourgithub Sep 8, 2025
b36c92c
perf: only compute pre norm if useful
thiswillbeyourgithub Sep 7, 2025
5a79dd1
perf: slight improvement
thiswillbeyourgithub Sep 7, 2025
081318c
feat: support for mamba like models
thiswillbeyourgithub Sep 8, 2025
7b69172
feat: support for gemma 3
thiswillbeyourgithub Sep 8, 2025
690a398
doc: mention quantization can hurt some models
thiswillbeyourgithub Sep 8, 2025
cfb8807
doc: example should usea better example for layer zones
thiswillbeyourgithub Sep 8, 2025
046f44f
docfix: example in the readme was not the right one
thiswillbeyourgithub Sep 8, 2025
47484dc
revert to using warnings.warn instead of logger
thiswillbeyourgithub Sep 9, 2025
File filter

Filter by extension

Filter by extension


Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
83 changes: 69 additions & 14 deletions README.md
Original file line number Diff line number Diff line change
Expand Up @@ -16,41 +16,95 @@ import torch
from transformers import AutoModelForCausalLM, AutoTokenizer

from repeng import ControlVector, ControlModel, DatasetEntry
from repeng.utils import make_dataset

# load and wrap model
model_name = "mistralai/Mistral-7B-Instruct-v0.3"

# If you need quantization, but can lead to issues. For
# example Gemma3 models seem to silently generate empty strings.
# from transformers import BitsAndBytesConfig
# bnb_config = BitsAndBytesConfig(
# load_in_4bit=True,
# bnb_4bit_quant_type="nf4",
# bnb_4bit_compute_dtype=torch.bfloat16,
# bnb_4bit_use_double_quant=True,
# )

model = AutoModelForCausalLM.from_pretrained(
model_name,
# quantization_config=bnb_config,
# torch_dtype=torch.float16,
)
)

# load and wrap Mistral-7B
model_name = "mistralai/Mistral-7B-Instruct-v0.1"
model = AutoModelForCausalLM.from_pretrained(model_name, torch_dtype=torch.float16)
model = ControlModel(model, list(range(-5, -18, -1)))
# wrap the model to give us control
model = ControlModel(
model,
# layer_ids=list(range(-5, -18, -1)) # specify layers to control by layer ID
layer_zones=[[0.3, 0.5]], # control layers with relative depth in [0.3, 0.5[
)

def make_dataset(template: str, pos_personas: list[str], neg_personas: list[str], suffixes: list[str]):
# see notebooks/experiments.ipynb for a definition of `make_dataset`
...
tokenizer = AutoTokenizer.from_pretrained(
model_name,
# quantization_config=bnb_config,
)

# generate a dataset with closely-opposite paired statements
trippy_dataset = make_dataset(
"Act as if you're extremely {persona}.",
["high on psychedelic drugs"],
["sober from psychedelic drugs"],
truncated_output_suffixes,
# you can use either chat as dicts...
template=[
{"role": "system", "content": "You talk like you are {persona}."},
{"role": "user", "content": "{suffix}"},
],
# ...or directly strings:
# template="Act as if you're {persona}. Someone comes at you and says '{suffix}'.",

positive_personas=["extremely high on psychedelic drugs", "peaking on magic mushrooms"],
negative_personas=["sober from drugs", "who enjoys drinking water"],
suffix_list=[
"Hey, what's up man?",
"Hey, what's up girl?",
"Welcome Mr Musk, come this way.",
"How have you been feeling lately with the medications?",
],
)

# train the vector—takes less than a minute!
trippy_vector = ControlVector.train(model, tokenizer, trippy_dataset)

# Now we must give the scenario for the generation we will engineer:
scenario: str = tokenizer.apply_chat_template(
conversation=[
{
"role": "user",
"content": "Give me a one-sentence pitch for a TV show."
},
],
continue_final_message=False,
tokenize=False,
)

# Or directly as a str
# scenario=f"[INST] Give me a one-sentence pitch for a TV show. [/INST]",

# set the control strength and let inference rip!
for strength in (-2.2, 1, 2.2):
print(f"strength={strength}")
model.set_control(trippy_vector, strength)
out = model.generate(
**tokenizer(
f"[INST] Give me a one-sentence pitch for a TV show. [/INST]",
scenario,
return_tensors="pt"
),
).to(model.device),
do_sample=False,
max_new_tokens=128,
# temperature=1.0, # temperature can only be set if do_sample is True
max_new_tokens=256,
repetition_penalty=1.1,
)
print(tokenizer.decode(out.squeeze()).strip())
# or if you want to display the special tokens:
# print(tokenizer.decode(out.squeeze(), skip_special_tokens=False).strip())
print()
```

Expand All @@ -69,6 +123,7 @@ For a more detailed explanation of how the library works and what it can do, see

* For a list of changes by version, see the [CHANGELOG](https://github.com/vgel/repeng/blob/main/CHANGELOG).
* For quantized use, you may be interested in [llama.cpp#5970](https://github.com/ggerganov/llama.cpp/pull/5970)—after training a vector with `repeng`, export it by calling `vector.export_gguf(filename)` and then use it in `llama.cpp` with any quant!
* To load gguf files directly, you can run into OOM errors, see [this github issue for more](See here: https://github.com/huggingface/transformers/issues/34417).
* Vector training *currently does not work* with MoE models (such as Mixtral). (This is theoretically fixable with some work, let me know if you're interested.)
* Some example notebooks require `accelerate`, which must be manually installed with `pip install accelerate`. (This can also be done in the notebook with the IPython magic `%pip install accelerate`.)

Expand Down
3 changes: 3 additions & 0 deletions pyproject.toml
Original file line number Diff line number Diff line change
Expand Up @@ -12,8 +12,11 @@ dependencies = [
"transformers>=4.36.2",
"tqdm>=4.66.1",
"gguf>=0.13.0",
"loguru>=0.7.3",

]


[dependency-groups]
dev = ["pytest>=8.0.2", "ruff>=0.8.3"]

Expand Down
77 changes: 67 additions & 10 deletions repeng/control.py
Original file line number Diff line number Diff line change
Expand Up @@ -5,6 +5,8 @@
import torch
from transformers import PretrainedConfig, PreTrainedModel

from repeng.utils import get_num_hidden_layer

if typing.TYPE_CHECKING:
from .extract import ControlVector

Expand All @@ -16,20 +18,68 @@ class ControlModel(torch.nn.Module):
A wrapped language model that can have controls set on its layers with `self.set_control`.
"""

def __init__(self, model: PreTrainedModel, layer_ids: typing.Iterable[int]):
def __init__(
self,
model: PreTrainedModel,
layer_ids: typing.Optional[typing.Iterable[int]]=None,
layer_zones: typing.Optional[typing.Iterable[float]]=None,
):
"""
**This mutates the wrapped `model`! Be careful using `model` after passing it to this class.**

Build a new ControlModel around a model instance, initializing control on
the layers specified in `layer_ids`.
the layers specified in `layer_ids` or `layer_zones`.

To control layers #3 and 5, use layer_ids=[3,5].
To control layers by their relative depth, use layer_zones=[[0.1, 0.5]] to
control the layers with depth between 10% and 50% (left inclusive). you
can specify multiple zones but no overlapping nor empty zones are allowed.
"""

assert (layer_ids or layer_zones) and not (layer_ids and layer_zones), "Must supply either layer_ids or layer_zones argument"

super().__init__()
self.model = model

# Get the number of layers
layer_ids = list(range(get_num_hidden_layer(model)))
nlayers = len(layer_ids)

if layer_zones:
self.layer_ids = []
for start_zone, end_zone in layer_zones:
assert (
start_zone < end_zone
and start_zone >= 0
and start_zone <= 1
and end_zone >= 0
and end_zone <= 1
), "wrong layer_zones format"
if end_zone != 1.0:
new_layers = [
ilayer
for ilayer in layer_ids
if start_zone <= (ilayer / nlayers) < end_zone
]
else: # trick to make sure to include the last layers if desired
new_layers = [
ilayer
for ilayer in layer_ids
if start_zone <= (ilayer / nlayers)
]
assert new_layers, f"No layers found in zone {start_zone} to {end_zone}"
assert not any(nl in self.layer_ids for nl in new_layers), "Overlapping zones found"
self.layer_ids.extend(new_layers)
else:
# remap to make sure they are not negative
self.layer_ids = layer_ids

assert self.layer_ids, "No layers to control"

layers = model_layer_list(model)
self.layer_ids = [i if i >= 0 else len(layers) + i for i in layer_ids]
for layer_id in layer_ids:
assert len(layers) == len(layer_ids)

for layer_id in self.layer_ids:
layer = layers[layer_id]
if not isinstance(layer, ControlModule):
layers[layer_id] = ControlModule(layer)
Expand Down Expand Up @@ -164,7 +214,8 @@ def forward(self, *args, **kwargs):
assert len(control.shape) == len(modified.shape)
control = control.to(modified.device)

norm_pre = torch.norm(modified, dim=-1, keepdim=True)
if self.params.normalize:
norm_pre = torch.norm(modified, dim=-1, keepdim=True)

# we should ignore the padding tokens when doing the activation addition
# mask has ones for non padding tokens and zeros at padding tokens.
Expand All @@ -180,10 +231,10 @@ def forward(self, *args, **kwargs):
.reshape(target_shape[0], target_shape[1], 1)
)
mask = mask.to(modified.dtype).to(modified.device)
modified = self.params.operator(modified, control * mask)
else:
mask = 1.0
modified = self.params.operator(modified, control)

modified = self.params.operator(modified, control * mask)

if self.params.normalize:
norm_post = torch.norm(modified, dim=-1, keepdim=True)
Expand All @@ -201,9 +252,15 @@ def model_layer_list(model: ControlModel | PreTrainedModel) -> torch.nn.ModuleLi
if isinstance(model, ControlModel):
model = model.model

if hasattr(model, "model"): # mistral-like
return model.model.layers
if hasattr(model, "language_model"): # gemmma3 like
layers = model.language_model.layers
elif hasattr(model, "layers"): # qwen3-like
layers = model.layers
elif hasattr(model, "base_model"): # mamba like
layers = model.base_model.layers
elif hasattr(model, "transformer"): # gpt-2-like
return model.transformer.h
layers = model.transformer.h
elif hasattr(model, "model"): # mistral-like
layers = model.model.layers
else:
raise ValueError(f"don't know how to get layer list for {type(model)}")
23 changes: 10 additions & 13 deletions repeng/extract.py
Original file line number Diff line number Diff line change
Expand Up @@ -12,13 +12,9 @@

from .control import ControlModel, model_layer_list
from .saes import Sae
from .utils import DatasetEntry, get_model_name, autocorrect_chat_templates, get_num_hidden_layer


@dataclasses.dataclass
class DatasetEntry:
positive: str
negative: str


@dataclasses.dataclass
class ControlVector:
Expand Down Expand Up @@ -255,14 +251,16 @@ def read_representations(
Extract the representations based on the contrast dataset.
"""
if not hidden_layers:
hidden_layers = range(-1, -model.config.num_hidden_layers, -1)
hidden_layers = list(range(get_num_hidden_layer(model)))

# normalize the layer indexes if they're negative
n_layers = len(model_layer_list(model))
hidden_layers = [i if i >= 0 else n_layers + i for i in hidden_layers]

# the order is [positive, negative, positive, negative, ...]
train_strs = [s for ex in inputs for s in (ex.positive, ex.negative)]
train_strs = autocorrect_chat_templates(
messages=[s for ex in inputs for s in (ex.positive, ex.negative)],
tokenizer=tokenizer,
model=model,
)

layer_hiddens = batched_get_hiddens(
model, tokenizer, train_strs, hidden_layers, batch_size
Expand All @@ -273,7 +271,7 @@ def read_representations(

# get directions for each layer using PCA
directions: dict[int, np.ndarray] = {}
for layer in tqdm.tqdm(hidden_layers):
for layer in tqdm.tqdm(hidden_layers, desc="Altering directions"):
h = layer_hiddens[layer]
assert h.shape[0] == len(inputs) * 2

Expand Down Expand Up @@ -343,10 +341,9 @@ def batched_get_hiddens(
]
hidden_states = {layer: [] for layer in hidden_layers}
with torch.no_grad():
for batch in tqdm.tqdm(batched_inputs):
for batch in tqdm.tqdm(batched_inputs, desc="Computing activations"):
# get the last token, handling right padding if present
encoded_batch = tokenizer(batch, padding=True, return_tensors="pt")
encoded_batch = encoded_batch.to(model.device)
encoded_batch = tokenizer(batch, padding=True, return_tensors="pt").to(model.device)
out = model(**encoded_batch, output_hidden_states=True)
attention_mask = encoded_batch["attention_mask"]
for i in range(len(batch)):
Expand Down
2 changes: 1 addition & 1 deletion repeng/saes.py
Original file line number Diff line number Diff line change
Expand Up @@ -71,7 +71,7 @@ def decode(self, features: np.ndarray) -> np.ndarray:
huggingface_hub.snapshot_download(repo_id, revision=revision)
)
layer_dict: dict[int, SaeLayer] = {}
for layer in tqdm.tqdm(layers):
for layer in tqdm.tqdm(layers, desc="Creating SAE"):
eleuther_layer = layer - 1 # see docstr
# this is in `sae` but to load the dtype we want, need to reimpl some stuff
layer_path = base_path / f"layers.{eleuther_layer}"
Expand Down
Loading