From b1ebe12011fd8862ac3bbf5cea7ca84fc1363867 Mon Sep 17 00:00:00 2001 From: Asjad Date: Sat, 1 Aug 2026 11:53:12 +0500 Subject: [PATCH 1/3] feat: Conv2d/BatchNorm/Pool ops + GitHub Actions CI - interpreter: add Conv2d, fused BatchNorm, MaxPool2d/AveragePool2d/GlobalAveragePool - ci: add build/lint/python-tests GitHub Actions workflow - tests: fail with nonzero exit code on check failure instead of always passing --- .github/workflows/ci.yml | 80 ++++++++++++++++++++++++++++++ .gitignore | 3 ++ src/interpreter/context.rs | 29 +++++++++-- src/interpreter/mod.rs | 98 ++++++++++++++++++++++++++++++------- src/interpreter/ops/conv.rs | 77 +++++++++++++++++++++++++++++ src/interpreter/ops/mod.rs | 2 + src/interpreter/ops/pool.rs | 52 ++++++++++++++++++++ test.py | 13 +++++ tests/compare_ort.py | 10 ++++ tests/test_bridge.py | 10 ++++ 10 files changed, 353 insertions(+), 21 deletions(-) create mode 100644 .github/workflows/ci.yml create mode 100644 src/interpreter/ops/conv.rs create mode 100644 src/interpreter/ops/pool.rs create mode 100644 test.py diff --git a/.github/workflows/ci.yml b/.github/workflows/ci.yml new file mode 100644 index 0000000..c77369c --- /dev/null +++ b/.github/workflows/ci.yml @@ -0,0 +1,80 @@ +name: CI + +on: + push: + branches: [main] + pull_request: + branches: [main] + +env: + CARGO_TERM_COLOR: always + +jobs: + build: + name: cargo build + runs-on: ubuntu-latest + steps: + - uses: actions/checkout@v4 + + - name: Install Rust + uses: dtolnay/rust-toolchain@stable + + - name: Cache cargo + uses: Swatinem/rust-cache@v2 + + - name: cargo build + run: cargo build --release + + lint: + name: cargo fmt / clippy + runs-on: ubuntu-latest + continue-on-error: true + steps: + - uses: actions/checkout@v4 + + - name: Install Rust + uses: dtolnay/rust-toolchain@stable + with: + components: rustfmt, clippy + + - name: Cache cargo + uses: Swatinem/rust-cache@v2 + + - name: cargo fmt --check + run: cargo fmt --check + + - name: cargo clippy + run: cargo clippy --release --all-targets -- -D warnings + + python-tests: + name: python tests (maturin) + runs-on: ubuntu-latest + needs: build + steps: + - uses: actions/checkout@v4 + + - name: Install Rust + uses: dtolnay/rust-toolchain@stable + + - name: Cache cargo + uses: Swatinem/rust-cache@v2 + + - name: Set up Python + uses: actions/setup-python@v5 + with: + python-version: "3.12" + + - name: Install Python deps + run: pip install maturin numpy onnx onnxruntime + + - name: Build extension + run: maturin develop --release + + - name: Generate test fixtures + run: python tests/make_test_model.py + + - name: Run tensor bridge tests + run: python tests/test_bridge.py + + - name: Run ORT comparison tests + run: python tests/compare_ort.py diff --git a/.gitignore b/.gitignore index 82c91e7..0a45f9b 100644 --- a/.gitignore +++ b/.gitignore @@ -6,3 +6,6 @@ dist/ __pycache__/ .venv/ Cargo.lock +.claude/ +tests/*.onnx +tests/*.npy diff --git a/src/interpreter/context.rs b/src/interpreter/context.rs index 6d94fa7..545e6a0 100644 --- a/src/interpreter/context.rs +++ b/src/interpreter/context.rs @@ -2,12 +2,14 @@ // weights (model params) live in the OnnxModel and are referenced here to avoid cloning. use std::collections::HashMap; -use crate::tensor::FloatPrim; + +use burn_backend::backend::ops::FloatTensorOps; +use onnx_ir::ir::{Argument, ValueSource}; + +use crate::tensor::{default_device, FloatPrim, B}; pub struct ExecutionContext<'w> { - // intermediate tensors produced during the forward pass tensors: HashMap, - // read-only reference to pre-loaded model weights weights: &'w HashMap, } @@ -25,4 +27,25 @@ impl<'w> ExecutionContext<'w> { pub fn insert(&mut self, name: String, tensor: FloatPrim) { self.tensors.insert(name, tensor); } + + /// Resolve any Argument to a FloatPrim: + /// - Dynamic/named Constant → look up by name + /// - Static with empty name → convert inline TensorData directly + /// - Optional/missing → None + pub fn resolve(&self, arg: &Argument) -> Option { + match arg.value_source { + ValueSource::Dynamic | ValueSource::Constant => self.get(&arg.name), + ValueSource::Static(_) => { + if !arg.name.is_empty() { + // pre-loaded by name + if let Some(t) = self.get(&arg.name) { + return Some(t); + } + } + // anonymous inline static — convert on the fly + arg.value().map(|d| B::float_from_data(d, &default_device())) + } + ValueSource::Optional => None, + } + } } diff --git a/src/interpreter/mod.rs b/src/interpreter/mod.rs index cab3175..d765f61 100644 --- a/src/interpreter/mod.rs +++ b/src/interpreter/mod.rs @@ -1,14 +1,15 @@ // ONNX runtime interpreter: -// parse -> pre-load weights -> dispatch nodes on each forward pass +// parse -> pre-load weights (with BN fusion) -> dispatch nodes on each forward pass mod context; pub mod ops; use std::collections::HashMap; -use burn_backend::{backend::ops::FloatTensorOps, DType}; +use burn_backend::{backend::ops::FloatTensorOps, DType, TensorData, TensorMetadata}; use onnx_ir::{ ir::{Argument, ValueSource}, + batch_norm::{BatchNormConfig, BatchNormalizationNode}, Node, OnnxGraphBuilder, }; @@ -34,6 +35,50 @@ fn load_weight(arg: &Argument) -> Option<(String, FloatPrim)> { } } +/// Pre-compute BN scale/bias as [1,C,1,1] tensors so runtime BN is just 2 elementwise ops. +/// key: "{node_name}::scale" and "{node_name}::bias" +/// inputs: [x, gamma, beta, running_mean, running_var] +fn precompute_bn(node: &BatchNormalizationNode, weights: &HashMap) + -> Option<(String, FloatPrim, FloatPrim)> +{ + let eps = match &node.config { + BatchNormConfig::Static(c) => c.epsilon as f32, + BatchNormConfig::Runtime(c) => c.epsilon as f32, + }; + + // all params must be pre-loaded (static/constant) + let get = |arg: &Argument| -> Option { + if !arg.name.is_empty() { + weights.get(&arg.name).cloned() + } else { + arg.value().map(|d| B::float_from_data(d, &default_device())) + } + }; + + let gamma = get(&node.inputs[1])?; + let beta = get(&node.inputs[2])?; + let mean = get(&node.inputs[3])?; + let var = get(&node.inputs[4])?; + + let c = gamma.shape().iter().next().copied().unwrap_or(1); + + // scale = gamma / sqrt(var + eps), shaped [1, C, 1, 1] + let eps_t = B::float_from_data(TensorData::from([eps]), &default_device()); + let scale_flat = B::float_div( + gamma, + B::float_sqrt(B::float_add(var, eps_t)), + ); + // offset = beta - mean * scale, shaped [1, C, 1, 1] + let offset_flat = B::float_sub(beta, B::float_mul(mean, scale_flat.clone())); + + // pre-expand to [1, C, 1, 1] so no reshape at inference time + let shape_4d: Vec = vec![1, c, 1, 1]; + let scale = B::float_reshape(scale_flat, shape_4d.clone().into()); + let offset = B::float_reshape(offset_flat, shape_4d.into()); + + Some((node.name.clone(), scale, offset)) +} + #[pyclass] pub struct OnnxModel { nodes: Vec, @@ -89,6 +134,8 @@ pub fn load_onnx(path: &str) -> Result { .map_err(|e| e.to_string())?; let mut weights = HashMap::new(); + + // pass 1: load all named static weights for node in &graph.nodes { for arg in node.inputs() { if let Some((name, tensor)) = load_weight(arg) { @@ -97,6 +144,16 @@ pub fn load_onnx(path: &str) -> Result { } } + // pass 2: pre-compute BN scale/bias so dispatch is just 2 ops + for node in &graph.nodes { + if let Node::BatchNormalization(bn) = node { + if let Some((name, scale, offset)) = precompute_bn(bn, &weights) { + weights.insert(format!("{name}::scale"), scale); + weights.insert(format!("{name}::offset"), offset); + } + } + } + Ok(OnnxModel { input_names: graph.inputs.iter().map(|a| a.name.clone()).collect(), output_names: graph.outputs.iter().map(|a| a.name.clone()).collect(), @@ -107,22 +164,27 @@ pub fn load_onnx(path: &str) -> Result { fn dispatch(node: &Node, ctx: &mut ExecutionContext) { match node { - Node::Relu(n) => ops::activation::relu(n, ctx), - Node::Sigmoid(n) => ops::activation::sigmoid(n, ctx), - Node::Tanh(n) => ops::activation::tanh(n, ctx), - Node::Gelu(n) => ops::activation::gelu(n, ctx), - Node::Softmax(n) => ops::activation::softmax(n, ctx), - Node::LogSoftmax(n) => ops::activation::log_softmax(n, ctx), - Node::Linear(n) => ops::linear::linear(n, ctx), - Node::Gemm(n) => ops::linear::gemm(n, ctx), - Node::Add(n) => ops::elementwise::add(n, ctx), - Node::Sub(n) => ops::elementwise::sub(n, ctx), - Node::Mul(n) => ops::elementwise::mul(n, ctx), - Node::Div(n) => ops::elementwise::div(n, ctx), - Node::Reshape(n) => ops::reshape::reshape(n, ctx), - Node::Flatten(n) => ops::reshape::flatten(n, ctx), - Node::Transpose(n) => ops::reshape::transpose(n, ctx), - Node::Constant(_) => {} // already loaded into weights + Node::Relu(n) => ops::activation::relu(n, ctx), + Node::Sigmoid(n) => ops::activation::sigmoid(n, ctx), + Node::Tanh(n) => ops::activation::tanh(n, ctx), + Node::Gelu(n) => ops::activation::gelu(n, ctx), + Node::Softmax(n) => ops::activation::softmax(n, ctx), + Node::LogSoftmax(n) => ops::activation::log_softmax(n, ctx), + Node::Linear(n) => ops::linear::linear(n, ctx), + Node::Gemm(n) => ops::linear::gemm(n, ctx), + Node::Add(n) => ops::elementwise::add(n, ctx), + Node::Sub(n) => ops::elementwise::sub(n, ctx), + Node::Mul(n) => ops::elementwise::mul(n, ctx), + Node::Div(n) => ops::elementwise::div(n, ctx), + Node::Reshape(n) => ops::reshape::reshape(n, ctx), + Node::Flatten(n) => ops::reshape::flatten(n, ctx), + Node::Transpose(n) => ops::reshape::transpose(n, ctx), + Node::Conv2d(n) => ops::conv::conv2d(n, ctx), + Node::BatchNormalization(n) => ops::conv::batch_norm_fused(n, ctx), + Node::MaxPool2d(n) => ops::pool::max_pool2d(n, ctx), + Node::AveragePool2d(n) => ops::pool::avg_pool2d(n, ctx), + Node::GlobalAveragePool(n) => ops::pool::global_avg_pool(n, ctx), + Node::Constant(_) => {} other => { eprintln!("warn: unimplemented op '{}' — skipping", other.name()); } diff --git a/src/interpreter/ops/conv.rs b/src/interpreter/ops/conv.rs new file mode 100644 index 0000000..bf64d2f --- /dev/null +++ b/src/interpreter/ops/conv.rs @@ -0,0 +1,77 @@ +use burn_backend::{ + backend::ops::{FloatTensorOps, ModuleOps}, + backend::ops::ConvOptions, + TensorData, TensorMetadata, +}; +use onnx_ir::{ + conv2d::Conv2dNode, + batch_norm::{BatchNormalizationNode, BatchNormConfig}, + node::padding::PaddingConfig2d, +}; +use crate::tensor::{B, default_device}; +use super::super::context::ExecutionContext; + +fn symmetric_padding(cfg: &PaddingConfig2d) -> [usize; 2] { + match cfg { + PaddingConfig2d::Valid => [0, 0], + PaddingConfig2d::Explicit(top, left, bottom, right) => { + [(top + bottom) / 2, (left + right) / 2] + } + } +} + +pub fn conv2d(node: &Conv2dNode, ctx: &mut ExecutionContext) { + let x = ctx.resolve(&node.inputs[0]).expect("conv2d: missing input"); + let w = ctx.resolve(&node.inputs[1]).expect("conv2d: missing weight"); + let bias = node.inputs.get(2).and_then(|a| ctx.resolve(a)); + + let pad = symmetric_padding(&node.config.padding); + let options = ConvOptions::new( + node.config.stride, + pad, + node.config.dilation, + node.config.groups, + ); + + ctx.insert(node.outputs[0].name.clone(), B::conv2d(x, w, bias, options)); +} + +// BatchNorm dispatch — uses pre-fused scale/offset computed at load time. +// Falls back to full computation if pre-fusion wasn't possible. +pub fn batch_norm_fused(node: &BatchNormalizationNode, ctx: &mut ExecutionContext) { + let x = ctx.resolve(&node.inputs[0]).expect("batch_norm: missing input"); + + let scale_key = format!("{}::scale", node.name); + let offset_key = format!("{}::offset", node.name); + + if let (Some(scale), Some(offset)) = (ctx.get(&scale_key), ctx.get(&offset_key)) { + // fast path: 2 ops, no reshape (scale/offset are already [1,C,1,1]) + let y = B::float_mul(x, scale); + let y = B::float_add(y, offset); + ctx.insert(node.outputs[0].name.clone(), y); + } else { + // fallback: full manual BN (slow, but correct) + let gamma = ctx.resolve(&node.inputs[1]).expect("batch_norm: missing gamma"); + let beta = ctx.resolve(&node.inputs[2]).expect("batch_norm: missing beta"); + let mean = ctx.resolve(&node.inputs[3]).expect("batch_norm: missing mean"); + let var = ctx.resolve(&node.inputs[4]).expect("batch_norm: missing var"); + + let eps = match &node.config { + BatchNormConfig::Static(c) => c.epsilon as f32, + BatchNormConfig::Runtime(c) => c.epsilon as f32, + }; + let rank = x.shape().num_dims(); + let c = gamma.shape().iter().next().copied().unwrap_or(1); + let bshape: Vec = std::iter::once(1) + .chain(std::iter::once(c)) + .chain(std::iter::repeat(1).take(rank - 2)) + .collect(); + let bcast = |t| B::float_reshape(t, bshape.clone().into()); + let eps_t = B::float_from_data(TensorData::from([eps]), &default_device()); + let denom = B::float_sqrt(B::float_add(bcast(var), eps_t)); + let y = B::float_div(B::float_sub(x, bcast(mean)), denom); + let y = B::float_mul(y, bcast(gamma)); + let y = B::float_add(y, bcast(beta)); + ctx.insert(node.outputs[0].name.clone(), y); + } +} diff --git a/src/interpreter/ops/mod.rs b/src/interpreter/ops/mod.rs index 27ba454..6532549 100644 --- a/src/interpreter/ops/mod.rs +++ b/src/interpreter/ops/mod.rs @@ -2,3 +2,5 @@ pub mod activation; pub mod linear; pub mod elementwise; pub mod reshape; +pub mod conv; +pub mod pool; diff --git a/src/interpreter/ops/pool.rs b/src/interpreter/ops/pool.rs new file mode 100644 index 0000000..acd09f5 --- /dev/null +++ b/src/interpreter/ops/pool.rs @@ -0,0 +1,52 @@ +use burn_backend::backend::ops::ModuleOps; +use onnx_ir::{ + max_pool2d::MaxPool2dNode, + avg_pool2d::AveragePool2dNode, + global_avg_pool::GlobalAveragePoolNode, + node::padding::PaddingConfig2d, +}; +use crate::tensor::B; +use super::super::context::ExecutionContext; + +fn symmetric_padding(cfg: &PaddingConfig2d) -> [usize; 2] { + match cfg { + PaddingConfig2d::Valid => [0, 0], + PaddingConfig2d::Explicit(top, left, bottom, right) => { + [(top + bottom) / 2, (left + right) / 2] + } + } +} + +pub fn max_pool2d(node: &MaxPool2dNode, ctx: &mut ExecutionContext) { + let x = ctx.resolve(&node.inputs[0]).expect("max_pool2d: missing input"); + let pad = symmetric_padding(&node.config.padding); + let y = B::max_pool2d( + x, + node.config.kernel_size, + node.config.strides, + pad, + node.config.dilation, + false, + ); + ctx.insert(node.outputs[0].name.clone(), y); +} + +pub fn avg_pool2d(node: &AveragePool2dNode, ctx: &mut ExecutionContext) { + let x = ctx.resolve(&node.inputs[0]).expect("avg_pool2d: missing input"); + let pad = symmetric_padding(&node.config.padding); + let y = B::avg_pool2d( + x, + node.config.kernel_size, + node.config.strides, + pad, + node.config.count_include_pad, + false, + ); + ctx.insert(node.outputs[0].name.clone(), y); +} + +pub fn global_avg_pool(node: &GlobalAveragePoolNode, ctx: &mut ExecutionContext) { + let x = ctx.resolve(&node.inputs[0]).expect("global_avg_pool: missing input"); + let y = B::adaptive_avg_pool2d(x, [1, 1]); + ctx.insert(node.outputs[0].name.clone(), y); +} diff --git a/test.py b/test.py new file mode 100644 index 0000000..74c8948 --- /dev/null +++ b/test.py @@ -0,0 +1,13 @@ +import burn_python as burn +import numpy as np + +# Load an ONNX model +model = burn.load_onnx("tests/mlp.onnx") + +# Generate sample input matching the model expected input shape (batch_size=1, features=4) +x = np.array([[0.5, -0.2, 0.1, 0.9]], dtype=np.float32) + +# Run inference +output = model([x])[0] + +print("Model Output:", output) diff --git a/tests/compare_ort.py b/tests/compare_ort.py index 2cd5a49..55b5ad8 100644 --- a/tests/compare_ort.py +++ b/tests/compare_ort.py @@ -6,6 +6,7 @@ Run: python tests/compare_ort.py """ +import sys import time import numpy as np import onnxruntime as ort @@ -17,8 +18,13 @@ PASS = "\033[92mPASS\033[0m" FAIL = "\033[91mFAIL\033[0m" +failures = 0 + def check(label, cond): + global failures print(f" [{'PASS' if cond else 'FAIL'}] {label}") + if not cond: + failures += 1 return cond # ── load both ────────────────────────────────────────────────────────────────── @@ -70,3 +76,7 @@ def check(label, cond): print(f" batch={batch:<4} burn {burn_ms:.3f} ms | ORT {ort_ms:.3f} ms | ratio {ratio:.1f}x") print() + +if failures: + print(f"{failures} check(s) failed") + sys.exit(1) diff --git a/tests/test_bridge.py b/tests/test_bridge.py index 0d14d38..55a9221 100644 --- a/tests/test_bridge.py +++ b/tests/test_bridge.py @@ -3,6 +3,7 @@ Run with: python tests/test_bridge.py """ +import sys import time import numpy as np import burn_python as burn @@ -10,8 +11,13 @@ PASS = "\033[92mPASS\033[0m" FAIL = "\033[91mFAIL\033[0m" +failures = 0 + def check(label, cond): + global failures print(f" {'[' + PASS + ']' if cond else '[' + FAIL + ']'} {label}") + if not cond: + failures += 1 return cond # ── correctness ────────────────────────────────────────────────────────────── @@ -85,3 +91,7 @@ def check(label, cond): print(f" {label:>6} f32 | burn {burn_ms:.3f} ms | np.copy {np_ms:.3f} ms | ratio {ratio:.1f}x") print() + +if failures: + print(f"{failures} check(s) failed") + sys.exit(1) From 6605b74b96b5dc75307281cae5956555ffb7838a Mon Sep 17 00:00:00 2001 From: Asjad Date: Sat, 1 Aug 2026 12:00:31 +0500 Subject: [PATCH 2/3] style: cargo fmt + clippy fixes --- src/interpreter/context.rs | 13 ++-- src/interpreter/mod.rs | 99 ++++++++++++++++-------------- src/interpreter/ops/activation.rs | 20 +++--- src/interpreter/ops/conv.rs | 42 ++++++++----- src/interpreter/ops/elementwise.rs | 6 +- src/interpreter/ops/linear.rs | 30 ++++++--- src/interpreter/ops/mod.rs | 6 +- src/interpreter/ops/pool.rs | 22 ++++--- src/interpreter/ops/reshape.rs | 68 ++++++++++++-------- src/lib.rs | 5 +- src/tensor.rs | 5 +- 11 files changed, 190 insertions(+), 126 deletions(-) diff --git a/src/interpreter/context.rs b/src/interpreter/context.rs index 545e6a0..4e80143 100644 --- a/src/interpreter/context.rs +++ b/src/interpreter/context.rs @@ -6,7 +6,7 @@ use std::collections::HashMap; use burn_backend::backend::ops::FloatTensorOps; use onnx_ir::ir::{Argument, ValueSource}; -use crate::tensor::{default_device, FloatPrim, B}; +use crate::tensor::{B, FloatPrim, default_device}; pub struct ExecutionContext<'w> { tensors: HashMap, @@ -15,11 +15,15 @@ pub struct ExecutionContext<'w> { impl<'w> ExecutionContext<'w> { pub fn new(weights: &'w HashMap) -> Self { - Self { tensors: HashMap::new(), weights } + Self { + tensors: HashMap::new(), + weights, + } } pub fn get(&self, name: &str) -> Option { - self.tensors.get(name) + self.tensors + .get(name) .or_else(|| self.weights.get(name)) .cloned() } @@ -43,7 +47,8 @@ impl<'w> ExecutionContext<'w> { } } // anonymous inline static — convert on the fly - arg.value().map(|d| B::float_from_data(d, &default_device())) + arg.value() + .map(|d| B::float_from_data(d, &default_device())) } ValueSource::Optional => None, } diff --git a/src/interpreter/mod.rs b/src/interpreter/mod.rs index d765f61..c1de952 100644 --- a/src/interpreter/mod.rs +++ b/src/interpreter/mod.rs @@ -6,17 +6,16 @@ pub mod ops; use std::collections::HashMap; -use burn_backend::{backend::ops::FloatTensorOps, DType, TensorData, TensorMetadata}; +use burn_backend::{DType, TensorData, TensorMetadata, backend::ops::FloatTensorOps}; +use numpy::{PyArrayDyn, PyReadonlyArrayDyn}; use onnx_ir::{ - ir::{Argument, ValueSource}, + Node, OnnxGraphBuilder, batch_norm::{BatchNormConfig, BatchNormalizationNode}, - Node, - OnnxGraphBuilder, + ir::{Argument, ValueSource}, }; -use numpy::{PyArrayDyn, PyReadonlyArrayDyn}; use pyo3::prelude::*; -use crate::tensor::{default_device, flex_to_numpy, numpy_to_flex, FloatPrim, B}; +use crate::tensor::{B, FloatPrim, default_device, flex_to_numpy, numpy_to_flex}; use context::ExecutionContext; fn load_weight(arg: &Argument) -> Option<(String, FloatPrim)> { @@ -29,7 +28,10 @@ fn load_weight(arg: &Argument) -> Option<(String, FloatPrim)> { if !matches!(data.dtype, DType::F32 | DType::F16 | DType::BF16) { return None; } - Some((arg.name.clone(), B::float_from_data(data, &default_device()))) + Some(( + arg.name.clone(), + B::float_from_data(data, &default_device()), + )) } _ => None, } @@ -38,11 +40,12 @@ fn load_weight(arg: &Argument) -> Option<(String, FloatPrim)> { /// Pre-compute BN scale/bias as [1,C,1,1] tensors so runtime BN is just 2 elementwise ops. /// key: "{node_name}::scale" and "{node_name}::bias" /// inputs: [x, gamma, beta, running_mean, running_var] -fn precompute_bn(node: &BatchNormalizationNode, weights: &HashMap) - -> Option<(String, FloatPrim, FloatPrim)> -{ +fn precompute_bn( + node: &BatchNormalizationNode, + weights: &HashMap, +) -> Option<(String, FloatPrim, FloatPrim)> { let eps = match &node.config { - BatchNormConfig::Static(c) => c.epsilon as f32, + BatchNormConfig::Static(c) => c.epsilon as f32, BatchNormConfig::Runtime(c) => c.epsilon as f32, }; @@ -51,29 +54,27 @@ fn precompute_bn(node: &BatchNormalizationNode, weights: &HashMap = vec![1, c, 1, 1]; - let scale = B::float_reshape(scale_flat, shape_4d.clone().into()); + let scale = B::float_reshape(scale_flat, shape_4d.clone().into()); let offset = B::float_reshape(offset_flat, shape_4d.into()); Some((node.name.clone(), scale, offset)) @@ -111,9 +112,11 @@ impl OnnxModel { dispatch(node, &mut ctx); } - self.output_names.iter() + self.output_names + .iter() .map(|name| { - let t = ctx.get(name) + let t = ctx + .get(name) .unwrap_or_else(|| panic!("output '{}' not found", name)); Ok(flex_to_numpy(py, t)) }) @@ -123,7 +126,9 @@ impl OnnxModel { fn __repr__(&self) -> String { format!( "OnnxModel(inputs={:?}, outputs={:?}, nodes={})", - self.input_names, self.output_names, self.nodes.len() + self.input_names, + self.output_names, + self.nodes.len() ) } } @@ -146,11 +151,11 @@ pub fn load_onnx(path: &str) -> Result { // pass 2: pre-compute BN scale/bias so dispatch is just 2 ops for node in &graph.nodes { - if let Node::BatchNormalization(bn) = node { - if let Some((name, scale, offset)) = precompute_bn(bn, &weights) { - weights.insert(format!("{name}::scale"), scale); - weights.insert(format!("{name}::offset"), offset); - } + if let Node::BatchNormalization(bn) = node + && let Some((name, scale, offset)) = precompute_bn(bn, &weights) + { + weights.insert(format!("{name}::scale"), scale); + weights.insert(format!("{name}::offset"), offset); } } @@ -164,27 +169,27 @@ pub fn load_onnx(path: &str) -> Result { fn dispatch(node: &Node, ctx: &mut ExecutionContext) { match node { - Node::Relu(n) => ops::activation::relu(n, ctx), - Node::Sigmoid(n) => ops::activation::sigmoid(n, ctx), - Node::Tanh(n) => ops::activation::tanh(n, ctx), - Node::Gelu(n) => ops::activation::gelu(n, ctx), - Node::Softmax(n) => ops::activation::softmax(n, ctx), - Node::LogSoftmax(n) => ops::activation::log_softmax(n, ctx), - Node::Linear(n) => ops::linear::linear(n, ctx), - Node::Gemm(n) => ops::linear::gemm(n, ctx), - Node::Add(n) => ops::elementwise::add(n, ctx), - Node::Sub(n) => ops::elementwise::sub(n, ctx), - Node::Mul(n) => ops::elementwise::mul(n, ctx), - Node::Div(n) => ops::elementwise::div(n, ctx), - Node::Reshape(n) => ops::reshape::reshape(n, ctx), - Node::Flatten(n) => ops::reshape::flatten(n, ctx), - Node::Transpose(n) => ops::reshape::transpose(n, ctx), - Node::Conv2d(n) => ops::conv::conv2d(n, ctx), + Node::Relu(n) => ops::activation::relu(n, ctx), + Node::Sigmoid(n) => ops::activation::sigmoid(n, ctx), + Node::Tanh(n) => ops::activation::tanh(n, ctx), + Node::Gelu(n) => ops::activation::gelu(n, ctx), + Node::Softmax(n) => ops::activation::softmax(n, ctx), + Node::LogSoftmax(n) => ops::activation::log_softmax(n, ctx), + Node::Linear(n) => ops::linear::linear(n, ctx), + Node::Gemm(n) => ops::linear::gemm(n, ctx), + Node::Add(n) => ops::elementwise::add(n, ctx), + Node::Sub(n) => ops::elementwise::sub(n, ctx), + Node::Mul(n) => ops::elementwise::mul(n, ctx), + Node::Div(n) => ops::elementwise::div(n, ctx), + Node::Reshape(n) => ops::reshape::reshape(n, ctx), + Node::Flatten(n) => ops::reshape::flatten(n, ctx), + Node::Transpose(n) => ops::reshape::transpose(n, ctx), + Node::Conv2d(n) => ops::conv::conv2d(n, ctx), Node::BatchNormalization(n) => ops::conv::batch_norm_fused(n, ctx), - Node::MaxPool2d(n) => ops::pool::max_pool2d(n, ctx), - Node::AveragePool2d(n) => ops::pool::avg_pool2d(n, ctx), + Node::MaxPool2d(n) => ops::pool::max_pool2d(n, ctx), + Node::AveragePool2d(n) => ops::pool::avg_pool2d(n, ctx), Node::GlobalAveragePool(n) => ops::pool::global_avg_pool(n, ctx), - Node::Constant(_) => {} + Node::Constant(_) => {} other => { eprintln!("warn: unimplemented op '{}' — skipping", other.name()); } diff --git a/src/interpreter/ops/activation.rs b/src/interpreter/ops/activation.rs index 528e349..b24ba77 100644 --- a/src/interpreter/ops/activation.rs +++ b/src/interpreter/ops/activation.rs @@ -1,10 +1,10 @@ +use super::super::context::ExecutionContext; +use crate::tensor::B; use burn_backend::backend::ops::{ActivationOps, FloatTensorOps}; use onnx_ir::{ gelu::GeluNode, log_softmax::LogSoftmaxNode, relu::ReluNode, sigmoid::SigmoidNode, softmax::SoftmaxNode, tanh::TanhNode, }; -use crate::tensor::B; -use super::super::context::ExecutionContext; pub fn relu(node: &ReluNode, ctx: &mut ExecutionContext) { let x = ctx.get(&node.inputs[0].name).expect("relu: missing input"); @@ -12,7 +12,9 @@ pub fn relu(node: &ReluNode, ctx: &mut ExecutionContext) { } pub fn sigmoid(node: &SigmoidNode, ctx: &mut ExecutionContext) { - let x = ctx.get(&node.inputs[0].name).expect("sigmoid: missing input"); + let x = ctx + .get(&node.inputs[0].name) + .expect("sigmoid: missing input"); ctx.insert(node.outputs[0].name.clone(), B::sigmoid(x)); } @@ -27,13 +29,17 @@ pub fn gelu(node: &GeluNode, ctx: &mut ExecutionContext) { } pub fn softmax(node: &SoftmaxNode, ctx: &mut ExecutionContext) { - let x = ctx.get(&node.inputs[0].name).expect("softmax: missing input"); - let dim = node.config.axis as usize; + let x = ctx + .get(&node.inputs[0].name) + .expect("softmax: missing input"); + let dim = node.config.axis; ctx.insert(node.outputs[0].name.clone(), B::softmax(x, dim)); } pub fn log_softmax(node: &LogSoftmaxNode, ctx: &mut ExecutionContext) { - let x = ctx.get(&node.inputs[0].name).expect("log_softmax: missing input"); - let dim = node.config.axis as usize; + let x = ctx + .get(&node.inputs[0].name) + .expect("log_softmax: missing input"); + let dim = node.config.axis; ctx.insert(node.outputs[0].name.clone(), B::log_softmax(x, dim)); } diff --git a/src/interpreter/ops/conv.rs b/src/interpreter/ops/conv.rs index bf64d2f..72c859e 100644 --- a/src/interpreter/ops/conv.rs +++ b/src/interpreter/ops/conv.rs @@ -1,15 +1,15 @@ +use super::super::context::ExecutionContext; +use crate::tensor::{B, default_device}; use burn_backend::{ - backend::ops::{FloatTensorOps, ModuleOps}, - backend::ops::ConvOptions, TensorData, TensorMetadata, + backend::ops::ConvOptions, + backend::ops::{FloatTensorOps, ModuleOps}, }; use onnx_ir::{ + batch_norm::{BatchNormConfig, BatchNormalizationNode}, conv2d::Conv2dNode, - batch_norm::{BatchNormalizationNode, BatchNormConfig}, node::padding::PaddingConfig2d, }; -use crate::tensor::{B, default_device}; -use super::super::context::ExecutionContext; fn symmetric_padding(cfg: &PaddingConfig2d) -> [usize; 2] { match cfg { @@ -21,8 +21,10 @@ fn symmetric_padding(cfg: &PaddingConfig2d) -> [usize; 2] { } pub fn conv2d(node: &Conv2dNode, ctx: &mut ExecutionContext) { - let x = ctx.resolve(&node.inputs[0]).expect("conv2d: missing input"); - let w = ctx.resolve(&node.inputs[1]).expect("conv2d: missing weight"); + let x = ctx.resolve(&node.inputs[0]).expect("conv2d: missing input"); + let w = ctx + .resolve(&node.inputs[1]) + .expect("conv2d: missing weight"); let bias = node.inputs.get(2).and_then(|a| ctx.resolve(a)); let pad = symmetric_padding(&node.config.padding); @@ -39,9 +41,11 @@ pub fn conv2d(node: &Conv2dNode, ctx: &mut ExecutionContext) { // BatchNorm dispatch — uses pre-fused scale/offset computed at load time. // Falls back to full computation if pre-fusion wasn't possible. pub fn batch_norm_fused(node: &BatchNormalizationNode, ctx: &mut ExecutionContext) { - let x = ctx.resolve(&node.inputs[0]).expect("batch_norm: missing input"); + let x = ctx + .resolve(&node.inputs[0]) + .expect("batch_norm: missing input"); - let scale_key = format!("{}::scale", node.name); + let scale_key = format!("{}::scale", node.name); let offset_key = format!("{}::offset", node.name); if let (Some(scale), Some(offset)) = (ctx.get(&scale_key), ctx.get(&offset_key)) { @@ -51,20 +55,28 @@ pub fn batch_norm_fused(node: &BatchNormalizationNode, ctx: &mut ExecutionContex ctx.insert(node.outputs[0].name.clone(), y); } else { // fallback: full manual BN (slow, but correct) - let gamma = ctx.resolve(&node.inputs[1]).expect("batch_norm: missing gamma"); - let beta = ctx.resolve(&node.inputs[2]).expect("batch_norm: missing beta"); - let mean = ctx.resolve(&node.inputs[3]).expect("batch_norm: missing mean"); - let var = ctx.resolve(&node.inputs[4]).expect("batch_norm: missing var"); + let gamma = ctx + .resolve(&node.inputs[1]) + .expect("batch_norm: missing gamma"); + let beta = ctx + .resolve(&node.inputs[2]) + .expect("batch_norm: missing beta"); + let mean = ctx + .resolve(&node.inputs[3]) + .expect("batch_norm: missing mean"); + let var = ctx + .resolve(&node.inputs[4]) + .expect("batch_norm: missing var"); let eps = match &node.config { - BatchNormConfig::Static(c) => c.epsilon as f32, + BatchNormConfig::Static(c) => c.epsilon as f32, BatchNormConfig::Runtime(c) => c.epsilon as f32, }; let rank = x.shape().num_dims(); let c = gamma.shape().iter().next().copied().unwrap_or(1); let bshape: Vec = std::iter::once(1) .chain(std::iter::once(c)) - .chain(std::iter::repeat(1).take(rank - 2)) + .chain(std::iter::repeat_n(1, rank - 2)) .collect(); let bcast = |t| B::float_reshape(t, bshape.clone().into()); let eps_t = B::float_from_data(TensorData::from([eps]), &default_device()); diff --git a/src/interpreter/ops/elementwise.rs b/src/interpreter/ops/elementwise.rs index 6a143aa..6b536b4 100644 --- a/src/interpreter/ops/elementwise.rs +++ b/src/interpreter/ops/elementwise.rs @@ -1,7 +1,7 @@ -use burn_backend::backend::ops::FloatTensorOps; -use onnx_ir::arithmetic::{AddNode, SubNode, MulNode, DivNode}; -use crate::tensor::B; use super::super::context::ExecutionContext; +use crate::tensor::B; +use burn_backend::backend::ops::FloatTensorOps; +use onnx_ir::arithmetic::{AddNode, DivNode, MulNode, SubNode}; pub fn add(node: &AddNode, ctx: &mut ExecutionContext) { let a = ctx.get(&node.inputs[0].name).expect("add: missing lhs"); diff --git a/src/interpreter/ops/linear.rs b/src/interpreter/ops/linear.rs index 1025c38..58a2a1a 100644 --- a/src/interpreter/ops/linear.rs +++ b/src/interpreter/ops/linear.rs @@ -1,11 +1,15 @@ -use burn_backend::backend::ops::FloatTensorOps; -use onnx_ir::{linear::LinearNode, gemm::GemmNode}; -use crate::tensor::{B, default_device}; use super::super::context::ExecutionContext; +use crate::tensor::{B, default_device}; +use burn_backend::backend::ops::FloatTensorOps; +use onnx_ir::{gemm::GemmNode, linear::LinearNode}; pub fn linear(node: &LinearNode, ctx: &mut ExecutionContext) { - let x = ctx.get(&node.inputs[0].name).expect("linear: missing input"); - let w = ctx.get(&node.inputs[1].name).expect("linear: missing weight"); + let x = ctx + .get(&node.inputs[0].name) + .expect("linear: missing input"); + let w = ctx + .get(&node.inputs[1].name) + .expect("linear: missing weight"); // Gemm layout: w is [out, in], needs transpose before matmul // MatMul layout: w is [in, out], use as-is @@ -47,8 +51,16 @@ pub fn gemm(node: &GemmNode, ctx: &mut ExecutionContext) { let a = ctx.get(&node.inputs[0].name).expect("gemm: missing A"); let b = ctx.get(&node.inputs[1].name).expect("gemm: missing B"); - let a = if node.config.trans_a != 0 { B::float_swap_dims(a, 0, 1) } else { a }; - let b = if node.config.trans_b != 0 { B::float_swap_dims(b, 0, 1) } else { b }; + let a = if node.config.trans_a != 0 { + B::float_swap_dims(a, 0, 1) + } else { + a + }; + let b = if node.config.trans_b != 0 { + B::float_swap_dims(b, 0, 1) + } else { + b + }; let mut y = B::float_matmul(a, b); @@ -64,7 +76,9 @@ pub fn gemm(node: &GemmNode, ctx: &mut ExecutionContext) { let c = if !c_arg.name.is_empty() { ctx.get(&c_arg.name) } else { - c_arg.value().map(|d| B::float_from_data(d, &default_device())) + c_arg + .value() + .map(|d| B::float_from_data(d, &default_device())) }; if let Some(c) = c { let c = if node.config.beta != 1.0 { diff --git a/src/interpreter/ops/mod.rs b/src/interpreter/ops/mod.rs index 6532549..7e098b5 100644 --- a/src/interpreter/ops/mod.rs +++ b/src/interpreter/ops/mod.rs @@ -1,6 +1,6 @@ pub mod activation; -pub mod linear; -pub mod elementwise; -pub mod reshape; pub mod conv; +pub mod elementwise; +pub mod linear; pub mod pool; +pub mod reshape; diff --git a/src/interpreter/ops/pool.rs b/src/interpreter/ops/pool.rs index acd09f5..38865e6 100644 --- a/src/interpreter/ops/pool.rs +++ b/src/interpreter/ops/pool.rs @@ -1,12 +1,10 @@ +use super::super::context::ExecutionContext; +use crate::tensor::B; use burn_backend::backend::ops::ModuleOps; use onnx_ir::{ - max_pool2d::MaxPool2dNode, - avg_pool2d::AveragePool2dNode, - global_avg_pool::GlobalAveragePoolNode, - node::padding::PaddingConfig2d, + avg_pool2d::AveragePool2dNode, global_avg_pool::GlobalAveragePoolNode, + max_pool2d::MaxPool2dNode, node::padding::PaddingConfig2d, }; -use crate::tensor::B; -use super::super::context::ExecutionContext; fn symmetric_padding(cfg: &PaddingConfig2d) -> [usize; 2] { match cfg { @@ -18,7 +16,9 @@ fn symmetric_padding(cfg: &PaddingConfig2d) -> [usize; 2] { } pub fn max_pool2d(node: &MaxPool2dNode, ctx: &mut ExecutionContext) { - let x = ctx.resolve(&node.inputs[0]).expect("max_pool2d: missing input"); + let x = ctx + .resolve(&node.inputs[0]) + .expect("max_pool2d: missing input"); let pad = symmetric_padding(&node.config.padding); let y = B::max_pool2d( x, @@ -32,7 +32,9 @@ pub fn max_pool2d(node: &MaxPool2dNode, ctx: &mut ExecutionContext) { } pub fn avg_pool2d(node: &AveragePool2dNode, ctx: &mut ExecutionContext) { - let x = ctx.resolve(&node.inputs[0]).expect("avg_pool2d: missing input"); + let x = ctx + .resolve(&node.inputs[0]) + .expect("avg_pool2d: missing input"); let pad = symmetric_padding(&node.config.padding); let y = B::avg_pool2d( x, @@ -46,7 +48,9 @@ pub fn avg_pool2d(node: &AveragePool2dNode, ctx: &mut ExecutionContext) { } pub fn global_avg_pool(node: &GlobalAveragePoolNode, ctx: &mut ExecutionContext) { - let x = ctx.resolve(&node.inputs[0]).expect("global_avg_pool: missing input"); + let x = ctx + .resolve(&node.inputs[0]) + .expect("global_avg_pool: missing input"); let y = B::adaptive_avg_pool2d(x, [1, 1]); ctx.insert(node.outputs[0].name.clone(), y); } diff --git a/src/interpreter/ops/reshape.rs b/src/interpreter/ops/reshape.rs index 74143be..1b92d01 100644 --- a/src/interpreter/ops/reshape.rs +++ b/src/interpreter/ops/reshape.rs @@ -1,44 +1,64 @@ -use burn_backend::{backend::ops::FloatTensorOps, TensorMetadata}; -use onnx_ir::{flatten::FlattenNode, reshape::ReshapeNode, transpose::TransposeNode}; -use crate::tensor::B; use super::super::context::ExecutionContext; +use crate::tensor::B; +use burn_backend::{TensorMetadata, backend::ops::FloatTensorOps}; +use onnx_ir::{flatten::FlattenNode, reshape::ReshapeNode, transpose::TransposeNode}; pub fn reshape(node: &ReshapeNode, ctx: &mut ExecutionContext) { - let x = ctx.get(&node.inputs[0].name).expect("reshape: missing input"); + let x = ctx + .get(&node.inputs[0].name) + .expect("reshape: missing input"); let orig_shape: Vec = x.shape().iter().copied().collect(); let total: usize = orig_shape.iter().product(); - let shape_data = node.inputs[1].value().expect("reshape: shape must be static"); + let shape_data = node.inputs[1] + .value() + .expect("reshape: shape must be static"); let raw: Vec = shape_data.to_vec().expect("reshape: shape must be i64"); - let shape: Vec = raw.iter().enumerate().map(|(i, &s)| { - if s == -1 { - let known: usize = raw.iter().enumerate() - .filter(|(j, v)| *j != i && **v != -1) - .map(|(_, &v)| v as usize) - .product(); - total / known - } else if s == 0 { - orig_shape[i] - } else { - s as usize - } - }).collect(); - - ctx.insert(node.outputs[0].name.clone(), B::float_reshape(x, shape.into())); + let shape: Vec = raw + .iter() + .enumerate() + .map(|(i, &s)| { + if s == -1 { + let known: usize = raw + .iter() + .enumerate() + .filter(|(j, v)| *j != i && **v != -1) + .map(|(_, &v)| v as usize) + .product(); + total / known + } else if s == 0 { + orig_shape[i] + } else { + s as usize + } + }) + .collect(); + + ctx.insert( + node.outputs[0].name.clone(), + B::float_reshape(x, shape.into()), + ); } pub fn flatten(node: &FlattenNode, ctx: &mut ExecutionContext) { - let x = ctx.get(&node.inputs[0].name).expect("flatten: missing input"); + let x = ctx + .get(&node.inputs[0].name) + .expect("flatten: missing input"); let shape: Vec = x.shape().iter().copied().collect(); - let axis = node.config.axis as usize; + let axis = node.config.axis; let outer: usize = shape[..axis].iter().product::().max(1); let inner: usize = shape[axis..].iter().product(); - ctx.insert(node.outputs[0].name.clone(), B::float_reshape(x, vec![outer, inner].into())); + ctx.insert( + node.outputs[0].name.clone(), + B::float_reshape(x, vec![outer, inner].into()), + ); } pub fn transpose(node: &TransposeNode, ctx: &mut ExecutionContext) { - let x = ctx.get(&node.inputs[0].name).expect("transpose: missing input"); + let x = ctx + .get(&node.inputs[0].name) + .expect("transpose: missing input"); let rank = x.shape().num_dims(); let perm: Vec = if node.config.perm.is_empty() { diff --git a/src/lib.rs b/src/lib.rs index 89c2991..21018c2 100644 --- a/src/lib.rs +++ b/src/lib.rs @@ -1,8 +1,8 @@ use numpy::{PyArrayDyn, PyReadonlyArrayDyn}; use pyo3::prelude::*; -mod tensor; mod interpreter; +mod tensor; use interpreter::OnnxModel; @@ -19,8 +19,7 @@ fn roundtrip<'py>( /// Load an ONNX model from a file path. #[pyfunction] fn load_onnx(path: &str) -> PyResult { - interpreter::load_onnx(path) - .map_err(|e| pyo3::exceptions::PyRuntimeError::new_err(e)) + interpreter::load_onnx(path).map_err(pyo3::exceptions::PyRuntimeError::new_err) } #[pymodule] diff --git a/src/tensor.rs b/src/tensor.rs index a43efe5..71f9836 100644 --- a/src/tensor.rs +++ b/src/tensor.rs @@ -3,7 +3,7 @@ // in: numpy f32 -> TensorData (one copy) -> FlexTensor // out: FlexTensor -> TensorData (sync CPU read) -> numpy f32 (one copy) -use burn_backend::{backend::ops::FloatTensorOps, TensorData, DType}; +use burn_backend::{DType, TensorData, backend::ops::FloatTensorOps}; use burn_flex::{Flex, FlexDevice, FlexTensor}; use cubecl_common::reader::read_sync; use numpy::{PyArray1, PyArrayDyn, PyArrayMethods, PyReadonlyArrayDyn, PyUntypedArrayMethods}; @@ -32,8 +32,7 @@ pub fn numpy_to_flex(arr: &PyReadonlyArrayDyn<'_, f32>) -> FloatPrim { // FlexTensor -> numpy f32 ndarray (sync read + one copy out) pub fn flex_to_numpy<'py>(py: Python<'py>, prim: FloatPrim) -> Bound<'py, PyArrayDyn> { // for CPU backends the future resolves immediately - let data: TensorData = read_sync(B::float_into_data(prim)) - .expect("float_into_data error"); + let data: TensorData = read_sync(B::float_into_data(prim)).expect("float_into_data error"); let shape: Vec = data.shape.iter().copied().collect(); let floats: &[f32] = bytemuck::cast_slice(data.as_bytes()); From 4e54a8564c17fa88b32b6effa9ca01110645fbe9 Mon Sep 17 00:00:00 2001 From: Asjad Date: Sat, 1 Aug 2026 12:05:43 +0500 Subject: [PATCH 3/3] ci: fix maturin develop failing without an active virtualenv --- .github/workflows/ci.yml | 13 ++++++++----- 1 file changed, 8 insertions(+), 5 deletions(-) diff --git a/.github/workflows/ci.yml b/.github/workflows/ci.yml index c77369c..2c98001 100644 --- a/.github/workflows/ci.yml +++ b/.github/workflows/ci.yml @@ -64,17 +64,20 @@ jobs: with: python-version: "3.12" + - name: Create virtualenv + run: python -m venv .venv + - name: Install Python deps - run: pip install maturin numpy onnx onnxruntime + run: .venv/bin/pip install maturin numpy onnx onnxruntime - name: Build extension - run: maturin develop --release + run: source .venv/bin/activate && maturin develop --release - name: Generate test fixtures - run: python tests/make_test_model.py + run: .venv/bin/python tests/make_test_model.py - name: Run tensor bridge tests - run: python tests/test_bridge.py + run: .venv/bin/python tests/test_bridge.py - name: Run ORT comparison tests - run: python tests/compare_ort.py + run: .venv/bin/python tests/compare_ort.py