1use 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
23pub fn interference_compute_lib_symbol() -> Symbol {
25 Symbol::qualified("sim", "interference-compute")
26}
27
28#[derive(Clone, Debug, PartialEq, Eq)]
30pub enum CpuFallbackReason {
31 NoExecutor,
33 ProviderCpuChoice,
35 UnsupportedDtype,
37 UnsupportedOperations,
39 BelowCrossover,
41 MultipleOutputTiles,
43}
44
45#[derive(Clone, Debug, PartialEq, Eq)]
47pub enum ProviderRoute {
48 Idle,
50 Cpu(CpuFallbackReason),
52 ProviderCompleted(Symbol),
54 ProviderFailed(Symbol),
56}
57
58#[derive(Clone, Debug, PartialEq, Eq)]
60pub struct TensorStudySnapshot {
61 pub cpu_fallbacks: u64,
63 pub provider_completions: u64,
65 pub provider_failures: u64,
67 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#[derive(Clone, Debug, PartialEq)]
84pub struct TensorStudyConfig {
85 pub min_accelerated_cells: u64,
87 pub phase_budget: PhaseBudget,
89 pub tile_profile: TileProfile,
91 pub tolerances: DifferentialTolerances,
93 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#[derive(Clone)]
111pub struct TensorStudySolver {
112 config: TensorStudyConfig,
113 snapshot: Arc<Mutex<TensorStudySnapshot>>,
114}
115
116impl TensorStudySolver {
117 pub fn new(config: TensorStudyConfig) -> Self {
119 Self {
120 config,
121 snapshot: Arc::new(Mutex::new(TensorStudySnapshot::default())),
122 }
123 }
124
125 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#[derive(Clone, Default)]
273pub struct InterferenceComputeLib {
274 solver: TensorStudySolver,
275}
276
277impl InterferenceComputeLib {
278 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}