Skip to main content

sim_lib_interference_compute/
provider.rs

1//! Provider routing from the runtime study seam to canonical Tensor execution.
2
3use std::sync::{Arc, Mutex};
4
5use sim_kernel::{
6    AbiVersion, Dependency, Error, Export, Lib, LibManifest, LibTarget, Linker, Result, Symbol,
7    Version,
8};
9use sim_lib_interference_core::RequestPreflight;
10use sim_lib_interference_runtime::{
11    InterferenceStudy, PhasorFieldDescriptor, PlaneDescriptor, ProblemDescriptor,
12    ReferenceStudySolver, SamplingCertificateDescriptor, SolveRequest, SolverProvider,
13    StudyDescriptor, StudyEvidenceDescriptor, StudySolver, WorkEstimateDescriptor,
14    tensor_study_solver_symbol,
15};
16use sim_lib_numbers_tensor::{TensorExecutorCard, active_tensor_executor, domains};
17
18use crate::{
19    DifferentialTolerances, LoweringPlan, PhaseBudget, TileProfile,
20    preflight::required_operation_symbols, resident::execute_resident,
21};
22
23/// Stable library symbol for the accelerated interference provider.
24pub fn interference_compute_lib_symbol() -> Symbol {
25    Symbol::qualified("sim", "interference-compute")
26}
27
28/// A reason the Tensor provider was declined before any submission.
29#[derive(Clone, Debug, PartialEq, Eq)]
30pub enum CpuFallbackReason {
31    /// No Tensor executor is active in the current environment.
32    NoExecutor,
33    /// An automatic executor selected its CPU route before submission.
34    ProviderCpuChoice,
35    /// The configured provider profile does not admit canonical f32.
36    UnsupportedDtype,
37    /// The executor card omits at least one operation required by the lowering.
38    UnsupportedOperations,
39    /// The admitted work is below the configured provider crossover.
40    BelowCrossover,
41    /// The provider limits would require host assembly of multiple output tiles.
42    MultipleOutputTiles,
43}
44
45/// Last routing outcome observed by a [`TensorStudySolver`].
46#[derive(Clone, Debug, PartialEq, Eq)]
47pub enum ProviderRoute {
48    /// No request has reached this solver.
49    Idle,
50    /// CPU was selected before provider submission.
51    Cpu(CpuFallbackReason),
52    /// The named executor completed a resident Study.
53    ProviderCompleted(Symbol),
54    /// The named executor was selected and the request failed without restart.
55    ProviderFailed(Symbol),
56}
57
58/// Stable solver counters and last route.
59#[derive(Clone, Debug, PartialEq, Eq)]
60pub struct TensorStudySnapshot {
61    /// Requests routed to the deterministic CPU reference before submission.
62    pub cpu_fallbacks: u64,
63    /// Requests completed through an active Tensor executor.
64    pub provider_completions: u64,
65    /// Selected provider requests that failed without CPU restart.
66    pub provider_failures: u64,
67    /// Most recent routing outcome.
68    pub last_route: ProviderRoute,
69}
70
71impl Default for TensorStudySnapshot {
72    fn default() -> Self {
73        Self {
74            cpu_fallbacks: 0,
75            provider_completions: 0,
76            provider_failures: 0,
77            last_route: ProviderRoute::Idle,
78        }
79    }
80}
81
82/// Admission, crossover, accuracy, and evidence settings for Tensor studies.
83#[derive(Clone, Debug, PartialEq)]
84pub struct TensorStudyConfig {
85    /// Smallest admitted cell count sent to an active provider.
86    pub min_accelerated_cells: u64,
87    /// Residual phase and predicted arithmetic limits.
88    pub phase_budget: PhaseBudget,
89    /// Provider allocation, segmentation, and dtype limits.
90    pub tile_profile: TileProfile,
91    /// Published fixed f32 comparison tolerances.
92    pub tolerances: DifferentialTolerances,
93    /// Stable adapter identity retained in Study evidence.
94    pub adapter: String,
95}
96
97impl Default for TensorStudyConfig {
98    fn default() -> Self {
99        Self {
100            min_accelerated_cells: 4_096,
101            phase_budget: PhaseBudget::default(),
102            tile_profile: TileProfile::default(),
103            tolerances: DifferentialTolerances::default(),
104            adapter: "interference-tensor-v1".to_owned(),
105        }
106    }
107}
108
109/// Study solver that preselects CPU or executes one resident Tensor plan.
110#[derive(Clone)]
111pub struct TensorStudySolver {
112    config: TensorStudyConfig,
113    snapshot: Arc<Mutex<TensorStudySnapshot>>,
114}
115
116impl TensorStudySolver {
117    /// Builds a Tensor study solver from explicit routing settings.
118    pub fn new(config: TensorStudyConfig) -> Self {
119        Self {
120            config,
121            snapshot: Arc::new(Mutex::new(TensorStudySnapshot::default())),
122        }
123    }
124
125    /// Returns a stable snapshot of routing outcomes.
126    pub fn snapshot(&self) -> TensorStudySnapshot {
127        self.snapshot
128            .lock()
129            .expect("Tensor study snapshot poisoned")
130            .clone()
131    }
132
133    fn fallback(
134        &self,
135        cx: &mut sim_kernel::Cx,
136        request: &SolveRequest<'_>,
137        reason: CpuFallbackReason,
138    ) -> Result<InterferenceStudy> {
139        {
140            let mut snapshot = self
141                .snapshot
142                .lock()
143                .expect("Tensor study snapshot poisoned");
144            snapshot.cpu_fallbacks = snapshot.cpu_fallbacks.saturating_add(1);
145            snapshot.last_route = ProviderRoute::Cpu(reason);
146        }
147        ReferenceStudySolver.solve(cx, request)
148    }
149
150    fn fail(&self, executor: Symbol, error: impl core::fmt::Display) -> Error {
151        let mut snapshot = self
152            .snapshot
153            .lock()
154            .expect("Tensor study snapshot poisoned");
155        snapshot.provider_failures = snapshot.provider_failures.saturating_add(1);
156        snapshot.last_route = ProviderRoute::ProviderFailed(executor.clone());
157        Error::Eval(format!(
158            "interference Tensor provider {executor} failed after selection; CPU restart is forbidden: {error}"
159        ))
160    }
161
162    fn card_is_eligible(&self, card: &TensorExecutorCard) -> bool {
163        required_operation_symbols()
164            .iter()
165            .all(|required| card.operations.iter().any(|actual| actual == required))
166    }
167}
168
169impl Default for TensorStudySolver {
170    fn default() -> Self {
171        Self::new(TensorStudyConfig::default())
172    }
173}
174
175impl StudySolver for TensorStudySolver {
176    fn solve(
177        &self,
178        cx: &mut sim_kernel::Cx,
179        request: &SolveRequest<'_>,
180    ) -> Result<InterferenceStudy> {
181        let admitted = RequestPreflight::admit(
182            request.problem(),
183            request.plane(),
184            request.sampling_policy(),
185            request.sampling_thresholds(),
186            request.work_budget(),
187        )
188        .map_err(|error| Error::Eval(format!("interference request admission failed: {error}")))?;
189        let Some(executor) = active_tensor_executor(cx) else {
190            return self.fallback(cx, request, CpuFallbackReason::NoExecutor);
191        };
192        let card = executor.card();
193        if card.locality == Symbol::qualified("compute", "auto") && card.provider == "auto/cpu" {
194            return self.fallback(cx, request, CpuFallbackReason::ProviderCpuChoice);
195        }
196        if !self.config.tile_profile.supports_f32 {
197            return self.fallback(cx, request, CpuFallbackReason::UnsupportedDtype);
198        }
199        if !self.card_is_eligible(&card) {
200            return self.fallback(cx, request, CpuFallbackReason::UnsupportedOperations);
201        }
202        if admitted.work_estimate.cells < self.config.min_accelerated_cells {
203            return self.fallback(cx, request, CpuFallbackReason::BelowCrossover);
204        }
205        let plan = LoweringPlan::preflight_with_executor(
206            cx,
207            request.problem(),
208            *request.plane(),
209            self.config.phase_budget,
210            self.config.tile_profile,
211            executor,
212        )
213        .map_err(|error| self.fail(card.symbol.clone(), error))?;
214        if plan.tile_plan().tiles().len() != 1 {
215            return self.fallback(cx, request, CpuFallbackReason::MultipleOutputTiles);
216        }
217        let execution =
218            execute_resident(cx, &plan).map_err(|error| self.fail(card.symbol.clone(), error))?;
219        let tolerances = self.config.tolerances;
220        let field = PhasorFieldDescriptor::new(
221            request.plane().rows(),
222            request.plane().columns(),
223            execution.real,
224            execution.imaginary,
225        )
226        .map_err(|error| self.fail(card.symbol.clone(), error))?;
227        let evidence = StudyEvidenceDescriptor {
228            sampling_policy: match request.sampling_policy() {
229                sim_lib_interference_core::SamplingPolicy::Strict => {
230                    Symbol::qualified("interference", "strict")
231                }
232                sim_lib_interference_core::SamplingPolicy::Annotate => {
233                    Symbol::qualified("interference", "annotate")
234                }
235            },
236            sampling: SamplingCertificateDescriptor::from_certificate(
237                admitted.sampling_certificate,
238            ),
239            work: WorkEstimateDescriptor::from_estimate(admitted.work_estimate),
240            provider: card.symbol.clone(),
241            dtype: domains::f32(),
242            component_absolute_tolerance: tolerances.component.absolute,
243            squared_magnitude_absolute_tolerance: tolerances.magnitude_squared.absolute,
244            completed_cells: admitted.work_estimate.cells,
245            completed_emitter_evaluations: admitted.work_estimate.emitter_evaluations,
246            uploads: execution.uploads,
247            submissions: execution.submissions,
248            intermediate_materializations: 0,
249            final_materializations: 2,
250            segments: execution.segments,
251            adapter: self.config.adapter.clone(),
252            profile: Some(card.provider),
253        };
254        let study = StudyDescriptor::new(
255            ProblemDescriptor::from_problem(request.problem()),
256            PlaneDescriptor::from_plane(*request.plane()),
257            field,
258            evidence,
259        )
260        .map_err(|error| self.fail(card.symbol.clone(), error))?;
261        let mut snapshot = self
262            .snapshot
263            .lock()
264            .expect("Tensor study snapshot poisoned");
265        snapshot.provider_completions = snapshot.provider_completions.saturating_add(1);
266        snapshot.last_route = ProviderRoute::ProviderCompleted(card.symbol);
267        Ok(study)
268    }
269}
270
271/// Loadable library that registers the Tensor study provider.
272#[derive(Clone, Default)]
273pub struct InterferenceComputeLib {
274    solver: TensorStudySolver,
275}
276
277impl InterferenceComputeLib {
278    /// Builds a provider library around a configured, observable solver.
279    pub fn new(solver: TensorStudySolver) -> Self {
280        Self { solver }
281    }
282}
283
284impl Lib for InterferenceComputeLib {
285    fn manifest(&self) -> LibManifest {
286        LibManifest {
287            id: interference_compute_lib_symbol(),
288            version: Version(env!("CARGO_PKG_VERSION").to_owned()),
289            abi: AbiVersion { major: 0, minor: 1 },
290            target: LibTarget::HostRegistered,
291            requires: vec![Dependency {
292                id: Symbol::qualified("sim", "interference"),
293                minimum_version: None,
294            }],
295            capabilities: Vec::new(),
296            exports: vec![Export::Value {
297                symbol: tensor_study_solver_symbol(),
298            }],
299        }
300    }
301
302    fn load(&self, _cx: &mut sim_kernel::LoadCx, linker: &mut Linker<'_>) -> Result<()> {
303        linker.value(
304            tensor_study_solver_symbol(),
305            SolverProvider::new(Arc::new(self.solver.clone())).into_value()?,
306        )
307    }
308}