from __future__ import annotations import hashlib import json from pathlib import Path import torch from safetensors.torch import load_file from modeling_supplychain_jepa import ( ModelConfig, SemanticConceptSlotPredictor, StatePoolPredictor, SupplyChainJEPACore, ) def _sha256(path: Path) -> str: digest = hashlib.sha256() with path.open("rb") as handle: for block in iter(lambda: handle.read(1024 * 1024), b""): digest.update(block) return digest.hexdigest() def verify_release(directory: str | Path) -> None: root = Path(directory) for line in (root / "CHECKSUMS.sha256").read_text().splitlines(): expected, name = line.split(" ", 1) actual = _sha256(root / name) if actual != expected: raise ValueError(f"Checksum mismatch for {name}: {actual} != {expected}") def load_supplychain_jepa(directory: str | Path, device: str = "cpu", verify: bool = True): root = Path(directory) if verify: verify_release(root) config = json.loads((root / "config.json").read_text()) model_config = ModelConfig(**config["model"]) model = SupplyChainJEPACore(model_config) model.load_state_dict(load_file(root / config["weights"]["core"], device="cpu"), strict=True) raw_aux = load_file(root / config["weights"]["state_auxiliary"], device="cpu") pool_cfg = config["state_pool_predictor"] state_pool_predictor = None if pool_cfg["enabled"]: state_pool_predictor = StatePoolPredictor( model_config.d_model, query_vocab_size=max(pool_cfg["query_vocabulary"].values()) + 1, mask_vocab_size=max(pool_cfg["mask_vocabulary"].values()) + 1, hidden_multiplier=pool_cfg["hidden_multiplier"], dropout=pool_cfg["dropout"], ) state_pool_predictor.load_state_dict({k.split(".", 1)[1]: v for k, v in raw_aux.items() if k.startswith("state_pool_predictor.")}, strict=True) slot_cfg = config["semantic_concept_slot_predictor"] semantic_concept_slot_predictor = None if slot_cfg["enabled"]: slot_state = {k.split(".", 1)[1]: v for k, v in raw_aux.items() if k.startswith("semantic_concept_slot_predictor.")} semantic_concept_slot_predictor = SemanticConceptSlotPredictor( model_config.d_model, concept_count=slot_state["concept_queries"].shape[1], nhead=model_config.nhead, decoder_layers=slot_cfg["decoder_layers"], dropout=slot_cfg["dropout"], ) semantic_concept_slot_predictor.load_state_dict(slot_state, strict=True) grounding_state = {k.split(".", 1)[1]: v for k, v in raw_aux.items() if k.startswith("grounding_head.")} grounding_head = torch.nn.Linear(model_config.d_model, grounding_state["weight"].shape[0]) grounding_head.load_state_dict(grounding_state, strict=True) modules = [model, state_pool_predictor, semantic_concept_slot_predictor, grounding_head] for module in modules: if module is not None: module.to(device).eval() for parameter in module.parameters(): parameter.requires_grad_(False) return { "model": model, "state_pool_predictor": state_pool_predictor, "semantic_concept_slot_predictor": semantic_concept_slot_predictor, "grounding_head": grounding_head, "config": config, "tensorizer": json.loads((root / "tensorizer.json").read_text()), "schema": json.loads((root / "schema.json").read_text()), "normalization": json.loads((root / "normalization.json").read_text()), }