diff --git a/python/sglang/srt/utils/model_file_verifier.py b/python/sglang/srt/utils/model_file_verifier.py index 535d2a941..dad7703be 100644 --- a/python/sglang/srt/utils/model_file_verifier.py +++ b/python/sglang/srt/utils/model_file_verifier.py @@ -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: diff --git a/test/registered/utils/test_model_file_verifier.py b/test/registered/utils/test_model_file_verifier.py index 9d7aa678b..4c2f3d4c8 100644 --- a/test/registered/utils/test_model_file_verifier.py +++ b/test/registered/utils/test_model_file_verifier.py @@ -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)