Skip to main content

surrealml_core/execution/
session.rs

1//! Defines the session module for the execution module.
2#[cfg(feature = "gpu")]
3use ort::execution_providers::CUDAExecutionProvider;
4#[cfg(feature = "gpu")]
5use ort::execution_providers::ExecutionProvider;
6use ort::session::Session;
7
8use crate::errors::error::{SurrealError, SurrealErrorStatus};
9use crate::safe_eject;
10
11/// Creates a session for a model.
12///
13/// # Arguments
14/// * `model_bytes` - The model bytes (usually extracted fromt the surml file)
15///
16/// # Returns
17/// A session object.
18pub fn get_session(model_bytes: Vec<u8>) -> Result<Session, SurrealError> {
19	#[cfg(feature = "gpu")]
20	let mut builder = safe_eject!(Session::builder(), SurrealErrorStatus::Unknown);
21
22	#[cfg(not(feature = "gpu"))]
23	let mut builder = safe_eject!(Session::builder(), SurrealErrorStatus::Unknown);
24
25	#[cfg(feature = "gpu")]
26	{
27		let cuda = CUDAExecutionProvider::default();
28		if let Err(e) = cuda.register(&mut builder) {
29			eprintln!("Failed to register CUDA: {:?}. Falling back to CPU.", e);
30		} else {
31			println!("CUDA registered successfully");
32		}
33	}
34	let session: Session =
35		safe_eject!(builder.commit_from_memory(&model_bytes), SurrealErrorStatus::Unknown);
36	Ok(session)
37}