pub struct VaeDecoder {
pub latent_channels: usize,
pub scaling_factor: f32,
pub shift_factor: f32,
/* private fields */
}Expand description
The full decoder: conv_in → mid(res, attn, res) → up-blocks → GroupNorm + SiLU → conv_out, plus the diffusers latent
de-normalization z/scaling_factor + shift_factor.
Fields§
§latent_channels: usize§scaling_factor: f32§shift_factor: f32Implementations§
Source§impl VaeDecoder
impl VaeDecoder
Sourcepub fn chain_args(&self) -> VaeChainArgs<'_>
pub fn chain_args(&self) -> VaeChainArgs<'_>
The borrowed device-chain view (see VaeChainArgs).
Sourcepub fn decode_fast(&self, z: &[f32], h: usize, w: usize) -> Vec<f32>
pub fn decode_fast(&self, z: &[f32], h: usize, w: usize) -> Vec<f32>
decode with the resident device chain first: de-normalise on the
host, try gpu::vae_decode_chain, and fall back to decode (the
unchanged Lumina path) when the backend declines. Same contract as
decode: model-scale latents in, RGB [3, 8h, 8w] in ≈[-1, 1] out.
Examples found in repository?
examples/zimage_vaecheck.rs (line 39)
28fn main() {
29 let a: Vec<String> = std::env::args().collect();
30 let model = cortiq_core::CmfModel::open(&a[1]).unwrap();
31 let o = read_st(&a[2]);
32 let reps: usize = a.get(3).and_then(|v| v.parse().ok()).unwrap_or(3);
33 let (zs, z) = &o["z"];
34 let (_, img) = &o["img"];
35 let (h, w) = (zs[zs.len() - 2], zs[zs.len() - 1]);
36 let vae = cortiq_engine::vae::VaeDecoder::from_cmf(&model).unwrap();
37 for r in 0..reps {
38 let t = std::time::Instant::now();
39 let got = vae.decode_fast(z, h, w);
40 let dt = t.elapsed().as_secs_f64();
41 let (mut d, mut n, mut se) = (0f64, 0f64, 0f64);
42 for (x, y) in got.iter().zip(img) {
43 d += (*x as f64 - *y as f64).powi(2);
44 n += (*y as f64).powi(2);
45 let q = |v: f32| ((v / 2.0 + 0.5).clamp(0.0, 1.0) * 255.0).round_ties_even() as f64;
46 se += (q(*x) - q(*y)).powi(2);
47 }
48 let psnr = 10.0 * (255f64.powi(2) / (se / got.len() as f64).max(1e-12)).log10();
49 println!("rep {r}: {:.3} s img rel {:.3e} u8 PSNR {:.2} dB", dt, (d / n).sqrt(), psnr);
50 }
51}Source§impl VaeDecoder
impl VaeDecoder
Sourcepub fn load_dir(dir: &Path) -> Result<Self, String>
pub fn load_dir(dir: &Path) -> Result<Self, String>
Load from a diffusers vae/ directory (config.json +
diffusion_pytorch_model.safetensors).
Sourcepub fn from_cmf(model: &CmfModel) -> Result<Self, String>
pub fn from_cmf(model: &CmfModel) -> Result<Self, String>
Load from a packaged imagegen .cmf (vae.* tensors, stored
f16/f32, + vae.config_json).
Examples found in repository?
examples/zimage_vaecheck.rs (line 36)
28fn main() {
29 let a: Vec<String> = std::env::args().collect();
30 let model = cortiq_core::CmfModel::open(&a[1]).unwrap();
31 let o = read_st(&a[2]);
32 let reps: usize = a.get(3).and_then(|v| v.parse().ok()).unwrap_or(3);
33 let (zs, z) = &o["z"];
34 let (_, img) = &o["img"];
35 let (h, w) = (zs[zs.len() - 2], zs[zs.len() - 1]);
36 let vae = cortiq_engine::vae::VaeDecoder::from_cmf(&model).unwrap();
37 for r in 0..reps {
38 let t = std::time::Instant::now();
39 let got = vae.decode_fast(z, h, w);
40 let dt = t.elapsed().as_secs_f64();
41 let (mut d, mut n, mut se) = (0f64, 0f64, 0f64);
42 for (x, y) in got.iter().zip(img) {
43 d += (*x as f64 - *y as f64).powi(2);
44 n += (*y as f64).powi(2);
45 let q = |v: f32| ((v / 2.0 + 0.5).clamp(0.0, 1.0) * 255.0).round_ties_even() as f64;
46 se += (q(*x) - q(*y)).powi(2);
47 }
48 let psnr = 10.0 * (255f64.powi(2) / (se / got.len() as f64).max(1e-12)).log10();
49 println!("rep {r}: {:.3} s img rel {:.3e} u8 PSNR {:.2} dB", dt, (d / n).sqrt(), psnr);
50 }
51}Auto Trait Implementations§
impl Freeze for VaeDecoder
impl RefUnwindSafe for VaeDecoder
impl Send for VaeDecoder
impl Sync for VaeDecoder
impl Unpin for VaeDecoder
impl UnsafeUnpin for VaeDecoder
impl UnwindSafe for VaeDecoder
Blanket Implementations§
Source§impl<T> BorrowMut<T> for Twhere
T: ?Sized,
impl<T> BorrowMut<T> for Twhere
T: ?Sized,
Source§fn borrow_mut(&mut self) -> &mut T
fn borrow_mut(&mut self) -> &mut T
Mutably borrows from an owned value. Read more