| 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, |