[diffusion] feat: add support for LoRA layers in transformer_2 within LoRAPipeline (#14606)
This commit is contained in:
@@ -46,6 +46,7 @@ class LoRAPipeline(ComposedPipelineBase):
|
|||||||
# [dit_layer_name] = wrapped_lora_layer
|
# [dit_layer_name] = wrapped_lora_layer
|
||||||
lora_layers: dict[str, BaseLayerWithLoRA] = {}
|
lora_layers: dict[str, BaseLayerWithLoRA] = {}
|
||||||
lora_layers_critic: dict[str, BaseLayerWithLoRA] = {}
|
lora_layers_critic: dict[str, BaseLayerWithLoRA] = {}
|
||||||
|
lora_layers_transformer_2: dict[str, BaseLayerWithLoRA] = {}
|
||||||
server_args: ServerArgs
|
server_args: ServerArgs
|
||||||
exclude_lora_layers: list[str] = []
|
exclude_lora_layers: list[str] = []
|
||||||
device: torch.device = get_local_torch_device()
|
device: torch.device = get_local_torch_device()
|
||||||
@@ -79,18 +80,31 @@ class LoRAPipeline(ComposedPipelineBase):
|
|||||||
target_name in module_name for target_name in self.lora_target_modules
|
target_name in module_name for target_name in self.lora_target_modules
|
||||||
)
|
)
|
||||||
|
|
||||||
def convert_to_lora_layers(self) -> None:
|
def convert_module_lora_layers(
|
||||||
|
self,
|
||||||
|
module: torch.nn.Module,
|
||||||
|
module_name: str,
|
||||||
|
target_lora_layers: dict[str, BaseLayerWithLoRA],
|
||||||
|
check_exclude: bool = True,
|
||||||
|
) -> int:
|
||||||
"""
|
"""
|
||||||
Unified method to convert the transformer to a LoRA transformer.
|
Convert layers in a module to LoRA layers.
|
||||||
|
|
||||||
|
Args:
|
||||||
|
module: The module to convert.
|
||||||
|
module_name: The name of the module (for replace_submodule).
|
||||||
|
target_lora_layers: The dictionary to store the converted LoRA layers.
|
||||||
|
check_exclude: Whether to check the exclude_lora_layers list.
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
The number of layers converted.
|
||||||
"""
|
"""
|
||||||
if self.lora_initialized:
|
|
||||||
return
|
|
||||||
self.lora_initialized = True
|
|
||||||
converted_count = 0
|
converted_count = 0
|
||||||
for name, layer in self.modules["transformer"].named_modules():
|
for name, layer in module.named_modules():
|
||||||
if not self.is_target_layer(name):
|
if not self.is_target_layer(name):
|
||||||
continue
|
continue
|
||||||
|
|
||||||
|
if check_exclude:
|
||||||
excluded = any(
|
excluded = any(
|
||||||
exclude_layer in name for exclude_layer in self.exclude_lora_layers
|
exclude_layer in name for exclude_layer in self.exclude_lora_layers
|
||||||
)
|
)
|
||||||
@@ -103,31 +117,99 @@ class LoRAPipeline(ComposedPipelineBase):
|
|||||||
lora_alpha=self.lora_alpha,
|
lora_alpha=self.lora_alpha,
|
||||||
)
|
)
|
||||||
if lora_layer is not None:
|
if lora_layer is not None:
|
||||||
self.lora_layers[name] = lora_layer
|
target_lora_layers[name] = lora_layer
|
||||||
replace_submodule(self.modules["transformer"], name, lora_layer)
|
replace_submodule(self.modules[module_name], name, lora_layer)
|
||||||
converted_count += 1
|
converted_count += 1
|
||||||
|
return converted_count
|
||||||
|
|
||||||
|
def convert_to_lora_layers(self) -> None:
|
||||||
|
"""
|
||||||
|
Unified method to convert the transformer to a LoRA transformer.
|
||||||
|
"""
|
||||||
|
if self.lora_initialized:
|
||||||
|
return
|
||||||
|
self.lora_initialized = True
|
||||||
|
|
||||||
|
# Convert transformer
|
||||||
|
converted_count = self.convert_module_lora_layers(
|
||||||
|
self.modules["transformer"],
|
||||||
|
"transformer",
|
||||||
|
self.lora_layers,
|
||||||
|
check_exclude=True,
|
||||||
|
)
|
||||||
logger.info("Converted %d layers to LoRA layers", converted_count)
|
logger.info("Converted %d layers to LoRA layers", converted_count)
|
||||||
|
|
||||||
|
# Convert transformer_2 if exists (e.g., Wan2.2 A14B dual-transformer)
|
||||||
|
if (
|
||||||
|
"transformer_2" in self.modules
|
||||||
|
and self.modules["transformer_2"] is not None
|
||||||
|
):
|
||||||
|
converted_count_2 = self.convert_module_lora_layers(
|
||||||
|
self.modules["transformer_2"],
|
||||||
|
"transformer_2",
|
||||||
|
self.lora_layers_transformer_2,
|
||||||
|
check_exclude=True,
|
||||||
|
)
|
||||||
|
logger.info(
|
||||||
|
"Converted %d layers to LoRA layers in transformer_2", converted_count_2
|
||||||
|
)
|
||||||
|
|
||||||
|
# Convert fake_score_transformer if exists
|
||||||
if "fake_score_transformer" in self.modules:
|
if "fake_score_transformer" in self.modules:
|
||||||
for name, layer in self.modules["fake_score_transformer"].named_modules():
|
converted_count_critic = self.convert_module_lora_layers(
|
||||||
if not self.is_target_layer(name):
|
self.modules["fake_score_transformer"],
|
||||||
continue
|
"fake_score_transformer",
|
||||||
layer = wrap_with_lora_layer(
|
self.lora_layers_critic,
|
||||||
layer,
|
check_exclude=False,
|
||||||
lora_rank=self.lora_rank,
|
|
||||||
lora_alpha=self.lora_alpha,
|
|
||||||
)
|
)
|
||||||
if layer is not None:
|
|
||||||
self.lora_layers_critic[name] = layer
|
|
||||||
replace_submodule(
|
|
||||||
self.modules["fake_score_transformer"], name, layer
|
|
||||||
)
|
|
||||||
converted_count += 1
|
|
||||||
logger.info(
|
logger.info(
|
||||||
"Converted %d layers to LoRA layers in the critic model",
|
"Converted %d layers to LoRA layers in the critic model",
|
||||||
converted_count,
|
converted_count_critic,
|
||||||
)
|
)
|
||||||
|
|
||||||
|
def _apply_lora_to_layers(
|
||||||
|
self,
|
||||||
|
lora_layers: dict[str, BaseLayerWithLoRA],
|
||||||
|
lora_nickname: str,
|
||||||
|
lora_path: str | None,
|
||||||
|
rank: int,
|
||||||
|
) -> int:
|
||||||
|
"""
|
||||||
|
Apply LoRA weights to the given lora_layers.
|
||||||
|
|
||||||
|
Args:
|
||||||
|
lora_layers: The dictionary of LoRA layers to apply weights to.
|
||||||
|
lora_nickname: The nickname of the LoRA adapter.
|
||||||
|
lora_path: The path to the LoRA adapter.
|
||||||
|
rank: The distributed rank (for logging).
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
The number of layers that had LoRA weights applied.
|
||||||
|
"""
|
||||||
|
adapted_count = 0
|
||||||
|
for name, layer in lora_layers.items():
|
||||||
|
lora_A_name = name + ".lora_A"
|
||||||
|
lora_B_name = name + ".lora_B"
|
||||||
|
if (
|
||||||
|
lora_A_name in self.lora_adapters[lora_nickname]
|
||||||
|
and lora_B_name in self.lora_adapters[lora_nickname]
|
||||||
|
):
|
||||||
|
layer.set_lora_weights(
|
||||||
|
self.lora_adapters[lora_nickname][lora_A_name],
|
||||||
|
self.lora_adapters[lora_nickname][lora_B_name],
|
||||||
|
lora_path=lora_path,
|
||||||
|
)
|
||||||
|
adapted_count += 1
|
||||||
|
else:
|
||||||
|
if rank == 0:
|
||||||
|
logger.warning(
|
||||||
|
"LoRA adapter %s does not contain the weights for layer '%s'. LoRA will not be applied to it.",
|
||||||
|
lora_path,
|
||||||
|
name,
|
||||||
|
)
|
||||||
|
layer.disable_lora = True
|
||||||
|
return adapted_count
|
||||||
|
|
||||||
def is_lora_effective(self):
|
def is_lora_effective(self):
|
||||||
return self.is_lora_merged
|
return self.is_lora_merged
|
||||||
|
|
||||||
@@ -234,28 +316,18 @@ class LoRAPipeline(ComposedPipelineBase):
|
|||||||
self.cur_adapter_name = lora_nickname
|
self.cur_adapter_name = lora_nickname
|
||||||
|
|
||||||
# Merge the new adapter
|
# Merge the new adapter
|
||||||
adapted_count = 0
|
adapted_count = self._apply_lora_to_layers(
|
||||||
for name, layer in self.lora_layers.items():
|
self.lora_layers, lora_nickname, lora_path, rank
|
||||||
lora_A_name = name + ".lora_A"
|
|
||||||
lora_B_name = name + ".lora_B"
|
|
||||||
if (
|
|
||||||
lora_A_name in self.lora_adapters[lora_nickname]
|
|
||||||
and lora_B_name in self.lora_adapters[lora_nickname]
|
|
||||||
):
|
|
||||||
layer.set_lora_weights(
|
|
||||||
self.lora_adapters[lora_nickname][lora_A_name],
|
|
||||||
self.lora_adapters[lora_nickname][lora_B_name],
|
|
||||||
lora_path=lora_path,
|
|
||||||
)
|
)
|
||||||
adapted_count += 1
|
# Apply LoRA to transformer_2 if exists
|
||||||
else:
|
adapted_count += self._apply_lora_to_layers(
|
||||||
if rank == 0:
|
self.lora_layers_transformer_2, lora_nickname, lora_path, rank
|
||||||
logger.warning(
|
|
||||||
"LoRA adapter %s does not contain the weights for layer '%s'. LoRA will not be applied to it.",
|
|
||||||
lora_path,
|
|
||||||
name,
|
|
||||||
)
|
)
|
||||||
layer.disable_lora = True
|
# Apply LoRA to fake_score_transformer (critic) if exists
|
||||||
|
adapted_count += self._apply_lora_to_layers(
|
||||||
|
self.lora_layers_critic, lora_nickname, lora_path, rank
|
||||||
|
)
|
||||||
|
|
||||||
self.is_lora_merged = True
|
self.is_lora_merged = True
|
||||||
logger.info(
|
logger.info(
|
||||||
"Rank %d: LoRA adapter %s applied to %d layers",
|
"Rank %d: LoRA adapter %s applied to %d layers",
|
||||||
@@ -271,6 +343,10 @@ class LoRAPipeline(ComposedPipelineBase):
|
|||||||
|
|
||||||
for name, layer in self.lora_layers.items():
|
for name, layer in self.lora_layers.items():
|
||||||
layer.merge_lora_weights()
|
layer.merge_lora_weights()
|
||||||
|
for name, layer in self.lora_layers_transformer_2.items():
|
||||||
|
layer.merge_lora_weights()
|
||||||
|
for name, layer in self.lora_layers_critic.items():
|
||||||
|
layer.merge_lora_weights()
|
||||||
logger.info("LoRA weights merged")
|
logger.info("LoRA weights merged")
|
||||||
self.is_lora_merged = True
|
self.is_lora_merged = True
|
||||||
|
|
||||||
@@ -281,4 +357,8 @@ class LoRAPipeline(ComposedPipelineBase):
|
|||||||
|
|
||||||
for name, layer in self.lora_layers.items():
|
for name, layer in self.lora_layers.items():
|
||||||
layer.unmerge_lora_weights()
|
layer.unmerge_lora_weights()
|
||||||
|
for name, layer in self.lora_layers_transformer_2.items():
|
||||||
|
layer.unmerge_lora_weights()
|
||||||
|
for name, layer in self.lora_layers_critic.items():
|
||||||
|
layer.unmerge_lora_weights()
|
||||||
self.is_lora_merged = False
|
self.is_lora_merged = False
|
||||||
|
|||||||
Reference in New Issue
Block a user