Skip to content

Commit 5d1483e

Browse files
committed
feat: expose HostMemoryOrVec via API
Since I needed the same primitive also for RustNN we could expose it in the trtx API. This helper is commonly when you either get owned engine bytes from TensorRT or from a cache.
1 parent 361eb08 commit 5d1483e

2 files changed

Lines changed: 68 additions & 35 deletions

File tree

trtexec-rs/src/main.rs

Lines changed: 1 addition & 35 deletions
Original file line numberDiff line numberDiff line change
@@ -10,14 +10,13 @@ use rustnn::load_graph_from_path;
1010
use std::ffi::{c_void, OsString};
1111
use std::fs::File;
1212
use std::io::{BufRead, Read, Write};
13-
use std::ops::Deref;
1413
use std::sync::atomic::AtomicU32;
1514
use std::sync::atomic::Ordering;
1615
use std::sync::Arc;
1716
use std::time::Instant;
1817
use tracing::level_filters::LevelFilter;
1918
use tracing_subscriber::{prelude::*, EnvFilter};
20-
use trtx::host_memory::HostMemory;
19+
use trtx::host_memory::HostMemoryOrVec;
2120
use trtx::{Builder, Logger, OnnxParser, ProfilingVerbosity};
2221
use trtx::{LayerInformationFormat, Runtime};
2322

@@ -75,39 +74,6 @@ fn digest_hex(digest: &md5::Digest) -> String {
7574
format!("{digest:x}")
7675
}
7776

78-
enum HostMemoryOrVec<'memory> {
79-
HostMemory(HostMemory<'memory>),
80-
Vec(Vec<u8>),
81-
}
82-
83-
impl<'memory> AsRef<[u8]> for HostMemoryOrVec<'memory> {
84-
fn as_ref(&self) -> &[u8] {
85-
match self {
86-
HostMemoryOrVec::HostMemory(host_memory) => host_memory.as_ref(),
87-
HostMemoryOrVec::Vec(items) => items.as_ref(),
88-
}
89-
}
90-
}
91-
92-
impl<'memory> Deref for HostMemoryOrVec<'memory> {
93-
type Target = [u8];
94-
95-
fn deref(&self) -> &Self::Target {
96-
self.as_ref()
97-
}
98-
}
99-
100-
impl<'buffer> From<HostMemory<'buffer>> for HostMemoryOrVec<'buffer> {
101-
fn from(value: HostMemory<'buffer>) -> Self {
102-
HostMemoryOrVec::HostMemory(value)
103-
}
104-
}
105-
impl From<Vec<u8>> for HostMemoryOrVec<'_> {
106-
fn from(value: Vec<u8>) -> Self {
107-
HostMemoryOrVec::Vec(value)
108-
}
109-
}
110-
11177
fn main() -> Result<()> {
11278
let args = Args::parse();
11379
if let Some(shell) = args.shell_completion {

trtx/src/host_memory.rs

Lines changed: 67 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -67,3 +67,70 @@ impl<'builder> Deref for HostMemory<'builder> {
6767
self.as_ref()
6868
}
6969
}
70+
71+
pub enum HostMemoryOrVec<'memory> {
72+
HostMemory(HostMemory<'memory>),
73+
Vec(Vec<u8>),
74+
}
75+
76+
impl<'memory> HostMemoryOrVec<'memory> {
77+
/// Returns `true` if the host memory or vec is [`HostMemory`].
78+
///
79+
/// [`HostMemory`]: HostMemoryOrVec::HostMemory
80+
#[must_use]
81+
pub fn is_host_memory(&self) -> bool {
82+
matches!(self, Self::HostMemory(..))
83+
}
84+
85+
pub fn as_host_memory(&self) -> Option<&HostMemory<'memory>> {
86+
if let Self::HostMemory(v) = self {
87+
Some(v)
88+
} else {
89+
None
90+
}
91+
}
92+
93+
/// Returns `true` if the host memory or vec is [`Vec`].
94+
///
95+
/// [`Vec`]: HostMemoryOrVec::Vec
96+
#[must_use]
97+
pub fn is_vec(&self) -> bool {
98+
matches!(self, Self::Vec(..))
99+
}
100+
101+
pub fn as_vec(&self) -> Option<&Vec<u8>> {
102+
if let Self::Vec(v) = self {
103+
Some(v)
104+
} else {
105+
None
106+
}
107+
}
108+
}
109+
110+
impl<'memory> AsRef<[u8]> for HostMemoryOrVec<'memory> {
111+
fn as_ref(&self) -> &[u8] {
112+
match self {
113+
HostMemoryOrVec::HostMemory(host_memory) => host_memory.as_ref(),
114+
HostMemoryOrVec::Vec(items) => items.as_ref(),
115+
}
116+
}
117+
}
118+
119+
impl<'memory> Deref for HostMemoryOrVec<'memory> {
120+
type Target = [u8];
121+
122+
fn deref(&self) -> &Self::Target {
123+
self.as_ref()
124+
}
125+
}
126+
127+
impl<'buffer> From<HostMemory<'buffer>> for HostMemoryOrVec<'buffer> {
128+
fn from(value: HostMemory<'buffer>) -> Self {
129+
HostMemoryOrVec::HostMemory(value)
130+
}
131+
}
132+
impl From<Vec<u8>> for HostMemoryOrVec<'_> {
133+
fn from(value: Vec<u8>) -> Self {
134+
HostMemoryOrVec::Vec(value)
135+
}
136+
}

0 commit comments

Comments
 (0)