fix: improve model commit handling

This commit is contained in:
Philipp Emanuel Weidmann
2026-04-17 08:35:37 +05:30
parent 0592090b40
commit b47a541472
4 changed files with 22 additions and 17 deletions
+5
View File
@@ -88,6 +88,11 @@ class BenchmarkSpecification(BaseModel):
class Settings(BaseSettings): class Settings(BaseSettings):
model: str = Field(description="Hugging Face model ID, or path to model on disk.") model: str = Field(description="Hugging Face model ID, or path to model on disk.")
model_commit: str | None = Field(
default=None,
description="Hugging Face commit hash of the model.",
)
evaluate_model: str | None = Field( evaluate_model: str | None = Field(
default=None, default=None,
description=( description=(
+5 -4
View File
@@ -66,7 +66,7 @@ from .utils import (
) )
def obtain_merge_strategy(settings: Settings) -> str | None: def obtain_merge_strategy(settings: Settings, model: Model) -> str | None:
""" """
Prompts the user for how to proceed with saving the model. Prompts the user for how to proceed with saving the model.
Provides info to the user if the model is quantized on memory use. Provides info to the user if the model is quantized on memory use.
@@ -95,7 +95,8 @@ def obtain_merge_strategy(settings: Settings) -> str | None:
settings.model, settings.model,
device_map="meta", device_map="meta",
torch_dtype=torch.bfloat16, torch_dtype=torch.bfloat16,
trust_remote_code=True, trust_remote_code=model.trusted_models.get(settings.model),
**model.revision_kwargs,
) )
footprint_bytes = meta_model.get_memory_footprint() footprint_bytes = meta_model.get_memory_footprint()
footprint_gb = footprint_bytes / (1024**3) footprint_gb = footprint_bytes / (1024**3)
@@ -752,7 +753,7 @@ def run():
if not save_directory: if not save_directory:
continue continue
strategy = obtain_merge_strategy(settings) strategy = obtain_merge_strategy(settings, model)
if strategy is None: if strategy is None:
continue continue
@@ -803,7 +804,7 @@ def run():
continue continue
private = visibility == "Private" private = visibility == "Private"
strategy = obtain_merge_strategy(settings) strategy = obtain_merge_strategy(settings, model)
if strategy is None: if strategy is None:
continue continue
+8
View File
@@ -62,12 +62,17 @@ class Model:
self.settings = settings self.settings = settings
self.needs_reload = False self.needs_reload = False
self.revision_kwargs = {}
if settings.model_commit is not None:
self.revision_kwargs["revision"] = settings.model_commit
print() print()
print(f"Loading model [bold]{settings.model}[/]...") print(f"Loading model [bold]{settings.model}[/]...")
self.tokenizer = AutoTokenizer.from_pretrained( self.tokenizer = AutoTokenizer.from_pretrained(
settings.model, settings.model,
trust_remote_code=settings.trust_remote_code, trust_remote_code=settings.trust_remote_code,
**self.revision_kwargs,
) )
# Fallback for tokenizers that don't declare a special pad token. # Fallback for tokenizers that don't declare a special pad token.
@@ -108,6 +113,7 @@ class Model:
device_map=settings.device_map, device_map=settings.device_map,
max_memory=self.max_memory, max_memory=self.max_memory,
trust_remote_code=self.trusted_models.get(settings.model), trust_remote_code=self.trusted_models.get(settings.model),
**self.revision_kwargs,
**extra_kwargs, **extra_kwargs,
) )
@@ -257,6 +263,7 @@ class Model:
torch_dtype=self.model.dtype, torch_dtype=self.model.dtype,
device_map="cpu", device_map="cpu",
trust_remote_code=self.trusted_models.get(self.settings.model), trust_remote_code=self.trusted_models.get(self.settings.model),
**self.revision_kwargs,
) )
# Apply LoRA adapters to the CPU model # Apply LoRA adapters to the CPU model
@@ -318,6 +325,7 @@ class Model:
device_map=self.settings.device_map, device_map=self.settings.device_map,
max_memory=self.max_memory, max_memory=self.max_memory,
trust_remote_code=self.trusted_models.get(self.settings.model), trust_remote_code=self.trusted_models.get(self.settings.model),
**self.revision_kwargs,
**extra_kwargs, **extra_kwargs,
) )
+4 -13
View File
@@ -357,7 +357,6 @@ def generate_reproduce_readme(
checkpoint_filename: str, checkpoint_filename: str,
trial: Trial, trial: Trial,
timestamp: str | None = None, timestamp: str | None = None,
base_model_commit: str | None = None,
) -> str: ) -> str:
"""Generates the contents of a README.md for the reproduce/ folder.""" """Generates the contents of a README.md for the reproduce/ folder."""
@@ -404,7 +403,7 @@ 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.**
""" """
model_link = format_hf_link(settings.model, base_model_commit) model_link = format_hf_link(settings.model, settings.model_commit)
dataset_info = f"""## Dataset Information dataset_info = f"""## Dataset Information
- **Good Prompts:** {format_hf_link(settings.good_prompts.dataset, settings.good_prompts.commit, is_dataset=True)} - **Good Prompts:** {format_hf_link(settings.good_prompts.dataset, settings.good_prompts.commit, is_dataset=True)}
@@ -518,7 +517,6 @@ def generate_reproduce_json(
settings: Settings, settings: Settings,
trial: Trial, trial: Trial,
timestamp: str | None = None, timestamp: 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 the contents of a reproduce.json file for the reproduce/ folder.""" """Generates the contents of a reproduce.json file for the reproduce/ folder."""
@@ -528,11 +526,6 @@ def generate_reproduce_json(
data = { data = {
"version": "1", # Version number of the reproduce.json file format, to allow for future changes. "version": "1", # Version number of the reproduce.json file format, to allow for future changes.
"timestamp": timestamp, "timestamp": timestamp,
# TODO: Remove this, it's redundant with settings!
"base_model": {
"id": settings.model,
"commit": base_model_commit,
},
"system": { "system": {
"python": get_python_env_info_dict(), "python": get_python_env_info_dict(),
"os": { "os": {
@@ -592,6 +585,9 @@ def create_reproduce_folder(
checkpoint_filename = Path(checkpoint_path).name checkpoint_filename = Path(checkpoint_path).name
# Fetch commit hash for the base model.
settings.model_commit = huggingface_hub.model_info(settings.model).sha
# Fetch commit hashes for all HF datasets to ensure reproducibility. # Fetch commit hashes for all HF datasets to ensure reproducibility.
for spec in [ for spec in [
settings.good_prompts, settings.good_prompts,
@@ -601,9 +597,6 @@ def create_reproduce_folder(
]: ]:
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.
base_model_commit = huggingface_hub.model_info(settings.model).sha
# Strip microseconds and timezone for a clean format. # Strip microseconds and timezone for a clean format.
timestamp = ( timestamp = (
datetime.now(timezone.utc).replace(microsecond=0, tzinfo=None).isoformat() datetime.now(timezone.utc).replace(microsecond=0, tzinfo=None).isoformat()
@@ -623,7 +616,6 @@ def create_reproduce_folder(
checkpoint_filename, checkpoint_filename,
trial, trial,
timestamp=timestamp, timestamp=timestamp,
base_model_commit=base_model_commit,
), ),
encoding="utf-8", encoding="utf-8",
) )
@@ -637,7 +629,6 @@ def create_reproduce_folder(
settings, settings,
trial, trial,
timestamp=timestamp, timestamp=timestamp,
base_model_commit=base_model_commit,
uploaded_model_hashes=uploaded_model_hashes, uploaded_model_hashes=uploaded_model_hashes,
), ),
encoding="utf-8", encoding="utf-8",