Skip to main content

wsi_rs/core/
decode_runtime.rs

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;