use super::*;
#[derive(Clone)]
pub(super) struct ProviderArtifactRequirement {
config: ExecutorArtifactConfig,
state: Arc<dyn onnx_runtime_ep_api::ExecutorArtifactRequirementState>,
}
impl std::fmt::Debug for ProviderArtifactRequirement {
fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
formatter
.debug_struct("ProviderArtifactRequirement")
.field("provider", &self.config.provider().get())
.field("executor", &self.config.executor().get())
.field("generation", &self.config.generation().get())
.finish_non_exhaustive()
}
}
impl ProviderArtifactRequirement {
pub(super) fn new(
config: ExecutorArtifactConfig,
state: Arc<dyn onnx_runtime_ep_api::ExecutorArtifactRequirementState>,
) -> Self {
Self { config, state }
}
pub(super) fn acquire_use(
&self,
config: ExecutorArtifactConfig,
) -> Result<Box<dyn onnx_runtime_ep_api::ExecutorArtifactUseGuard>> {
if self.config != config {
return Err(SessionError::Internal(format!(
"baked provider-artifact requirement names provider {} executor {} generation {}, \
but the active private session scope is provider {} executor {} generation {}",
self.config.provider().get(),
self.config.executor().get(),
self.config.generation().get(),
config.provider().get(),
config.executor().get(),
config.generation().get(),
)));
}
self.state.acquire_use().map_err(Into::into)
}
}
#[derive(Clone, Debug, Default)]
pub(super) enum CapturedProviderArtifactRequirement {
#[default]
Uncaptured,
NeverInstalled,
Required(ProviderArtifactRequirement),
}
#[derive(Default)]
pub(crate) struct SlotCaptureState {
pub(super) device_graph_token: Option<DeviceGraphToken>,
pub(super) provider_artifact_requirement: CapturedProviderArtifactRequirement,
pub(super) device_graph_signature: Option<Vec<DeviceBindingSignature>>,
pub(super) capture_schedule: Option<CaptureSchedule>,
pub(super) capture_segmentation: Vec<CaptureDecline>,
pub(super) capture_cf_shapes: HashMap<ValueId, Vec<usize>>,
pub(super) capture_warm_signature: Option<Vec<ExternalCaptureSig>>,
pub(super) capture_warm_shapes: HashMap<ValueId, Vec<usize>>,
pub(super) capture_warm_seeded: HashMap<ValueId, Vec<usize>>,
pub(super) capture_quarantine_ops: HashSet<(String, String)>,
}
#[derive(Clone, Debug)]
enum ProviderArtifactOutcome {
Unfinalized,
Pending(ExecutorArtifactPending),
Failed(String),
Complete {
route_residency: ExecutorRouteResidency,
requirement: Option<ProviderArtifactRequirement>,
},
}
#[derive(Clone, Copy, Debug, Default, PartialEq, Eq)]
pub(super) enum ExecutorRouteResidency {
#[default]
Disabled,
Declined,
Required {
owner: ExecutorInstanceId,
},
}
impl ExecutorRouteResidency {
pub(super) fn owner(self) -> Option<ExecutorInstanceId> {
match self {
Self::Required { owner } => Some(owner),
Self::Disabled | Self::Declined => None,
}
}
pub(super) fn is_required(self) -> bool {
self.owner().is_some()
}
}
#[derive(Clone, Debug)]
pub(super) struct ProviderArtifactReadiness {
epoch: ExecutorArtifactReadinessEpoch,
outcome: ProviderArtifactOutcome,
}
impl Default for ProviderArtifactReadiness {
fn default() -> Self {
Self {
epoch: ExecutorArtifactReadinessEpoch::INITIAL,
outcome: ProviderArtifactOutcome::Unfinalized,
}
}
}
impl ProviderArtifactReadiness {
pub(super) fn checked_next_epoch(&self) -> Result<ExecutorArtifactReadinessEpoch> {
self.epoch
.get()
.checked_add(1)
.map(ExecutorArtifactReadinessEpoch::new)
.ok_or_else(|| {
EpError::KernelFailed(
"executor artifact readiness epoch space exhausted; refusing to wrap and \
authorize an ABA-stale provider artifact"
.to_string(),
)
.into()
})
}
pub(super) fn advance_to(&mut self, epoch: ExecutorArtifactReadinessEpoch) {
if epoch > self.epoch {
self.epoch = epoch;
self.outcome = ProviderArtifactOutcome::Unfinalized;
}
}
#[cfg(test)]
pub(super) fn at_epoch_for_test(epoch: u64) -> Self {
Self {
epoch: ExecutorArtifactReadinessEpoch::new(epoch),
outcome: ProviderArtifactOutcome::Unfinalized,
}
}
pub(super) fn needs_finalization(&self) -> bool {
matches!(self.outcome, ProviderArtifactOutcome::Unfinalized)
}
pub(super) fn finalize_if_needed(
&mut self,
ep: &dyn ExecutionProvider,
config: ExecutorArtifactConfig,
graph: &Graph,
finalized_banks: &[FinalizedExpertBank],
) -> Result<()> {
let executor = config.executor();
if self.needs_finalization() {
match ep.inspect_executor_artifacts(
config.provider(),
config.executor(),
config.generation(),
self.epoch,
graph,
finalized_banks,
) {
Ok(report) => {
if report.provider() != config.provider()
|| report.executor() != config.executor()
|| report.generation() != config.generation()
|| report.readiness() != self.epoch
{
self.outcome = ProviderArtifactOutcome::Failed(format!(
"provider artifact report mismatch: returned provider {} executor {} \
generation {} epoch {}, expected provider {} executor {} generation {} \
epoch {}",
report.provider().get(),
report.executor().get(),
report.generation().get(),
report.readiness().get(),
config.provider().get(),
config.executor().get(),
config.generation().get(),
self.epoch.get(),
));
} else {
self.outcome = match (config.route_residency(), report.into_state()) {
(
ExecutorRouteResidencyConfig::Disabled,
ExecutorArtifactState::Disabled,
) => ProviderArtifactOutcome::Complete {
route_residency: ExecutorRouteResidency::Disabled,
requirement: None,
},
(
ExecutorRouteResidencyConfig::Enabled,
ExecutorArtifactState::Declined,
) => ProviderArtifactOutcome::Complete {
route_residency: ExecutorRouteResidency::Declined,
requirement: None,
},
(
ExecutorRouteResidencyConfig::Enabled,
ExecutorArtifactState::Required,
) => match ep.executor_artifact_requirement(
config.provider(),
config.executor(),
config.generation(),
) {
Ok(Some(state)) => ProviderArtifactOutcome::Complete {
route_residency: ExecutorRouteResidency::Required {
owner: executor,
},
requirement: Some(ProviderArtifactRequirement::new(
config, state,
)),
},
Ok(None) => ProviderArtifactOutcome::Failed(format!(
"{} reported required artifacts for executor {} generation {}, \
but retained no exact use requirement",
ep.name(),
config.executor().get(),
config.generation().get(),
)),
Err(error) => ProviderArtifactOutcome::Failed(error.to_string()),
},
(
ExecutorRouteResidencyConfig::Enabled,
ExecutorArtifactState::Pending(pending),
) => ProviderArtifactOutcome::Pending(pending),
(policy, state) => ProviderArtifactOutcome::Failed(format!(
"provider artifact state {state:?} is incompatible with immutable \
route-residency policy {policy:?}",
)),
};
}
}
Err(error) => {
self.outcome = ProviderArtifactOutcome::Failed(error.to_string());
}
}
}
self.require_complete(ep.name(), executor)
}
pub(super) fn acquire_use(
&self,
ep: &dyn ExecutionProvider,
config: ExecutorArtifactConfig,
exact_requirement: Option<&CapturedProviderArtifactRequirement>,
) -> Result<Option<Box<dyn onnx_runtime_ep_api::ExecutorArtifactUseGuard>>> {
self.require_complete(ep.name(), config.executor())?;
let current = match &self.outcome {
ProviderArtifactOutcome::Complete { requirement, .. } => requirement.as_ref(),
_ => None,
};
let requirement = match exact_requirement {
None => current,
Some(CapturedProviderArtifactRequirement::NeverInstalled) => None,
Some(CapturedProviderArtifactRequirement::Required(requirement)) => Some(requirement),
Some(CapturedProviderArtifactRequirement::Uncaptured) => {
return Err(SessionError::Internal(
"device graph has no captured provider-artifact requirement".into(),
));
}
};
requirement
.map(|requirement| requirement.acquire_use(config))
.transpose()
}
fn requirement(&self) -> Option<&ProviderArtifactRequirement> {
match &self.outcome {
ProviderArtifactOutcome::Complete { requirement, .. } => requirement.as_ref(),
_ => None,
}
}
pub(super) fn captured_requirement(&self) -> CapturedProviderArtifactRequirement {
match self.requirement() {
Some(requirement) => CapturedProviderArtifactRequirement::Required(requirement.clone()),
None => CapturedProviderArtifactRequirement::NeverInstalled,
}
}
pub(super) fn require_complete(
&self,
provider: &str,
executor: ExecutorInstanceId,
) -> Result<()> {
match &self.outcome {
ProviderArtifactOutcome::Complete { .. } => Ok(()),
ProviderArtifactOutcome::Unfinalized => {
Err(SessionError::ExecutionProviderArtifactsPending {
provider: provider.to_string(),
executor: executor.get(),
readiness_epoch: self.epoch.get(),
reason: "provider artifact finalization has not reached a terminal outcome"
.to_string(),
})
}
ProviderArtifactOutcome::Pending(pending) => {
Err(SessionError::ExecutionProviderArtifactsPending {
provider: provider.to_string(),
executor: executor.get(),
readiness_epoch: self.epoch.get(),
reason: pending.reason(),
})
}
ProviderArtifactOutcome::Failed(reason) => {
Err(SessionError::ExecutionProviderArtifactFinalizationFailed {
provider: provider.to_string(),
executor: executor.get(),
readiness_epoch: self.epoch.get(),
reason: reason.clone(),
})
}
}
}
pub(super) fn route_residency(
&self,
provider: &str,
executor: ExecutorInstanceId,
) -> Result<ExecutorRouteResidency> {
self.require_complete(provider, executor)?;
match self.outcome {
ProviderArtifactOutcome::Complete {
route_residency, ..
} => Ok(route_residency),
_ => unreachable!("require_complete accepted only a complete outcome"),
}
}
}
pub(crate) struct Executor {
pub(super) instance_id: ExecutorInstanceId,
pub(super) artifact_config: ExecutorArtifactConfig,
pub(super) artifact_teardown_armed: bool,
pub(super) graph: Graph,
pub(super) weights: Arc<WeightStore>,
pub(super) ep: Arc<dyn ExecutionProvider>,
pub(super) heterogeneous: Option<Box<crate::hetero::HeterogeneousExecutor>>,
pub(super) graph_slot: DeviceGraphSlot,
pub(super) graph_owner: DeviceGraphOwner,
pub(super) validation_registration: Option<DeviceValidationRegistration>,
pub(super) pending_device_validation: Option<DeviceValidationToken>,
pub(super) weight_handles: HashMap<ValueId, WeightHandle>,
#[allow(dead_code)]
pub(super) expert_region_candidates: HashMap<ValueId, onnx_runtime_loader::WeightRegionCatalog>,
pub(super) finalized_expert_banks: Vec<FinalizedExpertBank>,
#[allow(dead_code)]
pub(super) residency_plan: onnx_runtime_ep_api::ResidencyPlan,
pub(super) prefetch_issue_nodes: std::sync::Mutex<HashMap<ValueId, usize>>,
pub(super) prefetch_lookahead_nodes: usize,
pub(super) buffers: HashMap<ValueId, DeviceBuffer>,
pub(super) buffer_shapes: HashMap<ValueId, Vec<usize>>,
pub(super) parked_input_buffers: Vec<(ValueId, DeviceBuffer)>,
pub(super) capture_deferred_frees: Vec<DeviceBuffer>,
pub(super) value_shapes: HashMap<ValueId, Shape>,
pub(super) value_dtypes: HashMap<ValueId, DataType>,
pub(super) plan: Vec<NodePlan>,
pub(super) input_index: HashMap<String, ValueId>,
pub(super) required_inputs: Vec<ValueId>,
pub(super) has_symbols: bool,
pub(super) cache: KernelCache,
pub(super) name_index: HashMap<String, ValueId>,
pub(super) subgraph_execs: HashMap<(NodeId, String), ChildExecutor>,
pub(super) control_flow_stats: ControlFlowStats,
pub(super) if_last_predicate: HashMap<NodeId, bool>,
pub(super) slot_capture: [SlotCaptureState; DeviceGraphSlot::COUNT],
pub(super) control_flow_output_values: HashSet<ValueId>,
pub(super) capture_growing_symbols: HashSet<SymbolId>,
pub(super) capacity_pinned_kv_symbols: HashSet<SymbolId>,
pub(super) last_capture_failed_node: Option<NodeId>,
pub(super) views: HashMap<ValueId, ValueView>,
pub(super) pinned: HashSet<ValueId>,
pub(super) sequence_values: HashSet<ValueId>,
pub(super) activation_memory_plan: Option<ActivationMemoryPlanStats>,
pub(super) shared_buffers: HashMap<ValueId, Arc<SharedTensorBuffer>>,
pub(super) sequences: HashMap<ValueId, SequenceValue>,
pub(super) seq_elem_values: HashMap<ValueId, SeqTensor>,
pub(super) execution_provider_fallback_report: Option<ExecutionProviderFallbackReport>,
pub(super) trace: TraceContext,
pub(super) scratch_input_shapes: Vec<Vec<usize>>,
pub(super) scratch_input_infos: Vec<InInfo>,
pub(super) scratch_output_shapes: Vec<Vec<usize>>,
pub(super) scratch_output_strides: Vec<Vec<i64>>,
pub(super) scratch_materialized_inputs: Vec<Option<(Vec<u8>, Vec<i64>)>>,
pub(super) scratch_external_bindings: ExternalBindings,
pub(super) scratch_resolved_shapes: HashMap<ValueId, Vec<usize>>,
pub(super) all_value_ids: Vec<ValueId>,
pub(super) decode_memo_enabled: bool,
pub(super) decode_memo_verify: bool,
pub(super) decode_memo: Option<DecodePlanMemo>,
pub(super) decode_memo_prev_bindings: Option<HashMap<SymbolId, usize>>,
pub(super) decode_memo_last_action: DecodeMemoAction,
pub(super) decode_memo_resolved: HashMap<ValueId, Vec<usize>>,
pub(super) decode_memo_primed_count: u64,
pub(super) decode_memo_rebuilt_count: u64,
pub(super) decode_memo_replayed_count: u64,
pub(super) decode_memo_ineligible_count: u64,
pub(super) decode_view_plan: Option<DecodeViewPlan>,
pub(super) decode_views_reused_count: u64,
pub(super) decode_dispatch_elided_count: u64,
pub(super) decode_view_plan_sig_mismatch_streak: u32,
pub(super) decode_view_plan_disabled: bool,
pub(super) compute_in_place_enabled: bool,
pub(super) release_dead_values_enabled: bool,
pub(super) compute_in_place_alias_count: u64,
pub(super) scan_inline_single_trip_enabled: bool,
pub(super) scan_inline_single_trip_count: u64,
pub(super) kernel_bindings: Vec<Option<KernelKey>>,
pub(super) provider_artifact_readiness: ProviderArtifactReadiness,
pub(super) persistent_workspace: Option<PreparedWorkspace>,
pub(super) step_workspace: Option<PreparedWorkspace>,
pub(super) pin_step_workspace: bool,
pub(super) inherited_workspace: Option<(usize, usize)>,
pub(super) workspace_preparation_required: bool,
}
pub(super) struct PreparedWorkspace {
pub(super) buffer: WorkspaceAllocation,
pub(super) bytes: usize,
pub(super) alignment: usize,
}
pub(super) const STAGE2_SIG_MISMATCH_LIMIT: u32 = 2;
#[derive(Clone, Debug)]
pub(super) struct ValueView {
pub(super) source: ValueId,
pub(super) shape: Vec<usize>,
pub(super) strides: Vec<i64>,
pub(super) byte_offset: usize,
}
pub(super) struct DecodePlanMemo {
pub(super) reference_bindings: HashMap<SymbolId, usize>,
pub(super) decode_varying: HashSet<SymbolId>,
pub(super) invariant_shapes: HashMap<ValueId, Vec<usize>>,
pub(super) variant_values: Vec<ValueId>,
pub(super) canonical: HashSet<ValueId>,
pub(super) reference_external_sig: Vec<DecodeBindingSig>,
}
#[derive(Clone, PartialEq, Eq)]
pub(super) struct DecodeBindingSig {
pub(super) vid: ValueId,
pub(super) is_input: bool,
pub(super) dtype: DataType,
pub(super) decl_shape: Shape,
}
impl DecodePlanMemo {
pub(super) fn matches(
&self,
bindings: &HashMap<SymbolId, usize>,
external_sig: &[DecodeBindingSig],
) -> bool {
if external_sig != self.reference_external_sig {
return false;
}
if bindings.len() != self.reference_bindings.len() {
return false;
}
bindings.iter().all(|(sym, &val)| {
match self.reference_bindings.get(sym) {
Some(&ref_val) => val == ref_val || self.decode_varying.contains(sym),
None => false,
}
})
}
}
pub(super) struct DecodeViewPlan {
pub(super) elided_nodes: HashSet<usize>,
pub(super) retained_views: Vec<(ValueId, ValueView)>,
pub(super) pinned_sources: Vec<ValueId>,
pub(super) source_buffer_sig: Vec<(ValueId, usize, usize)>,
pub(super) validated: bool,
}
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub(super) enum DecodeMemoAction {
Disabled,
Primed,
Rebuilt,
Replayed,
}
pub(super) fn same_symbol_keys(a: &HashMap<SymbolId, usize>, b: &HashMap<SymbolId, usize>) -> bool {
a.len() == b.len() && a.keys().all(|k| b.contains_key(k))
}
pub(super) fn is_decode_growth_transition(
prev: &HashMap<SymbolId, usize>,
cur: &HashMap<SymbolId, usize>,
) -> bool {
if !same_symbol_keys(prev, cur) {
return false;
}
let mut any_grew = false;
for (sym, &c) in cur {
let p = prev[sym];
if c > p {
any_grew = true;
} else if c < p {
return false; }
}
any_grew
}
pub(super) fn shape_references_any(shape: &Shape, symbols: &HashSet<SymbolId>) -> bool {
shape
.iter()
.any(|d| matches!(d, Dim::Symbolic(s) if symbols.contains(s)))
}
pub(super) fn decode_memo_env_enabled() -> bool {
match std::env::var("ONNX_GENAI_DECODE_MEMO") {
Ok(value) => !matches!(
value.trim().to_ascii_lowercase().as_str(),
"0" | "false" | "off"
),
Err(_) => true,
}
}
pub(super) fn compute_in_place_env_enabled() -> bool {
match std::env::var("ONNX_GENAI_COMPUTE_IN_PLACE") {
Ok(value) => !matches!(
value.trim().to_ascii_lowercase().as_str(),
"0" | "false" | "off"
),
Err(_) => true,
}
}
pub(super) fn decode_memo_verify_env_enabled() -> bool {
matches!(
std::env::var("ONNX_GENAI_DECODE_MEMO_VERIFY")
.ok()
.as_deref(),
Some("1") | Some("true") | Some("on")
)
}
pub(super) fn scan_inline_single_trip_env_enabled() -> bool {
matches!(
std::env::var("ONNX_GENAI_SCAN_INLINE_SINGLE_TRIP")
.ok()
.as_deref()
.map(str::trim)
.map(str::to_ascii_lowercase)
.as_deref(),
Some("1") | Some("true") | Some("on")
)
}
pub(super) struct InInfo {
pub(super) present: bool,
pub(super) dtype: DataType,
pub(super) shape: Vec<usize>,
pub(super) strides: Vec<i64>,
pub(super) byte_offset: usize,
pub(super) base_ptr: usize,
pub(super) device: onnx_runtime_ir::DeviceId,
pub(super) backing: TensorBacking,
pub(super) root_len: usize,
pub(super) lazy_unresolved: bool,
pub(super) prepared_unresolved: bool,
}
#[derive(Clone)]
pub(super) struct ExternalValue {
pub(super) dtype: DataType,
pub(super) shape: Vec<usize>,
pub(super) accepts_subshape: bool,
pub(super) strides: Option<Vec<i64>>,
pub(super) fixed_stride_shape: Option<Vec<usize>>,
pub(super) ptr: usize,
pub(super) len: usize,
pub(super) alignment: usize,
pub(super) device: onnx_runtime_ir::DeviceId,
}
impl ExternalValue {
pub(super) fn accepts_output(&self, dtype: DataType, shape: &[usize], bytes: usize) -> bool {
self.dtype == dtype
&& self.len >= bytes
&& if self.accepts_subshape {
shape.len() == self.shape.len()
&& shape
.iter()
.zip(&self.shape)
.all(|(&required, &capacity)| required <= capacity)
} else {
self.shape == shape
}
}
pub(super) fn writable_buffer(&self) -> Result<DeviceBuffer> {
unsafe {
DeviceBuffer::from_borrowed_mut_parts(
self.ptr as *mut std::ffi::c_void,
self.device,
self.len,
self.alignment,
)
}
.ok_or_else(|| SessionError::Internal("external output binding has a null pointer".into()))
}
pub(super) fn readable_buffer(&self) -> Result<DeviceBuffer> {
if self.ptr == 0 {
return Err(SessionError::Internal(
"external input binding has a null pointer".into(),
));
}
Ok(unsafe {
DeviceBuffer::from_borrowed_parts(
self.ptr as *mut std::ffi::c_void,
self.device,
self.len,
self.alignment,
)
})
}
}
#[derive(Default)]
pub(super) struct ExternalBindings {
pub(super) inputs: HashMap<ValueId, ExternalValue>,
pub(super) outputs: HashMap<ValueId, ExternalValue>,
}
#[derive(Clone, PartialEq, Eq)]
pub(super) struct ExternalCaptureSig {
pub(super) vid: ValueId,
pub(super) is_input: bool,
pub(super) dtype: DataType,
pub(super) shape: Vec<usize>,
pub(super) ptr: usize,
pub(super) len: usize,
}
impl ExternalBindings {
fn capture_shape(value: &ExternalValue) -> &[usize] {
value.fixed_stride_shape.as_deref().unwrap_or(&value.shape)
}
pub(super) fn seed_capture_shapes(&self, resolved: &mut HashMap<ValueId, Vec<usize>>) {
for (&vid, value) in &self.inputs {
resolved.entry(vid).or_insert_with(|| value.shape.clone());
}
for (&vid, value) in &self.outputs {
if !value.accepts_subshape {
resolved.entry(vid).or_insert_with(|| value.shape.clone());
}
}
}
pub(super) fn capture_signature(&self) -> Vec<ExternalCaptureSig> {
let mut sig: Vec<ExternalCaptureSig> = self
.inputs
.iter()
.map(|(&vid, v)| (vid, true, v))
.chain(self.outputs.iter().map(|(&vid, v)| (vid, false, v)))
.map(|(vid, is_input, v)| ExternalCaptureSig {
vid,
is_input,
dtype: v.dtype,
shape: Self::capture_shape(v).to_vec(),
ptr: v.ptr,
len: v.len,
})
.collect();
sig.sort_by_key(|a| (a.vid.0, a.is_input));
sig
}
pub(super) fn refill_capture_signature(&self, sig: &mut Vec<ExternalCaptureSig>) {
sig.retain(|entry| {
if entry.is_input {
self.inputs.contains_key(&entry.vid)
} else {
self.outputs.contains_key(&entry.vid)
}
});
for (vid, is_input, value) in self
.inputs
.iter()
.map(|(&vid, value)| (vid, true, value))
.chain(self.outputs.iter().map(|(&vid, value)| (vid, false, value)))
{
if let Some(entry) = sig
.iter_mut()
.find(|entry| entry.vid == vid && entry.is_input == is_input)
{
entry.dtype = value.dtype;
entry.shape.clear();
entry.shape.extend_from_slice(Self::capture_shape(value));
entry.ptr = value.ptr;
entry.len = value.len;
} else {
sig.push(ExternalCaptureSig {
vid,
is_input,
dtype: value.dtype,
shape: Self::capture_shape(value).to_vec(),
ptr: value.ptr,
len: value.len,
});
}
}
sig.sort_by_key(|entry| (entry.vid.0, entry.is_input));
}
}
pub(super) struct CompiledChildPlan {
pub(super) exec: Executor,
pub(super) signature: Vec<ChildInputSignature>,
}
pub(super) const CHILD_EXECUTOR_CACHE_CAPACITY: usize = 4;
#[derive(Clone, Debug, PartialEq, Eq)]
pub(super) struct ChildInputSignature {
pub(super) dtype: DataType,
pub(super) shape: Vec<usize>,
}
pub(crate) struct ChildExecutor {
pub(super) name: String,
pub(super) template: Graph,
pub(super) inherited_opsets: HashMap<String, u64>,
pub(super) weights: Arc<WeightStore>,
pub(super) ep: Arc<dyn ExecutionProvider>,
pub(super) formal_names: Vec<String>,
pub(super) capture_names: Vec<String>,
pub(super) input_names: Vec<String>,
pub(super) compiled: Vec<CompiledChildPlan>,
pub(super) builds: u64,
pub(super) runs: u64,
pub(super) trace: TraceContext,
pub(super) release_dead_values: bool,
}
#[derive(Clone, Copy, Debug, Default, PartialEq, Eq)]
pub(crate) struct ChildExecutorStats {
pub builds: u64,
pub runs: u64,
}
pub(super) struct PreparedSubgraph {
pub(super) key: (NodeId, String),
pub(super) captures: HashMap<String, Tensor>,
}