Skip to content

Commit 361eb08

Browse files
authored
fix: manage RuntimeConfig, RuntimeCache with Arc/Rc (#102)
* fix: typo in argument of `create_execution_context_with_config` * refactor!: put RuntimeConfig for ExecutionContext into a Rc * refactor: put RuntimeCache into an Arc<Mutex>
1 parent 9975d23 commit 361eb08

3 files changed

Lines changed: 28 additions & 10 deletions

File tree

trtx/src/cuda_engine.rs

Lines changed: 7 additions & 6 deletions
Original file line numberDiff line numberDiff line change
@@ -3,6 +3,7 @@
33
//! [`CudaEngine`] wraps [`nvinfer1::ICudaEngine`] (C++ [`nvinfer1::ICudaEngine`](https://docs.nvidia.com/deeplearning/tensorrt-rtx/latest/_static/cpp-api/classnvinfer1_1_1_i_cuda_engine.html)).
44
//! [`SerializationConfig`] wraps [`nvinfer1::ISerializationConfig`] (C++ [`nvinfer1::ISerializationConfig`](https://docs.nvidia.com/deeplearning/tensorrt-rtx/latest/_static/cpp-api/classnvinfer1_1_1_i_serialization_config.html)).
55
6+
use std::rc::Rc;
67
use std::{ffi::CStr, marker::PhantomData};
78

89
use crate::engine_inspector::EngineInspector;
@@ -349,16 +350,16 @@ impl<'engine> CudaEngine<'engine> {
349350
.inner
350351
.pin_mut()
351352
.createExecutionContext(nvinfer1::ExecutionContextAllocationStrategy::kSTATIC);
352-
Ok(unsafe { ExecutionContext::from_ptr(context_ptr)? })
353+
Ok(unsafe { ExecutionContext::from_ptr(context_ptr, None)? })
353354
}
354355
#[cfg(feature = "mock_runtime")]
355-
Ok(unsafe { ExecutionContext::from_ptr(std::ptr::null_mut())? })
356+
Ok(unsafe { ExecutionContext::from_ptr(std::ptr::null_mut(), None)? })
356357
}
357358

358359
/// See [nvinfer1::ICudaEngine::createExecutionContext1]
359360
pub fn create_execution_context_with_config(
360361
&'_ mut self,
361-
runtime_conifg: &'engine RuntimeConfig<'engine>,
362+
runtime_config: Rc<RuntimeConfig<'engine>>,
362363
) -> Result<ExecutionContext<'engine>> {
363364
#[cfg(not(feature = "mock_runtime"))]
364365
{
@@ -367,12 +368,12 @@ impl<'engine> CudaEngine<'engine> {
367368
let context_ptr = unsafe {
368369
self.inner
369370
.pin_mut()
370-
.createExecutionContext1(runtime_conifg.inner.as_mut_ptr())
371+
.createExecutionContext1(runtime_config.inner.as_mut_ptr())
371372
};
372-
Ok(unsafe { ExecutionContext::from_ptr(context_ptr)? })
373+
Ok(unsafe { ExecutionContext::from_ptr(context_ptr, Some(runtime_config))? })
373374
}
374375
#[cfg(feature = "mock_runtime")]
375-
Ok(unsafe { ExecutionContext::from_ptr(std::ptr::null_mut())? })
376+
Ok(unsafe { ExecutionContext::from_ptr(std::ptr::null_mut(), None)? })
376377
}
377378

378379
/// See [nvinfer1::ICudaEngine::createSerializationConfig]

trtx/src/execution_context.rs

Lines changed: 8 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -1,5 +1,6 @@
11
use std::ffi::{CStr, CString};
22
use std::pin::Pin;
3+
use std::rc::Rc;
34

45
use cxx::UniquePtr;
56
use trtx_sys::nvinfer1;
@@ -10,14 +11,16 @@ use crate::error::{Error, PropertySetAttempt, Result};
1011
use crate::interfaces::{
1112
DebugListener, ErrorRecorder, ProcessDebugTensor, Profiler, RecordError, ReportLayerTime,
1213
};
14+
use crate::RuntimeConfig;
1315

1416
/// [`trtx_sys::nvinfer1::IExecutionContext`] — C++ [`nvinfer1::IExecutionContext`](https://docs.nvidia.com/deeplearning/tensorrt-rtx/latest/_static/cpp-api/classnvinfer1_1_1_i_execution_context.html).
1517
///
1618
/// `inner` is declared last so it is dropped first (see [`Drop`]): TensorRT must release
1719
/// [DebugListener] / [Profiler]
1820
/// pointers before their Rust wrappers run destructors.
19-
pub struct ExecutionContext<'a> {
20-
_engine: std::marker::PhantomData<&'a CudaEngine<'a>>,
21+
pub struct ExecutionContext<'engine> {
22+
_engine: std::marker::PhantomData<&'engine CudaEngine<'engine>>,
23+
_config: Option<Rc<RuntimeConfig<'engine>>>,
2124
debug_listener: Option<Pin<Box<DebugListener>>>,
2225
profiler: Option<Pin<Box<Profiler>>>,
2326
error_recorder: Option<Pin<Box<ErrorRecorder>>>,
@@ -32,9 +35,10 @@ impl std::fmt::Debug for ExecutionContext<'_> {
3235
}
3336
}
3437

35-
impl<'a> ExecutionContext<'a> {
38+
impl<'engine> ExecutionContext<'engine> {
3639
pub(crate) unsafe fn from_ptr(
3740
execution_context: *mut nvinfer1::IExecutionContext,
41+
config: Option<Rc<RuntimeConfig<'engine>>>,
3842
) -> Result<Self> {
3943
#[cfg(not(feature = "mock_runtime"))]
4044
if execution_context.is_null() {
@@ -44,6 +48,7 @@ impl<'a> ExecutionContext<'a> {
4448
}
4549
Ok(ExecutionContext {
4650
_engine: Default::default(),
51+
_config: config,
4752
debug_listener: None,
4853
error_recorder: None,
4954
profiler: None,

trtx/src/runtime_config.rs

Lines changed: 13 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -3,6 +3,8 @@
33
//! [`RuntimeConfig`] wraps [`trtx_sys::nvinfer1::IRuntimeConfig`] (C++ [`nvinfer1::IRuntimeConfig`](https://docs.nvidia.com/deeplearning/tensorrt-rtx/latest/_static/cpp-api/classnvinfer1_1_1_i_runtime_config.html)).
44
55
use std::marker::PhantomData;
6+
#[cfg(not(feature = "enterprise"))]
7+
use std::sync::{Arc, Mutex};
68

79
#[cfg(not(feature = "enterprise"))]
810
use crate::error::PropertySetAttempt;
@@ -21,6 +23,11 @@ use trtx_sys::{CudaGraphStrategy, DynamicShapesKernelSpecializationStrategy};
2123
pub struct RuntimeConfig<'engine> {
2224
pub(crate) inner: UniquePtr<IRuntimeConfig>,
2325
_engine: PhantomData<&'engine nvinfer1::ICudaEngine>,
26+
// actually IRuntimeCache has its mutex, so we could omit this if we made mut methods of RuntimeCache (e.g. deserialize &self)
27+
// this also makes it safe when we modify through our mutex, while cpp calls are made through
28+
// IExecution calls
29+
#[cfg(not(feature = "enterprise"))]
30+
_cache: Option<Arc<Mutex<RuntimeCache<'engine>>>>,
2431
}
2532

2633
impl std::fmt::Debug for RuntimeConfig<'_> {
@@ -40,6 +47,8 @@ impl<'engine> RuntimeConfig<'engine> {
4047
Ok(Self {
4148
inner: unsafe { UniquePtr::from_raw(runtime_config) },
4249
_engine: Default::default(),
50+
#[cfg(not(feature = "enterprise"))]
51+
_cache: None,
4352
})
4453
}
4554

@@ -75,14 +84,17 @@ impl<'engine> RuntimeConfig<'engine> {
7584

7685
#[cfg(not(feature = "enterprise"))]
7786
/// See [IRuntimeConfig::setRuntimeCache].
78-
pub fn set_runtime_cache(&mut self, cache: &RuntimeCache<'engine>) -> Result<()> {
87+
pub fn set_runtime_cache(&mut self, cache: Arc<Mutex<RuntimeCache<'engine>>>) -> Result<()> {
7988
if cfg!(not(feature = "mock")) {
8089
if self.inner.pin_mut().setRuntimeCache(
8190
cache
91+
.lock()
92+
.unwrap()
8293
.inner
8394
.as_ref()
8495
.expect("RuntimeCache inner must be non-null"),
8596
) {
97+
self._cache = Some(cache);
8698
Ok(())
8799
} else {
88100
Err(Error::FailedToSetProperty(

0 commit comments

Comments
 (0)