fix: mutable default arguments + wrong cfg var gating eva_clip text checkpoint download

Co-Authored-By: Claude <noreply@anthropic.com>
This commit is contained in:
QualiaRain
2026-06-12 05:08:03 -04:00
parent df568f9d4b
commit 6325f08be7
4 changed files with 21 additions and 11 deletions
+3 -1
View File
@@ -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']):
+1 -1
View File
@@ -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()
+14 -8
View File
@@ -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,
+3 -1
View File
@@ -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: