From 059a466ee4b040e0371cd085a92310a10a75787b Mon Sep 17 00:00:00 2001 From: Justin Chu Date: Thu, 30 Jul 2026 13:46:53 +0000 Subject: [PATCH 1/2] docs(pipeline): native multi-component decode refactor plan (#384) Co-authored-by: Copilot <223556219+Copilot@users.noreply.github.com> --- .../inbox/mary-native-pipeline-plan.md | 101 ++++++++++++++++++ 1 file changed, 101 insertions(+) create mode 100644 .squad/decisions/inbox/mary-native-pipeline-plan.md diff --git a/.squad/decisions/inbox/mary-native-pipeline-plan.md b/.squad/decisions/inbox/mary-native-pipeline-plan.md new file mode 100644 index 0000000000..2890399149 --- /dev/null +++ b/.squad/decisions/inbox/mary-native-pipeline-plan.md @@ -0,0 +1,101 @@ +### 2026-07-30: Native multi-component pipeline decode — refactor plan (issue #384) +**By:** Mary +**What:** Scoped plan to make `PipelineDecodeLoopBackend` drive **native** component +sessions (nxrt custom-EP), unblocking native decode of pipelined multi-component +models (Qwen3.6-35B-A3B = embedding + decoder + vision, and any multimodal pipeline). +This is the GAP 3 work referenced in `mary-35b-a3b-blocker.md`. + +#### 1. Where the pipeline loop hardcodes the concrete ORT `Session` / ORT `Value` + +Grepped `&'a Session`, `.run(`, `Session,`, `Value` across the pipeline module. +The decode loop that must become backend-neutral is **`flat_autoregressive.rs` + +`paged_decode.rs`** (only these two construct/use `PipelineDecodeLoopBackend`; +`nested_autoregressive.rs` and `iterative.rs` drive their own loops and are out of +scope for the AR-decoder path). + +Concrete-ORT couplings in `paged_decode.rs::PipelineDecodeLoopBackend`: + +| Field / call | Coupling | +|---|---| +| `decoder: &'a Session` | concrete ORT decoder session (GAP: Inc2) | +| `step_components: Vec<(StepComponentBinding, &'a Session)>` | concrete ORT **every_step** sessions (GAP: **Inc1**) | +| `pool: &'a mut PipelineTensors` = `HashMap` | pool holds ORT `Value` (the value-type seam) | +| `static_cross_kv: Vec<(String, Arc)>` | ORT `Value` cross-attn KV (Inc3) | +| `run_step_components`: builds `Vec<(String, Value)>`, calls `session.run(&refs)`, inserts `Value` | ORT value + ORT session run (Inc1) | +| `decoder_extras`: `clone_value`, `Value::alias_with_shape` | ORT `Value` decoder binding (Inc2) | +| `next_logits`: `run_decode_step_with_extra(self.decoder, ...)`, `mirror_present_kv_to_pages`, `extract_next_token_logits_with_io` | ORT decoder step + ORT KV mirror (Inc2/Inc3) | + +`flat_autoregressive.rs` pairs each binding with `self.models.session(name) -> &Session` +(lines ~152–169) and constructs the backend (line 187). `pipeline/mod.rs` imports the +concrete `Session, Value, DataType` from `onnx_genai_ort` and defines +`PipelineTensors = HashMap` (mod.rs:31–51). + +#### 2. Trait boundary — THE VALUE-TYPE SEAM (verdict) + +**Does the ORT `Session` implement `ComponentSession`?** No — but an adapter already +exists: `onnx_genai_ort::OrtComponentSession` wraps an **owned** `Session` and +implements `ComponentSession` (ort/src/component.rs:130). `NativeComponentSession` +implements the same trait (engine/src/native_component.rs:164). Both are object-safe +(`Box`). + +**Do `ComponentSession::run` inputs/outputs use a backend-neutral `Value` or ORT +`Value`?** **VERDICT: the seam is a backend-neutral, host-resident `ComponentTensor` +(raw little-endian element bytes + dtype + static shape), NOT ORT `Value` and NOT an +nxrt tensor.** `ComponentSession::run(&mut self, &[(&str, &ComponentTensor)]) -> +Vec<(String, ComponentTensor)>` (metadata/src/component.rs:283). Each backend adapter +translates at its own boundary: ORT via `Value::to_raw_bytes`/`from_raw_bytes` +(ort/src/component.rs `to_value`/`from_value`); native via `Tensor::from_raw`/`as_bytes` +(native_component.rs `to_native_tensor`/`from_native_tensor`). + +Consequence: the pipeline **pool holds ORT `Value`**, but the trait speaks +`ComponentTensor`. So routing step components through `dyn ComponentSession` requires a +**pool-`Value` ⇄ `ComponentTensor` conversion at the loop boundary** — this host +round-trip **is the crux of the work**. `DataType ⇄ ComponentDataType` `From` impls +already exist in `onnx_genai_ort`. The conversion is a host copy (`numel * dtype.size`); +for the decoder (Inc2) this is the KV-cache-sized cost and must eventually be avoided by +keeping tensors backend-native across the seam, but for small every_step embedding +outputs it is negligible (same order as the existing per-step `clone_value` the decoder +already pays). + +#### 3. Increment breakdown + +- **Inc1 (THIS task):** Route the `every_step` (step) components through + `dyn ComponentSession`. Change `step_components` to + `Vec<(StepComponentBinding, Box)>`. Rewrite + `run_step_components` to convert pool `Value → ComponentTensor`, call the trait, and + convert `ComponentTensor → Value` back into the pool. ORT default path wraps + `&Session` in a **borrowing** ORT adapter (`OrtComponentSessionRef`), behaviour + unchanged. Native path loads the component via `NativeComponentSession`. **Prove one + NATIVE every_step component (the gemma4-vlm `embedding`) runs inside the loop with + ORT-vs-native token parity** while the decoder stays ORT (hybrid). Deliverable = + solve the value seam for step components + wire embedding + parity test. +- **Inc2:** The decoder itself (`decoder: &Session`, `run_decode_step_with_extra`, + `decoder_extras`, logits extraction) becomes trait-driven. This forces the KV-cache + ownership question (see risks) and the decoder-sized value-seam copy — the real perf + seam. Reconcile with `NativeDecodeSession` (the single-graph native decode path). +- **Inc3:** Cross-component value handoff — `static_cross_kv` (encoder cross-attn KV), + device placement/transfer between components, and the vision `prompt_only` stage; the + full 35B-A3B embedding+decoder+vision chain end-to-end on native. + +#### 4. Risks + +- **Value-seam copies:** every seam crossing is a host copy. Fine for embeddings (Inc1), + costly for the decoder KV (Inc2) — Inc2 should keep tensors backend-native across the + seam or add a zero-copy fast path behind the trait rather than always round-tripping + through host bytes. +- **KV-cache ownership (Inc2):** the single-graph native path is `NativeDecodeSession`, + which owns its own persistent KV state. A pipeline-driven decoder needs the loop + (`DecodeState`, paged mirror) to own/advance KV. Reconciling these two ownership + models is the central Inc2 design decision. +- **Device placement:** native EP tensors may be device-resident; the `ComponentTensor` + seam is host-only (`as_raw_bytes`/`to_raw_bytes` reject device tensors). Cross-device + step→decoder handoff (Inc3) needs explicit host staging or a device-aware seam. +- **Feature gating (#436/#441 class):** native paths are `#[cfg(feature = + "native-backend")]`; imports must be cfg-correct so both cuda and non-cuda, + native and non-native builds compile without unused-import warnings. + +**Why:** `PipelineDecodeLoopBackend` owning ORT `Value`/`Session` is the last blocker to +native pipelined decode. Establishing the value-type seam verdict (neutral +`ComponentTensor`, not ORT `Value`) up front prevents Inc2/Inc3 from re-litigating the +boundary, and the Inc1 every_step slice proves the seam end-to-end with token parity on a +tiny CPU fixture before the heavier decoder/KV work. From 72d73748cc00c1a609f8bac16aaefdbf5e3f146e Mon Sep 17 00:00:00 2001 From: Justin Chu Date: Thu, 30 Jul 2026 14:17:54 +0000 Subject: [PATCH 2/2] feat(pipeline): drive every_step components via ComponentSession trait (native multi-component inc1) Route the pipeline decode loop's every_step (step) components through the backend-neutral ComponentSession trait instead of the concrete ORT Session, so the same run_step_components code path drives an ORT session or a native nxrt component with no forked native copy. - PipelineDecodeLoopBackend.step_components becomes Vec<(StepComponentBinding, Box)>. - run_step_components crosses the value-type seam: pool ORT Value -> neutral host ComponentTensor -> trait run -> ComponentTensor -> pool Value. - Add OrtComponentSessionRef, a borrowing ORT ComponentSession adapter, so the default path drives already-loaded sessions unchanged (behaviour-identical). - Select native every_step components at runtime via ONNX_GENAI_PIPELINE_NATIVE_STEP_COMPONENTS (empty/unset => all ORT). - Parity test: the tiny-gemma4-vlm embedding every_step component produces identical token ids [0,5,6,7] on ORT and native while the decoder stays ORT. The decoder itself and cross-component value handoff remain ORT-owned; that is inc2/inc3 (see .squad/decisions/inbox/mary-native-pipeline-plan.md). Refs #384. Co-authored-by: Copilot <223556219+Copilot@users.noreply.github.com> --- crates/onnx-genai-engine/Cargo.toml | 4 + .../src/pipeline/flat_autoregressive.rs | 29 +++- crates/onnx-genai-engine/src/pipeline/mod.rs | 111 ++++++++++++- .../src/pipeline/paged_decode.rs | 50 +++--- .../tests/native_step_component_parity.rs | 87 ++++++++++ crates/onnx-genai-ort/src/component.rs | 155 +++++++++++++++--- crates/onnx-genai-ort/src/lib.rs | 2 +- 7 files changed, 378 insertions(+), 60 deletions(-) create mode 100644 crates/onnx-genai-engine/tests/native_step_component_parity.rs diff --git a/crates/onnx-genai-engine/Cargo.toml b/crates/onnx-genai-engine/Cargo.toml index e6688834a4..3841367ce7 100644 --- a/crates/onnx-genai-engine/Cargo.toml +++ b/crates/onnx-genai-engine/Cargo.toml @@ -20,6 +20,10 @@ required-features = ["cuda", "native-backend"] name = "weight_offload_native_cuda_e2e" required-features = ["cuda", "native-backend"] +[[test]] +name = "native_step_component_parity" +required-features = ["native-backend"] + [features] default = [] native-backend = [ diff --git a/crates/onnx-genai-engine/src/pipeline/flat_autoregressive.rs b/crates/onnx-genai-engine/src/pipeline/flat_autoregressive.rs index 44c8986226..6e818c14b4 100644 --- a/crates/onnx-genai-engine/src/pipeline/flat_autoregressive.rs +++ b/crates/onnx-genai-engine/src/pipeline/flat_autoregressive.rs @@ -153,18 +153,23 @@ impl PipelineEngine { .models .session(&ar.decoder) .with_context(|| format!("pipeline decoder '{}' was not loaded", ar.decoder))?; - // Pair every `every_step` binding with its loaded session. This is the - // generic replacement for the old one-output `inputs_embeds` fusion. + // Pair every `every_step` binding with a backend-neutral component + // session. By default this borrows the already-loaded ORT session + // (behaviour unchanged); components named in + // `ONNX_GENAI_PIPELINE_NATIVE_STEP_COMPONENTS` are instead loaded and + // driven through the native nxrt backend, proving the value-type seam + // (native multi-component inc1) while the decoder stays on ORT. + let native_step_components = native_step_component_set(); let step_components = step_bindings .into_iter() .map(|binding| { - let session = self.models.session(&binding.component).with_context(|| { - format!( - "pipeline every_step component '{}' was not loaded", - binding.component - ) - })?; - Ok((binding, session)) + let component: Box = + build_step_component_session( + &self.models, + &binding.component, + &native_step_components, + )?; + Ok((binding, component)) }) .collect::>>()?; let tokenizer = self @@ -234,6 +239,12 @@ impl PipelineEngine { mirror.mirrored_tokens } }); + // The backend now owns its every_step component sessions behind + // `Box`, so it carries drop glue and its borrows of + // `tensors` / `self` would otherwise live to end of scope. Everything + // needed downstream has been copied out above, so release it explicitly + // before the paged-sequence retirement and the `tensors` move below. + drop(backend); let result = match result { Ok(result) => result, Err(error) => { diff --git a/crates/onnx-genai-engine/src/pipeline/mod.rs b/crates/onnx-genai-engine/src/pipeline/mod.rs index d822afe924..0c68a4ec81 100644 --- a/crates/onnx-genai-engine/src/pipeline/mod.rs +++ b/crates/onnx-genai-engine/src/pipeline/mod.rs @@ -365,7 +365,76 @@ fn build_native_pipeline_components( Ok(components) } -/// Components whose graphs contain only deterministic operators. +/// Every_step components the operator explicitly requested be run on the native +/// nxrt backend, from `ONNX_GENAI_PIPELINE_NATIVE_STEP_COMPONENTS` (a +/// comma-separated list of component names). Empty/unset means all every_step +/// components run on the default ORT backend, so the ORT decode path is +/// unchanged. This is the injection seam for the native multi-component inc1 +/// hybrid: it drives named every_step components natively inside an +/// otherwise-ORT pipeline decode loop, proving the value-type seam. +fn native_step_component_set() -> BTreeSet { + std::env::var("ONNX_GENAI_PIPELINE_NATIVE_STEP_COMPONENTS") + .ok() + .map(|list| { + list.split(',') + .map(str::trim) + .filter(|name| !name.is_empty()) + .map(str::to_string) + .collect() + }) + .unwrap_or_default() +} + +/// Build the backend-neutral [`ComponentSession`](onnx_genai_metadata::ComponentSession) +/// for one every_step component. +/// +/// By default this borrows the already-loaded ORT session +/// ([`OrtComponentSessionRef`]) so the ORT decode path is behaviour-identical. +/// A component named in `native_components` is instead loaded and driven through +/// the native nxrt backend, so the same decode loop drives both backends through +/// the trait with no forked code path. +fn build_step_component_session<'a>( + models: &'a PipelineModels, + component: &str, + native_components: &BTreeSet, +) -> anyhow::Result> { + if native_components.contains(component) { + #[cfg(feature = "native-backend")] + { + let path = models + .directory + .model_paths + .get(component) + .with_context(|| { + format!("native every_step component '{component}' has no model path") + })?; + // The every_step slice (inc1) stages tensors through the host + // `ComponentTensor` seam, so the native component runs on CPU; device + // placement of the pipeline decoder is the later (inc2/inc3) work. + let native = crate::native_component::NativeComponentSession::load( + path, + crate::native_decode::NativeDecodeDevice::Cpu, + ) + .with_context(|| format!("failed to load native every_step component '{component}'"))?; + return Ok(Box::new(native)); + } + #[cfg(not(feature = "native-backend"))] + { + anyhow::bail!( + "every_step component '{component}' was requested on the native backend via \ + ONNX_GENAI_PIPELINE_NATIVE_STEP_COMPONENTS, but this build was compiled without \ + the 'native-backend' feature. Rebuild with `--features native-backend`." + ); + } + } + let session = models + .session(component) + .with_context(|| format!("pipeline every_step component '{component}' was not loaded"))?; + Ok(Box::new(onnx_genai_ort::OrtComponentSessionRef::new( + session, + ))) +} + /// /// Read once at load, because a component's declared phase says when it runs, /// never that it is pure — a graph with `RandomNormal` in it would otherwise be @@ -1303,6 +1372,46 @@ fn coerce_value_to_dtype(value: &Value, target: DataType) -> anyhow::Result anyhow::Result { + let dtype = onnx_genai_metadata::ComponentDataType::from(value.dtype()); + let bytes = value.to_raw_bytes()?; + onnx_genai_metadata::ComponentTensor::from_raw(dtype, value.shape().to_vec(), bytes) + .map_err(Into::into) +} + +/// Convert a [`ComponentTensor`] produced by a component back into a pool ORT +/// [`Value`]. Inverse of [`value_to_component_tensor`]. +fn component_tensor_to_value( + tensor: &onnx_genai_metadata::ComponentTensor, +) -> anyhow::Result { + let dtype = DataType::from(tensor.dtype()); + Value::from_raw_bytes(tensor.as_bytes().to_vec(), tensor.shape(), dtype).map_err(Into::into) +} + +/// Build the running-token seed as a neutral `int64` [`ComponentTensor`] of shape +/// `[1, seq]`, matching the ORT `Value::from_slice_i64` the loop previously fed. +fn token_seed_component_tensor( + ids: &[i64], +) -> anyhow::Result { + let bytes: Vec = ids.iter().flat_map(|v| v.to_le_bytes()).collect(); + onnx_genai_metadata::ComponentTensor::from_raw( + onnx_genai_metadata::ComponentDataType::Int64, + vec![1, ids.len() as i64], + bytes, + ) + .map_err(Into::into) +} + #[derive(Debug, Clone)] struct IterativePlan { /// The component re-invoked once per step. diff --git a/crates/onnx-genai-engine/src/pipeline/paged_decode.rs b/crates/onnx-genai-engine/src/pipeline/paged_decode.rs index c299c74373..033ddeedc8 100644 --- a/crates/onnx-genai-engine/src/pipeline/paged_decode.rs +++ b/crates/onnx-genai-engine/src/pipeline/paged_decode.rs @@ -7,6 +7,7 @@ //! autoregressive driver constructs these adapters across the module boundary. use super::*; +use onnx_genai_metadata::ComponentSession; /// Whether this error is the KV page pool being full, rather than a fault. /// @@ -83,9 +84,12 @@ pub(crate) struct PipelineDecodeLoopBackend<'a> { /// Shared tensor pool: external inputs + prompt-phase outputs + the /// per-step outputs of the `every_step` components (refreshed each step). pub(crate) pool: &'a mut PipelineTensors, - /// Declared `every_step` components (with their loaded sessions), executed in - /// topological order on every step before the decoder runs. - pub(crate) step_components: Vec<(StepComponentBinding, &'a Session)>, + /// Declared `every_step` components (paired with their backend-neutral + /// [`ComponentSession`]), executed in topological order on every step before + /// the decoder runs. Boxed behind the trait so the same loop drives an ORT + /// session (via `OrtComponentSessionRef`) or a native nxrt component with no + /// forked code path — this is the value-type seam for the every_step slice. + pub(crate) step_components: Vec<(StepComponentBinding, Box)>, /// `(source_endpoint, decoder_input_port)` routing recomputed each step. pub(crate) decoder_in_edges: Vec<(String, String)>, /// Static encoder-produced cross-attention KV bound to the decoder every @@ -124,42 +128,44 @@ impl PipelineDecodeLoopBackend<'_> { return Ok(()); } let ids: Vec = seed.iter().map(|&t| i64::from(t)).collect(); - let seq = ids.len() as i64; - for (binding, session) in &self.step_components { - let mut inputs: Vec<(String, Value)> = + // Disjoint borrows: the step sessions run `&mut` while the pool is read + // for inputs and written for outputs. + let Self { + step_components, + pool, + .. + } = self; + for (binding, session) in step_components.iter_mut() { + let mut inputs: Vec<(String, onnx_genai_metadata::ComponentTensor)> = Vec::with_capacity(binding.routed_inputs.len() + 1); for routed in &binding.routed_inputs { - let value = self - .pool + let value = pool .get(&routed.endpoint) .or_else(|| { routed .routed_from .as_deref() - .and_then(|from| self.pool.get(from)) + .and_then(|from| pool.get(from)) }) .with_context(|| routed.missing_message.clone())?; - inputs.push(( - routed.port.clone(), - coerce_value_to_dtype(value, routed.dtype)?, - )); + let coerced = coerce_value_to_dtype(value, routed.dtype)?; + inputs.push((routed.port.clone(), value_to_component_tensor(&coerced)?)); } if let Some(port) = &binding.token_input { - inputs.push((port.clone(), Value::from_slice_i64(&ids, &[1, seq])?)); + inputs.push((port.clone(), token_seed_component_tensor(&ids)?)); } let refs = inputs .iter() - .map(|(name, value)| (name.as_str(), value)) + .map(|(name, tensor)| (name.as_str(), tensor)) .collect::>(); let outputs = session.run(&refs).map_err(|e| { - anyhow::anyhow!( - "ORT every_step component '{}' failed: {e}", - binding.component - ) + anyhow::anyhow!("every_step component '{}' failed: {e}", binding.component) })?; - for (name, value) in session.output_names().iter().zip(outputs) { - self.pool - .insert(format!("{}.{}", binding.component, name), value); + for (name, tensor) in outputs { + pool.insert( + format!("{}.{}", binding.component, name), + component_tensor_to_value(&tensor)?, + ); } } Ok(()) diff --git a/crates/onnx-genai-engine/tests/native_step_component_parity.rs b/crates/onnx-genai-engine/tests/native_step_component_parity.rs new file mode 100644 index 0000000000..bee2cabbb7 --- /dev/null +++ b/crates/onnx-genai-engine/tests/native_step_component_parity.rs @@ -0,0 +1,87 @@ +//! Native multi-component pipeline — increment 1 (issue #384). +//! +//! Proves the value-type seam: the `every_step` embedding component of the +//! Gemma4-style VLM composite pipeline produces **identical generated token +//! ids** whether it runs through ONNX Runtime (the default) or through the +//! native nxrt backend, while the decoder stays on ORT in both runs. +//! +//! The same decode loop drives both backends through the backend-neutral +//! [`ComponentSession`](onnx_genai_metadata::ComponentSession) trait — there is +//! no forked native copy of `run_step_components`. The native every_step +//! component is selected at runtime via +//! `ONNX_GENAI_PIPELINE_NATIVE_STEP_COMPONENTS=embedding`. +//! +//! Fixture: `scripts/build_tiny_gemma4_vlm.py`. Its closed-form head makes the +//! generated ids exact: prompt `[3, 7]` -> `[0, 5, 6, 7]`. The embedding graph +//! is an integer `Gather` plus exact-valued `Mul`/`Add` over a 0/1 table, so the +//! fused `inputs_embeds` is byte-identical across backends and the decoder — fed +//! identical inputs — samples identical tokens. + +use std::path::{Path, PathBuf}; + +use onnx_genai_engine::pipeline::PipelineGenerateRequest; +use onnx_genai_engine::{Engine, EngineConfig, GenerateOptions, GeneratePrompt, GenerateRequest}; +use onnx_genai_ort::Value; + +const NATIVE_STEP_ENV: &str = "ONNX_GENAI_PIPELINE_NATIVE_STEP_COMPONENTS"; + +fn tiny_gemma4_vlm_dir() -> PathBuf { + Path::new(env!("CARGO_MANIFEST_DIR")).join("../../tests/fixtures/tiny-gemma4-vlm") +} + +fn tiny_pixels() -> anyhow::Result { + // pixel_values[1,3,2,2] = i/12; the vision encoder means over channels. + Value::from_vec_f32((0..12).map(|i| i as f32 / 12.0).collect(), &[1, 3, 2, 2]) + .map_err(Into::into) +} + +/// One composite generation over the fixture. `native_embedding` selects whether +/// the `embedding` every_step component runs on the native backend. +fn generate_tokens(native_embedding: bool) -> anyhow::Result> { + // Process-global env: set/clear around the single engine construction that + // reads it, so the two runs in this test do not interleave. This test's + // integration binary owns the process. + if native_embedding { + unsafe { std::env::set_var(NATIVE_STEP_ENV, "embedding") }; + } else { + unsafe { std::env::remove_var(NATIVE_STEP_ENV) }; + } + + let result = (|| { + let mut engine = + Engine::from_pipeline_dir(&tiny_gemma4_vlm_dir(), EngineConfig::default())?; + let mut request = GenerateRequest::new(GeneratePrompt::TokenIds(vec![3, 7])); + request.options = GenerateOptions { + max_new_tokens: 4, + temperature: 0.0, + stop_on_eos: false, + ..GenerateOptions::default() + }; + let pipeline_request = PipelineGenerateRequest::new(request) + .with_input("vision_encoder.pixel_values", tiny_pixels()?); + let result = engine.generate_with_pipeline_request(pipeline_request)?; + Ok::<_, anyhow::Error>(result.token_ids) + })(); + + unsafe { std::env::remove_var(NATIVE_STEP_ENV) }; + result +} + +#[test] +fn native_every_step_embedding_matches_ort_token_ids() -> anyhow::Result<()> { + let ort_tokens = generate_tokens(false)?; + let native_tokens = generate_tokens(true)?; + + // Baseline the ORT path against the fixture's known closed-form ids so a + // regression that changes *both* backends identically still fails. + assert_eq!( + ort_tokens, + vec![0, 5, 6, 7], + "ORT every_step baseline drifted" + ); + assert_eq!( + native_tokens, ort_tokens, + "native every_step embedding diverged from the ORT baseline" + ); + Ok(()) +} diff --git a/crates/onnx-genai-ort/src/component.rs b/crates/onnx-genai-ort/src/component.rs index 794a9c3c04..47f52efa0e 100644 --- a/crates/onnx-genai-ort/src/component.rs +++ b/crates/onnx-genai-ort/src/component.rs @@ -127,6 +127,42 @@ fn from_value(component: &str, value: &Value) -> Result Result, ComponentError> { + // The component name is only available for diagnostics via the first + // declared output; fall back to a stable label otherwise. + let component = outputs_meta + .first() + .map(|io| io.name.as_str()) + .unwrap_or(""); + let values: Vec<(&str, Value)> = inputs + .iter() + .map(|(name, tensor)| Ok((*name, to_value(component, tensor)?))) + .collect::>()?; + let borrowed: Vec<(&str, &Value)> = values.iter().map(|(name, value)| (*name, value)).collect(); + let outputs = session + .run(&borrowed) + .map_err(|err: OrtError| ComponentError::Backend { + component: component.to_string(), + backend: BACKEND, + detail: err.to_string(), + })?; + session + .output_names() + .iter() + .zip(outputs.iter()) + .map(|(name, value)| Ok((name.clone(), from_value(component, value)?))) + .collect() +} + impl ComponentSession for OrtComponentSession { fn inputs(&self) -> &[ComponentIo] { &self.inputs @@ -140,33 +176,50 @@ impl ComponentSession for OrtComponentSession { &mut self, inputs: &[(&str, &ComponentTensor)], ) -> Result, ComponentError> { - // The component name is only available for diagnostics via the first - // declared output/input; fall back to a stable label otherwise. - let component = self - .outputs - .first() - .map(|io| io.name.as_str()) - .unwrap_or(""); - let values: Vec<(&str, Value)> = inputs - .iter() - .map(|(name, tensor)| Ok((*name, to_value(component, tensor)?))) - .collect::>()?; - let borrowed: Vec<(&str, &Value)> = - values.iter().map(|(name, value)| (*name, value)).collect(); - let outputs = - self.session - .run(&borrowed) - .map_err(|err: OrtError| ComponentError::Backend { - component: component.to_string(), - backend: BACKEND, - detail: err.to_string(), - })?; - self.session - .output_names() - .iter() - .zip(outputs.iter()) - .map(|(name, value)| Ok((name.clone(), from_value(component, value)?))) - .collect() + run_ort_component(&self.session, &self.outputs, inputs) + } +} + +/// A pipeline component backed by a **borrowed** ONNX Runtime [`Session`]. +/// +/// Behaviour-identical to [`OrtComponentSession`], but borrows a session owned +/// elsewhere (the pipeline model store keeps its sessions loaded and shared) so +/// the backend-neutral decode loop can drive an already-loaded ORT component +/// through the same [`ComponentSession`] seam it uses for native components, +/// without moving the session out of the store or reloading it per request. +pub struct OrtComponentSessionRef<'a> { + session: &'a Session, + inputs: Vec, + outputs: Vec, +} + +impl<'a> OrtComponentSessionRef<'a> { + /// Wrap a borrowed ORT [`Session`] as a backend-neutral component. + pub fn new(session: &'a Session) -> Self { + let inputs = session.inputs().iter().map(component_io).collect(); + let outputs = session.outputs().iter().map(component_io).collect(); + Self { + session, + inputs, + outputs, + } + } +} + +impl ComponentSession for OrtComponentSessionRef<'_> { + fn inputs(&self) -> &[ComponentIo] { + &self.inputs + } + + fn outputs(&self) -> &[ComponentIo] { + &self.outputs + } + + fn run( + &mut self, + inputs: &[(&str, &ComponentTensor)], + ) -> Result, ComponentError> { + run_ort_component(self.session, &self.outputs, inputs) } } @@ -258,4 +311,52 @@ mod tests { assert_eq!(tensor.shape(), reference[0].shape()); assert_eq!(tensor.as_bytes(), reference_bytes.as_slice()); } + + #[test] + fn borrowing_ref_adapter_matches_owning_adapter() { + let path = tiny_whisper_encoder_textproto(); + if !path.exists() { + eprintln!("borrowing_ref_adapter_matches_owning_adapter: fixture absent, skipping"); + return; + } + let session = Session::new( + test_environment(), + &path, + SessionOptions::default().with_intra_op_threads(1), + ) + .expect("session"); + + let input_bytes: Vec = vec![0.25f32; 80 * 8] + .iter() + .flat_map(|value| value.to_le_bytes()) + .collect(); + let input = + ComponentTensor::from_raw(ComponentDataType::Float32, vec![1, 80, 8], input_bytes) + .expect("component input"); + + // The borrowing adapter drives a session owned elsewhere; it must expose + // the same metadata and produce the same bytes as the owning adapter. + let mut borrowed = OrtComponentSessionRef::new(&session); + assert_eq!(borrowed.input_names(), vec!["input_features"]); + assert_eq!(borrowed.output_names(), vec!["encoder_hidden_states"]); + let borrowed_outputs = borrowed + .run(&[("input_features", &input)]) + .expect("borrowed run"); + + let reference = session + .run(&[( + "input_features", + &Value::from_slice_f32(&vec![0.25f32; 80 * 8], &[1, 80, 8]).expect("input"), + )]) + .expect("reference run"); + + assert_eq!(borrowed_outputs.len(), 1); + assert_eq!( + borrowed_outputs[0].1.as_bytes(), + reference[0] + .to_raw_bytes() + .expect("reference bytes") + .as_slice() + ); + } } diff --git a/crates/onnx-genai-ort/src/lib.rs b/crates/onnx-genai-ort/src/lib.rs index c514871e11..6a653b487b 100644 --- a/crates/onnx-genai-ort/src/lib.rs +++ b/crates/onnx-genai-ort/src/lib.rs @@ -33,7 +33,7 @@ pub mod value; pub use allocator::{Allocator, AllocatorType, MemoryInfo, MemoryType}; pub use binding::IoBinding; pub use chat_template::{ChatMessage, ChatRole, ChatTemplate}; -pub use component::OrtComponentSession; +pub use component::{OrtComponentSession, OrtComponentSessionRef}; pub use decode::{ BatchedDecodeSession, BatchedSharedBufferDecodeSession, BatchedStaticCacheDecodeSession, DecodeKvMode, DecodeSession, DecodeSessionOptions, DeviceSampleParams,