Use different keys for detection between original model and converted model
This commit is contained in:
@@ -17,7 +17,7 @@ conversion_map_ms = [ # for multi_scale_train (MS)
|
|||||||
]
|
]
|
||||||
|
|
||||||
def get_depth(state_dict):
|
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):
|
def get_conversion_map(state_dict):
|
||||||
conversion_map = [ # main SD conversion map (PixArt reference, HF Diffusers)
|
conversion_map = [ # main SD conversion map (PixArt reference, HF Diffusers)
|
||||||
|
|||||||
+1
-1
@@ -76,7 +76,7 @@ def load_pixart(model_path, model_conf):
|
|||||||
device=model_management.get_torch_device()
|
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":
|
if model_conf.model_target == "PixArtMS":
|
||||||
from .models.PixArtMS import PixArtMS
|
from .models.PixArtMS import PixArtMS
|
||||||
|
|||||||
Reference in New Issue
Block a user