candle_transformers/models/z_image/
scheduler.rs1use candle::{Result, Tensor};
6
7#[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 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#[derive(Debug, Clone)]
48pub struct FlowMatchEulerDiscreteScheduler {
49 pub config: SchedulerConfig,
51 pub timesteps: Vec<f64>,
53 pub sigmas: Vec<f64>,
55 pub sigma_min: f64,
57 pub sigma_max: f64,
59 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 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 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 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 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 if let Some(mu) = mu {
129 if self.config.use_dynamic_shifting {
130 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 sigmas.push(0.0);
153
154 self.timesteps = timesteps;
155 self.sigmas = sigmas;
156 self.step_index = 0;
157 }
158
159 pub fn current_sigma(&self) -> f64 {
161 self.sigmas[self.step_index]
162 }
163
164 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 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 let prev_sample = (sample + (model_output * dt)?)?;
187
188 self.step_index += 1;
189 Ok(prev_sample)
190 }
191
192 pub fn reset(&mut self) {
194 self.step_index = 0;
195 }
196
197 pub fn num_inference_steps(&self) -> usize {
199 self.timesteps.len()
200 }
201
202 pub fn step_index(&self) -> usize {
204 self.step_index
205 }
206
207 pub fn is_complete(&self) -> bool {
209 self.step_index >= self.timesteps.len()
210 }
211}
212
213pub 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
233pub 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;