1use catgrad::prelude::ops::*;
2use catgrad::prelude::*;
3
4use std::collections::HashMap;
5
6fn main() -> Result<(), Box<dyn std::error::Error>> {
8 let model = SimpleMNISTModel;
9
10 let typed_term = model.term().expect("Failed to create typed term");
12 save_svg(&typed_term.term, &format!("{}.svg", model.path()))?;
13
14 let parameters = load_param_types();
16
17 let mut env = stdlib();
19 env.declarations
20 .extend(to_load_ops(model.path(), parameters.keys()));
21
22 let check_result =
24 typecheck::check(&env, ¶meters, typed_term.clone()).expect("typecheck failed");
25
26 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 let backend = select_backend()?;
33
34 let results = run_interpreter(&backend, &typed_term, env)?;
36
37 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 let input_data: Vec<f32> = (0..2 * 28 * 28)
52 .map(|i| (i as f32 * 0.001) % 1.0) .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
72fn 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
97pub struct SimpleMNISTModel;
101
102impl Module<1, 1> for SimpleMNISTModel {
104 fn path(&self) -> Path {
107 Path::new(["model", "hidden"]).unwrap()
108 }
109
110 fn def(&self, builder: &Builder, [x]: [Var; 1]) -> [Var; 1] {
111 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 let x = matmul(builder, x, p);
123 let x = nn::Sigmoid.call(builder, [x]);
124
125 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 [x]
132 }
133
134 fn ty(&self) -> ([Type; 1], [Type; 1]) {
137 use catgrad::typecheck::*;
138
139 let batch_size = NatExpr::Var(0);
140
141 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 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
161fn 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 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 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
198fn 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 let layer1_data: Vec<f32> = (0..784 * 100)
207 .map(|i| (i as f32 * 0.01 % 2.0) - 1.0) .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 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#[cfg(test)]
279mod tests {
280 #[test]
281 fn main() {
282 super::main().unwrap();
283 }
284}