refactor scheduler import

This commit is contained in:
kijai
2025-07-17 15:56:55 +03:00
parent 2263d02a0a
commit 6bc53b771d
9 changed files with 15 additions and 18 deletions
+7 -6
View File
@@ -10,13 +10,14 @@ from tqdm import tqdm
from .wanvideo.modules.clip import CLIPModel
from .wanvideo.modules.model import rope_params
from .wanvideo.modules.t5 import T5EncoderModel
from .wanvideo.utils.fm_solvers import (FlowDPMSolverMultistepScheduler,
get_sampling_sigmas, retrieve_timesteps)
from .wanvideo.utils.fm_solvers_unipc import FlowUniPCMultistepScheduler
from .wanvideo.utils.basic_flowmatch import FlowMatchScheduler
from .wanvideo.utils.flowmatch_pusa import FlowMatchSchedulerPusa
from .wanvideo.schedulers import (
FlowDPMSolverMultistepScheduler, FlowUniPCMultistepScheduler,
FlowMatchScheduler, FlowMatchSchedulerPusa, FlowMatchLCMScheduler,
get_sampling_sigmas, retrieve_timesteps
)
from diffusers.schedulers import FlowMatchEulerDiscreteScheduler, DEISMultistepScheduler
from .wanvideo.utils.scheduling_flow_match_lcm import FlowMatchLCMScheduler
from .multitalk.multitalk import timestep_transform, add_noise
+2 -2
View File
@@ -6,9 +6,9 @@ import math
from tqdm import tqdm
from ..wanvideo.modules.model import rope_params
from ..wanvideo.utils.fm_solvers_unipc import FlowUniPCMultistepScheduler
from ..wanvideo.schedulers.fm_solvers_unipc import FlowUniPCMultistepScheduler
from diffusers.schedulers import FlowMatchEulerDiscreteScheduler
from ..wanvideo.utils.scheduling_flow_match_lcm import FlowMatchLCMScheduler
from ..wanvideo.schedulers.scheduling_flow_match_lcm import FlowMatchLCMScheduler
from ..nodes import optimized_scale
from einops import rearrange
+5
View File
@@ -0,0 +1,5 @@
from .fm_solvers import (FlowDPMSolverMultistepScheduler, get_sampling_sigmas, retrieve_timesteps)
from .fm_solvers_unipc import FlowUniPCMultistepScheduler
from .basic_flowmatch import FlowMatchScheduler
from .flowmatch_pusa import FlowMatchSchedulerPusa
from .scheduling_flow_match_lcm import FlowMatchLCMScheduler
@@ -38,7 +38,6 @@ class FlowMatchSchedulerPusa():
def step(self, model_output, timestep, sample, to_final=False, **kwargs):
print("timestep in scheduler", timestep)
if isinstance(timestep, torch.Tensor):
# timestep = timestep.cpu()
self.timesteps = self.timesteps.to(timestep.device)
@@ -75,7 +74,7 @@ class FlowMatchSchedulerPusa():
if torch.any(timestep == 0):
zero_indices = torch.where(timestep == 0)[1].to(torch.long)
sigma[:,:,zero_indices] = 0
print("sigma", sigma[0,0,:,0,0], '\n', "sigma_", sigma_[0,0,:,0,0])
#print("sigma", sigma[0,0,:,0,0], '\n', "sigma_", sigma_[0,0,:,0,0])
prev_sample = sample + model_output * (sigma_ - sigma)
return prev_sample
-8
View File
@@ -1,8 +0,0 @@
from .fm_solvers import (FlowDPMSolverMultistepScheduler, get_sampling_sigmas,
retrieve_timesteps)
from .fm_solvers_unipc import FlowUniPCMultistepScheduler
__all__ = [
'HuggingfaceTokenizer', 'get_sampling_sigmas', 'retrieve_timesteps',
'FlowDPMSolverMultistepScheduler', 'FlowUniPCMultistepScheduler'
]