1use zyx::{DType, Tensor, ZyxError};
5use zyx_derive::Module;
6
7#[derive(Debug, Module)]
9#[cfg_attr(feature = "py", pyo3::pyclass(get_all, set_all))]
10pub struct Linear {
11 pub weight: Tensor,
13 pub bias: Option<Tensor>,
15}
16
17impl Linear {
18 pub fn new(
20 in_features: u64,
21 out_features: u64,
22 bias: bool,
23 dtype: DType,
24 ) -> Result<Linear, ZyxError> {
25 let l = -(1.0 / (in_features as f32)).sqrt();
26 let u = (1.0 / (in_features as f32)).sqrt();
27 Ok(Linear {
28 weight: Tensor::uniform([out_features, in_features], l..u)?.cast(dtype),
29 bias: if bias {
30 Some(Tensor::uniform([out_features], l..u)?.cast(dtype))
31 } else {
32 None
33 },
34 })
35 }
36
37 pub fn forward(&self, x: impl Into<Tensor>) -> Result<Tensor, ZyxError> {
40 let x = x.into().dot(self.weight.t())?;
41 if let Some(bias) = &self.bias {
42 return Ok(x + bias);
43 }
44 Ok(x)
45 }
46}
47
48#[test]
49fn linear() -> Result<(), ZyxError> {
50 let l0 = Linear::new(4, 16, true, DType::F32)?;
51 println!("{}\n{}", l0.weight, l0.bias.as_ref().unwrap());
52 let x = Tensor::randn([8, 4], DType::F32)?;
53 let y = l0.forward(x)?.relu();
54
55 println!("{y}");
56
57 Ok(())
58}