[diffusion] feat: support native and peft minimax h3 loras (#34359)
This commit is contained in:
@@ -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,
|
||||
|
||||
Reference in New Issue
Block a user