mirror of
https://github.com/vladmandic/automatic
synced 2026-09-19 09:14:35 +02:00
cleanup
This commit is contained in:
@@ -13,16 +13,20 @@ def map_keys(key: str, key_mapping: dict) -> str:
|
||||
return new_key
|
||||
|
||||
|
||||
def load_safetensors(files: list[str], state_dict: dict, key_mapping: dict = None, device: torch.device = "cpu") -> dict:
|
||||
def load_safetensors(files: list[str], state_dict: dict = None, key_mapping: dict = None, device: torch.device = "cpu") -> dict:
|
||||
from safetensors.torch import safe_open
|
||||
if state_dict is None:
|
||||
state_dict = {}
|
||||
for fn in files:
|
||||
with safe_open(fn, framework="pt", device=str(device)) as f:
|
||||
for key in f.keys():
|
||||
state_dict[map_keys(key, key_mapping)] = f.get_tensor(key)
|
||||
|
||||
|
||||
def load_threaded(files: list[str], key_mapping: dict = None, device: torch.device = "cpu", state_dict: dict = {}) -> dict:
|
||||
def load_threaded(files: list[str], state_dict: dict = None, key_mapping: dict = None, device: torch.device = "cpu") -> dict:
|
||||
future_items = {}
|
||||
if state_dict is None:
|
||||
state_dict = {}
|
||||
with concurrent.futures.ThreadPoolExecutor(max_workers=4) as executor:
|
||||
for fn in files:
|
||||
future_items[executor.submit(load_safetensors, [fn], key_mapping=key_mapping, device=device, state_dict=state_dict)] = fn
|
||||
@@ -30,25 +34,29 @@ def load_threaded(files: list[str], key_mapping: dict = None, device: torch.devi
|
||||
future.result()
|
||||
|
||||
|
||||
def load_streamer(files: list[str], state_dict: dict, key_mapping: dict = None, device: torch.device = "cpu") -> dict:
|
||||
def load_streamer(files: list[str], state_dict: dict = None, key_mapping: dict = None, device: torch.device = "cpu") -> dict:
|
||||
# requires pip install runai_model_streamer
|
||||
from runai_model_streamer import SafetensorsStreamer
|
||||
if state_dict is None:
|
||||
state_dict = {}
|
||||
with SafetensorsStreamer() as streamer:
|
||||
streamer.stream_files(files)
|
||||
for key, tensor in streamer.get_tensors():
|
||||
state_dict[map_keys(key, key_mapping)] = tensor.to(device)
|
||||
|
||||
|
||||
def load_files(files: list[str], method: str = None, key_mapping: dict = None, device: torch.device = "cpu") -> dict:
|
||||
def load_files(files: list[str], state_dict: dict = None, key_mapping: dict = None, device: torch.device = "cpu", method: str = None) -> dict:
|
||||
# note: files is list-of-files within a module for chunked loading, not accross model
|
||||
method = method or 'safetensors'
|
||||
state_dict = {}
|
||||
if method is None:
|
||||
method = 'safetensors'
|
||||
if state_dict is None:
|
||||
state_dict = {}
|
||||
if method == 'safetensors':
|
||||
load_safetensors(files, key_mapping=key_mapping, device=device, state_dict=state_dict)
|
||||
load_safetensors(files, state_dict=state_dict, key_mapping=key_mapping, device=device)
|
||||
elif method == 'threaded':
|
||||
load_threaded(files, key_mapping=key_mapping, device=device, state_dict=state_dict)
|
||||
load_threaded(files, state_dict=state_dict, key_mapping=key_mapping, device=device)
|
||||
elif method == 'streamer':
|
||||
load_streamer(files, key_mapping=key_mapping, device=device, state_dict=state_dict)
|
||||
load_streamer(files, state_dict=state_dict, key_mapping=key_mapping, device=device)
|
||||
else:
|
||||
raise ValueError(f"Unsupported loading method: {method}")
|
||||
return state_dict
|
||||
|
||||
Reference in New Issue
Block a user