Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
Show all changes
132 commits
Select commit Hold shift + click to select a range
7f66f1c
[diffusion] feat: warmup-calibrated auto residency promotion
mickqian Aug 16, 2026
0c0f117
[diffusion] fix auto residency video estimation with two-point calibr…
mickqian Aug 16, 2026
2df7fee
[diffusion] narrate auto residency adjustment with equivalent server …
mickqian Aug 17, 2026
414b3b1
[diffusion] unify post-request residency hint with auto promotion logic
mickqian Aug 18, 2026
f856eb6
[diffusion] harden auto residency against review findings
mickqian Aug 18, 2026
17b51bf
[diffusion] extract control-request protocol into control_requests.py
mickqian Aug 19, 2026
7bd0b5c
[diffusion] migrate control requests from dataclass to msgspec.Struct
mickqian Aug 19, 2026
bfe142b
fix: preserve single-GPU warmup frame contract
mickqian Aug 18, 2026
e9db98a
fix: skip auto residency with Ulysses parallelism
mickqian Aug 18, 2026
274e9f6
[diffusion] fix: make auto residency reachable on a small card
mickqian Aug 20, 2026
a7ac4fa
fix host gather after residency promotion
mickqian Aug 20, 2026
f5a10fe
test(diffusion): preserve golden CI baselines
mickqian Aug 20, 2026
4a47e46
test(diffusion): preserve 1-GPU golden CI baselines
mickqian Aug 20, 2026
dfb2cfe
diffusion-ci: test auto residency by default
mickqian Aug 20, 2026
45eb7af
diffusion-ci: refresh LTX2 auto residency baseline
mickqian Aug 20, 2026
769597f
[diffusion] jointly plan serving memory residency
mickqian Aug 26, 2026
d4316d3
[diffusion] fix auto residency unit fixtures
mickqian Aug 26, 2026
d83fc87
[diffusion] calibrate residency at serving shape
mickqian Aug 26, 2026
481e404
[diffusion] calibrate prequantized residency
mickqian Aug 26, 2026
87ed38a
Merge remote-tracking branch 'origin/main' into codex/auto-residency-…
mickqian Aug 26, 2026
ae3dc54
[diffusion] harden residency workload calibration
mickqian Aug 26, 2026
de89ced
[diffusion] gate unsupported quantized offload
mickqian Aug 26, 2026
2f22951
[diffusion] keep runtime LoRA weights inference-only
mickqian Aug 26, 2026
992f4c1
[diffusion] avoid CUDA sync in timestep logs
mickqian Aug 26, 2026
c044a0b
[diffusion] preserve undeclared quantized offload capability
mickqian Aug 26, 2026
c497b0f
[diffusion] plan residency from live VRAM
mickqian Aug 26, 2026
ebff2f4
[diffusion] converge auto residency after warmup
mickqian Aug 26, 2026
9c8626e
[diffusion] keep LTX LoRA adapters versioned
mickqian Aug 26, 2026
9596c1c
Merge remote-tracking branch 'origin/main' into codex/auto-residency-…
mickqian Aug 26, 2026
9c4875a
[diffusion] make auto residency benefit-aware
mickqian Aug 26, 2026
1261612
[diffusion] let auto residency rebalance layerwise memory
mickqian Aug 26, 2026
2191868
[diffusion] preserve timing utility across ranks
mickqian Aug 26, 2026
ba9612d
[diffusion] solve complete residency target states
mickqian Aug 26, 2026
a268535
[diffusion] enumerate layerwise residency ratios
mickqian Aug 26, 2026
03b8922
[diffusion] document calibrated residency scope
mickqian Aug 26, 2026
6d03a14
[diffusion] constrain residency placement transitions
mickqian Aug 26, 2026
f8afe14
[diffusion] scale exact placement optimization
mickqian Aug 26, 2026
2a41251
[diffusion] add lazy layerwise placement targets
mickqian Aug 26, 2026
618f2ce
[diffusion] guard auto residency calibration regressions
mickqian Aug 27, 2026
0644252
[diffusion] restore layerwise pins after rollback
mickqian Aug 27, 2026
f73d112
[diffusion] constrain auto residency lifecycle phases
mickqian Aug 27, 2026
cfa1396
[diffusion] refresh auto residency CI baselines
mickqian Aug 27, 2026
95a4ee9
[diffusion] account for layerwise LoRA transition memory
mickqian Aug 27, 2026
7a4db98
[diffusion] record coarse LoRA weight transitions
mickqian Aug 27, 2026
c4ce470
[diffusion] account for unmaterialized LoRA transitions
mickqian Aug 27, 2026
ade6ab6
[diffusion] distinguish measured residency transitions
mickqian Aug 27, 2026
28fb5d3
Merge branch 'main' into mick/diffusion-auto-residency
mickqian Aug 27, 2026
137af03
[diffusion] score residency relative to calibrated placement
mickqian Aug 27, 2026
0aed0ab
[diffusion] calibrate residency with the CI workload
mickqian Aug 28, 2026
21743e2
fix(diffusion): stabilize residency calibration
mickqian Aug 27, 2026
3321551
[diffusion] bound auto residency validation
mickqian Aug 28, 2026
22edea8
[diffusion] commit validated resident placement
mickqian Aug 28, 2026
fee9e40
[diffusion] validate residency commit barrier
mickqian Aug 28, 2026
5f86c0a
[docs] define auto residency solve boundary
mickqian Aug 28, 2026
98f2e2d
[diffusion] validate residency without replanning
mickqian Aug 28, 2026
888f715
[diffusion] remove obsolete residency retry path
mickqian Aug 28, 2026
0a12e0f
[diffusion] auto-tune compatible custom residency
mickqian Aug 28, 2026
f843683
[diffusion] optimize layerwise host pin frontier
mickqian Aug 28, 2026
b8b2943
[diffusion] calibrate layerwise placement per group
mickqian Aug 28, 2026
d3878ad
[diffusion] scope fixed loading paths per component
mickqian Aug 28, 2026
c43a81f
[docs] explain calibrated layer-group placement
mickqian Aug 28, 2026
5eae402
[diffusion] infer repeated layer usage from warmup stages
mickqian Aug 28, 2026
6b005b7
[diffusion] scale all measured denoising stages
mickqian Aug 28, 2026
055b9b7
[diffusion] calibrate residency by stage workload
mickqian Aug 28, 2026
1eb326a
[diffusion] align 5090 Wan consistency guard
mickqian Aug 28, 2026
e923b79
[diffusion] align residency calibration stage names
mickqian Aug 28, 2026
d6e0a9e
[diffusion] use explicit residency warmup contracts
mickqian Aug 28, 2026
ca91558
[diffusion] shorten resident placement validation
mickqian Aug 28, 2026
713a16a
[docs] align residency validation policy
mickqian Aug 28, 2026
46d9b82
test(diffusion): remove duplicate residency case
mickqian Aug 28, 2026
ed21267
[diffusion] bound nonbinding residency optimization
mickqian Aug 28, 2026
07fece5
[diffusion] fix parent layerwise fixture
mickqian Aug 28, 2026
bf8c6b1
[diffusion] profile residency planning phases
mickqian Aug 28, 2026
4ea2dec
[diffusion] bound residency frontier construction
mickqian Aug 28, 2026
45a5210
[diffusion] preserve LTX2 original DiT residency
mickqian Aug 28, 2026
0834013
[diffusion] validate realized residency headroom
mickqian Aug 28, 2026
8785f08
Merge branch 'mick/diffusion-auto-residency' of https://github.com/sg…
mickqian Aug 28, 2026
052f27f
test(diffusion): align residency test fixtures
mickqian Aug 28, 2026
0df915c
test(diffusion): update LTX auto residency VRAM baseline
mickqian Aug 28, 2026
024138b
fix(diffusion): gate CI workload warmup by residency support
mickqian Aug 28, 2026
4549570
fix(diffusion): skip calibration hooks for unsupported pipelines
mickqian Aug 28, 2026
d0207f3
fix(diffusion): skip residency tracking for unsupported pipelines
mickqian Aug 28, 2026
c2f533e
fix(diffusion): keep SANA refiner encoder resident on H100
mickqian Aug 28, 2026
63f2f5f
Merge remote-tracking branch 'origin/main' into HEAD
mickqian Aug 28, 2026
d407145
fix(diffusion): keep SANA VAE resident on H100
mickqian Aug 29, 2026
137286c
test(diffusion): pin compatible SANA consistency data
mickqian Aug 29, 2026
f315c72
Merge origin/main into mick/diffusion-auto-residency
mickqian Aug 29, 2026
6f1adb6
Merge main into mick/diffusion-auto-residency
mickqian Aug 30, 2026
1af94a4
Fix transformer fallback test fixture
mickqian Aug 30, 2026
daf1e45
fix(diffusion): contain auto residency regressions
mickqian Aug 31, 2026
d54ea11
fix(diffusion): refresh 2-gpu residency baselines
mickqian Aug 31, 2026
715cc8f
Merge remote-tracking branch 'origin/main' into HEAD
mickqian Aug 31, 2026
20c12cb
test(diffusion): include component precision in selector fixtures
mickqian Aug 31, 2026
992c36b
test(diffusion): refresh auto residency VRAM baselines
mickqian Aug 31, 2026
29083a3
fix(diffusion): contain auto residency regressions
mickqian Aug 31, 2026
a16fc47
fix(diffusion): preserve legacy LTX-2 placement
mickqian Aug 31, 2026
a1b1e5d
Merge remote-tracking branch 'origin/main' into HEAD
mickqian Sep 1, 2026
6e2c9dd
Merge main into mick/diffusion-auto-residency
mickqian Sep 3, 2026
56596e0
Fix Ruff formatting after main merge
mickqian Sep 3, 2026
1d64090
Merge remote-tracking branch 'origin/main' into HEAD
mickqian Sep 3, 2026
b4d06ff
Merge remote-tracking branch 'origin/main' into mick/diffusion-auto-r…
mickqian Sep 3, 2026
60b537c
Merge remote-tracking branch 'origin/main' into t-merge-main
Sep 4, 2026
e1569d8
Merge feat/residency-apply-validate: layers 1-3 of this PR now live i…
Sep 4, 2026
f28b990
Merge branch 'feat/residency-apply-validate' into root-l4
Sep 4, 2026
a6838ae
Merge branch 'feat/residency-apply-validate' into root-l4
Sep 4, 2026
b9551a3
Merge remote-tracking branch 'origin/feat/residency-apply-validate' i…
mickqian Sep 4, 2026
e8c3f41
Merge remote-tracking branch 'origin/feat/residency-apply-validate' i…
mickqian Sep 4, 2026
6abbbc8
Merge branch 'feat/residency-apply-validate' into root-l4
Sep 4, 2026
38246fb
Merge remote-tracking branch 'origin/mick/diffusion-auto-residency' i…
Sep 5, 2026
e33833c
Merge branch 'feat/residency-apply-validate' into root-l4
Sep 5, 2026
f99038a
Merge branch 'feat/residency-apply-validate' into root-l4
Sep 5, 2026
250159b
Merge branch 'feat/residency-apply-validate' into root-l4
Sep 5, 2026
58004c6
Merge branch 'feat/residency-apply-validate' into root-l4
Sep 5, 2026
0762538
Merge branch 'feat/residency-apply-validate' into root-l4
Sep 5, 2026
e1885fd
Merge branch 'feat/residency-apply-validate' into root-l4
Sep 5, 2026
e323f82
Merge branch 'feat/residency-apply-validate' into root-l4
Sep 5, 2026
4df15d0
Merge branch 'feat/residency-apply-validate' into root-l4
Sep 5, 2026
3638835
Merge branch 'feat/residency-apply-validate' into root-l4
Sep 6, 2026
8bc3ffa
Merge branch 'feat/residency-apply-validate' into root-l4
Sep 6, 2026
4e92bad
Merge branch 'feat/residency-apply-validate' into root-l4
Sep 6, 2026
5f6286b
Merge direct base feat/residency-apply-validate
mickqian Sep 6, 2026
5e2bf3a
Merge branch 'feat/residency-apply-validate' into root-l4
Sep 6, 2026
40d8a50
Merge branch 'feat/residency-apply-validate' into root-l4
Sep 6, 2026
5406301
Merge branch 'feat/residency-apply-validate' into root-l4
Sep 7, 2026
a91ea72
Merge branch 'feat/residency-apply-validate' into root-l4
Sep 7, 2026
361041a
Merge branch 'feat/residency-apply-validate' into root-l4
Sep 8, 2026
64c518d
Merge branch 'feat/residency-apply-validate' into root-l4
Sep 8, 2026
b970ec4
Merge remote-tracking branch 'origin/feat/residency-apply-validate' i…
Sep 8, 2026
a68b2be
Merge remote-tracking branch 'origin/feat/residency-apply-validate' i…
Sep 8, 2026
e2152e9
Merge branch 'feat/residency-apply-validate' into mick/diffusion-auto…
Sep 11, 2026
5c9b3d0
Let a pipeline without a residency manager enter the LoRA offload con…
Sep 11, 2026
6f47376
Merge branch 'feat/residency-apply-validate' into mick/diffusion-auto…
Sep 12, 2026
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
Original file line number Diff line number Diff line change
Expand Up @@ -14,9 +14,9 @@ class ModelDeploymentConfig:
dit_layerwise_offload_modes: tuple[Literal["auto", "memory"], ...] = ()
auto_dit_offload_prefetch_size: float | None = None
keep_resident_min_available_gb: float | None = None
# Per-model resident defaults. Auto mode additionally keeps an image DiT
# resident above the image workload memory threshold; video DiT placement
# stays with the model's FSDP/layerwise policy.
# Fallback residency for offline/no-warmup deployments. Eligible server
# warmup paths start load-safe and replace these coarse thresholds with a
# measured per-phase placement plan before reporting ready.
keep_resident_components: tuple[OffloadComponentName, ...] = ("vae",)
fsdp_auto_min_available_memory_gb: float | None = None
fsdp_auto_requires_cfg: bool = True
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -165,11 +165,15 @@ def __post_init__(self):
def get_model_deployment_config(self) -> ModelDeploymentConfig:
return ModelDeploymentConfig(
dit_layerwise_offload_modes=("memory",),
# The two-stage path is not numerically invariant when its text
# encoders or VAE use layerwise offload. Keep them resident when
# there is enough headroom, while preserving the low-memory
# layerwise path on consumer GPUs.
keep_resident_min_available_gb=70,
keep_resident_components=("text_encoder", "vae"),
# Conservative auto-FSDP gate for the 720p world-model path. Users
# can still force FSDP explicitly on smaller cards.
fsdp_auto_min_available_memory_gb=60,
keep_resident_min_available_gb=120,
keep_resident_components=("dit", "vae"),
)

# --- Latent shape ---
Expand Down
10 changes: 6 additions & 4 deletions python/sglang/multimodal_gen/runtime/layers/lora/linear.py
Original file line number Diff line number Diff line change
Expand Up @@ -231,10 +231,12 @@ def set_lora_weights(
If False, append to existing list (for multi-LoRA support).
output_offset: Optional constant output term paired with this adapter
"""
lora_A_param = torch.nn.Parameter(
A
) # share storage with weights in the pipeline
lora_B_param = torch.nn.Parameter(B)
# This runtime never trains adapters. In particular, weights loaded
# during a warmup may be inference tensors; leaving Parameter's
# requires_grad default enabled makes a later dynamic-LoRA request try
# to read their unavailable version counters while sharding views.
lora_A_param = torch.nn.Parameter(A, requires_grad=False)
lora_B_param = torch.nn.Parameter(B, requires_grad=False)
output_offset_param = (
torch.nn.Parameter(output_offset, requires_grad=False)
if output_offset is not None
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -22,6 +22,7 @@
ComponentResidencyStrategy,
ComponentUse,
ResidencyState,
build_component_residency_strategy,
)
from sglang.multimodal_gen.runtime.models.schedulers.scheduling_flow_match_euler_discrete import (
FlowMatchEulerDiscreteScheduler,
Expand Down Expand Up @@ -456,6 +457,12 @@ class LTX2TwoStageResidencyStrategy(ComponentResidencyStrategy):

def __init__(self, manager: "LTX2TwoStageResidencyController") -> None:
self.manager = manager
self._placement_strategies: dict[
tuple[str, int, str], ComponentResidencyStrategy
] = {}

def supports_auto_residency(self) -> bool:
return True

@property
def pipeline(self) -> "LTX2TwoStagePipeline":
Expand All @@ -473,6 +480,19 @@ def _phase(self, use: ComponentUse) -> str:
def initialize(self) -> None:
pass

def _placement_strategy(
self, module: torch.nn.Module, use: ComponentUse
) -> ComponentResidencyStrategy:
mode = self.server_args.residency_mode(use.component_name)
key = (use.component_name, id(module), mode)
strategy = self._placement_strategies.get(key)
if strategy is None:
strategy = build_component_residency_strategy(
use.component_name, module, self.server_args
)
self._placement_strategies[key] = strategy
return strategy

def prepare_for_use(
self,
module: torch.nn.Module,
Expand All @@ -482,13 +502,28 @@ def prepare_for_use(
phase = self._phase(use)
if phase != self.manager._active_phase:
self.enter_phase(phase)
self._placement_strategy(module, use).prepare_for_use(module, use, state)

def prefetch_for_use(
self,
module: torch.nn.Module,
use: ComponentUse,
state: ResidencyState,
) -> bool:
phase = self._phase(use)
if phase != self.manager._active_phase:
self.enter_phase(phase)
return self._placement_strategy(module, use).prefetch_for_use(
module, use, state
)

def wait_for_use(
self,
module: torch.nn.Module,
use: ComponentUse,
state: ResidencyState,
) -> None:
self._placement_strategy(module, use).wait_for_use(module, use, state)
self.ensure_phase_ready(self._phase(use))

def finish_use(
Expand All @@ -497,6 +532,7 @@ def finish_use(
use: ComponentUse,
state: ResidencyState,
) -> None:
self._placement_strategy(module, use).finish_use(module, use, state)
self.exit_phase(self._phase(use))

def finish_request(
Expand All @@ -507,8 +543,11 @@ def finish_request(
*,
preferred: bool,
) -> None:
self._placement_strategy(module, use).finish_request(
module, use, state, preferred=preferred
)
if not preferred:
self.finish_use(module, use, state)
self.exit_phase(self._phase(use))
return
phase = self._phase(use)
if phase != self.manager._active_phase:
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -293,26 +293,48 @@ def _temporarily_disable_offload(
yield []
return

# clear device cache to free up unused memory
if torch.get_device_module().is_available():
torch.get_device_module().synchronize()
torch.get_device_module().empty_cache()
offloaded_module_names = [
module_name
for module_name in module_names
if (
(module := self.modules.get(module_name)) is not None
and is_layerwise_offloaded_module(module)
)
]
# Record every target, including coarse component offload. The current
# coarse path can already materialize the full component here, and the
# planner must preserve that phase if it later chooses layerwise mode.
residency_manager = getattr(self, "component_residency_manager", None)
residency_transition = (
residency_manager.full_weight_transition(module_names)
if residency_manager is not None
else nullcontext()
)

offload_disabled_modules = []
for module_name in module_names:
module = self.modules.get(module_name)
if isinstance(module, torch.nn.Module):
restore_weight_snapshot(module)
if module is not None and is_layerwise_offloaded_module(module):
with residency_transition:
# clear device cache to free up unused memory
if torch.get_device_module().is_available():
torch.get_device_module().synchronize()
torch.get_device_module().empty_cache()

# snapshot-offloaded components merge into their restored weights
for module_name in module_names:
module = self.modules.get(module_name)
if isinstance(module, torch.nn.Module):
restore_weight_snapshot(module)

offload_disabled_modules = []
for module_name in offloaded_module_names:
module = self.modules[module_name]
module.disable_offload()
offload_disabled_modules.append(module)

try:
yield offload_disabled_modules
finally:
# Re-enable layerwise offload: sync weights to CPU and restore hooks
for module in offload_disabled_modules:
module.enable_offload()
try:
yield offload_disabled_modules
finally:
# Re-enable layerwise offload: sync weights to CPU and restore hooks
for module in offload_disabled_modules:
module.enable_offload()

def _needs_lora_weight_update_context(
self,
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -246,7 +246,7 @@ def _encode_delight_prompts(
return torch.cat([positive, negative, negative])

@torch.no_grad()
def _run_delight(self, image: Image.Image) -> Image.Image:
def _run_delight(self, image: Image.Image, batch: Req) -> Image.Image:
assert self.delight_transformer is not None
assert self.delight_vae is not None
assert self.delight_scheduler is not None
Expand Down Expand Up @@ -278,7 +278,14 @@ def _run_delight(self, image: Image.Image) -> Image.Image:
)

scheduler = self.delight_scheduler
scheduler.set_timesteps(self.config.delight_num_inference_steps, device=device)
target_steps = self.config.delight_num_inference_steps
measured_steps = (
min(target_steps, max(1, int(batch.num_inference_steps)))
if batch.is_warmup
else target_steps
)
batch.record_stage_iterations(measured_steps, target_steps)
scheduler.set_timesteps(measured_steps, device=device)
generator = torch.Generator(device="cpu").manual_seed(42)
latent_channels = self.delight_transformer.config.out_channels
latents = randn_tensor(
Expand Down Expand Up @@ -350,11 +357,11 @@ def _run_delight(self, image: Image.Image) -> Image.Image:
(composited.clamp(0, 1).cpu().numpy() * 255).astype(np.uint8)
)

def _prepare_reference_image(self, image_path: str) -> Image.Image:
def _prepare_reference_image(self, image_path: str, batch: Req) -> Image.Image:
image = self._load_input_image(image_path)
if not self.config.delight_enable:
return image
return self._run_delight(image)
return self._run_delight(image, batch)

def _render_multiview(
self, mesh: Any
Expand Down Expand Up @@ -404,7 +411,9 @@ def forward(self, batch: Req, server_args: ServerArgs) -> Req:

with concurrent.futures.ThreadPoolExecutor(max_workers=2) as executor:
mesh_future = executor.submit(self._unwrap_mesh, mesh)
image_future = executor.submit(self._prepare_reference_image, image_path)
image_future = executor.submit(
self._prepare_reference_image, image_path, batch
)
paint_mesh = mesh_future.result()
delighted_image = image_future.result()

Expand Down Expand Up @@ -542,15 +551,27 @@ def _camera_index(azimuth: float, elevation: float) -> int:
base, divisor = 12, 1
return base + base_index // divisor

def _timesteps(self, device: torch.device) -> torch.Tensor:
def _timesteps(
self, device: torch.device, batch: Req | None = None
) -> torch.Tensor:
target_steps = (
10
if self.config.paint_turbo_mode
else (self.config.paint_num_inference_steps)
)
measured_steps = (
min(target_steps, max(1, int(batch.num_inference_steps)))
if batch is not None and batch.is_warmup
else target_steps
)
if batch is not None:
batch.record_stage_iterations(measured_steps, target_steps)
if not self.config.paint_turbo_mode:
self.scheduler.set_timesteps(
self.config.paint_num_inference_steps, device=device
)
self.scheduler.set_timesteps(measured_steps, device=device)
return self.scheduler.timesteps

self.scheduler.set_timesteps(
num_inference_steps=10,
num_inference_steps=measured_steps,
original_inference_steps=30,
device=device,
)
Expand Down Expand Up @@ -639,7 +660,7 @@ def _prepare_denoising_inputs(self, batch: Req) -> PaintDenoisingInputs:
if position_attention_mask is not None:
model_kwargs["position_attn_mask"] = position_attention_mask

timesteps = self._timesteps(device)
timesteps = self._timesteps(device, batch)
latent_channels = self.transformer.config.in_channels
latent_size = render_size // self.vae_scale_factor
latents = randn_tensor(
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -50,7 +50,10 @@ def forward(self, batch: Req, server_args: ServerArgs) -> Req:
if self.pipeline.should_skip_ltx2_lora_switch_stage():
batch.extra["ltx2_phase"] = self.phase
return batch
self.pipeline.switch_lora_phase(self.phase, batch=batch)
# The installed adapter survives this stage. Give its tensors version
# counters for later offload stages, which run outside inference mode.
with torch.inference_mode(False), torch.no_grad():
self.pipeline.switch_lora_phase(self.phase, batch=batch)
batch.extra["ltx2_phase"] = self.phase
return batch

Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -5,6 +5,12 @@
"ssim_threshold": 0.86,
"psnr_threshold": 19.5,
"mean_abs_diff_threshold": 8.5
},
"wan2_1_t2v_1.3b": {
"clip_threshold": 0.95,
"ssim_threshold": 0.85,
"psnr_threshold": 25.0,
"mean_abs_diff_threshold": 8.0
}
}
}
Original file line number Diff line number Diff line change
Expand Up @@ -2,6 +2,7 @@

import unittest
from types import SimpleNamespace
from unittest.mock import Mock

import torch
from diffusers import AutoencoderKL as DiffusersAutoencoderKL
Expand Down Expand Up @@ -222,6 +223,24 @@ def test_uses_standard_lcm_schedule_without_custom_timesteps(self):
)
self.assertFalse(stage.scheduler.custom_timesteps)

def test_warmup_caps_paint_steps_and_records_full_target(self):
stage = Hunyuan3DPaintTexGenStage.__new__(Hunyuan3DPaintTexGenStage)
stage.config = Hunyuan3D2PipelineConfig(paint_turbo_mode=True)
stage.scheduler = LCMScheduler(
num_train_timesteps=1000,
original_inference_steps=50,
)
batch = SimpleNamespace(
is_warmup=True,
num_inference_steps=4,
record_stage_iterations=Mock(),
)

timesteps = stage._timesteps(torch.device("cpu"), batch)

self.assertEqual(len(timesteps), 4)
batch.record_stage_iterations.assert_called_once_with(4, 10)


if __name__ == "__main__":
unittest.main()
Loading
Loading