diff --git a/CHANGELOG.md b/CHANGELOG.md index 7379c1b3a..a790d8a92 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -1,6 +1,6 @@ # Change Log for SD.Next -## Update for 2026-04-05 +## Update for 2026-04-06 - **Models** - [AiArtLab SDXS-1B](https://huggingface.co/AiArtLab/sdxs-1b) Simple Diffusion XS *(training still in progress)* @@ -11,6 +11,7 @@ see *main interface -> scripts -> rocm advanced config* - **Internal** - additional typing and typechecks, thanks @awsr + - wrap hf download methods - **Fixes** - Prohibit `python==3.14` unless `--experimental` - UI CSS fixes, thanks @awsr diff --git a/modules/sd_hijack_hfhub.py b/modules/sd_hijack_hfhub.py new file mode 100644 index 000000000..2c06c7c58 --- /dev/null +++ b/modules/sd_hijack_hfhub.py @@ -0,0 +1,49 @@ +import os +from modules.logger import log + + +debug = log.trace if os.environ.get('SD_DOWNLOAD_DEBUG', None) is not None else lambda *args, **kwargs: None +orig_http_get = None +orig_xet_get = None + + +def http_get_hijack(*args, **kwargs): + from modules.shared import state + if len(args) > 0 and isinstance(args[0], str) and args[0].endswith(".json"): + return orig_http_get(*args, **kwargs) + jobid = state.begin('Download') + fn = kwargs.get("displayed_filename", None) + size = kwargs.get("expected_size", None) + if fn: + log.debug(f'Download start: type=http fn="{fn}" size={size}') + debug(f'Download start: type=http args={args} kwargs={kwargs}') + res = orig_http_get(*args, **kwargs) + debug(f'Download end: type=http res={res}') + state.end(jobid) + return res + + +def xet_get_hijack(*args, **kwargs): + from modules.shared import state + if len(args) > 0 and isinstance(args[0], str) and args[0].endswith(".json"): + return orig_xet_get(*args, **kwargs) + jobid = state.begin('Download') + fn = kwargs.get("displayed_filename", None) + size = kwargs.get("expected_size", None) + if fn: + log.debug(f'Download start: type=xet fn="{fn}" size={size}') + debug(f'Download start: type=xet args={args} kwargs={kwargs}') + res = orig_xet_get(*args, **kwargs) + debug(f'Download end: type=xet res={res}') + state.end(jobid) + return res + + +def init_hijack(): + from huggingface_hub import file_download + global orig_http_get, orig_xet_get # pylint: disable=global-statement + if orig_http_get is None or orig_xet_get is None: + orig_http_get = file_download.http_get + orig_xet_get = file_download.xet_get + file_download.http_get = http_get_hijack + file_download.xet_get = xet_get_hijack diff --git a/modules/sd_models.py b/modules/sd_models.py index d33396b9b..b3ac8f8d8 100644 --- a/modules/sd_models.py +++ b/modules/sd_models.py @@ -11,7 +11,7 @@ import diffusers.loaders.single_file_utils import torch import huggingface_hub as hf from modules.logger import log -from modules import timer, paths, shared, shared_items, modelloader, devices, script_callbacks, sd_vae, sd_unet, errors, sd_models_compile, sd_detect, model_quant, sd_hijack_te, sd_hijack_accelerate, sd_hijack_safetensors, attention +from modules import timer, paths, shared, shared_items, modelloader, devices, script_callbacks, sd_vae, sd_unet, errors, sd_models_compile, sd_detect, model_quant, sd_hijack_te, sd_hijack_accelerate, sd_hijack_safetensors, sd_hijack_hfhub, attention from modules.memstats import memory_stats from modules.shared_helpers import walk_files from modules.modeldata import model_data @@ -71,6 +71,7 @@ def set_huggingface_options(quiet=False): sd_hijack_safetensors.hijack_safetensors(shared.opts.runai_streamer_diffusers, shared.opts.runai_streamer_transformers) else: sd_hijack_safetensors.restore_safetensors() + sd_hijack_hfhub.init_hijack() def set_caption_load_options(): @@ -85,6 +86,7 @@ def set_caption_load_options(): if shared.opts.caption_to_gpu: log.debug(f'Caption loader: to_gpu={shared.opts.caption_to_gpu}') sd_hijack_safetensors.restore_safetensors() + sd_hijack_hfhub.init_hijack() def set_vae_options(sd_model, vae=None, op:str='model', quiet:bool=False):