diff options
| author | Dennis Kobert <dennis@kobert.dev> | 2025-03-25 19:24:16 +0100 |
|---|---|---|
| committer | Dennis Kobert <dennis@kobert.dev> | 2025-03-25 19:24:16 +0100 |
| commit | e85e1cee0f662fe03ec2644330ae81b69502ca7b (patch) | |
| tree | 4d5f44b3819ff86da65bdd1d6ccc055607c8e9fd /src | |
| parent | a7da9c7140181ab4ef4af5e5f3222ca44741d66a (diff) | |
Build basic inference setup in rust
Diffstat (limited to 'src')
| -rw-r--r-- | src/main.rs | 17 | ||||
| -rw-r--r-- | src/model.rs | 59 |
2 files changed, 75 insertions, 1 deletions
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::<u64>("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<f32>; + +#[derive(Module, Debug)] +pub struct Net<B: Backend> { + input_lin: Linear<B>, + relu1: Relu, + lin2: Linear<B>, + relu2: Relu, + lin3: Linear<B>, +} + +impl<B: Backend> Net<B> { + /// 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<B, 1>) -> Tensor<B, 1> { + 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<ArrayBackend> { + let device = Default::default(); + let record: NetRecord<ArrayBackend> = PyTorchFileRecorder::<FullPrecisionSettings>::default() + .load("./perf.pt".into(), &device) + .expect("Failed to decode state"); + + Net::<ArrayBackend>::init(&device).load_record(record) +} |
