import torch from diffusers import FlowMatchEulerDiscreteScheduler def sampling_sigmas(steps): """Exact schedules used to distill the separately trained students.""" if steps == 1: return [1.0, 0.0] if steps == 2: return [1.0, 5.0 / 6.0, 0.0] raise ValueError("Only separately distilled 1-step and 2-step models are supported") class DMDFlowScheduler(FlowMatchEulerDiscreteScheduler): """Exact ODE times shared with DMD training; no second shift at inference.""" dmd_steps = 1 def set_timesteps(self, num_inference_steps=None, device=None, **kwargs): if num_inference_steps != self.dmd_steps: raise ValueError('Inference step count differs from distilled schedule') self.num_inference_steps = self.dmd_steps self.sigmas = torch.tensor(sampling_sigmas(self.dmd_steps),dtype=torch.float32,device=device) self.timesteps = self.sigmas[:-1]*self.config.num_train_timesteps self._step_index = None self._begin_index = None