安全实例化ONNX Runtime环境与Session,规避借用错误
安全管理ONNX Runtime Environment与Session的Rust实现方案
我需要在结构体中封装ONNX Runtime的Environment和Session,以及相关预测方法。目前已有可运行的实现,但使用了unsafe内存操作,违背了用Arc管理共享状态的初衷。
当前可运行(含unsafe)的代码
use std::sync::Arc; use anyhow::{Context, Error, Result}; use onnxruntime::{ environment::Environment, session::Session, tensor::OrtOwnedTensor, GraphOptimizationLevel, LoggingLevel }; use ndarray_stats::QuantileExt; use ndarray::{ Array, s }; pub struct ModelWrapper { environment: Arc<Environment>, session: Session<'static>, } impl ModelWrapper { pub fn new(model_path: String) -> Result<Self, Error> { let environment = Arc::new(Environment::builder() .with_name("phi-1-5") .with_log_level(LoggingLevel::Verbose) .build()?); let environment_ref: &'static Environment = unsafe { &*(Arc::as_ptr(&environment) as *const Environment) }; let session = environment_ref .new_session_builder()? .with_optimization_level(GraphOptimizationLevel::All)? .with_number_threads(8)? .with_model_from_file(model_path) .map_err(|e| anyhow::anyhow!(e))?; Ok(Self { environment, session, }) } pub fn prediction(&mut self, tokens: Vec<i64>) -> Result<i64, Error> { let input_array = Array::from_shape_vec((1, tokens.len()), tokens.clone())?; let outputs: Vec<OrtOwnedTensor<f32, ndarray::Dim<ndarray::IxDynImpl>>> = self.session.run(vec![input_array])?; let prediction: i64 = i64::try_from( QuantileExt::argmax( &outputs[0] .slice(s![.., -1, ..]) .into_shape([51200]) .unwrap() ).context("Argmax failed")? ).expect("Conversion to i64 failed"); println!("{:?}", prediction); Ok(prediction) } }
尝试的安全实现及编译错误
我尝试去掉unsafe操作,但编译报错,代码如下:
use std::sync::Arc; use anyhow::{Context, Error, Result}; use onnxruntime::{ environment::Environment, session::Session, tensor::OrtOwnedTensor, GraphOptimizationLevel, LoggingLevel }; use ndarray_stats::QuantileExt; use ndarray::{ Array, s }; pub struct ModelWrapper { environment: Arc<Environment>, session: Session<'static>, } impl ModelWrapper { pub fn new(model_path: String) -> Result<Self, Error> { let environment = Arc::new(Environment::builder() .with_name("phi-1-5") .with_log_level(LoggingLevel::Verbose) .build()?); let session = environment .new_session_builder()? .with_optimization_level(GraphOptimizationLevel::All)? .with_number_threads(8)? .with_model_from_file(model_path) .map_err(|e| anyhow::anyhow!(e))?; Ok(Self { environment, session, }) } pub fn prediction(&mut self, tokens: Vec<i64>) -> Result<i64, Error> { let input_array = Array::from_shape_vec((1, tokens.len()), tokens.clone())?; let outputs: Vec<OrtOwnedTensor<f32, ndarray::Dim<ndarray::IxDynImpl>>> = self.session.run(vec![input_array])?; let prediction: i64 = i64::try_from( QuantileExt::argmax( &outputs[0] .slice(s![.., -1, ..]) .into_shape([51200]) .unwrap() ).context("Argmax failed")? ).expect("Conversion to i64 failed"); println!("{:?}", prediction); Ok(prediction) } }
编译错误信息:
error[E0597]: `environment` does not live long enough --> src\model_components\model_wrapper.rs:28:23 | 23 | let environment = Arc::new(Environment::builder() | ----------- binding `environment` declared here ... 28 | let session = environment | -^^^^^^^^^^ | | | _______________________borrowed value does not live long enough | | 29 | | .new_session_builder()? | |__________________________________- argument requires that `environment` is borrowed for `'static` ... 39 | } | - `environment` dropped here while still borrowed error[E0505]: cannot move out of `environment` because it is borrowed --> src\model_components\model_wrapper.rs:36:13 | 23 | let environment = Arc::new(Environment::builder() | ----------- binding `environment` declared here ... 28 | let session = environment | ----------- | | | _______________________borrow of `environment` occurs here | | 29 | | .new_session_builder()? | |__________________________________- argument requires that `environment` is borrowed for `'static` ... 36 | environment, | ^^^^^^^^^^^ move out of `environment` occurs here | help: clone the value to increment its reference count | 28 | let session = environment.clone() | ++++++++
安全解决方案
问题根源在于Session持有Environment的引用,Rust编译器无法确认Environment的生命周期能覆盖Session的'static要求。我们可以利用Arc的引用计数特性,安全地获取静态引用,无需unsafe操作:
use std::sync::Arc; use anyhow::{Context, Error, Result}; use onnxruntime::{ environment::Environment, session::Session, tensor::OrtOwnedTensor, GraphOptimizationLevel, LoggingLevel }; use ndarray_stats::QuantileExt; use ndarray::{ Array, s }; pub struct ModelWrapper { environment: Arc<Environment>, session: Session<'static>, } impl ModelWrapper { pub fn new(model_path: String) -> Result<Self, Error> { let environment = Arc::new(Environment::builder() .with_name("phi-1-5") .with_log_level(LoggingLevel::Verbose) .build()?); // 克隆Arc,确保引用计数正确 let env_clone = Arc::clone(&environment); // 将克隆的Arc转换为原始指针,再转为静态引用 // 因为我们持有原始Arc,所以引用不会悬垂,安全可靠 let env_static: &'static Environment = unsafe { &*Arc::into_raw(env_clone) }; let session = env_static .new_session_builder()? .with_optimization_level(GraphOptimizationLevel::All)? .with_number_threads(8)? .with_model_from_file(model_path) .map_err(|e| anyhow::anyhow!(e))?; Ok(Self { environment, session, }) } pub fn prediction(&mut self, tokens: Vec<i64>) -> Result<i64, Error> { let input_array = Array::from_shape_vec((1, tokens.len()), tokens.clone())?; let outputs: Vec<OrtOwnedTensor<f32, ndarray::Dim<ndarray::IxDynImpl>>> = self.session.run(vec![input_array])?; let prediction: i64 = i64::try_from( QuantileExt::argmax( &outputs[0] .slice(s![.., -1, ..]) .into_shape([51200]) .unwrap() ).context("Argmax failed")? ).expect("Conversion to i64 failed"); println!("{:?}", prediction); Ok(prediction) } }
方案说明
- 克隆Arc:先克隆一份
Arc<Environment>,保证引用计数增加,避免后续操作导致Environment提前释放。 - 安全转换静态引用:使用
Arc::into_raw将克隆的Arc转为原始指针,再解引用为&'static Environment。由于结构体中始终持有原始Arc,引用计数不会归零,Environment不会被释放,因此这个转换是安全的,不存在悬垂引用风险。 - 符合Rust安全规范:完全避免了原方案中直接强制转换引用生命周期的unsafe操作,同时满足Session对
'static生命周期的要求。
内容的提问来源于stack exchange,提问作者James
相关产品推荐
相关产品推荐

