Skip to main content

hidden/
hidden.rs

1use catgrad::prelude::ops::*;
2use catgrad::prelude::*;
3
4use std::collections::HashMap;
5
6/// Construct, shapecheck, and interpret the `SimpleMNISTModel` using the ndarray backend.
7fn main() -> Result<(), Box<dyn std::error::Error>> {
8    let model = SimpleMNISTModel;
9
10    // Get the model as a typed term
11    let typed_term = model.term().expect("Failed to create typed term");
12    save_svg(&typed_term.term, &format!("{}.svg", model.path()))?;
13
14    // Create parameters for the model
15    let parameters = load_param_types();
16
17    // Get stdlib environment and extend with parameter declarations
18    let mut env = stdlib();
19    env.declarations
20        .extend(to_load_ops(model.path(), parameters.keys()));
21
22    // Shapecheck the model
23    let check_result =
24        typecheck::check(&env, &parameters, typed_term.clone()).expect("typecheck failed");
25
26    // Diagram of term with shapes inferred
27    let labeled_term = typed_term.term.clone().with_nodes(|_| check_result);
28    let filename = &format!("{}_typed.svg", model.path());
29    save_svg(&labeled_term.unwrap(), filename)?;
30
31    // Choose a backend from available features
32    let backend = select_backend()?;
33
34    // Run the interpreter with the selected backend
35    let results = run_interpreter(&backend, &typed_term, env)?;
36
37    // Print the `Value`s returned by the interpreter.
38    for value in results {
39        println!("{value:?}");
40    }
41
42    Ok(())
43}
44
45fn run_interpreter<B: interpreter::Backend>(
46    backend: &B,
47    typed_term: &TypedTerm,
48    env: Environment,
49) -> Result<Vec<interpreter::Value<B>>, Box<dyn std::error::Error>> {
50    // Create sample input data: batch of 2 MNIST-like images (28x28)
51    let input_data: Vec<f32> = (0..2 * 28 * 28)
52        .map(|i| (i as f32 * 0.001) % 1.0) // Simple pattern: values between 0 and 1
53        .collect();
54
55    let interpreter_params = load_param_data(backend);
56    let interpreter = interpreter::Interpreter::new(backend.clone(), env, interpreter_params);
57
58    let input_tensor = interpreter::tensor(
59        &interpreter.backend,
60        interpreter::Shape(vec![2, 28, 28]),
61        &input_data,
62    )
63    .expect("Failed to create input tensor");
64
65    let results = interpreter
66        .run(typed_term.term.clone(), vec![input_tensor])
67        .expect("Failed to run inference");
68
69    Ok(results)
70}
71
72/// Pick a backend depending on what features are available
73fn select_backend() -> Result<impl interpreter::Backend, Box<dyn std::error::Error>> {
74    #[cfg(feature = "candle-backend")]
75    {
76        println!("selected candle backend...");
77        use catgrad::interpreter::backend::candle::CandleBackend;
78        #[allow(clippy::needless_return)]
79        return Ok(CandleBackend::new());
80    }
81
82    #[cfg(all(feature = "ndarray-backend", not(feature = "candle-backend")))]
83    {
84        println!("selected ndarray backend...");
85        use catgrad::interpreter::backend::ndarray::NdArrayBackend;
86        #[allow(clippy::needless_return)]
87        return Ok(NdArrayBackend);
88    }
89
90    #[cfg(not(any(feature = "candle-backend", feature = "ndarray-backend")))]
91    {
92        println!("selected ShapeOnly backend (no tensors computed)");
93        return Ok(interpreter::backend::shape_only::ShapeOnlyBackend);
94    }
95}
96
97////////////////////////////////////////////////////////////////////////////////
98// Define the SimpleMNISTModel model
99
100pub struct SimpleMNISTModel;
101
102// Implement `Def`: this is like torch's `Module`.
103impl Module<1, 1> for SimpleMNISTModel {
104    // Model name
105    // TODO: NOTE: it's not clear how user is supposed to know how to choose this name!
106    fn path(&self) -> Path {
107        Path::new(["model", "hidden"]).unwrap()
108    }
109
110    fn def(&self, builder: &Builder, [x]: [Var; 1]) -> [Var; 1] {
111        // Flatten input from B×28×28 to B×784
112        let [batch_size, h, w] = unpack::<3>(builder, shape(builder, x.clone()));
113        let flat_size = h * w;
114        let flat_shape = pack::<2>(builder, [batch_size, flat_size]);
115        let x = reshape(builder, flat_shape, x);
116
117        let root = self.path();
118
119        let p = param(builder, &root.extend(["0", "weights"]).unwrap());
120
121        // layer 1: B×784 @ 784×100 = B×100
122        let x = matmul(builder, x, p);
123        let x = nn::Sigmoid.call(builder, [x]);
124
125        // layer 2: B×100 @ 100×10 = B×10
126        let p = param(builder, &root.extend(["1", "weights"]).unwrap());
127        let x = matmul(builder, x, p);
128        let x = nn::Sigmoid.call(builder, [x]);
129
130        // result
131        [x]
132    }
133
134    // This should return the *detailed* type of the model
135    // TODO: NOTE: API for writing types is still WIP. Lots of boilerplate here!
136    fn ty(&self) -> ([Type; 1], [Type; 1]) {
137        use catgrad::typecheck::*;
138
139        let batch_size = NatExpr::Var(0);
140
141        // Input shape B×28×28
142        let t_x = Value::Tensor(TypeExpr::NdArrayType(NdArrayType {
143            dtype: DtypeExpr::Constant(Dtype::F32),
144            shape: ShapeExpr::Shape(vec![
145                batch_size.clone(),
146                NatExpr::Constant(28),
147                NatExpr::Constant(28),
148            ]),
149        }));
150
151        // Output shape B×10
152        let t_y = Value::Tensor(TypeExpr::NdArrayType(NdArrayType {
153            dtype: DtypeExpr::Constant(Dtype::F32),
154            shape: ShapeExpr::Shape(vec![batch_size, NatExpr::Constant(10)]),
155        }));
156
157        ([t_x], [t_y])
158    }
159}
160
161////////////////////////////////////////////////////////////////////////////////
162// Parameter loading boilerplate
163// NOTE: in reality, this would be done by loading e.g. a safetensors file.
164
165// NOTE: you would normally create this data by reading the safetensors file!
166fn load_param_types() -> typecheck::Parameters {
167    use catgrad::category::core::Dtype;
168    use catgrad::typecheck::value_types::{DtypeExpr, NatExpr, NdArrayType, ShapeExpr, TypeExpr};
169
170    let mut map = HashMap::new();
171
172    // Layer 1: (28*28) → 100
173    let layer1_type = Value::Tensor(TypeExpr::NdArrayType(NdArrayType {
174        dtype: DtypeExpr::Constant(Dtype::F32),
175        shape: ShapeExpr::Shape(vec![
176            NatExpr::Mul(vec![NatExpr::Constant(28), NatExpr::Constant(28)]),
177            NatExpr::Constant(100),
178        ]),
179    }));
180    map.insert(
181        path(vec!["0", "weights"]).expect("invalid param path"),
182        layer1_type,
183    );
184
185    // Layer 2: 100 → 10
186    let layer2_type = Value::Tensor(TypeExpr::NdArrayType(NdArrayType {
187        dtype: DtypeExpr::Constant(Dtype::F32),
188        shape: ShapeExpr::Shape(vec![NatExpr::Constant(100), NatExpr::Constant(10)]),
189    }));
190    map.insert(
191        path(vec!["1", "weights"]).expect("invalid param path"),
192        layer2_type,
193    );
194
195    typecheck::Parameters::from(map)
196}
197
198// NOTE: you would normally create this data by reading the safetensors file!
199fn load_param_data<B: interpreter::Backend>(backend: &B) -> interpreter::Parameters<B> {
200    use catgrad::category::core::Shape;
201    use std::collections::HashMap;
202
203    let mut map = HashMap::new();
204
205    // Layer 1 weights: [784, 100] - initialize with small random-ish values
206    let layer1_data: Vec<f32> = (0..784 * 100)
207        .map(|i| (i as f32 * 0.01 % 2.0) - 1.0) // Simple pattern: values between -1 and 1
208        .collect();
209    let layer1_tensor =
210        interpreter::TaggedTensor::from_slice(backend, &layer1_data, Shape(vec![784, 100]))
211            .expect("Failed to create layer1 tensor");
212    map.insert(
213        path(vec!["0", "weights"]).expect("invalid param path"),
214        layer1_tensor,
215    );
216
217    // Layer 2 weights: [100, 10]
218    let layer2_data: Vec<f32> = (0..100 * 10)
219        .map(|i| (i as f32 * 0.01 % 2.0) - 1.0)
220        .collect();
221    let layer2_tensor =
222        interpreter::TaggedTensor::from_slice(backend, &layer2_data, Shape(vec![100, 10]))
223            .expect("Failed to create layer2 tensor");
224    map.insert(
225        path(vec!["1", "weights"]).expect("invalid param path"),
226        layer2_tensor,
227    );
228
229    interpreter::Parameters::from(map)
230}
231
232#[cfg(feature = "svg")]
233pub fn save_svg<
234    O: PartialEq + Clone + std::fmt::Display + std::fmt::Debug,
235    A: PartialEq + Clone + std::fmt::Display + std::fmt::Debug,
236>(
237    term: &open_hypergraphs::lax::OpenHypergraph<O, A>,
238    filename: &str,
239) -> Result<(), std::io::Error> {
240    use catgrad::svg::to_svg;
241    let bytes = match to_svg(term) {
242        Ok(bytes) => bytes,
243        Err(e) => {
244            eprintln!("Failed to generate SVG: {e}");
245            return Ok(());
246        }
247    };
248
249    let output_dir = std::path::Path::new(env!("CARGO_MANIFEST_DIR"))
250        .join("examples")
251        .join("images");
252
253    if let Err(e) = std::fs::create_dir_all(&output_dir) {
254        eprintln!("Failed to create directory {output_dir:?}: {e}");
255        return Ok(());
256    }
257
258    let output_path = output_dir.join(filename);
259    println!("saving svg to {output_path:?}");
260
261    if let Err(e) = std::fs::write(&output_path, bytes) {
262        eprintln!("Failed to write SVG file {output_path:?}: {e}");
263    }
264
265    Ok(())
266}
267
268#[cfg(not(feature = "svg"))]
269pub fn save_svg<O, A>(
270    _term: &open_hypergraphs::lax::OpenHypergraph<O, A>,
271    _filename: &str,
272) -> Result<(), std::io::Error> {
273    println!("SVG feature not enabled, skipping diagram generation");
274    Ok(())
275}
276
277// include this as a test
278#[cfg(test)]
279mod tests {
280    #[test]
281    fn main() {
282        super::main().unwrap();
283    }
284}