PixArt: fix OneTrainer LoRA

This commit is contained in:
City
2024-07-27 17:52:02 +02:00
parent d89de6e98e
commit 6fb8649168
+20 -7
View File
@@ -20,7 +20,14 @@ def get_depth(state_dict):
return sum(key.endswith('.attn1.to_k.bias') for key in state_dict.keys()) return sum(key.endswith('.attn1.to_k.bias') for key in state_dict.keys())
def get_lora_depth(state_dict): def get_lora_depth(state_dict):
return sum(key.endswith('.attn1.to_k.lora_A.weight') for key in state_dict.keys()) cnt = max([
sum(key.endswith('.attn1.to_k.lora_A.weight') for key in state_dict.keys()),
sum(key.endswith('_attn1_to_k.lora_A.weight') for key in state_dict.keys()),
sum(key.endswith('.attn1.to_k.lora_up.weight') for key in state_dict.keys()),
sum(key.endswith('_attn1_to_k.lora_up.weight') for key in state_dict.keys()),
])
assert cnt > 0, "Unable to detect model depth!"
return cnt
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)
@@ -191,13 +198,19 @@ def convert_lora_state_dict(state_dict, peft=True):
new_state_dict[fk(f"blocks.{depth}.cross_attn.proj.weight")] = state_dict[key('out.0')] new_state_dict[fk(f"blocks.{depth}.cross_attn.proj.weight")] = state_dict[key('out.0')]
matched += [key('out.0')] matched += [key('out.0')]
key = fp(f"transformer_blocks.{depth}.ff.net.0.proj.weight") try:
new_state_dict[fk(f"blocks.{depth}.mlp.fc1.weight")] = state_dict[key] key = fp(f"transformer_blocks.{depth}.ff.net.0.proj.weight")
matched += [key] new_state_dict[fk(f"blocks.{depth}.mlp.fc1.weight")] = state_dict[key]
matched += [key]
except KeyError:
pass
key = fp(f"transformer_blocks.{depth}.ff.net.2.weight") try:
new_state_dict[fk(f"blocks.{depth}.mlp.fc2.weight")] = state_dict[key] key = fp(f"transformer_blocks.{depth}.ff.net.2.weight")
matched += [key] new_state_dict[fk(f"blocks.{depth}.mlp.fc2.weight")] = state_dict[key]
matched += [key]
except KeyError:
pass
if len(matched) < len(state_dict): if len(matched) < len(state_dict):
print(f"PixArt: LoRA conversion has leftover keys! ({len(matched)} vs {len(state_dict)})") print(f"PixArt: LoRA conversion has leftover keys! ({len(matched)} vs {len(state_dict)})")