From a830c0a7e0fc26e3812844add27c013202f552c7 Mon Sep 17 00:00:00 2001 From: Disty0 Date: Mon, 27 Oct 2025 21:32:52 +0300 Subject: [PATCH] cleanup --- modules/sdnq/file_loader.py | 26 +++++++++++++++++--------- 1 file changed, 17 insertions(+), 9 deletions(-) diff --git a/modules/sdnq/file_loader.py b/modules/sdnq/file_loader.py index c026d015b..5a5796e6b 100644 --- a/modules/sdnq/file_loader.py +++ b/modules/sdnq/file_loader.py @@ -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