Skip to main content

candle_transformers/models/z_image/
scheduler.rs

1//! FlowMatch Euler Discrete Scheduler for Z-Image
2//!
3//! Implements the flow matching scheduler used in Z-Image generation.
4
5use candle::{Result, Tensor};
6
7/// FlowMatchEulerDiscreteScheduler configuration
8#[derive(Debug, Clone, serde::Deserialize)]
9pub struct SchedulerConfig {
10    #[serde(default = "default_num_train_timesteps")]
11    pub num_train_timesteps: usize,
12    #[serde(default = "default_shift")]
13    pub shift: f64,
14    #[serde(default)]
15    pub use_dynamic_shifting: bool,
16}
17
18fn default_num_train_timesteps() -> usize {
19    1000
20}
21fn default_shift() -> f64 {
22    3.0
23}
24
25impl Default for SchedulerConfig {
26    fn default() -> Self {
27        Self {
28            num_train_timesteps: default_num_train_timesteps(),
29            shift: default_shift(),
30            use_dynamic_shifting: false,
31        }
32    }
33}
34
35impl SchedulerConfig {
36    /// Create configuration for Z-Image Turbo
37    pub fn z_image_turbo() -> Self {
38        Self {
39            num_train_timesteps: 1000,
40            shift: 3.0,
41            use_dynamic_shifting: false,
42        }
43    }
44}
45
46/// FlowMatch Euler Discrete Scheduler
47#[derive(Debug, Clone)]
48pub struct FlowMatchEulerDiscreteScheduler {
49    /// Configuration
50    pub config: SchedulerConfig,
51    /// Timesteps for inference
52    pub timesteps: Vec<f64>,
53    /// Sigma values
54    pub sigmas: Vec<f64>,
55    /// Minimum sigma
56    pub sigma_min: f64,
57    /// Maximum sigma
58    pub sigma_max: f64,
59    /// Current step index
60    step_index: usize,
61}
62
63impl FlowMatchEulerDiscreteScheduler {
64    pub fn new(config: SchedulerConfig) -> Self {
65        let num_train_timesteps = config.num_train_timesteps;
66        let shift = config.shift;
67
68        // Generate initial sigmas
69        let timesteps: Vec<f64> = (1..=num_train_timesteps).rev().map(|t| t as f64).collect();
70
71        let sigmas: Vec<f64> = timesteps
72            .iter()
73            .map(|&t| t / num_train_timesteps as f64)
74            .collect();
75
76        // Apply shift
77        let sigmas: Vec<f64> = if !config.use_dynamic_shifting {
78            sigmas
79                .iter()
80                .map(|&s| shift * s / (1.0 + (shift - 1.0) * s))
81                .collect()
82        } else {
83            sigmas
84        };
85
86        let timesteps: Vec<f64> = sigmas
87            .iter()
88            .map(|&s| s * num_train_timesteps as f64)
89            .collect();
90
91        let sigma_max = sigmas[0];
92        let sigma_min = *sigmas.last().unwrap_or(&0.0);
93
94        Self {
95            config,
96            timesteps,
97            sigmas,
98            sigma_min,
99            sigma_max,
100            step_index: 0,
101        }
102    }
103
104    /// Set timesteps for inference
105    ///
106    /// # Arguments
107    /// * `num_inference_steps` - Number of denoising steps
108    /// * `mu` - Optional time shift parameter (from calculate_shift)
109    pub fn set_timesteps(&mut self, num_inference_steps: usize, mu: Option<f64>) {
110        let sigma_max = self.sigmas[0];
111        let sigma_min = *self.sigmas.last().unwrap_or(&0.0);
112
113        // Linear interpolation to generate timesteps
114        let timesteps: Vec<f64> = (0..num_inference_steps)
115            .map(|i| {
116                let t = i as f64 / num_inference_steps as f64;
117                sigma_max * (1.0 - t) + sigma_min * t
118            })
119            .map(|s| s * self.config.num_train_timesteps as f64)
120            .collect();
121
122        let mut sigmas: Vec<f64> = timesteps
123            .iter()
124            .map(|&t| t / self.config.num_train_timesteps as f64)
125            .collect();
126
127        // Apply shift
128        if let Some(mu) = mu {
129            if self.config.use_dynamic_shifting {
130                // time_shift: exp(mu) / (exp(mu) + (1/t - 1))
131                sigmas = sigmas
132                    .iter()
133                    .map(|&t| {
134                        if t <= 0.0 {
135                            0.0
136                        } else {
137                            let e_mu = mu.exp();
138                            e_mu / (e_mu + (1.0 / t - 1.0))
139                        }
140                    })
141                    .collect();
142            }
143        } else if !self.config.use_dynamic_shifting {
144            let shift = self.config.shift;
145            sigmas = sigmas
146                .iter()
147                .map(|&s| shift * s / (1.0 + (shift - 1.0) * s))
148                .collect();
149        }
150
151        // Add terminal sigma = 0
152        sigmas.push(0.0);
153
154        self.timesteps = timesteps;
155        self.sigmas = sigmas;
156        self.step_index = 0;
157    }
158
159    /// Get current sigma value
160    pub fn current_sigma(&self) -> f64 {
161        self.sigmas[self.step_index]
162    }
163
164    /// Get current timestep (for model input)
165    /// Converts scheduler timestep to model input format: (1000 - t) / 1000
166    pub fn current_timestep_normalized(&self) -> f64 {
167        let t = self.timesteps.get(self.step_index).copied().unwrap_or(0.0);
168        (1000.0 - t) / 1000.0
169    }
170
171    /// Euler step
172    ///
173    /// # Arguments
174    /// * `model_output` - Model predicted velocity field
175    /// * `sample` - Current sample x_t
176    ///
177    /// # Returns
178    /// Next sample x_{t-1}
179    pub fn step(&mut self, model_output: &Tensor, sample: &Tensor) -> Result<Tensor> {
180        let sigma = self.sigmas[self.step_index];
181        let sigma_next = self.sigmas[self.step_index + 1];
182
183        let dt = sigma_next - sigma;
184
185        // prev_sample = sample + dt * model_output
186        let prev_sample = (sample + (model_output * dt)?)?;
187
188        self.step_index += 1;
189        Ok(prev_sample)
190    }
191
192    /// Reset scheduler state
193    pub fn reset(&mut self) {
194        self.step_index = 0;
195    }
196
197    /// Get number of inference steps
198    pub fn num_inference_steps(&self) -> usize {
199        self.timesteps.len()
200    }
201
202    /// Get current step index
203    pub fn step_index(&self) -> usize {
204        self.step_index
205    }
206
207    /// Check if denoising is complete
208    pub fn is_complete(&self) -> bool {
209        self.step_index >= self.timesteps.len()
210    }
211}
212
213/// Calculate timestep shift parameter mu
214///
215/// # Arguments
216/// * `image_seq_len` - Image sequence length (after patchify)
217/// * `base_seq_len` - Base sequence length (typically 256)
218/// * `max_seq_len` - Maximum sequence length (typically 4096)
219/// * `base_shift` - Base shift value (typically 0.5)
220/// * `max_shift` - Maximum shift value (typically 1.15)
221pub fn calculate_shift(
222    image_seq_len: usize,
223    base_seq_len: usize,
224    max_seq_len: usize,
225    base_shift: f64,
226    max_shift: f64,
227) -> f64 {
228    let m = (max_shift - base_shift) / (max_seq_len - base_seq_len) as f64;
229    let b = base_shift - m * base_seq_len as f64;
230    image_seq_len as f64 * m + b
231}
232
233/// Constants for shift calculation
234pub const BASE_IMAGE_SEQ_LEN: usize = 256;
235pub const MAX_IMAGE_SEQ_LEN: usize = 4096;
236pub const BASE_SHIFT: f64 = 0.5;
237pub const MAX_SHIFT: f64 = 1.15;