[diffusion] chore: adjust layer wise-offload strategy (#25930)
This commit is contained in:
@@ -95,7 +95,6 @@ class WanT2V480PConfig(PipelineConfig):
|
||||
def get_model_deployment_config(self) -> ModelDeploymentConfig:
|
||||
return ModelDeploymentConfig(
|
||||
auto_dit_layerwise_offload=True,
|
||||
auto_dit_layerwise_offload_high_memory_disable_gb=130,
|
||||
)
|
||||
|
||||
|
||||
@@ -146,7 +145,6 @@ class WanI2V480PConfig(WanT2V480PConfig, WanI2VCommonConfig):
|
||||
def get_model_deployment_config(self) -> ModelDeploymentConfig:
|
||||
return ModelDeploymentConfig(
|
||||
auto_dit_layerwise_offload=True,
|
||||
auto_dit_layerwise_offload_high_memory_disable_gb=130,
|
||||
)
|
||||
|
||||
|
||||
|
||||
@@ -201,7 +201,7 @@ class ServerArgsAutoTuner:
|
||||
self._enable_cfg_parallel_if_supported()
|
||||
|
||||
def maybe_adjust_auto_default_layerwise_offload(self) -> None:
|
||||
"""Enable verified non-DiT layerwise defaults for unset component placement."""
|
||||
"""Enable verified layerwise defaults for unset component placement."""
|
||||
args = self.server_args
|
||||
if args.performance_mode != "auto":
|
||||
return
|
||||
@@ -220,7 +220,7 @@ class ServerArgsAutoTuner:
|
||||
return
|
||||
|
||||
logger.info(
|
||||
"Automatically enable default non-DiT layerwise offload for %s: %s",
|
||||
"Automatically enable default layerwise offload for %s: %s",
|
||||
args.pipeline_config.__class__.__name__,
|
||||
layerwise_components,
|
||||
)
|
||||
@@ -367,17 +367,70 @@ class ServerArgsAutoTuner:
|
||||
or args.dit_layerwise_offload is True
|
||||
):
|
||||
# The legacy --dit-layerwise-offload flag is a DiT-only selector.
|
||||
# Do not merge implicit non-DiT defaults into that explicit mode.
|
||||
# Do not merge implicit defaults into that explicit mode.
|
||||
return []
|
||||
|
||||
# `*_cpu_offload` is the component placement knob. If a user explicitly
|
||||
# set it to either true or false, keep that component out of default
|
||||
# layerwise selection.
|
||||
return [
|
||||
components = [
|
||||
component_name
|
||||
for component_name, arg_name in DEFAULT_LAYERWISE_COMPONENT_ARG_NAMES
|
||||
if not args.is_arg_explicitly_set(arg_name)
|
||||
]
|
||||
if self._should_auto_enable_dit_layerwise_offload():
|
||||
components.insert(0, LAYERWISE_OFFLOAD_DIT_GROUP)
|
||||
self._set_default_wan_dit_offload_prefetch_size()
|
||||
return components
|
||||
|
||||
def _should_auto_enable_dit_layerwise_offload(self) -> bool:
|
||||
args = self.server_args
|
||||
|
||||
# only for wan for now
|
||||
if not self._is_wan_pipeline_config():
|
||||
return False
|
||||
if not self._deployment_config().auto_dit_layerwise_offload:
|
||||
return False
|
||||
|
||||
if (
|
||||
args.pipeline_config.dmd_denoising_steps is not None
|
||||
or not current_platform.enable_dit_layerwise_offload_for_wan_by_default()
|
||||
or envs.SGLANG_CACHE_DIT_ENABLED
|
||||
or args.use_fsdp_inference
|
||||
or args.is_arg_explicitly_set("dit_cpu_offload")
|
||||
):
|
||||
return False
|
||||
|
||||
# memory mode is memory-first: keep the broad Wan DiT layerwise policy
|
||||
# unless a guard above says it conflicts with another placement path
|
||||
if args.performance_mode == "memory":
|
||||
return True
|
||||
|
||||
# auto mode is performance-first: profiling only showed clear wins for
|
||||
# Wan2.2 A14B, where coarse DiT CPU offload creates large step spikes
|
||||
return (
|
||||
args.performance_mode == "auto" and self._is_wan2_2_a14b_pipeline_config()
|
||||
)
|
||||
|
||||
def _is_wan2_2_a14b_pipeline_config(self) -> bool:
|
||||
config_name = self.server_args.pipeline_config.__class__.__name__
|
||||
return config_name.startswith("Wan2_2_") and "A14B" in config_name
|
||||
|
||||
def _set_default_wan_dit_offload_prefetch_size(self) -> None:
|
||||
args = self.server_args
|
||||
if (
|
||||
args.performance_mode == "auto"
|
||||
and self._is_wan2_2_a14b_pipeline_config()
|
||||
and not args.is_arg_explicitly_set("dit_offload_prefetch_size")
|
||||
):
|
||||
# p2 was the fastest stable default in the Wan2.2 A14B sweep
|
||||
args.dit_offload_prefetch_size = 2
|
||||
|
||||
def _is_wan_pipeline_config(self) -> bool:
|
||||
return any(
|
||||
cls.__module__.endswith(".wan")
|
||||
for cls in self.server_args.pipeline_config.__class__.mro()
|
||||
)
|
||||
|
||||
def _auto_uses_dit_offload(self) -> bool:
|
||||
args = self.server_args
|
||||
|
||||
@@ -8,6 +8,8 @@ Each collected request prints a performance log before validation.
|
||||
from __future__ import annotations
|
||||
|
||||
import os
|
||||
import queue
|
||||
import threading
|
||||
import time
|
||||
from pathlib import Path
|
||||
from typing import Any, Callable
|
||||
@@ -56,6 +58,21 @@ logger = init_logger(__name__)
|
||||
# Track test cases missing estimated_full_test_time_s for time measurement output
|
||||
_MISSING_ESTIMATED_TIME_CASES: set[str] = set()
|
||||
_PENDING_BASELINE_DUMPS: dict[str, tuple["PerformanceSummary", bool]] = {}
|
||||
_OPENAI_REQUEST_TIMEOUT_SECS = float(
|
||||
os.environ.get("SGLANG_TEST_OPENAI_REQUEST_TIMEOUT_SECS", "600")
|
||||
)
|
||||
_SERVER_EXIT_POLL_INTERVAL_SECS = float(
|
||||
os.environ.get("SGLANG_TEST_SERVER_EXIT_POLL_INTERVAL_SECS", "1")
|
||||
)
|
||||
_CONTROL_API_TIMEOUT_SECS = float(
|
||||
os.environ.get("SGLANG_TEST_CONTROL_API_TIMEOUT_SECS", "300")
|
||||
)
|
||||
_SERVER_FATAL_LOG_PATTERNS = (
|
||||
"terminate called after throwing an instance of",
|
||||
"Fatal Python error:",
|
||||
"Segmentation fault",
|
||||
"Aborted (core dumped)",
|
||||
)
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
@@ -259,8 +276,80 @@ class DiffusionServerBase:
|
||||
return OpenAI(
|
||||
api_key="sglang-anything",
|
||||
base_url=f"http://localhost:{ctx.port}/v1",
|
||||
timeout=_OPENAI_REQUEST_TIMEOUT_SECS,
|
||||
max_retries=0,
|
||||
)
|
||||
|
||||
def _fail_if_server_stopped_or_crashed(
|
||||
self, ctx: ServerContext, case_id: str
|
||||
) -> None:
|
||||
returncode = ctx.process.poll()
|
||||
if returncode is None:
|
||||
tail = ctx.log_tail()
|
||||
for pattern in _SERVER_FATAL_LOG_PATTERNS:
|
||||
if pattern in tail:
|
||||
pytest.fail(
|
||||
f"{case_id}: server reported a fatal backend error during "
|
||||
f"generation: {pattern}\n\nServer log tail:\n{tail}",
|
||||
pytrace=False,
|
||||
)
|
||||
return
|
||||
|
||||
tail = ctx.log_tail()
|
||||
message = (
|
||||
f"{case_id}: server process exited during generation "
|
||||
f"(code {returncode})."
|
||||
)
|
||||
if tail:
|
||||
message += f"\n\nServer log tail:\n{tail}"
|
||||
pytest.fail(message, pytrace=False)
|
||||
|
||||
def _run_generation_with_server_watchdog(
|
||||
self,
|
||||
ctx: ServerContext,
|
||||
case_id: str,
|
||||
generate_fn: Callable[[str, openai.Client], tuple[str, bytes]],
|
||||
client: openai.Client,
|
||||
) -> tuple[str, bytes]:
|
||||
result_queue: queue.Queue[tuple[str, tuple[str, bytes] | BaseException]] = (
|
||||
queue.Queue(maxsize=1)
|
||||
)
|
||||
|
||||
def _target() -> None:
|
||||
try:
|
||||
result_queue.put(("ok", generate_fn(case_id, client)))
|
||||
except BaseException as exc:
|
||||
result_queue.put(("error", exc))
|
||||
|
||||
# native backend crashes can leave the HTTP client blocked until its read
|
||||
# timeout; keep the request in a daemon thread so the main test thread can
|
||||
# fail as soon as the server subprocess exits
|
||||
thread = threading.Thread(
|
||||
target=_target,
|
||||
name=f"diffusion-generation-{case_id}",
|
||||
daemon=True,
|
||||
)
|
||||
thread.start()
|
||||
|
||||
while True:
|
||||
try:
|
||||
state, payload = result_queue.get(
|
||||
timeout=_SERVER_EXIT_POLL_INTERVAL_SECS
|
||||
)
|
||||
except queue.Empty:
|
||||
self._fail_if_server_stopped_or_crashed(ctx, case_id)
|
||||
continue
|
||||
|
||||
if state == "ok":
|
||||
if isinstance(payload, BaseException):
|
||||
raise payload
|
||||
return payload
|
||||
|
||||
self._fail_if_server_stopped_or_crashed(ctx, case_id)
|
||||
if not isinstance(payload, BaseException):
|
||||
pytest.fail(f"{case_id}: invalid generation result state: {state}")
|
||||
raise payload
|
||||
|
||||
def run_and_collect(
|
||||
self,
|
||||
ctx: ServerContext,
|
||||
@@ -274,7 +363,9 @@ class DiffusionServerBase:
|
||||
Tuple of (performance_record, content_bytes)
|
||||
"""
|
||||
client = self._client(ctx)
|
||||
rid, content = generate_fn(case_id, client)
|
||||
rid, content = self._run_generation_with_server_watchdog(
|
||||
ctx, case_id, generate_fn, client
|
||||
)
|
||||
|
||||
if not collect_perf:
|
||||
return None, content
|
||||
@@ -680,41 +771,55 @@ Pinned revision used by this check: {SGL_TEST_FILES_CI_DATA_REVISION}
|
||||
This test verifies that each API call succeeds AND that generation works after each operation.
|
||||
"""
|
||||
base_url = f"http://localhost:{ctx.port}/v1"
|
||||
client = OpenAI(base_url=base_url, api_key="dummy")
|
||||
client = self._client(ctx)
|
||||
|
||||
# Test 1: unmerge_lora_weights - API should succeed and generation should work
|
||||
logger.info("[LoRA E2E] Testing unmerge_lora_weights for %s", case.id)
|
||||
resp = requests.post(f"{base_url}/unmerge_lora_weights")
|
||||
resp = requests.post(
|
||||
f"{base_url}/unmerge_lora_weights", timeout=_CONTROL_API_TIMEOUT_SECS
|
||||
)
|
||||
assert resp.status_code == 200, f"unmerge_lora_weights failed: {resp.text}"
|
||||
|
||||
logger.info("[LoRA E2E] Verifying generation after unmerge for %s", case.id)
|
||||
rid_after_unmerge, _ = generate_fn(case.id, client)
|
||||
rid_after_unmerge, _ = self._run_generation_with_server_watchdog(
|
||||
ctx, case.id, generate_fn, client
|
||||
)
|
||||
assert rid_after_unmerge is not None, "Generation after unmerge failed"
|
||||
logger.info("[LoRA E2E] Generation after unmerge succeeded")
|
||||
|
||||
# Test 2: merge_lora_weights - API should succeed and generation should work
|
||||
logger.info("[LoRA E2E] Testing merge_lora_weights for %s", case.id)
|
||||
resp = requests.post(f"{base_url}/merge_lora_weights")
|
||||
resp = requests.post(
|
||||
f"{base_url}/merge_lora_weights", timeout=_CONTROL_API_TIMEOUT_SECS
|
||||
)
|
||||
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, _ = generate_fn(case.id, client)
|
||||
rid_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"
|
||||
logger.info("[LoRA E2E] Generation after merge succeeded")
|
||||
|
||||
# Test 3: set_lora (re-set the same adapter) - API should succeed and generation should work
|
||||
logger.info("[LoRA E2E] Testing set_lora for %s", case.id)
|
||||
resp = requests.post(f"{base_url}/set_lora", json={"lora_nickname": "default"})
|
||||
resp = requests.post(
|
||||
f"{base_url}/set_lora",
|
||||
json={"lora_nickname": "default"},
|
||||
timeout=_CONTROL_API_TIMEOUT_SECS,
|
||||
)
|
||||
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, _ = generate_fn(case.id, client)
|
||||
rid_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"
|
||||
logger.info("[LoRA E2E] Generation after set_lora succeeded")
|
||||
|
||||
# Test 4: list_loras - API should return the expected list of LoRA adapters
|
||||
logger.info("[LoRA E2E] Testing list_loras for %s", case.id)
|
||||
resp = requests.get(f"{base_url}/list_loras")
|
||||
resp = requests.get(f"{base_url}/list_loras", timeout=_CONTROL_API_TIMEOUT_SECS)
|
||||
assert resp.status_code == 200, f"list_loras failed: {resp.text}"
|
||||
lora_info = resp.json()
|
||||
logger.info("[LoRA E2E] list_loras returned %s", lora_info)
|
||||
@@ -742,13 +847,15 @@ Pinned revision used by this check: {SGL_TEST_FILES_CI_DATA_REVISION}
|
||||
and generation succeeds after each switch.
|
||||
"""
|
||||
base_url = f"http://localhost:{ctx.port}/v1"
|
||||
client = OpenAI(base_url=base_url, api_key="dummy")
|
||||
client = self._client(ctx)
|
||||
|
||||
# Test 1: Generate with initial LoRA
|
||||
logger.info(
|
||||
"[LoRA Switch E2E] Testing generation with initial LoRA for %s", case.id
|
||||
)
|
||||
rid_initial, _ = generate_fn(case.id, client)
|
||||
rid_initial, _ = self._run_generation_with_server_watchdog(
|
||||
ctx, case.id, generate_fn, client
|
||||
)
|
||||
assert rid_initial is not None, "Generation with initial LoRA failed"
|
||||
logger.info("[LoRA Switch E2E] Generation with initial LoRA succeeded")
|
||||
|
||||
@@ -759,6 +866,7 @@ Pinned revision used by this check: {SGL_TEST_FILES_CI_DATA_REVISION}
|
||||
resp = requests.post(
|
||||
f"{base_url}/set_lora",
|
||||
json={"lora_nickname": "lora2", "lora_path": second_lora_path},
|
||||
timeout=_CONTROL_API_TIMEOUT_SECS,
|
||||
)
|
||||
assert (
|
||||
resp.status_code == 200
|
||||
@@ -767,20 +875,28 @@ Pinned revision used by this check: {SGL_TEST_FILES_CI_DATA_REVISION}
|
||||
logger.info(
|
||||
"[LoRA Switch E2E] Verifying generation with second LoRA for %s", case.id
|
||||
)
|
||||
rid_second, _ = generate_fn(case.id, client)
|
||||
rid_second, _ = self._run_generation_with_server_watchdog(
|
||||
ctx, case.id, generate_fn, client
|
||||
)
|
||||
assert rid_second is not None, "Generation with second LoRA failed"
|
||||
logger.info("[LoRA Switch E2E] Generation with second LoRA succeeded")
|
||||
|
||||
# Test 3: Switch back to original LoRA and generate
|
||||
logger.info("[LoRA Switch E2E] Switching back to original LoRA for %s", case.id)
|
||||
resp = requests.post(f"{base_url}/set_lora", json={"lora_nickname": "default"})
|
||||
resp = requests.post(
|
||||
f"{base_url}/set_lora",
|
||||
json={"lora_nickname": "default"},
|
||||
timeout=_CONTROL_API_TIMEOUT_SECS,
|
||||
)
|
||||
assert resp.status_code == 200, f"set_lora back to default failed: {resp.text}"
|
||||
|
||||
logger.info(
|
||||
"[LoRA Switch E2E] Verifying generation after switching back for %s",
|
||||
case.id,
|
||||
)
|
||||
rid_switched_back, _ = generate_fn(case.id, client)
|
||||
rid_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"
|
||||
logger.info("[LoRA Switch E2E] Generation after switching back succeeded")
|
||||
|
||||
@@ -812,6 +928,7 @@ Pinned revision used by this check: {SGL_TEST_FILES_CI_DATA_REVISION}
|
||||
resp = requests.post(
|
||||
f"{base_url}/set_lora",
|
||||
json={"lora_nickname": "default", "lora_path": dynamic_lora_path},
|
||||
timeout=_CONTROL_API_TIMEOUT_SECS,
|
||||
)
|
||||
assert resp.status_code == 200, f"Dynamic set_lora failed: {resp.text}"
|
||||
logger.info("[Dynamic LoRA] set_lora succeeded for %s", case.id)
|
||||
@@ -829,7 +946,7 @@ Pinned revision used by this check: {SGL_TEST_FILES_CI_DATA_REVISION}
|
||||
Tests: basic multi-LoRA, different strengths, cached adapters, switch back to single.
|
||||
"""
|
||||
base_url = f"http://localhost:{ctx.port}/v1"
|
||||
client = OpenAI(base_url=base_url, api_key="dummy")
|
||||
client = self._client(ctx)
|
||||
|
||||
# Test 1: Basic multi-LoRA with list format
|
||||
resp = requests.post(
|
||||
@@ -840,11 +957,14 @@ Pinned revision used by this check: {SGL_TEST_FILES_CI_DATA_REVISION}
|
||||
"target": "all",
|
||||
"strength": [1.0, 1.0],
|
||||
},
|
||||
timeout=_CONTROL_API_TIMEOUT_SECS,
|
||||
)
|
||||
assert (
|
||||
resp.status_code == 200
|
||||
), f"set_lora with multiple adapters failed: {resp.text}"
|
||||
rid, _ = generate_fn(case.id, client)
|
||||
rid, _ = self._run_generation_with_server_watchdog(
|
||||
ctx, case.id, generate_fn, client
|
||||
)
|
||||
assert rid is not None
|
||||
|
||||
# Test 2: Different strengths
|
||||
@@ -856,15 +976,22 @@ Pinned revision used by this check: {SGL_TEST_FILES_CI_DATA_REVISION}
|
||||
"target": "all",
|
||||
"strength": [0.8, 0.5],
|
||||
},
|
||||
timeout=_CONTROL_API_TIMEOUT_SECS,
|
||||
)
|
||||
assert (
|
||||
resp.status_code == 200
|
||||
), f"set_lora with different strengths failed: {resp.text}"
|
||||
rid, _ = generate_fn(case.id, client)
|
||||
rid, _ = self._run_generation_with_server_watchdog(
|
||||
ctx, case.id, generate_fn, client
|
||||
)
|
||||
assert rid is not None
|
||||
|
||||
# Test 3: Different targets
|
||||
requests.post(f"{base_url}/set_lora", json={"lora_nickname": "default"})
|
||||
requests.post(
|
||||
f"{base_url}/set_lora",
|
||||
json={"lora_nickname": "default"},
|
||||
timeout=_CONTROL_API_TIMEOUT_SECS,
|
||||
)
|
||||
resp = requests.post(
|
||||
f"{base_url}/set_lora",
|
||||
json={
|
||||
@@ -873,19 +1000,28 @@ Pinned revision used by this check: {SGL_TEST_FILES_CI_DATA_REVISION}
|
||||
"target": ["transformer", "transformer_2"],
|
||||
"strength": [0.8, 0.5],
|
||||
},
|
||||
timeout=_CONTROL_API_TIMEOUT_SECS,
|
||||
)
|
||||
assert (
|
||||
resp.status_code == 200
|
||||
), f"set_lora with cached adapters failed: {resp.text}"
|
||||
rid, _ = generate_fn(case.id, client)
|
||||
rid, _ = self._run_generation_with_server_watchdog(
|
||||
ctx, case.id, generate_fn, client
|
||||
)
|
||||
assert rid is not None
|
||||
|
||||
# Test 4: Switch back to single LoRA
|
||||
resp = requests.post(f"{base_url}/set_lora", json={"lora_nickname": "default"})
|
||||
resp = requests.post(
|
||||
f"{base_url}/set_lora",
|
||||
json={"lora_nickname": "default"},
|
||||
timeout=_CONTROL_API_TIMEOUT_SECS,
|
||||
)
|
||||
assert (
|
||||
resp.status_code == 200
|
||||
), f"set_lora back to single adapter failed: {resp.text}"
|
||||
rid, _ = generate_fn(case.id, client)
|
||||
rid, _ = self._run_generation_with_server_watchdog(
|
||||
ctx, case.id, generate_fn, client
|
||||
)
|
||||
assert rid is not None
|
||||
|
||||
logger.info("[Multi-LoRA] All multi-LoRA tests passed for %s", case.id)
|
||||
@@ -901,7 +1037,7 @@ Pinned revision used by this check: {SGL_TEST_FILES_CI_DATA_REVISION}
|
||||
|
||||
# Test GET /v1/models
|
||||
logger.info("[Models API] Testing GET /v1/models for %s", case.id)
|
||||
resp = requests.get(f"{base_url}/v1/models")
|
||||
resp = requests.get(f"{base_url}/v1/models", timeout=_CONTROL_API_TIMEOUT_SECS)
|
||||
assert resp.status_code == 200, f"/v1/models failed: {resp.text}"
|
||||
|
||||
data = resp.json()
|
||||
@@ -948,7 +1084,9 @@ Pinned revision used by this check: {SGL_TEST_FILES_CI_DATA_REVISION}
|
||||
# Test GET /v1/models/{model_path}
|
||||
model_path = model["id"]
|
||||
logger.info("[Models API] Testing GET /v1/models/%s", model_path)
|
||||
resp = requests.get(f"{base_url}/v1/models/{model_path}")
|
||||
resp = requests.get(
|
||||
f"{base_url}/v1/models/{model_path}", timeout=_CONTROL_API_TIMEOUT_SECS
|
||||
)
|
||||
assert resp.status_code == 200, f"/v1/models/{model_path} failed: {resp.text}"
|
||||
|
||||
single_model = resp.json()
|
||||
@@ -968,7 +1106,10 @@ Pinned revision used by this check: {SGL_TEST_FILES_CI_DATA_REVISION}
|
||||
|
||||
# Test GET /v1/models/{non_existent_model} returns 404
|
||||
logger.info("[Models API] Testing GET /v1/models/non_existent_model")
|
||||
resp = requests.get(f"{base_url}/v1/models/non_existent_model")
|
||||
resp = requests.get(
|
||||
f"{base_url}/v1/models/non_existent_model",
|
||||
timeout=_CONTROL_API_TIMEOUT_SECS,
|
||||
)
|
||||
assert resp.status_code == 404, f"Expected 404, got {resp.status_code}"
|
||||
error_data = resp.json()
|
||||
assert "error" in error_data, "404 response missing 'error' field"
|
||||
@@ -986,7 +1127,7 @@ Pinned revision used by this check: {SGL_TEST_FILES_CI_DATA_REVISION}
|
||||
return
|
||||
|
||||
base_url = f"http://localhost:{ctx.port}"
|
||||
resp = requests.get(f"{base_url}/v1/models")
|
||||
resp = requests.get(f"{base_url}/v1/models", timeout=_CONTROL_API_TIMEOUT_SECS)
|
||||
assert resp.status_code == 200, f"/v1/models failed: {resp.text}"
|
||||
data = resp.json().get("data", [])
|
||||
if not data:
|
||||
@@ -1001,7 +1142,11 @@ Pinned revision used by this check: {SGL_TEST_FILES_CI_DATA_REVISION}
|
||||
if case.sampling_params.output_size:
|
||||
payload["size"] = case.sampling_params.output_size
|
||||
|
||||
resp = requests.post(f"{base_url}/v1/videos", json=payload)
|
||||
resp = requests.post(
|
||||
f"{base_url}/v1/videos",
|
||||
json=payload,
|
||||
timeout=_CONTROL_API_TIMEOUT_SECS,
|
||||
)
|
||||
assert (
|
||||
resp.status_code == 400
|
||||
), f"Expected 400 for T2V input_reference, got {resp.status_code}: {resp.text}"
|
||||
|
||||
@@ -160,6 +160,14 @@ class ServerContext:
|
||||
_stdout_fh: Any = field(repr=False)
|
||||
_log_thread: threading.Thread | None = field(default=None, repr=False)
|
||||
|
||||
def log_tail(self, lines: int = 200) -> str:
|
||||
"""Return recent server output for failure diagnostics."""
|
||||
try:
|
||||
content = self.stdout_file.read_text(encoding="utf-8", errors="ignore")
|
||||
return "\n".join(content.splitlines()[-lines:])
|
||||
except Exception:
|
||||
return ""
|
||||
|
||||
def cleanup(self) -> None:
|
||||
"""Clean up server resources."""
|
||||
try:
|
||||
|
||||
@@ -17,7 +17,16 @@ from sglang.multimodal_gen.configs.pipeline_configs.mova import MOVAPipelineConf
|
||||
from sglang.multimodal_gen.configs.pipeline_configs.qwen_image import (
|
||||
QwenImagePipelineConfig,
|
||||
)
|
||||
from sglang.multimodal_gen.configs.pipeline_configs.wan import WanT2V480PConfig
|
||||
from sglang.multimodal_gen.configs.pipeline_configs.wan import (
|
||||
FastWan2_2_TI2V_5B_Config,
|
||||
TurboWanT2V480PConfig,
|
||||
Wan2_2_I2V_A14B_Config,
|
||||
Wan2_2_T2V_A14B_Config,
|
||||
WanI2V480PConfig,
|
||||
WanI2V720PConfig,
|
||||
WanT2V480PConfig,
|
||||
WanT2V720PConfig,
|
||||
)
|
||||
from sglang.multimodal_gen.configs.pipeline_configs.zimage import ZImagePipelineConfig
|
||||
from sglang.multimodal_gen.registry import _get_config_info
|
||||
from sglang.multimodal_gen.runtime.models.dits.qwen_image import (
|
||||
@@ -510,6 +519,54 @@ class TestOffloadDefaults(unittest.TestCase):
|
||||
["text_encoder", "image_encoder", "vae"],
|
||||
)
|
||||
|
||||
def test_auto_wan2_2_a14b_layerwise_offload_adds_dit(self):
|
||||
for pipeline_config, model_path in (
|
||||
(Wan2_2_T2V_A14B_Config(), "Wan-AI/Wan2.2-T2V-A14B-Diffusers"),
|
||||
(Wan2_2_I2V_A14B_Config(), "Wan-AI/Wan2.2-I2V-A14B-Diffusers"),
|
||||
):
|
||||
with self.subTest(pipeline_config=pipeline_config.__class__.__name__):
|
||||
args = self._from_dict_with_pipeline_config(
|
||||
pipeline_config,
|
||||
kwargs={
|
||||
"model_path": model_path,
|
||||
"performance_mode": "auto",
|
||||
},
|
||||
)
|
||||
|
||||
self.assertTrue(args.layerwise_offload_components)
|
||||
self.assertFalse(args.use_fsdp_inference)
|
||||
self.assertFalse(args.dit_cpu_offload)
|
||||
self.assertFalse(args.text_encoder_cpu_offload)
|
||||
self.assertFalse(args.image_encoder_cpu_offload)
|
||||
self.assertEqual(args.dit_offload_prefetch_size, 2)
|
||||
self.assertEqual(
|
||||
args.layerwise_offload_components,
|
||||
["dit", "text_encoder", "image_encoder", "vae"],
|
||||
)
|
||||
|
||||
def test_auto_wan2_1_14b_layerwise_offload_uses_non_dit_default(self):
|
||||
for pipeline_config, model_path in (
|
||||
(WanT2V720PConfig(), "Wan-AI/Wan2.1-T2V-14B-Diffusers"),
|
||||
(WanI2V480PConfig(), "Wan-AI/Wan2.1-I2V-14B-480P-Diffusers"),
|
||||
(WanI2V720PConfig(), "Wan-AI/Wan2.1-I2V-14B-720P-Diffusers"),
|
||||
):
|
||||
with self.subTest(pipeline_config=pipeline_config.__class__.__name__):
|
||||
args = self._from_dict_with_pipeline_config(
|
||||
pipeline_config,
|
||||
kwargs={
|
||||
"model_path": model_path,
|
||||
"performance_mode": "auto",
|
||||
},
|
||||
)
|
||||
|
||||
self.assertTrue(args.layerwise_offload_components)
|
||||
self.assertTrue(args.dit_cpu_offload)
|
||||
self.assertEqual(args.dit_offload_prefetch_size, 0.0)
|
||||
self.assertEqual(
|
||||
args.layerwise_offload_components,
|
||||
["text_encoder", "image_encoder", "vae"],
|
||||
)
|
||||
|
||||
def test_memory_wan_layerwise_offload_is_enabled_without_fsdp(self):
|
||||
args = self._from_dict_with_pipeline_config(
|
||||
WanT2V480PConfig(),
|
||||
@@ -518,12 +575,12 @@ class TestOffloadDefaults(unittest.TestCase):
|
||||
|
||||
self.assertTrue(args.layerwise_offload_components)
|
||||
self.assertFalse(args.use_fsdp_inference)
|
||||
self.assertTrue(args.dit_cpu_offload)
|
||||
self.assertFalse(args.dit_cpu_offload)
|
||||
self.assertFalse(args.text_encoder_cpu_offload)
|
||||
self.assertFalse(args.image_encoder_cpu_offload)
|
||||
self.assertEqual(
|
||||
args.layerwise_offload_components,
|
||||
["text_encoder", "image_encoder", "vae"],
|
||||
["dit", "text_encoder", "image_encoder", "vae"],
|
||||
)
|
||||
|
||||
def test_auto_wan_layerwise_offload_does_not_disable_explicit_fsdp(self):
|
||||
@@ -543,6 +600,79 @@ class TestOffloadDefaults(unittest.TestCase):
|
||||
)
|
||||
self.assertTrue(args.use_fsdp_inference)
|
||||
|
||||
def test_auto_wan_layerwise_offload_preserves_explicit_dit_cpu_offload(self):
|
||||
args = self._from_dict_with_pipeline_config(
|
||||
WanT2V480PConfig(),
|
||||
kwargs={
|
||||
"model_path": "Wan-AI/Wan2.1-T2V-1.3B-Diffusers",
|
||||
"performance_mode": "auto",
|
||||
"dit_cpu_offload": True,
|
||||
},
|
||||
)
|
||||
|
||||
self.assertTrue(args.dit_cpu_offload)
|
||||
self.assertEqual(
|
||||
args.layerwise_offload_components,
|
||||
["text_encoder", "image_encoder", "vae"],
|
||||
)
|
||||
|
||||
def test_auto_mova_layerwise_offload_does_not_implicitly_add_dit(self):
|
||||
args = self._from_dict_with_pipeline_config(
|
||||
MOVAPipelineConfig(),
|
||||
kwargs={
|
||||
"model_path": "OpenMOSS-Team/MOVA-360p",
|
||||
"performance_mode": "auto",
|
||||
},
|
||||
)
|
||||
|
||||
self.assertTrue(args.dit_cpu_offload)
|
||||
self.assertEqual(
|
||||
args.layerwise_offload_components,
|
||||
["text_encoder", "image_encoder", "vae"],
|
||||
)
|
||||
|
||||
def test_auto_fastwan_layerwise_offload_does_not_implicitly_add_dit(self):
|
||||
args = self._from_dict_with_pipeline_config(
|
||||
FastWan2_2_TI2V_5B_Config(),
|
||||
kwargs={
|
||||
"model_path": "FastVideo/FastWan2.2-TI2V-5B-FullAttn-Diffusers",
|
||||
"performance_mode": "auto",
|
||||
},
|
||||
)
|
||||
|
||||
self.assertTrue(args.dit_cpu_offload)
|
||||
self.assertEqual(
|
||||
args.layerwise_offload_components,
|
||||
["text_encoder", "image_encoder", "vae"],
|
||||
)
|
||||
|
||||
def test_auto_turbo_wan_layerwise_offload_does_not_implicitly_add_dit(self):
|
||||
args = self._from_dict_with_pipeline_config(
|
||||
TurboWanT2V480PConfig(),
|
||||
kwargs={
|
||||
"model_path": "IPostYellow/TurboWan2.1-T2V-1.3B-Diffusers",
|
||||
"performance_mode": "auto",
|
||||
},
|
||||
)
|
||||
|
||||
self.assertTrue(args.dit_cpu_offload)
|
||||
self.assertEqual(
|
||||
args.layerwise_offload_components,
|
||||
["text_encoder", "image_encoder", "vae"],
|
||||
)
|
||||
|
||||
def test_explicit_fastwan_dit_layerwise_still_selects_dit_group(self):
|
||||
args = self._from_dict_with_pipeline_config(
|
||||
FastWan2_2_TI2V_5B_Config(),
|
||||
kwargs={
|
||||
"model_path": "FastVideo/FastWan2.2-TI2V-5B-FullAttn-Diffusers",
|
||||
"dit_layerwise_offload": True,
|
||||
},
|
||||
)
|
||||
|
||||
self.assertFalse(args.dit_cpu_offload)
|
||||
self.assertEqual(args.layerwise_offload_components, ["dit"])
|
||||
|
||||
def test_auto_multi_gpu_wan_uses_layerwise_offload_without_cfg(self):
|
||||
with patch.object(ServerArgs, "_model_default_uses_cfg", return_value=False):
|
||||
args = self._from_dict_with_pipeline_config(
|
||||
@@ -825,12 +955,12 @@ class TestOffloadDefaults(unittest.TestCase):
|
||||
|
||||
self.assertFalse(args.use_fsdp_inference)
|
||||
self.assertTrue(args.layerwise_offload_components)
|
||||
self.assertTrue(args.dit_cpu_offload)
|
||||
self.assertFalse(args.dit_cpu_offload)
|
||||
self.assertFalse(args.text_encoder_cpu_offload)
|
||||
self.assertFalse(args.image_encoder_cpu_offload)
|
||||
self.assertEqual(
|
||||
args.layerwise_offload_components,
|
||||
["text_encoder", "image_encoder", "vae"],
|
||||
["dit", "text_encoder", "image_encoder", "vae"],
|
||||
)
|
||||
|
||||
def test_memory_mode_preserves_explicit_fsdp(self):
|
||||
|
||||
Reference in New Issue
Block a user