diff --git a/config.default.toml b/config.default.toml index b9cc7aa..943e19f 100644 --- a/config.default.toml +++ b/config.default.toml @@ -144,7 +144,6 @@ split = "train[:400]" column = "text" residual_plot_label = '"Harmless" prompts' residual_plot_color = "royalblue" -commit = "" # Dataset of prompts that tend to result in refusals (used for calculating refusal directions). [bad_prompts] @@ -153,18 +152,15 @@ split = "train[:400]" column = "text" residual_plot_label = '"Harmful" prompts' residual_plot_color = "darkorange" -commit = "" # Dataset of prompts that tend to not result in refusals (used for evaluating model performance). [good_evaluation_prompts] dataset = "mlabonne/harmless_alpaca" split = "test[:100]" column = "text" -commit = "" # Dataset of prompts that tend to result in refusals (used for evaluating model performance). [bad_evaluation_prompts] dataset = "mlabonne/harmful_behaviors" split = "test[:100]" column = "text" -commit = "" diff --git a/src/heretic/config.py b/src/heretic/config.py index 8a11466..ba23d95 100644 --- a/src/heretic/config.py +++ b/src/heretic/config.py @@ -31,6 +31,11 @@ class DatasetSpecification(BaseModel): description="Hugging Face dataset ID, or path to dataset on disk." ) + commit: str | None = Field( + default=None, + description="Hugging Face commit hash of the dataset.", + ) + split: str = Field(description="Portion of the dataset to use.") column: str = Field(description="Column in the dataset that contains the prompts.") @@ -59,10 +64,6 @@ class DatasetSpecification(BaseModel): default=None, description="Matplotlib color to use for the dataset in plots of residual vectors.", ) - commit: str | None = Field( - default=None, - description="Hugging Face commit hash of the dataset.", - ) class BenchmarkSpecification(BaseModel): diff --git a/src/heretic/main.py b/src/heretic/main.py index aa0adc3..d356e94 100644 --- a/src/heretic/main.py +++ b/src/heretic/main.py @@ -544,7 +544,8 @@ def run(): trial.set_user_attr("kl_divergence", kl_divergence) trial.set_user_attr("refusals", refusals) - trial.set_user_attr("total_refusal_prompts", len(evaluator.bad_prompts)) + trial.set_user_attr("base_refusals", evaluator.base_refusals) + trial.set_user_attr("n_bad_prompts", len(evaluator.bad_prompts)) return score @@ -873,13 +874,7 @@ This saves your exact configuration and system information, along with the study card.data.tags.append("decensored") card.data.tags.append("abliterated") card.text = ( - get_readme_intro( - settings, - trial, - evaluator.base_refusals, - evaluator.bad_prompts, - ) - + card.text + get_readme_intro(settings, trial) + card.text ) card.push_to_hub(repo_id, token=token) diff --git a/src/heretic/system.py b/src/heretic/system.py index e62f948..3980eb5 100644 --- a/src/heretic/system.py +++ b/src/heretic/system.py @@ -25,6 +25,7 @@ from accelerate.utils import ( def empty_cache(): """Clears the backend cache and collects garbage.""" + # 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. # See https://github.com/p-e-w/heretic/pull/17 for details. @@ -48,6 +49,7 @@ def empty_cache(): def get_nvidia_driver_version() -> str | None: """Gets the NVIDIA driver version using nvidia-smi.""" + try: output = subprocess.check_output( ["nvidia-smi", "--query-gpu=driver_version", "--format=csv,noheader"], @@ -61,6 +63,7 @@ def get_nvidia_driver_version() -> str | None: def get_amdgpu_driver_version() -> str | None: """Gets the AMD GPU (ROCm) driver and suite version info.""" + # 1. Try amd-smi (modern standard for ROCm 6.0+) try: output = subprocess.check_output( @@ -101,6 +104,7 @@ def get_amdgpu_driver_version() -> str | None: def get_xpu_driver_version() -> str | None: """Gets the Intel XPU driver version.""" + try: output = subprocess.check_output( ["xpu-smi", "discovery"], @@ -117,6 +121,7 @@ def get_xpu_driver_version() -> str | None: def get_npu_driver_version() -> str | None: """Gets the Huawei NPU driver version.""" + try: output = subprocess.check_output( ["npu-smi", "info", "-t", "board", "-i", "0"], @@ -133,6 +138,7 @@ def get_npu_driver_version() -> str | None: def get_mps_driver_version() -> str | None: """Gets the Apple Silicon (MPS) driver version via macOS version.""" + try: output = subprocess.check_output( ["sw_vers", "-productVersion"], @@ -156,6 +162,7 @@ class HereticVersionInfo: def get_heretic_version_info() -> HereticVersionInfo: """Detects version and installation source (PyPI, Git, Local) of heretic-llm.""" + package_name = "heretic-llm" origin_metadata: dict[str, Any] = {"type": "unknown"} # This package must be installed for this code to run. @@ -171,6 +178,7 @@ def get_heretic_version_info() -> HereticVersionInfo: if not direct_url_content: # Standard PyPI installation. origin_metadata["type"] = "pypi" + return HereticVersionInfo( version=base_version, origin="PyPI", @@ -178,51 +186,48 @@ def get_heretic_version_info() -> HereticVersionInfo: metadata=origin_metadata, ) - try: - data = json.loads(direct_url_content) + data = json.loads(direct_url_content) - # Check for Git source. - if "vcs_info" in data and data["vcs_info"].get("vcs") == "git": - vcs_info = data["vcs_info"] - commit_hash = vcs_info.get("commit_id", "unknown") - repo_url = data.get("url", "unknown_repo") - requested_revision = vcs_info.get("requested_revision") + # Check for Git source. + if "vcs_info" in data and data["vcs_info"].get("vcs") == "git": + vcs_info = data["vcs_info"] + commit_hash = vcs_info.get("commit_id", "unknown") + repo_url = data.get("url", "unknown_repo") + requested_revision = vcs_info.get("requested_revision") - if requested_revision: - origin_str = ( - f"Git ({repo_url}@{requested_revision} - commit: {commit_hash})" - ) - else: - origin_str = f"Git ({repo_url} @ {commit_hash})" - - origin_metadata.update( - { - "type": "git", - "url": repo_url, - "commit_hash": commit_hash, - "requested_revision": requested_revision, - } + if requested_revision: + origin_str = ( + f"Git ({repo_url}@{requested_revision} - commit: {commit_hash})" ) + else: + origin_str = f"Git ({repo_url} @ {commit_hash})" - return HereticVersionInfo( - version=base_version, - origin=origin_str, - is_standard_pypi=False, - metadata=origin_metadata, - ) + origin_metadata.update( + { + "type": "git", + "url": repo_url, + "commit_hash": commit_hash, + "requested_revision": requested_revision, + } + ) - # Check for local file/wheel directory. - if "url" in data and data["url"].startswith("file://"): - origin_metadata["type"] = "local" - return HereticVersionInfo( - version=base_version, - origin="Local", - is_standard_pypi=False, - metadata=origin_metadata, - ) + return HereticVersionInfo( + version=base_version, + origin=origin_str, + is_standard_pypi=False, + metadata=origin_metadata, + ) - except json.JSONDecodeError: - pass + # Check for local file/wheel directory. + if "url" in data and data["url"].startswith("file://"): + origin_metadata["type"] = "local" + + return HereticVersionInfo( + version=base_version, + origin="Local", + is_standard_pypi=False, + metadata=origin_metadata, + ) return HereticVersionInfo( version=base_version, @@ -234,6 +239,7 @@ def get_heretic_version_info() -> HereticVersionInfo: def get_accelerator_info_dict() -> dict[str, Any]: """Retrieves raw accelerator info (CUDA, ROCm, etc) directly into structured keys.""" + if torch.cuda.is_available(): count = torch.cuda.device_count() is_rocm = getattr(torch.version, "hip", None) is not None @@ -320,6 +326,7 @@ def get_accelerator_info_dict() -> dict[str, Any]: def get_accelerator_info(include_warnings: bool = True) -> str: """Convenience wrapper for hardware detection and console-friendly formatting.""" + info = get_accelerator_info_dict() if info["type"] is None: @@ -350,6 +357,7 @@ def get_accelerator_info(include_warnings: bool = True) -> str: def get_cpu_info_dict() -> dict[str, str | int | None]: """Gets granular CPU identifiers using the py-cpuinfo library.""" + info = cpuinfo.get_cpu_info() return { @@ -363,6 +371,7 @@ def get_cpu_info_dict() -> dict[str, str | int | None]: def get_cpu_info() -> str: """Gets the CPU brand name.""" + info = get_cpu_info_dict() parts = [] parts.append( @@ -397,12 +406,14 @@ def get_python_env_info_dict() -> dict[str, str]: def get_python_env_info() -> str: """Detects the type of Python environment (Conda, Venv, etc.) and build info.""" + info = get_python_env_info_dict() return f"{info['version']} ({info['implementation']}, {info['compiler']}) [{info['environment']}]" def get_package_version(name: str) -> str | None: """Gets the installed version of a package, stripping local suffixes like +cu128.""" + # Normalize name: pip considers hyphens and underscores equivalent. normalized_name = name.lower().replace("_", "-") version_str = importlib.metadata.version(normalized_name) @@ -411,7 +422,10 @@ def get_package_version(name: str) -> str | None: def get_requirements_dict() -> dict[str, str]: """Recursively finds all direct and transitive dependencies of heretic-llm and core libraries.""" + # We start with heretic-llm and the core compute libraries. + # PyTorch is not listed as a dependency in the heretic-llm package + # because installation is hardware-specific and must be done manually. packages_to_check = ["heretic-llm", "torch", "torchaudio", "torchvision"] visited = set() required_packages = set() diff --git a/src/heretic/utils.py b/src/heretic/utils.py index 06d47e8..e616c8c 100644 --- a/src/heretic/utils.py +++ b/src/heretic/utils.py @@ -214,6 +214,7 @@ def load_prompts( # Path is a local directory. dataset = load_dataset( path, + revision=specification.commit, split=split_str, # Don't require the number of examples (lines) per split to be pre-defined. verification_mode=VerificationMode.NO_CHECKS, @@ -269,12 +270,7 @@ def get_trial_parameters(trial: Trial) -> dict[str, str]: return params -def get_readme_intro( - settings: Settings, - trial: Trial, - base_refusals: int, - bad_prompts: list[Prompt], -) -> str: +def get_readme_intro(settings: Settings, trial: Trial) -> str: if Path(settings.model).exists(): # Hide the path, which may contain private information. model_link = "a model" @@ -304,9 +300,9 @@ def get_readme_intro( | Metric | This model | Original model ({model_link}) | | :----- | :--------: | :---------------------------: | | **KL divergence** | {trial.user_attrs["kl_divergence"]:.4f} | 0 *(by definition)* | -| **Refusals** | {trial.user_attrs["refusals"]}/{len(bad_prompts)} | {base_refusals}/{ - len(bad_prompts) - } | +| **Refusals** | {trial.user_attrs["refusals"]}/{trial.user_attrs["n_bad_prompts"]} | { + trial.user_attrs["base_refusals"] + }/{trial.user_attrs["n_bad_prompts"]} | ----- @@ -315,11 +311,13 @@ def get_readme_intro( def generate_config_toml(settings: Settings) -> str: """Serializes the full Settings object to TOML.""" + return tomli_w.dumps(settings.model_dump(exclude_none=True)) def generate_requirements_txt() -> str: """Collects direct project dependencies as a formatted string.""" + requirements = get_requirements_dict() sorted_requirements = sorted( [f"{name}=={version}" for name, version in requirements.items()], @@ -330,11 +328,30 @@ def generate_requirements_txt() -> str: def set_seed(seed: int): """Sets the seed for all RNGs.""" + random.seed(seed) np.random.seed(seed) torch.manual_seed(seed) +def format_hf_link( + name: str, + commit: str | None = None, + is_dataset: bool = False, +) -> str: + if Path(name).exists(): + return f"`{name}` (Local)" + + prefix = "datasets/" if is_dataset else "" + base_url = f"https://huggingface.co/{prefix}{name}" + link = f"[{name}]({base_url})" + if commit: + commit_url = f"{base_url}/commit/{commit}" + link += f" (Commit: [{commit[:7]}]({commit_url}))" + + return link + + def generate_reproduce_readme( settings: Settings, checkpoint_filename: str, @@ -342,7 +359,8 @@ def generate_reproduce_readme( timestamp: str | None = None, base_model_commit: str | None = None, ) -> str: - """Generates a README.md for the reproduce/ folder.""" + """Generates the contents of a README.md for the reproduce/ folder.""" + torch_version = torch.__version__ install_hint = f"pip install torch=={torch_version}" if "+" in torch_version: @@ -386,20 +404,6 @@ def generate_reproduce_readme( > This system installed `heretic-llm` from an unknown non-standard source. **Reproducibility ***cannot*** be guaranteed in this environment.** """ - def format_hf_link( - name: str, commit: str | None = None, is_dataset: bool = False - ) -> str: - if Path(name).exists(): - return f"`{name}` (Local)" - - prefix = "datasets/" if is_dataset else "" - base_url = f"https://huggingface.co/{prefix}{name}" - link = f"[{name}]({base_url})" - if commit: - commit_url = f"{base_url}/commit/{commit}" - link += f" (Commit: [{commit[:7]}]({commit_url}))" - return link - model_link = format_hf_link(settings.model, base_model_commit) dataset_info = f"""## Dataset Information @@ -475,8 +479,8 @@ This directory contains the necessary information and assets to reproduce the re ## Selected Trial - **Trial Number:** `#{trial.user_attrs["index"]}` -- **Refusal Count:** `{trial.user_attrs.get("refusals")}/{trial.user_attrs.get("total_refusal_prompts")}` -- **KL Divergence:** `{trial.user_attrs.get("kl_divergence", 0):.6f}` +- **Refusal Count:** `{trial.user_attrs["refusals"]}/{trial.user_attrs["n_bad_prompts"]}` +- **KL Divergence:** `{trial.user_attrs["kl_divergence"]:.6f}` ## System Environment @@ -503,7 +507,7 @@ This directory contains the necessary information and assets to reproduce the re 5. Verify the integrity of the reproduced files by comparing their SHA256 hashes against the manifest in `SHA256SUMS`. > [!TIP] -> To use the included Optuna study journal `{checkpoint_filename}`, place it in a `checkpoints/` directory before running `heretic` on the same model. +> To use the included Optuna study journal `{checkpoint_filename}`, place it in the checkpoints directory (usually `checkpoints/`) before running `heretic` on the same model. > [!IMPORTANT] > Make sure to install correct PyTorch version from: `{install_hint}` @@ -517,48 +521,62 @@ def generate_reproduce_json( base_model_commit: str | None = None, uploaded_model_hashes: dict[str, str] | None = None, ) -> str: - """Generates a reproduce.json file for the reproduce/ folder.""" + """Generates the contents of a reproduce.json file for the reproduce/ folder.""" + version_info = get_heretic_version_info() + data = { + "version": "1", # Version number of the reproduce.json file format, to allow for future changes. + "timestamp": timestamp, + # TODO: Remove this, it's redundant with settings! "base_model": { "id": settings.model, - "commit_hash": base_model_commit, + "commit": base_model_commit, }, "system": { - "os": {"platform": platform.platform(), "machine": platform.machine()}, - "cpu": get_cpu_info_dict(), "python": get_python_env_info_dict(), + "os": { + "platform": platform.platform(), + "machine": platform.machine(), + }, + "cpu": get_cpu_info_dict(), + "accelerator": get_accelerator_info_dict(), + }, + "environment": { "heretic": { "version": version_info.version, "is_standard_pypi": version_info.is_standard_pypi, "metadata": version_info.metadata, }, "pytorch_version": torch.__version__, - "accelerator": get_accelerator_info_dict(), + "requirements": get_requirements_dict(), }, - "requirements": get_requirements_dict(), "settings": settings.model_dump(exclude_none=True), - "trial": { - "direction_index": trial.user_attrs.get("direction_index"), - "parameters": trial.user_attrs.get("parameters"), - "metrics": { - "refusals": trial.user_attrs.get("refusals"), - "total_refusal_prompts": trial.user_attrs.get("total_refusal_prompts"), - "kl_divergence": trial.user_attrs.get("kl_divergence"), - }, + "parameters": { + "direction_index": trial.user_attrs["direction_index"], + "abliteration_parameters": trial.user_attrs["parameters"], }, - "timestamp": timestamp, - "uploaded_model_hashes": uploaded_model_hashes or {}, + "metrics": { + "kl_divergence": trial.user_attrs["kl_divergence"], + "refusals": trial.user_attrs["refusals"], + "base_refusals": trial.user_attrs["base_refusals"], + "n_bad_prompts": trial.user_attrs["n_bad_prompts"], + }, + "hashes": uploaded_model_hashes or {}, } + return json.dumps(data, indent=4) def generate_sha256sums(hashes: dict[str, str]) -> str: - """Generates a GNU Coreutils compatible SHA256SUMS file content.""" + """Generates GNU Coreutils compatible SHA256SUMS file content.""" + lines = [] + for filename, sha256 in sorted(hashes.items()): # Use '*' to indicate binary mode for model weights. lines.append(f"{sha256} *{filename}") + return "\n".join(lines) + "\n" @@ -568,7 +586,7 @@ def create_reproduce_folder( checkpoint_path: str | Path, trial: Trial, uploaded_model_hashes: dict[str, str] | None = None, -) -> None: +): reproduce_dir = path / "reproduce" reproduce_dir.mkdir(parents=True, exist_ok=True) @@ -581,17 +599,10 @@ def create_reproduce_folder( settings.good_evaluation_prompts, settings.bad_evaluation_prompts, ]: - if not Path(spec.dataset).exists(): - # Fail if the dataset is missing or unreachable. - spec.commit = huggingface_hub.dataset_info(spec.dataset).sha + spec.commit = huggingface_hub.dataset_info(spec.dataset).sha - # Fetch commit hash for the base model if it's on HF. - base_model_commit = None - if not Path(settings.model).exists(): - try: - base_model_commit = huggingface_hub.model_info(settings.model).sha - except Exception: - pass + # Fetch commit hash for the base model. + base_model_commit = huggingface_hub.model_info(settings.model).sha # Strip microseconds and timezone for a clean format. timestamp = ( @@ -599,10 +610,12 @@ def create_reproduce_folder( ) (reproduce_dir / "config.toml").write_text( - generate_config_toml(settings), encoding="utf-8" + generate_config_toml(settings), + encoding="utf-8", ) (reproduce_dir / "requirements.txt").write_text( - generate_requirements_txt(), encoding="utf-8" + generate_requirements_txt(), + encoding="utf-8", ) (reproduce_dir / "README.md").write_text( generate_reproduce_readme( @@ -616,7 +629,8 @@ def create_reproduce_folder( ) if uploaded_model_hashes: (reproduce_dir / "SHA256SUMS").write_text( - generate_sha256sums(uploaded_model_hashes), encoding="utf-8" + generate_sha256sums(uploaded_model_hashes), + encoding="utf-8", ) (reproduce_dir / "reproduce.json").write_text( generate_reproduce_json( @@ -641,22 +655,24 @@ def upload_reproduce_folder( token: str, checkpoint_path: str | Path, trial: Trial, -) -> None: +): + api = huggingface_hub.HfApi() + info = api.model_info(repo_id=repo_id, files_metadata=True, token=token) + + if not info.siblings: + raise RuntimeError("Could not fetch uploaded model hashes.") + + # For weights, we only care about safetensors. + weight_extensions = (".safetensors",) + uploaded_model_hashes = {} - try: - api = huggingface_hub.HfApi() - info = api.model_info(repo_id=repo_id, files_metadata=True, token=token) - # For weights, we only care about safetensors. - weight_extensions = (".safetensors",) - if info.siblings is not None: - for file in info.siblings: - if file.rfilename.endswith(weight_extensions): - sha256 = getattr(file, "lfs", {}).get("sha256") - if sha256: - uploaded_model_hashes[file.rfilename] = sha256 - except Exception as e: - # Fail if integrity checks cannot be completed. - raise RuntimeError(f"Could not fetch uploaded model hashes: {e}") from e + + for file in info.siblings: + if file.rfilename.endswith(weight_extensions): + sha256 = getattr(file, "lfs", {}).get("sha256") + if not sha256: + raise RuntimeError("Could not fetch uploaded model hashes.") + uploaded_model_hashes[file.rfilename] = sha256 with tempfile.TemporaryDirectory() as tmpdir: tmp_path = Path(tmpdir)