返回 JoyAI-Echo
sft_loader.py
根目录 / ltx-core / src / ltx_core / loader / sft_loader.py
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
67 lines PYTHON