1use super::tensor::Tensor;
4use super::model::*;
5use std::io::Read;
6use std::time::Instant;
7
8#[derive(Debug, Clone, Copy, PartialEq)]
10pub enum Device {
11 CPU,
12 GPUCompute,
13}
14
15pub struct InferenceEngine {
17 pub model: Model,
18 pub device: Device,
19 pub stats: InferenceStats,
20}
21
22#[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 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 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 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 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 flops += current_size * c_out * kh * kw * 2 / c_in.max(1);
90 }
91 Layer::Attention(a) => {
92 flops += 4 * a.d_model * a.d_model * 2;
94 }
95 _ => {
96 flops += current_size;
98 }
99 }
100 }
101 flops
102 }
103}
104
105#[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#[derive(Debug, Clone)]
123struct OnnxNode {
124 op: OnnxOp,
125 inputs: Vec<String>,
126 outputs: Vec<String>,
127}
128
129pub struct OnnxLoader;
131
132impl OnnxLoader {
133 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 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 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 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 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 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, 6 => Layer::Softmax(0),
229 7 | 8 => {
230 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 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![]), };
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
283pub 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 }
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 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 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()); bytes.push(op);
412 bytes.extend_from_slice(&0u32.to_le_bytes()); 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}