1use std::collections::VecDeque;
4use std::sync::{Arc, Mutex};
5
6use sim_kernel::Symbol;
7use sim_lib_numbers_tensor::{
8 CpuTensorExecutor, SubmissionEvidence, Tensor, TensorExecError, TensorExecution,
9 TensorExecutor, TensorExecutorCard, TensorRequest,
10};
11
12use crate::storage::{ModeledResidentDescriptor, ModeledResidentStorage, ResidentHandle};
13
14const DEFAULT_SEGMENT_TILE_BYTES: u64 = 256 * 1024;
15const DEFAULT_STORAGE_BINDING_BYTES: u64 = 1024 * 1024;
16const DEFAULT_RESIDENT_BYTES: u64 = 8 * 1024 * 1024;
17const DEFAULT_QUEUE_BYTES: u64 = 2 * 1024 * 1024;
18const DEFAULT_DEADLINE_TICKS: u64 = 8;
19const MODELED_CELL_BYTES: u64 = 8;
20
21pub fn modeled_executor_symbol() -> Symbol {
23 Symbol::qualified("compute", "executor/model")
24}
25
26#[derive(Clone, Debug, PartialEq, Eq)]
28pub enum ModeledComputeFault {
29 OomBeforeSubmit,
31 DeviceLostDuringExecute,
33 ReadbackFailure,
35}
36
37#[derive(Clone, Debug, PartialEq, Eq)]
39pub struct ModeledComputeProfile {
40 pub provider: String,
42 pub max_queue_depth: usize,
44 pub max_queue_bytes: u64,
46 pub max_resident_bytes: u64,
48 pub segment_tile_bytes: u64,
50 pub max_storage_binding_bytes: u64,
52 pub submission_deadline_ticks: u64,
54 pub fault: Option<ModeledComputeFault>,
56 pub auto_flush_batches: bool,
58}
59
60impl Default for ModeledComputeProfile {
61 fn default() -> Self {
62 Self {
63 provider: "modeled-compute".to_owned(),
64 max_queue_depth: 8,
65 max_queue_bytes: DEFAULT_QUEUE_BYTES,
66 max_resident_bytes: DEFAULT_RESIDENT_BYTES,
67 segment_tile_bytes: DEFAULT_SEGMENT_TILE_BYTES,
68 max_storage_binding_bytes: DEFAULT_STORAGE_BINDING_BYTES,
69 submission_deadline_ticks: DEFAULT_DEADLINE_TICKS,
70 fault: None,
71 auto_flush_batches: false,
72 }
73 }
74}
75
76#[derive(Clone, Debug, PartialEq, Eq)]
78pub struct ModeledResidentSegment {
79 pub index: usize,
81 pub offset: u64,
83 pub bytes: u64,
85}
86
87#[derive(Clone, Debug, Default, PartialEq, Eq)]
89pub struct ModeledComputeSnapshot {
90 pub accepted: usize,
92 pub completed: usize,
94 pub queued: usize,
96 pub queued_bytes: u64,
98 pub resident_allocations: usize,
100 pub live_allocations: usize,
102 pub resident_bytes: u64,
104 pub evictions: usize,
106 pub segments: usize,
108 pub readbacks: usize,
110 pub materialization_failures: usize,
112 pub batch_flushes: usize,
114}
115
116#[derive(Default)]
117pub(crate) struct ModeledCounters {
118 accepted: usize,
119 completed: usize,
120 queued: usize,
121 queued_bytes: u64,
122 resident_allocations: usize,
123 live_allocations: VecDeque<ResidentAllocation>,
124 resident_bytes: u64,
125 evictions: usize,
126 segments: usize,
127 readbacks: usize,
128 materialization_failures: usize,
129 batch_flushes: usize,
130 internal_materializations: usize,
131 tick: u64,
132}
133
134#[derive(Clone)]
135struct ResidentAllocation {
136 handle: ResidentHandle,
137 bytes: u64,
138 active: bool,
139}
140
141#[derive(Clone)]
143pub struct ModeledTensorExecutor {
144 profile: ModeledComputeProfile,
145 counters: Arc<Mutex<ModeledCounters>>,
146}
147
148impl ModeledTensorExecutor {
149 pub fn new(profile: ModeledComputeProfile) -> Self {
151 Self {
152 profile,
153 counters: Arc::new(Mutex::new(ModeledCounters::default())),
154 }
155 }
156
157 pub fn default_profile() -> Self {
159 Self::new(ModeledComputeProfile::default())
160 }
161
162 pub fn snapshot(&self) -> ModeledComputeSnapshot {
164 let counters = self.counters.lock().expect("modeled counters poisoned");
165 ModeledComputeSnapshot {
166 accepted: counters.accepted,
167 completed: counters.completed,
168 queued: counters.queued,
169 queued_bytes: counters.queued_bytes,
170 resident_allocations: counters.resident_allocations,
171 live_allocations: counters
172 .live_allocations
173 .iter()
174 .filter(|allocation| allocation.active)
175 .count(),
176 resident_bytes: counters.resident_bytes,
177 evictions: counters.evictions,
178 segments: counters.segments,
179 readbacks: counters.readbacks,
180 materialization_failures: counters.materialization_failures,
181 batch_flushes: counters.batch_flushes,
182 }
183 }
184
185 pub(crate) fn increment_readbacks(&self) {
186 let mut counters = self.counters.lock().expect("modeled counters poisoned");
187 if counters.internal_materializations > 0 {
188 return;
189 }
190 counters.readbacks += 1;
191 }
192
193 pub(crate) fn increment_materialization_failures(&self) {
194 let mut counters = self.counters.lock().expect("modeled counters poisoned");
195 counters.materialization_failures += 1;
196 }
197
198 pub(crate) fn is_resident_active(&self, handle: &ResidentHandle) -> bool {
199 let counters = self.counters.lock().expect("modeled counters poisoned");
200 counters
201 .live_allocations
202 .iter()
203 .any(|allocation| allocation.active && allocation.handle == *handle)
204 }
205
206 pub(crate) fn begin_internal_materialization(&self) {
207 let mut counters = self.counters.lock().expect("modeled counters poisoned");
208 counters.internal_materializations += 1;
209 }
210
211 pub(crate) fn end_internal_materialization(&self) {
212 let mut counters = self.counters.lock().expect("modeled counters poisoned");
213 counters.internal_materializations = counters.internal_materializations.saturating_sub(1);
214 }
215
216 fn prepare_inputs(&self, request: TensorRequest) -> TensorRequest {
217 let inputs = request.inputs.iter().map(Self::prepare_tensor).collect();
218 TensorRequest::new(request.operation, inputs, request.output)
219 }
220
221 fn prepare_tensor(tensor: &Tensor) -> Tensor {
222 tensor
223 .storage()
224 .as_any()
225 .downcast_ref::<ModeledResidentStorage>()
226 .and_then(ModeledResidentStorage::resident_tensor)
227 .unwrap_or_else(|| tensor.clone())
228 }
229
230 fn request_bytes(request: &TensorRequest) -> std::result::Result<u64, TensorExecError> {
231 let output_bytes = tensor_bytes(request.output.shape())?;
232 request
233 .inputs
234 .iter()
235 .try_fold(output_bytes, |bytes, tensor| {
236 Ok(bytes.saturating_add(tensor_bytes(tensor.shape())?))
237 })
238 }
239
240 fn segment_layout(&self, bytes: u64) -> Vec<ModeledResidentSegment> {
241 let boundary = self
242 .profile
243 .segment_tile_bytes
244 .min(self.profile.max_storage_binding_bytes)
245 .max(MODELED_CELL_BYTES);
246 let mut segments = Vec::new();
247 let mut offset = 0;
248 while offset < bytes {
249 let segment_bytes = (bytes - offset).min(boundary);
250 segments.push(ModeledResidentSegment {
251 index: segments.len(),
252 offset,
253 bytes: segment_bytes,
254 });
255 offset += segment_bytes;
256 }
257 segments
258 }
259
260 fn reserve_submission(&self, bytes: u64) -> std::result::Result<(), TensorExecError> {
261 let mut counters = self.counters.lock().expect("modeled counters poisoned");
262 counters.tick = counters.tick.saturating_add(1);
263 if self.profile.auto_flush_batches
264 && (counters.queued >= self.profile.max_queue_depth
265 || counters.queued_bytes.saturating_add(bytes) > self.profile.max_queue_bytes)
266 && counters.queued > 0
267 {
268 counters.queued = 0;
269 counters.queued_bytes = 0;
270 counters.batch_flushes += 1;
271 }
272 if counters.queued >= self.profile.max_queue_depth {
273 return Err(TensorExecError::InvalidRequest {
274 message: Arc::from("modeled compute queue is full"),
275 });
276 }
277 if counters.queued_bytes.saturating_add(bytes) > self.profile.max_queue_bytes {
278 return Err(TensorExecError::InvalidRequest {
279 message: Arc::from("modeled compute queue byte budget is full"),
280 });
281 }
282 if self.profile.submission_deadline_ticks == 0 {
283 return Err(TensorExecError::InvalidRequest {
284 message: Arc::from("modeled compute submission deadline expired"),
285 });
286 }
287 counters.accepted += 1;
288 counters.queued += 1;
289 counters.queued_bytes += bytes;
290 Ok(())
291 }
292
293 fn release_submission(&self, bytes: u64) {
294 let mut counters = self.counters.lock().expect("modeled counters poisoned");
295 counters.queued = counters.queued.saturating_sub(1);
296 counters.queued_bytes = counters.queued_bytes.saturating_sub(bytes);
297 }
298
299 fn allocate_resident(
300 &self,
301 bytes: u64,
302 segments: usize,
303 ) -> std::result::Result<ResidentHandle, TensorExecError> {
304 if bytes > self.profile.max_resident_bytes {
305 return Err(TensorExecError::InvalidRequest {
306 message: Arc::from("modeled compute resident allocation exceeds pool"),
307 });
308 }
309 let mut counters = self.counters.lock().expect("modeled counters poisoned");
310 while counters.resident_bytes.saturating_add(bytes) > self.profile.max_resident_bytes {
311 let Some(mut allocation) = counters.live_allocations.pop_front() else {
312 break;
313 };
314 if allocation.active {
315 allocation.active = false;
316 counters.resident_bytes = counters.resident_bytes.saturating_sub(allocation.bytes);
317 counters.evictions += 1;
318 }
319 counters.live_allocations.push_back(allocation);
320 }
321 counters.completed += 1;
322 counters.resident_allocations += 1;
323 counters.resident_bytes += bytes;
324 counters.segments += segments;
325 let handle = ResidentHandle::new(counters.resident_allocations);
326 counters.live_allocations.push_back(ResidentAllocation {
327 handle: handle.clone(),
328 bytes,
329 active: true,
330 });
331 Ok(handle)
332 }
333}
334
335fn tensor_bytes(shape: &[usize]) -> std::result::Result<u64, TensorExecError> {
336 let cells = shape.iter().try_fold(1_u64, |count, extent| {
337 count
338 .checked_mul(
339 u64::try_from(*extent).map_err(|_| TensorExecError::InvalidRequest {
340 message: Arc::from("modeled compute tensor extent exceeds u64"),
341 })?,
342 )
343 .ok_or_else(|| TensorExecError::InvalidRequest {
344 message: Arc::from("modeled compute tensor byte count overflowed"),
345 })
346 })?;
347 cells
348 .checked_mul(MODELED_CELL_BYTES)
349 .ok_or_else(|| TensorExecError::InvalidRequest {
350 message: Arc::from("modeled compute tensor byte count overflowed"),
351 })
352}
353
354impl Default for ModeledTensorExecutor {
355 fn default() -> Self {
356 Self::default_profile()
357 }
358}
359
360impl TensorExecutor for ModeledTensorExecutor {
361 fn card(&self) -> TensorExecutorCard {
362 let cpu = CpuTensorExecutor::new().card();
363 TensorExecutorCard::new(
364 modeled_executor_symbol(),
365 self.profile.provider.clone(),
366 Symbol::qualified("compute", "modeled-resident"),
367 cpu.operations.to_vec(),
368 None,
369 )
370 }
371
372 fn execute(
373 &self,
374 cx: &mut sim_kernel::Cx,
375 request: TensorRequest,
376 ) -> std::result::Result<TensorExecution, TensorExecError> {
377 let request_bytes = Self::request_bytes(&request)?;
378 if self.profile.fault == Some(ModeledComputeFault::OomBeforeSubmit) {
379 return Err(TensorExecError::InvalidRequest {
380 message: Arc::from("modeled compute out of memory before submission"),
381 });
382 }
383 self.reserve_submission(request_bytes)?;
384 if self.profile.fault == Some(ModeledComputeFault::DeviceLostDuringExecute) {
385 self.release_submission(request_bytes);
386 return Err(TensorExecError::Eval {
387 message: Arc::from("modeled compute device lost during execution"),
388 });
389 }
390
391 let request = self.prepare_inputs(request);
392 self.begin_internal_materialization();
393 let result = CpuTensorExecutor::new().execute(cx, request);
394 self.end_internal_materialization();
395 let result = match result {
396 Ok(result) => result,
397 Err(error) => {
398 self.release_submission(request_bytes);
399 return Err(error);
400 }
401 };
402 let TensorExecution::Complete(tensor) = result else {
403 return Ok(result);
404 };
405 let host_tensor = Self::prepare_tensor(&tensor);
406 self.begin_internal_materialization();
407 let host_cells = host_tensor.cells().map_err(TensorExecError::from);
408 self.end_internal_materialization();
409 let host_cells = host_cells?;
410 let resident_bytes = tensor_bytes(tensor.shape())?;
411 let segments = self.segment_layout(resident_bytes);
412 let allocation = self.allocate_resident(resident_bytes, segments.len())?;
413 let storage = ModeledResidentStorage::new(
414 ModeledResidentDescriptor {
415 site: Symbol::new("site/compute/model"),
416 allocation,
417 segments,
418 shape: tensor.shape().to_vec(),
419 dtype: tensor.dtype().clone(),
420 },
421 host_cells,
422 self.clone(),
423 self.profile.fault.clone(),
424 );
425 Ok(TensorExecution::Complete(Tensor::from_storage(
426 tensor.shape().to_vec(),
427 tensor.dtype().clone(),
428 Arc::new(storage),
429 )?))
430 }
431
432 fn flush(&self) -> std::result::Result<SubmissionEvidence, TensorExecError> {
433 let mut counters = self.counters.lock().expect("modeled counters poisoned");
434 let accepted = counters.queued;
435 counters.queued = 0;
436 counters.queued_bytes = 0;
437 Ok(SubmissionEvidence::new(modeled_executor_symbol(), accepted))
438 }
439}