| 39 | |
| 40 | impl ConvNeXtBlock { |
| 41 | fn load(dim: usize, ff_mult: usize, kernel_size: usize, vb: VarBuilder, backend: Arc<dyn ComputeBackend>) -> Result<Self> { |
| 42 | let ff_dim = dim * ff_mult; |
| 43 | let dwconv_weight = vb.get((dim, 1, kernel_size), "dwconv.weight")?; |
| 44 | let dwconv_bias = vb.get(dim, "dwconv.bias")?; |
| 45 | let gamma = vb.get(dim, "gamma")?; |
| 46 | let norm_weight = vb.get(dim, "norm.weight")?; |
| 47 | let norm_bias = vb.get(dim, "norm.bias")?; |
| 48 | let pwconv1_weight = vb.pp("pwconv1").get((ff_dim, dim), "weight")?; |
| 49 | let pwconv1_bias = Some(vb.pp("pwconv1").get(ff_dim, "bias")?); |
| 50 | let pwconv2_weight = vb.pp("pwconv2").get((dim, ff_dim), "weight")?; |
| 51 | let pwconv2_bias = Some(vb.pp("pwconv2").get(dim, "bias")?); |
| 52 | Ok(Self { |
| 53 | dwconv_weight, |
| 54 | dwconv_bias, |
| 55 | gamma, |
| 56 | norm_weight, |
| 57 | norm_bias, |
| 58 | pwconv1_weight, |
| 59 | pwconv1_bias, |
| 60 | pwconv2_weight, |
| 61 | pwconv2_bias, |
| 62 | kernel_size, |
| 63 | dim, |
| 64 | backend, |
| 65 | }) |
| 66 | } |
| 67 | |
| 68 | fn forward(&self, x: &Tensor) -> Result<Tensor> { |
| 69 | // x: [batch, dim, seq] |