You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

安全实例化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)
    }
}

方案说明

  1. 克隆Arc:先克隆一份Arc<Environment>,保证引用计数增加,避免后续操作导致Environment提前释放。
  2. 安全转换静态引用:使用Arc::into_raw将克隆的Arc转为原始指针,再解引用为&'static Environment。由于结构体中始终持有原始Arc,引用计数不会归零,Environment不会被释放,因此这个转换是安全的,不存在悬垂引用风险。
  3. 符合Rust安全规范:完全避免了原方案中直接强制转换引用生命周期的unsafe操作,同时满足Session对'static生命周期的要求。

内容的提问来源于stack exchange,提问作者James

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.06.19 09:04:52