[diffusion] fix: fix h3 swap peft SwiGLU lora_B halves when loading FFN Lora (#34940)
This commit is contained in:
@@ -39,6 +39,22 @@ os.environ["TOKENIZERS_PARALLELISM"] = "false"
|
|||||||
logger = init_logger(__name__)
|
logger = init_logger(__name__)
|
||||||
|
|
||||||
|
|
||||||
|
def _swap_peft_swiglu_fc1_lora_b(
|
||||||
|
source_name: str, target_name: str, weight: torch.Tensor
|
||||||
|
) -> torch.Tensor:
|
||||||
|
# Only the PEFT -> native H3 FFN rewrite: ff.net.0.proj [value; gate]
|
||||||
|
# onto mlp.fc1 [gate; value]. Native fused mlp.fc1 and other models'
|
||||||
|
# ff.net.0.proj (e.g. Flux) must not match.
|
||||||
|
if (
|
||||||
|
weight.dim() != 2
|
||||||
|
or ".ff.net.0.proj.lora_B" not in source_name
|
||||||
|
or not target_name.endswith(".mlp.fc1.lora_B")
|
||||||
|
):
|
||||||
|
return weight
|
||||||
|
value, gate = weight.chunk(2, dim=0)
|
||||||
|
return torch.cat([gate, value], dim=0)
|
||||||
|
|
||||||
|
|
||||||
class LoRAPipeline(ComposedPipelineBase):
|
class LoRAPipeline(ComposedPipelineBase):
|
||||||
"""
|
"""
|
||||||
Pipeline that supports injecting LoRA adapters into the diffusion transformer.
|
Pipeline that supports injecting LoRA adapters into the diffusion transformer.
|
||||||
@@ -775,6 +791,7 @@ class LoRAPipeline(ComposedPipelineBase):
|
|||||||
else:
|
else:
|
||||||
continue
|
continue
|
||||||
|
|
||||||
|
weight = _swap_peft_swiglu_fc1_lora_b(name, target_name, weight)
|
||||||
if target_name in self.lora_adapters[lora_nickname]:
|
if target_name in self.lora_adapters[lora_nickname]:
|
||||||
raise ValueError(
|
raise ValueError(
|
||||||
f"Dit target weight name {target_name} already exists in lora_adapters[{lora_nickname}]"
|
f"Dit target weight name {target_name} already exists in lora_adapters[{lora_nickname}]"
|
||||||
|
|||||||
Reference in New Issue
Block a user