[diffusion] feat: support lora strength (#15691)
This commit is contained in:
@@ -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}'
|
||||
```
|
||||
|
||||
|
||||
|
||||
@@ -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",
|
||||
)
|
||||
|
||||
|
||||
@@ -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",
|
||||
)
|
||||
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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:
|
||||
"""
|
||||
|
||||
@@ -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]):
|
||||
|
||||
@@ -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)
|
||||
|
||||
Reference in New Issue
Block a user