MCPcopy Create free account
hub / github.com/evilsocket/cake / forward

Method forward

cake-core/src/models/vibevoice/vae_encoder.rs:59–76  ·  view source on GitHub ↗
(&self, x: &Tensor)

Source from the content-addressed store, hash-verified

57 }
58
59 fn forward(&self, x: &Tensor) -> Result<Tensor> {
60 let channels = x.dim(1)?;
61
62 let h = self.backend.rms_norm_channel(x, &self.norm_weight, self.eps)?;
63 let zeros = Tensor::zeros((h.dim(0)?, channels, 6), h.dtype(), h.device())?;
64 let h = self.backend.depthwise_conv1d_bias_ctx(
65 &zeros, &h, &self.mixer_weight, &self.mixer_bias, 7, channels,
66 )?;
67 let x = self.backend.add_scaled(x, &h, &self.gamma)?;
68
69 let h = self.backend.rms_norm_channel(&x, &self.ffn_norm_weight, self.eps)?;
70 let h = h.transpose(1, 2)?;
71 let h = self.backend.linear_forward(&h, &self.ffn_linear1_weight, self.ffn_linear1_bias.as_ref())?;
72 let h = self.backend.gelu(&h)?;
73 let h = self.backend.linear_forward(&h, &self.ffn_linear2_weight, self.ffn_linear2_bias.as_ref())?;
74 let h = h.transpose(1, 2)?;
75 self.backend.add_scaled(&x, &h, &self.ffn_gamma)
76 }
77
78 fn forward_cached(
79 &self,

Callers 3

encodeMethod · 0.45

Calls 8

dtypeMethod · 0.80
rms_norm_channelMethod · 0.45
deviceMethod · 0.45
add_scaledMethod · 0.45
linear_forwardMethod · 0.45
geluMethod · 0.45
cloneMethod · 0.45

Tested by 2