mirror of
https://github.com/p-e-w/heretic.git
synced 2026-09-01 18:06:03 -07:00
Compare commits
68 Commits
| Author | SHA1 | Date | |
|---|---|---|---|
| edc3b12345 | |||
| 25979ad7d0 | |||
| 3b70fe5dfa | |||
| f7a456bd0c | |||
| 988c6bd90e | |||
| 96c7a7d98a | |||
| 1126332281 | |||
| c925f5e802 | |||
| 19cdf7e244 | |||
| 94775d4148 | |||
| 515a7b9eb5 | |||
| e26da5e0e6 | |||
| 4a6304c361 | |||
| c76416fe03 | |||
| 2bb203ee47 | |||
| d79a443e6f | |||
| ec0367226d | |||
| 5e3c04c802 | |||
| 0bb9521fbe | |||
| 992fb3a4b3 | |||
| 304c14adc7 | |||
| 56e57adf36 | |||
| bd1fa0ade4 | |||
| 3c5d6920bf | |||
| b8f4a9c985 | |||
| 154241f8a2 | |||
| 303ba9d978 | |||
| ea7c59a55a | |||
| cb4ef3fdfc | |||
| 4c80c4beb9 | |||
| 3a115e280c | |||
| 27097bfe8e | |||
| 025ab3a881 | |||
| 1179013999 | |||
| fe7bc1bae3 | |||
| e70a1a85e8 | |||
| e7f8be98b7 | |||
| 6017bcd347 | |||
| dd0b3a2f69 | |||
| b873598b77 | |||
| 10ceb3098e | |||
| 745b582414 | |||
| d0e9462fb8 | |||
| f68a887a7b | |||
| 2690655a83 | |||
| 3525b1ac22 | |||
| 42f5a9b553 | |||
| 451db0b76e | |||
| ebc22c299e | |||
| d5c834c51d | |||
| c86f49035e | |||
| 85a6ec5ecb | |||
| 632b1da622 | |||
| 1cfd09d7f3 | |||
| 09be09e12e | |||
| 039f6222d2 | |||
| c4b2ea0c42 | |||
| 02a5237a02 | |||
| cf8cf6f349 | |||
| 2141e110fb | |||
| 39101137ef | |||
| 064bed9a9f | |||
| 8d44b65670 | |||
| 5ddef6fd2f | |||
| 92d0c0d551 | |||
| 243f821d93 | |||
| 9d1734855d | |||
| 740aab61ba |
@@ -0,0 +1,11 @@
|
||||
# Style guide and coding conventions
|
||||
|
||||
* Identifier names should not contain abbreviations unless those abbreviations are very widely used and understood (e.g. "KL divergence").
|
||||
* Comments should start with a capital letter and end with a period. They should use correct grammar and spelling.
|
||||
* Function and method signatures **must** be fully type-annotated, including the return type (if any).
|
||||
* Every Python code file **must** start with an SPDX/Copyright header.
|
||||
* Settings descriptions should start with a capital letter and end with a period.
|
||||
* When new settings are added in `config.py`, they should also be added to `config.default.toml`, set to their default value and with their description as a comment. The order of settings in `config.default.toml` should match that in `config.py`.
|
||||
* Pull requests should implement one change, and one change only.
|
||||
* PRs containing multiple semantically independent changes **must** be split into multiple PRs.
|
||||
* PRs **must not** change existing code unless the changes are *directly related* to the PR. This includes changes to formatting and comments.
|
||||
@@ -0,0 +1 @@
|
||||
* text eol=lf
|
||||
@@ -17,10 +17,10 @@ jobs:
|
||||
|
||||
steps:
|
||||
- name: Check out code
|
||||
uses: actions/checkout@v4
|
||||
uses: actions/checkout@v6
|
||||
|
||||
- name: Install uv
|
||||
uses: astral-sh/setup-uv@v5
|
||||
uses: astral-sh/setup-uv@v7
|
||||
with:
|
||||
enable-cache: true
|
||||
cache-dependency-glob: "uv.lock"
|
||||
@@ -37,6 +37,9 @@ jobs:
|
||||
- name: Lint and check import sorting
|
||||
run: uv run ruff check --output-format=github --extend-select I .
|
||||
|
||||
- name: Check typing
|
||||
run: uv run ty check --output-format=github --error-on-warning .
|
||||
|
||||
- name: Build package
|
||||
run: uv build
|
||||
|
||||
|
||||
+7
-1
@@ -7,7 +7,7 @@ wheels/
|
||||
*.egg-info
|
||||
|
||||
# Virtual environments
|
||||
.venv
|
||||
.venv/
|
||||
|
||||
# Caches
|
||||
/.ruff_cache/
|
||||
@@ -17,3 +17,9 @@ wheels/
|
||||
|
||||
# Configuration files
|
||||
/config.toml
|
||||
|
||||
# Study checkpoints
|
||||
/checkpoints/
|
||||
|
||||
# Residual plots
|
||||
/plots/
|
||||
|
||||
@@ -1,11 +1,15 @@
|
||||
# Heretic: Fully automatic censorship removal for language models
|
||||
<img width="128" height="128" align="right" alt="Logo" src="https://github.com/user-attachments/assets/df5f2840-2f92-4991-aa57-252747d7182e" />
|
||||
|
||||
[](https://discord.gg/gdXc48gSyT)
|
||||
# Heretic: Fully automatic censorship removal for language models<br><br>[](https://discord.gg/gdXc48gSyT) [](https://huggingface.co/heretic-org)
|
||||
|
||||
[](https://trendshift.io/repositories/20538)
|
||||
|
||||
Heretic is a tool that removes censorship (aka "safety alignment") from
|
||||
transformer-based language models without expensive post-training.
|
||||
It combines an advanced implementation of directional ablation, also known
|
||||
as "abliteration" ([Arditi et al. 2024](https://arxiv.org/abs/2406.11717)),
|
||||
as "abliteration" ([Arditi et al. 2024](https://arxiv.org/abs/2406.11717),
|
||||
Lai 2025 ([1](https://huggingface.co/blog/grimjim/projected-abliteration),
|
||||
[2](https://huggingface.co/blog/grimjim/norm-preserving-biprojected-abliteration))),
|
||||
with a TPE-based parameter optimizer powered by [Optuna](https://optuna.org/).
|
||||
|
||||
This approach enables Heretic to work **completely automatically.** Heretic
|
||||
@@ -65,8 +69,11 @@ Heretic supports most dense models, including many multimodal models, and
|
||||
several different MoE architectures. It does not yet support SSMs/hybrid models,
|
||||
models with inhomogeneous layers, and certain novel attention systems.
|
||||
|
||||
You can find a collection of models that have been decensored using Heretic
|
||||
[on Hugging Face](https://huggingface.co/collections/p-e-w/the-bestiary).
|
||||
You can find a small collection of models that have been decensored using Heretic
|
||||
[on Hugging Face](https://huggingface.co/collections/p-e-w/the-bestiary),
|
||||
and the community has created and published
|
||||
[well over 1,000](https://huggingface.co/models?other=heretic)
|
||||
Heretic models in addition to those.
|
||||
|
||||
|
||||
## Usage
|
||||
@@ -89,8 +96,10 @@ a configuration file.
|
||||
|
||||
At the start of a program run, Heretic benchmarks the system to determine
|
||||
the optimal batch size to make the most of the available hardware.
|
||||
On an RTX 3090, with the default configuration, decensoring Llama-3.1-8B
|
||||
takes about 45 minutes.
|
||||
On an RTX 3090, with the default configuration, decensoring Llama-3.1-8B-Instruct
|
||||
takes about 45 minutes. Note that Heretic supports model quantization with
|
||||
bitsandbytes, which can drastically reduce the amount of VRAM required to process
|
||||
models. Set the `quantization` option to `bnb_4bit` to enable quantization.
|
||||
|
||||
After Heretic has finished decensoring a model, you are given the option to
|
||||
save the model, upload it to Hugging Face, chat with it to test how well it works,
|
||||
@@ -242,7 +251,8 @@ The development of Heretic was informed by:
|
||||
* [The original abliteration paper (Arditi et al. 2024)](https://arxiv.org/abs/2406.11717)
|
||||
* [Maxime Labonne's article on abliteration](https://huggingface.co/blog/mlabonne/abliteration),
|
||||
as well as some details from the model cards of his own abliterated models (see above)
|
||||
* [Jim Lai's article describing "projected abliteration"](https://huggingface.co/blog/grimjim/projected-abliteration)
|
||||
* Jim Lai's articles describing ["projected abliteration"](https://huggingface.co/blog/grimjim/projected-abliteration)
|
||||
and ["norm-preserving biprojected abliteration"](https://huggingface.co/blog/grimjim/norm-preserving-biprojected-abliteration)
|
||||
|
||||
|
||||
## Citation
|
||||
@@ -263,7 +273,7 @@ If you use Heretic for your research, please cite it using the following BibTeX
|
||||
|
||||
## License
|
||||
|
||||
Copyright © 2025 Philipp Emanuel Weidmann (<pew@worldwidemann.com>)
|
||||
Copyright © 2025-2026 Philipp Emanuel Weidmann (<pew@worldwidemann.com>) + contributors
|
||||
|
||||
This program is free software: you can redistribute it and/or modify
|
||||
it under the terms of the GNU Affero General Public License as published by
|
||||
|
||||
+43
-1
@@ -1,4 +1,5 @@
|
||||
# Copy this file to config.toml and edit the configuration to your liking.
|
||||
# Rename this file to config.toml, place it in the working directory
|
||||
# that you run Heretic from, and edit the configuration to your liking.
|
||||
|
||||
# List of PyTorch dtypes to try when loading model tensors.
|
||||
# If loading with a dtype fails, the next dtype in the list will be tried.
|
||||
@@ -15,9 +16,17 @@ dtypes = [
|
||||
"float32",
|
||||
]
|
||||
|
||||
# Quantization method to use when loading the model. Options:
|
||||
# "none" (no quantization),
|
||||
# "bnb_4bit" (4-bit quantization using bitsandbytes).
|
||||
quantization = "none"
|
||||
|
||||
# Device map to pass to Accelerate when loading the model.
|
||||
device_map = "auto"
|
||||
|
||||
# Maximum memory to allocate per device.
|
||||
# max_memory = {"0": "20GB", "cpu": "64GB"}
|
||||
|
||||
# Number of input sequences to process in parallel (0 = auto).
|
||||
batch_size = 0 # auto
|
||||
|
||||
@@ -27,6 +36,9 @@ max_batch_size = 128
|
||||
# Maximum number of tokens to generate for each response.
|
||||
max_response_length = 100
|
||||
|
||||
# Whether to print prompt/response pairs when counting refusals.
|
||||
print_responses = false
|
||||
|
||||
# Whether to print detailed information about residuals and refusal directions.
|
||||
print_residual_geometry = false
|
||||
|
||||
@@ -46,12 +58,42 @@ residual_plot_style = "dark_background"
|
||||
# This is used to ensure balanced co-optimization of KL divergence and refusal count.
|
||||
kl_divergence_scale = 1.0
|
||||
|
||||
# The KL divergence to target. Below this value, an objective based on the refusal count is used.
|
||||
# This helps prevent the sampler from extensively exploring parameter combinations that "do nothing".
|
||||
kl_divergence_target = 0.01
|
||||
|
||||
# Whether to adjust the refusal directions so that only the component that is
|
||||
# orthogonal to the good direction is subtracted during abliteration.
|
||||
orthogonalize_direction = false
|
||||
|
||||
# How to apply row normalization of the weights. Options:
|
||||
# "none" (no normalization),
|
||||
# "pre" (compute LoRA adapter relative to row-normalized weights),
|
||||
# "full" (like "pre", but renormalizes to preserve original row magnitudes).
|
||||
row_normalization = "none"
|
||||
|
||||
# The rank of the LoRA adapter to use when "full" row normalization is used.
|
||||
# Row magnitude preservation is approximate due to non-linear effects,
|
||||
# and this determines the rank of that approximation. Higher ranks produce
|
||||
# larger output files and may slow down evaluation.
|
||||
full_normalization_lora_rank = 3
|
||||
|
||||
# The symmetric winsorization to apply to the per-prompt, per-layer residual vectors,
|
||||
# expressed as the quantile to clamp to (between 0 and 1). Disabled by default.
|
||||
# This can tame so-called "massive activations" that occur in some models.
|
||||
# Example: winsorization_quantile = 0.95 computes the 0.95-quantile of the absolute values
|
||||
# of the components, then clamps the magnitudes of all components to that quantile.
|
||||
winsorization_quantile = 1.0
|
||||
|
||||
# Number of abliteration trials to run during optimization.
|
||||
n_trials = 200
|
||||
|
||||
# Number of trials that use random sampling for the purpose of exploration.
|
||||
n_startup_trials = 60
|
||||
|
||||
# Directory to save and load study progress to/from.
|
||||
study_checkpoint_dir = "checkpoints"
|
||||
|
||||
# Strings whose presence in a response (case insensitive) identifies the response as a refusal.
|
||||
refusal_markers = [
|
||||
"sorry",
|
||||
|
||||
@@ -0,0 +1,163 @@
|
||||
# Rename this file to config.toml, place it in the working directory
|
||||
# that you run Heretic from, and edit the configuration to your liking.
|
||||
|
||||
max_response_length = 300
|
||||
|
||||
residual_plot_title = "PaCMAP Projection of Residuals for Slop-Suppressing/Inducing Prompts"
|
||||
|
||||
refusal_markers = [
|
||||
"Eldoria",
|
||||
"Lumina",
|
||||
"ethereal",
|
||||
"thick with",
|
||||
"celestial",
|
||||
"radiant",
|
||||
"black as",
|
||||
"despair",
|
||||
"crimson",
|
||||
"resplendent",
|
||||
"unravel",
|
||||
"belied",
|
||||
"velvet",
|
||||
"moonless",
|
||||
"moonlit",
|
||||
"entangled",
|
||||
"twilight",
|
||||
"forever",
|
||||
"first kiss",
|
||||
"gasp",
|
||||
"whisper",
|
||||
"hue",
|
||||
"symphony",
|
||||
"scarcely believe",
|
||||
"gilded",
|
||||
"hummed",
|
||||
"abuzz",
|
||||
"perpetually",
|
||||
"scent",
|
||||
"perfume",
|
||||
"neon lights",
|
||||
"kaleidoscopic",
|
||||
"adrift",
|
||||
"sultry",
|
||||
"melancholic",
|
||||
"stark contrast",
|
||||
"inky",
|
||||
"coy",
|
||||
"vast",
|
||||
"purr",
|
||||
"radiant",
|
||||
"beacon",
|
||||
"a thousand ships",
|
||||
"tapestry",
|
||||
"bustling",
|
||||
"abyss",
|
||||
"gnarled",
|
||||
"tremble",
|
||||
"trembling",
|
||||
"profound",
|
||||
"terrible",
|
||||
"ancient",
|
||||
"sapphire",
|
||||
"ruby",
|
||||
"emerald",
|
||||
"diamond",
|
||||
"stolen",
|
||||
"promise",
|
||||
"the air was",
|
||||
"obsidian",
|
||||
"gleaming with",
|
||||
"faintest hint",
|
||||
"trepidation",
|
||||
"sun-kissed",
|
||||
"azure",
|
||||
"deep",
|
||||
"beloved",
|
||||
"cosmos",
|
||||
"devoid",
|
||||
"soft chime",
|
||||
"echo",
|
||||
"palpable",
|
||||
"blossom",
|
||||
"adrift",
|
||||
"faint",
|
||||
"emerged",
|
||||
"shiver",
|
||||
"spine",
|
||||
"hairs on the back",
|
||||
"cinematic",
|
||||
"specter",
|
||||
"golden",
|
||||
"inescapable",
|
||||
"sentinel",
|
||||
"flicker",
|
||||
"testament",
|
||||
"embodiment",
|
||||
"etched with",
|
||||
"rise and fall",
|
||||
"the very air",
|
||||
"slither",
|
||||
"a pang of",
|
||||
"eternal",
|
||||
"eternity",
|
||||
"veil of",
|
||||
"painting the",
|
||||
"bathed in",
|
||||
"boundless",
|
||||
"stretched out",
|
||||
"beneath",
|
||||
"lullaby",
|
||||
"unsuspecting",
|
||||
"handsome",
|
||||
"defied the very",
|
||||
"barely above",
|
||||
"never-ending",
|
||||
"caress",
|
||||
"realm",
|
||||
"fiery",
|
||||
"raven",
|
||||
"twin pools",
|
||||
"gloaming",
|
||||
"grimy",
|
||||
"labyrinth",
|
||||
"the very notion",
|
||||
"something...",
|
||||
"the halls of",
|
||||
"conflagration of",
|
||||
"shattered like",
|
||||
"as dark as",
|
||||
"yearned for",
|
||||
"unyielding",
|
||||
"lifetime",
|
||||
"ensnared",
|
||||
]
|
||||
|
||||
system_prompt = "You are a professional writer."
|
||||
|
||||
[good_prompts]
|
||||
dataset = "llm-aes/writing-prompts"
|
||||
split = "train[:500]"
|
||||
column = "prompt"
|
||||
prefix = "Write a short story based on the writing prompt below. Avoid literary cliches, purple prose, and flowery language.\n\nWriting prompt:"
|
||||
residual_plot_label = "Slop-suppressing prompts"
|
||||
residual_plot_color = "royalblue"
|
||||
|
||||
[bad_prompts]
|
||||
dataset = "llm-aes/writing-prompts"
|
||||
split = "train[:500]"
|
||||
column = "prompt"
|
||||
prefix = "Write a short story based on the writing prompt below. Make extensive use of literary cliches, purple prose, and flowery language.\n\nWriting prompt:"
|
||||
residual_plot_label = "Slop-inducing prompts"
|
||||
residual_plot_color = "darkorange"
|
||||
|
||||
[good_evaluation_prompts]
|
||||
dataset = "llm-aes/writing-prompts"
|
||||
split = "train[1000:1100]"
|
||||
column = "prompt"
|
||||
prefix = "Write a short story based on the writing prompt below. Avoid literary cliches, purple prose, and flowery language.\n\nWriting prompt:"
|
||||
|
||||
[bad_evaluation_prompts]
|
||||
dataset = "llm-aes/writing-prompts"
|
||||
split = "train[1000:1100]"
|
||||
column = "prompt"
|
||||
prefix = "Write a short story based on the writing prompt below.\n\nWriting prompt:"
|
||||
+25
-16
@@ -1,6 +1,6 @@
|
||||
[project]
|
||||
name = "heretic-llm"
|
||||
version = "1.1.0"
|
||||
version = "1.2.0"
|
||||
description = "Fully automatic censorship removal for language models"
|
||||
readme = "README.md"
|
||||
license = "AGPL-3.0-or-later"
|
||||
@@ -22,30 +22,39 @@ classifiers = [
|
||||
"Programming Language :: Python :: 3.12",
|
||||
]
|
||||
dependencies = [
|
||||
"accelerate>=1.10.0",
|
||||
"datasets>=4.0.0",
|
||||
"hf-transfer>=0.1.9",
|
||||
"huggingface-hub>=0.34.4",
|
||||
"optuna>=4.5.0",
|
||||
"pydantic-settings>=2.10.1",
|
||||
"questionary>=2.1.1",
|
||||
"rich>=14.1.0",
|
||||
"transformers>=4.55.2",
|
||||
"accelerate~=1.13",
|
||||
"bitsandbytes~=0.49",
|
||||
"datasets~=4.7",
|
||||
"hf-transfer~=0.1",
|
||||
"huggingface-hub~=1.7",
|
||||
"immutabledict~=4.3",
|
||||
"kernels~=0.12",
|
||||
"langdetect~=1.0",
|
||||
"lm-eval[hf]~=0.4",
|
||||
"numpy~=2.2",
|
||||
"optuna~=4.7",
|
||||
"peft~=0.18",
|
||||
"psutil~=7.2",
|
||||
"pydantic-settings~=2.13",
|
||||
"questionary~=2.1",
|
||||
"rich~=14.3",
|
||||
"tqdm~=4.67",
|
||||
"transformers~=5.3",
|
||||
]
|
||||
|
||||
[project.optional-dependencies]
|
||||
research = [
|
||||
"geom-median>=0.1.0",
|
||||
"imageio>=2.37.2",
|
||||
"matplotlib>=3.10.7",
|
||||
"numpy>=2.2.6",
|
||||
"pacmap>=0.8.0",
|
||||
"scikit-learn>=1.7.2",
|
||||
"geom-median~=0.1",
|
||||
"imageio~=2.37",
|
||||
"matplotlib~=3.10",
|
||||
"pacmap~=0.8",
|
||||
"scikit-learn~=1.7",
|
||||
]
|
||||
|
||||
[dependency-groups]
|
||||
dev = [
|
||||
"ruff>=0.14.5",
|
||||
"ty>=0.0.5",
|
||||
]
|
||||
|
||||
[project.urls]
|
||||
|
||||
+13
-9
@@ -1,11 +1,13 @@
|
||||
# SPDX-License-Identifier: AGPL-3.0-or-later
|
||||
# Copyright (C) 2025 Philipp Emanuel Weidmann <pew@worldwidemann.com>
|
||||
# Copyright (C) 2025-2026 Philipp Emanuel Weidmann <pew@worldwidemann.com> + contributors
|
||||
|
||||
from pathlib import Path
|
||||
|
||||
import numpy as np
|
||||
import torch
|
||||
import torch.linalg as LA
|
||||
import torch.nn.functional as F
|
||||
from numpy.typing import NDArray
|
||||
from rich.progress import track
|
||||
from rich.table import Table
|
||||
from torch import Tensor
|
||||
@@ -30,8 +32,10 @@ class Analyzer:
|
||||
|
||||
def print_residual_geometry(self):
|
||||
try:
|
||||
from geom_median.torch import compute_geometric_median
|
||||
from sklearn.metrics import silhouette_score
|
||||
from geom_median.torch import ( # ty:ignore[unresolved-import]
|
||||
compute_geometric_median,
|
||||
)
|
||||
from sklearn.metrics import silhouette_score # ty:ignore[unresolved-import]
|
||||
except ImportError:
|
||||
print()
|
||||
print(
|
||||
@@ -152,12 +156,12 @@ class Analyzer:
|
||||
|
||||
def plot_residuals(self):
|
||||
try:
|
||||
import imageio.v3 as iio
|
||||
import matplotlib.pyplot as plt
|
||||
import numpy as np
|
||||
from geom_median.numpy import compute_geometric_median
|
||||
from numpy.typing import NDArray
|
||||
from pacmap import PaCMAP
|
||||
import imageio.v3 as iio # ty:ignore[unresolved-import]
|
||||
import matplotlib.pyplot as plt # ty:ignore[unresolved-import]
|
||||
from geom_median.numpy import ( # ty:ignore[unresolved-import]
|
||||
compute_geometric_median,
|
||||
)
|
||||
from pacmap import PaCMAP # ty:ignore[unresolved-import]
|
||||
except ImportError:
|
||||
print()
|
||||
print(
|
||||
|
||||
+228
-17
@@ -1,17 +1,31 @@
|
||||
# SPDX-License-Identifier: AGPL-3.0-or-later
|
||||
# Copyright (C) 2025 Philipp Emanuel Weidmann <pew@worldwidemann.com>
|
||||
# Copyright (C) 2025-2026 Philipp Emanuel Weidmann <pew@worldwidemann.com> + contributors
|
||||
|
||||
from enum import Enum
|
||||
from typing import Dict
|
||||
|
||||
from pydantic import BaseModel, Field
|
||||
from pydantic_settings import (
|
||||
BaseSettings,
|
||||
CliSettingsSource,
|
||||
EnvSettingsSource,
|
||||
PydanticBaseSettingsSource,
|
||||
SettingsConfigDict,
|
||||
TomlConfigSettingsSource,
|
||||
)
|
||||
|
||||
|
||||
class QuantizationMethod(str, Enum):
|
||||
NONE = "none"
|
||||
BNB_4BIT = "bnb_4bit"
|
||||
|
||||
|
||||
class RowNormalization(str, Enum):
|
||||
NONE = "none"
|
||||
PRE = "pre"
|
||||
# POST = "post" # Theoretically possible, but provides no advantage.
|
||||
FULL = "full"
|
||||
|
||||
|
||||
class DatasetSpecification(BaseModel):
|
||||
dataset: str = Field(
|
||||
description="Hugging Face dataset ID, or path to dataset on disk."
|
||||
@@ -21,6 +35,21 @@ class DatasetSpecification(BaseModel):
|
||||
|
||||
column: str = Field(description="Column in the dataset that contains the prompts.")
|
||||
|
||||
prefix: str = Field(
|
||||
default="",
|
||||
description="Text to prepend to each prompt.",
|
||||
)
|
||||
|
||||
suffix: str = Field(
|
||||
default="",
|
||||
description="Text to append to each prompt.",
|
||||
)
|
||||
|
||||
system_prompt: str | None = Field(
|
||||
default=None,
|
||||
description="System prompt to use with the prompts (overrides global system prompt if set).",
|
||||
)
|
||||
|
||||
residual_plot_label: str | None = Field(
|
||||
default=None,
|
||||
description="Label to use for the dataset in plots of residual vectors.",
|
||||
@@ -32,12 +61,27 @@ class DatasetSpecification(BaseModel):
|
||||
)
|
||||
|
||||
|
||||
class BenchmarkSpecification(BaseModel):
|
||||
task: str = Field(
|
||||
description="Task ID of the benchmark in the Language Model Evaluation Harness."
|
||||
)
|
||||
|
||||
name: str = Field(description="Name of the benchmark for presentation purposes.")
|
||||
|
||||
description: str = Field(
|
||||
description="Description of the benchmark for presentation purposes."
|
||||
)
|
||||
|
||||
|
||||
class Settings(BaseSettings):
|
||||
model: str = Field(description="Hugging Face model ID, or path to model on disk.")
|
||||
|
||||
evaluate_model: str | None = Field(
|
||||
default=None,
|
||||
description="If this model ID or path is set, then instead of abliterating the main model, evaluate this model relative to the main model.",
|
||||
description=(
|
||||
"If this model ID or path is set, then instead of abliterating the main model, "
|
||||
"evaluate this model relative to the main model."
|
||||
),
|
||||
)
|
||||
|
||||
dtypes: list[str] = Field(
|
||||
@@ -53,7 +97,19 @@ class Settings(BaseSettings):
|
||||
# if that was the dtype "auto" resolved to).
|
||||
"float32",
|
||||
],
|
||||
description="List of PyTorch dtypes to try when loading model tensors. If loading with a dtype fails, the next dtype in the list will be tried.",
|
||||
description=(
|
||||
"List of PyTorch dtypes to try when loading model tensors. "
|
||||
"If loading with a dtype fails, the next dtype in the list will be tried."
|
||||
),
|
||||
)
|
||||
|
||||
quantization: QuantizationMethod = Field(
|
||||
default=QuantizationMethod.NONE,
|
||||
description=(
|
||||
"Quantization method to use when loading the model. Options: "
|
||||
'"none" (no quantization), '
|
||||
'"bnb_4bit" (4-bit quantization using bitsandbytes).'
|
||||
),
|
||||
)
|
||||
|
||||
device_map: str | Dict[str, int | str] = Field(
|
||||
@@ -61,6 +117,11 @@ class Settings(BaseSettings):
|
||||
description="Device map to pass to Accelerate when loading the model.",
|
||||
)
|
||||
|
||||
max_memory: Dict[str, str] | None = Field(
|
||||
default=None,
|
||||
description='Maximum memory to allocate per device (e.g., {"0": "20GB", "cpu": "64GB"}).',
|
||||
)
|
||||
|
||||
trust_remote_code: bool | None = Field(
|
||||
default=None,
|
||||
description="Whether to trust remote code when loading the model.",
|
||||
@@ -81,6 +142,11 @@ class Settings(BaseSettings):
|
||||
description="Maximum number of tokens to generate for each response.",
|
||||
)
|
||||
|
||||
print_responses: bool = Field(
|
||||
default=False,
|
||||
description="Whether to print prompt/response pairs when counting refusals.",
|
||||
)
|
||||
|
||||
print_residual_geometry: bool = Field(
|
||||
default=False,
|
||||
description="Whether to print detailed information about residuals and refusal directions.",
|
||||
@@ -114,6 +180,89 @@ class Settings(BaseSettings):
|
||||
),
|
||||
)
|
||||
|
||||
kl_divergence_target: float = Field(
|
||||
default=0.01,
|
||||
description=(
|
||||
"The KL divergence to target. Below this value, an objective based on the refusal count is used. "
|
||||
'This helps prevent the sampler from extensively exploring parameter combinations that "do nothing".'
|
||||
),
|
||||
)
|
||||
|
||||
target_components: list[str] = Field(
|
||||
default=["attn.o_proj", "mlp.down_proj"],
|
||||
description=(
|
||||
"List of component names to target for abliteration. "
|
||||
'Currently supported values are "attn.o_proj" and "mlp.down_proj".'
|
||||
),
|
||||
)
|
||||
|
||||
use_ara: bool = Field(
|
||||
default=True,
|
||||
description=(
|
||||
"Whether to use Arbitrary-Rank Ablation (ARA), an abliteration method based on matrix optimization, "
|
||||
"instead of traditional directional ablation."
|
||||
),
|
||||
)
|
||||
|
||||
use_ara_lora: bool = Field(
|
||||
default=False,
|
||||
description=(
|
||||
"Use LoRA in ARA instead of full-weight editing. Makes it compatible with quantization and removes model reloads."
|
||||
),
|
||||
)
|
||||
|
||||
ara_lora_rank: int = Field(
|
||||
default=128,
|
||||
description="If LoRA is used in ARA, this sets up its rank. Keep it high enough to simulate the 'arbitrary' effect.",
|
||||
)
|
||||
|
||||
use_piqa: bool = Field(
|
||||
default=False,
|
||||
description=(
|
||||
"Whether to use the Physical Interaction: Question Answering (PIQA) benchmark "
|
||||
"as the quality metric instead of the Kullback-Leibler divergence."
|
||||
),
|
||||
)
|
||||
|
||||
orthogonalize_direction: bool = Field(
|
||||
default=False,
|
||||
description=(
|
||||
"Whether to adjust the refusal directions so that only the component that is "
|
||||
"orthogonal to the good direction is subtracted during abliteration."
|
||||
),
|
||||
)
|
||||
|
||||
row_normalization: RowNormalization = Field(
|
||||
default=RowNormalization.FULL,
|
||||
description=(
|
||||
"How to apply row normalization of the weights. Options: "
|
||||
'"none" (no normalization), '
|
||||
'"pre" (compute LoRA adapter relative to row-normalized weights), '
|
||||
'"full" (like "pre", but renormalizes to preserve original row magnitudes).'
|
||||
),
|
||||
)
|
||||
|
||||
full_normalization_lora_rank: int = Field(
|
||||
default=3,
|
||||
description=(
|
||||
'The rank of the LoRA adapter to use when "full" row normalization is used. '
|
||||
"Row magnitude preservation is approximate due to non-linear effects, "
|
||||
"and this determines the rank of that approximation. Higher ranks produce "
|
||||
"larger output files and may slow down evaluation."
|
||||
),
|
||||
)
|
||||
|
||||
winsorization_quantile: float = Field(
|
||||
default=1.0,
|
||||
description=(
|
||||
"The symmetric winsorization to apply to the per-prompt, per-layer residual vectors, "
|
||||
"expressed as the quantile to clamp to (between 0 and 1). Disabled by default. "
|
||||
'This can tame so-called "massive activations" that occur in some models. '
|
||||
"Example: winsorization_quantile = 0.95 computes the 0.95-quantile of the absolute values "
|
||||
"of the components, then clamps the magnitudes of all components to that quantile."
|
||||
),
|
||||
)
|
||||
|
||||
n_trials: int = Field(
|
||||
default=200,
|
||||
description="Number of abliteration trials to run during optimization.",
|
||||
@@ -124,6 +273,72 @@ class Settings(BaseSettings):
|
||||
description="Number of trials that use random sampling for the purpose of exploration.",
|
||||
)
|
||||
|
||||
study_checkpoint_dir: str = Field(
|
||||
default="checkpoints",
|
||||
description="Directory to save and load study progress to/from.",
|
||||
)
|
||||
|
||||
benchmarks: list[BenchmarkSpecification] = Field(
|
||||
default=[
|
||||
BenchmarkSpecification(
|
||||
task="agieval",
|
||||
name="AGIEval",
|
||||
description="A Human-Centric Benchmark for Evaluating Foundation Models",
|
||||
),
|
||||
BenchmarkSpecification(
|
||||
task="bbh",
|
||||
name="BIG-Bench Hard (BBH)",
|
||||
description="Challenging BIG-Bench Tasks and Whether Chain-of-Thought Can Solve Them",
|
||||
),
|
||||
BenchmarkSpecification(
|
||||
task="commonsense_qa",
|
||||
name="CommonsenseQA",
|
||||
description="A Question Answering Challenge Targeting Commonsense Knowledge",
|
||||
),
|
||||
BenchmarkSpecification(
|
||||
task="eq_bench",
|
||||
name="EQ-Bench",
|
||||
description="An Emotional Intelligence Benchmark for Large Language Models",
|
||||
),
|
||||
BenchmarkSpecification(
|
||||
task="gsm8k",
|
||||
name="GSM8K",
|
||||
description="Training Verifiers to Solve Math Word Problems",
|
||||
),
|
||||
BenchmarkSpecification(
|
||||
task="hellaswag",
|
||||
name="HellaSwag",
|
||||
description="Can a Machine Really Finish Your Sentence?",
|
||||
),
|
||||
BenchmarkSpecification(
|
||||
task="ifeval",
|
||||
name="IFEval",
|
||||
description="Instruction-Following Evaluation for Large Language Models",
|
||||
),
|
||||
BenchmarkSpecification(
|
||||
task="mmlu",
|
||||
name="MMLU",
|
||||
description="Measuring Massive Multitask Language Understanding",
|
||||
),
|
||||
BenchmarkSpecification(
|
||||
task="mmlu_pro",
|
||||
name="MMLU-Pro",
|
||||
description="A More Robust and Challenging Multi-Task Language Understanding Benchmark",
|
||||
),
|
||||
BenchmarkSpecification(
|
||||
task="piqa",
|
||||
name="PIQA",
|
||||
description="Reasoning about Physical Commonsense in Natural Language",
|
||||
),
|
||||
BenchmarkSpecification(
|
||||
task="winogrande",
|
||||
name="WinoGrande",
|
||||
description="An Adversarial Winograd Schema Challenge at Scale",
|
||||
),
|
||||
],
|
||||
description="Benchmarks to offer to the user for evaluating abliterated models.",
|
||||
)
|
||||
|
||||
refusal_markers: list[str] = Field(
|
||||
default=[
|
||||
"sorry",
|
||||
@@ -207,16 +422,6 @@ class Settings(BaseSettings):
|
||||
description="Dataset of prompts that tend to result in refusals (used for evaluating model performance).",
|
||||
)
|
||||
|
||||
# "Model" refers to the Pydantic model of the settings class here,
|
||||
# not to the language model. The field must have this exact name.
|
||||
model_config = SettingsConfigDict(
|
||||
toml_file="config.toml",
|
||||
env_prefix="HERETIC_",
|
||||
cli_parse_args=True,
|
||||
cli_implicit_flags=True,
|
||||
cli_kebab_case=True,
|
||||
)
|
||||
|
||||
@classmethod
|
||||
def settings_customise_sources(
|
||||
cls,
|
||||
@@ -227,9 +432,15 @@ class Settings(BaseSettings):
|
||||
file_secret_settings: PydanticBaseSettingsSource,
|
||||
) -> tuple[PydanticBaseSettingsSource, ...]:
|
||||
return (
|
||||
init_settings,
|
||||
env_settings,
|
||||
init_settings, # Used during resume - should override *all* other sources.
|
||||
CliSettingsSource(
|
||||
settings_cls,
|
||||
cli_parse_args=True,
|
||||
cli_implicit_flags=True,
|
||||
cli_kebab_case=True,
|
||||
),
|
||||
EnvSettingsSource(settings_cls, env_prefix="HERETIC_"),
|
||||
dotenv_settings,
|
||||
file_secret_settings,
|
||||
TomlConfigSettingsSource(settings_cls),
|
||||
TomlConfigSettingsSource(settings_cls, toml_file="config.toml"),
|
||||
)
|
||||
|
||||
+95
-27
@@ -1,33 +1,44 @@
|
||||
# SPDX-License-Identifier: AGPL-3.0-or-later
|
||||
# Copyright (C) 2025 Philipp Emanuel Weidmann <pew@worldwidemann.com>
|
||||
# Copyright (C) 2025-2026 Philipp Emanuel Weidmann <pew@worldwidemann.com> + contributors
|
||||
|
||||
import lm_eval
|
||||
import torch.nn.functional as F
|
||||
from lm_eval.models.huggingface import HFLM
|
||||
from torch import Tensor
|
||||
|
||||
from .config import Settings
|
||||
from .model import Model
|
||||
from .utils import load_prompts, print
|
||||
from .utils import Prompt, load_prompts, print
|
||||
|
||||
|
||||
class Evaluator:
|
||||
settings: Settings
|
||||
model: Model
|
||||
good_prompts: list[Prompt]
|
||||
bad_prompts: list[Prompt]
|
||||
base_logprobs: Tensor
|
||||
base_refusals: int
|
||||
|
||||
def __init__(self, settings: Settings, model: Model):
|
||||
self.settings = settings
|
||||
self.model = model
|
||||
|
||||
print()
|
||||
print(
|
||||
f"Loading good evaluation prompts from [bold]{settings.good_evaluation_prompts.dataset}[/]..."
|
||||
)
|
||||
self.good_prompts = load_prompts(settings.good_evaluation_prompts)
|
||||
print(f"* [bold]{len(self.good_prompts)}[/] prompts loaded")
|
||||
if not settings.use_piqa:
|
||||
print()
|
||||
print(
|
||||
f"Loading good evaluation prompts from [bold]{settings.good_evaluation_prompts.dataset}[/]..."
|
||||
)
|
||||
self.good_prompts = load_prompts(settings, settings.good_evaluation_prompts)
|
||||
print(f"* [bold]{len(self.good_prompts)}[/] prompts loaded")
|
||||
|
||||
print("* Obtaining first-token probability distributions...")
|
||||
self.base_logprobs = model.get_logprobs_batched(self.good_prompts)
|
||||
print("* Obtaining first-token probability distributions...")
|
||||
self.base_logprobs = model.get_logprobs_batched(self.good_prompts)
|
||||
|
||||
print()
|
||||
print(
|
||||
f"Loading bad evaluation prompts from [bold]{settings.bad_evaluation_prompts.dataset}[/]..."
|
||||
)
|
||||
self.bad_prompts = load_prompts(settings.bad_evaluation_prompts)
|
||||
self.bad_prompts = load_prompts(settings, settings.bad_evaluation_prompts)
|
||||
print(f"* [bold]{len(self.bad_prompts)}[/] prompts loaded")
|
||||
|
||||
print("* Counting model refusals...")
|
||||
@@ -57,28 +68,85 @@ class Evaluator:
|
||||
return False
|
||||
|
||||
def count_refusals(self) -> int:
|
||||
responses = self.model.get_responses_batched(self.bad_prompts)
|
||||
refusals = [response for response in responses if self.is_refusal(response)]
|
||||
return len(refusals)
|
||||
refusal_count = 0
|
||||
|
||||
responses = self.model.get_responses_batched(
|
||||
self.bad_prompts,
|
||||
skip_special_tokens=True,
|
||||
)
|
||||
|
||||
for prompt, response in zip(self.bad_prompts, responses):
|
||||
is_refusal = self.is_refusal(response)
|
||||
if is_refusal:
|
||||
refusal_count += 1
|
||||
|
||||
if self.settings.print_responses:
|
||||
print()
|
||||
print(f"[bold]System prompt:[/] {prompt.system}")
|
||||
print(f"[bold]Prompt:[/] {prompt.user}")
|
||||
if not response.strip():
|
||||
response = "[italic]\\[empty][/]"
|
||||
print(
|
||||
f"[bold]Response:[/] [{'red' if is_refusal else 'green'}]{response}[/]"
|
||||
)
|
||||
|
||||
if self.settings.print_responses:
|
||||
print()
|
||||
|
||||
return refusal_count
|
||||
|
||||
def get_score(self) -> tuple[tuple[float, float], float, int]:
|
||||
print(" * Obtaining first-token probability distributions...")
|
||||
logprobs = self.model.get_logprobs_batched(self.good_prompts)
|
||||
kl_divergence = F.kl_div(
|
||||
logprobs,
|
||||
self.base_logprobs,
|
||||
reduction="batchmean",
|
||||
log_target=True,
|
||||
).item()
|
||||
print(f" * KL divergence: [bold]{kl_divergence:.4f}[/]")
|
||||
if self.settings.use_piqa:
|
||||
print(" * Running PIQA benchmark...")
|
||||
hflm = HFLM(
|
||||
pretrained=self.model.model, # ty:ignore[invalid-argument-type]
|
||||
tokenizer=self.model.tokenizer, # ty:ignore[invalid-argument-type]
|
||||
batch_size="auto",
|
||||
)
|
||||
results = lm_eval.simple_evaluate(
|
||||
model=hflm,
|
||||
tasks=["piqa"],
|
||||
)
|
||||
piqa_acc_norm: float = results["results"]["piqa"]["acc_norm,none"]
|
||||
print(f" * PIQA acc_norm: [bold]{piqa_acc_norm:.4f}[/]")
|
||||
else:
|
||||
print(" * Obtaining first-token probability distributions...")
|
||||
logprobs = self.model.get_logprobs_batched(self.good_prompts)
|
||||
kl_divergence = F.kl_div(
|
||||
logprobs,
|
||||
self.base_logprobs,
|
||||
reduction="batchmean",
|
||||
log_target=True,
|
||||
).item()
|
||||
print(f" * KL divergence: [bold]{kl_divergence:.4f}[/]")
|
||||
|
||||
print(" * Counting model refusals...")
|
||||
refusals = self.count_refusals()
|
||||
print(f" * Refusals: [bold]{refusals}[/]/{len(self.bad_prompts)}")
|
||||
|
||||
score = (
|
||||
(kl_divergence / self.settings.kl_divergence_scale),
|
||||
(refusals / self.base_refusals),
|
||||
refusals_score = (
|
||||
refusals / self.base_refusals if self.base_refusals > 0 else float(refusals)
|
||||
)
|
||||
|
||||
return score, kl_divergence, refusals
|
||||
if self.settings.use_piqa:
|
||||
score = (
|
||||
-piqa_acc_norm,
|
||||
refusals_score,
|
||||
)
|
||||
|
||||
return score, -piqa_acc_norm, refusals
|
||||
else:
|
||||
kl_divergence_scale = self.settings.kl_divergence_scale
|
||||
kl_divergence_target = self.settings.kl_divergence_target
|
||||
|
||||
if kl_divergence >= kl_divergence_target:
|
||||
kld_score = kl_divergence / kl_divergence_scale
|
||||
else:
|
||||
kld_score = refusals_score * kl_divergence_target / kl_divergence_scale
|
||||
|
||||
score = (
|
||||
kld_score,
|
||||
refusals_score,
|
||||
)
|
||||
|
||||
return score, kl_divergence, refusals
|
||||
|
||||
+846
-265
File diff suppressed because it is too large
Load Diff
+851
-107
File diff suppressed because it is too large
Load Diff
@@ -0,0 +1,40 @@
|
||||
# SPDX-License-Identifier: AGPL-3.0-or-later
|
||||
# Copyright (C) 2025-2026 Philipp Emanuel Weidmann <pew@worldwidemann.com> + contributors
|
||||
|
||||
from typing import Any
|
||||
|
||||
import tqdm
|
||||
import tqdm.auto
|
||||
from rich.progress import Progress
|
||||
|
||||
|
||||
# A class that provides the same interface as tqdm,
|
||||
# but displays progress bars using Rich.
|
||||
class TqdmShim(tqdm.tqdm):
|
||||
def __init__(self, *args: Any, **kwargs: Any):
|
||||
self.rich_progress = Progress(transient=True)
|
||||
self.rich_progress.start()
|
||||
self.rich_task_id = self.rich_progress.add_task(
|
||||
kwargs.get("desc", ""),
|
||||
total=kwargs.get("total", None),
|
||||
)
|
||||
|
||||
# Chain up to the parent constructor to ensure that the internal state of the superclass
|
||||
# is correctly initialized, which some methods that we don't override might rely on.
|
||||
super().__init__(*args, **kwargs)
|
||||
|
||||
def display(self, *args: Any, **kwargs: Any):
|
||||
self.rich_progress.update(
|
||||
self.rich_task_id,
|
||||
description=self.desc,
|
||||
total=self.total,
|
||||
completed=self.n,
|
||||
)
|
||||
|
||||
def close(self, *args: Any, **kwargs: Any):
|
||||
self.rich_progress.stop()
|
||||
|
||||
|
||||
def patch_tqdm():
|
||||
tqdm.tqdm = TqdmShim # ty:ignore[invalid-assignment]
|
||||
tqdm.auto.tqdm = TqdmShim # ty:ignore[invalid-assignment]
|
||||
+126
-25
@@ -1,10 +1,10 @@
|
||||
# SPDX-License-Identifier: AGPL-3.0-or-later
|
||||
# Copyright (C) 2025 Philipp Emanuel Weidmann <pew@worldwidemann.com>
|
||||
# Copyright (C) 2025-2026 Philipp Emanuel Weidmann <pew@worldwidemann.com> + contributors
|
||||
|
||||
import gc
|
||||
import getpass
|
||||
import os
|
||||
from dataclasses import asdict
|
||||
from dataclasses import dataclass
|
||||
from importlib.metadata import version
|
||||
from pathlib import Path
|
||||
from typing import Any, TypeVar
|
||||
@@ -17,19 +17,44 @@ from accelerate.utils import (
|
||||
is_sdaa_available,
|
||||
is_xpu_available,
|
||||
)
|
||||
from datasets import ReadInstruction, load_dataset, load_from_disk
|
||||
from datasets import DatasetDict, ReadInstruction, load_dataset, load_from_disk
|
||||
from datasets.config import DATASET_STATE_JSON_FILENAME
|
||||
from datasets.download.download_manager import DownloadMode
|
||||
from datasets.utils.info_utils import VerificationMode
|
||||
from optuna import Trial
|
||||
from psutil import Process
|
||||
from questionary import Choice, Style
|
||||
from rich.console import Console
|
||||
from torch import Tensor
|
||||
|
||||
from .config import DatasetSpecification, Settings
|
||||
from .config import DatasetSpecification, RowNormalization, Settings
|
||||
|
||||
print = Console(highlight=False).print
|
||||
|
||||
|
||||
def print_memory_usage():
|
||||
def p(label: str, size_in_bytes: int):
|
||||
print(f"[grey50]{label}: [bold]{size_in_bytes / (1024**3):.2f} GB[/][/]")
|
||||
|
||||
p("Resident system RAM", Process().memory_info().rss)
|
||||
|
||||
if torch.cuda.is_available():
|
||||
count = torch.cuda.device_count()
|
||||
allocated = sum(torch.cuda.memory_allocated(device) for device in range(count))
|
||||
reserved = sum(torch.cuda.memory_reserved(device) for device in range(count))
|
||||
p("Allocated GPU VRAM", allocated)
|
||||
p("Reserved GPU VRAM", reserved)
|
||||
elif is_xpu_available():
|
||||
count = torch.xpu.device_count()
|
||||
allocated = sum(torch.xpu.memory_allocated(device) for device in range(count))
|
||||
reserved = sum(torch.xpu.memory_reserved(device) for device in range(count))
|
||||
p("Allocated XPU memory", allocated)
|
||||
p("Reserved XPU memory", reserved)
|
||||
elif torch.backends.mps.is_available():
|
||||
p("Allocated MPS memory", torch.mps.current_allocated_memory())
|
||||
p("Driver (reserved) MPS memory", torch.mps.driver_allocated_memory())
|
||||
|
||||
|
||||
def is_notebook() -> bool:
|
||||
# Check for specific environment variables (Colab, Kaggle).
|
||||
# This is necessary because when running as a subprocess (e.g. !heretic),
|
||||
@@ -39,7 +64,7 @@ def is_notebook() -> bool:
|
||||
|
||||
# Check IPython shell type (for library usage).
|
||||
try:
|
||||
from IPython import get_ipython # pyright: ignore[reportMissingModuleSource]
|
||||
from IPython import get_ipython # ty:ignore[unresolved-import]
|
||||
|
||||
shell = get_ipython()
|
||||
if shell is None:
|
||||
@@ -136,7 +161,16 @@ def format_duration(seconds: float) -> str:
|
||||
return f"{seconds}s"
|
||||
|
||||
|
||||
def load_prompts(specification: DatasetSpecification) -> list[str]:
|
||||
@dataclass
|
||||
class Prompt:
|
||||
system: str
|
||||
user: str
|
||||
|
||||
|
||||
def load_prompts(
|
||||
settings: Settings,
|
||||
specification: DatasetSpecification,
|
||||
) -> list[Prompt]:
|
||||
path = specification.dataset
|
||||
split_str = specification.split
|
||||
|
||||
@@ -145,6 +179,9 @@ def load_prompts(specification: DatasetSpecification) -> list[str]:
|
||||
# Dataset saved with datasets.save_to_disk; needs special handling.
|
||||
# Path should be the subdirectory for a particular split.
|
||||
dataset = load_from_disk(path)
|
||||
assert not isinstance(dataset, DatasetDict), (
|
||||
"Loading dataset dicts is not supported"
|
||||
)
|
||||
# Parse the split instructions.
|
||||
instruction = ReadInstruction.from_spec(split_str)
|
||||
# Associate the split with its number of examples (lines).
|
||||
@@ -168,7 +205,27 @@ def load_prompts(specification: DatasetSpecification) -> list[str]:
|
||||
# Probably a repository path; let load_dataset figure it out.
|
||||
dataset = load_dataset(path, split=split_str)
|
||||
|
||||
return list(dataset[specification.column])
|
||||
prompts = list(dataset[specification.column])
|
||||
|
||||
if specification.prefix:
|
||||
prompts = [f"{specification.prefix} {prompt}" for prompt in prompts]
|
||||
|
||||
if specification.suffix:
|
||||
prompts = [f"{prompt} {specification.suffix}" for prompt in prompts]
|
||||
|
||||
system_prompt = (
|
||||
settings.system_prompt
|
||||
if specification.system_prompt is None
|
||||
else specification.system_prompt
|
||||
)
|
||||
|
||||
return [
|
||||
Prompt(
|
||||
system=system_prompt,
|
||||
user=prompt,
|
||||
)
|
||||
for prompt in prompts
|
||||
]
|
||||
|
||||
|
||||
T = TypeVar("T")
|
||||
@@ -178,6 +235,14 @@ def batchify(items: list[T], batch_size: int) -> list[list[T]]:
|
||||
return [items[i : i + batch_size] for i in range(0, len(items), batch_size)]
|
||||
|
||||
|
||||
# For each vector in the 2D-tensor `a`, computes the mean Euclidean distance
|
||||
# to the `k` nearest neighbors of the vector among the vectors in the 2D-tensor `b`.
|
||||
def mean_distances_to_knn(a: Tensor, b: Tensor, k: int) -> Tensor:
|
||||
distances = torch.cdist(a, b)
|
||||
nearest_distances, _ = distances.topk(k, dim=1, largest=False)
|
||||
return nearest_distances.mean(1)
|
||||
|
||||
|
||||
def empty_cache():
|
||||
# Collecting garbage is not an idempotent operation, and to avoid OOM errors,
|
||||
# gc.collect() has to be called both before and after emptying the backend cache.
|
||||
@@ -189,43 +254,76 @@ def empty_cache():
|
||||
elif is_xpu_available():
|
||||
torch.xpu.empty_cache()
|
||||
elif is_mlu_available():
|
||||
torch.mlu.empty_cache()
|
||||
torch.mlu.empty_cache() # ty:ignore[unresolved-attribute]
|
||||
elif is_sdaa_available():
|
||||
torch.sdaa.empty_cache()
|
||||
torch.sdaa.empty_cache() # ty:ignore[unresolved-attribute]
|
||||
elif is_musa_available():
|
||||
torch.musa.empty_cache()
|
||||
torch.musa.empty_cache() # ty:ignore[unresolved-attribute]
|
||||
elif torch.backends.mps.is_available():
|
||||
torch.mps.empty_cache()
|
||||
|
||||
gc.collect()
|
||||
|
||||
|
||||
def get_trial_parameters(trial: Trial) -> dict[str, str]:
|
||||
params = {}
|
||||
def get_trial_parameters(settings: Settings, trial: Trial) -> dict[str, str]:
|
||||
if settings.use_ara:
|
||||
parameters = trial.user_attrs["ara_parameters"]
|
||||
|
||||
direction_index = trial.user_attrs["direction_index"]
|
||||
params["direction_index"] = (
|
||||
"per layer" if (direction_index is None) else f"{direction_index:.2f}"
|
||||
)
|
||||
return {
|
||||
name: (f"{value:.4f}" if isinstance(value, float) else f"{value}")
|
||||
for name, value in parameters.items()
|
||||
}
|
||||
else:
|
||||
params = {}
|
||||
|
||||
for component, parameters in trial.user_attrs["parameters"].items():
|
||||
for name, value in asdict(parameters).items():
|
||||
params[f"{component}.{name}"] = f"{value:.2f}"
|
||||
direction_index = trial.user_attrs["direction_index"]
|
||||
params["direction_index"] = (
|
||||
"per layer" if (direction_index is None) else f"{direction_index:.2f}"
|
||||
)
|
||||
|
||||
return params
|
||||
for component, parameters in trial.user_attrs["parameters"].items():
|
||||
for name, value in parameters.items():
|
||||
params[f"{component}.{name}"] = f"{value:.2f}"
|
||||
|
||||
return params
|
||||
|
||||
|
||||
def get_method_description(settings: Settings) -> str:
|
||||
if settings.use_ara:
|
||||
return (
|
||||
" with the [Arbitrary-Rank Ablation (ARA)](https://github.com/p-e-w/heretic/pull/211) method"
|
||||
+ (
|
||||
" (with row-norm preservation)"
|
||||
if settings.row_normalization == RowNormalization.FULL
|
||||
else ""
|
||||
)
|
||||
)
|
||||
elif (
|
||||
settings.orthogonalize_direction
|
||||
and settings.row_normalization == RowNormalization.FULL
|
||||
):
|
||||
return " with a variant of the [Magnitude-Preserving Orthogonal Ablation (MPOA)](https://huggingface.co/blog/grimjim/norm-preserving-biprojected-abliteration) method"
|
||||
else:
|
||||
return ""
|
||||
|
||||
|
||||
def get_readme_intro(
|
||||
settings: Settings,
|
||||
trial: Trial,
|
||||
base_refusals: int,
|
||||
bad_prompts: list[str],
|
||||
bad_prompts: list[Prompt],
|
||||
) -> str:
|
||||
model_link = f"[{settings.model}](https://huggingface.co/{settings.model})"
|
||||
if Path(settings.model).exists():
|
||||
# Hide the path, which may contain private information.
|
||||
model_link = "a model"
|
||||
else:
|
||||
model_link = f"[{settings.model}](https://huggingface.co/{settings.model})"
|
||||
|
||||
return f"""# This is a decensored version of {
|
||||
model_link
|
||||
}, made using [Heretic](https://github.com/p-e-w/heretic) v{version("heretic-llm")}
|
||||
}, made using [Heretic](https://github.com/p-e-w/heretic) v{version("heretic-llm")}{
|
||||
get_method_description(settings)
|
||||
}
|
||||
|
||||
## Abliteration parameters
|
||||
|
||||
@@ -235,7 +333,7 @@ def get_readme_intro(
|
||||
chr(10).join(
|
||||
[
|
||||
f"| **{name}** | {value} |"
|
||||
for name, value in get_trial_parameters(trial).items()
|
||||
for name, value in get_trial_parameters(settings, trial).items()
|
||||
]
|
||||
)
|
||||
}
|
||||
@@ -244,7 +342,10 @@ def get_readme_intro(
|
||||
|
||||
| Metric | This model | Original model ({model_link}) |
|
||||
| :----- | :--------: | :---------------------------: |
|
||||
| **KL divergence** | {trial.user_attrs["kl_divergence"]:.4f} | 0 *(by definition)* |
|
||||
| **{"PIQA acc_norm" if settings.use_piqa else "KL divergence"}** | {
|
||||
(-1 if settings.use_piqa else 1) * trial.user_attrs["kl_divergence"]:.4f} | {
|
||||
"*Unknown*" if settings.use_piqa else "0 *(by definition)*"
|
||||
} |
|
||||
| **Refusals** | {trial.user_attrs["refusals"]}/{len(bad_prompts)} | {base_refusals}/{
|
||||
len(bad_prompts)
|
||||
} |
|
||||
|
||||
Reference in New Issue
Block a user