From 644d55ebfae18914e8d0ac7bbc67fa425bb87c8e Mon Sep 17 00:00:00 2001 From: Mick Date: Wed, 12 Aug 2026 17:52:40 +0800 Subject: [PATCH] [diffusion] feat: support native and peft minimax h3 loras (#34359) --- .../cookbook/diffusion/MiniMax/MiniMax-H3.mdx | 73 ++++++++++++++++--- docs/docs/sglang-diffusion/api/cli.mdx | 2 + .../sglang-diffusion/compatibility_matrix.mdx | 2 +- .../configs/models/dits/minimax_h3.py | 55 +++++++++++++- .../entrypoints/diffusion_generator.py | 5 +- .../runtime/entrypoints/openai/common_api.py | 3 + .../runtime/entrypoints/utils.py | 1 + .../runtime/layers/lora/linear.py | 33 ++++++++- .../runtime/managers/gpu_worker.py | 8 +- .../runtime/managers/scheduler.py | 1 + .../runtime/pipelines_core/lora_pipeline.py | 43 ++++++++--- .../runtime/server_args/server_args.py | 12 +++ .../runtime/utils/hf_diffusers_utils.py | 9 ++- .../test/unit/test_lora_inference_mode.py | 15 +++- .../test/unit/test_lora_pipeline.py | 39 ++++++++++ .../test/unit/test_minimax_h3_dit_contract.py | 30 +++++++- 16 files changed, 300 insertions(+), 31 deletions(-) diff --git a/docs/cookbook/diffusion/MiniMax/MiniMax-H3.mdx b/docs/cookbook/diffusion/MiniMax/MiniMax-H3.mdx index d6eaf3eb0..77711c9c5 100644 --- a/docs/cookbook/diffusion/MiniMax/MiniMax-H3.mdx +++ b/docs/cookbook/diffusion/MiniMax/MiniMax-H3.mdx @@ -383,24 +383,75 @@ Poll and download any conditioned request with the same job-status and content endpoints used in the T2VA example. Server-local `file://` URIs must refer to files visible inside the SGLang server environment. -## 5. Turbo LoRA for few-step generation +## 5. LoRA recipes -[`larryvrh/MiniMax-H3-Turbo-Lora`](https://huggingface.co/larryvrh/MiniMax-H3-Turbo-Lora) distills the native **FL2VA** DiT for usable **4–8 step** generation. On `--model-variant fl2va`, use it with **`t2va`** or **`fl2va`** from section 4 and set `"num_inference_steps": 4` or `8` instead of `50`. **`ref2va`** uses a separate checkpoint partition and is not validated with this LoRA. +H3 accepts both native fused adapters and standard Diffusers/PEFT adapters. +Native adapters target modules such as `blocks.*.attn.qkv_proj`; PEFT adapters +may instead provide separate `to_q`, `to_k`, and `to_v` projections and the +`default` adapter namespace. SGLang normalizes both layouts. + +The following FL2VA adapters have distinct purposes: + +| Recipe | Repository and pinned file | Request setting | Prompt requirement | +| --- | --- | --- | --- | +| Recommended speed/quality balance | [`larryvrh/MiniMax-H3-Turbo-Lora`](https://huggingface.co/larryvrh/MiniMax-H3-Turbo-Lora), `minimax_h3_turbo_v4_step600_ema.safetensors` | `num_inference_steps: 9` (8 denoiser evaluations), `lora_scale: 1.0` | None | +| Most aggressive speed preset (standard PEFT layout) | [`lightx2v/Minimax-h3-Turbo`](https://huggingface.co/lightx2v/Minimax-h3-Turbo), `minimax_h3_fl2v_turbo_4step_v0.1.safetensors` | `num_inference_steps: 5` (4 denoiser evaluations), `lora_scale: 1.0`, `lora_alpha: 8` | None | +| Realistic people style | [`fal/MiniMax-H3-Realism-People-LoRA`](https://huggingface.co/fal/MiniMax-H3-Realism-People-LoRA), `h3-realism-people-t2v-i2v-r2v.safetensors` | Keep the normal `num_inference_steps: 50` schedule; start with `lora_scale: 0.7` | Include `r34l1sm` in the prompt | + +The H3 request field controls the number of sigma grid points, including the +terminal zero; the denoising loop therefore runs one fewer model evaluation. +This is why an adapter described as 8-step uses `9`, and a 4-step adapter uses +`5`, in the request. + +All three use the same launch shape. Pinning the filename is required for +repositories that publish multiple revisions, and is also recommended for a +reproducible single-file recipe: ```bash Command -curl -sS -X POST http://127.0.0.1:30010/v1/set_lora \ - -H "Content-Type: application/json" \ - -d '{ - "lora_nickname": "h3-turbo", - "lora_path": "larryvrh/MiniMax-H3-Turbo-Lora", - "strength": 1.0 - }' +LORA_REPO=larryvrh/MiniMax-H3-Turbo-Lora +LORA_FILE=minimax_h3_turbo_v4_step600_ema.safetensors +LORA_NAME=h3-turbo-v4 +LORA_SCALE=1.0 +LORA_ALPHA_ARGS=() +# LightX2V only: LORA_ALPHA_ARGS=(--lora-alpha 8) + +sglang serve \ + --model-path MiniMaxAI/MiniMax-H3 \ + --model-variant fl2va \ + --num-gpus 4 \ + --ulysses-degree 4 \ + --performance-mode speed \ + --lora-path "$LORA_REPO" \ + --lora-weight-name "$LORA_FILE" \ + --lora-nickname "$LORA_NAME" \ + --lora-scale "$LORA_SCALE" \ + "${LORA_ALPHA_ARGS[@]}" \ + --lora-merge-mode auto \ + --port 30010 ``` -When the repository contains multiple safetensors files, prefer `minimax_h3_turbo_4step_ckpt500.safetensors` (upstream default) via `--lora-path` and `--lora-weight-name` on `sglang serve`, or pass the local path to that file as `lora_path`. +`auto` merges an adapter into ordinary resident weights to avoid per-step LoRA +matmuls, but keeps the dynamic path for FSDP-sharded weights where a full +gather can increase peak memory. Use `dynamic` when one resident server must +switch repeatedly between base and LoRA output. + +Use the filename, scale, and request schedule from the table together. The +4-evaluation LightX2V recipe is the more aggressive latency/quality tradeoff. +Its checkpoint has rank 128 but omits the training alpha from both the file and +repository metadata, so `--lora-alpha 8` is required to reproduce the author's +reference implementation. Start with the Larry 8-evaluation recipe when +preserving fine visual detail is more important than minimum latency. + +These adapters were trained for the **FL2VA** partition and apply to `t2va` or +`fl2va` requests. Do not use them with the separate `ref2va` weights unless +the adapter author explicitly provides Ref2VA-compatible weights. Also avoid +stacking a distilled adapter with `quality: "high"`: both alter denoising, and +that combination has not been quality-validated. -LoRAs for the ComfyUI pruned MiniMax-H3 graph (for example H3-GalaxyAce) are not compatible with SGLang's native FL2VA weights. +LoRAs trained for a pruned or structurally modified ComfyUI graph are not +automatically compatible with the native H3 weights. Use only adapters whose +architecture and target modules match the full native H3 checkpoint. ## 6. Sampling and output controls diff --git a/docs/docs/sglang-diffusion/api/cli.mdx b/docs/docs/sglang-diffusion/api/cli.mdx index ba14864bd..2e150b3f8 100644 --- a/docs/docs/sglang-diffusion/api/cli.mdx +++ b/docs/docs/sglang-diffusion/api/cli.mdx @@ -79,6 +79,8 @@ Use `sglang generate --help` and `sglang serve --help` for the full argument lis - `--model-variant {NAME}`: semantic checkpoint variant to load when one model repository contains multiple weight partitions. The pipeline maps this stable name to the repository layout before loading; for example, MiniMax-H3 accepts `fl2va` and `ref2va`. This is a server/load-time choice, unlike a request's `task`. - `--model-subfolder {PATH}`: advanced direct override for a component subfolder inside the model repository. Prefer `--model-variant` when the pipeline exposes semantic routing. If both are supplied, they must resolve to the same weight partition. - `--lora-path {PATH}` and `--lora-nickname {NAME}`: load a LoRA adapter +- `--lora-weight-name {FILE}`: select one adapter file from a repository that contains multiple LoRA revisions. The Hub download is filtered to that file plus JSON metadata, so unused weights are not downloaded. +- `--lora-alpha {N}`: supply the training alpha when a single-file adapter omits both per-layer alpha tensors and `adapter_config.json`. Do not set it when the adapter already records alpha metadata. - `--lora-merge-mode {auto|merge|dynamic}`: choose how LoRA is applied. `auto` statically merges regular weights and uses dynamic LoRA for FSDP-sharded weights to avoid full-gather peaks. - `--num-gpus {N}`: number of GPUs to use - `--performance-mode {manual|auto|speed|memory}` / `--mode`: preset for latency/throughput and memory defaults. `auto` is the default and dispatches residency from selected-GPU headroom and workload type: image DiTs stay resident above the 45 GiB threshold, while video DiT placement remains model-specific. It uses FSDP only for validated DiT-offload replacement paths. `speed` keeps `torch.compile` disabled unless a model-specific deployment config opts in after validation; pass `--enable-torch-compile true` to enable it explicitly. Use `manual` to keep performance-related server args under explicit user control. Explicit offload, FSDP, and parallelism flags take precedence in all modes. diff --git a/docs/docs/sglang-diffusion/compatibility_matrix.mdx b/docs/docs/sglang-diffusion/compatibility_matrix.mdx index 8d7ef4c2d..1fcad4a8d 100644 --- a/docs/docs/sglang-diffusion/compatibility_matrix.mdx +++ b/docs/docs/sglang-diffusion/compatibility_matrix.mdx @@ -820,7 +820,7 @@ The entries below simply reflect configurations that have been manually validate MiniMax-H3 - `larryvrh/MiniMax-H3-Turbo-Lora` + `larryvrh/MiniMax-H3-Turbo-Lora`
`lightx2v/Minimax-h3-Turbo`
`fal/MiniMax-H3-Realism-People-LoRA` Wan2.2 diff --git a/python/sglang/multimodal_gen/configs/models/dits/minimax_h3.py b/python/sglang/multimodal_gen/configs/models/dits/minimax_h3.py index 606583552..07487013f 100644 --- a/python/sglang/multimodal_gen/configs/models/dits/minimax_h3.py +++ b/python/sglang/multimodal_gen/configs/models/dits/minimax_h3.py @@ -9,7 +9,60 @@ MINIMAX_H3_ADALN_MODALITY_NUM = 3 @dataclass class MiniMaxH3DiTArchConfig(DiTArchConfig): - lora_param_names_mapping: dict = field(default_factory=dict) + # accept Diffusers/PEFT aliases in the source-to-native model mapping + # H3 fuses Q/K/V, so split projections are stacked for the fused LoRA layer + param_names_mapping: dict = field( + default_factory=lambda: { + r"^(.*\.lora_[AB])\.[^.]+$": r"\1", + r"^base_model\.model\.(.*\.lora_[AB])$": r"\1", + r"^transformer\.(.*\.lora_[AB])$": r"\1", + r"^proj_in\.(lora_[AB])$": r"video_patch_proj.\1", + r"^audio_proj_in\.(lora_[AB])$": r"audio_patch_proj.\1", + r"^context_embedder\.(lora_[AB])$": r"condition_proj.\1", + r"^time_embedder\.linear_1\.(lora_[AB])$": r"time_embedder.proj_in.\1", + r"^time_embedder\.linear_2\.(lora_[AB])$": r"time_embedder.proj_out.\1", + r"^norm_out\.linear\.(lora_[AB])$": r"final_layer.adaln_proj.linear.\1", + r"^proj_out\.(lora_[AB])$": r"final_layer.video_out.\1", + r"^audio_proj_out\.(lora_[AB])$": r"final_layer.audio_out.\1", + r"^transformer_blocks\.(\d+)\.adaln_proj\.linear\.(lora_[AB])$": r"blocks.\1.adaln_proj.linear.\2", + r"^transformer_blocks\.(\d+)\.attn\.to_q\.(lora_[AB])$": ( + r"blocks.\1.attn.qkv_proj.\2", + 0, + 3, + ), + r"^transformer_blocks\.(\d+)\.attn\.to_k\.(lora_[AB])$": ( + r"blocks.\1.attn.qkv_proj.\2", + 1, + 3, + ), + r"^transformer_blocks\.(\d+)\.attn\.to_v\.(lora_[AB])$": ( + r"blocks.\1.attn.qkv_proj.\2", + 2, + 3, + ), + r"^transformer_blocks\.(\d+)\.attn\.to_out\.0\.(lora_[AB])$": r"blocks.\1.attn.out_proj.\2", + r"^transformer_blocks\.(\d+)\.ff\.net\.0\.proj\.(lora_[AB])$": r"blocks.\1.mlp.fc1.\2", + r"^transformer_blocks\.(\d+)\.ff\.net\.2\.(lora_[AB])$": r"blocks.\1.mlp.fc2.\2", + r"^token_refiner\.refiner_blocks\.(\d+)\.attn\.to_q\.(lora_[AB])$": ( + r"token_refiner.blocks.\1.attn.qkv_proj.\2", + 0, + 3, + ), + r"^token_refiner\.refiner_blocks\.(\d+)\.attn\.to_k\.(lora_[AB])$": ( + r"token_refiner.blocks.\1.attn.qkv_proj.\2", + 1, + 3, + ), + r"^token_refiner\.refiner_blocks\.(\d+)\.attn\.to_v\.(lora_[AB])$": ( + r"token_refiner.blocks.\1.attn.qkv_proj.\2", + 2, + 3, + ), + r"^token_refiner\.refiner_blocks\.(\d+)\.attn\.to_out\.0\.(lora_[AB])$": r"token_refiner.blocks.\1.attn.out_proj.\2", + r"^token_refiner\.refiner_blocks\.(\d+)\.ff\.net\.0\.proj\.(lora_[AB])$": r"token_refiner.blocks.\1.mlp.fc1.\2", + r"^token_refiner\.refiner_blocks\.(\d+)\.ff\.net\.2\.(lora_[AB])$": r"token_refiner.blocks.\1.mlp.fc2.\2", + } + ) num_layers: int = 50 token_refiner_num_layers: int = 2 diff --git a/python/sglang/multimodal_gen/runtime/entrypoints/diffusion_generator.py b/python/sglang/multimodal_gen/runtime/entrypoints/diffusion_generator.py index 084307e1d..e026ae895 100644 --- a/python/sglang/multimodal_gen/runtime/entrypoints/diffusion_generator.py +++ b/python/sglang/multimodal_gen/runtime/entrypoints/diffusion_generator.py @@ -13,7 +13,7 @@ import multiprocessing as mp import os import time from contextlib import ExitStack -from typing import Any, List, Union +from typing import Any, List, Optional, Union from sglang.multimodal_gen.configs.sample.sampling_params import ( DataType, @@ -547,6 +547,7 @@ class DiffGenerator: target: Union[str, List[str]] = "all", strength: Union[float, List[float]] = 1.0, merge_mode: str | None = None, + lora_alpha: Optional[Union[int, List[Optional[int]]]] = None, ) -> None: """ Set LoRA adapter(s) for the specified transformer(s). @@ -563,6 +564,7 @@ class DiffGenerator: - "critic": Apply only to the critic model strength: LoRA strength(s) for merge, default 1.0. Can be a float or a list of floats. merge_mode: Optional LoRA merge mode: "auto", "merge", or "dynamic". + lora_alpha: Training alpha override for adapters that omit it from metadata. """ req = SetLoraReq( lora_nickname=lora_nickname, @@ -570,6 +572,7 @@ class DiffGenerator: target=target, strength=strength, merge_mode=merge_mode, + lora_alpha=lora_alpha, ) nickname_str, target_str, strength_str = format_lora_message( lora_nickname, target, strength 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 a4a68545a..9a3e6d831 100644 --- a/python/sglang/multimodal_gen/runtime/entrypoints/openai/common_api.py +++ b/python/sglang/multimodal_gen/runtime/entrypoints/openai/common_api.py @@ -89,6 +89,7 @@ async def set_lora( target: Union[str, List[str]] = Body("all", embed=True), strength: Union[float, List[float]] = Body(1.0, embed=True), merge_mode: Optional[str] = Body(None, embed=True), + lora_alpha: Optional[Union[int, List[Optional[int]]]] = Body(None, embed=True), ): """ Set LoRA adapter(s) for the specified transformer(s). @@ -108,6 +109,7 @@ async def set_lora( If a list, must match the length of lora_nickname. Values < 1.0 reduce the effect, values > 1.0 amplify the effect. merge_mode: Optional LoRA merge mode: "auto", "merge", or "dynamic". + lora_alpha: Training alpha override for adapters that omit it from metadata. """ req = SetLoraReq( lora_nickname=lora_nickname, @@ -115,6 +117,7 @@ async def set_lora( target=target, strength=strength, merge_mode=merge_mode, + lora_alpha=lora_alpha, ) nickname_str, target_str, strength_str = format_lora_message( lora_nickname, target, strength diff --git a/python/sglang/multimodal_gen/runtime/entrypoints/utils.py b/python/sglang/multimodal_gen/runtime/entrypoints/utils.py index 1494f5847..ab51abf93 100644 --- a/python/sglang/multimodal_gen/runtime/entrypoints/utils.py +++ b/python/sglang/multimodal_gen/runtime/entrypoints/utils.py @@ -163,6 +163,7 @@ class SetLoraReq: target: Union[str, List[str]] = "all" strength: Union[float, List[float]] = 1.0 merge_mode: Optional[str] = None + lora_alpha: Optional[Union[int, List[Optional[int]]]] = None @dataclass diff --git a/python/sglang/multimodal_gen/runtime/layers/lora/linear.py b/python/sglang/multimodal_gen/runtime/layers/lora/linear.py index 47023e8f6..272331fce 100644 --- a/python/sglang/multimodal_gen/runtime/layers/lora/linear.py +++ b/python/sglang/multimodal_gen/runtime/layers/lora/linear.py @@ -47,6 +47,27 @@ LoRAWeightEntry = tuple[ ] +def _compute_lora_delta( + x: torch.Tensor, lora_A: torch.Tensor, lora_B: torch.Tensor +) -> torch.Tensor: + """Apply a regular or stacked LoRA projection to the last dimension.""" + if lora_A.dim() == 2 and lora_B.dim() == 2: + return x @ lora_A.T @ lora_B.T + if lora_A.dim() == 3 and lora_B.dim() == 3: + if lora_A.shape[0] != lora_B.shape[0]: + raise ValueError( + "Stacked LoRA A/B projections must have the same group count, got " + f"{lora_A.shape[0]} and {lora_B.shape[0]}" + ) + hidden = torch.einsum("...i,nri->...nr", x, lora_A) + delta = torch.einsum("...nr,nor->...no", hidden, lora_B) + return delta.flatten(start_dim=-2) + raise ValueError( + "LoRA A/B projections must both be 2D or both be 3D, got " + f"{tuple(lora_A.shape)} and {tuple(lora_B.shape)}" + ) + + class BaseLayerWithLoRA(nn.Module): def __init__( self, @@ -100,7 +121,7 @@ class BaseLayerWithLoRA(nn.Module): lora_B_sliced = self.slice_lora_b_weights( lora_B.to(device=x.device, non_blocking=True) ) - delta = x_lora @ lora_A_sliced.T @ lora_B_sliced.T + delta = _compute_lora_delta(x_lora, lora_A_sliced, lora_B_sliced) if self.lora_alpha != self.lora_rank: delta = delta * ( self.lora_alpha / self.lora_rank # type: ignore @@ -481,7 +502,9 @@ class ColumnParallelLinearWithLoRA(BaseLayerWithLoRA): lora_B_sliced = self.slice_lora_b_weights( lora_B.to(device=input_.device, non_blocking=True) ) - delta_parallel = input_lora @ lora_A_sliced.T @ lora_B_sliced.T + delta_parallel = _compute_lora_delta( + input_lora, lora_A_sliced, lora_B_sliced + ) if self.lora_alpha != self.lora_rank: delta_parallel = delta_parallel * ( self.lora_alpha / self.lora_rank # type: ignore @@ -616,7 +639,9 @@ class RowParallelLinearWithLoRA(BaseLayerWithLoRA): lora_B_sliced = self.slice_lora_b_weights( lora_B.to(device=input_parallel.device, non_blocking=True) ) - delta_parallel = input_parallel_lora @ lora_A_sliced.T @ lora_B_sliced.T + delta_parallel = _compute_lora_delta( + input_parallel_lora, lora_A_sliced, lora_B_sliced + ) if self.lora_alpha != self.lora_rank: delta_parallel = delta_parallel * ( self.lora_alpha / self.lora_rank # type: ignore @@ -688,7 +713,7 @@ class LinearWithLoRA(BaseLayerWithLoRA): lora_B_sliced = self.slice_lora_b_weights( lora_B.to(device=x.device, non_blocking=True) ) - delta = x_lora @ lora_A_sliced.T @ lora_B_sliced.T + delta = _compute_lora_delta(x_lora, lora_A_sliced, lora_B_sliced) if self.lora_alpha != self.lora_rank: delta = delta * ( self.lora_alpha / self.lora_rank # type: ignore diff --git a/python/sglang/multimodal_gen/runtime/managers/gpu_worker.py b/python/sglang/multimodal_gen/runtime/managers/gpu_worker.py index 1063d6a83..66f23b604 100644 --- a/python/sglang/multimodal_gen/runtime/managers/gpu_worker.py +++ b/python/sglang/multimodal_gen/runtime/managers/gpu_worker.py @@ -987,6 +987,7 @@ class GPUWorker(GPUWorkerPostTrainingMixin): target: Union[str, List[str]] = "all", strength: Union[float, List[float]] = 1.0, merge_mode: str | None = None, + lora_alpha: int | None | list[int | None] = None, ) -> OutputBatch: """ Set the LoRA adapter(s) for the pipeline. @@ -1002,7 +1003,12 @@ class GPUWorker(GPUWorkerPostTrainingMixin): if not isinstance(self.pipeline, LoRAPipeline): return OutputBatch(error="Lora is not enabled") self.pipeline.set_lora( - lora_nickname, lora_path, target, strength, merge_mode=merge_mode + lora_nickname, + lora_path, + target, + strength, + merge_mode=merge_mode, + lora_alpha=lora_alpha, ) return OutputBatch() diff --git a/python/sglang/multimodal_gen/runtime/managers/scheduler.py b/python/sglang/multimodal_gen/runtime/managers/scheduler.py index ec7e2cb37..c85af07e9 100644 --- a/python/sglang/multimodal_gen/runtime/managers/scheduler.py +++ b/python/sglang/multimodal_gen/runtime/managers/scheduler.py @@ -211,6 +211,7 @@ class Scheduler(SchedulerWarmupMixin, SchedulerPostTrainingMixin, SchedulerDisag req.target, req.strength, req.merge_mode, + req.lora_alpha, ) def _handle_merge_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 8f3b26aff..4b4ab0851 100644 --- a/python/sglang/multimodal_gen/runtime/pipelines_core/lora_pipeline.py +++ b/python/sglang/multimodal_gen/runtime/pipelines_core/lora_pipeline.py @@ -110,6 +110,7 @@ class LoRAPipeline(ComposedPipelineBase): self.lora_nickname, self.lora_path, strength=self.server_args.lora_scale, # type: ignore + lora_alpha=self.server_args.lora_alpha, ) # type: ignore def is_target_layer(self, module_name: str) -> bool: @@ -328,7 +329,8 @@ class LoRAPipeline(ComposedPipelineBase): lora_path: str | None | list[str | None], strength: float | list[float], target: str | list[str], - ) -> tuple[list[str], list[str | None], list[float], list[str]]: + lora_alpha: int | None | list[int | None], + ) -> tuple[list[str], list[str | None], list[float], list[str], list[int | None]]: """ Normalize LoRA parameters to lists for multi-LoRA support. @@ -374,7 +376,20 @@ class LoRAPipeline(ComposedPipelineBase): f"Length mismatch: lora_nickname has {len(lora_nicknames)} items, " f"but target has {len(targets)} items" ) - return lora_nicknames, lora_paths, strengths, targets + + lora_alphas = ( + lora_alpha + if isinstance(lora_alpha, list) + else [lora_alpha] * len(lora_nicknames) + ) + if len(lora_alphas) != len(lora_nicknames): + raise ValueError( + f"Length mismatch: lora_nickname has {len(lora_nicknames)} items, " + f"but lora_alpha has {len(lora_alphas)} items" + ) + if any(alpha is not None and alpha <= 0 for alpha in lora_alphas): + raise ValueError("lora_alpha values must be positive integers or null") + return lora_nicknames, lora_paths, strengths, targets, lora_alphas def _check_lora_config_matches( self, @@ -525,7 +540,7 @@ class LoRAPipeline(ComposedPipelineBase): and lora_B_name in self.lora_adapters[nickname] ): inferred_rank = int( - self.lora_adapters[nickname][lora_A_name].shape[0] + self.lora_adapters[nickname][lora_A_name].shape[-2] ) alpha_key = name + ".alpha" adapter_lora_alpha = self.loaded_adapter_alphas.get(nickname) @@ -686,6 +701,7 @@ class LoRAPipeline(ComposedPipelineBase): lora_nickname: str, rank: int, weight_name: str | None = None, + lora_alpha: int | None = None, ): """ Load the LoRA, and setup the lora_adapters for later weight replacement @@ -712,11 +728,11 @@ class LoRAPipeline(ComposedPipelineBase): raw_state_dict = load_file(lora_local_path) lora_state_dict = normalize_lora_state_dict(raw_state_dict, logger=logger) - adapter_lora_alpha = None + adapter_lora_alpha = lora_alpha adapter_config_path = os.path.join( os.path.dirname(lora_local_path), "adapter_config.json" ) - if os.path.isfile(adapter_config_path): + if adapter_lora_alpha is None and os.path.isfile(adapter_config_path): with open(adapter_config_path, encoding="utf-8") as f: adapter_config = json.load(f) if adapter_config.get("lora_alpha") is not None: @@ -764,6 +780,7 @@ class LoRAPipeline(ComposedPipelineBase): f"Dit target weight name {target_name} already exists in lora_adapters[{lora_nickname}]" ) self.lora_adapters[lora_nickname][target_name] = weight.to(self.device) + self.loaded_adapter_paths[lora_nickname] = lora_path self.loaded_adapter_alphas[lora_nickname] = adapter_lora_alpha logger.info("Rank %d: loaded LoRA adapter %s", rank, lora_path) @@ -776,6 +793,7 @@ class LoRAPipeline(ComposedPipelineBase): strength: float | list[float] = 1.0, merge_weights: bool | None = None, merge_mode: str | None = None, + lora_alpha: int | None | list[int | None] = None, ): # type: ignore """ Load LoRA adapter(s) into the pipeline and apply them to the specified transformer(s). @@ -784,8 +802,10 @@ class LoRAPipeline(ComposedPipelineBase): merge_mode = self._resolve_lora_merge_mode(merge_weights, merge_mode) # Normalize inputs to lists for multi-LoRA support - lora_nicknames, lora_paths, strengths, targets = self._normalize_lora_params( - lora_nickname, lora_path, strength, target + lora_nicknames, lora_paths, strengths, targets, lora_alphas = ( + self._normalize_lora_params( + lora_nickname, lora_path, strength, target, lora_alpha + ) ) # Validate targets @@ -809,7 +829,7 @@ class LoRAPipeline(ComposedPipelineBase): rank = dist.get_rank() # load required adapters - for nickname, path in zip(lora_nicknames, lora_paths): + for nickname, path, alpha in zip(lora_nicknames, lora_paths, lora_alphas): if nickname not in self.lora_adapters and path is None: raise ValueError( f"Adapter {nickname} not found in the pipeline. Please provide lora_path to load it." @@ -823,7 +843,12 @@ class LoRAPipeline(ComposedPipelineBase): should_load = True if should_load: adapter_updated = True - self.load_lora_adapter(path, nickname, rank) + self.load_lora_adapter(path, nickname, rank, lora_alpha=alpha) + elif ( + alpha is not None and self.loaded_adapter_alphas.get(nickname) != alpha + ): + self.loaded_adapter_alphas[nickname] = alpha + adapter_updated = True # Group by target to apply separately target_to_indices = {} diff --git a/python/sglang/multimodal_gen/runtime/server_args/server_args.py b/python/sglang/multimodal_gen/runtime/server_args/server_args.py index 93e5a79c5..5036b3cfc 100644 --- a/python/sglang/multimodal_gen/runtime/server_args/server_args.py +++ b/python/sglang/multimodal_gen/runtime/server_args/server_args.py @@ -282,6 +282,7 @@ class ServerArgs(DisaggServerArgsMixin): lora_path: str | None = None lora_nickname: str = "default" # for swapping adapters in the pipeline lora_scale: float = 1.0 # LoRA scale for merging (e.g., 0.125 for Hyper-SD) + lora_alpha: int | None = None # Override training alpha when metadata omits it lora_merge_mode: str = "auto" lora_weight_name: str | None = None @@ -525,6 +526,8 @@ class ServerArgs(DisaggServerArgsMixin): self._validate_pipeline() self._validate_offload() self._validate_direct_gpu_weight_loading() + if self.lora_alpha is not None and self.lora_alpha <= 0: + raise ValueError("lora_alpha must be a positive integer") if not current_platform.is_cpu(): self._validate_parallelism() self._validate_cfg_parallel() @@ -2170,6 +2173,15 @@ class ServerArgs(DisaggServerArgsMixin): default=ServerArgs.lora_scale, help="LoRA scale for merging (e.g., 0.125 for Hyper-SD). Same as lora_scale in Diffusers", ) + parser.add_argument( + "--lora-alpha", + type=int, + default=ServerArgs.lora_alpha, + help=( + "Override the LoRA training alpha when neither the checkpoint nor " + "adapter_config.json records it" + ), + ) parser.add_argument( "--lora-merge-mode", type=str, diff --git a/python/sglang/multimodal_gen/runtime/utils/hf_diffusers_utils.py b/python/sglang/multimodal_gen/runtime/utils/hf_diffusers_utils.py index c216e930a..817c84d19 100644 --- a/python/sglang/multimodal_gen/runtime/utils/hf_diffusers_utils.py +++ b/python/sglang/multimodal_gen/runtime/utils/hf_diffusers_utils.py @@ -594,7 +594,14 @@ def maybe_download_lora( Returns: Local path to the model """ - allow_patterns = ["*.json", "*.safetensors", "*.bin"] + # Repositories often publish several adapter revisions side by side. If a + # filename is pinned, do not download every weight before selecting it. + # Keep JSON metadata so PEFT's lora_alpha remains available. + allow_patterns = ( + ["*.json", weight_name, f"**/{weight_name}"] + if weight_name is not None + else ["*.json", "*.safetensors", "*.bin"] + ) local_path = maybe_download_model( model_name_or_path, diff --git a/python/sglang/multimodal_gen/test/unit/test_lora_inference_mode.py b/python/sglang/multimodal_gen/test/unit/test_lora_inference_mode.py index bac541163..ae1b7db42 100644 --- a/python/sglang/multimodal_gen/test/unit/test_lora_inference_mode.py +++ b/python/sglang/multimodal_gen/test/unit/test_lora_inference_mode.py @@ -1,7 +1,20 @@ import torch from torch import nn -from sglang.multimodal_gen.runtime.layers.lora.linear import LinearWithLoRA +from sglang.multimodal_gen.runtime.layers.lora.linear import ( + LinearWithLoRA, + _compute_lora_delta, +) + + +def test_stacked_lora_delta_preserves_projection_order(): + x = torch.tensor([[2.0, 3.0]]) + lora_a = torch.tensor([[[1.0, 0.0]], [[0.0, 1.0]]]) + lora_b = torch.tensor([[[1.0], [2.0]], [[3.0], [4.0]]]) + + actual = _compute_lora_delta(x, lora_a, lora_b) + + torch.testing.assert_close(actual, torch.tensor([[2.0, 4.0, 9.0, 12.0]])) def test_lora_merge_unmerge_handles_inference_base_weight(): diff --git a/python/sglang/multimodal_gen/test/unit/test_lora_pipeline.py b/python/sglang/multimodal_gen/test/unit/test_lora_pipeline.py index 8a55546b3..73a489258 100644 --- a/python/sglang/multimodal_gen/test/unit/test_lora_pipeline.py +++ b/python/sglang/multimodal_gen/test/unit/test_lora_pipeline.py @@ -7,6 +7,7 @@ import torch from sglang.multimodal_gen.runtime.layers.lora.linear import BaseLayerWithLoRA from sglang.multimodal_gen.runtime.pipelines_core.lora_pipeline import LoRAPipeline +from sglang.multimodal_gen.runtime.utils.hf_diffusers_utils import maybe_download_lora _RANK_PATCH = "sglang.multimodal_gen.runtime.pipelines_core.lora_pipeline.dist.get_rank" @@ -120,3 +121,41 @@ def test_merged_lora_still_uses_weight_update_context(): assert context_calls == 1 assert layer.merged assert pipeline.is_lora_merged["transformer"] + + +def test_lora_alpha_override_updates_cached_adapter_scale(): + layer = _make_layer() + pipeline = _make_pipeline(layer) + + with patch(_RANK_PATCH, return_value=0): + pipeline.set_lora( + "adapter", + None, + target="transformer", + strength=1.0, + merge_mode="dynamic", + lora_alpha=8, + ) + + assert pipeline.loaded_adapter_alphas["adapter"] == 8 + assert layer.lora_rank == 1 + assert layer.lora_alpha == 8 + + +def test_pinned_lora_weight_limits_snapshot_download(tmp_path): + weight_name = "adapter-v4.safetensors" + weight_path = tmp_path / weight_name + weight_path.touch() + + download_target = ( + "sglang.multimodal_gen.runtime.utils.hf_diffusers_utils.maybe_download_model" + ) + with patch(download_target, return_value=str(tmp_path)) as download: + actual = maybe_download_lora("org/multi-adapter", weight_name=weight_name) + + assert actual == str(weight_path) + assert download.call_args.kwargs["allow_patterns"] == [ + "*.json", + weight_name, + f"**/{weight_name}", + ] diff --git a/python/sglang/multimodal_gen/test/unit/test_minimax_h3_dit_contract.py b/python/sglang/multimodal_gen/test/unit/test_minimax_h3_dit_contract.py index d77414914..4a722ef48 100644 --- a/python/sglang/multimodal_gen/test/unit/test_minimax_h3_dit_contract.py +++ b/python/sglang/multimodal_gen/test/unit/test_minimax_h3_dit_contract.py @@ -43,7 +43,6 @@ def _ensure_single_process_parallel_runtime() -> None: def test_native_weight_names_and_grouped_qkv_reorder(): arch = MiniMaxH3DiTArchConfig() - assert arch.param_names_mapping == {} assert arch.reverse_param_names_mapping == {} mapping = get_param_names_mapping(arch.param_names_mapping) for key in ( @@ -54,6 +53,35 @@ def test_native_weight_names_and_grouped_qkv_reorder(): ): assert mapping(key) == (key, None, None) + assert mapping( + "base_model.model.transformer.transformer_blocks.7.attn.to_k.lora_A.default" + ) == ("blocks.7.attn.qkv_proj.lora_A", 1, 3) + assert mapping("token_refiner.refiner_blocks.1.ff.net.0.proj.lora_B") == ( + "token_refiner.blocks.1.mlp.fc1.lora_B", + None, + None, + ) + assert mapping("transformer.transformer_blocks.3.adaln_proj.linear.lora_A") == ( + "blocks.3.adaln_proj.linear.lora_A", + None, + None, + ) + assert mapping("transformer.audio_proj_out.lora_B") == ( + "final_layer.audio_out.lora_B", + None, + None, + ) + assert mapping("blocks.3.attn.out_proj.lora_A") == ( + "blocks.3.attn.out_proj.lora_A", + None, + None, + ) + assert mapping("transformer.blocks.0.attn.qkv_proj.weight") == ( + "transformer.blocks.0.attn.qkv_proj.weight", + None, + None, + ) + weight = torch.arange(12, dtype=torch.float32).reshape(12, 1) actual = _reorder_grouped_qkv_to_qkv( weight,