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