Download model/check_checkpoint.py from OneScience-Group/W-MAE: direct link, hf CLI and curl.
- Browser
- Download file 4.19 kB
-
https://huggingface.co/OneScience-Group/W-MAE/resolve/main/model/check_checkpoint.py
- Command line
-
hf download hf://OneScience-Group/W-MAE/model/check_checkpoint.py
-
curl -L -o check_checkpoint.py https://huggingface.co/OneScience-Group/W-MAE/resolve/main/model/check_checkpoint.py
4.19 kB
| from __future__ import annotations | |
| import argparse | |
| import json | |
| from pathlib import Path | |
| import re | |
| import subprocess | |
| import sys | |
| import zipfile | |
| PROJECT_ROOT = Path(__file__).resolve().parents[1] | |
| if str(PROJECT_ROOT) not in sys.path: | |
| sys.path.insert(0, str(PROJECT_ROOT)) | |
| REQUIRED_KEYS = { | |
| "cls_token", | |
| "pos_embed", | |
| "mask_token", | |
| "decoder_pos_embed", | |
| "patch_embed.proj.weight", | |
| "decoder_embed.weight", | |
| "decoder_norm.weight", | |
| "decoder_pred.weight", | |
| } | |
| PARAMETER_KEY_PATTERN = re.compile( | |
| r"(?:decoder_blocks|blocks)\.\d+\.[A-Za-z0-9_.]*?(?:weight|bias|w1|b1|w2|b2)" | |
| r"|(?:patch_embed\.proj|decoder_embed|decoder_norm|decoder_pred)\.(?:weight|bias)" | |
| r"|(?:cls_token|pos_embed|mask_token|decoder_pos_embed)" | |
| ) | |
| def inspect_archive(path: Path) -> dict: | |
| if not zipfile.is_zipfile(path): | |
| raise ValueError(f"{path} is not a zip-based PyTorch checkpoint.") | |
| with zipfile.ZipFile(path) as archive: | |
| names = archive.namelist() | |
| data_pickle = next((name for name in names if name.endswith("/data.pkl")), None) | |
| if data_pickle is None: | |
| raise ValueError("Checkpoint archive does not contain data.pkl.") | |
| payload = archive.read(data_pickle) | |
| strings = subprocess.run( | |
| ["strings"], input=payload, stdout=subprocess.PIPE, stderr=subprocess.PIPE, check=True | |
| ).stdout.decode("utf-8", errors="ignore").splitlines() | |
| parameter_keys = { | |
| match.group(0) | |
| for line in strings | |
| for match in [PARAMETER_KEY_PATTERN.search(line)] | |
| if match | |
| } | |
| encoder_indices = sorted( | |
| {int(match.group(1)) for key in parameter_keys if (match := re.match(r"blocks\.(\d+)\.", key))} | |
| ) | |
| decoder_indices = sorted( | |
| {int(match.group(1)) for key in parameter_keys if (match := re.match(r"decoder_blocks\.(\d+)\.", key))} | |
| ) | |
| missing = sorted(REQUIRED_KEYS - parameter_keys) | |
| return { | |
| "path": str(path), | |
| "archive_entries": len(names), | |
| "parameter_key_count": len(parameter_keys), | |
| "encoder_blocks": len(encoder_indices), | |
| "decoder_blocks": len(decoder_indices), | |
| "missing_required_keys": missing, | |
| "status": "compatible_structure" if not missing and len(encoder_indices) == 12 and len(decoder_indices) == 6 else "mismatch", | |
| } | |
| def inspect_with_torch(path: Path) -> dict: | |
| try: | |
| import torch | |
| from model.w_mae import w_mae_base | |
| except (ImportError, OSError) as error: | |
| return {"status": "unavailable", "reason": str(error)} | |
| try: | |
| checkpoint = torch.load(path, map_location="cpu") | |
| except Exception as error: | |
| return {"status": "unavailable", "reason": f"checkpoint deserialization failed: {error}"} | |
| state_dict = checkpoint.get("model", checkpoint) | |
| model = w_mae_base() | |
| model_state = model.state_dict() | |
| missing = sorted(set(model_state) - set(state_dict)) | |
| unexpected = sorted(set(state_dict) - set(model_state)) | |
| shape_mismatches = { | |
| key: {"model": list(model_state[key].shape), "checkpoint": list(state_dict[key].shape)} | |
| for key in model_state.keys() & state_dict.keys() | |
| if tuple(model_state[key].shape) != tuple(state_dict[key].shape) | |
| } | |
| return { | |
| "status": "compatible" if not missing and not unexpected and not shape_mismatches else "mismatch", | |
| "missing_keys": missing, | |
| "unexpected_keys": unexpected, | |
| "shape_mismatches": shape_mismatches, | |
| } | |
| def main() -> None: | |
| parser = argparse.ArgumentParser(description="Inspect W-MAE checkpoint structure and model compatibility.") | |
| parser.add_argument("checkpoint", type=Path) | |
| parser.add_argument("--torch-check", action="store_true") | |
| args = parser.parse_args() | |
| report = {"archive": inspect_archive(args.checkpoint)} | |
| if args.torch_check: | |
| report["torch"] = inspect_with_torch(args.checkpoint) | |
| print(json.dumps(report, indent=2)) | |
| archive_ok = report["archive"]["status"] == "compatible_structure" | |
| torch_ok = not args.torch_check or report["torch"]["status"] == "compatible" | |
| raise SystemExit(0 if archive_ok and torch_ok else 1) | |
| if __name__ == "__main__": | |
| main() | |