From e9d74cea7294b7874700d8cee9a448c9b2e2e0c8 Mon Sep 17 00:00:00 2001 From: awsr <43862868+awsr@users.noreply.github.com> Date: Thu, 23 Apr 2026 16:49:50 -0700 Subject: [PATCH] Fix typing --- modules/sdnq/file_loader.py | 8 ++++---- 1 file changed, 4 insertions(+), 4 deletions(-) diff --git a/modules/sdnq/file_loader.py b/modules/sdnq/file_loader.py index 4390745b7..dc9317130 100644 --- a/modules/sdnq/file_loader.py +++ b/modules/sdnq/file_loader.py @@ -3,7 +3,7 @@ import concurrent.futures import torch -def map_keys(key: str, key_mapping: dict) -> str: +def map_keys(key: str, key_mapping: dict | None) -> str: new_key = key if key_mapping: for pattern, replacement in key_mapping.items(): @@ -13,7 +13,7 @@ def map_keys(key: str, key_mapping: dict) -> str: return new_key -def load_safetensors(files: list[str], state_dict: dict | None = None, key_mapping: dict | None = None, device: torch.device = "cpu") -> dict: +def load_safetensors(files: list[str], state_dict: dict | None = None, key_mapping: dict | None = None, device: torch.device = "cpu") -> None: from safetensors.torch import safe_open if state_dict is None: state_dict = {} @@ -23,7 +23,7 @@ def load_safetensors(files: list[str], state_dict: dict | None = None, key_mappi state_dict[map_keys(key, key_mapping)] = f.get_tensor(key) -def load_threaded(files: list[str], state_dict: dict | None = None, key_mapping: dict | None = None, device: torch.device = "cpu") -> dict: +def load_threaded(files: list[str], state_dict: dict | None = None, key_mapping: dict | None = None, device: torch.device = "cpu") -> None: future_items = {} if state_dict is None: state_dict = {} @@ -34,7 +34,7 @@ def load_threaded(files: list[str], state_dict: dict | None = None, key_mapping: future.result() -def load_streamer(files: list[str], state_dict: dict | None = None, key_mapping: dict | None = None, device: torch.device = "cpu") -> dict: +def load_streamer(files: list[str], state_dict: dict | None = None, key_mapping: dict | None = None, device: torch.device = "cpu") -> None: # requires pip install runai_model_streamer from runai_model_streamer import SafetensorsStreamer if state_dict is None: