| 31 | |
| 32 | impl 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> { |