| 1 | import json |
| 2 | |
| 3 | import safetensors |
| 4 | import torch |
| 5 | |
| 6 | from ltx_core.loader.primitives import StateDict, StateDictLoader |
| 7 | from ltx_core.loader.sd_ops import SDOps |
| 8 | |
| 9 | |
| 10 | class SafetensorsStateDictLoader(StateDictLoader): |
| 11 | """ |
| 12 | Loads weights from safetensors files without metadata support. |
| 13 | Use this for loading raw weight files. For model files that include |
| 14 | configuration metadata, use SafetensorsModelStateDictLoader instead. |
| 15 | """ |
| 16 | |
| 17 | def metadata(self, path: str) -> dict: |
| 18 | raise NotImplementedError("Not implemented") |
| 19 | |
| 20 | def load(self, path: str | list[str], sd_ops: SDOps, device: torch.device | None = None) -> StateDict: |
| 21 | """ |
| 22 | Load state dict from path or paths (for sharded model storage) and apply sd_ops |
| 23 | """ |
| 24 | sd = {} |
| 25 | size = 0 |
| 26 | dtype = set() |
| 27 | device = device or torch.device("cpu") |
| 28 | model_paths = path if isinstance(path, list) else [path] |
| 29 | for shard_path in model_paths: |
| 30 | with safetensors.safe_open(shard_path, framework="pt", device=str(device)) as f: |
| 31 | safetensor_keys = f.keys() |
| 32 | for name in safetensor_keys: |
| 33 | expected_name = name if sd_ops is None else sd_ops.apply_to_key(name) |
| 34 | if expected_name is None: |
| 35 | continue |
| 36 | value = f.get_tensor(name).to(device=device, non_blocking=True, copy=False) |
| 37 | key_value_pairs = ((expected_name, value),) |
| 38 | if sd_ops is not None: |
| 39 | key_value_pairs = sd_ops.apply_to_key_value(expected_name, value) |
| 40 | for key, value in key_value_pairs: |
| 41 | size += value.nbytes |
| 42 | dtype.add(value.dtype) |
| 43 | sd[key] = value |
| 44 | |
| 45 | return StateDict(sd=sd, device=device, size=size, dtype=dtype) |
| 46 | |
| 47 | |
| 48 | class SafetensorsModelStateDictLoader(StateDictLoader): |
| 49 | """ |
| 50 | Loads weights and configuration metadata from safetensors model files. |
| 51 | Unlike SafetensorsStateDictLoader, this loader can read model configuration |
| 52 | from the safetensors file metadata via the metadata() method. |
| 53 | """ |
| 54 | |
| 55 | def __init__(self, weight_loader: SafetensorsStateDictLoader | None = None): |
| 56 | self.weight_loader = weight_loader if weight_loader is not None else SafetensorsStateDictLoader() |
| 57 | |
| 58 | def metadata(self, path: str) -> dict: |
| 59 | with safetensors.safe_open(path, framework="pt") as f: |
| 60 | meta = f.metadata() |
| 61 | if meta is None or "config" not in meta: |
| 62 | return {} |
| 63 | return json.loads(meta["config"]) |
| 64 | |
| 65 | def load(self, path: str | list[str], sd_ops: SDOps | None = None, device: torch.device | None = None) -> StateDict: |
| 66 | return self.weight_loader.load(path, sd_ops, device) |
| 67 |