mirror of
https://github.com/p-e-w/heretic.git
synced 2026-10-02 00:41:26 -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"
|
||||
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 = ""
|
||||
|
||||
@@ -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):
|
||||
|
||||
+3
-8
@@ -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)
|
||||
|
||||
|
||||
+53
-39
@@ -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()
|
||||
|
||||
+89
-73
@@ -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)
|
||||
|
||||
Reference in New Issue
Block a user