Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension


Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
83 changes: 83 additions & 0 deletions .github/workflows/ci.yml
Original file line number Diff line number Diff line change
@@ -0,0 +1,83 @@
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: Create virtualenv
run: python -m venv .venv

- name: Install Python deps
run: .venv/bin/pip install maturin numpy onnx onnxruntime

- name: Build extension
run: source .venv/bin/activate && maturin develop --release

- name: Generate test fixtures
run: .venv/bin/python tests/make_test_model.py

- name: Run tensor bridge tests
run: .venv/bin/python tests/test_bridge.py

- name: Run ORT comparison tests
run: .venv/bin/python tests/compare_ort.py
3 changes: 3 additions & 0 deletions .gitignore
Original file line number Diff line number Diff line change
Expand Up @@ -6,3 +6,6 @@ dist/
__pycache__/
.venv/
Cargo.lock
.claude/
tests/*.onnx
tests/*.npy
38 changes: 33 additions & 5 deletions src/interpreter/context.rs
Original file line number Diff line number Diff line change
Expand Up @@ -2,27 +2,55 @@
// 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::{B, FloatPrim, default_device};

pub struct ExecutionContext<'w> {
// intermediate tensors produced during the forward pass
tensors: HashMap<String, FloatPrim>,
// read-only reference to pre-loaded model weights
weights: &'w HashMap<String, FloatPrim>,
}

impl<'w> ExecutionContext<'w> {
pub fn new(weights: &'w HashMap<String, FloatPrim>) -> Self {
Self { tensors: HashMap::new(), weights }
Self {
tensors: HashMap::new(),
weights,
}
}

pub fn get(&self, name: &str) -> Option<FloatPrim> {
self.tensors.get(name)
self.tensors
.get(name)
.or_else(|| self.weights.get(name))
.cloned()
}

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<FloatPrim> {
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,
}
}
}
117 changes: 92 additions & 25 deletions src/interpreter/mod.rs
Original file line number Diff line number Diff line change
@@ -1,21 +1,21 @@
// 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::{DType, TensorData, TensorMetadata, backend::ops::FloatTensorOps};
use numpy::{PyArrayDyn, PyReadonlyArrayDyn};
use onnx_ir::{
Node, OnnxGraphBuilder,
batch_norm::{BatchNormConfig, BatchNormalizationNode},
ir::{Argument, ValueSource},
Node,
OnnxGraphBuilder,
};
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)> {
Expand All @@ -28,12 +28,58 @@ 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,
}
}

/// 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<String, FloatPrim>,
) -> 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<FloatPrim> {
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<usize> = 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<Node>,
Expand Down Expand Up @@ -66,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))
})
Expand All @@ -78,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()
)
}
}
Expand All @@ -89,6 +139,8 @@ pub fn load_onnx(path: &str) -> Result<OnnxModel, String> {
.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) {
Expand All @@ -97,6 +149,16 @@ pub fn load_onnx(path: &str) -> Result<OnnxModel, String> {
}
}

// pass 2: pre-compute BN scale/bias so dispatch is just 2 ops
for node in &graph.nodes {
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);
}
}

Ok(OnnxModel {
input_names: graph.inputs.iter().map(|a| a.name.clone()).collect(),
output_names: graph.outputs.iter().map(|a| a.name.clone()).collect(),
Expand All @@ -107,22 +169,27 @@ pub fn load_onnx(path: &str) -> Result<OnnxModel, String> {

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::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::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());
}
Expand Down
20 changes: 13 additions & 7 deletions src/interpreter/ops/activation.rs
Original file line number Diff line number Diff line change
@@ -1,18 +1,20 @@
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");
ctx.insert(node.outputs[0].name.clone(), B::relu(x));
}

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));
}

Expand All @@ -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));
}
Loading
Loading