diff --git a/python/sglang/multimodal_gen/docs/openai_api.md b/python/sglang/multimodal_gen/docs/openai_api.md index eae8ef6cf..ff896855d 100644 --- a/python/sglang/multimodal_gen/docs/openai_api.md +++ b/python/sglang/multimodal_gen/docs/openai_api.md @@ -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}' ``` diff --git a/python/sglang/multimodal_gen/runtime/entrypoints/diffusion_generator.py b/python/sglang/multimodal_gen/runtime/entrypoints/diffusion_generator.py index 993e8b6aa..65ac0abcb 100644 --- a/python/sglang/multimodal_gen/runtime/entrypoints/diffusion_generator.py +++ b/python/sglang/multimodal_gen/runtime/entrypoints/diffusion_generator.py @@ -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", ) diff --git a/python/sglang/multimodal_gen/runtime/entrypoints/openai/common_api.py b/python/sglang/multimodal_gen/runtime/entrypoints/openai/common_api.py index fc7623a42..2a9dc2892 100644 --- a/python/sglang/multimodal_gen/runtime/entrypoints/openai/common_api.py +++ b/python/sglang/multimodal_gen/runtime/entrypoints/openai/common_api.py @@ -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", ) diff --git a/python/sglang/multimodal_gen/runtime/entrypoints/openai/utils.py b/python/sglang/multimodal_gen/runtime/entrypoints/openai/utils.py index 839b40d0a..aa6eb626d 100644 --- a/python/sglang/multimodal_gen/runtime/entrypoints/openai/utils.py +++ b/python/sglang/multimodal_gen/runtime/entrypoints/openai/utils.py @@ -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 diff --git a/python/sglang/multimodal_gen/runtime/layers/lora/linear.py b/python/sglang/multimodal_gen/runtime/layers/lora/linear.py index c507e9ea2..496af4fd9 100644 --- a/python/sglang/multimodal_gen/runtime/layers/lora/linear.py +++ b/python/sglang/multimodal_gen/runtime/layers/lora/linear.py @@ -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 diff --git a/python/sglang/multimodal_gen/runtime/managers/gpu_worker.py b/python/sglang/multimodal_gen/runtime/managers/gpu_worker.py index 8358b35b3..bddcdffe5 100644 --- a/python/sglang/multimodal_gen/runtime/managers/gpu_worker.py +++ b/python/sglang/multimodal_gen/runtime/managers/gpu_worker.py @@ -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: """ diff --git a/python/sglang/multimodal_gen/runtime/managers/scheduler.py b/python/sglang/multimodal_gen/runtime/managers/scheduler.py index 7fb9a2e81..c3e0c2890 100644 --- a/python/sglang/multimodal_gen/runtime/managers/scheduler.py +++ b/python/sglang/multimodal_gen/runtime/managers/scheduler.py @@ -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]): diff --git a/python/sglang/multimodal_gen/runtime/pipelines_core/lora_pipeline.py b/python/sglang/multimodal_gen/runtime/pipelines_core/lora_pipeline.py index c61cb793f..2a1fc258e 100644 --- a/python/sglang/multimodal_gen/runtime/pipelines_core/lora_pipeline.py +++ b/python/sglang/multimodal_gen/runtime/pipelines_core/lora_pipeline.py @@ -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)