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,