From e85e1cee0f662fe03ec2644330ae81b69502ca7b Mon Sep 17 00:00:00 2001 From: Dennis Kobert Date: Tue, 25 Mar 2025 19:24:16 +0100 Subject: Build basic inference setup in rust --- src/main.rs | 17 ++++++++++++++++- src/model.rs | 59 +++++++++++++++++++++++++++++++++++++++++++++++++++++++++++ 2 files changed, 75 insertions(+), 1 deletion(-) create mode 100644 src/model.rs (limited to 'src') diff --git a/src/main.rs b/src/main.rs index f3d4992..4f7d7a4 100644 --- a/src/main.rs +++ b/src/main.rs @@ -7,12 +7,14 @@ mod benchmark; mod e_core_selector; mod energy; mod freq; +mod model; mod scheduler; #[rustfmt::skip] mod bpf; use anyhow::Result; +use burn::tensor::Tensor; use clap::{Arg, ArgAction, Command}; use scheduler::Scheduler; use std::mem::MaybeUninit; @@ -46,13 +48,26 @@ fn main() -> Result<()> { ) .get_matches(); + let model = model::load_model(); + let device = Default::default(); + let result = model.forward(Tensor::from_floats( + [ + // 2899.97, 21420886., 59226., 148084., 301003., 36244800., 115862107., 43766905., + 1533, 1077473, 9448, 4269, 52805, 2456984, 5867215, 3954587, + ], + // [1., 1., 1., 1., 1., 1., 1., 1.], + &device, + )); + + println!("result: {}", result); + let power_cap = *matches.get_one::("power_cap").unwrap_or(&u64::MAX); let use_mocking = matches.get_flag("mock"); let benchmark = matches.get_flag("benchmark"); // Initialize and load the scheduler. let mut open_object = MaybeUninit::uninit(); - let log_path = "logs.csv"; + let log_path = "/tmp/logs.csv"; if benchmark { let mut sched = BenchmarkScheduler::init(&mut open_object, log_path)?; sched.run(); diff --git a/src/model.rs b/src/model.rs new file mode 100644 index 0000000..04eb6c8 --- /dev/null +++ b/src/model.rs @@ -0,0 +1,59 @@ +use burn::{ + nn::conv::{Conv2d, Conv2dConfig}, + nn::{Linear, Relu}, + prelude::*, +}; +use nn::{LeakyReluConfig, LinearConfig}; + +use burn::record::{FullPrecisionSettings, NamedMpkFileRecorder, Recorder}; +use burn_import::pytorch::PyTorchFileRecorder; + +type ArrayBackend = burn_ndarray::NdArray; + +#[derive(Module, Debug)] +pub struct Net { + input_lin: Linear, + relu1: Relu, + lin2: Linear, + relu2: Relu, + lin3: Linear, +} + +impl Net { + /// Create a new model. + pub fn init(device: &B::Device) -> Self { + let input_size = 8; + let input_lin = LinearConfig::new(input_size, 16).init(device); + let relu1 = Relu::new(); + let lin2 = LinearConfig::new(16, 4).init(device); + let relu2 = Relu::new(); + let lin3 = LinearConfig::new(4, 1).init(device); + Self { + input_lin, + relu1, + lin2, + relu2, + lin3, + } + } + + /// Forward pass of the model. + pub fn forward(&self, x: Tensor) -> Tensor { + let x = self.input_lin.forward(x); + let x = self.relu1.forward(x); + let x = self.lin2.forward(x); + let x = self.relu2.forward(x); + let x = self.lin3.forward(x); + x + } +} + +/// Load the model from the file in your source code (not in build.rs or script). +pub fn load_model() -> Net { + let device = Default::default(); + let record: NetRecord = PyTorchFileRecorder::::default() + .load("./perf.pt".into(), &device) + .expect("Failed to decode state"); + + Net::::init(&device).load_record(record) +} -- cgit v1.3.1