[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

View File

@@ -248,6 +248,8 @@ Loads a LoRA adapter and merges its weights into the model.
**Parameters:**
- `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
- `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:**
@@ -256,7 +258,8 @@ curl -X POST http://localhost:30010/v1/set_lora \
-H "Content-Type: application/json" \
-d '{
"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`
**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:**
```bash
curl -X POST http://localhost:30010/v1/merge_lora_weights \
-H "Content-Type: application/json"
-H "Content-Type: application/json" \
-d '{"strength": 0.8}'
```

View File

@@ -303,7 +303,11 @@ class DiffGenerator:
raise RuntimeError(f"{failure_msg}: {error_msg}")
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:
"""
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_2": Apply only to transformer_2 (low noise for Wan2.2)
- "critic": Apply only to the critic model
strength: LoRA strength for merge, default 1.0.
"""
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(
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",
)
@@ -340,17 +348,18 @@ class DiffGenerator:
"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.
Args:
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(
req,
f"Successfully merged LoRA weights (target: {target})",
f"Successfully merged LoRA weights (target: {target}, strength: {strength})",
"Failed to merge LoRA weights",
)

View File

@@ -39,6 +39,7 @@ async def set_lora(
lora_nickname: str = Body(..., embed=True),
lora_path: Optional[str] = Body(None, embed=True),
target: str = Body("all", embed=True),
strength: float = Body(1.0, embed=True),
):
"""
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_2": Apply only to transformer_2 (low noise for Wan2.2)
- "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(
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",
)
@@ -63,6 +71,7 @@ async def set_lora(
@router.post("/merge_lora_weights")
async def merge_lora_weights(
target: str = Body("all", embed=True),
strength: float = Body(1.0, embed=True),
):
"""
Merge LoRA weights into the base model.
@@ -70,11 +79,13 @@ async def merge_lora_weights(
Args:
target: Which transformer(s) to merge. One of "all", "transformer",
"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(
req,
f"Successfully merged LoRA weights (target: {target})",
f"Successfully merged LoRA weights (target: {target}, strength: {strength})",
"Failed to merge LoRA weights",
)

View File

@@ -25,11 +25,13 @@ class SetLoraReq:
lora_nickname: str
lora_path: Optional[str] = None
target: str = "all" # "all", "transformer", "transformer_2", "critic"
strength: float = 1.0 # LoRA strength for merge, default 1.0
@dataclasses.dataclass
class MergeLoraWeightsReq:
target: str = "all" # "all", "transformer", "transformer_2", "critic"
strength: float = 1.0 # LoRA strength for merge, default 1.0
@dataclasses.dataclass

View File

@@ -56,6 +56,7 @@ class BaseLayerWithLoRA(nn.Module):
self.lora_rank = lora_rank
self.lora_alpha = lora_alpha
self.lora_path: str | None = None
self.strength: float = 1.0
self.lora_A = None
self.lora_B = None
@@ -84,6 +85,7 @@ class BaseLayerWithLoRA(nn.Module):
delta = delta * (
self.lora_alpha / self.lora_rank # type: ignore
) # type: ignore
delta = delta * self.strength
out, output_bias = self.base_layer(x)
return out + delta, output_bias
else:
@@ -101,17 +103,22 @@ class BaseLayerWithLoRA(nn.Module):
A: torch.Tensor,
B: torch.Tensor,
lora_path: str | None = None,
strength: float = 1.0,
) -> None:
self.lora_A = torch.nn.Parameter(
A
) # share storage with weights in the pipeline
self.lora_B = torch.nn.Parameter(B)
self.disable_lora = False
self.strength = strength
self.merge_lora_weights()
self.lora_path = lora_path
@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:
return
@@ -136,9 +143,14 @@ class BaseLayerWithLoRA(nn.Module):
data = self.base_layer.weight.data.to(
get_local_torch_device()
).full_tensor()
data += self.slice_lora_b_weights(self.lora_B).to(
lora_delta = self.slice_lora_b_weights(self.lora_B).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))
if isinstance(getattr(self.base_layer, "bias", None), DTensor):
unsharded_base_layer.bias = nn.Parameter(
@@ -161,9 +173,14 @@ class BaseLayerWithLoRA(nn.Module):
else:
current_device = self.base_layer.weight.data.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.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.merged = True
@@ -391,6 +408,7 @@ class LinearWithLoRA(BaseLayerWithLoRA):
delta = delta * (
self.lora_alpha / self.lora_rank # type: ignore
) # type: ignore
delta = delta * self.strength
# nn.Linear.forward() returns a single tensor, not a tuple
out = self.base_layer(x)
return out + delta

View File

@@ -129,7 +129,11 @@ class GPUWorker:
return output_batch
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:
"""
Set the LoRA adapter for the pipeline.
@@ -138,19 +142,21 @@ class GPUWorker:
lora_nickname: The nickname of the adapter.
lora_path: Path to the LoRA adapter.
target: Which transformer(s) to apply the LoRA to.
strength: LoRA strength for merge, default 1.0.
"""
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.
Args:
target: Which transformer(s) to merge.
strength: LoRA strength for merge, default 1.0.
"""
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:
"""

View File

@@ -85,12 +85,12 @@ class Scheduler:
def _handle_set_lora(self, reqs: List[Any]):
# TODO: return set status
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"}
def _handle_merge_lora(self, reqs: List[Any]):
req = reqs[0]
self.worker.merge_lora_weights(req.target)
self.worker.merge_lora_weights(req.target, req.strength)
return {"status": "ok"}
def _handle_unmerge_lora(self, reqs: List[Any]):

View File

@@ -46,6 +46,7 @@ class LoRAPipeline(ComposedPipelineBase):
# Track current adapter per module: {"transformer": "high_lora", "transformer_2": "low_lora"}
cur_adapter_name: 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
lora_layers: dict[str, BaseLayerWithLoRA]
lora_layers_critic: dict[str, BaseLayerWithLoRA]
@@ -71,6 +72,7 @@ class LoRAPipeline(ComposedPipelineBase):
self.loaded_adapter_paths = {}
self.cur_adapter_name = {}
self.cur_adapter_path = {}
self.cur_adapter_strength = {}
self.lora_layers = {}
self.lora_layers_critic = {}
self.lora_layers_transformer_2 = {}
@@ -234,6 +236,7 @@ class LoRAPipeline(ComposedPipelineBase):
lora_nickname: str,
lora_path: str | None,
rank: int,
strength: float = 1.0,
) -> int:
"""
Apply LoRA weights to the given lora_layers.
@@ -243,6 +246,7 @@ class LoRAPipeline(ComposedPipelineBase):
lora_nickname: The nickname of the LoRA adapter.
lora_path: The path to the LoRA adapter.
rank: The distributed rank (for logging).
strength: LoRA strength for merge, default 1.0.
Returns:
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_B_name],
lora_path=lora_path,
strength=strength,
)
adapted_count += 1
else:
@@ -351,7 +356,11 @@ class LoRAPipeline(ComposedPipelineBase):
logger.info("Rank %d: loaded LoRA adapter %s", rank, lora_path)
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
"""
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_2": Apply only to transformer_2 (low noise for Wan2.2)
- "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:
raise ValueError(
@@ -410,11 +420,12 @@ class LoRAPipeline(ComposedPipelineBase):
adapter_updated = True
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(
not adapter_updated
and self.cur_adapter_name.get(module_name) == lora_nickname
and self.is_lora_merged.get(module_name, False)
and self.cur_adapter_strength.get(module_name) == strength
for module_name, _ in target_modules
)
if all_already_applied:
@@ -424,7 +435,7 @@ class LoRAPipeline(ComposedPipelineBase):
adapted_count = 0
for module_name, lora_layers_dict in target_modules:
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
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, "")
)
self.is_lora_merged[module_name] = True
self.cur_adapter_strength[module_name] = strength
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,
lora_path,
adapted_count,
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.
@@ -450,6 +463,7 @@ class LoRAPipeline(ComposedPipelineBase):
Args:
target: Which transformer(s) to merge. One of "all", "transformer",
"transformer_2", "critic".
strength: LoRA strength for merge, default 1.0.
"""
target_modules, error = self._get_target_lora_layers(target)
if error:
@@ -459,8 +473,19 @@ class LoRAPipeline(ComposedPipelineBase):
for module_name, lora_layers_dict in target_modules:
if self.is_lora_merged.get(module_name, False):
logger.warning("LoRA weights are already merged for %s", module_name)
continue
# Check if strength is the same - if so, skip (idempotent)
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():
# 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
@@ -469,12 +494,15 @@ class LoRAPipeline(ComposedPipelineBase):
if hasattr(layer, "disable_lora"):
layer.disable_lora = False
try:
layer.merge_lora_weights()
layer.merge_lora_weights(strength=strength)
except Exception as e:
logger.warning("Could not merge layer %s: %s", name, e)
continue
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:
"""
@@ -519,4 +547,5 @@ class LoRAPipeline(ComposedPipelineBase):
layer.disable_lora = True
continue
self.is_lora_merged[module_name] = False
self.cur_adapter_strength.pop(module_name, None)
logger.info("LoRA weights unmerged for %s", module_name)