From e43ecbeb64946ac6eeed0b858509a97b27c2472d Mon Sep 17 00:00:00 2001 From: gchapman Date: Tue, 18 Jun 2024 23:24:30 +0100 Subject: [PATCH] Use correct key to get the number of layers, I accidentally used a layer with a 'final' layer giving me an out by one error --- PixArt/diffusers_convert.py | 2 +- PixArt/loader.py | 2 ++ 2 files changed, 3 insertions(+), 1 deletion(-) diff --git a/PixArt/diffusers_convert.py b/PixArt/diffusers_convert.py index 19852a4..1630106 100644 --- a/PixArt/diffusers_convert.py +++ b/PixArt/diffusers_convert.py @@ -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('cross_attn.proj.weight') for key in state_dict.keys()) def get_conversion_map(state_dict): conversion_map = [ # main SD conversion map (PixArt reference, HF Diffusers) diff --git a/PixArt/loader.py b/PixArt/loader.py index 3e26711..d1581d9 100644 --- a/PixArt/loader.py +++ b/PixArt/loader.py @@ -76,6 +76,8 @@ def load_pixart(model_path, model_conf): device=model_management.get_torch_device() ) + model_conf.unet_config['depth'] = sum(key.endswith('cross_attn.proj.weight') for key in state_dict.keys()) + if model_conf.model_target == "PixArtMS": from .models.PixArtMS import PixArtMS model.diffusion_model = PixArtMS(**model_conf.unet_config)