[diffusion] feat: support lora strength (#15691)

This commit is contained in:
Prozac614
2025-12-24 13:31:23 +08:00
committed by GitHub
parent e254cdf326
commit eee3700d84
8 changed files with 113 additions and 30 deletions
@@ -248,6 +248,8 @@ Loads a LoRA adapter and merges its weights into the model.
**Parameters:** **Parameters:**
- `lora_nickname` (string, required): A unique identifier for this LoRA - `lora_nickname` (string, required): A unique identifier for this LoRA
- `lora_path` (string, optional): Path to the `.safetensors` file or Hugging Face repo ID. Required for the first load; optional if re-activating a cached nickname - `lora_path` (string, optional): Path to the `.safetensors` file or Hugging Face repo ID. Required for the first load; optional if re-activating a cached nickname
- `target` (string, optional): Which transformer(s) to apply the LoRA to. One of "all" (default), "transformer", "transformer_2", "critic"
- `strength` (float, optional): LoRA strength for merge, default 1.0. Values < 1.0 reduce the effect, values > 1.0 amplify the effect
**Curl Example:** **Curl Example:**
@@ -256,7 +258,8 @@ curl -X POST http://localhost:30010/v1/set_lora \
-H "Content-Type: application/json" \ -H "Content-Type: application/json" \
-d '{ -d '{
"lora_nickname": "lora_name", "lora_nickname": "lora_name",
"lora_path": "/path/to/lora.safetensors" "lora_path": "/path/to/lora.safetensors",
"strength": 0.8
}' }'
``` ```
@@ -270,11 +273,16 @@ Manually merges the currently set LoRA weights into the base model.
**Endpoint:** `POST /v1/merge_lora_weights` **Endpoint:** `POST /v1/merge_lora_weights`
**Parameters:**
- `target` (string, optional): Which transformer(s) to merge. One of "all" (default), "transformer", "transformer_2", "critic"
- `strength` (float, optional): LoRA strength for merge, default 1.0. Values < 1.0 reduce the effect, values > 1.0 amplify the effect
**Curl Example:** **Curl Example:**
```bash ```bash
curl -X POST http://localhost:30010/v1/merge_lora_weights \ curl -X POST http://localhost:30010/v1/merge_lora_weights \
-H "Content-Type: application/json" -H "Content-Type: application/json" \
-d '{"strength": 0.8}'
``` ```
@@ -303,7 +303,11 @@ class DiffGenerator:
raise RuntimeError(f"{failure_msg}: {error_msg}") raise RuntimeError(f"{failure_msg}: {error_msg}")
def set_lora( def set_lora(
self, lora_nickname: str, lora_path: str | None = None, target: str = "all" self,
lora_nickname: str,
lora_path: str | None = None,
target: str = "all",
strength: float = 1.0,
) -> None: ) -> None:
""" """
Set a LoRA adapter for the specified transformer(s). Set a LoRA adapter for the specified transformer(s).
@@ -316,13 +320,17 @@ class DiffGenerator:
- "transformer": Apply only to the primary transformer (high noise for Wan2.2) - "transformer": Apply only to the primary transformer (high noise for Wan2.2)
- "transformer_2": Apply only to transformer_2 (low noise for Wan2.2) - "transformer_2": Apply only to transformer_2 (low noise for Wan2.2)
- "critic": Apply only to the critic model - "critic": Apply only to the critic model
strength: LoRA strength for merge, default 1.0.
""" """
req = SetLoraReq( req = SetLoraReq(
lora_nickname=lora_nickname, lora_path=lora_path, target=target lora_nickname=lora_nickname,
lora_path=lora_path,
target=target,
strength=strength,
) )
self._send_lora_request( self._send_lora_request(
req, req,
f"Successfully set LoRA adapter: {lora_nickname} (target: {target})", f"Successfully set LoRA adapter: {lora_nickname} (target: {target}, strength: {strength})",
"Failed to set LoRA adapter", "Failed to set LoRA adapter",
) )
@@ -340,17 +348,18 @@ class DiffGenerator:
"Failed to unmerge LoRA weights", "Failed to unmerge LoRA weights",
) )
def merge_lora_weights(self, target: str = "all") -> None: def merge_lora_weights(self, target: str = "all", strength: float = 1.0) -> None:
""" """
Merge LoRA weights into the base model. Merge LoRA weights into the base model.
Args: Args:
target: Which transformer(s) to merge. target: Which transformer(s) to merge.
strength: LoRA strength for merge, default 1.0.
""" """
req = MergeLoraWeightsReq(target=target) req = MergeLoraWeightsReq(target=target, strength=strength)
self._send_lora_request( self._send_lora_request(
req, req,
f"Successfully merged LoRA weights (target: {target})", f"Successfully merged LoRA weights (target: {target}, strength: {strength})",
"Failed to merge LoRA weights", "Failed to merge LoRA weights",
) )
@@ -39,6 +39,7 @@ async def set_lora(
lora_nickname: str = Body(..., embed=True), lora_nickname: str = Body(..., embed=True),
lora_path: Optional[str] = Body(None, embed=True), lora_path: Optional[str] = Body(None, embed=True),
target: str = Body("all", embed=True), target: str = Body("all", embed=True),
strength: float = Body(1.0, embed=True),
): ):
""" """
Set a LoRA adapter for the specified transformer(s). Set a LoRA adapter for the specified transformer(s).
@@ -51,11 +52,18 @@ async def set_lora(
- "transformer": Apply only to the primary transformer (high noise for Wan2.2) - "transformer": Apply only to the primary transformer (high noise for Wan2.2)
- "transformer_2": Apply only to transformer_2 (low noise for Wan2.2) - "transformer_2": Apply only to transformer_2 (low noise for Wan2.2)
- "critic": Apply only to the critic model - "critic": Apply only to the critic model
strength: LoRA strength for merge, default 1.0. Values < 1.0 reduce the effect,
values > 1.0 amplify the effect.
""" """
req = SetLoraReq(lora_nickname=lora_nickname, lora_path=lora_path, target=target) req = SetLoraReq(
lora_nickname=lora_nickname,
lora_path=lora_path,
target=target,
strength=strength,
)
return await _handle_lora_request( return await _handle_lora_request(
req, req,
f"Successfully set LoRA adapter: {lora_nickname} (target: {target})", f"Successfully set LoRA adapter: {lora_nickname} (target: {target}, strength: {strength})",
"Failed to set LoRA adapter", "Failed to set LoRA adapter",
) )
@@ -63,6 +71,7 @@ async def set_lora(
@router.post("/merge_lora_weights") @router.post("/merge_lora_weights")
async def merge_lora_weights( async def merge_lora_weights(
target: str = Body("all", embed=True), target: str = Body("all", embed=True),
strength: float = Body(1.0, embed=True),
): ):
""" """
Merge LoRA weights into the base model. Merge LoRA weights into the base model.
@@ -70,11 +79,13 @@ async def merge_lora_weights(
Args: Args:
target: Which transformer(s) to merge. One of "all", "transformer", target: Which transformer(s) to merge. One of "all", "transformer",
"transformer_2", "critic". "transformer_2", "critic".
strength: LoRA strength for merge, default 1.0. Values < 1.0 reduce the effect,
values > 1.0 amplify the effect.
""" """
req = MergeLoraWeightsReq(target=target) req = MergeLoraWeightsReq(target=target, strength=strength)
return await _handle_lora_request( return await _handle_lora_request(
req, req,
f"Successfully merged LoRA weights (target: {target})", f"Successfully merged LoRA weights (target: {target}, strength: {strength})",
"Failed to merge LoRA weights", "Failed to merge LoRA weights",
) )
@@ -25,11 +25,13 @@ class SetLoraReq:
lora_nickname: str lora_nickname: str
lora_path: Optional[str] = None lora_path: Optional[str] = None
target: str = "all" # "all", "transformer", "transformer_2", "critic" target: str = "all" # "all", "transformer", "transformer_2", "critic"
strength: float = 1.0 # LoRA strength for merge, default 1.0
@dataclasses.dataclass @dataclasses.dataclass
class MergeLoraWeightsReq: class MergeLoraWeightsReq:
target: str = "all" # "all", "transformer", "transformer_2", "critic" target: str = "all" # "all", "transformer", "transformer_2", "critic"
strength: float = 1.0 # LoRA strength for merge, default 1.0
@dataclasses.dataclass @dataclasses.dataclass
@@ -56,6 +56,7 @@ class BaseLayerWithLoRA(nn.Module):
self.lora_rank = lora_rank self.lora_rank = lora_rank
self.lora_alpha = lora_alpha self.lora_alpha = lora_alpha
self.lora_path: str | None = None self.lora_path: str | None = None
self.strength: float = 1.0
self.lora_A = None self.lora_A = None
self.lora_B = None self.lora_B = None
@@ -84,6 +85,7 @@ class BaseLayerWithLoRA(nn.Module):
delta = delta * ( delta = delta * (
self.lora_alpha / self.lora_rank # type: ignore self.lora_alpha / self.lora_rank # type: ignore
) # type: ignore ) # type: ignore
delta = delta * self.strength
out, output_bias = self.base_layer(x) out, output_bias = self.base_layer(x)
return out + delta, output_bias return out + delta, output_bias
else: else:
@@ -101,17 +103,22 @@ class BaseLayerWithLoRA(nn.Module):
A: torch.Tensor, A: torch.Tensor,
B: torch.Tensor, B: torch.Tensor,
lora_path: str | None = None, lora_path: str | None = None,
strength: float = 1.0,
) -> None: ) -> None:
self.lora_A = torch.nn.Parameter( self.lora_A = torch.nn.Parameter(
A A
) # share storage with weights in the pipeline ) # share storage with weights in the pipeline
self.lora_B = torch.nn.Parameter(B) self.lora_B = torch.nn.Parameter(B)
self.disable_lora = False self.disable_lora = False
self.strength = strength
self.merge_lora_weights() self.merge_lora_weights()
self.lora_path = lora_path self.lora_path = lora_path
@torch.no_grad() @torch.no_grad()
def merge_lora_weights(self) -> None: def merge_lora_weights(self, strength: float | None = None) -> None:
if strength is not None:
self.strength = strength
if self.disable_lora: if self.disable_lora:
return return
@@ -136,9 +143,14 @@ class BaseLayerWithLoRA(nn.Module):
data = self.base_layer.weight.data.to( data = self.base_layer.weight.data.to(
get_local_torch_device() get_local_torch_device()
).full_tensor() ).full_tensor()
data += self.slice_lora_b_weights(self.lora_B).to( lora_delta = self.slice_lora_b_weights(self.lora_B).to(
data data
) @ self.slice_lora_a_weights(self.lora_A).to(data) ) @ self.slice_lora_a_weights(self.lora_A).to(data)
# Apply lora_alpha / lora_rank scaling for consistency with forward()
if self.lora_alpha is not None and self.lora_rank is not None:
if self.lora_alpha != self.lora_rank:
lora_delta = lora_delta * (self.lora_alpha / self.lora_rank)
data += self.strength * lora_delta
unsharded_base_layer.weight = nn.Parameter(data.to(current_device)) unsharded_base_layer.weight = nn.Parameter(data.to(current_device))
if isinstance(getattr(self.base_layer, "bias", None), DTensor): if isinstance(getattr(self.base_layer, "bias", None), DTensor):
unsharded_base_layer.bias = nn.Parameter( unsharded_base_layer.bias = nn.Parameter(
@@ -161,9 +173,14 @@ class BaseLayerWithLoRA(nn.Module):
else: else:
current_device = self.base_layer.weight.data.device current_device = self.base_layer.weight.data.device
data = self.base_layer.weight.data.to(get_local_torch_device()) data = self.base_layer.weight.data.to(get_local_torch_device())
data += self.slice_lora_b_weights( lora_delta = self.slice_lora_b_weights(
self.lora_B.to(data) self.lora_B.to(data)
) @ self.slice_lora_a_weights(self.lora_A.to(data)) ) @ self.slice_lora_a_weights(self.lora_A.to(data))
# Apply lora_alpha / lora_rank scaling for consistency with forward()
if self.lora_alpha is not None and self.lora_rank is not None:
if self.lora_alpha != self.lora_rank:
lora_delta = lora_delta * (self.lora_alpha / self.lora_rank)
data += self.strength * lora_delta
self.base_layer.weight.data = data.to(current_device, non_blocking=True) self.base_layer.weight.data = data.to(current_device, non_blocking=True)
self.merged = True self.merged = True
@@ -391,6 +408,7 @@ class LinearWithLoRA(BaseLayerWithLoRA):
delta = delta * ( delta = delta * (
self.lora_alpha / self.lora_rank # type: ignore self.lora_alpha / self.lora_rank # type: ignore
) # type: ignore ) # type: ignore
delta = delta * self.strength
# nn.Linear.forward() returns a single tensor, not a tuple # nn.Linear.forward() returns a single tensor, not a tuple
out = self.base_layer(x) out = self.base_layer(x)
return out + delta return out + delta
@@ -129,7 +129,11 @@ class GPUWorker:
return output_batch return output_batch
def set_lora( def set_lora(
self, lora_nickname: str, lora_path: str | None = None, target: str = "all" self,
lora_nickname: str,
lora_path: str | None = None,
target: str = "all",
strength: float = 1.0,
) -> None: ) -> None:
""" """
Set the LoRA adapter for the pipeline. Set the LoRA adapter for the pipeline.
@@ -138,19 +142,21 @@ class GPUWorker:
lora_nickname: The nickname of the adapter. lora_nickname: The nickname of the adapter.
lora_path: Path to the LoRA adapter. lora_path: Path to the LoRA adapter.
target: Which transformer(s) to apply the LoRA to. target: Which transformer(s) to apply the LoRA to.
strength: LoRA strength for merge, default 1.0.
""" """
assert self.pipeline is not None assert self.pipeline is not None
self.pipeline.set_lora(lora_nickname, lora_path, target) self.pipeline.set_lora(lora_nickname, lora_path, target, strength)
def merge_lora_weights(self, target: str = "all") -> None: def merge_lora_weights(self, target: str = "all", strength: float = 1.0) -> None:
""" """
Merge LoRA weights. Merge LoRA weights.
Args: Args:
target: Which transformer(s) to merge. target: Which transformer(s) to merge.
strength: LoRA strength for merge, default 1.0.
""" """
assert self.pipeline is not None assert self.pipeline is not None
self.pipeline.merge_lora_weights(target) self.pipeline.merge_lora_weights(target, strength)
def unmerge_lora_weights(self, target: str = "all") -> None: def unmerge_lora_weights(self, target: str = "all") -> None:
""" """
@@ -85,12 +85,12 @@ class Scheduler:
def _handle_set_lora(self, reqs: List[Any]): def _handle_set_lora(self, reqs: List[Any]):
# TODO: return set status # TODO: return set status
req = reqs[0] req = reqs[0]
self.worker.set_lora(req.lora_nickname, req.lora_path, req.target) self.worker.set_lora(req.lora_nickname, req.lora_path, req.target, req.strength)
return {"status": "ok"} return {"status": "ok"}
def _handle_merge_lora(self, reqs: List[Any]): def _handle_merge_lora(self, reqs: List[Any]):
req = reqs[0] req = reqs[0]
self.worker.merge_lora_weights(req.target) self.worker.merge_lora_weights(req.target, req.strength)
return {"status": "ok"} return {"status": "ok"}
def _handle_unmerge_lora(self, reqs: List[Any]): def _handle_unmerge_lora(self, reqs: List[Any]):
@@ -46,6 +46,7 @@ class LoRAPipeline(ComposedPipelineBase):
# Track current adapter per module: {"transformer": "high_lora", "transformer_2": "low_lora"} # Track current adapter per module: {"transformer": "high_lora", "transformer_2": "low_lora"}
cur_adapter_name: dict[str, str] cur_adapter_name: dict[str, str]
cur_adapter_path: dict[str, str] cur_adapter_path: dict[str, str]
cur_adapter_strength: dict[str, float] # Track current strength per module
# [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]
@@ -71,6 +72,7 @@ class LoRAPipeline(ComposedPipelineBase):
self.loaded_adapter_paths = {} self.loaded_adapter_paths = {}
self.cur_adapter_name = {} self.cur_adapter_name = {}
self.cur_adapter_path = {} self.cur_adapter_path = {}
self.cur_adapter_strength = {}
self.lora_layers = {} self.lora_layers = {}
self.lora_layers_critic = {} self.lora_layers_critic = {}
self.lora_layers_transformer_2 = {} self.lora_layers_transformer_2 = {}
@@ -234,6 +236,7 @@ class LoRAPipeline(ComposedPipelineBase):
lora_nickname: str, lora_nickname: str,
lora_path: str | None, lora_path: str | None,
rank: int, rank: int,
strength: float = 1.0,
) -> int: ) -> int:
""" """
Apply LoRA weights to the given lora_layers. Apply LoRA weights to the given lora_layers.
@@ -243,6 +246,7 @@ class LoRAPipeline(ComposedPipelineBase):
lora_nickname: The nickname of the LoRA adapter. lora_nickname: The nickname of the LoRA adapter.
lora_path: The path to the LoRA adapter. lora_path: The path to the LoRA adapter.
rank: The distributed rank (for logging). rank: The distributed rank (for logging).
strength: LoRA strength for merge, default 1.0.
Returns: Returns:
The number of layers that had LoRA weights applied. The number of layers that had LoRA weights applied.
@@ -259,6 +263,7 @@ class LoRAPipeline(ComposedPipelineBase):
self.lora_adapters[lora_nickname][lora_A_name], self.lora_adapters[lora_nickname][lora_A_name],
self.lora_adapters[lora_nickname][lora_B_name], self.lora_adapters[lora_nickname][lora_B_name],
lora_path=lora_path, lora_path=lora_path,
strength=strength,
) )
adapted_count += 1 adapted_count += 1
else: else:
@@ -351,7 +356,11 @@ class LoRAPipeline(ComposedPipelineBase):
logger.info("Rank %d: loaded LoRA adapter %s", rank, lora_path) logger.info("Rank %d: loaded LoRA adapter %s", rank, lora_path)
def set_lora( def set_lora(
self, lora_nickname: str, lora_path: str | None = None, target: str = "all" self,
lora_nickname: str,
lora_path: str | None = None,
target: str = "all",
strength: float = 1.0,
): # type: ignore ): # type: ignore
""" """
Load a LoRA adapter into the pipeline and apply it to the specified transformer(s). Load a LoRA adapter into the pipeline and apply it to the specified transformer(s).
@@ -364,6 +373,7 @@ class LoRAPipeline(ComposedPipelineBase):
- "transformer": Apply only to the primary transformer (high noise for Wan2.2) - "transformer": Apply only to the primary transformer (high noise for Wan2.2)
- "transformer_2": Apply only to transformer_2 (low noise for Wan2.2) - "transformer_2": Apply only to transformer_2 (low noise for Wan2.2)
- "critic": Apply only to the critic model (fake_score_transformer) - "critic": Apply only to the critic model (fake_score_transformer)
strength: LoRA strength for merge, default 1.0.
""" """
if target not in self.VALID_TARGETS: if target not in self.VALID_TARGETS:
raise ValueError( raise ValueError(
@@ -410,11 +420,12 @@ class LoRAPipeline(ComposedPipelineBase):
adapter_updated = True adapter_updated = True
self.load_lora_adapter(lora_path, lora_nickname, rank) self.load_lora_adapter(lora_path, lora_nickname, rank)
# Check if we can skip (same adapter already applied to all target modules) # Check if we can skip (same adapter already applied to all target modules with same strength)
all_already_applied = all( all_already_applied = all(
not adapter_updated not adapter_updated
and self.cur_adapter_name.get(module_name) == lora_nickname and self.cur_adapter_name.get(module_name) == lora_nickname
and self.is_lora_merged.get(module_name, False) and self.is_lora_merged.get(module_name, False)
and self.cur_adapter_strength.get(module_name) == strength
for module_name, _ in target_modules for module_name, _ in target_modules
) )
if all_already_applied: if all_already_applied:
@@ -424,7 +435,7 @@ class LoRAPipeline(ComposedPipelineBase):
adapted_count = 0 adapted_count = 0
for module_name, lora_layers_dict in target_modules: for module_name, lora_layers_dict in target_modules:
count = self._apply_lora_to_layers( count = self._apply_lora_to_layers(
lora_layers_dict, lora_nickname, lora_path, rank lora_layers_dict, lora_nickname, lora_path, rank, strength
) )
adapted_count += count adapted_count += count
self.cur_adapter_name[module_name] = lora_nickname self.cur_adapter_name[module_name] = lora_nickname
@@ -432,16 +443,18 @@ class LoRAPipeline(ComposedPipelineBase):
lora_path or self.loaded_adapter_paths.get(lora_nickname, "") lora_path or self.loaded_adapter_paths.get(lora_nickname, "")
) )
self.is_lora_merged[module_name] = True self.is_lora_merged[module_name] = True
self.cur_adapter_strength[module_name] = strength
logger.info( logger.info(
"Rank %d: LoRA adapter %s applied to %d layers (target: %s)", "Rank %d: LoRA adapter %s applied to %d layers (target: %s, strength: %s)",
rank, rank,
lora_path, lora_path,
adapted_count, adapted_count,
target, target,
strength,
) )
def merge_lora_weights(self, target: str = "all") -> None: def merge_lora_weights(self, target: str = "all", strength: float = 1.0) -> None:
""" """
Merge LoRA weights into the base model for the specified target. Merge LoRA weights into the base model for the specified target.
@@ -450,6 +463,7 @@ class LoRAPipeline(ComposedPipelineBase):
Args: Args:
target: Which transformer(s) to merge. One of "all", "transformer", target: Which transformer(s) to merge. One of "all", "transformer",
"transformer_2", "critic". "transformer_2", "critic".
strength: LoRA strength for merge, default 1.0.
""" """
target_modules, error = self._get_target_lora_layers(target) target_modules, error = self._get_target_lora_layers(target)
if error: if error:
@@ -459,8 +473,19 @@ class LoRAPipeline(ComposedPipelineBase):
for module_name, lora_layers_dict in target_modules: for module_name, lora_layers_dict in target_modules:
if self.is_lora_merged.get(module_name, False): if self.is_lora_merged.get(module_name, False):
logger.warning("LoRA weights are already merged for %s", module_name) # Check if strength is the same - if so, skip (idempotent)
continue if self.cur_adapter_strength.get(module_name) == strength:
logger.warning(
"LoRA weights are already merged for %s with same strength",
module_name,
)
continue
# Different strength requested - allow re-merge (layer handles unmerge internally)
logger.info(
"Re-merging LoRA weights for %s with new strength %s",
module_name,
strength,
)
for name, layer in lora_layers_dict.items(): for name, layer in lora_layers_dict.items():
# Only re-enable LoRA for layers that actually have LoRA weights # Only re-enable LoRA for layers that actually have LoRA weights
has_lora_weights = hasattr(layer, "lora_A") and layer.lora_A is not None has_lora_weights = hasattr(layer, "lora_A") and layer.lora_A is not None
@@ -469,12 +494,15 @@ class LoRAPipeline(ComposedPipelineBase):
if hasattr(layer, "disable_lora"): if hasattr(layer, "disable_lora"):
layer.disable_lora = False layer.disable_lora = False
try: try:
layer.merge_lora_weights() layer.merge_lora_weights(strength=strength)
except Exception as e: except Exception as e:
logger.warning("Could not merge layer %s: %s", name, e) logger.warning("Could not merge layer %s: %s", name, e)
continue continue
self.is_lora_merged[module_name] = True self.is_lora_merged[module_name] = True
logger.info("LoRA weights merged for %s", module_name) self.cur_adapter_strength[module_name] = strength
logger.info(
"LoRA weights merged for %s (strength: %s)", module_name, strength
)
def unmerge_lora_weights(self, target: str = "all") -> None: def unmerge_lora_weights(self, target: str = "all") -> None:
""" """
@@ -519,4 +547,5 @@ class LoRAPipeline(ComposedPipelineBase):
layer.disable_lora = True layer.disable_lora = True
continue continue
self.is_lora_merged[module_name] = False self.is_lora_merged[module_name] = False
self.cur_adapter_strength.pop(module_name, None)
logger.info("LoRA weights unmerged for %s", module_name) logger.info("LoRA weights unmerged for %s", module_name)