1use crate::core::registry::SlideReader;
2use crate::core::types::{
3 CpuTile, Dataset, Level, OutputBackendRequest, TileCodecKind, TileLayout, TileOutputPreference,
4 TilePixels, TileRequest,
5};
6use crate::error::WsiError;
7use rayon::ThreadPool;
8use std::cell::RefCell;
9use std::collections::{HashMap, VecDeque};
10use std::num::NonZeroUsize;
11use std::sync::{Arc, Mutex, OnceLock};
12use std::time::{Duration, Instant};
13
14const DEFAULT_ROUTE_SAMPLE_SIZE: usize = 32;
15const DIRECT_DEVICE_BATCH_THRESHOLD: usize = 8;
16const DEVICE_WIN_RATIO: f64 = 0.85;
17const ROUTE_CACHE_MAX_ENTRIES: usize = 1024;
18
19thread_local! {
20 static CURRENT_DECODE_RUNTIME: RefCell<Option<Arc<DecodeRuntime>>> = const { RefCell::new(None) };
21}
22
23#[derive(Debug, Clone, Copy, PartialEq, Eq)]
24#[non_exhaustive]
25pub struct DecodeExecutionOptions {
26 jp2k_cpu_threads: Option<NonZeroUsize>,
27 route_sample_size: usize,
28}
29
30impl DecodeExecutionOptions {
31 pub fn with_jp2k_cpu_threads(mut self, threads: NonZeroUsize) -> Self {
32 self.jp2k_cpu_threads = Some(threads);
33 self
34 }
35
36 pub fn with_route_sample_size(mut self, sample_size: usize) -> Self {
37 self.route_sample_size = sample_size.max(1);
38 self
39 }
40
41 pub fn jp2k_cpu_threads(&self) -> Option<NonZeroUsize> {
42 self.jp2k_cpu_threads
43 }
44
45 pub fn route_sample_size(&self) -> usize {
46 self.route_sample_size
47 }
48}
49
50impl Default for DecodeExecutionOptions {
51 fn default() -> Self {
52 Self {
53 jp2k_cpu_threads: None,
54 route_sample_size: DEFAULT_ROUTE_SAMPLE_SIZE,
55 }
56 }
57}
58
59#[derive(Debug, Clone, Copy, PartialEq, Eq)]
60#[non_exhaustive]
61pub enum DecodeRoute {
62 Cpu,
63 Device,
64}
65
66#[derive(Debug, Clone, PartialEq)]
67#[non_exhaustive]
68pub struct DecodeRouteDecision {
69 pub winner: DecodeRoute,
70 pub sample_tile_count: usize,
71 pub cpu_elapsed: Duration,
72 pub device_elapsed: Duration,
73 pub device_tile_count: usize,
74}
75
76impl DecodeRouteDecision {
77 pub fn measured(
78 sample_tile_count: usize,
79 cpu_elapsed: Duration,
80 device_elapsed: Duration,
81 device_tile_count: usize,
82 ) -> Self {
83 Self {
84 winner: Self::winner_for_measurement(cpu_elapsed, device_elapsed, device_tile_count),
85 sample_tile_count,
86 cpu_elapsed,
87 device_elapsed,
88 device_tile_count,
89 }
90 }
91
92 pub fn winner_for_measurement(
93 cpu_elapsed: Duration,
94 device_elapsed: Duration,
95 device_tile_count: usize,
96 ) -> DecodeRoute {
97 let cpu_ms = cpu_elapsed.as_secs_f64() * 1000.0;
98 let device_ms = device_elapsed.as_secs_f64() * 1000.0;
99 if device_tile_count > 0 && cpu_ms > 0.0 && device_ms <= cpu_ms * DEVICE_WIN_RATIO {
100 DecodeRoute::Device
101 } else {
102 DecodeRoute::Cpu
103 }
104 }
105}
106
107struct MeasuredDecodeRoute {
108 decision: DecodeRouteDecision,
109 sample_tiles: Vec<TilePixels>,
110}
111
112#[derive(Debug)]
113pub(crate) struct DecodeRuntime {
114 options: DecodeExecutionOptions,
115 jp2k_cpu_pool: Option<ThreadPool>,
116 route_cache: Mutex<DecodeRouteCache>,
117}
118
119impl DecodeRuntime {
120 pub(crate) fn new(options: DecodeExecutionOptions) -> Result<Self, WsiError> {
121 Self::build(options, true)
122 }
123
124 pub(crate) fn arc_for_options(options: DecodeExecutionOptions) -> Result<Arc<Self>, WsiError> {
125 if options == DecodeExecutionOptions::default() {
126 Ok(Self::default_arc())
127 } else {
128 Ok(Arc::new(Self::new(options)?))
129 }
130 }
131
132 fn build(options: DecodeExecutionOptions, fail_on_pool_error: bool) -> Result<Self, WsiError> {
133 let threads = options
134 .jp2k_cpu_threads
135 .map_or_else(default_jp2k_cpu_threads, NonZeroUsize::get);
136 let jp2k_cpu_pool = match rayon::ThreadPoolBuilder::new()
137 .num_threads(threads)
138 .thread_name(|index| format!("wsi_rs-jp2k-cpu-{index}"))
139 .build()
140 {
141 Ok(pool) => Some(pool),
142 Err(err) if fail_on_pool_error => {
143 return Err(WsiError::Unsupported {
144 reason: format!("failed to initialize JP2K CPU decode pool: {err}"),
145 });
146 }
147 Err(err) => {
148 tracing::error!(
149 error = %err,
150 "failed to initialize default JP2K CPU decode pool; falling back to inline decode"
151 );
152 None
153 }
154 };
155 Ok(Self {
156 options,
157 jp2k_cpu_pool,
158 route_cache: Mutex::new(DecodeRouteCache::new()),
159 })
160 }
161
162 pub(crate) fn default_arc() -> Arc<Self> {
163 static DEFAULT_RUNTIME: OnceLock<Arc<DecodeRuntime>> = OnceLock::new();
164 DEFAULT_RUNTIME
165 .get_or_init(|| {
166 Arc::new(match Self::build(DecodeExecutionOptions::default(), false) {
167 Ok(runtime) => runtime,
168 Err(err) => {
169 tracing::error!(
170 error = %err,
171 "failed to initialize default decode runtime; falling back to inline decode"
172 );
173 Self::inline(DecodeExecutionOptions::default())
174 }
175 })
176 })
177 .clone()
178 }
179
180 fn inline(options: DecodeExecutionOptions) -> Self {
181 Self {
182 options,
183 jp2k_cpu_pool: None,
184 route_cache: Mutex::new(DecodeRouteCache::new()),
185 }
186 }
187
188 pub(crate) fn install_jp2k_cpu<R: Send>(&self, op: impl FnOnce() -> R + Send) -> R {
189 if let Some(pool) = &self.jp2k_cpu_pool {
190 pool.install(op)
191 } else {
192 op()
193 }
194 }
195
196 pub(crate) fn has_jp2k_cpu_pool(&self) -> bool {
197 self.jp2k_cpu_pool.is_some()
198 }
199
200 pub(crate) fn options(&self) -> DecodeExecutionOptions {
201 self.options
202 }
203
204 pub(crate) fn with_current<T>(self: &Arc<Self>, f: impl FnOnce() -> T) -> T {
205 struct Restore(Option<Arc<DecodeRuntime>>);
206 impl Drop for Restore {
207 fn drop(&mut self) {
208 let previous = self.0.take();
209 CURRENT_DECODE_RUNTIME.with(|slot| {
210 *slot.borrow_mut() = previous;
211 });
212 }
213 }
214
215 let previous = CURRENT_DECODE_RUNTIME.with(|slot| slot.replace(Some(self.clone())));
216 let _restore = Restore(previous);
217 f()
218 }
219
220 fn cached_route(&self, key: &DecodeRouteKey) -> Option<DecodeRouteDecision> {
221 self.route_cache
222 .lock()
223 .unwrap_or_else(|err| err.into_inner())
224 .get(key)
225 }
226
227 fn store_route(&self, key: DecodeRouteKey, decision: DecodeRouteDecision) {
228 self.route_cache
229 .lock()
230 .unwrap_or_else(|err| err.into_inner())
231 .insert(key, decision);
232 }
233}
234
235#[derive(Debug)]
236struct DecodeRouteCache {
237 entries: HashMap<DecodeRouteKey, DecodeRouteDecision>,
238 insertion_order: VecDeque<DecodeRouteKey>,
239}
240
241impl DecodeRouteCache {
242 fn new() -> Self {
243 Self {
244 entries: HashMap::new(),
245 insertion_order: VecDeque::new(),
246 }
247 }
248
249 fn get(&self, key: &DecodeRouteKey) -> Option<DecodeRouteDecision> {
250 self.entries.get(key).cloned()
251 }
252
253 fn insert(&mut self, key: DecodeRouteKey, decision: DecodeRouteDecision) {
254 if !self.entries.contains_key(&key) {
255 while self.entries.len() >= ROUTE_CACHE_MAX_ENTRIES {
256 let Some(evicted) = self.insertion_order.pop_front() else {
257 break;
258 };
259 self.entries.remove(&evicted);
260 }
261 self.insertion_order.push_back(key.clone());
262 }
263 self.entries.insert(key, decision);
264 }
265
266 #[cfg(test)]
267 fn len(&self) -> usize {
268 self.entries.len()
269 }
270}
271
272pub(crate) fn current_decode_runtime() -> Option<Arc<DecodeRuntime>> {
273 CURRENT_DECODE_RUNTIME.with(|slot| slot.borrow().clone())
274}
275
276fn default_jp2k_cpu_threads() -> usize {
277 std::thread::available_parallelism()
278 .map_or(1, NonZeroUsize::get)
279 .max(1)
280}
281
282#[derive(Debug, Clone, PartialEq, Eq, Hash)]
283struct DecodeRouteKey {
284 dataset_id: u128,
285 scene: usize,
286 series: usize,
287 level: u32,
288 tile_grid: RouteTileGrid,
289 codec_kind: TileCodecKind,
290 output_backend: OutputBackendRequest,
291 device_backend_identity: String,
292 sample_tile_count: usize,
293}
294
295#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash)]
296struct RouteTileGrid {
297 tile_width: u32,
298 tile_height: u32,
299 tiles_across: u64,
300 tiles_down: u64,
301}
302
303pub(crate) struct AdaptiveDecodeReader {
304 inner: Box<dyn SlideReader>,
305 runtime: Arc<DecodeRuntime>,
306}
307
308impl AdaptiveDecodeReader {
309 pub(crate) fn new(inner: Box<dyn SlideReader>, runtime: Arc<DecodeRuntime>) -> Self {
310 Self { inner, runtime }
311 }
312
313 fn read_tiles_adaptive(
314 &self,
315 reqs: &[TileRequest],
316 output: TileOutputPreference,
317 ) -> Result<Vec<TilePixels>, WsiError> {
318 if !should_adapt_output(&output) {
319 tracing::debug!(
320 requested_tiles = reqs.len(),
321 adaptive_decode = false,
322 "wsi tile batch routed without adaptive decode"
323 );
324 return self
325 .runtime
326 .with_current(|| self.inner.read_tiles(reqs, output));
327 }
328 let route_sample_size = self.runtime.options.route_sample_size();
329 let Some(key) = route_key_for_batch(self.inner.as_ref(), reqs, &output, route_sample_size)
330 else {
331 tracing::debug!(
332 requested_tiles = reqs.len(),
333 route_sample_size,
334 adaptive_decode = true,
335 route_key_available = false,
336 "wsi adaptive decode fell back to requested output"
337 );
338 return self
339 .runtime
340 .with_current(|| self.inner.read_tiles(reqs, output));
341 };
342 if reqs.len() >= DIRECT_DEVICE_BATCH_THRESHOLD {
343 tracing::debug!(
344 requested_tiles = reqs.len(),
345 route_sample_size,
346 direct_device_batch_threshold = DIRECT_DEVICE_BATCH_THRESHOLD,
347 adaptive_decode = true,
348 route_key_available = true,
349 "wsi adaptive decode sent large batch through requested output"
350 );
351 return self
352 .runtime
353 .with_current(|| self.inner.read_tiles(reqs, output));
354 }
355 let (route, measured) = match self.runtime.cached_route(&key) {
356 Some(decision) => {
357 tracing::debug!(
358 requested_tiles = reqs.len(),
359 route_sample_size,
360 route_cache_hit = true,
361 route = ?decision.winner,
362 sample_tile_count = decision.sample_tile_count,
363 cpu_elapsed_ms = decision.cpu_elapsed.as_secs_f64() * 1000.0,
364 device_elapsed_ms = decision.device_elapsed.as_secs_f64() * 1000.0,
365 device_tile_count = decision.device_tile_count,
366 "wsi adaptive decode reused cached route"
367 );
368 (decision.winner, None)
369 }
370 None => {
371 let measured = self.measure_route(reqs, output.clone())?;
372 let winner = measured.decision.winner;
373 tracing::debug!(
374 requested_tiles = reqs.len(),
375 route_sample_size,
376 route_cache_hit = false,
377 route = ?winner,
378 sample_tile_count = measured.decision.sample_tile_count,
379 cpu_elapsed_ms = measured.decision.cpu_elapsed.as_secs_f64() * 1000.0,
380 device_elapsed_ms = measured.decision.device_elapsed.as_secs_f64() * 1000.0,
381 device_tile_count = measured.decision.device_tile_count,
382 "wsi adaptive decode measured route"
383 );
384 self.runtime.store_route(key, measured.decision.clone());
385 (winner, Some(measured.sample_tiles))
386 }
387 };
388 let routed_output = match route {
389 DecodeRoute::Cpu => TileOutputPreference::cpu(),
390 DecodeRoute::Device => output,
391 };
392 if let Some(mut measured) = measured {
393 let sample_len = reqs.len().min(self.runtime.options.route_sample_size());
394 if measured.len() == sample_len {
395 if sample_len == reqs.len() {
396 return Ok(measured);
397 }
398 let mut rest = self
399 .runtime
400 .with_current(|| self.inner.read_tiles(&reqs[sample_len..], routed_output))?;
401 measured.append(&mut rest);
402 return Ok(measured);
403 }
404 }
405 self.runtime
406 .with_current(|| self.inner.read_tiles(reqs, routed_output))
407 }
408
409 fn measure_route(
410 &self,
411 reqs: &[TileRequest],
412 device_output: TileOutputPreference,
413 ) -> Result<MeasuredDecodeRoute, WsiError> {
414 let sample_len = reqs.len().min(self.runtime.options.route_sample_size());
415 let sample = &reqs[..sample_len];
416
417 let device_started = Instant::now();
418 let device_result = self
419 .runtime
420 .with_current(|| self.inner.read_tiles(sample, device_output));
421 let device_elapsed = device_started.elapsed();
422 let device_tile_count = device_result
423 .as_ref()
424 .map(|tiles| {
425 tiles
426 .iter()
427 .filter(|tile| matches!(tile, TilePixels::Device(_)))
428 .count()
429 })
430 .unwrap_or(0);
431 let device_result = match device_result {
432 Ok(device_tiles) if device_tile_count == 0 => {
433 return Ok(MeasuredDecodeRoute {
434 decision: DecodeRouteDecision::measured(
435 device_tiles.len(),
436 device_elapsed,
437 device_elapsed,
438 device_tile_count,
439 ),
440 sample_tiles: device_tiles,
441 });
442 }
443 other => other,
444 };
445
446 let cpu_started = Instant::now();
447 let cpu_tiles = self
448 .runtime
449 .with_current(|| self.inner.read_tiles(sample, TileOutputPreference::cpu()))?;
450 let cpu_elapsed = cpu_started.elapsed();
451
452 let decision = DecodeRouteDecision::measured(
453 cpu_tiles.len(),
454 cpu_elapsed,
455 device_elapsed,
456 device_tile_count,
457 );
458 let sample_tiles = match decision.winner {
459 DecodeRoute::Cpu => cpu_tiles,
460 DecodeRoute::Device => device_result?,
461 };
462
463 Ok(MeasuredDecodeRoute {
464 decision,
465 sample_tiles,
466 })
467 }
468}
469
470impl SlideReader for AdaptiveDecodeReader {
471 fn dataset(&self) -> &Dataset {
472 self.inner.dataset()
473 }
474
475 fn tile_codec_kind(&self, req: &TileRequest) -> TileCodecKind {
476 self.inner.tile_codec_kind(req)
477 }
478
479 fn level_source_kind(
480 &self,
481 scene: crate::core::types::SceneId,
482 series: crate::core::types::SeriesId,
483 level: crate::core::types::LevelIdx,
484 ) -> Result<crate::core::types::LevelSourceKind, WsiError> {
485 self.inner.level_source_kind(scene, series, level)
486 }
487
488 fn read_tiles(
489 &self,
490 reqs: &[TileRequest],
491 output: TileOutputPreference,
492 ) -> Result<Vec<TilePixels>, WsiError> {
493 self.read_tiles_adaptive(reqs, output)
494 }
495
496 fn read_tiles_controlled(
497 &self,
498 reqs: &[TileRequest],
499 output: TileOutputPreference,
500 control: &crate::ReadControl,
501 ) -> Result<Vec<TilePixels>, WsiError> {
502 control.check_cancelled()?;
503 let tiles = self
504 .runtime
505 .with_current(|| self.inner.read_tiles_controlled(reqs, output, control))?;
506 control.check_cancelled()?;
507 Ok(tiles)
508 }
509
510 fn read_tile_cpu(&self, req: &TileRequest) -> Result<CpuTile, WsiError> {
511 self.runtime.with_current(|| self.inner.read_tile_cpu(req))
512 }
513
514 fn read_raw_compressed_tile(
515 &self,
516 req: &TileRequest,
517 ) -> Result<crate::core::types::RawCompressedTile, WsiError> {
518 self.inner.read_raw_compressed_tile(req)
519 }
520
521 fn read_raw_compressed_display_tile(
522 &self,
523 req: &crate::core::types::TileViewRequest,
524 ) -> Result<crate::core::types::RawCompressedTile, WsiError> {
525 self.inner.read_raw_compressed_display_tile(req)
526 }
527
528 fn read_tiles_cpu(&self, reqs: &[TileRequest]) -> Result<Vec<CpuTile>, WsiError> {
529 self.runtime
530 .with_current(|| self.inner.read_tiles_cpu(reqs))
531 }
532
533 fn use_display_tile_cache(&self, req: &crate::core::types::TileViewRequest) -> bool {
534 self.inner.use_display_tile_cache(req)
535 }
536
537 fn read_region_fastpath(
538 &self,
539 ctx: &mut crate::core::registry::SlideReadContext<'_>,
540 req: &crate::core::types::RegionRequest,
541 ) -> Option<Result<CpuTile, WsiError>> {
542 self.runtime
543 .with_current(|| self.inner.read_region_fastpath(ctx, req))
544 }
545
546 fn read_region(
547 &self,
548 req: &crate::core::types::RegionRequest,
549 output: TileOutputPreference,
550 ) -> Result<TilePixels, WsiError> {
551 self.runtime
552 .with_current(|| self.inner.read_region(req, output))
553 }
554
555 fn read_display_tile(
556 &self,
557 req: &crate::core::types::TileViewRequest,
558 ) -> Result<CpuTile, WsiError> {
559 self.runtime
560 .with_current(|| self.inner.read_display_tile(req))
561 }
562
563 fn read_associated(&self, name: &str) -> Result<CpuTile, WsiError> {
564 self.inner.read_associated(name)
565 }
566
567 fn recommended_shared_cache_bytes(&self) -> Option<u64> {
568 self.inner.recommended_shared_cache_bytes()
569 }
570}
571
572fn should_adapt_output(output: &TileOutputPreference) -> bool {
573 matches!(output, TileOutputPreference::PreferDevice { .. })
574 && output.compressed_device_decode_enabled()
575 && output.adaptive_decode_route_enabled()
576}
577
578fn route_key_for_batch(
579 reader: &dyn SlideReader,
580 reqs: &[TileRequest],
581 output: &TileOutputPreference,
582 route_sample_size: usize,
583) -> Option<DecodeRouteKey> {
584 let first = reqs.first()?;
585 if !reqs.iter().all(|req| {
586 req.scene == first.scene && req.series == first.series && req.level == first.level
587 }) {
588 return None;
589 }
590 let codec_kind = reader.tile_codec_kind(first);
591 if !matches!(codec_kind, TileCodecKind::Jp2k | TileCodecKind::Htj2k) {
592 return None;
593 }
594 if !reqs
595 .iter()
596 .all(|req| reader.tile_codec_kind(req) == codec_kind)
597 {
598 return None;
599 }
600 let level = dataset_level(
601 reader.dataset(),
602 first.scene.get(),
603 first.series.get(),
604 first.level.get(),
605 )?;
606 let tile_grid = route_tile_grid(level)?;
607 Some(DecodeRouteKey {
608 dataset_id: reader.dataset().id.0,
609 scene: first.scene.get(),
610 series: first.series.get(),
611 level: first.level.get(),
612 tile_grid,
613 codec_kind,
614 output_backend: output.backend(),
615 device_backend_identity: device_backend_identity(output),
616 sample_tile_count: reqs.len().min(route_sample_size.max(1)),
617 })
618}
619
620fn dataset_level(dataset: &Dataset, scene: usize, series: usize, level: u32) -> Option<&Level> {
621 dataset
622 .scenes
623 .get(scene)?
624 .series
625 .get(series)?
626 .levels
627 .get(level as usize)
628}
629
630fn route_tile_grid(level: &Level) -> Option<RouteTileGrid> {
631 match &level.tile_layout {
632 TileLayout::Regular {
633 tile_width,
634 tile_height,
635 tiles_across,
636 tiles_down,
637 } => Some(RouteTileGrid {
638 tile_width: *tile_width,
639 tile_height: *tile_height,
640 tiles_across: *tiles_across,
641 tiles_down: *tiles_down,
642 }),
643 _ => None,
644 }
645}
646
647fn device_backend_identity(output: &TileOutputPreference) -> String {
648 #[cfg(feature = "metal")]
649 if let Some(metal) = output.metal_sessions() {
650 return format!("{:?}:{}", output.backend(), metal.device_identity());
651 }
652 #[cfg(feature = "cuda")]
653 if let Some(cuda) = output.cuda_sessions() {
654 return format!("{:?}:{}", output.backend(), cuda.device_identity());
655 }
656 format!("{:?}", output.backend())
657}
658
659#[cfg(test)]
660#[path = "decode_runtime/tests.rs"]
661mod tests;