diff --git a/modules/video_models/video_minimax.py b/modules/video_models/video_minimax.py index 247112496..cd0d17e2a 100644 --- a/modules/video_models/video_minimax.py +++ b/modules/video_models/video_minimax.py @@ -102,14 +102,18 @@ def calculate_audio_shift(steps: int, value: float = 0.15, max_shift: float = 6. def set_sampler_shift(pipe, steps: int, video_shift: float = 12.0, audio_shift: float = 3.0): - import diffusers - if getattr(pipe, 'scheduler', None) is None or getattr(pipe.scheduler, 'config', None) is None: - log.warning(f'Pipeline: cls={pipe.__class__.__name__} scheduler is missing') + scheduler = getattr(pipe, 'scheduler', None) + audio_scheduler = getattr(pipe, 'audio_scheduler', None) + if not hasattr(scheduler, 'set_shift') or not hasattr(audio_scheduler, 'set_shift'): + log.warning(f'Pipeline: cls={pipe.__class__.__name__} scheduler={scheduler.__class__.__name__} audio={audio_scheduler.__class__.__name__} shift unsupported') return video_calc_shift = calculate_video_shift(steps=steps, value=video_shift) audio_calc_shift = calculate_audio_shift(steps=steps, value=audio_shift) - dct_video = { 'shift': video_shift, 'steps': video_calc_shift } - dct_audio = { 'shift': audio_shift, 'steps': audio_calc_shift } - pipe.scheduler = diffusers.MiniMaxH3Scheduler(shift=video_calc_shift) - pipe.audio_scheduler.config.shift = diffusers.MiniMaxH3Scheduler(shift=audio_calc_shift) - log.debug(f'Pipeline: scheduler={pipe.scheduler.__class__.__name__} video={dct_video} audio={dct_audio} ') + dct_video = { 'value': video_shift, 'shift': video_calc_shift } + dct_audio = { 'value': audio_shift, 'shift': audio_calc_shift } + # set_shift keeps the shipped value in config.shift; the default sampler restores scheduler from default_scheduler every generation, so that copy carries the shift too + for target in (scheduler, getattr(pipe, 'default_scheduler', None)): + if hasattr(target, 'set_shift'): + target.set_shift(video_calc_shift) + audio_scheduler.set_shift(audio_calc_shift) + log.debug(f'Pipeline: scheduler={scheduler.__class__.__name__} video={dct_video} audio={dct_audio}')