Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
2 changes: 2 additions & 0 deletions src/bgem3_embedding/impl.rs
Original file line number Diff line number Diff line change
Expand Up @@ -79,6 +79,7 @@ impl Bgem3Embedding {
execution_providers,
max_length,
intra_threads,
..
} = options;

let session = init_session_builder(execution_providers, intra_threads)?
Expand All @@ -99,6 +100,7 @@ impl Bgem3Embedding {
execution_providers,
max_length,
intra_threads,
..
} = options;

let session = init_session_builder(execution_providers, intra_threads)?
Expand Down
15 changes: 15 additions & 0 deletions src/reranking/impl.rs
Original file line number Diff line number Diff line change
Expand Up @@ -93,9 +93,24 @@ impl TextRerank {
execution_providers,
max_length,
intra_threads,
disable_cpu_fallback,
dimension_overrides,
} = options;

let mut session_builder = init_session_builder(execution_providers, intra_threads)?;
let builder_error = |err: ort::Error<ort::session::builder::SessionBuilder>| {
Error::OrtBuilder(err.to_string())
};
if disable_cpu_fallback {
session_builder = session_builder
.with_disable_cpu_fallback()
.map_err(builder_error)?;
}
for (name, size) in dimension_overrides {
session_builder = session_builder
.with_dimension_override(name, size)
.map_err(builder_error)?;
}
let session = match &model.onnx_source {
OnnxSource::Memory(bytes) => session_builder.commit_from_memory(bytes)?,
OnnxSource::File(path) => session_builder.commit_from_file(path)?,
Expand Down
26 changes: 25 additions & 1 deletion src/reranking/init.rs
Original file line number Diff line number Diff line change
Expand Up @@ -33,6 +33,12 @@ pub struct RerankInitOptionsUserDefined {
/// every available CPU core via `std::thread::available_parallelism`.
/// Set this to cap CPU usage (e.g. on laptops) at the cost of throughput.
pub intra_threads: Option<usize>,
/// Refuse session creation when any graph node would fall back to ORT's
/// default CPU execution provider.
pub disable_cpu_fallback: bool,
/// Override named free dimensions before ORT optimizes and places the
/// user-defined reranker graph.
pub dimension_overrides: Vec<(String, i64)>,
}

impl Default for RerankInitOptionsUserDefined {
Expand All @@ -41,6 +47,8 @@ impl Default for RerankInitOptionsUserDefined {
execution_providers: Default::default(),
max_length: DEFAULT_MAX_LENGTH,
intra_threads: None,
disable_cpu_fallback: false,
dimension_overrides: Vec::new(),
}
}
}
Expand Down Expand Up @@ -72,6 +80,16 @@ impl RerankInitOptionsUserDefined {
self.intra_threads = Some(intra_threads);
self
}

pub fn with_disable_cpu_fallback(mut self, disable: bool) -> Self {
self.disable_cpu_fallback = disable;
self
}

pub fn with_dimension_override(mut self, name: impl Into<String>, size: i64) -> Self {
self.dimension_overrides.push((name.into(), size));
self
}
}

/// Convert RerankInitOptions to RerankInitOptionsUserDefined
Expand All @@ -83,6 +101,8 @@ impl From<RerankInitOptions> for RerankInitOptionsUserDefined {
execution_providers: options.execution_providers,
max_length: options.max_length,
intra_threads: options.intra_threads,
disable_cpu_fallback: false,
dimension_overrides: Vec::new(),
}
}
}
Expand Down Expand Up @@ -143,8 +163,12 @@ mod tests {
fn userdefined_builders_set_fields() {
let o = RerankInitOptionsUserDefined::new()
.with_max_length(128)
.with_intra_threads(2);
.with_intra_threads(2)
.with_disable_cpu_fallback(true)
.with_dimension_override("sequence_length", 128);
assert_eq!(o.max_length, 128);
assert_eq!(o.intra_threads, Some(2));
assert!(o.disable_cpu_fallback);
assert_eq!(o.dimension_overrides, vec![("sequence_length".into(), 128)]);
}
}
13 changes: 13 additions & 0 deletions src/text_embedding/impl.rs
Original file line number Diff line number Diff line change
Expand Up @@ -92,6 +92,8 @@ impl TextEmbedding {
execution_providers,
max_length,
intra_threads,
disable_cpu_fallback,
dimension_overrides,
} = options;

let session = {
Expand All @@ -100,6 +102,17 @@ impl TextEmbedding {
};
let mut session_builder = init_session_builder(execution_providers, intra_threads)?;

if disable_cpu_fallback {
session_builder = session_builder
.with_disable_cpu_fallback()
.map_err(builder_error)?;
}
for (name, size) in dimension_overrides {
session_builder = session_builder
.with_dimension_override(name, size)
.map_err(builder_error)?;
}

for external_initializer_file in model.external_initializers {
session_builder = session_builder
.with_external_initializer_file_in_memory(
Expand Down
47 changes: 47 additions & 0 deletions src/text_embedding/init.rs
Original file line number Diff line number Diff line change
Expand Up @@ -31,6 +31,14 @@ pub struct InitOptionsUserDefined {
/// every available CPU core via `std::thread::available_parallelism`.
/// Set this to cap CPU usage (e.g. on laptops) at the cost of throughput.
pub intra_threads: Option<usize>,
/// Refuse session creation when any graph node would fall back to the CPU
/// execution provider. This is useful for accelerator conformance tests
/// and for applications where silent partial placement is incorrect.
pub disable_cpu_fallback: bool,
/// Override named free dimensions before ORT optimizes and places the
/// graph. Static-shape execution providers such as QNN can use this with
/// one user-defined model session per admitted shape.
pub dimension_overrides: Vec<(String, i64)>,
}

impl InitOptionsUserDefined {
Expand Down Expand Up @@ -60,6 +68,16 @@ impl InitOptionsUserDefined {
self.intra_threads = Some(intra_threads);
self
}

pub fn with_disable_cpu_fallback(mut self, disable: bool) -> Self {
self.disable_cpu_fallback = disable;
self
}

pub fn with_dimension_override(mut self, name: impl Into<String>, size: i64) -> Self {
self.dimension_overrides.push((name.into(), size));
self
}
}

impl Default for InitOptionsUserDefined {
Expand All @@ -68,6 +86,8 @@ impl Default for InitOptionsUserDefined {
execution_providers: Default::default(),
max_length: DEFAULT_MAX_LENGTH,
intra_threads: None,
disable_cpu_fallback: false,
dimension_overrides: Vec::new(),
}
}
}
Expand All @@ -81,6 +101,8 @@ impl From<TextInitOptions> for InitOptionsUserDefined {
execution_providers: options.execution_providers,
max_length: options.max_length,
intra_threads: options.intra_threads,
disable_cpu_fallback: false,
dimension_overrides: Vec::new(),
}
}
}
Expand Down Expand Up @@ -146,3 +168,28 @@ pub struct TextEmbedding {
pub(crate) quantization: QuantizationMode,
pub(crate) output_key: Option<OutputKey>,
}

#[cfg(test)]
mod tests {
use super::*;

#[test]
fn user_defined_session_controls_are_opt_in_and_composable() {
let defaults = InitOptionsUserDefined::default();
assert!(!defaults.disable_cpu_fallback);
assert!(defaults.dimension_overrides.is_empty());

let configured = InitOptionsUserDefined::new()
.with_disable_cpu_fallback(true)
.with_dimension_override("batch_size", 1)
.with_dimension_override("sequence_length", 512);
assert!(configured.disable_cpu_fallback);
assert_eq!(
configured.dimension_overrides,
[
("batch_size".to_string(), 1),
("sequence_length".to_string(), 512)
]
);
}
}