Add file size hints in bitwise model file verifier (#16735)

This commit is contained in:
fzyzcjy
2026-01-08 22:28:52 +08:00
committed by GitHub
parent fb7609f1dd
commit b6e8a0d851
2 changed files with 135 additions and 45 deletions

View File

@@ -14,12 +14,49 @@ import argparse
import fnmatch
import hashlib
import json
import warnings
from concurrent.futures import ThreadPoolExecutor
from dataclasses import asdict, dataclass
from pathlib import Path
from typing import Dict, List, Optional, Tuple
# ======== Data Format ========
@dataclass
class FileInfo:
sha256: str
size: int
@dataclass
class Manifest:
files: Dict[str, FileInfo]
@classmethod
def from_dict(cls, data: dict) -> "Manifest":
if "checksums" in data:
warnings.warn(
"The 'checksums' format is deprecated. "
"Please regenerate with the latest version to use the new 'files' format.",
DeprecationWarning,
stacklevel=3,
)
return cls(
files={
k: FileInfo(sha256=v, size=-1) for k, v in data["checksums"].items()
}
)
return cls(files={k: FileInfo(**v) for k, v in data["files"].items()})
def to_dict(self) -> dict:
return asdict(self)
# ======== Constants ========
IGNORE_PATTERNS = [
"checksums.json",
".DS_Store",
"*.lock",
".gitattributes",
@@ -37,21 +74,24 @@ IGNORE_PATTERNS = [
def verify(*, model_path: str, checksums_source: str, max_workers: int = 4) -> None:
model_path = Path(model_path).resolve()
expected = _load_checksums(checksums_source)
actual = _compute_checksums_from_folder(
model_path=model_path, filenames=list(expected.keys()), max_workers=max_workers
actual = _compute_manifest_from_folder(
model_path=model_path,
filenames=list(expected.files.keys()),
max_workers=max_workers,
)
_compare_checksums(expected=expected, actual=actual)
print(f"[ModelFileVerifier] All {len(expected)} files verified successfully.")
_compare_manifests(expected=expected, actual=actual)
print(f"[ModelFileVerifier] All {len(expected.files)} files verified successfully.")
def _compare_checksums(*, expected: Dict[str, str], actual: Dict[str, str]) -> None:
def _compare_manifests(*, expected: Manifest, actual: Manifest) -> None:
errors = []
for filename, expected_hash in expected.items():
if filename not in actual:
errors.append(f"{filename}: missing")
elif actual[filename] != expected_hash:
for filename, exp in expected.files.items():
if filename not in actual.files:
errors.append(f"{filename}: missing (expected size={exp.size})")
elif actual.files[filename].sha256 != exp.sha256:
act = actual.files[filename]
errors.append(
f"{filename}: mismatch (expected={expected_hash[:16]}..., actual={actual[filename][:16]}...)"
f"{filename}: mismatch (expected={exp.sha256[:16]}... size={exp.size}, actual={act.sha256[:16]}... size={act.size})"
)
if errors:
@@ -63,25 +103,26 @@ def _compare_checksums(*, expected: Dict[str, str], actual: Dict[str, str]) -> N
def generate_checksums(
*, source: str, output_path: str, max_workers: int = 4
) -> Dict[str, str]:
) -> Manifest:
if Path(source).is_dir():
model_path = Path(source).resolve()
files = _discover_files(model_path)
if not files:
raise IntegrityError(f"No model files found in {model_path}")
checksums = _compute_checksums_from_folder(
manifest = _compute_manifest_from_folder(
model_path=model_path, filenames=files, max_workers=max_workers
)
else:
checksums = _load_checksums_from_hf(repo_id=source)
manifest = Manifest(files=_load_file_infos_from_hf(repo_id=source))
output = {"checksums": checksums}
Path(output_path).write_text(json.dumps(output, indent=2, sort_keys=True))
Path(output_path).write_text(
json.dumps(manifest.to_dict(), indent=2, sort_keys=True)
)
print(
f"[ModelFileVerifier] Generated checksums for {len(checksums)} files -> {output_path}"
f"[ModelFileVerifier] Generated checksums for {len(manifest.files)} files -> {output_path}"
)
return checksums
return manifest
def _discover_files(model_path: Path) -> List[str]:
@@ -97,31 +138,31 @@ def _discover_files(model_path: Path) -> List[str]:
# ======== Load Checksums ========
def _load_checksums(source: str) -> Dict[str, str]:
def _load_checksums(source: str) -> Manifest:
if Path(source).is_file():
data = json.loads(Path(source).read_text())
return data["checksums"]
return _load_checksums_from_hf(repo_id=source)
return Manifest.from_dict(data)
return Manifest(files=_load_file_infos_from_hf(repo_id=source))
def _load_checksums_from_hf(*, repo_id: str) -> Dict[str, str]:
def _load_file_infos_from_hf(*, repo_id: str) -> Dict[str, FileInfo]:
from huggingface_hub import HfFileSystem
fs = HfFileSystem()
files = fs.ls(repo_id, detail=True)
checksums = dict(
r
for r in map(lambda f: _get_filename_and_checksum_from_hf_file(fs, f), files)
if r
file_infos = dict(
r for r in map(lambda f: _get_filename_and_info_from_hf_file(fs, f), files) if r
)
if not checksums:
if not file_infos:
raise IntegrityError(f"No files found in HF repo {repo_id}.")
return checksums
return file_infos
def _get_filename_and_checksum_from_hf_file(fs, file_info):
def _get_filename_and_info_from_hf_file(
fs, file_info
) -> Optional[Tuple[str, FileInfo]]:
if file_info.get("type") != "file":
return None
@@ -129,31 +170,35 @@ def _get_filename_and_checksum_from_hf_file(fs, file_info):
if any(fnmatch.fnmatch(filename, pat) for pat in IGNORE_PATTERNS):
return None
size = file_info.get("size", -1)
lfs_info = file_info.get("lfs")
if lfs_info and "sha256" in lfs_info:
return filename, lfs_info["sha256"]
return filename, FileInfo(sha256=lfs_info["sha256"], size=size)
if "sha256" in file_info:
return filename, file_info["sha256"]
return filename, FileInfo(sha256=file_info["sha256"], size=size)
content = fs.read_bytes(file_info.get("name", ""))
return filename, hashlib.sha256(content).hexdigest()
return filename, FileInfo(
sha256=hashlib.sha256(content).hexdigest(), size=len(content)
)
# ======== Compute Checksums ========
def _compute_checksums_from_folder(
def _compute_manifest_from_folder(
*, model_path: Path, filenames: List[str], max_workers: int
) -> Dict[str, str]:
) -> Manifest:
from tqdm import tqdm
def compute_one(filename: str) -> Tuple[str, Optional[str]]:
def compute_one(filename: str) -> Tuple[str, Optional[FileInfo]]:
full_path = model_path / filename
if not full_path.exists():
return filename, None
sha256 = compute_sha256(file_path=full_path)
return filename, sha256
size = full_path.stat().st_size
return filename, FileInfo(sha256=sha256, size=size)
with ThreadPoolExecutor(max_workers=max_workers) as executor:
results = list(
@@ -164,7 +209,7 @@ def _compute_checksums_from_folder(
)
)
return {k: v for k, v in results if v is not None}
return Manifest(files={k: v for k, v in results if v is not None})
def compute_sha256(*, file_path) -> str:

View File

@@ -6,6 +6,7 @@ import subprocess
import sys
import tempfile
import unittest
import warnings
from contextlib import nullcontext
from io import StringIO
@@ -111,11 +112,55 @@ class TestModelFileVerifier(_FakeModelTestCase):
)
checksums_file = os.path.join(self.test_dir, "checksums.json")
checksums = generate_checksums(
result = generate_checksums(
source=self.test_dir, output_path=checksums_file, max_workers=4
)
self.assertGreaterEqual(len(checksums), 10)
self.assertGreaterEqual(len(result.files), 10)
def test_generated_json_snapshot(self):
checksums_file = os.path.join(self.test_dir, "checksums.json")
generate_checksums(source=self.test_dir, output_path=checksums_file)
with open(checksums_file) as f:
data = json.load(f)
expected = {
"files": {
"config.json": {
"sha256": "81dddc8c379baae137d99d24c5fa081d3a5ce52b6a221ddc22fe364711f8beaf",
"size": 23,
},
"model.safetensors": {
"sha256": "eb0c73a48a89fefb6b68dd41af830d75610c885135eac99139373b04705d05f3",
"size": 2500,
},
"tokenizer.json": {
"sha256": "4e3043229142b64d998563bc543ce034e0a2251af5d404995e3afcb8ce8850df",
"size": 18,
},
}
}
self.assertEqual(data, expected)
def test_legacy_checksums_format_deprecated(self):
legacy_data = {
"checksums": {
"model.safetensors": "eb0c73a48a89fefb6b68dd41af830d75610c885135eac99139373b04705d05f3",
"config.json": "81dddc8c379baae137d99d24c5fa081d3a5ce52b6a221ddc22fe364711f8beaf",
"tokenizer.json": "4e3043229142b64d998563bc543ce034e0a2251af5d404995e3afcb8ce8850df",
}
}
legacy_file = os.path.join(self.test_dir, "legacy_checksums.json")
with open(legacy_file, "w") as f:
json.dump(legacy_data, f)
with warnings.catch_warnings(record=True) as w:
warnings.simplefilter("always")
verify(model_path=self.test_dir, checksums_source=legacy_file)
self.assertEqual(len(w), 1)
self.assertTrue(issubclass(w[0].category, DeprecationWarning))
self.assertIn("deprecated", str(w[0].message).lower())
# ======== CLI Tests ========
@@ -144,8 +189,8 @@ class TestModelFileVerifierCLI(_FakeModelTestCase):
with open(checksums_file) as f:
data = json.load(f)
self.assertIn("checksums", data)
self.assertEqual(len(data["checksums"]), 3)
self.assertIn("files", data)
self.assertEqual(len(data["files"]), 3)
def test_cli_verify_success(self):
checksums_file = os.path.join(self.test_dir, "checksums.json")
@@ -204,12 +249,12 @@ class TestModelFileVerifierHF(_RealModelTestCase):
def test_generate_checksums_from_hf(self):
checksums_file = os.path.join(self.test_dir, "checksums.json")
checksums = generate_checksums(source=MODEL_NAME, output_path=checksums_file)
result = generate_checksums(source=MODEL_NAME, output_path=checksums_file)
self.assertTrue(os.path.exists(checksums_file))
self.assertGreater(len(checksums), 0)
for filename, sha256 in checksums.items():
self.assertEqual(len(sha256), 64)
self.assertGreater(len(result.files), 0)
for filename, file_info in result.files.items():
self.assertEqual(len(file_info.sha256), 64)
def test_verify_with_hf_checksums_source(self):
verify(model_path=self.test_dir, checksums_source=MODEL_NAME)