diff --git a/src/modelinfo/parsers/gguf.py b/src/modelinfo/parsers/gguf.py index 3af2fb4..c33eb06 100644 --- a/src/modelinfo/parsers/gguf.py +++ b/src/modelinfo/parsers/gguf.py @@ -10,36 +10,46 @@ 29: "IQ1_M", 30: "BF16", 31: "Q4_0_4_4", 32: "Q4_0_4_8", 33: "Q4_0_8_8", } +def _read_exact(f: Any, size: int) -> bytes: + data = bytearray() + while len(data) < size: + chunk = f.read(size - len(data)) + if not chunk: + raise EOFError("Unexpected end of GGUF header") + data.extend(chunk) + return bytes(data) + + def _read_gguf_value(f: Any, val_type: int) -> Any: if val_type == 0: - return struct.unpack(" Dict[str, Any]: def _parse_gguf_header_from_stream(f: Any) -> Dict[str, Any]: tensors: Dict[str, Any] = {} - magic = f.read(4) + magic = _read_exact(f, 4) if magic != b"GGUF": raise ValueError("Invalid GGUF file: Magic bytes missing.") - version = struct.unpack(" DummyStorage: + # ZIP checkpoints refer to tensor storage via a five-element ID. + # Only reconstruct metadata using our existing dummy allowlist; never + # load storage bytes, import torch, or invoke a pickle-provided callable. + if not isinstance(saved_id, tuple) or len(saved_id) != 5: + raise pickle.UnpicklingError("Invalid storage persistent ID") + kind, storage_type, _, _, _ = saved_id + if (kind != "storage" or not isinstance(storage_type, type) + or not issubclass(storage_type, DummyStorage) + or storage_type is DummyStorage): + raise pickle.UnpicklingError("Unsupported storage persistent ID") + return storage_type() + def find_class(self, module: str, name: str) -> Any: if module in self.ALLOWED_MODULES and name in self.ALLOWED_MODULES[module]: if name == "OrderedDict": @@ -84,6 +112,11 @@ def parse_pytorch_header(path: str) -> Dict[str, Any]: data = unpickler.load() if isinstance(data, dict): + for key in ("state_dict", "model_state_dict"): + candidate = data.get(key) + if isinstance(candidate, dict) and "shape" not in candidate: + data = candidate + break for k, v in data.items(): if isinstance(v, dict) and "shape" in v: tensors[k] = v diff --git a/src/modelinfo/parsers/safetensors.py b/src/modelinfo/parsers/safetensors.py index 2e7d705..28bb39b 100644 --- a/src/modelinfo/parsers/safetensors.py +++ b/src/modelinfo/parsers/safetensors.py @@ -63,9 +63,12 @@ def parse_safetensors_header(path: str) -> dict[str, Any]: total_size += os.path.getsize(shard_path) try: shard_header = _read_single_header(shard_path) - for k, v in shard_header.items(): - if k != "__metadata__": - tensors[k] = v + for name, assigned_shard in weight_map.items(): + if assigned_shard != shard: + continue + if name not in shard_header: + raise ValueError(f"Indexed tensor {name!r} is missing from shard {shard!r}") + tensors[name] = shard_header[name] except FileNotFoundError: missing_shards += 1 diff --git a/tests/test_binary_parsers.py b/tests/test_binary_parsers.py new file mode 100644 index 0000000..9ec0bf4 --- /dev/null +++ b/tests/test_binary_parsers.py @@ -0,0 +1,125 @@ +"""Small generated checkpoint fixtures; no model downloads or torch dependency.""" + +import json +import struct +import zipfile +from pickle import UnpicklingError + +import pytest + +from modelinfo.parsers.pytorch import parse_pytorch_header +from modelinfo.parsers.safetensors import parse_safetensors_header + + +def write_safetensors(path, header): + payload = json.dumps(header).encode() + path.write_bytes(struct.pack("