From ced1ddaa1af69e58feec4a36a6f0cc7344d1a6ee Mon Sep 17 00:00:00 2001 From: kijai <40791699+kijai@users.noreply.github.com> Date: Thu, 12 Jun 2025 16:29:53 +0300 Subject: [PATCH] Update basic_flowmatch.py --- wanvideo/utils/basic_flowmatch.py | 14 +++++++++----- 1 file changed, 9 insertions(+), 5 deletions(-) diff --git a/wanvideo/utils/basic_flowmatch.py b/wanvideo/utils/basic_flowmatch.py index 591510b..4b184c1 100644 --- a/wanvideo/utils/basic_flowmatch.py +++ b/wanvideo/utils/basic_flowmatch.py @@ -42,10 +42,12 @@ class FlowMatchScheduler(): self.linear_timesteps_weights = bsmntw_weighing def step(self, model_output, timestep, sample, to_final=False): + if timestep.ndim == 2: + timestep = timestep.flatten(0, 1) self.sigmas = self.sigmas.to(model_output.device) self.timesteps = self.timesteps.to(model_output.device) timestep_id = torch.argmin( - (self.timesteps - timestep).abs(), dim=0) + (self.timesteps.unsqueeze(0) - timestep.unsqueeze(1)).abs(), dim=1) sigma = self.sigmas[timestep_id].reshape(-1, 1, 1, 1) if to_final or (timestep_id + 1 >= len(self.timesteps)).any(): sigma_ = 1 if ( @@ -59,11 +61,13 @@ class FlowMatchScheduler(): """ Diffusion forward corruption process. Input: - - clean_latent: the clean latent with shape [B, C, H, W] - - noise: the noise with shape [B, C, H, W] - - timestep: the timestep with shape [B] - Output: the corrupted latent with shape [B, C, H, W] + - clean_latent: the clean latent with shape [B*T, C, H, W] + - noise: the noise with shape [B*T, C, H, W] + - timestep: the timestep with shape [B*T] + Output: the corrupted latent with shape [B*T, C, H, W] """ + if timestep.ndim == 2: + timestep = timestep.flatten(0, 1) self.sigmas = self.sigmas.to(noise.device) self.timesteps = self.timesteps.to(noise.device) timestep_id = torch.argmin(