refactor scheduler import
This commit is contained in:
@@ -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
@@ -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
|
||||
|
||||
|
||||
@@ -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
|
||||
@@ -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'
|
||||
]
|
||||
Reference in New Issue
Block a user