summaryrefslogtreecommitdiff
path: root/src
diff options
context:
space:
mode:
authorDennis Kobert <dennis@kobert.dev>2025-03-25 19:24:16 +0100
committerDennis Kobert <dennis@kobert.dev>2025-03-25 19:24:16 +0100
commite85e1cee0f662fe03ec2644330ae81b69502ca7b (patch)
tree4d5f44b3819ff86da65bdd1d6ccc055607c8e9fd /src
parenta7da9c7140181ab4ef4af5e5f3222ca44741d66a (diff)
Build basic inference setup in rust
Diffstat (limited to 'src')
-rw-r--r--src/main.rs17
-rw-r--r--src/model.rs59
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)
+}