fix: various cleanups and improvements for the reproducibility system

This commit is contained in:
Philipp Emanuel Weidmann
2026-04-15 11:13:58 +05:30
parent e2c74bfb3c
commit b46396b785
5 changed files with 150 additions and 128 deletions
-4
View File
@@ -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 = ""
+5 -4
View File
@@ -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
View File
@@ -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
View File
@@ -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
View File
@@ -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)