File size: 1,026 Bytes
9264c1c
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
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