Skip to main content

proof_engine/ml/
inference.rs

1//! Inference engine: run models, batch inference, ONNX loading, quantization.
2
3use super::tensor::Tensor;
4use super::model::*;
5use std::io::Read;
6use std::time::Instant;
7
8/// Compute device target.
9#[derive(Debug, Clone, Copy, PartialEq)]
10pub enum Device {
11    CPU,
12    GPUCompute,
13}
14
15/// Inference engine wrapping a model and target device.
16pub struct InferenceEngine {
17    pub model: Model,
18    pub device: Device,
19    pub stats: InferenceStats,
20}
21
22/// Statistics from a forward pass.
23#[derive(Debug, Clone, Default)]
24pub struct InferenceStats {
25    pub latency_ms: f64,
26    pub memory_bytes: usize,
27    pub flops: usize,
28}
29
30impl InferenceEngine {
31    pub fn new(model: Model, device: Device) -> Self {
32        Self {
33            model,
34            device,
35            stats: InferenceStats::default(),
36        }
37    }
38
39    /// Run a single inference pass.
40    pub fn infer(&mut self, input: &Tensor) -> Tensor {
41        let start = Instant::now();
42        let result = self.model.forward(input);
43        let elapsed = start.elapsed();
44        self.stats.latency_ms = elapsed.as_secs_f64() * 1000.0;
45        self.stats.memory_bytes = result.data.len() * 4 + input.data.len() * 4;
46        self.stats.flops = self.estimate_flops(input);
47        result
48    }
49
50    /// Batched inference: run each input through the model.
51    pub fn batch_infer(&mut self, inputs: &[Tensor]) -> Vec<Tensor> {
52        let start = Instant::now();
53        let results: Vec<Tensor> = inputs.iter().map(|inp| self.model.forward(inp)).collect();
54        let elapsed = start.elapsed();
55        self.stats.latency_ms = elapsed.as_secs_f64() * 1000.0;
56        self.stats.memory_bytes = results.iter().map(|r| r.data.len() * 4).sum::<usize>()
57            + inputs.iter().map(|i| i.data.len() * 4).sum::<usize>();
58        self.stats.flops = inputs.iter().map(|i| self.estimate_flops(i)).sum();
59        results
60    }
61
62    /// Warm up the inference pipeline by running dummy inputs.
63    pub fn warm_up(&mut self, input_shape: Vec<usize>, runs: usize) {
64        let dummy = Tensor::zeros(input_shape);
65        for _ in 0..runs {
66            let _ = self.model.forward(&dummy);
67        }
68    }
69
70    /// Rough FLOPs estimation based on layer types.
71    fn estimate_flops(&self, input: &Tensor) -> usize {
72        let mut flops = 0usize;
73        let mut current_size: usize = input.data.len();
74        for layer in &self.model.layers {
75            match layer {
76                Layer::Dense(d) => {
77                    let m = current_size / d.weights.shape[0];
78                    let k = d.weights.shape[0];
79                    let n = d.weights.shape[1];
80                    flops += 2 * m * k * n;
81                    current_size = m * n;
82                }
83                Layer::Conv2D(c) => {
84                    let c_out = c.filters.shape[0];
85                    let c_in = c.filters.shape[1];
86                    let kh = c.filters.shape[2];
87                    let kw = c.filters.shape[3];
88                    // rough: output_spatial * c_out * c_in * kh * kw * 2
89                    flops += current_size * c_out * kh * kw * 2 / c_in.max(1);
90                }
91                Layer::Attention(a) => {
92                    // Q,K,V projections + attention + output projection
93                    flops += 4 * a.d_model * a.d_model * 2;
94                }
95                _ => {
96                    // element-wise ops: ~N flops
97                    flops += current_size;
98                }
99            }
100        }
101        flops
102    }
103}
104
105// ── ONNX Loader ─────────────────────────────────────────────────────────
106
107/// Supported ONNX operation types (simplified).
108#[derive(Debug, Clone)]
109enum OnnxOp {
110    Gemm { transA: bool, transB: bool, alpha: f32, beta: f32 },
111    Conv { strides: Vec<usize>, pads: Vec<usize> },
112    Relu,
113    MaxPool { kernel_shape: Vec<usize>, strides: Vec<usize> },
114    BatchNorm { eps: f32 },
115    Reshape,
116    Softmax { axis: i32 },
117    Add,
118    Mul,
119}
120
121/// Minimal ONNX-like graph node.
122#[derive(Debug, Clone)]
123struct OnnxNode {
124    op: OnnxOp,
125    inputs: Vec<String>,
126    outputs: Vec<String>,
127}
128
129/// ONNX model loader.
130pub struct OnnxLoader;
131
132impl OnnxLoader {
133    /// Load a model from this engine's own simple binary layout.
134    ///
135    /// Despite the name and the `ONNX` magic bytes, this is **not** the ONNX
136    /// protobuf format: real `.onnx` files from PyTorch or ONNX Runtime will
137    /// be rejected. The layout:
138    /// - magic: b"ONNX" (4 bytes)
139    /// - num_nodes: u32 LE
140    /// - For each node:
141    ///   - op_type: u8 (0=Gemm,1=Conv,2=Relu,3=MaxPool,4=BatchNorm,5=Reshape,6=Softmax,7=Add,8=Mul)
142    ///   - num_weights: u32 LE
143    ///   - For each weight tensor: [ndim: u32] [shape...] [data as f32 LE]
144    pub fn load_onnx(path: &str) -> Result<Model, String> {
145        let mut file = std::fs::File::open(path).map_err(|e| format!("cannot open {path}: {e}"))?;
146        let mut buf4 = [0u8; 4];
147        let mut buf1 = [0u8; 1];
148
149        // magic
150        file.read_exact(&mut buf4).map_err(|e| e.to_string())?;
151        if &buf4 != b"ONNX" {
152            return Err("invalid ONNX magic".into());
153        }
154
155        // num_nodes
156        file.read_exact(&mut buf4).map_err(|e| e.to_string())?;
157        let num_nodes = u32::from_le_bytes(buf4) as usize;
158
159        let mut layers = Vec::new();
160
161        for _ in 0..num_nodes {
162            file.read_exact(&mut buf1).map_err(|e| e.to_string())?;
163            let op_type = buf1[0];
164
165            file.read_exact(&mut buf4).map_err(|e| e.to_string())?;
166            let num_weights = u32::from_le_bytes(buf4) as usize;
167
168            let mut tensors = Vec::new();
169            for _ in 0..num_weights {
170                file.read_exact(&mut buf4).map_err(|e| e.to_string())?;
171                let ndim = u32::from_le_bytes(buf4) as usize;
172                let mut shape = Vec::with_capacity(ndim);
173                for _ in 0..ndim {
174                    file.read_exact(&mut buf4).map_err(|e| e.to_string())?;
175                    shape.push(u32::from_le_bytes(buf4) as usize);
176                }
177                let n: usize = shape.iter().product();
178                let mut data = Vec::with_capacity(n);
179                for _ in 0..n {
180                    file.read_exact(&mut buf4).map_err(|e| e.to_string())?;
181                    data.push(f32::from_le_bytes(buf4));
182                }
183                tensors.push(Tensor { shape, data });
184            }
185
186            let layer = match op_type {
187                0 => {
188                    // Gemm -> Dense
189                    if tensors.len() >= 2 {
190                        Layer::Dense(DenseLayer {
191                            weights: tensors[0].clone(),
192                            bias: tensors[1].clone(),
193                        })
194                    } else {
195                        return Err("Gemm requires 2 weight tensors".into());
196                    }
197                }
198                1 => {
199                    // Conv
200                    if tensors.len() >= 2 {
201                        Layer::Conv2D(Conv2DLayer {
202                            filters: tensors[0].clone(),
203                            bias: tensors[1].clone(),
204                            stride: 1,
205                            padding: 0,
206                        })
207                    } else {
208                        return Err("Conv requires 2 weight tensors".into());
209                    }
210                }
211                2 => Layer::ReLU,
212                3 => Layer::MaxPool(MaxPoolLayer { kernel_size: 2, stride: 2 }),
213                4 => {
214                    // BatchNorm
215                    if tensors.len() >= 4 {
216                        Layer::BatchNorm(BatchNormLayer {
217                            gamma: tensors[0].clone(),
218                            beta: tensors[1].clone(),
219                            running_mean: tensors[2].clone(),
220                            running_var: tensors[3].clone(),
221                            eps: 1e-5,
222                        })
223                    } else {
224                        return Err("BatchNorm requires 4 tensors".into());
225                    }
226                }
227                5 => Layer::Flatten, // Reshape treated as flatten
228                6 => Layer::Softmax(0),
229                7 | 8 => {
230                    // Element-wise Add / Mul need graph connections this
231                    // sequential model does not have. They used to be loaded
232                    // as ReLU, which silently zeroed negative values.
233                    return Err(format!(
234                        "op type {op_type} ({}) is not supported by the sequential loader",
235                        if op_type == 7 { "Add" } else { "Mul" }
236                    ));
237                }
238                _ => return Err(format!("unknown op type {op_type}")),
239            };
240            layers.push(layer);
241        }
242
243        Ok(Model { layers, name: "onnx_model".to_string() })
244    }
245
246    /// Write a model in our simplified ONNX binary format.
247    pub fn save_onnx(model: &Model, path: &str) -> Result<(), String> {
248        use std::io::Write;
249        let mut file = std::fs::File::create(path).map_err(|e| e.to_string())?;
250        file.write_all(b"ONNX").map_err(|e| e.to_string())?;
251        let num_nodes = model.layers.len() as u32;
252        file.write_all(&num_nodes.to_le_bytes()).map_err(|e| e.to_string())?;
253
254        for layer in &model.layers {
255            let (op_type, tensors): (u8, Vec<&Tensor>) = match layer {
256                Layer::Dense(l) => (0, vec![&l.weights, &l.bias]),
257                Layer::Conv2D(l) => (1, vec![&l.filters, &l.bias]),
258                Layer::ReLU => (2, vec![]),
259                Layer::MaxPool(_) => (3, vec![]),
260                Layer::BatchNorm(l) => (4, vec![&l.gamma, &l.beta, &l.running_mean, &l.running_var]),
261                Layer::Flatten => (5, vec![]),
262                Layer::Softmax(_) => (6, vec![]),
263                _ => (2, vec![]), // default to relu-like
264            };
265            file.write_all(&[op_type]).map_err(|e| e.to_string())?;
266            let nw = tensors.len() as u32;
267            file.write_all(&nw.to_le_bytes()).map_err(|e| e.to_string())?;
268            for t in tensors {
269                let ndim = t.shape.len() as u32;
270                file.write_all(&ndim.to_le_bytes()).map_err(|e| e.to_string())?;
271                for &d in &t.shape {
272                    file.write_all(&(d as u32).to_le_bytes()).map_err(|e| e.to_string())?;
273                }
274                for &v in &t.data {
275                    file.write_all(&v.to_le_bytes()).map_err(|e| e.to_string())?;
276                }
277            }
278        }
279        Ok(())
280    }
281}
282
283// ── Quantization ────────────────────────────────────────────────────────
284
285/// Simple weight quantization: clamp weights to int8 range then dequantize.
286/// This simulates the effect of lower-precision storage.
287pub fn quantize_model(model: &Model, bits: u32) -> Model {
288    let max_val = (1 << (bits - 1)) as f32 - 1.0;
289    let min_val = -max_val - 1.0;
290
291    let mut new_layers = Vec::new();
292    for layer in &model.layers {
293        let new_layer = match layer {
294            Layer::Dense(l) => {
295                let (qw, qb) = (quantize_tensor(&l.weights, min_val, max_val),
296                                 quantize_tensor(&l.bias, min_val, max_val));
297                Layer::Dense(DenseLayer { weights: qw, bias: qb })
298            }
299            Layer::Conv2D(l) => {
300                let qf = quantize_tensor(&l.filters, min_val, max_val);
301                let qb = quantize_tensor(&l.bias, min_val, max_val);
302                Layer::Conv2D(Conv2DLayer { filters: qf, bias: qb, stride: l.stride, padding: l.padding })
303            }
304            other => other.clone(),
305        };
306        new_layers.push(new_layer);
307    }
308    Model { layers: new_layers, name: format!("{}_q{}", model.name, bits) }
309}
310
311fn quantize_tensor(t: &Tensor, min_val: f32, max_val: f32) -> Tensor {
312    let abs_max = t.data.iter().map(|v| v.abs()).fold(0.0f32, f32::max);
313    if abs_max == 0.0 {
314        return t.clone();
315    }
316    let scale = max_val / abs_max;
317    let inv_scale = abs_max / max_val;
318    let data: Vec<f32> = t.data.iter().map(|&v| {
319        let q = (v * scale).round().clamp(min_val, max_val);
320        q * inv_scale
321    }).collect();
322    Tensor { shape: t.shape.clone(), data }
323}
324
325#[cfg(test)]
326mod tests {
327    use super::*;
328
329    #[test]
330    fn test_infer() {
331        let model = Sequential::new("test")
332            .dense(4, 3)
333            .relu()
334            .build();
335        let mut engine = InferenceEngine::new(model, Device::CPU);
336        let input = Tensor::ones(vec![1, 4]);
337        let out = engine.infer(&input);
338        assert_eq!(out.shape, vec![1, 3]);
339        assert!(engine.stats.latency_ms >= 0.0);
340    }
341
342    #[test]
343    fn test_batch_infer() {
344        let model = Sequential::new("test")
345            .dense(3, 2)
346            .build();
347        let mut engine = InferenceEngine::new(model, Device::CPU);
348        let inputs = vec![
349            Tensor::ones(vec![1, 3]),
350            Tensor::zeros(vec![1, 3]),
351        ];
352        let outputs = engine.batch_infer(&inputs);
353        assert_eq!(outputs.len(), 2);
354        assert_eq!(outputs[0].shape, vec![1, 2]);
355        assert_eq!(outputs[1].shape, vec![1, 2]);
356    }
357
358    #[test]
359    fn test_warm_up() {
360        let model = Sequential::new("test").dense(4, 2).build();
361        let mut engine = InferenceEngine::new(model, Device::CPU);
362        engine.warm_up(vec![1, 4], 5);
363        // just verify it doesn't panic
364    }
365
366    #[test]
367    fn test_quantize_model() {
368        let model = Sequential::new("test")
369            .dense(4, 3)
370            .relu()
371            .build();
372        let qmodel = quantize_model(&model, 8);
373        assert!(qmodel.name.contains("q8"));
374        // forward still works
375        let input = Tensor::ones(vec![1, 4]);
376        let out = qmodel.forward(&input);
377        assert_eq!(out.shape, vec![1, 3]);
378    }
379
380    #[test]
381    fn test_onnx_save_load_roundtrip() {
382        let model = Sequential::new("onnx_test")
383            .dense(4, 3)
384            .relu()
385            .dense(3, 2)
386            .softmax()
387            .build();
388
389        let path = std::env::temp_dir().join("proof_engine_test.onnx");
390        let path_str = path.to_str().unwrap();
391
392        OnnxLoader::save_onnx(&model, path_str).unwrap();
393        let loaded = OnnxLoader::load_onnx(path_str).unwrap();
394
395        assert_eq!(loaded.layers.len(), model.layers.len());
396
397        // Verify dense weights match
398        if let (Layer::Dense(orig), Layer::Dense(loaded_l)) = (&model.layers[0], &loaded.layers[0]) {
399            assert_eq!(orig.weights.data, loaded_l.weights.data);
400        }
401
402        let _ = std::fs::remove_file(path);
403    }
404
405    #[test]
406    fn test_add_mul_ops_are_rejected_not_loaded_as_relu() {
407        for op in [7u8, 8] {
408            let path = std::env::temp_dir().join(format!("proof_engine_op{op}.onnx"));
409            let mut bytes = b"ONNX".to_vec();
410            bytes.extend_from_slice(&1u32.to_le_bytes()); // one node
411            bytes.push(op);
412            bytes.extend_from_slice(&0u32.to_le_bytes()); // no weights
413            std::fs::write(&path, &bytes).unwrap();
414            let result = OnnxLoader::load_onnx(path.to_str().unwrap());
415            assert!(result.is_err(), "op {op} loaded: {:?}", result.map(|m| m.layers.len()));
416            let _ = std::fs::remove_file(path);
417        }
418    }
419
420    #[test]
421    fn test_onnx_load_bad_magic() {
422        let path = std::env::temp_dir().join("proof_engine_bad.onnx");
423        std::fs::write(&path, b"NOPE1234").unwrap();
424        let result = OnnxLoader::load_onnx(path.to_str().unwrap());
425        assert!(result.is_err());
426        let _ = std::fs::remove_file(path);
427    }
428
429    #[test]
430    fn test_inference_stats() {
431        let model = Sequential::new("s").dense(2, 2).build();
432        let mut engine = InferenceEngine::new(model, Device::CPU);
433        let _ = engine.infer(&Tensor::ones(vec![1, 2]));
434        assert!(engine.stats.flops > 0);
435        assert!(engine.stats.memory_bytes > 0);
436    }
437}