use rlx_ir::op::{Activation, ReduceOp};
use rlx_ir::{DType, Graph, NodeId, Op, Shape};
use crate::model::{Layer, Model};
pub struct IrModel {
pub graph: Graph,
pub node_for_layer: Vec<NodeId>,
}
pub fn to_graph(model: &Model) -> IrModel {
let mut g = Graph::new(model.name.clone());
let f32_dt = DType::F32;
let i8_dt = DType::I8;
let input_node = g.input("model_input", Shape::new(&[model.input_len], i8_dt));
let mut prev: NodeId = input_node;
let mut node_for_layer = Vec::with_capacity(model.layers.len());
for layer in &model.layers {
let node = match layer {
Layer::Conv2d {
name,
h_in: _,
w_in: _,
c_in,
c_out,
kh,
kw,
pad_h,
pad_w,
stride_h,
stride_w,
x_zp,
w_zp,
out_zp,
weight_bits: _,
requant,
weights,
bias,
..
} => {
let w_param = g.param(
format!("{name}_w"),
Shape::new(&[*c_out, *c_in, *kh, *kw], i8_dt),
);
let b_shape = match bias {
Some(b) => Shape::new(&[b.len()], DType::I32),
None => Shape::new(&[*c_out], DType::I32),
};
let b_param = g.param(format!("{name}_b"), b_shape);
let _ = (weights, requant); let scalar_mult = requant
.first()
.map(|&(m0, sh)| q31_to_f32_mult(m0, sh))
.unwrap_or(1.0);
g.q_conv2d(
prev,
w_param,
b_param,
vec![*kh, *kw],
vec![*stride_h, *stride_w],
vec![*pad_h, *pad_w],
vec![1, 1],
1,
*x_zp,
*w_zp,
*out_zp,
scalar_mult,
Shape::new(&[layer.out_len()], i8_dt),
)
}
Layer::Relu { len, .. } => {
g.activation(Activation::Relu, prev, Shape::new(&[*len], i8_dt))
}
Layer::MaxPool2d {
kh,
kw,
stride_h,
stride_w,
..
} => g.add_node(
Op::Pool {
kind: ReduceOp::Max,
kernel_size: vec![*kh, *kw],
stride: vec![*stride_h, *stride_w],
padding: vec![0, 0],
},
vec![prev],
Shape::new(&[layer.out_len()], i8_dt),
),
Layer::Dense {
name,
in_features,
out_features,
x_zp,
w_zp,
out_zp,
weight_bits: _,
requant,
weights: _,
bias,
..
} => {
let w_param = g.param(
format!("{name}_w"),
Shape::new(&[*in_features, *out_features], i8_dt),
);
let b_shape = match bias {
Some(b) => Shape::new(&[b.len()], DType::I32),
None => Shape::new(&[*out_features], DType::I32),
};
let b_param = g.param(format!("{name}_b"), b_shape);
let scalar_mult = requant
.first()
.map(|&(m0, sh)| q31_to_f32_mult(m0, sh))
.unwrap_or(1.0);
g.q_matmul(
prev,
w_param,
b_param,
*x_zp,
*w_zp,
*out_zp,
scalar_mult,
Shape::new(&[*out_features], i8_dt),
)
}
Layer::Argmax { len: _, .. } => {
g.add_node(
Op::TopK { k: 1 },
vec![prev],
Shape::new(&[1], f32_dt), )
}
};
node_for_layer.push(node);
prev = node;
}
g.set_outputs(vec![prev]);
IrModel {
graph: g,
node_for_layer,
}
}
fn q31_to_f32_mult(m0: i32, shift: i32) -> f32 {
let s = m0 as f64 / (1u64 << 31) as f64;
let scale = 2f64.powi(-shift);
(s * scale) as f32
}
#[cfg(test)]
mod tests {
use super::*;
use crate::model::tinyconv_mnist_from_cortexm;
#[test]
fn graph_has_one_node_per_layer_plus_params_plus_input() {
let m = tinyconv_mnist_from_cortexm();
let ir = to_graph(&m);
let mut expected = 1; for l in &m.layers {
expected += 1; if matches!(l, Layer::Conv2d { .. } | Layer::Dense { .. }) {
expected += 2; }
}
assert_eq!(ir.graph.len(), expected);
assert_eq!(ir.node_for_layer.len(), m.layers.len());
}
#[test]
fn ir_graph_passes_verifier() {
let m = tinyconv_mnist_from_cortexm();
let ir = to_graph(&m);
let errors = rlx_ir::verify::verify(&ir.graph);
assert!(
errors.is_empty(),
"verifier reported errors: {:?}",
errors.iter().map(|e| e.to_string()).collect::<Vec<_>>()
);
}
#[test]
fn graph_output_is_argmax() {
let m = tinyconv_mnist_from_cortexm();
let ir = to_graph(&m);
assert_eq!(ir.graph.outputs.len(), 1);
let out_node = ir.graph.node(ir.graph.outputs[0]);
assert!(
matches!(out_node.op, Op::TopK { .. }),
"expected TopK output, got {:?}",
out_node.op
);
}
}