use crate::vae::attention::{AttentionBlock, Linear};
use crate::vae::conv::Conv2d;
use crate::vae::error::{VaeError, VaeResult};
use crate::vae::norm::GroupNorm;
use crate::vae::ops::{bn_denorm, silu_inplace, unpatchify, upsample_nearest2x};
use crate::vae::resnet::ResnetBlock2D;
use crate::vae::weights::VaeWeights;
const NUM_GROUPS: usize = 32;
const GN_EPS: f32 = 1e-6;
const BN_EPS: f32 = 1e-4;
const LATENT_CH: usize = 32;
#[derive(Clone)]
pub struct Map {
pub data: Vec<f32>,
pub c: usize,
pub h: usize,
pub w: usize,
}
impl Map {
fn new(data: Vec<f32>, c: usize, h: usize, w: usize) -> Self {
Self { data, c, h, w }
}
}
#[derive(Default)]
pub struct DecodeTaps {
pub bn_denorm: Option<Vec<f32>>,
pub unpatchified: Option<Vec<f32>>,
pub post_quant_conv: Option<Vec<f32>>,
pub conv_in: Option<Vec<f32>>,
pub mid: Option<Vec<f32>>,
pub up: Vec<Vec<f32>>,
pub conv_norm_out: Option<Vec<f32>>,
}
struct UpBlock {
resnets: Vec<ResnetBlock2D>,
upsampler: Option<Conv2d>,
}
struct MidBlock {
resnet0: ResnetBlock2D,
attention: AttentionBlock,
resnet1: ResnetBlock2D,
}
pub struct VaeDecoder {
bn_mean: Vec<f32>,
bn_var: Vec<f32>,
post_quant_conv: Conv2d,
conv_in: Conv2d,
mid_block: MidBlock,
up_blocks: Vec<UpBlock>,
conv_norm_out: GroupNorm,
conv_out: Conv2d,
}
impl VaeDecoder {
pub fn from_weights(w: &VaeWeights) -> VaeResult<Self> {
let bn_mean = w.vec1("bn.running_mean")?.data.clone();
let bn_var = w.vec1("bn.running_var")?.data.clone();
let post_quant_conv = load_conv(w, "post_quant_conv", 0)?;
let conv_in = load_conv(w, "decoder.conv_in", 1)?;
let mid_block = load_mid_block(w, "decoder.mid_block")?;
let mut up_blocks = Vec::with_capacity(4);
for i in 0..4 {
let prefix = format!("decoder.up_blocks.{i}");
let has_upsample = i < 3;
up_blocks.push(load_up_block(w, &prefix, 3, has_upsample)?);
}
let conv_norm_out = load_group_norm(w, "decoder.conv_norm_out")?;
let conv_out = load_conv(w, "decoder.conv_out", 1)?;
Ok(Self {
bn_mean,
bn_var,
post_quant_conv,
conv_in,
mid_block,
up_blocks,
conv_norm_out,
conv_out,
})
}
pub fn decode_packed_latents(
&self,
packed: &[f32],
ph: usize,
pw: usize,
mut taps: Option<&mut DecodeTaps>,
) -> VaeResult<Map> {
let packed_ch = 4 * LATENT_CH; if packed.len() != packed_ch * ph * pw {
return Err(VaeError::Shape(format!(
"decode_packed input len {} != 128*{ph}*{pw}",
packed.len()
)));
}
let denorm = bn_denorm(
packed,
&self.bn_mean,
&self.bn_var,
packed_ch,
ph,
pw,
BN_EPS,
)?;
if let Some(t) = taps.as_deref_mut() {
t.bn_denorm = Some(denorm.clone());
}
let up = unpatchify(&denorm, packed_ch, ph, pw)?;
if let Some(t) = taps.as_deref_mut() {
t.unpatchified = Some(up.data.clone());
}
let latents = Map::new(up.data, up.c, up.h, up.w);
self.decode(&latents, taps)
}
pub fn decode(&self, latents: &Map, mut taps: Option<&mut DecodeTaps>) -> VaeResult<Map> {
let pqc = self
.post_quant_conv
.forward(&latents.data, latents.h, latents.w)?;
let mut cur = Map::new(pqc.data, self.post_quant_conv.out_ch, pqc.h, pqc.w);
if let Some(t) = taps.as_deref_mut() {
t.post_quant_conv = Some(cur.data.clone());
}
let ci = self.conv_in.forward(&cur.data, cur.h, cur.w)?;
cur = Map::new(ci.data, self.conv_in.out_ch, ci.h, ci.w);
if let Some(t) = taps.as_deref_mut() {
t.conv_in = Some(cur.data.clone());
}
cur = self.mid_block.forward(&cur)?;
if let Some(t) = taps.as_deref_mut() {
t.mid = Some(cur.data.clone());
}
for ub in &self.up_blocks {
cur = ub.forward(&cur)?;
if let Some(t) = taps.as_deref_mut() {
t.up.push(cur.data.clone());
}
}
self.conv_norm_out
.forward_inplace(&mut cur.data, cur.h, cur.w)?;
if let Some(t) = taps {
t.conv_norm_out = Some(cur.data.clone());
}
silu_inplace(&mut cur.data);
let out = self.conv_out.forward(&cur.data, cur.h, cur.w)?;
Ok(Map::new(out.data, self.conv_out.out_ch, out.h, out.w))
}
}
impl MidBlock {
fn forward(&self, x: &Map) -> VaeResult<Map> {
let r0 = self.resnet0.forward(&x.data, x.h, x.w)?;
let cur = Map::new(r0, self.resnet0.out_ch, x.h, x.w);
let attn = self.attention.forward(&cur.data, cur.h, cur.w)?;
let cur = Map::new(attn, cur.c, cur.h, cur.w);
let r1 = self.resnet1.forward(&cur.data, cur.h, cur.w)?;
Ok(Map::new(r1, self.resnet1.out_ch, cur.h, cur.w))
}
}
impl UpBlock {
fn forward(&self, x: &Map) -> VaeResult<Map> {
let mut cur = x.clone();
for resnet in &self.resnets {
let data = resnet.forward(&cur.data, cur.h, cur.w)?;
cur = Map::new(data, resnet.out_ch, cur.h, cur.w);
}
if let Some(conv) = self.upsampler.as_ref() {
let upsampled = upsample_nearest2x(&cur.data, cur.c, cur.h, cur.w)?;
let out = conv.forward(&upsampled.data, upsampled.h, upsampled.w)?;
cur = Map::new(out.data, conv.out_ch, out.h, out.w);
}
Ok(cur)
}
}
fn load_conv(w: &VaeWeights, prefix: &str, pad: usize) -> VaeResult<Conv2d> {
let weight = w.get(&format!("{prefix}.weight"))?;
let bias = w.get(&format!("{prefix}.bias"))?;
Conv2d::from_weights(&weight.data, &weight.shape, &bias.data, pad)
}
fn load_linear(w: &VaeWeights, prefix: &str) -> VaeResult<Linear> {
let weight = w.get(&format!("{prefix}.weight"))?;
let bias = w.get(&format!("{prefix}.bias"))?;
Linear::from_weights(&weight.data, &weight.shape, &bias.data)
}
fn load_group_norm(w: &VaeWeights, prefix: &str) -> VaeResult<GroupNorm> {
let weight = w.vec1(&format!("{prefix}.weight"))?;
let bias = w.vec1(&format!("{prefix}.bias"))?;
GroupNorm::new(&weight.data, &bias.data, NUM_GROUPS, GN_EPS)
}
fn load_resnet(w: &VaeWeights, prefix: &str) -> VaeResult<ResnetBlock2D> {
let norm1 = load_group_norm(w, &format!("{prefix}.norm1"))?;
let conv1 = load_conv(w, &format!("{prefix}.conv1"), 1)?;
let norm2 = load_group_norm(w, &format!("{prefix}.norm2"))?;
let conv2 = load_conv(w, &format!("{prefix}.conv2"), 1)?;
let in_ch = conv1.in_ch;
let out_ch = conv2.out_ch;
let conv_shortcut = if in_ch != out_ch {
Some(load_conv(w, &format!("{prefix}.conv_shortcut"), 0)?)
} else {
None
};
Ok(ResnetBlock2D {
norm1,
conv1,
norm2,
conv2,
conv_shortcut,
in_ch,
out_ch,
})
}
fn load_mid_block(w: &VaeWeights, prefix: &str) -> VaeResult<MidBlock> {
let resnet0 = load_resnet(w, &format!("{prefix}.resnets.0"))?;
let resnet1 = load_resnet(w, &format!("{prefix}.resnets.1"))?;
let attention = load_attention(w, &format!("{prefix}.attentions.0"))?;
Ok(MidBlock {
resnet0,
attention,
resnet1,
})
}
fn load_attention(w: &VaeWeights, prefix: &str) -> VaeResult<AttentionBlock> {
let group_norm = load_group_norm(w, &format!("{prefix}.group_norm"))?;
let to_q = load_linear(w, &format!("{prefix}.to_q"))?;
let to_k = load_linear(w, &format!("{prefix}.to_k"))?;
let to_v = load_linear(w, &format!("{prefix}.to_v"))?;
let to_out = load_linear(w, &format!("{prefix}.to_out"))?;
let channels = group_norm.channels;
Ok(AttentionBlock {
group_norm,
to_q,
to_k,
to_v,
to_out,
channels,
})
}
fn load_up_block(
w: &VaeWeights,
prefix: &str,
num_layers: usize,
has_upsample: bool,
) -> VaeResult<UpBlock> {
let mut resnets = Vec::with_capacity(num_layers);
for i in 0..num_layers {
resnets.push(load_resnet(w, &format!("{prefix}.resnets.{i}"))?);
}
let upsampler = if has_upsample {
Some(load_conv(w, &format!("{prefix}.upsamplers.0.conv"), 1)?)
} else {
None
};
Ok(UpBlock { resnets, upsampler })
}