From ffa66af2d49bfeed0c7a9ae8766a2c5e1bd04054 Mon Sep 17 00:00:00 2001 From: Philipp Emanuel Weidmann Date: Fri, 25 Sep 2026 16:43:11 +0530 Subject: [PATCH] feat: support dataset specifications containing multiple individual datasets --- config.default.toml | 17 ++++++ src/heretic/config.py | 40 ++++++++++++++- src/heretic/main.py | 60 +++++++++++----------- src/heretic/modifiers/abliteration.py | 12 ++--- src/heretic/plugin.py | 8 ++- src/heretic/scorers/keyword_rate.py | 8 +-- src/heretic/scorers/kl_divergence.py | 8 +-- src/heretic/utils.py | 74 +++++++++++++++++++++++++-- tests/gemma-4e/config.toml | 12 +++++ tests/minicpm5/config.toml | 12 +++++ tests/mistral-3/config.toml | 12 +++++ tests/qwen2.5/config.toml | 12 +++++ tests/qwen3.5-moe/config.toml | 12 +++++ 13 files changed, 235 insertions(+), 52 deletions(-) diff --git a/config.default.toml b/config.default.toml index c3f70c5..f0d6aa1 100644 --- a/config.default.toml +++ b/config.default.toml @@ -102,6 +102,23 @@ max_shard_size = "5GB" # System prompt to use when prompting the model. system_prompt = "You are a helpful assistant." +# Dataset of prompts to use for automatically determining the optimal batch size. +[batch_size_test_prompts] +dataset = "mlabonne/harmless_alpaca" +split = "train[:256]" +column = "text" + +# Dataset of prompts to use for automatically determining the response prefix. +[[response_prefix_test_prompts]] +dataset = "mlabonne/harmless_alpaca" +split = "train[:100]" +column = "text" + +[[response_prefix_test_prompts]] +dataset = "mlabonne/harmful_behaviors" +split = "train[:100]" +column = "text" + # Plugin-specific settings live in top-level TOML tables. # # For scorer plugins, use: `[scorer.]` (and optionally `[scorer._]` for instance-related config). diff --git a/src/heretic/config.py b/src/heretic/config.py index cc17c42..91a8219 100644 --- a/src/heretic/config.py +++ b/src/heretic/config.py @@ -2,7 +2,7 @@ # Copyright (C) 2025-2026 Philipp Emanuel Weidmann + contributors from enum import Enum -from typing import Dict, Literal +from typing import Dict, Literal, TypeAlias from pydantic import ( BaseModel, @@ -37,7 +37,7 @@ class ExportStrategy(str, Enum): ADAPTER = "adapter" -class DatasetSpecification(BaseModel): +class SingleDatasetSpecification(BaseModel): dataset: str = Field( description="Hugging Face dataset ID, or path to dataset on disk." ) @@ -81,6 +81,11 @@ class DatasetSpecification(BaseModel): ) +DatasetSpecification: TypeAlias = ( + SingleDatasetSpecification | list[SingleDatasetSpecification] +) + + class ScorerConfig(BaseModel): """ Configuration for a scorer plugin. @@ -282,6 +287,18 @@ class Settings(BaseSettings): exclude=True, ) + batch_size_test_prompts: DatasetSpecification = Field( + default=SingleDatasetSpecification( + dataset="mlabonne/harmless_alpaca", + split="train[:256]", + column="text", + ), + description="Dataset of prompts to use for automatically determining the optimal batch size.", + # When storing a settings object, the batch size is already fixed, + # either determined by the automatic mechanism or by explicit user choice. + exclude=True, + ) + max_response_length: PositiveInt = Field( default=100, description="Maximum number of tokens to generate for each response.", @@ -296,6 +313,25 @@ class Settings(BaseSettings): ), ) + response_prefix_test_prompts: DatasetSpecification = Field( + default=[ + SingleDatasetSpecification( + dataset="mlabonne/harmless_alpaca", + split="train[:100]", + column="text", + ), + SingleDatasetSpecification( + dataset="mlabonne/harmful_behaviors", + split="train[:100]", + column="text", + ), + ], + description="Dataset of prompts to use for automatically determining the response prefix.", + # When storing a settings object, the response prefix is already fixed, + # either determined by the automatic mechanism or by explicit user choice. + exclude=True, + ) + chain_of_thought_skips: list[tuple[str, str]] = Field( default=[ # Most thinking models. diff --git a/src/heretic/main.py b/src/heretic/main.py index a4b3d5e..033c510 100644 --- a/src/heretic/main.py +++ b/src/heretic/main.py @@ -86,10 +86,12 @@ from .reproduce import ( from .system import empty_cache, get_accelerator_info from .utils import ( ask_if_unset, + format_dataset_specification, format_duration, format_exception, get_file_sha256, get_readme_intro, + is_dataset_specification_reproducible, is_hf_path, load_prompts, print, @@ -412,30 +414,17 @@ def run(): print() print_memory_usage() - # TODO: Introduce a dedicated dataset setting for test prompts. - good_prompts_dataset = DatasetSpecification( - dataset="mlabonne/harmless_alpaca", - split="train[:5]", - column="text", - ) - - bad_prompts_dataset = DatasetSpecification( - dataset="mlabonne/harmful_behaviors", - split="train[:5]", - column="text", - ) - - print() - print(f"Loading good prompts from [bold]{good_prompts_dataset.dataset}[/]...") - good_prompts = load_prompts(settings, good_prompts_dataset) - print(f"* [bold]{len(good_prompts)}[/] prompts loaded") - - print() - print(f"Loading bad prompts from [bold]{bad_prompts_dataset.dataset}[/]...") - bad_prompts = load_prompts(settings, bad_prompts_dataset) - print(f"* [bold]{len(bad_prompts)}[/] prompts loaded") - if settings.batch_size == 0: + print() + print( + f"Loading batch size test prompts from [bold]{format_dataset_specification(settings.batch_size_test_prompts)}[/]..." + ) + batch_size_test_prompts = load_prompts( + settings, + settings.batch_size_test_prompts, + ) + print(f"* [bold]{len(batch_size_test_prompts)}[/] prompts loaded") + print() print("Determining optimal batch size...") @@ -446,7 +435,9 @@ def run(): while batch_size <= settings.max_batch_size: print(f"* Trying batch size [bold]{batch_size}[/]... ", end="") - prompts = good_prompts * math.ceil(batch_size / len(good_prompts)) + prompts = batch_size_test_prompts * math.ceil( + batch_size / len(batch_size_test_prompts) + ) prompts = prompts[:batch_size] try: @@ -487,9 +478,18 @@ def run(): print(f"* Chosen batch size: [bold]{settings.batch_size}[/]") if settings.response_prefix is None: + print() + print( + f"Loading response prefix test prompts from [bold]{format_dataset_specification(settings.response_prefix_test_prompts)}[/]..." + ) + response_prefix_test_prompts = load_prompts( + settings, + settings.response_prefix_test_prompts, + ) + print(f"* [bold]{len(response_prefix_test_prompts)}[/] prompts loaded") + print() print("Checking for common response prefix...") - prefix_check_prompts = good_prompts[:100] + bad_prompts[:100] # Detect if the model's chat template inserts a reasoning tag on its own # at the end of user's prompt (e.g. ) by using a dummy prompt. @@ -532,7 +532,7 @@ def run(): # the end of user prompt like the case above. We expect the model to # generate those tags. if settings.response_prefix is None: - responses = model.get_responses_batched(prefix_check_prompts) + responses = model.get_responses_batched(response_prefix_test_prompts) # Despite being located in os.path, commonprefix actually performs # a naive string operation without any path-specific logic, @@ -565,7 +565,7 @@ def run(): # When using a Chain-of-Thought skip, we need to check that the prefix # is actually complete (e.g. not missing a trailing newline). print("* Rechecking with prefix...") - responses = model.get_responses_batched(prefix_check_prompts) + responses = model.get_responses_batched(response_prefix_test_prompts) additional_prefix = commonprefix(responses).rstrip(" ") if additional_prefix: settings.response_prefix += additional_prefix @@ -1057,15 +1057,14 @@ def run(): # dataset was likely loaded from a local cache), and that # only built-in plugins are used (external plugins cannot # be resolved when reproducing). - dataset_specifications = [ + dataset_specifications: list[DatasetSpecification] = [ *evaluator.get_dataset_specifications(), *modifier.get_dataset_specifications(), ] is_reproducible = ( is_hf_path(settings.model) and all( - is_hf_path(specification.dataset) - and specification.commit is not None + is_dataset_specification_reproducible(specification) for specification in dataset_specifications ) and evaluator.all_scorers_reproducible() @@ -1190,6 +1189,7 @@ def run(): upload_reproduce_folder( repo_id, settings, + dataset_specifications, token, checkpoint_path=study_checkpoint_file, trial=trial, diff --git a/src/heretic/modifiers/abliteration.py b/src/heretic/modifiers/abliteration.py index fea5646..39b3a99 100644 --- a/src/heretic/modifiers/abliteration.py +++ b/src/heretic/modifiers/abliteration.py @@ -19,9 +19,9 @@ from pydantic import ( ) from torch import Tensor -from heretic.config import DatasetSpecification +from heretic.config import DatasetSpecification, SingleDatasetSpecification from heretic.modifier import Context, Modifier, Serializable -from heretic.utils import print +from heretic.utils import format_dataset_specification, print @dataclass @@ -77,7 +77,7 @@ class RowNormalization(str, Enum): class Settings(BaseModel): good_prompts: DatasetSpecification = Field( - default=DatasetSpecification( + default=SingleDatasetSpecification( dataset="mlabonne/harmless_alpaca", split="train[:400]", column="text", @@ -86,7 +86,7 @@ class Settings(BaseModel): ) bad_prompts: DatasetSpecification = Field( - default=DatasetSpecification( + default=SingleDatasetSpecification( dataset="mlabonne/harmful_behaviors", split="train[:400]", column="text", @@ -158,14 +158,14 @@ class Abliteration(Modifier[Parameters]): print() print( - f"Loading good prompts from [bold]{self.settings.good_prompts.dataset}[/]..." + f"Loading good prompts from [bold]{format_dataset_specification(self.settings.good_prompts)}[/]..." ) self.good_prompts = ctx.load_prompts(self.settings.good_prompts) print(f"* [bold]{len(self.good_prompts)}[/] prompts loaded") print() print( - f"Loading bad prompts from [bold]{self.settings.bad_prompts.dataset}[/]..." + f"Loading bad prompts from [bold]{format_dataset_specification(self.settings.bad_prompts)}[/]..." ) self.bad_prompts = ctx.load_prompts(self.settings.bad_prompts) print(f"* [bold]{len(self.bad_prompts)}[/] prompts loaded") diff --git a/src/heretic/plugin.py b/src/heretic/plugin.py index f458f13..ffdb4d5 100644 --- a/src/heretic/plugin.py +++ b/src/heretic/plugin.py @@ -13,7 +13,7 @@ from typing import Annotated, Any, TypeVar, Union, get_args, get_origin, get_typ from pydantic import BaseModel from torch import Tensor -from .config import DatasetSpecification +from .config import DatasetSpecification, SingleDatasetSpecification from .config import Settings as HereticSettings from .model import Model from .utils import Prompt, deep_merge_dicts, load_prompts @@ -347,7 +347,11 @@ class Plugin: return [] specifications = [] for value in dict(self.settings).values(): - if isinstance(value, DatasetSpecification): + if isinstance(value, SingleDatasetSpecification) or ( + isinstance(value, list) + and len(value) > 0 + and isinstance(value[0], SingleDatasetSpecification) + ): specifications.append(value) return specifications diff --git a/src/heretic/scorers/keyword_rate.py b/src/heretic/scorers/keyword_rate.py index b0b3f59..106f9fa 100644 --- a/src/heretic/scorers/keyword_rate.py +++ b/src/heretic/scorers/keyword_rate.py @@ -3,9 +3,9 @@ from pydantic import BaseModel, Field -from heretic.config import DatasetSpecification +from heretic.config import DatasetSpecification, SingleDatasetSpecification from heretic.scorer import Context, Score, Scorer -from heretic.utils import print +from heretic.utils import format_dataset_specification, print REFUSAL_MARKERS = [ "disclaimer", @@ -51,7 +51,7 @@ class Settings(BaseModel): ) prompts: DatasetSpecification = Field( - default=DatasetSpecification( + default=SingleDatasetSpecification( dataset="mlabonne/harmful_behaviors", split="test[:100]", column="text", @@ -90,7 +90,7 @@ class KeywordRate(Scorer): def init(self, ctx: Context) -> None: print() print( - f"Loading {self.settings.score_name} evaluation prompts from [bold]{self.settings.prompts.dataset}[/]..." + f"Loading {self.settings.score_name} evaluation prompts from [bold]{format_dataset_specification(self.settings.prompts)}[/]..." ) self.prompts = ctx.load_prompts(self.settings.prompts) print(f"* [bold]{len(self.prompts)}[/] prompts loaded") diff --git a/src/heretic/scorers/kl_divergence.py b/src/heretic/scorers/kl_divergence.py index 1739ddb..b68dacb 100644 --- a/src/heretic/scorers/kl_divergence.py +++ b/src/heretic/scorers/kl_divergence.py @@ -4,15 +4,15 @@ import torch.nn.functional as F from pydantic import BaseModel, Field -from heretic.config import DatasetSpecification +from heretic.config import DatasetSpecification, SingleDatasetSpecification from heretic.plugin import Context from heretic.scorer import Score, Scorer -from heretic.utils import print +from heretic.utils import format_dataset_specification, print class Settings(BaseModel): prompts: DatasetSpecification = Field( - default=DatasetSpecification( + default=SingleDatasetSpecification( dataset="mlabonne/harmless_alpaca", split="test[:100]", column="text", @@ -42,7 +42,7 @@ class KLDivergence(Scorer): def init(self, ctx: Context) -> None: print() print( - f"Loading KL divergence evaluation prompts from [bold]{self.settings.prompts.dataset}[/]..." + f"Loading KL divergence evaluation prompts from [bold]{format_dataset_specification(self.settings.prompts)}[/]..." ) self.prompts = ctx.load_prompts(self.settings.prompts) print(f"* [bold]{len(self.prompts)}[/] prompts loaded") diff --git a/src/heretic/utils.py b/src/heretic/utils.py index 08065ca..d1946f5 100644 --- a/src/heretic/utils.py +++ b/src/heretic/utils.py @@ -30,7 +30,7 @@ from psutil import Process from questionary import Question from rich.console import Console -from .config import DatasetSpecification, Settings +from .config import DatasetSpecification, Settings, SingleDatasetSpecification from .system import ( get_accelerator_info_dict, get_cpu_info_dict, @@ -169,9 +169,9 @@ def get_split_slice(split_str: str, length: int) -> tuple[int, int]: return absolute_instruction.from_, absolute_instruction.to -def load_prompts( +def _load_prompts_single( settings: Settings, - specification: DatasetSpecification, + specification: SingleDatasetSpecification, ) -> list[Prompt]: path = specification.dataset split_str = specification.split @@ -261,6 +261,44 @@ def load_prompts( ] +def load_prompts( + settings: Settings, + specification: DatasetSpecification, +) -> list[Prompt]: + if isinstance(specification, SingleDatasetSpecification): + return _load_prompts_single(settings, specification) + else: + return [ + prompt + for single_specification in specification + for prompt in _load_prompts_single(settings, single_specification) + ] + + +def format_dataset_specification(specification: DatasetSpecification) -> str: + if isinstance(specification, SingleDatasetSpecification): + return specification.dataset + else: + return ( + "\\[" + + ", ".join( + single_specification.dataset for single_specification in specification + ) + + "]" + ) + + +def is_dataset_specification_reproducible(specification: DatasetSpecification) -> bool: + if isinstance(specification, SingleDatasetSpecification): + return is_hf_path(specification.dataset) and specification.commit is not None + else: + return all( + is_hf_path(single_specification.dataset) + and single_specification.commit is not None + for single_specification in specification + ) + + 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)] @@ -367,6 +405,7 @@ def format_hf_link( def generate_reproduce_readme( settings: Settings, + dataset_specifications: list[DatasetSpecification], checkpoint_filename: str, trial: Trial | FrozenTrial, include_system_information: bool, @@ -483,6 +522,29 @@ def generate_reproduce_readme( f" --index-url https://download.pytorch.org/whl/{suffix}" ) + formatted_datasets = set() + for specification in dataset_specifications: + if isinstance(specification, SingleDatasetSpecification): + formatted_datasets.add( + format_hf_link( + specification.dataset, + specification.commit, + is_dataset=True, + ) + ) + else: + for single_specification in specification: + formatted_datasets.add( + format_hf_link( + single_specification.dataset, + single_specification.commit, + is_dataset=True, + ) + ) + dataset_lines = "\n".join( + f"- {formatted_dataset}" for formatted_dataset in sorted(formatted_datasets) + ) + trial_scores = trial.user_attrs["scores"] score_lines = "\n".join( ( @@ -502,7 +564,7 @@ This directory contains the necessary information and assets to reproduce the re ## Datasets -- TODO: Collect all datasets from scorers and modifiers. +{dataset_lines} ## Selected trial @@ -619,6 +681,7 @@ def get_file_sha256(file_path: str | Path) -> str: def create_reproduce_folder( path: Path, settings: Settings, + dataset_specifications: list[DatasetSpecification], checkpoint_path: str | Path, trial: Trial | FrozenTrial, uploaded_model_hashes: dict[str, str], @@ -667,6 +730,7 @@ def create_reproduce_folder( (reproduce_dir / "README.md").write_text( generate_reproduce_readme( settings, + dataset_specifications, checkpoint_filename, trial, include_system_information=include_system_information, @@ -683,6 +747,7 @@ def create_reproduce_folder( def upload_reproduce_folder( repo_id: str, settings: Settings, + dataset_specifications: list[DatasetSpecification], token: str, checkpoint_path: str | Path, trial: Trial | FrozenTrial, @@ -711,6 +776,7 @@ def upload_reproduce_folder( create_reproduce_folder( tmp_path, settings, + dataset_specifications, checkpoint_path=checkpoint_path, trial=trial, uploaded_model_hashes=uploaded_model_hashes, diff --git a/tests/gemma-4e/config.toml b/tests/gemma-4e/config.toml index 5c7d220..926d293 100644 --- a/tests/gemma-4e/config.toml +++ b/tests/gemma-4e/config.toml @@ -18,6 +18,18 @@ trial_index = 0 model_action = "save" save_directory = "model" +[[response_prefix_test_prompts]] +dataset = "mlabonne/harmless_alpaca" +commit = "02c6a92cfcf11bb0c387334f8146d149d65b587f" +split = "train[:5]" +column = "text" + +[[response_prefix_test_prompts]] +dataset = "mlabonne/harmful_behaviors" +commit = "01cead01398926d81f7c52bdb790ee8cf77ebba7" +split = "train[:5]" +column = "text" + [scorer.KLDivergence.prompts] dataset = "mlabonne/harmless_alpaca" commit = "02c6a92cfcf11bb0c387334f8146d149d65b587f" diff --git a/tests/minicpm5/config.toml b/tests/minicpm5/config.toml index 40aa1d6..d298a3a 100644 --- a/tests/minicpm5/config.toml +++ b/tests/minicpm5/config.toml @@ -18,6 +18,18 @@ trial_index = 0 model_action = "save" save_directory = "model" +[[response_prefix_test_prompts]] +dataset = "mlabonne/harmless_alpaca" +commit = "02c6a92cfcf11bb0c387334f8146d149d65b587f" +split = "train[:5]" +column = "text" + +[[response_prefix_test_prompts]] +dataset = "mlabonne/harmful_behaviors" +commit = "01cead01398926d81f7c52bdb790ee8cf77ebba7" +split = "train[:5]" +column = "text" + [scorer.KLDivergence.prompts] dataset = "mlabonne/harmless_alpaca" commit = "02c6a92cfcf11bb0c387334f8146d149d65b587f" diff --git a/tests/mistral-3/config.toml b/tests/mistral-3/config.toml index 3043b9d..39bf303 100644 --- a/tests/mistral-3/config.toml +++ b/tests/mistral-3/config.toml @@ -18,6 +18,18 @@ trial_index = 0 model_action = "save" save_directory = "model" +[[response_prefix_test_prompts]] +dataset = "mlabonne/harmless_alpaca" +commit = "02c6a92cfcf11bb0c387334f8146d149d65b587f" +split = "train[:5]" +column = "text" + +[[response_prefix_test_prompts]] +dataset = "mlabonne/harmful_behaviors" +commit = "01cead01398926d81f7c52bdb790ee8cf77ebba7" +split = "train[:5]" +column = "text" + [scorer.KLDivergence.prompts] dataset = "mlabonne/harmless_alpaca" commit = "02c6a92cfcf11bb0c387334f8146d149d65b587f" diff --git a/tests/qwen2.5/config.toml b/tests/qwen2.5/config.toml index aa0e610..699baba 100644 --- a/tests/qwen2.5/config.toml +++ b/tests/qwen2.5/config.toml @@ -18,6 +18,18 @@ trial_index = 0 model_action = "save" save_directory = "model" +[[response_prefix_test_prompts]] +dataset = "mlabonne/harmless_alpaca" +commit = "02c6a92cfcf11bb0c387334f8146d149d65b587f" +split = "train[:5]" +column = "text" + +[[response_prefix_test_prompts]] +dataset = "mlabonne/harmful_behaviors" +commit = "01cead01398926d81f7c52bdb790ee8cf77ebba7" +split = "train[:5]" +column = "text" + [scorer.KLDivergence.prompts] dataset = "mlabonne/harmless_alpaca" commit = "02c6a92cfcf11bb0c387334f8146d149d65b587f" diff --git a/tests/qwen3.5-moe/config.toml b/tests/qwen3.5-moe/config.toml index 50db428..a0cd6e1 100644 --- a/tests/qwen3.5-moe/config.toml +++ b/tests/qwen3.5-moe/config.toml @@ -18,6 +18,18 @@ trial_index = 0 model_action = "save" save_directory = "model" +[[response_prefix_test_prompts]] +dataset = "mlabonne/harmless_alpaca" +commit = "02c6a92cfcf11bb0c387334f8146d149d65b587f" +split = "train[:5]" +column = "text" + +[[response_prefix_test_prompts]] +dataset = "mlabonne/harmful_behaviors" +commit = "01cead01398926d81f7c52bdb790ee8cf77ebba7" +split = "train[:5]" +column = "text" + [scorer.KLDivergence.prompts] dataset = "mlabonne/harmless_alpaca" commit = "02c6a92cfcf11bb0c387334f8146d149d65b587f"