From 6325f08be7c4a712b0c96d0c5e8b453ca443440d Mon Sep 17 00:00:00 2001 From: QualiaRain <44004657+QualiaRain@users.noreply.github.com> Date: Fri, 12 Jun 2026 05:08:03 -0400 Subject: [PATCH] fix: mutable default arguments + wrong cfg var gating eva_clip text checkpoint download Co-Authored-By: Claude --- pipelines/generic_vae.py | 4 +++- pipelines/model_google.py | 2 +- scripts/pulid/eva_clip/factory.py | 22 ++++++++++++++-------- scripts/pulid/eva_clip/utils.py | 4 +++- 4 files changed, 21 insertions(+), 11 deletions(-) diff --git a/pipelines/generic_vae.py b/pipelines/generic_vae.py index 27f48fac4..4b26ef1df 100644 --- a/pipelines/generic_vae.py +++ b/pipelines/generic_vae.py @@ -6,7 +6,9 @@ from modules.logger import log debug = os.environ.get('SD_LOAD_DEBUG', None) is not None -def load_vae_override(pipe, load_config=None, override_cls=None, override_args={}): +def load_vae_override(pipe, load_config=None, override_cls=None, override_args=None): + if override_args is None: + override_args = {} if shared.state.interrupted: return if (shared.opts.sd_vae in [None, 'None', 'Default', 'Automatic']): diff --git a/pipelines/model_google.py b/pipelines/model_google.py index 31f3fc4d5..624c77339 100644 --- a/pipelines/model_google.py +++ b/pipelines/model_google.py @@ -110,7 +110,7 @@ class GoogleNanoBananaPipeline(): log.debug(f'Cloud: model="{self.model}" args={args_log}') return args - def __call__(self, prompt: list[str], width: int, height: int, images: list[Image.Image] = []): + def __call__(self, prompt: list[str], width: int, height: int, images: list[Image.Image] = None): from google import genai # pylint: disable=no-name-in-module if self.client is None: args = self.get_args() diff --git a/scripts/pulid/eva_clip/factory.py b/scripts/pulid/eva_clip/factory.py index 3a466d605..134593a6d 100644 --- a/scripts/pulid/eva_clip/factory.py +++ b/scripts/pulid/eva_clip/factory.py @@ -77,7 +77,7 @@ def get_tokenizer(model_name): # loading openai CLIP weights when is_openai=True for training -def load_state_dict(checkpoint_path: str, map_location: str='cpu', model_key: str='model|module|state_dict', is_openai: bool=False, skip_list: list=[]): +def load_state_dict(checkpoint_path: str, map_location: str='cpu', model_key: str='model|module|state_dict', is_openai: bool=False, skip_list: list=None): if is_openai: model = torch.jit.load(checkpoint_path, map_location="cpu").eval() state_dict = model.state_dict() @@ -94,6 +94,8 @@ def load_state_dict(checkpoint_path: str, map_location: str='cpu', model_key: st if next(iter(state_dict.items()))[0].startswith('module'): state_dict = {k[7:]: v for k, v in state_dict.items()} + if skip_list is None: + skip_list = [] for k in skip_list: if k in list(state_dict.keys()): logging.info(f"Removing key {k} from pretrained checkpoint") @@ -128,7 +130,7 @@ def load_checkpoint(model, checkpoint_path, model_key="model|module|state_dict", logging.info(f"incompatible_keys.missing_keys: {incompatible_keys.missing_keys}") return incompatible_keys -def load_clip_visual_state_dict(checkpoint_path: str, map_location: str='cpu', is_openai: bool=False, skip_list:list=[]): +def load_clip_visual_state_dict(checkpoint_path: str, map_location: str='cpu', is_openai: bool=False, skip_list:list=None): state_dict = load_state_dict(checkpoint_path, map_location=map_location, is_openai=is_openai, skip_list=skip_list) for k in list(state_dict.keys()): @@ -141,7 +143,7 @@ def load_clip_visual_state_dict(checkpoint_path: str, map_location: str='cpu', i del state_dict[k] return state_dict -def load_clip_text_state_dict(checkpoint_path: str, map_location: str='cpu', is_openai: bool=False, skip_list:list=[]): +def load_clip_text_state_dict(checkpoint_path: str, map_location: str='cpu', is_openai: bool=False, skip_list:list=None): state_dict = load_state_dict(checkpoint_path, map_location=map_location, is_openai=is_openai, skip_list=skip_list) for k in list(state_dict.keys()): @@ -168,7 +170,9 @@ def load_pretrained_checkpoint( visual_model=None, text_model=None, model_key="model|module|state_dict", - skip_list=[]): + skip_list=None): + if skip_list is None: + skip_list = [] visual_tag = get_pretrained_tag(visual_model) text_tag = get_pretrained_tag(text_model) @@ -223,8 +227,10 @@ def create_model( pretrained_visual_model: str | None = None, pretrained_text_model: str | None = None, cache_dir: Optional[str] = None, - skip_list: list = [], + skip_list: list = None, ): + if skip_list is None: + skip_list = [] model_name = model_name.replace('/', '-') # for callers using old naming with / in ViT names if isinstance(device, str): device = torch.device(device) @@ -314,7 +320,7 @@ def create_model( if pretrained_text: pretrained_text_model = pretrained_text_model.replace('/', '-') # for callers using old naming with / in ViT names pretrained_text_cfg = get_pretrained_cfg(pretrained_text_model, pretrained_text) - if pretrained_image_cfg: + if pretrained_text_cfg: text_checkpoint_path = download_pretrained(pretrained_text_cfg, cache_dir=cache_dir) elif os.path.exists(pretrained_text): text_checkpoint_path = pretrained_text @@ -372,7 +378,7 @@ def create_model_and_transforms( image_mean: Optional[Tuple[float, ...]] = None, image_std: Optional[Tuple[float, ...]] = None, cache_dir: Optional[str] = None, - skip_list: list = [], + skip_list: list = None, ): model = create_model( model_name, @@ -427,7 +433,7 @@ def create_transforms( image_mean: Optional[Tuple[float, ...]] = None, image_std: Optional[Tuple[float, ...]] = None, cache_dir: Optional[str] = None, - skip_list: list = [], + skip_list: list = None, ): model = create_model( model_name, diff --git a/scripts/pulid/eva_clip/utils.py b/scripts/pulid/eva_clip/utils.py index 398982438..27c7c9ab0 100644 --- a/scripts/pulid/eva_clip/utils.py +++ b/scripts/pulid/eva_clip/utils.py @@ -234,7 +234,7 @@ def resize_rel_pos_embed(state_dict, model, interpolation: str = 'bicubic', seq_ patch_embed_proj.float(), size=patch_size, mode='bicubic', align_corners=False) -def freeze_batch_norm_2d(module, module_match={}, name=''): +def freeze_batch_norm_2d(module, module_match=None, name=''): """ Converts all `BatchNorm2d` and `SyncBatchNorm` layers of provided module into `FrozenBatchNorm2d`. If `module` is itself an instance of either `BatchNorm2d` or `SyncBatchNorm`, it is converted into `FrozenBatchNorm2d` and @@ -250,6 +250,8 @@ def freeze_batch_norm_2d(module, module_match={}, name=''): Inspired by https://github.com/pytorch/pytorch/blob/a5895f85be0f10212791145bfedc0261d364f103/torch/nn/modules/batchnorm.py#L762 """ + if module_match is None: + module_match = {} res = module is_match = True if module_match: