PixArt: fix OneTrainer LoRA
This commit is contained in:
@@ -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)})")
|
||||||
|
|||||||
Reference in New Issue
Block a user