[diffusion] fix: fix diffusion LoRA consistency cases (#26327)

This commit is contained in:
Mick
2026-05-28 06:29:35 +08:00
committed by GitHub
parent deaba74745
commit 24bcb37efb
6 changed files with 74 additions and 13 deletions
@@ -246,7 +246,7 @@ class BaseLayerWithLoRA(nn.Module):
self,
lora_list: list[LoRAWeightEntry],
) -> bool:
if os.getenv("SGLANG_DIFFUSION_LORA_MERGE_FP32", "0") != "1":
if os.getenv("SGLANG_DIFFUSION_LORA_MERGE_FP32", "1") != "1":
return False
for _, _, lora_path, _, _, _ in lora_list:
if lora_path and "distilled-lora" in lora_path.lower():
@@ -257,6 +257,18 @@ class BaseLayerWithLoRA(nn.Module):
def merge_lora_weights(self, strength: float | None = None) -> None:
if strength is not None:
self.strength = strength
if self.lora_weights_list:
self.lora_weights_list = [
(lora_A, lora_B, lora_path, strength, lora_rank, lora_alpha)
for (
lora_A,
lora_B,
lora_path,
_,
lora_rank,
lora_alpha,
) in self.lora_weights_list
]
if self.disable_lora:
return
@@ -1,6 +1,7 @@
# 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
@@ -48,6 +49,8 @@ class LoRAPipeline(ComposedPipelineBase):
# e.g., [jinx][transformer_blocks.0.attn.to_v.lora_A]
lora_adapters: dict[str, dict[str, torch.Tensor]]
loaded_adapter_paths: dict[str, str] # nickname -> lora_path
loaded_adapter_alphas: dict[str, int | None]
# nickname -> adapter_config lora_alpha
# Track current adapter per module: {"transformer": "high_lora", "transformer_2": "low_lora"}
cur_adapter_name: dict[str, str]
cur_adapter_path: dict[str, str]
@@ -77,6 +80,7 @@ class LoRAPipeline(ComposedPipelineBase):
# Initialize all mutable instance attributes to avoid sharing across instances
self.lora_adapters = defaultdict(dict)
self.loaded_adapter_paths = {}
self.loaded_adapter_alphas = {}
self.cur_adapter_name = {}
self.cur_adapter_path = {}
self.cur_adapter_strength = {}
@@ -524,10 +528,13 @@ class LoRAPipeline(ComposedPipelineBase):
self.lora_adapters[nickname][lora_A_name].shape[0]
)
alpha_key = name + ".alpha"
adapter_lora_alpha = self.loaded_adapter_alphas.get(nickname)
if alpha_key in self.lora_adapters[nickname]:
inferred_alpha = int(
self.lora_adapters[nickname][alpha_key].item()
)
elif adapter_lora_alpha is not None:
inferred_alpha = adapter_lora_alpha
else:
# Some distilled LoRAs omit per-layer alpha and rely on the
# default LoRA scale of alpha == rank. Falling back to rank
@@ -705,6 +712,15 @@ 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_config_path = os.path.join(
os.path.dirname(lora_local_path), "adapter_config.json"
)
if 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"])
if lora_nickname in self.lora_adapters:
self.lora_adapters[lora_nickname].clear()
@@ -749,6 +765,7 @@ class LoRAPipeline(ComposedPipelineBase):
)
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)
def set_lora(
@@ -824,7 +841,10 @@ class LoRAPipeline(ComposedPipelineBase):
continue
tgt_nicknames = [lora_nicknames[i] for i in idx_list]
tgt_paths = [lora_paths[i] for i in idx_list]
tgt_paths = [
lora_paths[i] or self.loaded_adapter_paths.get(lora_nicknames[i])
for i in idx_list
]
tgt_strengths = [strengths[i] for i in idx_list]
merged_name = (
@@ -546,6 +546,8 @@ TWO_GPU_CASES = [
extras=[
"--lora-weight-name",
"985347-wan22_14B-low-Nfj1nx-e65.safetensors",
"--lora-merge-mode",
"dynamic",
],
),
DiffusionSamplingParams(
@@ -757,6 +757,22 @@ Pinned revision used by this check: {SGL_TEST_FILES_CI_DATA_REVISION}
output_path.write_bytes(content)
logger.info(f"Saved GT image: {output_path} (format: {detected_format})")
def _validate_lora_consistency(
self, case: DiffusionTestCase, content: bytes, operation: str
) -> None:
if not case.run_consistency_check:
logger.info(
"[LoRA Consistency] Skipping %s consistency for %s: disabled for case",
operation,
case.id,
)
return
logger.info(
"[LoRA Consistency] Validating %s output for %s", operation, case.id
)
self._validate_consistency(case, content)
def _test_lora_api_functionality(
self,
ctx: ServerContext,
@@ -792,10 +808,11 @@ Pinned revision used by this check: {SGL_TEST_FILES_CI_DATA_REVISION}
assert resp.status_code == 200, f"merge_lora_weights failed: {resp.text}"
logger.info("[LoRA E2E] Verifying generation after re-merge for %s", case.id)
rid_after_merge, _ = self._run_generation_with_server_watchdog(
ctx, case.id, generate_fn, client
rid_after_merge, content_after_merge = (
self._run_generation_with_server_watchdog(ctx, case.id, generate_fn, client)
)
assert rid_after_merge is not None, "Generation after merge failed"
self._validate_lora_consistency(case, content_after_merge, "merge_lora_weights")
logger.info("[LoRA E2E] Generation after merge succeeded")
# Test 3: set_lora (re-set the same adapter) - API should succeed and generation should work
@@ -808,10 +825,11 @@ Pinned revision used by this check: {SGL_TEST_FILES_CI_DATA_REVISION}
assert resp.status_code == 200, f"set_lora failed: {resp.text}"
logger.info("[LoRA E2E] Verifying generation after set_lora for %s", case.id)
rid_after_set, _ = self._run_generation_with_server_watchdog(
rid_after_set, content_after_set = self._run_generation_with_server_watchdog(
ctx, case.id, generate_fn, client
)
assert rid_after_set is not None, "Generation after set_lora failed"
self._validate_lora_consistency(case, content_after_set, "set_lora")
logger.info("[LoRA E2E] Generation after set_lora succeeded")
# Test 4: list_loras - API should return the expected list of LoRA adapters
@@ -850,10 +868,13 @@ Pinned revision used by this check: {SGL_TEST_FILES_CI_DATA_REVISION}
logger.info(
"[LoRA Switch E2E] Testing generation with initial LoRA for %s", case.id
)
rid_initial, _ = self._run_generation_with_server_watchdog(
rid_initial, content_initial = self._run_generation_with_server_watchdog(
ctx, case.id, generate_fn, client
)
assert rid_initial is not None, "Generation with initial LoRA failed"
self._validate_lora_consistency(
case, content_initial, "dynamic switch initial LoRA"
)
logger.info("[LoRA Switch E2E] Generation with initial LoRA succeeded")
# Test 2: Switch to second LoRA and generate
@@ -891,10 +912,13 @@ Pinned revision used by this check: {SGL_TEST_FILES_CI_DATA_REVISION}
"[LoRA Switch E2E] Verifying generation after switching back for %s",
case.id,
)
rid_switched_back, _ = self._run_generation_with_server_watchdog(
ctx, case.id, generate_fn, client
rid_switched_back, content_switched_back = (
self._run_generation_with_server_watchdog(ctx, case.id, generate_fn, client)
)
assert rid_switched_back is not None, "Generation after switching back failed"
self._validate_lora_consistency(
case, content_switched_back, "dynamic switch default LoRA"
)
logger.info("[LoRA Switch E2E] Generation after switching back succeeded")
logger.info(
@@ -952,7 +976,7 @@ Pinned revision used by this check: {SGL_TEST_FILES_CI_DATA_REVISION}
"lora_nickname": ["default", "lora2"],
"lora_path": [first_lora_path, second_lora_path],
"target": "all",
"strength": [1.0, 1.0],
"strength": [0.5, 0.5],
},
timeout=_CONTROL_API_TIMEOUT_SECS,
)
@@ -971,7 +995,7 @@ Pinned revision used by this check: {SGL_TEST_FILES_CI_DATA_REVISION}
"lora_nickname": ["default", "lora2"],
"lora_path": [first_lora_path, second_lora_path],
"target": "all",
"strength": [0.8, 0.5],
"strength": [0.6, 0.35],
},
timeout=_CONTROL_API_TIMEOUT_SECS,
)
@@ -995,7 +1019,7 @@ Pinned revision used by this check: {SGL_TEST_FILES_CI_DATA_REVISION}
"lora_nickname": ["default", "lora2"],
"lora_path": [first_lora_path, second_lora_path],
"target": ["transformer", "transformer_2"],
"strength": [0.8, 0.5],
"strength": [0.6, 0.35],
},
timeout=_CONTROL_API_TIMEOUT_SECS,
)
@@ -1016,10 +1040,11 @@ Pinned revision used by this check: {SGL_TEST_FILES_CI_DATA_REVISION}
assert (
resp.status_code == 200
), f"set_lora back to single adapter failed: {resp.text}"
rid, _ = self._run_generation_with_server_watchdog(
rid, content = self._run_generation_with_server_watchdog(
ctx, case.id, generate_fn, client
)
assert rid is not None
self._validate_lora_consistency(case, content, "multi-LoRA default adapter")
logger.info("[Multi-LoRA] All multi-LoRA tests passed for %s", case.id)
@@ -33,7 +33,7 @@ if TYPE_CHECKING:
logger = init_logger(__name__)
SGL_TEST_FILES_CI_DATA_REVISION = "b7455318873fc5af399c8447b3bb0d9471a5084c"
SGL_TEST_FILES_CI_DATA_REVISION = "68a42bf2229361497e6c4ff22ce89a157bf2f531"
SGL_TEST_FILES_CONSISTENCY_GT_ROOT = (
"https://raw.githubusercontent.com/"
f"sgl-project/ci-data/{SGL_TEST_FILES_CI_DATA_REVISION}/"
@@ -58,6 +58,7 @@ SGL_TEST_FILES_OFFICIAL_CONSISTENCY_GT_CASES = frozenset(
"ltx_2_3_two_stage_ti2v_2gpus",
}
)
CONSISTENCY_THRESHOLD_JSON_PATH = (
Path(__file__).resolve().parent / "server" / "consistency_threshold.json"
)
@@ -27,6 +27,7 @@ def _make_pipeline(layer: BaseLayerWithLoRA) -> _TestLoRAPipeline:
pipeline.lora_initialized = True
pipeline.lora_adapters = defaultdict(dict)
pipeline.loaded_adapter_paths = {"adapter": "/adapter"}
pipeline.loaded_adapter_alphas = {"adapter": None}
pipeline.cur_adapter_name = {}
pipeline.cur_adapter_path = {}
pipeline.cur_adapter_strength = {}