[diffusion] feat: support native and peft minimax h3 loras (#34359)

This commit is contained in:
Mick
2026-08-12 17:52:40 +08:00
committed by GitHub
parent 00e57d74f0
commit 644d55ebfa
16 changed files with 300 additions and 31 deletions
@@ -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
@@ -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
@@ -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
@@ -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
@@ -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
@@ -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()
@@ -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]):
@@ -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 = {}
@@ -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,
@@ -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,
@@ -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():
@@ -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}",
]
@@ -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,