[diffusion] feat: support loading peft lora (#35868)
This commit is contained in:
@@ -82,6 +82,7 @@ Use `sglang generate --help` and `sglang serve --help` for the full argument lis
|
||||
- `--lora-path {PATH}` and `--lora-nickname {NAME}`: load a LoRA adapter from a local path, Hugging Face repo/subfolder, or exact Hub file URL
|
||||
- `--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.
|
||||
- PEFT `adapter_config.json` semantics are applied automatically for named adapter slots, RSLoRA, and per-layer `alpha_pattern`. Adapters that require unsupported auxiliary parameters or runtime behavior, such as DoRA, fail before weight injection instead of silently using ordinary LoRA math.
|
||||
- `--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.
|
||||
|
||||
@@ -1093,7 +1093,7 @@ class GPUWorker(GPUWorkerPostTrainingMixin):
|
||||
"""
|
||||
List loaded LoRA adapters and current application status per module.
|
||||
"""
|
||||
from sglang.multimodal_gen.runtime.pipelines_core.lora_pipeline import (
|
||||
from sglang.multimodal_gen.runtime.pipelines_core.lora.pipeline import (
|
||||
LoRAPipeline,
|
||||
)
|
||||
|
||||
|
||||
@@ -11,7 +11,7 @@ import os
|
||||
from sglang.multimodal_gen.runtime.pipelines_core.composed_pipeline_base import (
|
||||
ComposedPipelineBase,
|
||||
)
|
||||
from sglang.multimodal_gen.runtime.pipelines_core.lora_pipeline import LoRAPipeline
|
||||
from sglang.multimodal_gen.runtime.pipelines_core.lora.pipeline import LoRAPipeline
|
||||
from sglang.multimodal_gen.runtime.pipelines_core.stages.model_specific_stages.cosmos3 import (
|
||||
Cosmos3DecodingStage,
|
||||
Cosmos3DenoisingStage,
|
||||
|
||||
@@ -7,7 +7,7 @@ import os
|
||||
from sglang.multimodal_gen.runtime.pipelines_core.composed_pipeline_base import (
|
||||
ComposedPipelineBase,
|
||||
)
|
||||
from sglang.multimodal_gen.runtime.pipelines_core.lora_pipeline import LoRAPipeline
|
||||
from sglang.multimodal_gen.runtime.pipelines_core.lora.pipeline import LoRAPipeline
|
||||
from sglang.multimodal_gen.runtime.pipelines_core.stages.input_validation import (
|
||||
InputValidationStage,
|
||||
)
|
||||
|
||||
@@ -9,7 +9,7 @@ using the modular pipeline architecture. Phase 1: T2V only.
|
||||
from sglang.multimodal_gen.runtime.pipelines_core.composed_pipeline_base import (
|
||||
ComposedPipelineBase,
|
||||
)
|
||||
from sglang.multimodal_gen.runtime.pipelines_core.lora_pipeline import LoRAPipeline
|
||||
from sglang.multimodal_gen.runtime.pipelines_core.lora.pipeline import LoRAPipeline
|
||||
from sglang.multimodal_gen.runtime.pipelines_core.stages import (
|
||||
InputValidationStage,
|
||||
)
|
||||
|
||||
@@ -3,7 +3,7 @@
|
||||
from sglang.multimodal_gen.runtime.pipelines_core.composed_pipeline_base import (
|
||||
ComposedPipelineBase,
|
||||
)
|
||||
from sglang.multimodal_gen.runtime.pipelines_core.lora_pipeline import LoRAPipeline
|
||||
from sglang.multimodal_gen.runtime.pipelines_core.lora.pipeline import LoRAPipeline
|
||||
from sglang.multimodal_gen.runtime.pipelines_core.stages import (
|
||||
DenoisingStage,
|
||||
InputValidationStage,
|
||||
|
||||
@@ -12,7 +12,7 @@ from sglang.multimodal_gen.runtime.models.schedulers.scheduling_self_forcing_flo
|
||||
from sglang.multimodal_gen.runtime.pipelines_core.composed_pipeline_base import (
|
||||
ComposedPipelineBase,
|
||||
)
|
||||
from sglang.multimodal_gen.runtime.pipelines_core.lora_pipeline import LoRAPipeline
|
||||
from sglang.multimodal_gen.runtime.pipelines_core.lora.pipeline import LoRAPipeline
|
||||
from sglang.multimodal_gen.runtime.pipelines_core.stages import (
|
||||
AuxiliaryConditionEncodingStage,
|
||||
DMDTimestepPreparationStage,
|
||||
|
||||
@@ -29,7 +29,7 @@ from sglang.multimodal_gen.runtime.models.schedulers.scheduling_flow_match_euler
|
||||
from sglang.multimodal_gen.runtime.pipelines_core.composed_pipeline_base import (
|
||||
ComposedPipelineBase,
|
||||
)
|
||||
from sglang.multimodal_gen.runtime.pipelines_core.lora_pipeline import LoRAPipeline
|
||||
from sglang.multimodal_gen.runtime.pipelines_core.lora.pipeline import LoRAPipeline
|
||||
from sglang.multimodal_gen.runtime.pipelines_core.schedule_batch import Req
|
||||
from sglang.multimodal_gen.runtime.pipelines_core.stages import (
|
||||
InputValidationStage,
|
||||
|
||||
@@ -9,7 +9,7 @@ from sglang.multimodal_gen.runtime.disaggregation.roles import RoleType
|
||||
from sglang.multimodal_gen.runtime.pipelines_core.composed_pipeline_base import (
|
||||
ComposedPipelineBase,
|
||||
)
|
||||
from sglang.multimodal_gen.runtime.pipelines_core.lora_pipeline import LoRAPipeline
|
||||
from sglang.multimodal_gen.runtime.pipelines_core.lora.pipeline import LoRAPipeline
|
||||
from sglang.multimodal_gen.runtime.pipelines_core.stages import InputValidationStage
|
||||
from sglang.multimodal_gen.runtime.pipelines_core.stages.model_specific_stages.minimax_h3 import (
|
||||
MiniMaxH3AudioEncodingStage,
|
||||
|
||||
@@ -8,7 +8,7 @@ import torch
|
||||
from sglang.multimodal_gen.runtime.pipelines_core.composed_pipeline_base import (
|
||||
ComposedPipelineBase,
|
||||
)
|
||||
from sglang.multimodal_gen.runtime.pipelines_core.lora_pipeline import LoRAPipeline
|
||||
from sglang.multimodal_gen.runtime.pipelines_core.lora.pipeline import LoRAPipeline
|
||||
from sglang.multimodal_gen.runtime.pipelines_core.schedule_batch import Req
|
||||
from sglang.multimodal_gen.runtime.pipelines_core.stages import (
|
||||
InputValidationStage,
|
||||
|
||||
@@ -10,7 +10,7 @@ This module wires the causal DMD denoising stage into the modular pipeline.
|
||||
from sglang.multimodal_gen.runtime.pipelines_core.composed_pipeline_base import (
|
||||
ComposedPipelineBase,
|
||||
)
|
||||
from sglang.multimodal_gen.runtime.pipelines_core.lora_pipeline import LoRAPipeline
|
||||
from sglang.multimodal_gen.runtime.pipelines_core.lora.pipeline import LoRAPipeline
|
||||
|
||||
# isort: off
|
||||
from sglang.multimodal_gen.runtime.pipelines_core.stages import (
|
||||
|
||||
@@ -14,7 +14,7 @@ from sglang.multimodal_gen.runtime.models.schedulers.scheduling_flow_match_euler
|
||||
from sglang.multimodal_gen.runtime.pipelines_core.composed_pipeline_base import (
|
||||
ComposedPipelineBase,
|
||||
)
|
||||
from sglang.multimodal_gen.runtime.pipelines_core.lora_pipeline import LoRAPipeline
|
||||
from sglang.multimodal_gen.runtime.pipelines_core.lora.pipeline import LoRAPipeline
|
||||
from sglang.multimodal_gen.runtime.server_args import ServerArgs
|
||||
|
||||
# isort: off
|
||||
|
||||
@@ -14,7 +14,7 @@ from sglang.multimodal_gen.runtime.models.schedulers.scheduling_flow_match_euler
|
||||
from sglang.multimodal_gen.runtime.pipelines_core.composed_pipeline_base import (
|
||||
ComposedPipelineBase,
|
||||
)
|
||||
from sglang.multimodal_gen.runtime.pipelines_core.lora_pipeline import LoRAPipeline
|
||||
from sglang.multimodal_gen.runtime.pipelines_core.lora.pipeline import LoRAPipeline
|
||||
from sglang.multimodal_gen.runtime.pipelines_core.stages import DmdDenoisingStage
|
||||
from sglang.multimodal_gen.runtime.server_args import ServerArgs
|
||||
|
||||
|
||||
@@ -14,7 +14,7 @@ from sglang.multimodal_gen.runtime.models.schedulers.scheduling_flow_unipc_multi
|
||||
from sglang.multimodal_gen.runtime.pipelines_core.composed_pipeline_base import (
|
||||
ComposedPipelineBase,
|
||||
)
|
||||
from sglang.multimodal_gen.runtime.pipelines_core.lora_pipeline import LoRAPipeline
|
||||
from sglang.multimodal_gen.runtime.pipelines_core.lora.pipeline import LoRAPipeline
|
||||
from sglang.multimodal_gen.runtime.server_args import ServerArgs
|
||||
|
||||
|
||||
|
||||
@@ -14,7 +14,7 @@ from sglang.multimodal_gen.runtime.models.schedulers.scheduling_flow_unipc_multi
|
||||
from sglang.multimodal_gen.runtime.pipelines_core.composed_pipeline_base import (
|
||||
ComposedPipelineBase,
|
||||
)
|
||||
from sglang.multimodal_gen.runtime.pipelines_core.lora_pipeline import LoRAPipeline
|
||||
from sglang.multimodal_gen.runtime.pipelines_core.lora.pipeline import LoRAPipeline
|
||||
from sglang.multimodal_gen.runtime.pipelines_core.stages import (
|
||||
InputValidationStage,
|
||||
)
|
||||
|
||||
@@ -13,7 +13,7 @@ from sglang.multimodal_gen.registry import get_model_info
|
||||
from sglang.multimodal_gen.runtime.pipelines_core.composed_pipeline_base import (
|
||||
ComposedPipelineBase,
|
||||
)
|
||||
from sglang.multimodal_gen.runtime.pipelines_core.lora_pipeline import LoRAPipeline
|
||||
from sglang.multimodal_gen.runtime.pipelines_core.lora.pipeline import LoRAPipeline
|
||||
from sglang.multimodal_gen.runtime.pipelines_core.schedule_batch import Req
|
||||
from sglang.multimodal_gen.runtime.server_args import ServerArgs
|
||||
from sglang.multimodal_gen.runtime.utils.hf_diffusers_utils import (
|
||||
|
||||
@@ -0,0 +1 @@
|
||||
"""LoRA loading and pipeline integration."""
|
||||
+10
-1
@@ -2,11 +2,16 @@ from __future__ import annotations
|
||||
|
||||
import logging
|
||||
from enum import Enum
|
||||
from typing import Dict, Iterable, Mapping, Optional
|
||||
from typing import Any, Dict, Iterable, Mapping, Optional
|
||||
|
||||
import torch
|
||||
from diffusers.loaders import lora_conversion_utils as lcu
|
||||
|
||||
from sglang.multimodal_gen.runtime.pipelines_core.lora.peft_adapter import (
|
||||
apply_peft_config,
|
||||
normalize_peft_keys,
|
||||
)
|
||||
|
||||
logger = logging.getLogger("LoRAFormatAdapter")
|
||||
|
||||
|
||||
@@ -539,10 +544,13 @@ def convert_lora_state_dict_by_format(
|
||||
def normalize_lora_state_dict(
|
||||
state_dict: Mapping[str, torch.Tensor],
|
||||
logger: Optional[logging.Logger] = None,
|
||||
*,
|
||||
adapter_config: Mapping[str, Any] | None = None,
|
||||
) -> Dict[str, torch.Tensor]:
|
||||
"""Normalize any supported LoRA format into a single canonical layout."""
|
||||
log = logger or globals()["logger"]
|
||||
|
||||
state_dict = normalize_peft_keys(state_dict)
|
||||
keys = list(state_dict.keys())
|
||||
log.info(
|
||||
"[LoRAFormatAdapter] normalize_lora_state_dict called, #keys=%d",
|
||||
@@ -558,6 +566,7 @@ def normalize_lora_state_dict(
|
||||
log.info("[LoRAFormatAdapter] detected format: %s", fmt)
|
||||
|
||||
normalized = convert_lora_state_dict_by_format(state_dict, fmt, log)
|
||||
normalized = apply_peft_config(normalized, adapter_config or {})
|
||||
|
||||
norm_keys = list(normalized.keys())
|
||||
if norm_keys:
|
||||
@@ -0,0 +1,211 @@
|
||||
"""Adapt PEFT checkpoint semantics to native diffusion LoRA layers."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
import math
|
||||
import re
|
||||
from pathlib import Path
|
||||
from typing import Any, Mapping
|
||||
|
||||
import torch
|
||||
|
||||
_ADAPTER_SLOT = re.compile(r"(\.lora_[AB])\.([^.]+)\.weight$")
|
||||
_WRAPPER_PREFIXES = ("peft_model.base_model.model.", "base_model.model.")
|
||||
_UNSUPPORTED_CONFIG_FIELDS = (
|
||||
"alora_invocation_tokens",
|
||||
"layer_replication",
|
||||
"modules_to_save",
|
||||
"target_parameters",
|
||||
"trainable_token_indices",
|
||||
"use_bdlora",
|
||||
"use_qalora",
|
||||
)
|
||||
|
||||
|
||||
def load_peft_config(weight_path: str) -> dict[str, Any]:
|
||||
path = Path(weight_path).with_name("adapter_config.json")
|
||||
if not path.is_file():
|
||||
return {}
|
||||
with path.open(encoding="utf-8") as file:
|
||||
config = json.load(file)
|
||||
if not isinstance(config, dict):
|
||||
raise ValueError("PEFT adapter_config.json must contain a JSON object")
|
||||
return config
|
||||
|
||||
|
||||
def get_peft_lora_alpha(config: Mapping[str, Any]) -> int | None:
|
||||
alpha = config.get("lora_alpha")
|
||||
if alpha is None:
|
||||
return None
|
||||
if (
|
||||
isinstance(alpha, bool)
|
||||
or not isinstance(alpha, (int, float))
|
||||
or alpha <= 0
|
||||
or isinstance(alpha, float)
|
||||
and not alpha.is_integer()
|
||||
):
|
||||
raise ValueError("PEFT lora_alpha must be a positive integer")
|
||||
return int(alpha)
|
||||
|
||||
|
||||
def normalize_peft_keys(
|
||||
state_dict: Mapping[str, torch.Tensor],
|
||||
) -> dict[str, torch.Tensor]:
|
||||
"""Remove a uniform PEFT model wrapper and named adapter slot."""
|
||||
prefix = next(
|
||||
(
|
||||
prefix
|
||||
for prefix in _WRAPPER_PREFIXES
|
||||
if state_dict and all(name.startswith(prefix) for name in state_dict)
|
||||
),
|
||||
"",
|
||||
)
|
||||
normalized: dict[str, torch.Tensor] = {}
|
||||
slots = set()
|
||||
has_bare_weights = False
|
||||
for name, tensor in state_dict.items():
|
||||
name = name.removeprefix(prefix)
|
||||
match = _ADAPTER_SLOT.search(name)
|
||||
if match is not None:
|
||||
slots.add(match.group(2))
|
||||
elif name.endswith((".lora_A.weight", ".lora_B.weight")):
|
||||
has_bare_weights = True
|
||||
target = _ADAPTER_SLOT.sub(r"\1.weight", name)
|
||||
if target in normalized:
|
||||
raise ValueError(
|
||||
"LoRA checkpoint contains multiple PEFT adapter slots for "
|
||||
f"the same tensor: {target!r}"
|
||||
)
|
||||
normalized[target] = tensor
|
||||
if len(slots) > 1:
|
||||
raise ValueError(
|
||||
f"LoRA checkpoint contains multiple PEFT adapter slots: {sorted(slots)}"
|
||||
)
|
||||
if slots and has_bare_weights:
|
||||
raise ValueError("LoRA checkpoint mixes named and unnamed PEFT adapter slots")
|
||||
return normalized
|
||||
|
||||
|
||||
def _validate_peft_features(
|
||||
state_dict: Mapping[str, torch.Tensor], config: Mapping[str, Any]
|
||||
) -> None:
|
||||
unsupported = {name for name in _UNSUPPORTED_CONFIG_FIELDS if config.get(name)}
|
||||
if config.get("bias") not in (None, "none"):
|
||||
unsupported.add("bias")
|
||||
unsupported.update(
|
||||
name for name in ("fan_in_fan_out", "lora_bias") if config.get(name, False)
|
||||
)
|
||||
if unsupported:
|
||||
raise ValueError(
|
||||
"PEFT adapter requires unsupported auxiliary/runtime features: "
|
||||
f"{sorted(unsupported)}"
|
||||
)
|
||||
if config.get("use_dora", False) or any(
|
||||
"lora_magnitude_vector" in name or "dora_scale" in name for name in state_dict
|
||||
):
|
||||
raise ValueError(
|
||||
"DoRA adapters are not supported by the native diffusion LoRA layers"
|
||||
)
|
||||
auxiliary = [
|
||||
name
|
||||
for name in state_dict
|
||||
if any(
|
||||
marker in name
|
||||
for marker in ("lora_embedding_", "modules_to_save", "trainable_tokens")
|
||||
)
|
||||
]
|
||||
if auxiliary:
|
||||
raise ValueError(
|
||||
f"PEFT adapter contains unsupported auxiliary tensors: {auxiliary[:8]}"
|
||||
)
|
||||
|
||||
|
||||
def apply_peft_config(
|
||||
state_dict: dict[str, torch.Tensor], config: Mapping[str, Any]
|
||||
) -> dict[str, torch.Tensor]:
|
||||
"""Represent PEFT scaling variants with native LoRA tensors and alpha."""
|
||||
_validate_peft_features(state_dict, config)
|
||||
normalized = dict(state_dict)
|
||||
patterns = config.get("alpha_pattern", {})
|
||||
if not isinstance(patterns, dict):
|
||||
raise ValueError("PEFT alpha_pattern must be an object")
|
||||
compiled_patterns = []
|
||||
for pattern, value in patterns.items():
|
||||
if not isinstance(pattern, str):
|
||||
raise ValueError("PEFT alpha_pattern keys must be strings")
|
||||
try:
|
||||
expression = re.compile(rf"(.*\.)?({pattern})$")
|
||||
except re.error as error:
|
||||
raise ValueError(
|
||||
f"Invalid PEFT alpha_pattern expression {pattern!r}: {error}"
|
||||
) from error
|
||||
alpha = get_peft_lora_alpha({"lora_alpha": value})
|
||||
if alpha is None:
|
||||
raise ValueError(
|
||||
f"PEFT alpha_pattern value for {pattern!r} must be a positive integer"
|
||||
)
|
||||
compiled_patterns.append((expression, alpha))
|
||||
|
||||
suffix = ".lora_A.weight"
|
||||
for name in state_dict:
|
||||
if not name.endswith(suffix):
|
||||
continue
|
||||
base = name[: -len(suffix)]
|
||||
alpha = next(
|
||||
(alpha for pattern, alpha in compiled_patterns if pattern.match(base)), None
|
||||
)
|
||||
if alpha is None:
|
||||
continue
|
||||
alpha_key = f"{base}.alpha"
|
||||
existing = normalized.get(alpha_key)
|
||||
if existing is not None and (
|
||||
existing.numel() != 1 or int(existing.item()) != alpha
|
||||
):
|
||||
raise ValueError(f"PEFT alpha_pattern conflicts with tensor {alpha_key!r}")
|
||||
normalized[alpha_key] = torch.tensor(alpha)
|
||||
|
||||
use_rslora = config.get("use_rslora", False)
|
||||
if not isinstance(use_rslora, bool):
|
||||
raise ValueError("PEFT use_rslora must be boolean")
|
||||
if not use_rslora:
|
||||
return normalized
|
||||
|
||||
suffix = ".lora_B.weight"
|
||||
pairs = 0
|
||||
for name, weight in state_dict.items():
|
||||
if not name.endswith(suffix):
|
||||
continue
|
||||
lora_a = state_dict.get(f"{name[: -len(suffix)]}.lora_A.weight")
|
||||
if lora_a is None or lora_a.ndim < 2 or lora_a.shape[-2] <= 0:
|
||||
raise ValueError(f"RSLoRA weight {name!r} has no valid rank-bearing A")
|
||||
normalized[name] = weight * math.sqrt(lora_a.shape[-2])
|
||||
pairs += 1
|
||||
if not pairs:
|
||||
raise ValueError("PEFT use_rslora is true, but no LoRA A/B pairs were found")
|
||||
return normalized
|
||||
|
||||
|
||||
def scale_fused_sections(
|
||||
a_parts: Mapping[int, torch.Tensor],
|
||||
b_parts: Mapping[int, torch.Tensor],
|
||||
alpha_parts: Mapping[int, torch.Tensor],
|
||||
default_alpha: int | None,
|
||||
) -> list[torch.Tensor] | None:
|
||||
"""Fold per-section PEFT alpha values into fused LoRA B weights."""
|
||||
if not alpha_parts:
|
||||
return None
|
||||
scaled = []
|
||||
for index in range(len(a_parts)):
|
||||
rank = a_parts[index].shape[0]
|
||||
alpha = float(
|
||||
alpha_parts[index].item()
|
||||
if index in alpha_parts
|
||||
else default_alpha if default_alpha is not None else rank
|
||||
)
|
||||
scale = alpha / rank
|
||||
weight = b_parts[index]
|
||||
scaled.append(
|
||||
weight if scale == 1.0 else (weight.float() * scale).to(weight.dtype)
|
||||
)
|
||||
return scaled
|
||||
+25
-14
@@ -1,7 +1,6 @@
|
||||
# Copied and adapted from: https://github.com/hao-ai-lab/FastVideo
|
||||
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
import json
|
||||
import os
|
||||
from collections import defaultdict
|
||||
from collections.abc import Hashable
|
||||
@@ -26,9 +25,14 @@ from sglang.multimodal_gen.runtime.managers.memory_managers.layerwise_offload im
|
||||
from sglang.multimodal_gen.runtime.pipelines_core.composed_pipeline_base import (
|
||||
ComposedPipelineBase,
|
||||
)
|
||||
from sglang.multimodal_gen.runtime.pipelines_core.lora_format_adapter import (
|
||||
from sglang.multimodal_gen.runtime.pipelines_core.lora.format_adapter import (
|
||||
normalize_lora_state_dict,
|
||||
)
|
||||
from sglang.multimodal_gen.runtime.pipelines_core.lora.peft_adapter import (
|
||||
get_peft_lora_alpha,
|
||||
load_peft_config,
|
||||
scale_fused_sections,
|
||||
)
|
||||
from sglang.multimodal_gen.runtime.server_args import LORA_MERGE_MODES, ServerArgs
|
||||
from sglang.multimodal_gen.runtime.utils.hf_diffusers_utils import maybe_download_lora
|
||||
from sglang.multimodal_gen.runtime.utils.logging_utils import init_logger
|
||||
@@ -98,11 +102,19 @@ def _store_fused_lora_groups(
|
||||
or set(b_parts) != set(range(n))
|
||||
):
|
||||
continue
|
||||
a, b, fused_alpha = stack_or_compose_fused_lora(
|
||||
[a_parts[i] for i in range(n)],
|
||||
[b_parts[i] for i in range(n)],
|
||||
a_list = [a_parts[i] for i in range(n)]
|
||||
b_list = [b_parts[i] for i in range(n)]
|
||||
scaled_b = scale_fused_sections(
|
||||
a_parts,
|
||||
b_parts,
|
||||
to_merge_params.get(f"{base}.alpha", {}),
|
||||
adapter_alpha,
|
||||
)
|
||||
a, b, fused_alpha = stack_or_compose_fused_lora(
|
||||
a_list, scaled_b or b_list, None if scaled_b else adapter_alpha
|
||||
)
|
||||
if scaled_b:
|
||||
fused_alpha = a.shape[-2]
|
||||
adapter[str(a_key)] = a.to(device)
|
||||
adapter[b_key] = b.to(device)
|
||||
if fused_alpha is not None:
|
||||
@@ -823,16 +835,15 @@ class LoRAPipeline(ComposedPipelineBase):
|
||||
lora_local_path = maybe_download_lora(lora_path, weight_name=weight_name)
|
||||
|
||||
raw_state_dict = load_file(lora_local_path)
|
||||
lora_state_dict = normalize_lora_state_dict(raw_state_dict, logger=logger)
|
||||
adapter_lora_alpha = lora_alpha
|
||||
adapter_config_path = os.path.join(
|
||||
os.path.dirname(lora_local_path), "adapter_config.json"
|
||||
adapter_config = load_peft_config(lora_local_path)
|
||||
lora_state_dict = normalize_lora_state_dict(
|
||||
raw_state_dict,
|
||||
logger=logger,
|
||||
adapter_config=adapter_config,
|
||||
)
|
||||
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:
|
||||
adapter_lora_alpha = int(adapter_config["lora_alpha"])
|
||||
adapter_lora_alpha = lora_alpha
|
||||
if adapter_lora_alpha is None:
|
||||
adapter_lora_alpha = get_peft_lora_alpha(adapter_config)
|
||||
|
||||
if lora_nickname in self.lora_adapters:
|
||||
self.lora_adapters[lora_nickname].clear()
|
||||
@@ -58,7 +58,7 @@ from sglang.multimodal_gen.runtime.managers.memory_managers.layerwise_offload im
|
||||
is_layerwise_offloaded_module,
|
||||
)
|
||||
from sglang.multimodal_gen.runtime.pipelines.diffusers_pipeline import DiffusersPipeline
|
||||
from sglang.multimodal_gen.runtime.pipelines_core.lora_pipeline import (
|
||||
from sglang.multimodal_gen.runtime.pipelines_core.lora.pipeline import (
|
||||
LoRAPipeline,
|
||||
stack_or_compose_fused_lora,
|
||||
)
|
||||
|
||||
@@ -11,8 +11,9 @@ import torch.nn.functional as F
|
||||
from sglang.multimodal_gen.runtime.layers.lora.linear import (
|
||||
MergedColumnParallelLinearWithLoRA,
|
||||
)
|
||||
from sglang.multimodal_gen.runtime.pipelines_core.lora_pipeline import (
|
||||
from sglang.multimodal_gen.runtime.pipelines_core.lora.pipeline import (
|
||||
LoRAPipeline,
|
||||
_store_fused_lora_groups,
|
||||
stack_or_compose_fused_lora,
|
||||
)
|
||||
from sglang.multimodal_gen.runtime.post_training.weights_updater import (
|
||||
@@ -64,6 +65,31 @@ def test_compose_unequal_sections_matches_reference():
|
||||
)
|
||||
|
||||
|
||||
def test_fused_sections_preserve_per_layer_alpha():
|
||||
a_list, b_list = _make_ab_lists([2, 3, 1])
|
||||
alphas = [2, 6, 1]
|
||||
pending = defaultdict(dict)
|
||||
for index, (lora_a, lora_b, alpha) in enumerate(zip(a_list, b_list, alphas)):
|
||||
pending["attn.qkv.lora_A"][index] = lora_a
|
||||
pending["attn.qkv.lora_B"][index] = lora_b
|
||||
pending["attn.qkv.alpha"][index] = torch.tensor(alpha)
|
||||
|
||||
adapter = {}
|
||||
_store_fused_lora_groups(adapter, pending, adapter_alpha=4, device="cpu")
|
||||
|
||||
x = torch.randn(5, IN_DIM)
|
||||
actual = x @ adapter["attn.qkv.lora_A"].T @ adapter["attn.qkv.lora_B"].T
|
||||
expected = torch.cat(
|
||||
[
|
||||
(x @ lora_a.T @ lora_b.T) * (alpha / lora_a.shape[0])
|
||||
for lora_a, lora_b, alpha in zip(a_list, b_list, alphas)
|
||||
],
|
||||
dim=-1,
|
||||
)
|
||||
assert adapter["attn.qkv.alpha"].item() == sum(a.shape[0] for a in a_list)
|
||||
torch.testing.assert_close(actual, expected, rtol=1e-5, atol=1e-5)
|
||||
|
||||
|
||||
def test_stack_kept_for_equal_sections():
|
||||
torch.manual_seed(0)
|
||||
a_list = [torch.randn(2, IN_DIM) for _ in range(2)]
|
||||
@@ -272,7 +298,7 @@ def _make_loader_pipeline() -> _TestLoRAPipeline:
|
||||
|
||||
|
||||
def _load_adapter(pipeline, state_dict, lora_alpha=None):
|
||||
loader_mod = "sglang.multimodal_gen.runtime.pipelines_core.lora_pipeline"
|
||||
loader_mod = "sglang.multimodal_gen.runtime.pipelines_core.lora.pipeline"
|
||||
with (
|
||||
patch(f"{loader_mod}.maybe_download_lora", return_value="/adapter"),
|
||||
patch(f"{loader_mod}.load_file", return_value=state_dict),
|
||||
@@ -344,7 +370,7 @@ def test_apply_composed_adapter_end_to_end():
|
||||
|
||||
strength = 2.0
|
||||
with patch(
|
||||
"sglang.multimodal_gen.runtime.pipelines_core.lora_pipeline.dist.get_rank",
|
||||
"sglang.multimodal_gen.runtime.pipelines_core.lora.pipeline.dist.get_rank",
|
||||
return_value=0,
|
||||
):
|
||||
applied = pipeline._apply_lora_to_layers(
|
||||
|
||||
@@ -702,7 +702,7 @@ class TestGGUFRejectsLoraConversion(unittest.TestCase):
|
||||
"""The dynamic set_lora path must refuse before replacing any layer."""
|
||||
|
||||
def _pipeline_with(self, layer):
|
||||
from sglang.multimodal_gen.runtime.pipelines_core.lora_pipeline import (
|
||||
from sglang.multimodal_gen.runtime.pipelines_core.lora.pipeline import (
|
||||
LoRAPipeline,
|
||||
)
|
||||
|
||||
|
||||
@@ -16,7 +16,7 @@ import torch
|
||||
from huggingface_hub import hf_hub_download
|
||||
from safetensors.torch import load_file
|
||||
|
||||
from sglang.multimodal_gen.runtime.pipelines_core.lora_format_adapter import (
|
||||
from sglang.multimodal_gen.runtime.pipelines_core.lora.format_adapter import (
|
||||
LoRAFormat,
|
||||
detect_lora_format_from_state_dict,
|
||||
normalize_lora_state_dict,
|
||||
|
||||
@@ -0,0 +1,55 @@
|
||||
"""PEFT LoRA normalization and fail-closed capability tests."""
|
||||
|
||||
import math
|
||||
|
||||
import pytest
|
||||
import torch
|
||||
|
||||
from sglang.multimodal_gen.runtime.pipelines_core.lora.format_adapter import (
|
||||
normalize_lora_state_dict,
|
||||
)
|
||||
from sglang.multimodal_gen.runtime.pipelines_core.lora.peft_adapter import (
|
||||
get_peft_lora_alpha,
|
||||
)
|
||||
|
||||
|
||||
def test_peft_wrapper_slot_and_rslora_scaling_preserve_delta():
|
||||
lora_a = torch.randn(4, 8)
|
||||
lora_b = torch.randn(16, 4)
|
||||
normalized = normalize_lora_state_dict(
|
||||
{
|
||||
"base_model.model.transformer.proj.lora_A.default.weight": lora_a,
|
||||
"base_model.model.transformer.proj.lora_B.default.weight": lora_b,
|
||||
},
|
||||
adapter_config={
|
||||
"use_rslora": True,
|
||||
"lora_alpha": 8,
|
||||
"alpha_pattern": {"transformer.proj": 8},
|
||||
},
|
||||
)
|
||||
|
||||
normalized_b = normalized["transformer.proj.lora_B.weight"]
|
||||
ordinary_delta = (8 / 4) * normalized_b @ lora_a
|
||||
expected_delta = (8 / math.sqrt(4)) * lora_b @ lora_a
|
||||
torch.testing.assert_close(ordinary_delta, expected_delta)
|
||||
assert normalized["transformer.proj.alpha"].item() == 8
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
("state_dict", "adapter_config"),
|
||||
[
|
||||
({}, {"use_dora": True}),
|
||||
({}, {"modules_to_save": ["head"]}),
|
||||
({}, {"target_parameters": ["experts.weight"]}),
|
||||
({}, {"fan_in_fan_out": True}),
|
||||
({"encoder.lora_embedding_A.weight": torch.ones(2, 2)}, {}),
|
||||
],
|
||||
)
|
||||
def test_unsupported_peft_runtime_semantics_fail_closed(state_dict, adapter_config):
|
||||
with pytest.raises(ValueError, match="not supported|unsupported"):
|
||||
normalize_lora_state_dict(state_dict, adapter_config=adapter_config)
|
||||
|
||||
|
||||
def test_invalid_peft_lora_alpha_fails_closed():
|
||||
with pytest.raises(ValueError, match="positive integer"):
|
||||
get_peft_lora_alpha({"lora_alpha": 8.5})
|
||||
@@ -6,10 +6,10 @@ from unittest.mock import patch
|
||||
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.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"
|
||||
_RANK_PATCH = "sglang.multimodal_gen.runtime.pipelines_core.lora.pipeline.dist.get_rank"
|
||||
|
||||
|
||||
class _TestLoRAPipeline(LoRAPipeline):
|
||||
|
||||
@@ -564,7 +564,7 @@ class TestLTX25OptionalDecoderLoading(unittest.TestCase):
|
||||
|
||||
def test_decoder_is_not_loaded_by_default(self):
|
||||
from sglang.multimodal_gen.runtime.pipelines.ltx_2_pipeline import LTX2Pipeline
|
||||
from sglang.multimodal_gen.runtime.pipelines_core.lora_pipeline import (
|
||||
from sglang.multimodal_gen.runtime.pipelines_core.lora.pipeline import (
|
||||
LoRAPipeline,
|
||||
)
|
||||
|
||||
@@ -578,7 +578,7 @@ class TestLTX25OptionalDecoderLoading(unittest.TestCase):
|
||||
|
||||
def test_decoder_load_is_explicit_and_validated(self):
|
||||
from sglang.multimodal_gen.runtime.pipelines.ltx_2_pipeline import LTX2Pipeline
|
||||
from sglang.multimodal_gen.runtime.pipelines_core.lora_pipeline import (
|
||||
from sglang.multimodal_gen.runtime.pipelines_core.lora.pipeline import (
|
||||
LoRAPipeline,
|
||||
)
|
||||
|
||||
|
||||
Reference in New Issue
Block a user