Skip to main content

optirs_gpu/memory/management/
prefetching.rs

1// Memory prefetching for GPU memory management
2//
3// This module provides advanced prefetching strategies to improve GPU memory
4// performance by anticipating future memory access patterns and proactively
5// loading data before it's needed.
6
7use std::collections::{BTreeMap, HashMap, VecDeque};
8use std::sync::{Arc, Mutex};
9use std::time::{Duration, Instant};
10
11/// Main prefetching engine
12pub struct PrefetchingEngine {
13    /// Configuration
14    config: PrefetchConfig,
15    /// Statistics
16    stats: PrefetchStats,
17    /// Active prefetching strategies
18    strategies: Vec<Box<dyn PrefetchStrategy>>,
19    /// Access pattern history
20    access_history: AccessHistoryTracker,
21    /// Prefetch requests queue
22    prefetch_queue: VecDeque<PrefetchRequest>,
23    /// Cache of prefetched data
24    prefetch_cache: PrefetchCache,
25    /// Performance monitoring
26    performance_monitor: PerformanceMonitor,
27}
28
29/// Prefetching configuration
30#[derive(Debug, Clone)]
31pub struct PrefetchConfig {
32    /// Enable automatic prefetching
33    pub auto_prefetch: bool,
34    /// Maximum prefetch distance (bytes)
35    pub max_prefetch_distance: usize,
36    /// Prefetch window size
37    pub prefetch_window: usize,
38    /// Minimum access frequency for prefetching
39    pub min_access_frequency: f64,
40    /// Enable adaptive prefetching
41    pub enable_adaptive: bool,
42    /// Enable pattern-based prefetching
43    pub enable_pattern_based: bool,
44    /// Enable stride-based prefetching
45    pub enable_stride_based: bool,
46    /// Enable ML-based prefetching
47    pub enable_ml_based: bool,
48    /// Prefetch aggressiveness (0.0 to 1.0)
49    pub aggressiveness: f64,
50    /// Cache size for prefetched data
51    pub cache_size: usize,
52    /// Enable performance monitoring
53    pub enable_monitoring: bool,
54    /// History window size
55    pub history_window: usize,
56}
57
58impl Default for PrefetchConfig {
59    fn default() -> Self {
60        Self {
61            auto_prefetch: true,
62            max_prefetch_distance: 1024 * 1024, // 1MB
63            prefetch_window: 64,
64            min_access_frequency: 0.1,
65            enable_adaptive: true,
66            enable_pattern_based: true,
67            enable_stride_based: true,
68            enable_ml_based: false,
69            aggressiveness: 0.5,
70            cache_size: 16 * 1024 * 1024, // 16MB
71            enable_monitoring: true,
72            history_window: 1000,
73        }
74    }
75}
76
77/// Prefetching statistics
78#[derive(Debug, Clone, Default)]
79pub struct PrefetchStats {
80    /// Total prefetch requests
81    pub total_requests: u64,
82    /// Successful prefetches (used)
83    pub successful_prefetches: u64,
84    /// Failed prefetches (unused)
85    pub failed_prefetches: u64,
86    /// Prefetch accuracy ratio
87    pub accuracy_ratio: f64,
88    /// Total bytes prefetched
89    pub total_bytes_prefetched: u64,
90    /// Useful bytes prefetched
91    pub useful_bytes_prefetched: u64,
92    /// Cache hit rate
93    pub cache_hit_rate: f64,
94    /// Average prefetch latency
95    pub average_latency: Duration,
96    /// Bandwidth saved by prefetching
97    pub bandwidth_saved: u64,
98    /// Strategy performance
99    pub strategy_stats: HashMap<String, StrategyStats>,
100}
101
102/// Individual strategy statistics
103#[derive(Debug, Clone, Default)]
104pub struct StrategyStats {
105    pub requests: u64,
106    pub hits: u64,
107    pub misses: u64,
108    pub accuracy: f64,
109    pub latency: Duration,
110}
111
112/// Memory access tracking
113pub struct AccessHistoryTracker {
114    /// Recent access history
115    access_history: VecDeque<MemoryAccess>,
116    /// Access patterns
117    patterns: HashMap<AccessPattern, PatternFrequency>,
118    /// Stride patterns
119    stride_patterns: HashMap<usize, StrideInfo>,
120    /// Sequential access tracking
121    sequential_tracking: HashMap<usize, SequentialInfo>,
122    /// Access frequency map
123    frequency_map: HashMap<usize, AccessFrequency>,
124}
125
126/// Memory access record
127#[derive(Debug, Clone)]
128pub struct MemoryAccess {
129    /// Memory address accessed
130    pub address: usize,
131    /// Access size
132    pub size: usize,
133    /// Access timestamp
134    pub timestamp: Instant,
135    /// Access type (read/write)
136    pub access_type: AccessType,
137    /// Thread/context ID
138    pub context_id: u32,
139    /// GPU kernel ID
140    pub kernel_id: Option<u32>,
141}
142
143/// Access type enumeration
144#[derive(Debug, Clone, PartialEq)]
145pub enum AccessType {
146    Read,
147    Write,
148    ReadWrite,
149}
150
151/// Access pattern representation
152#[derive(Debug, Clone, Hash, PartialEq, Eq)]
153pub struct AccessPattern {
154    /// Pattern type
155    pub pattern_type: PatternType,
156    /// Address deltas
157    pub deltas: Vec<isize>,
158    /// Pattern size
159    pub size: usize,
160}
161
162/// Pattern type enumeration
163#[derive(Debug, Clone, Hash, PartialEq, Eq)]
164pub enum PatternType {
165    Sequential,
166    Strided,
167    Random,
168    Irregular,
169    Custom(String),
170}
171
172/// Pattern frequency tracking
173#[derive(Debug, Clone)]
174pub struct PatternFrequency {
175    pub count: u32,
176    pub last_seen: Instant,
177    pub confidence: f64,
178    pub prediction_accuracy: f64,
179}
180
181/// Stride pattern information
182#[derive(Debug, Clone)]
183pub struct StrideInfo {
184    pub stride: isize,
185    pub frequency: u32,
186    pub last_address: usize,
187    pub confidence: f64,
188    pub start_time: Instant,
189}
190
191/// Sequential access information
192#[derive(Debug, Clone)]
193pub struct SequentialInfo {
194    pub start_address: usize,
195    pub current_address: usize,
196    pub length: usize,
197    pub direction: i8, // 1 for forward, -1 for backward
198    pub last_access: Instant,
199}
200
201/// Access frequency tracking
202#[derive(Debug, Clone)]
203pub struct AccessFrequency {
204    pub count: u32,
205    pub first_access: Instant,
206    pub last_access: Instant,
207    pub average_interval: Duration,
208}
209
210impl AccessHistoryTracker {
211    pub fn new(capacity: usize) -> Self {
212        Self {
213            access_history: VecDeque::with_capacity(capacity),
214            patterns: HashMap::new(),
215            stride_patterns: HashMap::new(),
216            sequential_tracking: HashMap::new(),
217            frequency_map: HashMap::new(),
218        }
219    }
220
221    /// Record a memory access
222    pub fn record_access(&mut self, access: MemoryAccess) {
223        // Add to history
224        self.access_history.push_back(access.clone());
225        if self.access_history.len() > self.access_history.capacity() {
226            self.access_history.pop_front();
227        }
228
229        // Update frequency map
230        let freq = self
231            .frequency_map
232            .entry(access.address)
233            .or_insert_with(|| AccessFrequency {
234                count: 0,
235                first_access: access.timestamp,
236                last_access: access.timestamp,
237                average_interval: Duration::from_secs(0),
238            });
239
240        let interval = if freq.count > 0 {
241            access.timestamp.duration_since(freq.last_access)
242        } else {
243            Duration::from_secs(0)
244        };
245
246        freq.count += 1;
247        freq.last_access = access.timestamp;
248        freq.average_interval = if freq.count > 1 {
249            Duration::from_nanos(
250                (freq.average_interval.as_nanos() as u64 * (freq.count - 1) as u64
251                    + interval.as_nanos() as u64)
252                    / freq.count as u64,
253            )
254        } else {
255            interval
256        };
257
258        // Detect patterns
259        self.detect_patterns(&access);
260        self.detect_strides(&access);
261        self.track_sequential_access(&access);
262    }
263
264    fn detect_patterns(&mut self, current_access: &MemoryAccess) {
265        let window_size = 8;
266        if self.access_history.len() < window_size {
267            return;
268        }
269
270        let recent: Vec<&MemoryAccess> =
271            self.access_history.iter().rev().take(window_size).collect();
272        let mut deltas = Vec::new();
273
274        for i in 1..recent.len() {
275            let delta = recent[i - 1].address as isize - recent[i].address as isize;
276            deltas.push(delta);
277        }
278
279        // Classify pattern type
280        let pattern_type = if deltas.iter().all(|&d| d == deltas[0]) {
281            if deltas[0] == 0 {
282                PatternType::Random
283            } else if deltas[0].abs() < 128 {
284                PatternType::Sequential
285            } else {
286                PatternType::Strided
287            }
288        } else {
289            PatternType::Irregular
290        };
291
292        let pattern = AccessPattern {
293            pattern_type,
294            deltas,
295            size: window_size,
296        };
297
298        // Update pattern frequency
299        let freq = self
300            .patterns
301            .entry(pattern)
302            .or_insert_with(|| PatternFrequency {
303                count: 0,
304                last_seen: current_access.timestamp,
305                confidence: 0.0,
306                prediction_accuracy: 0.0,
307            });
308
309        freq.count += 1;
310        freq.last_seen = current_access.timestamp;
311        freq.confidence = (freq.count as f64 / 100.0).min(1.0);
312    }
313
314    fn detect_strides(&mut self, current_access: &MemoryAccess) {
315        if self.access_history.len() < 2 {
316            return;
317        }
318
319        let prev_access = &self.access_history[self.access_history.len() - 2];
320        let stride = current_access.address as isize - prev_access.address as isize;
321
322        let stride_info = self
323            .stride_patterns
324            .entry(current_access.context_id as usize)
325            .or_insert_with(|| StrideInfo {
326                stride: 0,
327                frequency: 0,
328                last_address: prev_access.address,
329                confidence: 0.0,
330                start_time: current_access.timestamp,
331            });
332
333        if stride == stride_info.stride {
334            stride_info.frequency += 1;
335            stride_info.confidence = (stride_info.frequency as f64 / 10.0).min(1.0);
336        } else {
337            stride_info.stride = stride;
338            stride_info.frequency = 1;
339            stride_info.confidence = 0.1;
340            stride_info.start_time = current_access.timestamp;
341        }
342
343        stride_info.last_address = current_access.address;
344    }
345
346    fn track_sequential_access(&mut self, current_access: &MemoryAccess) {
347        let seq_info = self
348            .sequential_tracking
349            .entry(current_access.context_id as usize)
350            .or_insert_with(|| SequentialInfo {
351                start_address: current_access.address,
352                current_address: current_access.address,
353                length: 1,
354                direction: 0,
355                last_access: current_access.timestamp,
356            });
357
358        let address_diff = current_access.address as isize - seq_info.current_address as isize;
359
360        if address_diff.abs() <= current_access.size as isize * 2 {
361            // Likely sequential
362            if seq_info.direction == 0 {
363                seq_info.direction = if address_diff > 0 { 1 } else { -1 };
364            }
365
366            if (seq_info.direction > 0 && address_diff > 0)
367                || (seq_info.direction < 0 && address_diff < 0)
368            {
369                seq_info.length += 1;
370                seq_info.current_address = current_access.address;
371                seq_info.last_access = current_access.timestamp;
372            } else {
373                // Reset sequence
374                seq_info.start_address = current_access.address;
375                seq_info.current_address = current_access.address;
376                seq_info.length = 1;
377                seq_info.direction = 0;
378                seq_info.last_access = current_access.timestamp;
379            }
380        } else {
381            // Non-sequential, reset
382            seq_info.start_address = current_access.address;
383            seq_info.current_address = current_access.address;
384            seq_info.length = 1;
385            seq_info.direction = 0;
386            seq_info.last_access = current_access.timestamp;
387        }
388    }
389
390    /// Get predicted next accesses
391    pub fn predict_next_accesses(&self, count: usize) -> Vec<PredictedAccess> {
392        let mut predictions = Vec::new();
393
394        // Sequential predictions
395        for seq_info in self.sequential_tracking.values() {
396            if seq_info.length >= 3 && seq_info.last_access.elapsed() < Duration::from_millis(100) {
397                let next_addr = if seq_info.direction > 0 {
398                    seq_info.current_address + 64 // Typical cache line size
399                } else {
400                    seq_info.current_address.saturating_sub(64)
401                };
402
403                predictions.push(PredictedAccess {
404                    address: next_addr,
405                    size: 64,
406                    confidence: 0.8,
407                    strategy: "Sequential".to_string(),
408                    estimated_time: Duration::from_micros(100),
409                });
410            }
411        }
412
413        // Stride predictions
414        for stride_info in self.stride_patterns.values() {
415            if stride_info.confidence > 0.5 && stride_info.frequency >= 3 {
416                let next_addr = (stride_info.last_address as isize + stride_info.stride) as usize;
417                predictions.push(PredictedAccess {
418                    address: next_addr,
419                    size: 64,
420                    confidence: stride_info.confidence,
421                    strategy: "Stride".to_string(),
422                    estimated_time: Duration::from_micros(150),
423                });
424            }
425        }
426
427        predictions.truncate(count);
428        predictions
429    }
430}
431
432/// Predicted memory access
433#[derive(Debug, Clone)]
434pub struct PredictedAccess {
435    pub address: usize,
436    pub size: usize,
437    pub confidence: f64,
438    pub strategy: String,
439    pub estimated_time: Duration,
440}
441
442/// Prefetch request
443#[derive(Debug, Clone)]
444pub struct PrefetchRequest {
445    /// Target address to prefetch
446    pub address: usize,
447    /// Size to prefetch
448    pub size: usize,
449    /// Priority level
450    pub priority: PrefetchPriority,
451    /// Strategy that generated this request
452    pub strategy: String,
453    /// Confidence in this prefetch
454    pub confidence: f64,
455    /// Request timestamp
456    pub timestamp: Instant,
457    /// Deadline for prefetch completion
458    pub deadline: Option<Instant>,
459}
460
461/// Prefetch priority levels
462#[derive(Debug, Clone, PartialEq, Ord, PartialOrd, Eq)]
463pub enum PrefetchPriority {
464    Low,
465    Normal,
466    High,
467    Critical,
468}
469
470/// Prefetch cache for storing prefetched data
471pub struct PrefetchCache {
472    /// Cache entries
473    entries: BTreeMap<usize, CacheEntry>,
474    /// Cache size limit
475    size_limit: usize,
476    /// Current cache size
477    current_size: usize,
478    /// LRU tracking
479    lru_order: VecDeque<usize>,
480    /// Cache statistics
481    stats: CacheStats,
482}
483
484/// Cache entry
485#[derive(Debug, Clone)]
486pub struct CacheEntry {
487    pub address: usize,
488    pub size: usize,
489    pub data: Vec<u8>,
490    pub prefetch_time: Instant,
491    pub last_access: Option<Instant>,
492    pub access_count: u32,
493    pub strategy: String,
494}
495
496/// Cache statistics
497#[derive(Debug, Clone, Default)]
498pub struct CacheStats {
499    pub hits: u64,
500    pub misses: u64,
501    pub evictions: u64,
502    pub total_size: usize,
503    pub utilization: f64,
504}
505
506impl PrefetchCache {
507    pub fn new(size_limit: usize) -> Self {
508        Self {
509            entries: BTreeMap::new(),
510            size_limit,
511            current_size: 0,
512            lru_order: VecDeque::new(),
513            stats: CacheStats::default(),
514        }
515    }
516
517    /// Insert prefetched data into cache
518    pub fn insert(&mut self, address: usize, data: Vec<u8>, strategy: String) -> bool {
519        let size = data.len();
520
521        // Check if we need to evict entries
522        while self.current_size + size > self.size_limit && !self.entries.is_empty() {
523            self.evict_lru();
524        }
525
526        if self.current_size + size <= self.size_limit {
527            let entry = CacheEntry {
528                address,
529                size,
530                data,
531                prefetch_time: Instant::now(),
532                last_access: None,
533                access_count: 0,
534                strategy,
535            };
536
537            self.entries.insert(address, entry);
538            self.lru_order.push_back(address);
539            self.current_size += size;
540            true
541        } else {
542            false
543        }
544    }
545
546    /// Check if data is in cache and mark as accessed
547    pub fn get(&mut self, address: usize, size: usize) -> Option<&[u8]> {
548        if let Some(entry) = self.entries.get_mut(&address) {
549            if entry.size >= size {
550                entry.last_access = Some(Instant::now());
551                entry.access_count += 1;
552
553                // Update LRU order
554                if let Some(pos) = self.lru_order.iter().position(|&addr| addr == address) {
555                    self.lru_order.remove(pos);
556                    self.lru_order.push_back(address);
557                }
558
559                self.stats.hits += 1;
560                return Some(&entry.data[..size]);
561            }
562        }
563
564        self.stats.misses += 1;
565        None
566    }
567
568    fn evict_lru(&mut self) {
569        if let Some(address) = self.lru_order.pop_front() {
570            if let Some(entry) = self.entries.remove(&address) {
571                self.current_size -= entry.size;
572                self.stats.evictions += 1;
573            }
574        }
575    }
576
577    /// Get cache statistics
578    pub fn get_stats(&self) -> &CacheStats {
579        &self.stats
580    }
581}
582
583/// Prefetch strategy trait
584pub trait PrefetchStrategy: Send + Sync {
585    fn name(&self) -> &str;
586    fn can_prefetch(&self, access: &MemoryAccess, history: &AccessHistoryTracker) -> bool;
587    fn generate_requests(
588        &self,
589        access: &MemoryAccess,
590        history: &AccessHistoryTracker,
591    ) -> Vec<PrefetchRequest>;
592    fn get_statistics(&self) -> StrategyStats;
593    fn configure(&mut self, config: &PrefetchConfig);
594}
595
596/// Sequential prefetching strategy
597pub struct SequentialPrefetcher {
598    stats: StrategyStats,
599    config: SequentialConfig,
600}
601
602/// Sequential prefetcher configuration
603#[derive(Debug, Clone)]
604pub struct SequentialConfig {
605    pub prefetch_distance: usize,
606    pub min_sequence_length: usize,
607    pub max_prefetch_count: usize,
608}
609
610impl Default for SequentialConfig {
611    fn default() -> Self {
612        Self {
613            prefetch_distance: 1024,
614            min_sequence_length: 3,
615            max_prefetch_count: 8,
616        }
617    }
618}
619
620impl SequentialPrefetcher {
621    pub fn new(config: SequentialConfig) -> Self {
622        Self {
623            stats: StrategyStats::default(),
624            config,
625        }
626    }
627}
628
629impl PrefetchStrategy for SequentialPrefetcher {
630    fn name(&self) -> &str {
631        "Sequential"
632    }
633
634    fn can_prefetch(&self, access: &MemoryAccess, history: &AccessHistoryTracker) -> bool {
635        if let Some(seq_info) = history
636            .sequential_tracking
637            .get(&(access.context_id as usize))
638        {
639            seq_info.length >= self.config.min_sequence_length
640                && seq_info.last_access.elapsed() < Duration::from_millis(50)
641        } else {
642            false
643        }
644    }
645
646    fn generate_requests(
647        &self,
648        access: &MemoryAccess,
649        history: &AccessHistoryTracker,
650    ) -> Vec<PrefetchRequest> {
651        let mut requests = Vec::new();
652
653        if let Some(seq_info) = history
654            .sequential_tracking
655            .get(&(access.context_id as usize))
656        {
657            let mut next_addr = access.address;
658            let step = 64; // Cache line size
659
660            for i in 0..self.config.max_prefetch_count {
661                next_addr = if seq_info.direction > 0 {
662                    next_addr + step
663                } else {
664                    next_addr.saturating_sub(step)
665                };
666
667                if (next_addr as isize - access.address as isize).abs()
668                    > self.config.prefetch_distance as isize
669                {
670                    break;
671                }
672
673                let confidence = (1.0 - i as f64 * 0.1).max(0.1);
674
675                requests.push(PrefetchRequest {
676                    address: next_addr,
677                    size: 64,
678                    priority: PrefetchPriority::Normal,
679                    strategy: self.name().to_string(),
680                    confidence,
681                    timestamp: Instant::now(),
682                    deadline: Some(Instant::now() + Duration::from_millis(10)),
683                });
684            }
685        }
686
687        requests
688    }
689
690    fn get_statistics(&self) -> StrategyStats {
691        self.stats.clone()
692    }
693
694    fn configure(&mut self, config: &PrefetchConfig) {
695        self.config.prefetch_distance = config.max_prefetch_distance;
696        self.config.max_prefetch_count = config.prefetch_window;
697    }
698}
699
700/// Stride-based prefetching strategy
701pub struct StridePrefetcher {
702    stats: StrategyStats,
703    config: StrideConfig,
704}
705
706/// Stride prefetcher configuration
707#[derive(Debug, Clone)]
708pub struct StrideConfig {
709    pub min_confidence: f64,
710    pub max_stride: isize,
711    pub prefetch_degree: usize,
712}
713
714impl Default for StrideConfig {
715    fn default() -> Self {
716        Self {
717            min_confidence: 0.6,
718            max_stride: 4096,
719            prefetch_degree: 4,
720        }
721    }
722}
723
724impl StridePrefetcher {
725    pub fn new(config: StrideConfig) -> Self {
726        Self {
727            stats: StrategyStats::default(),
728            config,
729        }
730    }
731}
732
733impl PrefetchStrategy for StridePrefetcher {
734    fn name(&self) -> &str {
735        "Stride"
736    }
737
738    fn can_prefetch(&self, access: &MemoryAccess, history: &AccessHistoryTracker) -> bool {
739        if let Some(stride_info) = history.stride_patterns.get(&(access.context_id as usize)) {
740            stride_info.confidence >= self.config.min_confidence
741                && stride_info.stride.abs() <= self.config.max_stride
742                && stride_info.stride != 0
743        } else {
744            false
745        }
746    }
747
748    fn generate_requests(
749        &self,
750        access: &MemoryAccess,
751        history: &AccessHistoryTracker,
752    ) -> Vec<PrefetchRequest> {
753        let mut requests = Vec::new();
754
755        if let Some(stride_info) = history.stride_patterns.get(&(access.context_id as usize)) {
756            let mut next_addr = access.address;
757
758            for i in 0..self.config.prefetch_degree {
759                next_addr = (next_addr as isize + stride_info.stride) as usize;
760                let confidence = stride_info.confidence * (1.0 - i as f64 * 0.15);
761
762                requests.push(PrefetchRequest {
763                    address: next_addr,
764                    size: access.size,
765                    priority: PrefetchPriority::Normal,
766                    strategy: self.name().to_string(),
767                    confidence,
768                    timestamp: Instant::now(),
769                    deadline: Some(Instant::now() + Duration::from_millis(15)),
770                });
771            }
772        }
773
774        requests
775    }
776
777    fn get_statistics(&self) -> StrategyStats {
778        self.stats.clone()
779    }
780
781    fn configure(&mut self, config: &PrefetchConfig) {
782        self.config.prefetch_degree = config.prefetch_window;
783    }
784}
785
786/// Performance monitoring for prefetching
787pub struct PerformanceMonitor {
788    /// Performance history
789    history: VecDeque<PerfSample>,
790    /// Current metrics
791    current_metrics: PerfMetrics,
792    /// Monitoring configuration
793    config: MonitorConfig,
794}
795
796/// Performance sample
797#[derive(Debug, Clone)]
798pub struct PerfSample {
799    pub timestamp: Instant,
800    pub cache_hit_rate: f64,
801    pub prefetch_accuracy: f64,
802    pub bandwidth_utilization: f64,
803    pub latency: Duration,
804}
805
806/// Performance metrics
807#[derive(Debug, Clone, Default)]
808pub struct PerfMetrics {
809    pub average_hit_rate: f64,
810    pub average_accuracy: f64,
811    pub average_bandwidth: f64,
812    pub average_latency: Duration,
813    pub trend_hit_rate: f64,
814    pub trend_accuracy: f64,
815}
816
817/// Monitor configuration
818#[derive(Debug, Clone)]
819pub struct MonitorConfig {
820    pub sample_interval: Duration,
821    pub history_size: usize,
822    pub enable_trends: bool,
823}
824
825impl Default for MonitorConfig {
826    fn default() -> Self {
827        Self {
828            sample_interval: Duration::from_secs(1),
829            history_size: 100,
830            enable_trends: true,
831        }
832    }
833}
834
835impl PerformanceMonitor {
836    pub fn new(config: MonitorConfig) -> Self {
837        Self {
838            history: VecDeque::with_capacity(config.history_size),
839            current_metrics: PerfMetrics::default(),
840            config,
841        }
842    }
843
844    /// Record a performance sample
845    pub fn record_sample(&mut self, sample: PerfSample) {
846        self.history.push_back(sample);
847        if self.history.len() > self.config.history_size {
848            self.history.pop_front();
849        }
850
851        self.update_metrics();
852    }
853
854    fn update_metrics(&mut self) {
855        if self.history.is_empty() {
856            return;
857        }
858
859        let count = self.history.len() as f64;
860        self.current_metrics.average_hit_rate =
861            self.history.iter().map(|s| s.cache_hit_rate).sum::<f64>() / count;
862        self.current_metrics.average_accuracy = self
863            .history
864            .iter()
865            .map(|s| s.prefetch_accuracy)
866            .sum::<f64>()
867            / count;
868        self.current_metrics.average_bandwidth = self
869            .history
870            .iter()
871            .map(|s| s.bandwidth_utilization)
872            .sum::<f64>()
873            / count;
874
875        let total_latency_nanos: u64 = self
876            .history
877            .iter()
878            .map(|s| s.latency.as_nanos() as u64)
879            .sum();
880        self.current_metrics.average_latency =
881            Duration::from_nanos(total_latency_nanos / count as u64);
882
883        // Calculate trends
884        if self.config.enable_trends && self.history.len() >= 10 {
885            let recent_hit_rate: f64 = self
886                .history
887                .iter()
888                .rev()
889                .take(5)
890                .map(|s| s.cache_hit_rate)
891                .sum::<f64>()
892                / 5.0;
893            let older_hit_rate: f64 = self
894                .history
895                .iter()
896                .rev()
897                .skip(5)
898                .take(5)
899                .map(|s| s.cache_hit_rate)
900                .sum::<f64>()
901                / 5.0;
902            self.current_metrics.trend_hit_rate = recent_hit_rate - older_hit_rate;
903
904            let recent_accuracy: f64 = self
905                .history
906                .iter()
907                .rev()
908                .take(5)
909                .map(|s| s.prefetch_accuracy)
910                .sum::<f64>()
911                / 5.0;
912            let older_accuracy: f64 = self
913                .history
914                .iter()
915                .rev()
916                .skip(5)
917                .take(5)
918                .map(|s| s.prefetch_accuracy)
919                .sum::<f64>()
920                / 5.0;
921            self.current_metrics.trend_accuracy = recent_accuracy - older_accuracy;
922        }
923    }
924
925    /// Get current performance metrics
926    pub fn get_metrics(&self) -> &PerfMetrics {
927        &self.current_metrics
928    }
929}
930
931impl PrefetchingEngine {
932    pub fn new(config: PrefetchConfig) -> Self {
933        let mut strategies: Vec<Box<dyn PrefetchStrategy>> = Vec::new();
934
935        if config.enable_pattern_based {
936            strategies.push(Box::new(SequentialPrefetcher::new(
937                SequentialConfig::default(),
938            )));
939        }
940
941        if config.enable_stride_based {
942            strategies.push(Box::new(StridePrefetcher::new(StrideConfig::default())));
943        }
944
945        let access_history = AccessHistoryTracker::new(config.history_window);
946        let prefetch_cache = PrefetchCache::new(config.cache_size);
947        let performance_monitor = PerformanceMonitor::new(MonitorConfig::default());
948
949        Self {
950            config,
951            stats: PrefetchStats::default(),
952            strategies,
953            access_history,
954            prefetch_queue: VecDeque::new(),
955            prefetch_cache,
956            performance_monitor,
957        }
958    }
959
960    /// Record a memory access and potentially trigger prefetching
961    pub fn record_access(&mut self, access: MemoryAccess) -> Vec<PrefetchRequest> {
962        // Record access in history
963        self.access_history.record_access(access.clone());
964
965        // Check cache for hit/miss
966        let cache_hit = self
967            .prefetch_cache
968            .get(access.address, access.size)
969            .is_some();
970        if cache_hit {
971            self.stats.successful_prefetches += 1;
972        }
973
974        let mut new_requests = Vec::new();
975
976        if self.config.auto_prefetch {
977            // Generate prefetch requests from strategies
978            for strategy in &self.strategies {
979                if strategy.can_prefetch(&access, &self.access_history) {
980                    let requests = strategy.generate_requests(&access, &self.access_history);
981                    for request in requests {
982                        if self.should_issue_prefetch(&request) {
983                            new_requests.push(request);
984                        }
985                    }
986                }
987            }
988
989            // Add requests to queue
990            for request in &new_requests {
991                self.prefetch_queue.push_back(request.clone());
992                self.stats.total_requests += 1;
993            }
994        }
995
996        new_requests
997    }
998
999    fn should_issue_prefetch(&self, request: &PrefetchRequest) -> bool {
1000        // Check if already in cache
1001        if self.prefetch_cache.entries.contains_key(&request.address) {
1002            return false;
1003        }
1004
1005        // Check confidence threshold
1006        if request.confidence < self.config.min_access_frequency {
1007            return false;
1008        }
1009
1010        // Check prefetch distance
1011        if request.size > self.config.max_prefetch_distance {
1012            return false;
1013        }
1014
1015        true
1016    }
1017
1018    /// Process prefetch queue and issue prefetches
1019    pub fn process_prefetch_queue(&mut self) -> Vec<PrefetchRequest> {
1020        let mut issued_requests = Vec::new();
1021        let max_concurrent = (self.config.aggressiveness * 10.0) as usize + 1;
1022
1023        // Sort by priority and confidence
1024        let mut pending: Vec<PrefetchRequest> = self.prefetch_queue.drain(..).collect();
1025        pending.sort_by(|a, b| {
1026            b.priority.cmp(&a.priority).then_with(|| {
1027                b.confidence
1028                    .partial_cmp(&a.confidence)
1029                    .unwrap_or(std::cmp::Ordering::Equal)
1030            })
1031        });
1032
1033        for request in pending.into_iter().take(max_concurrent) {
1034            // Simulate prefetch (in real implementation, this would trigger actual memory load)
1035            let dummy_data = vec![0u8; request.size];
1036            if self
1037                .prefetch_cache
1038                .insert(request.address, dummy_data, request.strategy.clone())
1039            {
1040                issued_requests.push(request);
1041            }
1042        }
1043
1044        issued_requests
1045    }
1046
1047    /// Update statistics and performance metrics
1048    pub fn update_performance(&mut self) {
1049        let cache_stats = self.prefetch_cache.get_stats();
1050
1051        // Update accuracy ratio
1052        let total_prefetches = self.stats.successful_prefetches + self.stats.failed_prefetches;
1053        if total_prefetches > 0 {
1054            self.stats.accuracy_ratio =
1055                self.stats.successful_prefetches as f64 / total_prefetches as f64;
1056        }
1057
1058        // Update cache hit rate
1059        let total_accesses = cache_stats.hits + cache_stats.misses;
1060        if total_accesses > 0 {
1061            self.stats.cache_hit_rate = cache_stats.hits as f64 / total_accesses as f64;
1062        }
1063
1064        // Record performance sample
1065        let sample = PerfSample {
1066            timestamp: Instant::now(),
1067            cache_hit_rate: self.stats.cache_hit_rate,
1068            prefetch_accuracy: self.stats.accuracy_ratio,
1069            bandwidth_utilization: 0.8, // Would be calculated from actual usage
1070            latency: self.stats.average_latency,
1071        };
1072
1073        self.performance_monitor.record_sample(sample);
1074    }
1075
1076    /// Get prefetching statistics
1077    pub fn get_stats(&self) -> &PrefetchStats {
1078        &self.stats
1079    }
1080
1081    /// Get cache statistics
1082    pub fn get_cache_stats(&self) -> &CacheStats {
1083        self.prefetch_cache.get_stats()
1084    }
1085
1086    /// Get performance metrics
1087    pub fn get_performance_metrics(&self) -> &PerfMetrics {
1088        self.performance_monitor.get_metrics()
1089    }
1090
1091    /// Get access history
1092    pub fn get_access_history(&self) -> &AccessHistoryTracker {
1093        &self.access_history
1094    }
1095
1096    /// Configure prefetching engine
1097    pub fn configure(&mut self, config: PrefetchConfig) {
1098        self.config = config.clone();
1099        for strategy in &mut self.strategies {
1100            strategy.configure(&config);
1101        }
1102    }
1103}
1104
1105/// Thread-safe prefetching engine wrapper
1106pub struct ThreadSafePrefetchingEngine {
1107    engine: Arc<Mutex<PrefetchingEngine>>,
1108}
1109
1110impl ThreadSafePrefetchingEngine {
1111    pub fn new(config: PrefetchConfig) -> Self {
1112        Self {
1113            engine: Arc::new(Mutex::new(PrefetchingEngine::new(config))),
1114        }
1115    }
1116
1117    pub fn record_access(&self, access: MemoryAccess) -> Vec<PrefetchRequest> {
1118        let mut engine = self.engine.lock().unwrap_or_else(|e| e.into_inner());
1119        engine.record_access(access)
1120    }
1121
1122    pub fn process_prefetch_queue(&self) -> Vec<PrefetchRequest> {
1123        let mut engine = self.engine.lock().unwrap_or_else(|e| e.into_inner());
1124        engine.process_prefetch_queue()
1125    }
1126
1127    pub fn get_stats(&self) -> PrefetchStats {
1128        let engine = self.engine.lock().unwrap_or_else(|e| e.into_inner());
1129        engine.get_stats().clone()
1130    }
1131
1132    pub fn update_performance(&self) {
1133        let mut engine = self.engine.lock().unwrap_or_else(|e| e.into_inner());
1134        engine.update_performance();
1135    }
1136}
1137
1138#[cfg(test)]
1139mod tests {
1140    use super::*;
1141
1142    #[test]
1143    fn test_prefetch_engine_creation() {
1144        let config = PrefetchConfig::default();
1145        let engine = PrefetchingEngine::new(config);
1146        assert!(!engine.strategies.is_empty());
1147    }
1148
1149    #[test]
1150    fn test_access_history_tracking() {
1151        let mut tracker = AccessHistoryTracker::new(100);
1152
1153        let access = MemoryAccess {
1154            address: 0x1000,
1155            size: 64,
1156            timestamp: Instant::now(),
1157            access_type: AccessType::Read,
1158            context_id: 1,
1159            kernel_id: Some(100),
1160        };
1161
1162        tracker.record_access(access);
1163        assert_eq!(tracker.access_history.len(), 1);
1164    }
1165
1166    #[test]
1167    fn test_prefetch_cache() {
1168        let mut cache = PrefetchCache::new(1024);
1169
1170        let data = vec![1, 2, 3, 4];
1171        assert!(cache.insert(0x1000, data, "Test".to_string()));
1172
1173        let retrieved = cache.get(0x1000, 4);
1174        assert!(retrieved.is_some());
1175        assert_eq!(retrieved.expect("unwrap failed"), &[1, 2, 3, 4]);
1176    }
1177
1178    #[test]
1179    fn test_sequential_prefetcher() {
1180        let config = SequentialConfig::default();
1181        let prefetcher = SequentialPrefetcher::new(config);
1182        assert_eq!(prefetcher.name(), "Sequential");
1183    }
1184
1185    #[test]
1186    fn test_stride_prefetcher() {
1187        let config = StrideConfig::default();
1188        let prefetcher = StridePrefetcher::new(config);
1189        assert_eq!(prefetcher.name(), "Stride");
1190    }
1191
1192    #[test]
1193    fn test_performance_monitor() {
1194        let config = MonitorConfig::default();
1195        let mut monitor = PerformanceMonitor::new(config);
1196
1197        let sample = PerfSample {
1198            timestamp: Instant::now(),
1199            cache_hit_rate: 0.8,
1200            prefetch_accuracy: 0.7,
1201            bandwidth_utilization: 0.9,
1202            latency: Duration::from_millis(5),
1203        };
1204
1205        monitor.record_sample(sample);
1206        let metrics = monitor.get_metrics();
1207        assert!(metrics.average_hit_rate > 0.0);
1208    }
1209
1210    #[test]
1211    fn test_thread_safe_engine() {
1212        let config = PrefetchConfig::default();
1213        let engine = ThreadSafePrefetchingEngine::new(config);
1214
1215        let access = MemoryAccess {
1216            address: 0x2000,
1217            size: 128,
1218            timestamp: Instant::now(),
1219            access_type: AccessType::Read,
1220            context_id: 2,
1221            kernel_id: Some(200),
1222        };
1223
1224        // A single access to a brand-new context has no established
1225        // pattern yet (every strategy requires a minimum run/frequency
1226        // before it will fire), so this must not panic and must not
1227        // fabricate a prefetch request out of thin air.
1228        let requests = engine.record_access(access);
1229        assert!(requests.is_empty());
1230    }
1231}