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

Method load

cake-core/src/models/luxtts/text_encoder.rs:33–77  ·  view source on GitHub ↗
(config: &LuxTTSConfig, embed_vb: VarBuilder, enc_vb: VarBuilder, backend: Arc<dyn crate::backends::ComputeBackend>)

Source from the content-addressed store, hash-verified

31
32impl TextEncoder {
33 pub fn load(config: &LuxTTSConfig, embed_vb: VarBuilder, enc_vb: VarBuilder, backend: Arc<dyn crate::backends::ComputeBackend>) -> Result<Self> {
34 let m = &config.model;
35 let dim = m.text_encoder_dim;
36
37 // embed.weight is at top level
38 let embed_weight = embed_vb.pp("embed").get((m.vocab_size, dim), "weight")?;
39
40 // text_encoder.in_proj
41 let in_proj_weight = enc_vb.pp("in_proj").get((dim, dim), "weight")?;
42 let in_proj_bias = Some(enc_vb.pp("in_proj").get(dim, "bias")?);
43
44 // text_encoder.layers
45 let mut layers = Vec::new();
46 for i in 0..m.text_encoder_num_layers {
47 let layer = ZipformerEncoderLayer::load(
48 dim,
49 m.text_encoder_feedforward_dim,
50 m.text_encoder_num_heads,
51 m.query_head_dim,
52 m.value_head_dim,
53 m.pos_dim,
54 m.pos_head_dim,
55 m.text_encoder_cnn_module_kernel,
56 enc_vb.pp(format!("layers.{i}")),
57 backend.clone(),
58 )?;
59 layers.push(layer);
60 }
61
62 // text_encoder.out_proj
63 let out_proj_weight = enc_vb.pp("out_proj").get((m.feat_dim, dim), "weight")?;
64 let out_proj_bias = Some(enc_vb.pp("out_proj").get(m.feat_dim, "bias")?);
65
66 Ok(Self {
67 embed_weight,
68 in_proj_weight,
69 in_proj_bias,
70 layers,
71 out_proj_weight,
72 out_proj_bias,
73 dim,
74 pos_dim: m.pos_dim,
75 backend,
76 })
77 }
78
79 /// Forward pass: token_ids [batch, seq] -> [batch, seq, feat_dim].
80 pub fn forward(&self, token_ids: &Tensor) -> Result<Tensor> {

Callers

nothing calls this directly

Calls 3

getMethod · 0.45
cloneMethod · 0.45
pushMethod · 0.45

Tested by

no test coverage detected