diff --git a/modules/ltx/ltx_process.py b/modules/ltx/ltx_process.py index 18aab0d44..330efcc9f 100644 --- a/modules/ltx/ltx_process.py +++ b/modules/ltx/ltx_process.py @@ -6,7 +6,7 @@ from PIL import Image from modules import shared, errors, timer, memstats, progress, processing, sd_models, sd_samplers, devices, extra_networks, call_queue from modules.logger import log from modules.ltx import ltx_capabilities -from modules.ltx.ltx_util import get_bucket, get_frames, load_model, load_upsample, load_upsample_2x, get_conditions, get_generator, get_prompts, vae_decode +from modules.ltx.ltx_util import get_bucket, get_frames, load_model, load_upsample, load_upsample_2x, get_conditions, get_generator, get_prompts, temp_scheduler_opts, vae_decode from modules.processing_callbacks import diffusers_callback from modules.video_models.video_vae import set_vae_params from modules.video_models.video_save import save_video @@ -290,30 +290,10 @@ def run_ltx(task_id, framewise = caps.family == '0.9' set_vae_params(p, framewise=framewise) - # Snapshot scheduler + shared.opts before mutation so the try/finally restores on every exit - # path (abort, interrupt, Stage 2 scheduler swap). Without this, run-specific sampler settings - # leak into shared.opts.data and across runs/tabs, and the default_scheduler snapshot from - # video_load.py:171 gets clobbered by a deepcopy of the mutated scheduler on every run. - orig_dynamic_shift = shared.opts.schedulers_dynamic_shift - orig_sampler_shift = shared.opts.schedulers_shift - orig_scheduler = shared.sd_model.scheduler - orig_default_scheduler = getattr(shared.sd_model, 'default_scheduler', None) - orig_use_dynamic_shifting = getattr(orig_scheduler.config, 'use_dynamic_shifting', None) if hasattr(orig_scheduler, 'config') else None - orig_flow_shift = getattr(orig_scheduler.config, 'flow_shift', None) if hasattr(orig_scheduler, 'config') else None - - try: - shared.opts.data['schedulers_dynamic_shift'] = dynamic_shift - shared.opts.data['schedulers_shift'] = sampler_shift - if hasattr(shared.sd_model, 'scheduler') and hasattr(shared.sd_model.scheduler, 'config') and hasattr(shared.sd_model.scheduler, 'register_to_config'): - if hasattr(shared.sd_model.scheduler.config, 'use_dynamic_shifting'): - shared.sd_model.scheduler.config.use_dynamic_shifting = dynamic_shift - shared.sd_model.scheduler.register_to_config(use_dynamic_shifting=dynamic_shift) - if hasattr(shared.sd_model.scheduler.config, 'flow_shift') and sampler_shift is not None and sampler_shift >= 0: - shared.sd_model.scheduler.config.flow_shift = sampler_shift - shared.sd_model.scheduler.register_to_config(flow_shift=sampler_shift) - # Do NOT re-snapshot default_scheduler; that overwrites video_load.py:171's load-time - # snapshot with the run-mutated config, so reset_scheduler then carries the last run's choice. - + # Scheduler + shared.opts mutation is wrapped in temp_scheduler_opts so restore runs on + # every exit path (normal return, abort, interrupt, Stage 2 scheduler swap). See the + # helper's docstring for the five pieces of state it snapshots. + with temp_scheduler_opts(shared.sd_model, dynamic_shift=dynamic_shift, sampler_shift=sampler_shift): if selected is not None: video_overrides.set_overrides(p, selected) @@ -606,18 +586,3 @@ def run_ltx(task_id, log.info(f'Processed: fn="{video_file}" frames={num_frames} fps={fps} its={its} resolution={resolution} time={t_end-t0:.2f} timers={timer.process.dct()} memory={memstats.memory_stats()}') yield video_file, f'LTX: Generation completed | File {video_file} | Frames {num_frames} | Resolution {resolution} | f/s {fps} | it/s {its} ' + f"

{summary} {memory}

" - finally: - shared.opts.data['schedulers_dynamic_shift'] = orig_dynamic_shift - shared.opts.data['schedulers_shift'] = orig_sampler_shift - if shared.sd_model.scheduler is not orig_scheduler: - shared.sd_model.scheduler = orig_scheduler - if orig_default_scheduler is not None and shared.sd_model.default_scheduler is not orig_default_scheduler: - shared.sd_model.default_scheduler = orig_default_scheduler - if hasattr(shared.sd_model.scheduler, 'config') and hasattr(shared.sd_model.scheduler, 'register_to_config'): - if orig_use_dynamic_shifting is not None and hasattr(shared.sd_model.scheduler.config, 'use_dynamic_shifting'): - shared.sd_model.scheduler.config.use_dynamic_shifting = orig_use_dynamic_shifting - shared.sd_model.scheduler.register_to_config(use_dynamic_shifting=orig_use_dynamic_shifting) - if orig_flow_shift is not None and hasattr(shared.sd_model.scheduler.config, 'flow_shift'): - shared.sd_model.scheduler.config.flow_shift = orig_flow_shift - shared.sd_model.scheduler.register_to_config(flow_shift=orig_flow_shift) - log.debug(f'LTX: scheduler/opts restored dynamic_shift={orig_dynamic_shift} sampler_shift={orig_sampler_shift}') diff --git a/modules/ltx/ltx_util.py b/modules/ltx/ltx_util.py index b404879f7..099ad94bd 100644 --- a/modules/ltx/ltx_util.py +++ b/modules/ltx/ltx_util.py @@ -1,4 +1,5 @@ import time +from contextlib import contextmanager import torch from PIL import Image from modules import devices, shared, sd_models, timer, extra_networks @@ -80,6 +81,51 @@ def load_upsample_2x(upsample_pipe, upsample_repo_id): return upsample_pipe +@contextmanager +def temp_scheduler_opts(sd_model, *, dynamic_shift=None, sampler_shift=None): + # Run-scoped override of shared.opts scheduler settings and scheduler.config. Snapshots + # five pieces of state (shared.opts dynamic_shift + shift, scheduler object, default_scheduler + # snapshot, and scheduler.config use_dynamic_shifting + flow_shift) and restores every one on + # exit. Keeps run-specific sampler settings out of config.json and prevents default_scheduler + # from getting clobbered by a deepcopy of the mutated scheduler at video_load.py:171. The + # scheduler-object restore matters for Stage 2 refine, which swaps the scheduler entirely. + orig_dynamic_shift = shared.opts.schedulers_dynamic_shift + orig_sampler_shift = shared.opts.schedulers_shift + orig_scheduler = sd_model.scheduler + orig_default_scheduler = getattr(sd_model, 'default_scheduler', None) + orig_use_dynamic_shifting = getattr(orig_scheduler.config, 'use_dynamic_shifting', None) if hasattr(orig_scheduler, 'config') else None + orig_flow_shift = getattr(orig_scheduler.config, 'flow_shift', None) if hasattr(orig_scheduler, 'config') else None + + try: + if dynamic_shift is not None: + shared.opts.data['schedulers_dynamic_shift'] = dynamic_shift + if sampler_shift is not None: + shared.opts.data['schedulers_shift'] = sampler_shift + if hasattr(sd_model, 'scheduler') and hasattr(sd_model.scheduler, 'config') and hasattr(sd_model.scheduler, 'register_to_config'): + if dynamic_shift is not None and hasattr(sd_model.scheduler.config, 'use_dynamic_shifting'): + sd_model.scheduler.config.use_dynamic_shifting = dynamic_shift + sd_model.scheduler.register_to_config(use_dynamic_shifting=dynamic_shift) + if sampler_shift is not None and sampler_shift >= 0 and hasattr(sd_model.scheduler.config, 'flow_shift'): + sd_model.scheduler.config.flow_shift = sampler_shift + sd_model.scheduler.register_to_config(flow_shift=sampler_shift) + yield + finally: + shared.opts.data['schedulers_dynamic_shift'] = orig_dynamic_shift + shared.opts.data['schedulers_shift'] = orig_sampler_shift + if sd_model.scheduler is not orig_scheduler: + sd_model.scheduler = orig_scheduler + if orig_default_scheduler is not None and sd_model.default_scheduler is not orig_default_scheduler: + sd_model.default_scheduler = orig_default_scheduler + if hasattr(sd_model.scheduler, 'config') and hasattr(sd_model.scheduler, 'register_to_config'): + if orig_use_dynamic_shifting is not None and hasattr(sd_model.scheduler.config, 'use_dynamic_shifting'): + sd_model.scheduler.config.use_dynamic_shifting = orig_use_dynamic_shifting + sd_model.scheduler.register_to_config(use_dynamic_shifting=orig_use_dynamic_shifting) + if orig_flow_shift is not None and hasattr(sd_model.scheduler.config, 'flow_shift'): + sd_model.scheduler.config.flow_shift = orig_flow_shift + sd_model.scheduler.register_to_config(flow_shift=orig_flow_shift) + log.debug(f'LTX: scheduler/opts restored dynamic_shift={orig_dynamic_shift} sampler_shift={orig_sampler_shift}') + + def _condition_cls(family: str): if family == '2.x': try: