Skip to main content

VaeDecoder

Struct VaeDecoder 

Source
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: f32

Implementations§

Source§

impl VaeDecoder

Source

pub fn chain_args(&self) -> VaeChainArgs<'_>

The borrowed device-chain view (see VaeChainArgs).

Source

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

Source

pub fn load_dir(dir: &Path) -> Result<Self, String>

Load from a diffusers vae/ directory (config.json + diffusion_pytorch_model.safetensors).

Source

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}
Source

pub fn decode(&self, z: &[f32], h: usize, w: usize) -> Vec<f32>

Decode latents [latent_channels, h, w] (model scale, i.e. as produced by the diffusion loop) into RGB [3, 8h, 8w] in [-1, 1].

Auto Trait Implementations§

Blanket Implementations§

Source§

impl<T> Any for T
where T: 'static + ?Sized,

Source§

fn type_id(&self) -> TypeId

Gets the TypeId of self. Read more
Source§

impl<T> Borrow<T> for T
where T: ?Sized,

Source§

fn borrow(&self) -> &T

Immutably borrows from an owned value. Read more
Source§

impl<T> BorrowMut<T> for T
where T: ?Sized,

Source§

fn borrow_mut(&mut self) -> &mut T

Mutably borrows from an owned value. Read more
Source§

impl<T> From<T> for T

Source§

fn from(t: T) -> T

Returns the argument unchanged.

Source§

impl<T> Instrument for T

Source§

fn instrument(self, span: Span) -> Instrumented<Self> ⓘ

Instruments this type with the provided Span, returning an Instrumented wrapper. Read more
Source§

fn in_current_span(self) -> Instrumented<Self> ⓘ

Instruments this type with the current Span, returning an Instrumented wrapper. Read more
Source§

impl<T, U> Into<U> for T
where U: From<T>,

Source§

fn into(self) -> U

Calls U::from(self).

That is, this conversion is whatever the implementation of From<T> for U chooses to do.

Source§

impl<T, U> TryFrom<U> for T
where U: Into<T>,

Source§

type Error = !

The type returned in the event of a conversion error.
Source§

fn try_from(value: U) -> Result<T, !>

Performs the conversion.
Source§

impl<T, U> TryInto<U> for T
where U: TryFrom<T>,

Source§

type Error = <U as TryFrom<T>>::Error

The type returned in the event of a conversion error.
Source§

fn try_into(self) -> Result<U, <U as TryFrom<T>>::Error>

Performs the conversion.
Source§

impl<T> WithSubscriber for T

Source§

fn with_subscriber<S>(self, subscriber: S) -> WithDispatch<Self> ⓘ
where S: Into<Dispatch>,

Attaches the provided Subscriber to this type, returning a WithDispatch wrapper. Read more
Source§

fn with_current_subscriber(self) -> WithDispatch<Self> ⓘ

Attaches the current default Subscriber to this type, returning a WithDispatch wrapper. Read more