| 25 | |
| 26 | impl ConvolutionModule { |
| 27 | pub fn load(dim: usize, kernel_size: usize, vb: VarBuilder, backend: Arc<dyn ComputeBackend>) -> Result<Self> { |
| 28 | // in_proj projects to 2*dim for GLU gating |
| 29 | let in_proj_weight = vb.pp("in_proj").get((2 * dim, dim), "weight")?; |
| 30 | let in_proj_bias = Some(vb.pp("in_proj").get(2 * dim, "bias")?); |
| 31 | let depthwise_weight = vb.get((dim, 1, kernel_size), "depthwise_conv.weight")?; |
| 32 | let depthwise_bias = vb.get(dim, "depthwise_conv.bias")?; |
| 33 | let out_proj_weight = vb.pp("out_proj").get((dim, dim), "weight")?; |
| 34 | let out_proj_bias = Some(vb.pp("out_proj").get(dim, "bias")?); |
| 35 | Ok(Self { |
| 36 | in_proj_weight, |
| 37 | in_proj_bias, |
| 38 | depthwise_weight, |
| 39 | depthwise_bias, |
| 40 | out_proj_weight, |
| 41 | out_proj_bias, |
| 42 | kernel_size, |
| 43 | dim, |
| 44 | backend, |
| 45 | }) |
| 46 | } |
| 47 | |
| 48 | pub fn forward(&self, x: &Tensor) -> Result<Tensor> { |
| 49 | // x: [batch, seq, dim] |