use crate::dit::Proj;
use crate::pool::Pool;
use cortiq_core::CmfModel;
use std::sync::Arc;
const IMAGENET_MEAN: [f32; 3] = [0.485, 0.456, 0.406];
const IMAGENET_STD: [f32; 3] = [0.229, 0.224, 0.225];
const TILE: usize = 256;
const TILE_OVERLAP_MIN: usize = 64;
const CLIP_LENGTH: usize = 17;
const TOKEN_DROP: usize = 3;
struct Block {
norm1: Vec<f32>,
norm2: Vec<f32>,
scale1: Vec<f32>,
scale2: Vec<f32>,
qkv: Proj, qkv_b: Vec<f32>,
out: Proj,
out_b: Vec<f32>,
w1: Proj, w1_b: Vec<f32>,
w2: Proj, w2_b: Vec<f32>,
}
pub struct VideoVae {
post_quant: Proj,
post_quant_b: Vec<f32>,
x_embed: Proj,
x_embed_b: Vec<f32>,
registers: Vec<f32>, blocks: Vec<Block>,
norm_out_w: Vec<f32>,
norm_out_b: Vec<f32>,
proj_out: Proj,
proj_out_b: Vec<f32>,
latents_mean: Vec<f32>,
latents_std: Vec<f32>,
pool: Option<Arc<Pool>>,
dim: usize,
heads: usize,
head_dim: usize,
z_channels: usize,
patch: usize,
patch_t: usize,
n_reg: usize,
rope_theta: f32,
rope_dim: usize,
eps: f64,
}
fn layer_norm(x: &[f32], w: &[f32], b: &[f32], eps: f64, dst: &mut [f32]) {
let n = x.len() as f64;
let mean = x.iter().map(|&v| v as f64).sum::<f64>() / n;
let var = x.iter().map(|&v| (v as f64 - mean).powi(2)).sum::<f64>() / n;
let inv = 1.0 / (var + eps).sqrt();
for (((d, &v), &g), &bb) in dst.iter_mut().zip(x).zip(w).zip(b) {
*d = ((v as f64 - mean) * inv) as f32 * g + bb;
}
}
fn rms_norm_into(x: &[f32], w: &[f32], eps: f64, dst: &mut [f32]) {
let ss = x.iter().map(|&v| (v as f64) * (v as f64)).sum::<f64>() / x.len() as f64;
let inv = 1.0 / (ss + eps).sqrt();
for ((d, &v), &g) in dst.iter_mut().zip(x).zip(w) {
*d = (v as f64 * inv) as f32 * g;
}
}
fn rms_norm_plain(x: &mut [f32], eps: f64) {
let ss = x.iter().map(|&v| (v as f64) * (v as f64)).sum::<f64>() / x.len() as f64;
let inv = 1.0 / (ss + eps).sqrt();
for v in x.iter_mut() {
*v = (*v as f64 * inv) as f32;
}
}
fn silu(v: f32) -> f32 {
v / (1.0 + (-v).exp())
}
struct SendPtr(*mut f32);
unsafe impl Send for SendPtr {}
unsafe impl Sync for SendPtr {}
impl SendPtr {
#[allow(clippy::mut_from_ref)]
unsafe fn row(&self, off: usize, len: usize) -> &mut [f32] {
unsafe { std::slice::from_raw_parts_mut(self.0.add(off), len) }
}
}
fn softmax_inplace(row: &mut [f32]) {
let mx = row.iter().cloned().fold(f32::MIN, f32::max);
let mut den = 0f32;
for r in row.iter_mut() {
*r = (*r - mx).exp();
den += *r;
}
if den > 0.0 {
let inv = 1.0 / den;
for r in row.iter_mut() {
*r *= inv;
}
}
}
fn split_tiles(input_len: usize, ratio: usize) -> (Vec<usize>, Vec<usize>, Vec<usize>) {
if TILE >= input_len {
return (vec![0], vec![input_len], Vec::new());
}
let mut n = input_len.div_ceil(TILE);
let (mut overlaps, mut remaining);
loop {
overlaps = vec![TILE_OVERLAP_MIN; n - 1];
let total: usize = overlaps.iter().sum();
if TILE * n < total + input_len {
n += 1;
continue;
}
remaining = TILE * n - total - input_len;
break;
}
for i in 0..remaining / ratio {
overlaps[i % (n - 1)] += ratio;
}
let mut starts = vec![0usize];
for i in 0..n - 1 {
starts.push(starts[i] + TILE - overlaps[i]);
}
(starts, vec![TILE; n], overlaps)
}
impl VideoVae {
pub fn from_cmf(model: &Arc<CmfModel>) -> Result<Self, String> {
let cfg: serde_json::Value = serde_json::from_slice(
model.tensor_bytes("vvae.config_json").map_err(|e| e.to_string())?,
)
.map_err(|e| format!("vvae.config_json: {e}"))?;
let u = |k: &str, d: usize| cfg[k].as_u64().map(|v| v as usize).unwrap_or(d);
let f32v = |n: &str| crate::dit::cmf_f32(model, n);
let n = u("num_layers", 36);
let mut blocks = Vec::with_capacity(n);
for l in 0..n {
let p = format!("vvae.blocks.{l}");
blocks.push(Block {
norm1: f32v(&format!("{p}.norm1"))?,
norm2: f32v(&format!("{p}.norm2"))?,
scale1: f32v(&format!("{p}.scale1"))?,
scale2: f32v(&format!("{p}.scale2"))?,
qkv: Proj::from_model(model, &format!("{p}.attn.to_qkv.weight"))?,
qkv_b: f32v(&format!("{p}.attn.to_qkv.bias"))?,
out: Proj::from_model(model, &format!("{p}.attn.to_out.weight"))?,
out_b: f32v(&format!("{p}.attn.to_out.bias"))?,
w1: Proj::from_model(model, &format!("{p}.ff.w1.weight"))?,
w1_b: f32v(&format!("{p}.ff.w1.bias"))?,
w2: Proj::from_model(model, &format!("{p}.ff.w2.weight"))?,
w2_b: f32v(&format!("{p}.ff.w2.bias"))?,
});
}
let dim = u("dim", 2048);
let heads = u("heads", 32);
Ok(Self {
post_quant: Proj::from_model(model, "vvae.post_quant_conv.weight")?,
post_quant_b: f32v("vvae.post_quant_conv.bias")?,
x_embed: Proj::from_model(model, "vvae.x_embedder.weight")?,
x_embed_b: f32v("vvae.x_embedder.bias")?,
registers: f32v("vvae.register_tokens")?,
blocks,
norm_out_w: f32v("vvae.norm_out.weight")?,
norm_out_b: f32v("vvae.norm_out.bias")?,
proj_out: Proj::from_model(model, "vvae.proj_out.weight")?,
proj_out_b: f32v("vvae.proj_out.bias")?,
latents_mean: f32v("vvae.latents_mean")?,
latents_std: f32v("vvae.latents_std")?,
pool: Pool::from_env(),
dim,
heads,
head_dim: dim / heads,
z_channels: u("z_channels", 24),
patch: u("patch_size", 16),
patch_t: u("patch_size_t", 4),
n_reg: u("num_register_tokens", 4),
rope_theta: cfg["rope_theta"].as_f64().unwrap_or(100.0) as f32,
rope_dim: ((dim / heads) as f64 * cfg["rope_dim_ratio"].as_f64().unwrap_or(0.75))
as usize,
eps: cfg["eps"].as_f64().unwrap_or(1e-5),
})
}
fn axis(n: usize) -> Vec<f32> {
(0..n)
.map(|i| 2.0 * ((i as f32 + 0.5) / n as f32) - 1.0)
.collect()
}
fn rope_angles(&self, t: usize, h: usize, w: usize, suffix: usize) -> Vec<f32> {
let k = self.rope_dim / 6; let inv: Vec<f32> = (0..k)
.map(|i| 1.0 / self.rope_theta.powf(i as f32 * 2.0 * 3.0 / self.rope_dim as f32))
.collect();
let (ta, ha, wa) = (Self::axis(t), Self::axis(h), Self::axis(w));
let tau = 2.0 * std::f32::consts::PI;
let mut out = Vec::with_capacity((t * h * w + suffix) * 3 * k);
for ti in 0..t {
for hi in 0..h {
for wi in 0..w {
for &c in &[ta[ti], ha[hi], wa[wi]] {
for &f in &inv {
out.push(tau * c * f);
}
}
}
}
}
out.extend(std::iter::repeat_n(0.0, suffix * 3 * k));
out
}
fn attention(&self, qkv: &[f32], n: usize, attn: &mut [f32], angles: &[f32]) {
let (nh, hd, dim) = (self.heads, self.head_dim, self.dim);
let pairs = angles.len() / n; let scale = 1.0 / (hd as f32).sqrt();
let pool = self.pool.as_deref();
let mut qh = vec![0f32; n * hd];
let mut kh = vec![0f32; n * hd];
let mut vt = vec![0f32; hd * n];
let mut scores = vec![0f32; n * n];
let mut oh = vec![0f32; n * hd];
for h in 0..nh {
for p in 0..n {
let base = p * 3 * dim + h * 3 * hd;
qh[p * hd..(p + 1) * hd].copy_from_slice(&qkv[base..base + hd]);
kh[p * hd..(p + 1) * hd].copy_from_slice(&qkv[base + hd..base + 2 * hd]);
for d in 0..hd {
vt[d * n + p] = qkv[base + 2 * hd + d];
}
}
for (buf, _) in [(&mut qh, 0), (&mut kh, 1)] {
for p in 0..n {
let x = &mut buf[p * hd..(p + 1) * hd];
rms_norm_plain(x, self.eps);
for j in 0..pairs {
let (s, c) = angles[p * pairs + j].sin_cos();
let (a, b) = (x[j], x[j + pairs]);
x[j] = a * c - b * s;
x[j + pairs] = a * s + b * c;
}
}
}
for v in qh.iter_mut() {
*v *= scale;
}
crate::fcd_ops::gemm_nt(&qh, &kh, &mut scores, n, hd, n, pool);
let sp = SendPtr(scores.as_mut_ptr());
let soft = |lo: usize, hi: usize| {
for r in lo..hi {
softmax_inplace(unsafe { sp.row(r * n, n) });
}
};
match pool {
Some(pl) => pl.run_rows(n, &soft),
None => soft(0, n),
}
crate::fcd_ops::gemm_nt(&scores, &vt, &mut oh, n, n, hd, pool);
for p in 0..n {
attn[p * dim + h * hd..p * dim + (h + 1) * hd]
.copy_from_slice(&oh[p * hd..(p + 1) * hd]);
}
}
}
fn decode_tile(&self, z: &[f32], t: usize, h: usize, w: usize) -> Vec<f32> {
let pool = self.pool.as_deref();
let dim = self.dim;
let np = t * h * w;
let suffix = 1 + self.n_reg;
let n = np + suffix;
let zc = self.z_channels;
let mut rows = vec![0f32; np * zc];
for c in 0..zc {
for i in 0..np {
rows[i * zc + c] = z[c * np + i];
}
}
let mut pq = vec![0f32; np * zc];
self.post_quant.matmat(&rows, np, &mut pq, pool);
for r in pq.chunks_exact_mut(zc) {
for (v, &b) in r.iter_mut().zip(&self.post_quant_b) {
*v += b;
}
}
let mut x = vec![0f32; n * dim];
self.x_embed.matmat(&pq, np, &mut x[..np * dim], pool);
for r in x[..np * dim].chunks_exact_mut(dim) {
for (v, &b) in r.iter_mut().zip(&self.x_embed_b) {
*v += b;
}
}
x[np * dim..(np + self.n_reg) * dim].copy_from_slice(&self.registers);
let angles = self.rope_angles(t, h, w, suffix);
let mut xn = vec![0f32; n * dim];
let mut qkv = vec![0f32; n * 3 * dim];
let mut attn = vec![0f32; n * dim];
let mut proj = vec![0f32; n * dim];
let inner = 4 * dim;
for blk in &self.blocks {
for (o, src) in xn.chunks_exact_mut(dim).zip(x.chunks_exact(dim)) {
rms_norm_into(src, &blk.norm1, self.eps, o);
}
blk.qkv.matmat(&xn, n, &mut qkv, pool);
for r in qkv.chunks_exact_mut(3 * dim) {
for (v, &b) in r.iter_mut().zip(&blk.qkv_b) {
*v += b;
}
}
self.attention(&qkv, n, &mut attn, &angles);
blk.out.matmat(&attn, n, &mut proj, pool);
for (p, r) in proj.chunks_exact_mut(dim).enumerate() {
for (i, v) in r.iter_mut().enumerate() {
*v += blk.out_b[i];
x[p * dim + i] += *v * blk.scale1[i];
}
}
for (o, src) in xn.chunks_exact_mut(dim).zip(x.chunks_exact(dim)) {
rms_norm_into(src, &blk.norm2, self.eps, o);
}
let mut gu = vec![0f32; n * 2 * inner];
blk.w1.matmat(&xn, n, &mut gu, pool);
let mut act = vec![0f32; n * inner];
for p in 0..n {
let r = &gu[p * 2 * inner..(p + 1) * 2 * inner];
for i in 0..inner {
act[p * inner + i] =
silu(r[i] + blk.w1_b[i]) * (r[inner + i] + blk.w1_b[inner + i]);
}
}
blk.w2.matmat(&act, n, &mut proj, pool);
for (p, r) in proj.chunks_exact_mut(dim).enumerate() {
for (i, v) in r.iter_mut().enumerate() {
*v += blk.w2_b[i];
x[p * dim + i] += *v * blk.scale2[i];
}
}
}
let (pt, ps) = (self.patch_t, self.patch);
let od = 3 * pt * ps * ps;
let mut head = vec![0f32; np * dim];
for (o, src) in head.chunks_exact_mut(dim).zip(x[..np * dim].chunks_exact(dim)) {
layer_norm(src, &self.norm_out_w, &self.norm_out_b, self.eps, o);
}
let mut out = vec![0f32; np * od];
self.proj_out.matmat(&head, np, &mut out, pool);
for r in out.chunks_exact_mut(od) {
for (v, &b) in r.iter_mut().zip(&self.proj_out_b) {
*v += b;
}
}
let (ot, oh, ow) = (t * pt, h * ps, w * ps);
let mut px = vec![0f32; 3 * ot * oh * ow];
for ti in 0..t {
for hi in 0..h {
for wi in 0..w {
let src = &out[((ti * h + hi) * w + wi) * od..((ti * h + hi) * w + wi + 1) * od];
for c in 0..3 {
for a in 0..pt {
for b in 0..ps {
for d in 0..ps {
px[((c * ot + ti * pt + a) * oh + hi * ps + b) * ow
+ wi * ps + d] =
src[((c * pt + a) * ps + b) * ps + d];
}
}
}
}
}
}
}
px
}
fn decode_clip(&self, z: &[f32], t: usize, zh: usize, zw: usize) -> Vec<f32> {
let r = self.patch;
let (height, width) = (zh * r, zw * r);
let (ys, yl, yo) = split_tiles(height, r);
let (xs, xl, xo) = split_tiles(width, r);
let frames = t * self.patch_t;
let mut canvas = vec![0f32; 3 * frames * height * width];
let mut row_tails: Vec<Vec<f32>> = Vec::new();
let mut out_y = 0usize;
for (i, (&ip, &il)) in ys.iter().zip(&yl).enumerate() {
let (zi, zl) = (ip / r, il / r);
let mut new_tails: Vec<Vec<f32>> = Vec::new();
let mut left_tail: Option<Vec<f32>> = None;
let mut out_x = 0usize;
let mut row_h = 0usize;
for (j, (&jp, &jl)) in xs.iter().zip(&xl).enumerate() {
let (zj, zw_t) = (jp / r, jl / r);
let mut sub = vec![0f32; self.z_channels * t * zl * zw_t];
for c in 0..self.z_channels {
for ti in 0..t {
for hh in 0..zl {
for ww in 0..zw_t {
sub[((c * t + ti) * zl + hh) * zw_t + ww] =
z[((c * t + ti) * zh + zi + hh) * zw + zj + ww];
}
}
}
}
let mut tile = self.decode_tile(&sub, t, zl, zw_t);
let (mut th, mut tw) = (zl * r, zw_t * r);
if i + 1 < ys.len() {
new_tails.push(crop(&tile, frames, th, tw, th - yo[i], th, 0, tw));
}
let next_left = if j + 1 < xs.len() {
Some(crop(&tile, frames, th, tw, 0, th, tw - xo[j], tw))
} else {
None
};
if i > 0 {
tile = blend(&row_tails[j], &tile, frames, th, tw, yo[i - 1], 2);
}
if j > 0 {
let lt = left_tail.as_ref().unwrap();
tile = blend(lt, &tile, frames, th, tw, xo[j - 1], 3);
}
left_tail = next_left;
if i + 1 < ys.len() {
tile = crop(&tile, frames, th, tw, 0, th - yo[i], 0, tw);
th -= yo[i];
}
if j + 1 < xs.len() {
tile = crop(&tile, frames, th, tw, 0, th, 0, tw - xo[j]);
tw -= xo[j];
}
for c in 0..3 {
for f in 0..frames {
for hh in 0..th {
let dst = ((c * frames + f) * height + out_y + hh) * width + out_x;
let src = ((c * frames + f) * th + hh) * tw;
canvas[dst..dst + tw].copy_from_slice(&tile[src..src + tw]);
}
}
}
out_x += tw;
row_h = th;
}
row_tails = new_tails;
out_y += row_h;
}
canvas
}
pub fn decode(&self, z: &[f32], t_lat: usize, zh: usize, zw: usize) -> (Vec<f32>, usize) {
let zc = self.z_channels;
let np = t_lat * zh * zw;
let mut zz = vec![0f32; zc * np];
for c in 0..zc {
let (m, s) = (self.latents_mean[c], self.latents_std[c]);
for i in 0..np {
zz[c * np + i] = z[c * np + i] * s + m;
}
}
let ratio_t = self.patch_t;
let chunk_tokens = CLIP_LENGTH.div_ceil(ratio_t); let token_overlap = (chunk_tokens - TOKEN_DROP % chunk_tokens) % chunk_tokens; let frame_pre_pad = (ratio_t - CLIP_LENGTH % ratio_t) % ratio_t; let frame_overlap = (token_overlap * ratio_t).saturating_sub(frame_pre_pad); let chunk_dec = chunk_tokens * ratio_t;
let mut pseudo = t_lat + TOKEN_DROP;
let mut pad = 0usize;
if pseudo % chunk_tokens != 0 {
pad = chunk_tokens - pseudo % chunk_tokens;
pseudo += pad;
}
let mut chunks = pseudo / chunk_tokens - usize::from(TOKEN_DROP > 0);
if chunks < 1 {
pad += chunk_tokens;
chunks += 1;
}
let t_pad = t_lat + pad;
if pad > 0 {
let mut grown = vec![0f32; zc * t_pad * zh * zw];
for c in 0..zc {
for ti in 0..t_pad {
let src = ti.min(t_lat - 1);
let a = (c * t_lat + src) * zh * zw;
let b = (c * t_pad + ti) * zh * zw;
grown[b..b + zh * zw].copy_from_slice(&zz[a..a + zh * zw]);
}
}
zz = grown;
}
let (h, w) = (zh * self.patch, zw * self.patch);
let mut out: Vec<f32> = Vec::new();
let mut carry: Option<(Vec<f32>, usize)> = None;
for i in 0..chunks {
let a = i * chunk_tokens;
let b = (a + chunk_tokens + token_overlap).min(t_pad);
let n = b.saturating_sub(a.min(t_pad));
if n == 0 {
continue;
}
let mut sub = vec![0f32; zc * n * zh * zw];
for c in 0..zc {
let src = (c * t_pad + a) * zh * zw;
let dst = c * n * zh * zw;
sub[dst..dst + n * zh * zw].copy_from_slice(&zz[src..src + n * zh * zw]);
}
let dec = self.decode_clip(&sub, n, zh, zw);
let dec_frames = n * ratio_t;
for j in 0..2 {
let fa = j * chunk_dec;
let fb = (fa + chunk_dec).min(dec_frames);
if fb <= fa + frame_pre_pad {
continue;
}
let mut part = frames_of(&dec, dec_frames, h, w, fa + frame_pre_pad, fb);
let mut pn = fb - fa - frame_pre_pad;
if j == 0 {
if let Some((tail, tn)) = carry.take() {
part = blend_frames(&tail, tn, &part, pn, h, w, frame_overlap);
pn = part.len() / (3 * h * w);
}
append_frames(&mut out, &part, pn, h, w);
} else {
carry = Some((part, pn));
}
}
if i + 1 == chunks {
if let Some((tail, tn)) = carry.take() {
append_frames(&mut out, &tail, tn, h, w);
}
}
}
let frames = out.len() / (3 * h * w);
let want = t_lat * ratio_t - pad_frames(t_lat, pad, chunk_tokens, ratio_t);
for c in 0..3 {
let base = c * frames * h * w;
for v in out[base..base + frames * h * w].iter_mut() {
*v = (*v * IMAGENET_STD[c] + IMAGENET_MEAN[c]).clamp(0.0, 1.0);
}
}
let keep = want.min(frames);
if keep < frames {
let mut trimmed = vec![0f32; 3 * keep * h * w];
for c in 0..3 {
let src = c * frames * h * w;
let dst = c * keep * h * w;
trimmed[dst..dst + keep * h * w].copy_from_slice(&out[src..src + keep * h * w]);
}
out = trimmed;
}
(out, keep)
}
pub fn spatial_ratio(&self) -> usize {
self.patch
}
pub fn temporal_ratio(&self) -> usize {
self.patch_t
}
}
fn pad_frames(t_lat: usize, pad: usize, chunk_tokens: usize, ratio_t: usize) -> usize {
if pad == 0 {
return 0;
}
let intra = CLIP_LENGTH % ratio_t;
if intra == 0 {
return pad * ratio_t;
}
(0..pad)
.map(|k| {
if (t_lat + k) % chunk_tokens == 0 {
intra
} else {
ratio_t
}
})
.sum()
}
#[allow(clippy::too_many_arguments)]
fn crop(x: &[f32], f: usize, h: usize, w: usize, y0: usize, y1: usize, x0: usize, x1: usize) -> Vec<f32> {
let (nh, nw) = (y1 - y0, x1 - x0);
let mut out = vec![0f32; 3 * f * nh * nw];
for c in 0..3 {
for fi in 0..f {
for yy in 0..nh {
let s = ((c * f + fi) * h + y0 + yy) * w + x0;
let d = ((c * f + fi) * nh + yy) * nw;
out[d..d + nw].copy_from_slice(&x[s..s + nw]);
}
}
}
out
}
#[allow(clippy::too_many_arguments)]
fn blend(a: &[f32], b: &[f32], f: usize, h: usize, w: usize, extent: usize, dim: usize) -> Vec<f32> {
let ah = a.len() / (3 * f * w);
let aw = a.len() / (3 * f * h);
let mut out = b.to_vec();
let e = if dim == 2 {
extent.min(ah).min(h)
} else {
extent.min(aw).min(w)
};
for c in 0..3 {
for fi in 0..f {
for k in 0..e {
let wb = k as f32 / e as f32;
let wa = 1.0 - wb;
if dim == 2 {
let sa = ((c * f + fi) * ah + ah - e + k) * w;
let sb = ((c * f + fi) * h + k) * w;
for x in 0..w {
out[sb + x] = a[sa + x] * wa + b[sb + x] * wb;
}
} else {
for y in 0..h {
let sa = ((c * f + fi) * h + y) * aw + aw - e + k;
let sb = ((c * f + fi) * h + y) * w + k;
out[sb] = a[sa] * wa + b[sb] * wb;
}
}
}
}
}
out
}
fn frames_of(x: &[f32], f: usize, h: usize, w: usize, a: usize, b: usize) -> Vec<f32> {
let n = b - a;
let mut out = vec![0f32; 3 * n * h * w];
for c in 0..3 {
let s = (c * f + a) * h * w;
let d = c * n * h * w;
out[d..d + n * h * w].copy_from_slice(&x[s..s + n * h * w]);
}
out
}
#[allow(clippy::too_many_arguments)]
fn blend_frames(
tail: &[f32],
tn: usize,
part: &[f32],
pn: usize,
h: usize,
w: usize,
extent: usize,
) -> Vec<f32> {
let e = extent.min(tn).min(pn);
let mut out = part.to_vec();
for c in 0..3 {
for k in 0..e {
let wb = k as f32 / e as f32;
let wa = 1.0 - wb;
let s = ((c * tn) + tn - e + k) * h * w;
let d = ((c * pn) + k) * h * w;
for i in 0..h * w {
out[d + i] = tail[s + i] * wa + part[d + i] * wb;
}
}
}
out
}
fn append_frames(out: &mut Vec<f32>, part: &[f32], n: usize, h: usize, w: usize) {
let old = out.len() / (3 * h * w);
let total = old + n;
let mut grown = vec![0f32; 3 * total * h * w];
for c in 0..3 {
if old > 0 {
let s = c * old * h * w;
let d = c * total * h * w;
grown[d..d + old * h * w].copy_from_slice(&out[s..s + old * h * w]);
}
let s = c * n * h * w;
let d = (c * total + old) * h * w;
grown[d..d + n * h * w].copy_from_slice(&part[s..s + n * h * w]);
}
*out = grown;
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn tile_schedule_matches_the_reference() {
let (s, l, o) = split_tiles(288, 16);
assert_eq!(s, vec![0, 32]);
assert_eq!(l, vec![256, 256]);
assert_eq!(o, vec![224]);
let (s, l, o) = split_tiles(512, 16);
assert_eq!(s, vec![0, 128, 256]);
assert_eq!(l, vec![256; 3]);
assert_eq!(o, vec![128, 128]);
assert_eq!(split_tiles(256, 16).0, vec![0]);
assert_eq!(split_tiles(128, 16).1, vec![128]);
}
#[test]
fn temporal_constants_come_out_as_the_reference_computes_them() {
let ratio_t = 4usize;
let chunk = CLIP_LENGTH.div_ceil(ratio_t);
assert_eq!(chunk, 5);
assert_eq!((chunk - TOKEN_DROP % chunk) % chunk, 2);
assert_eq!((ratio_t - CLIP_LENGTH % ratio_t) % ratio_t, 3);
}
}
struct EncConv {
w: Vec<f32>, b: Vec<f32>,
out_ch: usize,
in_ch: usize,
k: usize,
stride: usize,
pad: usize,
}
impl EncConv {
fn load(model: &Arc<CmfModel>, name: &str, stride: usize, pad: usize) -> Result<Self, String> {
let e = model
.tensor(&format!("{name}.weight"))
.ok_or_else(|| format!("missing {name}.weight"))?;
let (out_ch, in_ch, k) = (e.shape[0], e.shape[1], e.shape[2]);
Ok(Self {
w: crate::dit::cmf_f32(model, &format!("{name}.weight"))?,
b: crate::dit::cmf_f32(model, &format!("{name}.bias"))?,
out_ch,
in_ch,
k,
stride,
pad,
})
}
fn apply(&self, x: &[f32], h: usize, w: usize, pool: Option<&Pool>) -> (Vec<f32>, usize, usize) {
let (ph, pw) = (h + 2 * self.pad, w + 2 * self.pad);
let oh = (ph - self.k) / self.stride + 1;
let ow = (pw - self.k) / self.stride + 1;
let refl = |i: isize, n: usize| -> usize {
let n = n as isize;
let mut i = i;
while i < 0 || i >= n {
if i < 0 {
i = -i;
}
if i >= n {
i = 2 * (n - 1) - i;
}
}
i as usize
};
let mut out = vec![0f32; self.out_ch * oh * ow];
let ptr = SendPtr(out.as_mut_ptr());
let work = |lo: usize, hi: usize| {
for o in lo..hi {
let dst = unsafe { ptr.row(o * oh * ow, oh * ow) };
dst.fill(self.b[o]);
for i in 0..self.in_ch {
let ker = &self.w[(o * self.in_ch + i) * self.k * self.k
..(o * self.in_ch + i + 1) * self.k * self.k];
let src = &x[i * h * w..(i + 1) * h * w];
for oy in 0..oh {
for ox in 0..ow {
let mut acc = 0f32;
for ky in 0..self.k {
let sy = (oy * self.stride + ky) as isize - self.pad as isize;
let sy = refl(sy, h);
for kx in 0..self.k {
let sx = (ox * self.stride + kx) as isize - self.pad as isize;
acc += ker[ky * self.k + kx] * src[sy * w + refl(sx, w)];
}
}
dst[oy * ow + ox] += acc;
}
}
}
}
};
match pool {
Some(p) => p.run_rows(self.out_ch, &work),
None => work(0, self.out_ch),
}
(out, oh, ow)
}
}
struct GroupNorm {
w: Vec<f32>,
b: Vec<f32>,
}
impl GroupNorm {
fn load(model: &Arc<CmfModel>, name: &str) -> Result<Self, String> {
Ok(Self {
w: crate::dit::cmf_f32(model, &format!("{name}.weight"))?,
b: crate::dit::cmf_f32(model, &format!("{name}.bias"))?,
})
}
fn apply(&self, x: &mut [f32], ch: usize, hw: usize) {
let groups = 32;
let per = ch / groups;
for g in 0..groups {
let seg = &mut x[g * per * hw..(g + 1) * per * hw];
let n = seg.len() as f64;
let mean = seg.iter().map(|&v| v as f64).sum::<f64>() / n;
let var = seg.iter().map(|&v| (v as f64 - mean).powi(2)).sum::<f64>() / n;
let inv = 1.0 / (var + 1e-6).sqrt();
for (i, v) in seg.iter_mut().enumerate() {
let c = g * per + i / hw;
*v = ((*v as f64 - mean) * inv) as f32 * self.w[c] + self.b[c];
}
}
}
}
struct ResBlock {
norm1: GroupNorm,
norm2: GroupNorm,
conv1: EncConv,
conv2: EncConv,
shortcut: Option<EncConv>,
}
pub struct VideoVaeEncoder {
conv_in: EncConv,
levels: Vec<(Vec<ResBlock>, Option<EncConv>)>,
norm_out: GroupNorm,
conv_out: EncConv,
quant: Vec<f32>, quant_b: Vec<f32>,
latents_mean: Vec<f32>,
latents_std: Vec<f32>,
pool: Option<Arc<Pool>>,
z_channels: usize,
ratio: usize,
}
impl VideoVaeEncoder {
pub fn from_cmf(model: &Arc<CmfModel>) -> Result<Self, String> {
let cfg: serde_json::Value = serde_json::from_slice(
model
.tensor_bytes("vvae.config_json")
.map_err(|e| e.to_string())?,
)
.map_err(|e| format!("vvae.config_json: {e}"))?;
let space_down: Vec<usize> = cfg["space_down"]
.as_array()
.map(|a| a.iter().map(|v| v.as_u64().unwrap_or(1) as usize).collect())
.unwrap_or_else(|| vec![2, 2, 2, 2, 1, 1]);
let n_res = cfg["num_res_blocks"].as_u64().unwrap_or(2) as usize;
let mut levels = Vec::new();
for (i, &sd) in space_down.iter().enumerate() {
let mut blocks = Vec::new();
for j in 0..n_res {
let p = format!("vvae.enc.down.{i}.block.{j}");
let shortcut = model
.tensor(&format!("{p}.nin_shortcut.weight"))
.map(|_| EncConv::load(model, &format!("{p}.nin_shortcut"), 1, 0))
.transpose()?;
blocks.push(ResBlock {
norm1: GroupNorm::load(model, &format!("{p}.norm1"))?,
norm2: GroupNorm::load(model, &format!("{p}.norm2"))?,
conv1: EncConv::load(model, &format!("{p}.conv1"), 1, 1)?,
conv2: EncConv::load(model, &format!("{p}.conv2"), 1, 1)?,
shortcut,
});
}
let down = if sd > 1 {
Some(EncConv::load(
model,
&format!("vvae.enc.down.{i}.downsample.conv"),
sd,
0,
)?)
} else {
None
};
levels.push((blocks, down));
}
Ok(Self {
conv_in: EncConv::load(model, "vvae.enc.conv_in", 1, 1)?,
levels,
norm_out: GroupNorm::load(model, "vvae.enc.norm_out")?,
conv_out: EncConv::load(model, "vvae.enc.conv_out", 1, 1)?,
quant: crate::dit::cmf_f32(model, "vvae.quant_conv.weight")?,
quant_b: crate::dit::cmf_f32(model, "vvae.quant_conv.bias")?,
latents_mean: crate::dit::cmf_f32(model, "vvae.latents_mean")?,
latents_std: crate::dit::cmf_f32(model, "vvae.latents_std")?,
pool: Pool::from_env(),
z_channels: cfg["z_channels"].as_u64().unwrap_or(24) as usize,
ratio: cfg["patch_size"].as_u64().unwrap_or(16) as usize,
})
}
fn encode_tile(&self, x: &[f32], h: usize, w: usize) -> (Vec<f32>, usize, usize) {
let pool = self.pool.as_deref();
let (mut cur, mut ch, mut cw) = self.conv_in.apply(x, h, w, pool);
let mut c = self.conv_in.out_ch;
for (blocks, down) in &self.levels {
for b in blocks {
let mut hh = cur.clone();
self.norm_out_like(&b.norm1, &mut hh, c, ch * cw);
for v in hh.iter_mut() {
*v = silu(*v);
}
let (mut t, th, tw) = b.conv1.apply(&hh, ch, cw, pool);
let oc = b.conv1.out_ch;
self.norm_out_like(&b.norm2, &mut t, oc, th * tw);
for v in t.iter_mut() {
*v = silu(*v);
}
let (t2, th2, tw2) = b.conv2.apply(&t, th, tw, pool);
let skip = match &b.shortcut {
Some(s) => s.apply(&cur, ch, cw, pool).0,
None => cur.clone(),
};
cur = t2.iter().zip(&skip).map(|(&a, &b)| a + b).collect();
(ch, cw, c) = (th2, tw2, oc);
}
if let Some(d) = down {
let (padded, ph, pw) = reflect_pad_far(&cur, c, ch, cw);
let (t, th, tw) = d.apply(&padded, ph, pw, pool);
cur = t;
(ch, cw, c) = (th, tw, d.out_ch);
}
}
self.norm_out_like(&self.norm_out, &mut cur, c, ch * cw);
for v in cur.iter_mut() {
*v = silu(*v);
}
let (moments, mh, mw) = self.conv_out.apply(&cur, ch, cw, pool);
let n = 2 * self.z_channels;
let mut out = vec![0f32; n * mh * mw];
for o in 0..n {
for p in 0..mh * mw {
let mut acc = self.quant_b[o];
for i in 0..n {
acc += self.quant[o * n + i] * moments[i * mh * mw + p];
}
out[o * mh * mw + p] = acc;
}
}
(out, mh, mw)
}
fn norm_out_like(&self, g: &GroupNorm, x: &mut [f32], ch: usize, hw: usize) {
g.apply(x, ch, hw);
}
pub fn encode_frame(&self, rgb: &[f32], h: usize, w: usize) -> (Vec<f32>, usize, usize) {
let mut x = vec![0f32; 3 * h * w];
for c in 0..3 {
for p in 0..h * w {
let v = (rgb[c * h * w + p] + 1.0) * 0.5;
x[c * h * w + p] = (v - IMAGENET_MEAN[c]) / IMAGENET_STD[c];
}
}
let r = self.ratio;
let (ys, yl, yo) = split_tiles(h, r);
let (xs, xl, xo) = split_tiles(w, r);
let (zh, zw) = (h / r, w / r);
let zc = 2 * self.z_channels;
let mut rows: Vec<Vec<Vec<f32>>> = Vec::new();
let mut dims: Vec<Vec<(usize, usize)>> = Vec::new();
for (&ip, &il) in ys.iter().zip(&yl) {
let mut row = Vec::new();
let mut rd = Vec::new();
for (&jp, &jl) in xs.iter().zip(&xl) {
let mut sub = vec![0f32; 3 * il * jl];
for c in 0..3 {
for yy in 0..il {
let s = (c * h + ip + yy) * w + jp;
let d = (c * il + yy) * jl;
sub[d..d + jl].copy_from_slice(&x[s..s + jl]);
}
}
let (t, th, tw) = self.encode_tile(&sub, il, jl);
row.push(t);
rd.push((th, tw));
}
rows.push(row);
dims.push(rd);
}
let mut canvas = vec![0f32; zc * zh * zw];
let mut out_y = 0usize;
for i in 0..rows.len() {
let mut out_x = 0usize;
let mut row_h = 0usize;
for j in 0..rows[i].len() {
let (th, tw) = dims[i][j];
let mut tile = rows[i][j].clone();
let (mut ch_, mut cw_) = (th, tw);
if i > 0 {
tile = blend_plane(&rows[i - 1][j], &tile, zc, dims[i - 1][j], (ch_, cw_), yo[i - 1] / r, 0);
}
if j > 0 {
tile = blend_plane(&rows[i][j - 1], &tile, zc, dims[i][j - 1], (ch_, cw_), xo[j - 1] / r, 1);
}
if i + 1 < rows.len() {
tile = crop_plane(&tile, zc, ch_, cw_, 0, ch_ - yo[i] / r, 0, cw_);
ch_ -= yo[i] / r;
}
if j + 1 < rows[i].len() {
tile = crop_plane(&tile, zc, ch_, cw_, 0, ch_, 0, cw_ - xo[j] / r);
cw_ -= xo[j] / r;
}
for c in 0..zc {
for yy in 0..ch_ {
let d = (c * zh + out_y + yy) * zw + out_x;
let s = (c * ch_ + yy) * cw_;
canvas[d..d + cw_].copy_from_slice(&tile[s..s + cw_]);
}
}
out_x += cw_;
row_h = ch_;
}
out_y += row_h;
}
let mut z = vec![0f32; self.z_channels * zh * zw];
for c in 0..self.z_channels {
let (m, s) = (self.latents_mean[c], self.latents_std[c]);
for p in 0..zh * zw {
z[c * zh * zw + p] = (canvas[c * zh * zw + p] - m) / s;
}
}
(z, zh, zw)
}
}
fn reflect_pad_far(x: &[f32], c: usize, h: usize, w: usize) -> (Vec<f32>, usize, usize) {
let (ph, pw) = (h + 1, w + 1);
let mut out = vec![0f32; c * ph * pw];
for ci in 0..c {
for y in 0..ph {
let sy = if y < h { y } else { h - 2 };
for x2 in 0..pw {
let sx = if x2 < w { x2 } else { w - 2 };
out[(ci * ph + y) * pw + x2] = x[(ci * h + sy) * w + sx];
}
}
}
(out, ph, pw)
}
fn blend_plane(
a: &[f32],
b: &[f32],
c: usize,
ad: (usize, usize),
bd: (usize, usize),
extent: usize,
dim: usize,
) -> Vec<f32> {
let (ah, aw) = ad;
let (bh, bw) = bd;
let mut out = b.to_vec();
let e = if dim == 0 { extent.min(ah).min(bh) } else { extent.min(aw).min(bw) };
if e == 0 {
return out;
}
for ci in 0..c {
for k in 0..e {
let wb = k as f32 / e as f32;
let wa = 1.0 - wb;
if dim == 0 {
for x in 0..bw.min(aw) {
let sa = (ci * ah + ah - e + k) * aw + x;
let sb = (ci * bh + k) * bw + x;
out[sb] = a[sa] * wa + b[sb] * wb;
}
} else {
for y in 0..bh.min(ah) {
let sa = (ci * ah + y) * aw + aw - e + k;
let sb = (ci * bh + y) * bw + k;
out[sb] = a[sa] * wa + b[sb] * wb;
}
}
}
}
out
}
#[allow(clippy::too_many_arguments)]
fn crop_plane(x: &[f32], c: usize, h: usize, w: usize, y0: usize, y1: usize, x0: usize, x1: usize) -> Vec<f32> {
let (nh, nw) = (y1 - y0, x1 - x0);
let mut out = vec![0f32; c * nh * nw];
for ci in 0..c {
for y in 0..nh {
let s = (ci * h + y0 + y) * w + x0;
let d = (ci * nh + y) * nw;
out[d..d + nw].copy_from_slice(&x[s..s + nw]);
}
}
out
}