fix: fix issues in automatic reproduction system (#352)

* fix: Check if a model is gated / accessible

* fix: handle unknown gated models

* feat: Auto install requirements

* simplify

* Revert "simplify"

This reverts commit 10287926e9.

* Revert "feat: Auto install requirements"

This reverts commit f4be1abd04.

* fix: Seed pytorch method

* reference, style

* simplify token

* feat: Export strategy in reproduce.json, v2

* style: Name

* simplify export strategy

* style: Rename

* enumeration

* maybe remove seed as well

* fix: don't lock settings with permanent strategy

* simplify no choice, use try/finally block
This commit is contained in:
Vinay-Umrethe
2026-06-07 14:10:45 +05:30
committed by GitHub
parent bd3bf2c726
commit b9d5b45bbf
6 changed files with 59 additions and 27 deletions
-4
View File
@@ -123,10 +123,6 @@ n_trials = 200
# Number of trials that use random sampling for the purpose of exploration.
n_startup_trials = 60
# Random seed for reproducible optimization. Set to an integer to enable.
# Applies to Python's random module, NumPy, PyTorch, and Optuna.
# seed = 75
# Directory to save and load study progress to/from.
study_checkpoint_dir = "checkpoints"
+10
View File
@@ -32,6 +32,11 @@ class RowNormalization(str, Enum):
FULL = "full"
class ExportStrategy(str, Enum):
MERGE = "merge"
ADAPTER = "adapter"
class DatasetSpecification(BaseModel):
dataset: str = Field(
description="Hugging Face dataset ID, or path to dataset on disk."
@@ -412,6 +417,11 @@ class Settings(BaseSettings):
description="Maximum size for individual safetensors files generated when exporting a model.",
)
export_strategy: ExportStrategy | None = Field(
default=None,
description='How to export the model: "merge", "adapter", or unset to prompt the user.',
)
refusal_markers: list[str] = Field(
default=[
"sorry",
+32 -21
View File
@@ -62,7 +62,7 @@ from rich.table import Table
from rich.traceback import install
from .analyzer import Analyzer
from .config import QuantizationMethod
from .config import ExportStrategy, QuantizationMethod
from .evaluator import Evaluator
from .model import AbliterationParameters, Model, get_model_class
from .reproduce import (
@@ -88,13 +88,19 @@ from .utils import (
)
def obtain_merge_strategy(settings: Settings, model: Model) -> str | None:
def obtain_export_strategy(
settings: Settings,
model: Model,
) -> ExportStrategy | None:
"""
Prompts the user for how to proceed with saving the model.
Gets the export strategy from settings or prompts the user.
Provides info to the user if the model is quantized on memory use.
Returns "merge", "adapter", or None (if cancelled/invalid).
Returns an export strategy, or None if cancelled.
"""
if settings.export_strategy is not None:
return settings.export_strategy
if settings.quantization == QuantizationMethod.BNB_4BIT:
print()
print(
@@ -148,11 +154,11 @@ def obtain_merge_strategy(settings: Settings, model: Model) -> str | None:
if settings.quantization == QuantizationMethod.NONE
else " (requires sufficient RAM)"
),
value="merge",
value=ExportStrategy.MERGE,
),
Choice(
title="Save LoRA adapter only (can be merged later)",
value="adapter",
value=ExportStrategy.ADAPTER,
),
],
)
@@ -224,7 +230,7 @@ def run():
# FIXME: "Reproduction"/"reproducibility" name inconsistency!
reproduction_information = load_reproduction_information(settings.reproduce)
if reproduction_information["version"] not in ["1"]:
if reproduction_information["version"] not in ["1", "2"]:
print(
(
f"[red]Unsupported file format version: [bold]{reproduction_information['version']}[/].[/] "
@@ -865,11 +871,11 @@ def run():
if not save_directory:
continue
strategy = obtain_merge_strategy(settings, model)
strategy = obtain_export_strategy(settings, model)
if strategy is None:
continue
if strategy == "adapter":
if strategy == ExportStrategy.ADAPTER:
print("Saving LoRA adapter...")
model.model.save_pretrained(
save_directory,
@@ -923,7 +929,7 @@ def run():
continue
private = visibility == "Private"
strategy = obtain_merge_strategy(settings, model)
strategy = obtain_export_strategy(settings, model)
if strategy is None:
continue
@@ -973,7 +979,7 @@ def run():
else:
reproducibility_information = "none"
if strategy == "adapter":
if strategy == ExportStrategy.ADAPTER:
print("Uploading LoRA adapter...")
model.model.push_to_hub(
repo_id,
@@ -1036,17 +1042,22 @@ def run():
# Set the number of trials to the number of actual completed trials
# for the reproduction configuration.
settings.n_trials = count_completed_trials()
current_export_strategy = settings.export_strategy
settings.export_strategy = strategy
upload_reproduce_folder(
repo_id,
settings,
token,
checkpoint_path=study_checkpoint_file,
trial=trial,
include_system_information=(
reproducibility_information == "full"
),
)
try:
upload_reproduce_folder(
repo_id,
settings,
token,
checkpoint_path=study_checkpoint_file,
trial=trial,
include_system_information=(
reproducibility_information == "full"
),
)
finally:
settings.export_strategy = current_export_strategy
print(f"Model uploaded to [bold]{repo_id}[/].")
+4
View File
@@ -539,6 +539,10 @@ class Model:
W = W - W_org
# Use a low-rank SVD to get an approximation of the matrix.
r = self.peft_config.r
# svd_lowrank is randomized:
# https://github.com/pytorch/pytorch/blob/20919052303c0b5ba87f8bf7e19237dc33ab09d3/torch/_lowrank.py#L108-L109
# Reseed immediately before the call so restoring a trial is independent of RNG history.
torch.manual_seed(self.settings.seed)
U, S, Vh = torch.svd_lowrank(W, q=2 * r + 4, niter=6)
# Truncate it to the part we want to store in the LoRA adapter.
# Note: svd_lowrank actually returns V, so transpose it to get Vh.
+12 -1
View File
@@ -14,7 +14,11 @@ from urllib.request import urlopen
import cpuinfo
import torch
from huggingface_hub import HfApi, hf_hub_download
from huggingface_hub.utils import disable_progress_bars, enable_progress_bars
from huggingface_hub.utils import (
GatedRepoError,
disable_progress_bars,
enable_progress_bars,
)
from questionary import Choice
from rich.table import Table
@@ -37,6 +41,7 @@ def collect_reproducibles(path: str):
models = api.list_models(
filter=["heretic", "reproducible"],
sort="created_at",
expand=["gated", "tags"],
)
found = 0
@@ -51,6 +56,12 @@ def collect_reproducibles(path: str):
if model.tags is not None and "gguf" in model.tags:
continue
if model.gated:
try:
api.auth_check(model.id, repo_type="model")
except GatedRepoError:
continue
print(f"[bold]{model.id}[/]...", end="")
user, repository = model.id.split("/")
+1 -1
View File
@@ -547,7 +547,7 @@ def generate_reproduce_json(
version_info = get_heretic_version_info()
data = {
"version": "1", # Version number of the reproduce.json file format, to allow for future changes.
"version": "2", # Version number of the reproduce.json file format, to allow for future changes.
"timestamp": timestamp,
"system": None, # Defined here to preserve insertion order.
"environment": {