Use different keys for detection between original model and converted model

This commit is contained in:
gchapman
2024-06-19 09:45:33 +01:00
parent 91340dedc7
commit 712f57c915
2 changed files with 2 additions and 2 deletions
+1 -1
View File
@@ -17,7 +17,7 @@ conversion_map_ms = [ # for multi_scale_train (MS)
]
def get_depth(state_dict):
return sum(key.endswith('.scale_shift_table') for key in state_dict.keys())
return sum(key.endswith('.attn1.to_k.bias') for key in state_dict.keys())
def get_conversion_map(state_dict):
conversion_map = [ # main SD conversion map (PixArt reference, HF Diffusers)
+1 -1
View File
@@ -76,7 +76,7 @@ def load_pixart(model_path, model_conf):
device=model_management.get_torch_device()
)
model_conf.unet_config['depth'] = sum(key.endswith('.scale_shift_table') for key in state_dict.keys())
model_conf.unet_config['depth'] = sum(key.endswith('mlp.fc1.weight') for key in state_dict.keys())
if model_conf.model_target == "PixArtMS":
from .models.PixArtMS import PixArtMS