1pub mod optimizer;
2
3use crate::errors::{Result, TrustformersError};
4use crate::tensor::Tensor;
5use scirs2_core::ndarray::{s, IxDyn};
6use serde::{Deserialize, Serialize};
7use std::collections::HashMap;
8use std::fs::File;
9use std::io::{Read, Seek, SeekFrom};
10use std::sync::{Arc, Mutex, RwLock};
11use std::time::{Duration, Instant};
12
13#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize)]
23pub enum MemoryEvictionPolicy {
24 LRU,
26 LFU,
28 SizeBased,
30 ARC,
32 Hybrid,
34}
35
36#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize)]
38pub enum AdaptiveStrategy {
39 Fixed,
41 MemoryPressure,
43 HitRate,
45 Predictive,
47}
48
49#[derive(Debug, Clone)]
51pub struct MemoryConfig {
52 pub enable_memory_pool: bool,
54 pub max_pool_size: usize,
56 pub min_pool_size: usize,
58 pub enable_zero_copy: bool,
60 pub enable_mmap: bool,
62 pub mmap_threshold: usize,
64 pub cleanup_interval: Duration,
66 pub eviction_policy: MemoryEvictionPolicy,
68 pub adaptive_strategy: AdaptiveStrategy,
70 pub target_hit_rate: f64,
72 pub enable_prefetching: bool,
74 pub enable_defragmentation: bool,
76}
77
78impl Default for MemoryConfig {
79 fn default() -> Self {
80 Self {
81 enable_memory_pool: true,
82 max_pool_size: 1024 * 1024 * 1024, min_pool_size: 64 * 1024 * 1024, enable_zero_copy: true,
85 enable_mmap: true,
86 mmap_threshold: 100 * 1024 * 1024, cleanup_interval: Duration::from_secs(60),
88 eviction_policy: MemoryEvictionPolicy::Hybrid,
89 adaptive_strategy: AdaptiveStrategy::HitRate,
90 target_hit_rate: 0.85, enable_prefetching: true,
92 enable_defragmentation: true,
93 }
94 }
95}
96
97#[derive(Debug, Clone)]
99struct PoolEntry {
100 tensor: Tensor,
101 last_used: Instant,
102 ref_count: usize,
103 access_count: usize,
105 #[allow(dead_code)]
107 created_at: Instant,
108 #[allow(dead_code)]
110 pool_time: Duration,
111 size_bytes: usize,
113}
114
115impl PoolEntry {
116 fn new(tensor: Tensor, size_bytes: usize) -> Self {
117 let now = Instant::now();
118 Self {
119 tensor,
120 last_used: now,
121 ref_count: 0,
122 access_count: 0,
123 created_at: now,
124 pool_time: Duration::ZERO,
125 size_bytes,
126 }
127 }
128
129 fn mark_accessed(&mut self) {
130 self.last_used = Instant::now();
131 self.access_count += 1;
132 }
133
134 fn eviction_priority(&self, policy: MemoryEvictionPolicy) -> f64 {
136 match policy {
137 MemoryEvictionPolicy::LRU => {
138 -(self.last_used.elapsed().as_secs_f64())
140 },
141 MemoryEvictionPolicy::LFU => {
142 -(self.access_count as f64)
144 },
145 MemoryEvictionPolicy::SizeBased => {
146 -(self.size_bytes as f64)
148 },
149 MemoryEvictionPolicy::ARC => {
150 let recency_score = 1.0 / (1.0 + self.last_used.elapsed().as_secs_f64());
152 let frequency_score = self.access_count as f64;
153 -(recency_score + frequency_score)
154 },
155 MemoryEvictionPolicy::Hybrid => {
156 let recency = 1.0 / (1.0 + self.last_used.elapsed().as_secs_f64());
158 let frequency = self.access_count as f64;
159 let size_factor = 1.0 / (1.0 + (self.size_bytes as f64 / 1_000_000.0));
160 -(recency * 0.4 + frequency * 0.4 + size_factor * 0.2)
161 },
162 }
163 }
164}
165
166#[derive(Debug)]
168pub struct TensorView {
169 original: Arc<Tensor>,
171 offset: usize,
173 shape: Vec<usize>,
175 #[allow(dead_code)]
177 strides: Vec<usize>,
178}
179
180impl TensorView {
181 pub fn slice(tensor: Arc<Tensor>, start: usize, end: usize) -> Result<Self> {
183 let original_shape = tensor.shape();
184 if start >= end || end > original_shape.iter().product::<usize>() {
185 return Err(TrustformersError::invalid_input(
186 "Invalid slice bounds".to_string(),
187 ));
188 }
189
190 let slice_len = end - start;
191 Ok(Self {
192 original: tensor,
193 offset: start,
194 shape: vec![slice_len],
195 strides: vec![1],
196 })
197 }
198
199 pub fn shape(&self) -> &[usize] {
201 &self.shape
202 }
203
204 pub fn as_tensor(&self) -> Result<Tensor> {
206 match &*self.original {
209 Tensor::F32(arr) => {
210 let flat = arr
211 .view()
212 .into_shape_with_order(arr.len())
213 .map_err(|e| TrustformersError::shape_error(e.to_string()))?;
214 let slice = flat.slice(s![
215 self.offset..self.offset + self.shape.iter().product::<usize>()
216 ]);
217 let sliced_arr = slice
218 .to_owned()
219 .into_shape_with_order(IxDyn(&self.shape))
220 .map_err(|e| TrustformersError::shape_error(e.to_string()))?;
221 Ok(Tensor::F32(sliced_arr))
222 },
223 _ => Err(TrustformersError::tensor_op_error(
224 "Zero-copy slicing not implemented for this tensor type",
225 "zero_copy_slice",
226 )),
227 }
228 }
229}
230
231#[derive(Debug, Clone)]
233struct PoolStatistics {
234 total_requests: usize,
235 cache_hits: usize,
236 cache_misses: usize,
237 total_evictions: usize,
238 evictions_by_policy: HashMap<String, usize>,
239 total_allocated_bytes: usize,
240 peak_memory_usage: usize,
241 #[allow(dead_code)]
242 average_tensor_lifetime: Duration,
243 #[allow(dead_code)]
244 last_reset: Instant,
245}
246
247impl Default for PoolStatistics {
248 fn default() -> Self {
249 Self {
250 total_requests: 0,
251 cache_hits: 0,
252 cache_misses: 0,
253 total_evictions: 0,
254 evictions_by_policy: HashMap::new(),
255 total_allocated_bytes: 0,
256 peak_memory_usage: 0,
257 average_tensor_lifetime: Duration::ZERO,
258 last_reset: Instant::now(),
259 }
260 }
261}
262
263impl PoolStatistics {
264 fn hit_rate(&self) -> f64 {
265 if self.total_requests == 0 {
266 0.0
267 } else {
268 self.cache_hits as f64 / self.total_requests as f64
269 }
270 }
271
272 fn miss_rate(&self) -> f64 {
273 if self.total_requests == 0 {
274 0.0
275 } else {
276 self.cache_misses as f64 / self.total_requests as f64
277 }
278 }
279}
280
281pub struct TensorMemoryPool {
283 config: MemoryConfig,
284 pool: Arc<RwLock<HashMap<Vec<usize>, Vec<PoolEntry>>>>,
285 current_size: Arc<Mutex<usize>>,
286 last_cleanup: Arc<Mutex<Instant>>,
287 statistics: Arc<Mutex<PoolStatistics>>,
289 access_patterns: Arc<Mutex<HashMap<Vec<usize>, Vec<Instant>>>>,
291 dynamic_max_size: Arc<Mutex<usize>>,
293}
294
295impl TensorMemoryPool {
296 pub fn new(config: MemoryConfig) -> Self {
298 let dynamic_max_size = config.max_pool_size;
299 Self {
300 config,
301 pool: Arc::new(RwLock::new(HashMap::new())),
302 current_size: Arc::new(Mutex::new(0)),
303 last_cleanup: Arc::new(Mutex::new(Instant::now())),
304 statistics: Arc::new(Mutex::new(PoolStatistics::default())),
305 access_patterns: Arc::new(Mutex::new(HashMap::new())),
306 dynamic_max_size: Arc::new(Mutex::new(dynamic_max_size)),
307 }
308 }
309
310 pub fn get_tensor(&self, shape: &[usize], dtype: crate::tensor::DType) -> Result<Tensor> {
312 if self.config.enable_prefetching {
314 let mut patterns =
315 self.access_patterns.lock().unwrap_or_else(|poisoned| poisoned.into_inner());
316 patterns.entry(shape.to_vec()).or_default().push(Instant::now());
317 }
318
319 {
321 let mut stats = self.statistics.lock().unwrap_or_else(|poisoned| poisoned.into_inner());
322 stats.total_requests += 1;
323 }
324
325 if !self.config.enable_memory_pool {
326 return self.create_tensor(shape, dtype);
327 }
328
329 if let Some(tensor) = self.try_get_from_pool(shape)? {
331 let mut stats = self.statistics.lock().unwrap_or_else(|poisoned| poisoned.into_inner());
333 stats.cache_hits += 1;
334 return Ok(tensor);
335 }
336
337 {
339 let mut stats = self.statistics.lock().unwrap_or_else(|poisoned| poisoned.into_inner());
340 stats.cache_misses += 1;
341 }
342
343 self.apply_adaptive_sizing()?;
345
346 self.create_tensor(shape, dtype)
348 }
349
350 pub fn return_tensor(&self, tensor: Tensor) -> Result<()> {
352 if !self.config.enable_memory_pool {
353 return Ok(()); }
355
356 let shape = tensor.shape().to_vec();
357
358 let tensor_size = self.estimate_tensor_size(&tensor);
360
361 let entry = PoolEntry::new(tensor, tensor_size);
363
364 let mut pool = self.pool.write().unwrap_or_else(|poisoned| poisoned.into_inner());
365 pool.entry(shape).or_default().push(entry);
366
367 {
369 let mut current =
370 self.current_size.lock().unwrap_or_else(|poisoned| poisoned.into_inner());
371 *current += tensor_size;
372
373 let mut stats = self.statistics.lock().unwrap_or_else(|poisoned| poisoned.into_inner());
374 if *current > stats.peak_memory_usage {
375 stats.peak_memory_usage = *current;
376 }
377 stats.total_allocated_bytes += tensor_size;
378 }
379
380 self.cleanup_if_needed()?;
382
383 Ok(())
384 }
385
386 fn try_get_from_pool(&self, shape: &[usize]) -> Result<Option<Tensor>> {
388 let mut pool = self.pool.write().unwrap_or_else(|poisoned| poisoned.into_inner());
389
390 if let Some(entries) = pool.get_mut(shape) {
391 if let Some(mut entry) = entries.pop() {
392 entry.mark_accessed();
394
395 let tensor_size = entry.size_bytes;
396 *self.current_size.lock().unwrap_or_else(|poisoned| poisoned.into_inner()) -=
397 tensor_size;
398 return Ok(Some(entry.tensor));
399 }
400 }
401
402 Ok(None)
403 }
404
405 fn create_tensor(&self, shape: &[usize], dtype: crate::tensor::DType) -> Result<Tensor> {
407 match dtype {
408 crate::tensor::DType::F32 => Tensor::zeros(shape),
409 crate::tensor::DType::F64 => Tensor::zeros_f64(shape),
410 crate::tensor::DType::F16 => Tensor::zeros_f16(shape),
411 crate::tensor::DType::BF16 => Tensor::zeros_bf16(shape),
412 crate::tensor::DType::I64 => Tensor::zeros_i64(shape),
413 crate::tensor::DType::C32 => Tensor::zeros_c32(shape),
414 crate::tensor::DType::C64 => Tensor::zeros_c64(shape),
415 crate::tensor::DType::CF16 => Tensor::zeros_cf16(shape),
416 crate::tensor::DType::CBF16 => Tensor::zeros_cbf16(shape),
417 _ => Err(TrustformersError::tensor_op_error(
418 &format!("Tensor creation not implemented for dtype: {:?} - only supported types are F32, F64, F16, BF16, I64, C32, C64, CF16, CBF16", dtype),
419 "create_tensor"
420 )),
421 }
422 }
423
424 fn estimate_tensor_size(&self, tensor: &Tensor) -> usize {
426 let elements = tensor.shape().iter().product::<usize>();
427 match tensor {
428 Tensor::F32(_) => elements * 4, Tensor::F64(_) => elements * 8, Tensor::F16(_) => elements * 2, Tensor::BF16(_) => elements * 2, Tensor::I64(_) => elements * 8, Tensor::C32(_) => elements * 8, Tensor::C64(_) => elements * 16, Tensor::CF16(_) => elements * 4, Tensor::CBF16(_) => elements * 4, #[cfg(feature = "candle")]
438 Tensor::Candle(_) => elements * 4, #[cfg(all(target_os = "macos", feature = "metal"))]
440 Tensor::Metal(data) => elements * data.dtype.size_in_bytes(),
441 #[cfg(feature = "cuda")]
442 Tensor::CUDA(data) => elements * data.dtype.size_in_bytes(),
443 Tensor::Sparse(sparse) => {
444 let nnz = sparse.nnz();
446 nnz * 4 + nnz * std::mem::size_of::<usize>() },
448 }
449 }
450
451 fn cleanup_if_needed(&self) -> Result<()> {
453 let mut last_cleanup =
454 self.last_cleanup.lock().unwrap_or_else(|poisoned| poisoned.into_inner());
455 let should_cleanup_time = last_cleanup.elapsed() >= self.config.cleanup_interval;
456
457 let current_size =
458 *self.current_size.lock().unwrap_or_else(|poisoned| poisoned.into_inner());
459 let dynamic_max =
460 *self.dynamic_max_size.lock().unwrap_or_else(|poisoned| poisoned.into_inner());
461 let should_cleanup_size = current_size > dynamic_max;
462
463 if !should_cleanup_time && !should_cleanup_size {
464 return Ok(());
465 }
466
467 let mut pool = self.pool.write().unwrap_or_else(|poisoned| poisoned.into_inner());
469 let mut total_freed = 0;
470 let mut eviction_count = 0;
471 let policy = self.config.eviction_policy;
472
473 let target_size = (dynamic_max as f64 * 0.85) as usize; let need_to_free = current_size.saturating_sub(target_size);
476
477 let mut all_entries: Vec<(Vec<usize>, usize, f64)> = Vec::new();
479
480 for (shape, entries) in pool.iter() {
481 for (idx, entry) in entries.iter().enumerate() {
482 if entry.ref_count == 0 {
483 let priority = entry.eviction_priority(policy);
484 all_entries.push((shape.clone(), idx, priority));
485 }
486 }
487 }
488
489 all_entries.sort_by(|a, b| a.2.partial_cmp(&b.2).unwrap_or(std::cmp::Ordering::Equal));
491
492 let mut freed_so_far = 0;
494 let mut shapes_to_remove: Vec<Vec<usize>> = Vec::new();
495
496 for (shape, _, _) in all_entries.iter() {
497 if freed_so_far >= need_to_free {
498 break;
499 }
500
501 if let Some(entries) = pool.get_mut(shape) {
502 if let Some(entry) = entries.first() {
503 if entry.ref_count == 0 {
504 let size = entry.size_bytes;
505 freed_so_far += size;
506 total_freed += size;
507 eviction_count += 1;
508 shapes_to_remove.push(shape.clone());
509 }
510 }
511 }
512 }
513
514 for shape in shapes_to_remove {
516 if let Some(entries) = pool.get_mut(&shape) {
517 if !entries.is_empty() {
518 entries.remove(0);
519 }
520 }
521 }
522
523 pool.retain(|_, entries| !entries.is_empty());
525
526 drop(pool); {
530 let mut stats = self.statistics.lock().unwrap_or_else(|poisoned| poisoned.into_inner());
531 stats.total_evictions += eviction_count;
532 *stats.evictions_by_policy.entry(format!("{:?}", policy)).or_insert(0) +=
533 eviction_count;
534 }
535
536 *self.current_size.lock().unwrap_or_else(|poisoned| poisoned.into_inner()) -= total_freed;
538 *last_cleanup = Instant::now();
539
540 if self.config.enable_defragmentation {
542 self.defragment_pool()?;
543 }
544
545 Ok(())
546 }
547
548 fn apply_adaptive_sizing(&self) -> Result<()> {
550 match self.config.adaptive_strategy {
551 AdaptiveStrategy::Fixed => Ok(()), AdaptiveStrategy::HitRate => self.adapt_by_hit_rate(),
553 AdaptiveStrategy::MemoryPressure => self.adapt_by_memory_pressure(),
554 AdaptiveStrategy::Predictive => self.adapt_by_prediction(),
555 }
556 }
557
558 fn adapt_by_hit_rate(&self) -> Result<()> {
560 let stats = self.statistics.lock().unwrap_or_else(|poisoned| poisoned.into_inner());
561 let hit_rate = stats.hit_rate();
562 drop(stats);
563
564 let mut dynamic_max =
565 self.dynamic_max_size.lock().unwrap_or_else(|poisoned| poisoned.into_inner());
566 let target_rate = self.config.target_hit_rate;
567
568 if hit_rate < target_rate {
569 let increase = (*dynamic_max as f64 * 0.1) as usize;
571 let new_size = (*dynamic_max + increase).min(self.config.max_pool_size);
572 if new_size > *dynamic_max {
573 *dynamic_max = new_size;
574 }
575 } else if hit_rate > target_rate + 0.1 {
576 let decrease = (*dynamic_max as f64 * 0.05) as usize;
578 let new_size = (*dynamic_max - decrease).max(self.config.min_pool_size);
579 if new_size < *dynamic_max {
580 *dynamic_max = new_size;
581 }
582 }
583
584 Ok(())
585 }
586
587 fn adapt_by_memory_pressure(&self) -> Result<()> {
589 let current_size =
592 *self.current_size.lock().unwrap_or_else(|poisoned| poisoned.into_inner());
593 let mut dynamic_max =
594 self.dynamic_max_size.lock().unwrap_or_else(|poisoned| poisoned.into_inner());
595
596 let utilization = current_size as f64 / *dynamic_max as f64;
597
598 if utilization > 0.9 {
599 let new_size = (*dynamic_max as f64 * 0.9) as usize;
601 *dynamic_max = new_size.max(self.config.min_pool_size);
602 } else if utilization < 0.5 {
603 let new_size = (*dynamic_max as f64 * 1.1) as usize;
605 *dynamic_max = new_size.min(self.config.max_pool_size);
606 }
607
608 Ok(())
609 }
610
611 fn adapt_by_prediction(&self) -> Result<()> {
613 let patterns = self.access_patterns.lock().unwrap_or_else(|poisoned| poisoned.into_inner());
614
615 let mut total_recent_accesses = 0;
617 let recent_window = Duration::from_secs(60);
618 let now = Instant::now();
619
620 for timestamps in patterns.values() {
621 total_recent_accesses +=
622 timestamps.iter().filter(|t| now.duration_since(**t) < recent_window).count();
623 }
624
625 drop(patterns);
626
627 let mut dynamic_max =
629 self.dynamic_max_size.lock().unwrap_or_else(|poisoned| poisoned.into_inner());
630
631 if total_recent_accesses > 1000 {
632 let new_size = (*dynamic_max as f64 * 1.15) as usize;
634 *dynamic_max = new_size.min(self.config.max_pool_size);
635 } else if total_recent_accesses < 100 {
636 let new_size = (*dynamic_max as f64 * 0.9) as usize;
638 *dynamic_max = new_size.max(self.config.min_pool_size);
639 }
640
641 Ok(())
642 }
643
644 fn defragment_pool(&self) -> Result<()> {
646 let mut pool = self.pool.write().unwrap_or_else(|poisoned| poisoned.into_inner());
648
649 for entries in pool.values_mut() {
650 entries.sort_by_key(|entry| std::cmp::Reverse(entry.access_count));
652 }
653
654 Ok(())
655 }
656
657 pub fn get_stats(&self) -> MemoryPoolStats {
659 let pool = self.pool.read().unwrap_or_else(|poisoned| poisoned.into_inner());
660 let current_size =
661 *self.current_size.lock().unwrap_or_else(|poisoned| poisoned.into_inner());
662 let stats = self.statistics.lock().unwrap_or_else(|poisoned| poisoned.into_inner());
663 let dynamic_max =
664 *self.dynamic_max_size.lock().unwrap_or_else(|poisoned| poisoned.into_inner());
665
666 let total_tensors = pool.values().map(|v| v.len()).sum();
667 let total_shapes = pool.len();
668
669 MemoryPoolStats {
670 total_tensors,
671 total_shapes,
672 current_size_bytes: current_size,
673 max_size_bytes: self.config.max_pool_size,
674 dynamic_max_size_bytes: dynamic_max,
675 utilization: current_size as f64 / dynamic_max as f64,
676 hit_rate: stats.hit_rate(),
677 miss_rate: stats.miss_rate(),
678 total_requests: stats.total_requests,
679 cache_hits: stats.cache_hits,
680 cache_misses: stats.cache_misses,
681 total_evictions: stats.total_evictions,
682 peak_memory_usage_bytes: stats.peak_memory_usage,
683 eviction_policy: self.config.eviction_policy,
684 adaptive_strategy: self.config.adaptive_strategy,
685 }
686 }
687
688 pub fn reset_statistics(&self) {
690 let mut stats = self.statistics.lock().unwrap_or_else(|poisoned| poisoned.into_inner());
691 *stats = PoolStatistics::default();
692 }
693
694 pub fn hit_rate(&self) -> f64 {
696 let stats = self.statistics.lock().unwrap_or_else(|poisoned| poisoned.into_inner());
697 stats.hit_rate()
698 }
699
700 pub fn eviction_policy(&self) -> MemoryEvictionPolicy {
702 self.config.eviction_policy
703 }
704
705 pub fn adaptive_strategy(&self) -> AdaptiveStrategy {
707 self.config.adaptive_strategy
708 }
709
710 pub fn get_predicted_shapes(&self, window: Duration) -> Vec<Vec<usize>> {
712 let patterns = self.access_patterns.lock().unwrap_or_else(|poisoned| poisoned.into_inner());
713 let now = Instant::now();
714
715 let mut frequent_shapes: Vec<(Vec<usize>, usize)> = patterns
716 .iter()
717 .map(|(shape, timestamps)| {
718 let count = timestamps.iter().filter(|t| now.duration_since(**t) < window).count();
719 (shape.clone(), count)
720 })
721 .filter(|(_, count)| *count > 0)
722 .collect();
723
724 frequent_shapes.sort_by_key(|item| std::cmp::Reverse(item.1));
725 frequent_shapes.into_iter().map(|(shape, _)| shape).collect()
726 }
727}
728
729#[derive(Debug, Clone)]
731pub struct MemoryPoolStats {
732 pub total_tensors: usize,
734 pub total_shapes: usize,
736 pub current_size_bytes: usize,
738 pub max_size_bytes: usize,
740 pub dynamic_max_size_bytes: usize,
742 pub utilization: f64,
744 pub hit_rate: f64,
746 pub miss_rate: f64,
748 pub total_requests: usize,
750 pub cache_hits: usize,
752 pub cache_misses: usize,
754 pub total_evictions: usize,
756 pub peak_memory_usage_bytes: usize,
758 pub eviction_policy: MemoryEvictionPolicy,
760 pub adaptive_strategy: AdaptiveStrategy,
762}
763
764pub struct MemoryMappedTensor {
766 file_path: String,
768 shape: Vec<usize>,
770 dtype: crate::tensor::DType,
772 _file: Option<File>,
774 file_size: u64,
776}
777
778impl MemoryMappedTensor {
779 pub fn new(file_path: String, shape: Vec<usize>, dtype: crate::tensor::DType) -> Result<Self> {
781 let mut file = File::open(&file_path).map_err(|e| {
783 TrustformersError::tensor_op_error(
784 &format!("Failed to open file for memory mapping: {}", e),
785 "mmap_new",
786 )
787 })?;
788
789 let file_size = file.seek(SeekFrom::End(0)).map_err(|e| {
791 TrustformersError::tensor_op_error(
792 &format!("Failed to get file size: {}", e),
793 "mmap_new",
794 )
795 })?;
796
797 let element_size = dtype.size_in_bytes();
799 let total_elements: usize = shape.iter().product();
800 let expected_size = total_elements * element_size;
801
802 if file_size != expected_size as u64 {
803 return Err(TrustformersError::tensor_op_error(
804 &format!(
805 "File size {} doesn't match expected tensor size {}",
806 file_size, expected_size
807 ),
808 "mmap_new",
809 ));
810 }
811
812 Ok(Self {
813 file_path,
814 shape,
815 dtype,
816 _file: Some(file),
817 file_size,
818 })
819 }
820
821 pub fn load(&self) -> Result<Tensor> {
823 let mut file = File::open(&self.file_path).map_err(|e| {
825 TrustformersError::tensor_op_error(
826 &format!("Failed to open file for reading: {}", e),
827 "mmap_load",
828 )
829 })?;
830
831 let mut buffer = vec![0u8; self.file_size as usize];
832 file.read_exact(&mut buffer).map_err(|e| {
833 TrustformersError::tensor_op_error(
834 &format!("Failed to read file data: {}", e),
835 "mmap_load",
836 )
837 })?;
838
839 match self.dtype {
841 crate::tensor::DType::F32 => {
842 let float_data = buffer
843 .chunks_exact(4)
844 .map(|chunk| f32::from_le_bytes([chunk[0], chunk[1], chunk[2], chunk[3]]))
845 .collect::<Vec<f32>>();
846 Tensor::from_slice(&float_data, &self.shape)
847 },
848 crate::tensor::DType::F64 => {
849 let float_data = buffer
850 .chunks_exact(8)
851 .map(|chunk| {
852 f64::from_le_bytes([
853 chunk[0], chunk[1], chunk[2], chunk[3], chunk[4], chunk[5], chunk[6],
854 chunk[7],
855 ])
856 })
857 .collect::<Vec<f64>>();
858 Tensor::from_slice_f64(&float_data, &self.shape)
859 },
860 crate::tensor::DType::I64 => {
861 let int_data = buffer
862 .chunks_exact(8)
863 .map(|chunk| {
864 i64::from_le_bytes([
865 chunk[0], chunk[1], chunk[2], chunk[3], chunk[4], chunk[5], chunk[6],
866 chunk[7],
867 ])
868 })
869 .collect::<Vec<i64>>();
870 Tensor::from_slice_i64(&int_data, &self.shape)
871 },
872 crate::tensor::DType::I32 => {
873 let int_data = buffer
874 .chunks_exact(4)
875 .map(|chunk| i32::from_le_bytes([chunk[0], chunk[1], chunk[2], chunk[3]]))
876 .collect::<Vec<i32>>();
877 Tensor::from_slice_i32(&int_data, &self.shape)
878 },
879 _ => Err(TrustformersError::tensor_op_error(
880 "Unsupported dtype for memory mapped tensor",
881 "mmap_load",
882 )),
883 }
884 }
885
886 pub fn shape(&self) -> &[usize] {
888 &self.shape
889 }
890
891 pub fn file_path(&self) -> &str {
893 &self.file_path
894 }
895}
896
897static MEMORY_MANAGER: std::sync::OnceLock<TensorMemoryPool> = std::sync::OnceLock::new();
899
900pub fn init_memory_manager(config: MemoryConfig) -> Result<()> {
902 let pool = TensorMemoryPool::new(config);
903 MEMORY_MANAGER.set(pool).map_err(|_| {
904 TrustformersError::invalid_input("Memory manager already initialized".to_string())
905 })?;
906 Ok(())
907}
908
909pub fn get_memory_manager() -> Option<&'static TensorMemoryPool> {
911 MEMORY_MANAGER.get()
912}
913
914pub fn get_tensor(shape: &[usize], dtype: crate::tensor::DType) -> Result<Tensor> {
916 if let Some(manager) = get_memory_manager() {
917 manager.get_tensor(shape, dtype)
918 } else {
919 match dtype {
921 crate::tensor::DType::F32 => Tensor::zeros(shape),
922 crate::tensor::DType::F64 => Tensor::zeros_f64(shape),
923 crate::tensor::DType::I64 => Tensor::zeros_i64(shape),
924 _ => Err(TrustformersError::tensor_op_error(
925 "Unsupported dtype",
926 "get_tensor",
927 )),
928 }
929 }
930}
931
932pub fn return_tensor(tensor: Tensor) -> Result<()> {
934 if let Some(manager) = get_memory_manager() {
935 manager.return_tensor(tensor)
936 } else {
937 Ok(()) }
939}
940
941#[cfg(test)]
942mod tests {
943 use super::*;
944
945 #[test]
946 fn test_memory_config_default() {
947 let config = MemoryConfig::default();
948 assert!(config.enable_memory_pool);
949 assert!(config.enable_zero_copy);
950 assert!(config.enable_mmap);
951 assert_eq!(config.max_pool_size, 1024 * 1024 * 1024);
952 }
953
954 #[test]
955 fn test_tensor_pool_creation() {
956 let config = MemoryConfig::default();
957 let pool = TensorMemoryPool::new(config);
958 let stats = pool.get_stats();
959 assert_eq!(stats.total_tensors, 0);
960 assert_eq!(stats.current_size_bytes, 0);
961 }
962
963 #[test]
964 fn test_tensor_pool_get_and_return() -> Result<()> {
965 let config = MemoryConfig::default();
966 let pool = TensorMemoryPool::new(config);
967
968 let shape = vec![2, 3];
970 let tensor = pool.get_tensor(&shape, crate::tensor::DType::F32)?;
971 assert_eq!(tensor.shape(), shape.as_slice());
972
973 pool.return_tensor(tensor)?;
975
976 let tensor2 = pool.get_tensor(&shape, crate::tensor::DType::F32)?;
978 assert_eq!(tensor2.shape(), shape.as_slice());
979
980 Ok(())
981 }
982
983 #[test]
984 fn test_zero_copy_tensor_view() -> Result<()> {
985 let tensor = Arc::new(Tensor::ones(&[10])?);
986 let view = TensorView::slice(tensor, 2, 8)?;
987 assert_eq!(view.shape(), &[6]);
988
989 let viewed_tensor = view.as_tensor()?;
990 assert_eq!(viewed_tensor.shape(), &[6]);
991
992 Ok(())
993 }
994
995 #[test]
996 fn test_memory_mapped_tensor() -> Result<()> {
997 use std::fs::File;
998 use std::io::Write;
999
1000 let temp_file = "test_temp.bin";
1002 let data_size = 100 * 100 * std::mem::size_of::<f32>();
1003 let data: Vec<u8> = vec![0; data_size];
1004
1005 {
1006 let mut file = File::create(temp_file).map_err(|e| {
1007 TrustformersError::tensor_op_error(
1008 &format!("Failed to create test file: {}", e),
1009 "test_setup",
1010 )
1011 })?;
1012 file.write_all(&data).map_err(|e| {
1013 TrustformersError::tensor_op_error(
1014 &format!("Failed to write test data: {}", e),
1015 "test_setup",
1016 )
1017 })?;
1018 }
1019
1020 let mmap_tensor = MemoryMappedTensor::new(
1021 temp_file.to_string(),
1022 vec![100, 100],
1023 crate::tensor::DType::F32,
1024 )?;
1025
1026 assert_eq!(mmap_tensor.shape(), &[100, 100]);
1027 assert_eq!(mmap_tensor.file_path(), temp_file);
1028
1029 let loaded = mmap_tensor.load()?;
1030 assert_eq!(loaded.shape(), &[100, 100]);
1031
1032 std::fs::remove_file(temp_file).ok();
1034
1035 Ok(())
1036 }
1037
1038 #[test]
1039 fn test_global_memory_manager() -> Result<()> {
1040 let config = MemoryConfig::default();
1041 init_memory_manager(config)?;
1042
1043 let tensor = get_tensor(&[5, 5], crate::tensor::DType::F32)?;
1044 assert_eq!(tensor.shape(), [5, 5].as_slice());
1045
1046 return_tensor(tensor)?;
1047
1048 Ok(())
1049 }
1050
1051 #[test]
1054 fn test_memory_config_custom_values() {
1055 let config = MemoryConfig {
1056 enable_memory_pool: false,
1057 max_pool_size: 512 * 1024 * 1024,
1058 min_pool_size: 32 * 1024 * 1024,
1059 enable_zero_copy: false,
1060 enable_mmap: false,
1061 mmap_threshold: 50 * 1024 * 1024,
1062 cleanup_interval: Duration::from_secs(30),
1063 eviction_policy: MemoryEvictionPolicy::LRU,
1064 adaptive_strategy: AdaptiveStrategy::Fixed,
1065 target_hit_rate: 0.9,
1066 enable_prefetching: false,
1067 enable_defragmentation: false,
1068 };
1069 assert!(!config.enable_memory_pool);
1070 assert_eq!(config.max_pool_size, 512 * 1024 * 1024);
1071 assert_eq!(config.eviction_policy, MemoryEvictionPolicy::LRU);
1072 assert_eq!(config.adaptive_strategy, AdaptiveStrategy::Fixed);
1073 }
1074
1075 #[test]
1076 fn test_memory_eviction_policy_lru() {
1077 let config = MemoryConfig {
1078 eviction_policy: MemoryEvictionPolicy::LRU,
1079 ..Default::default()
1080 };
1081 let pool = TensorMemoryPool::new(config);
1082 assert_eq!(pool.eviction_policy(), MemoryEvictionPolicy::LRU);
1083 }
1084
1085 #[test]
1086 fn test_memory_eviction_policy_lfu() {
1087 let config = MemoryConfig {
1088 eviction_policy: MemoryEvictionPolicy::LFU,
1089 ..Default::default()
1090 };
1091 let pool = TensorMemoryPool::new(config);
1092 assert_eq!(pool.eviction_policy(), MemoryEvictionPolicy::LFU);
1093 }
1094
1095 #[test]
1096 fn test_memory_eviction_policy_size_based() {
1097 let config = MemoryConfig {
1098 eviction_policy: MemoryEvictionPolicy::SizeBased,
1099 ..Default::default()
1100 };
1101 let pool = TensorMemoryPool::new(config);
1102 assert_eq!(pool.eviction_policy(), MemoryEvictionPolicy::SizeBased);
1103 }
1104
1105 #[test]
1106 fn test_memory_eviction_policy_arc() {
1107 let config = MemoryConfig {
1108 eviction_policy: MemoryEvictionPolicy::ARC,
1109 ..Default::default()
1110 };
1111 let pool = TensorMemoryPool::new(config);
1112 assert_eq!(pool.eviction_policy(), MemoryEvictionPolicy::ARC);
1113 }
1114
1115 #[test]
1116 fn test_adaptive_strategy_fixed() {
1117 let config = MemoryConfig {
1118 adaptive_strategy: AdaptiveStrategy::Fixed,
1119 ..Default::default()
1120 };
1121 let pool = TensorMemoryPool::new(config);
1122 assert_eq!(pool.adaptive_strategy(), AdaptiveStrategy::Fixed);
1123 }
1124
1125 #[test]
1126 fn test_adaptive_strategy_memory_pressure() {
1127 let config = MemoryConfig {
1128 adaptive_strategy: AdaptiveStrategy::MemoryPressure,
1129 ..Default::default()
1130 };
1131 let pool = TensorMemoryPool::new(config);
1132 assert_eq!(pool.adaptive_strategy(), AdaptiveStrategy::MemoryPressure);
1133 }
1134
1135 #[test]
1136 fn test_adaptive_strategy_predictive() {
1137 let config = MemoryConfig {
1138 adaptive_strategy: AdaptiveStrategy::Predictive,
1139 ..Default::default()
1140 };
1141 let pool = TensorMemoryPool::new(config);
1142 assert_eq!(pool.adaptive_strategy(), AdaptiveStrategy::Predictive);
1143 }
1144
1145 #[test]
1146 fn test_pool_stats_initial_zero() {
1147 let pool = TensorMemoryPool::new(MemoryConfig::default());
1148 let stats = pool.get_stats();
1149 assert_eq!(stats.total_tensors, 0);
1150 assert_eq!(stats.current_size_bytes, 0);
1151 assert_eq!(stats.cache_hits, 0);
1152 assert_eq!(stats.cache_misses, 0);
1153 }
1154
1155 #[test]
1156 fn test_pool_initial_hit_rate() {
1157 let pool = TensorMemoryPool::new(MemoryConfig::default());
1158 let hr = pool.hit_rate();
1160 assert!(
1161 hr == 0.0 || hr.is_nan(),
1162 "initial hit rate should be 0.0 or NaN, got {hr}"
1163 );
1164 }
1165
1166 #[test]
1167 fn test_pool_multiple_shapes() -> Result<()> {
1168 let pool = TensorMemoryPool::new(MemoryConfig::default());
1169 let t1 = pool.get_tensor(&[2, 3], crate::tensor::DType::F32)?;
1170 let t2 = pool.get_tensor(&[4, 5], crate::tensor::DType::F32)?;
1171 assert_eq!(t1.shape(), &[2, 3]);
1172 assert_eq!(t2.shape(), &[4, 5]);
1173 Ok(())
1174 }
1175
1176 #[test]
1177 fn test_pool_f64_dtype() -> Result<()> {
1178 let pool = TensorMemoryPool::new(MemoryConfig::default());
1179 let t = pool.get_tensor(&[3, 3], crate::tensor::DType::F64)?;
1180 assert_eq!(t.shape(), &[3, 3]);
1181 Ok(())
1182 }
1183
1184 #[test]
1185 fn test_pool_i64_dtype() -> Result<()> {
1186 let pool = TensorMemoryPool::new(MemoryConfig::default());
1187 let t = pool.get_tensor(&[5], crate::tensor::DType::I64)?;
1188 assert_eq!(t.shape(), &[5]);
1189 Ok(())
1190 }
1191
1192 #[test]
1193 fn test_pool_reset_statistics() -> Result<()> {
1194 let pool = TensorMemoryPool::new(MemoryConfig::default());
1195 let t1 = pool.get_tensor(&[2, 2], crate::tensor::DType::F32)?;
1197 pool.return_tensor(t1)?;
1198 let _t2 = pool.get_tensor(&[2, 2], crate::tensor::DType::F32)?;
1199 pool.reset_statistics();
1201 let stats = pool.get_stats();
1202 assert_eq!(stats.cache_hits, 0);
1203 assert_eq!(stats.cache_misses, 0);
1204 Ok(())
1205 }
1206
1207 #[test]
1208 fn test_tensor_view_slice_middle() -> Result<()> {
1209 let tensor = Arc::new(Tensor::ones(&[10])?);
1210 let view = TensorView::slice(tensor, 3, 7)?;
1211 assert_eq!(view.shape(), &[4]);
1212 Ok(())
1213 }
1214
1215 #[test]
1216 fn test_tensor_view_as_tensor_values() -> Result<()> {
1217 let tensor = Arc::new(Tensor::ones(&[10])?);
1218 let view = TensorView::slice(tensor, 0, 5)?;
1219 let viewed = view.as_tensor()?;
1220 assert_eq!(viewed.shape(), &[5]);
1221 if let Tensor::F32(arr) = &viewed {
1223 for v in arr.iter() {
1224 assert!((*v - 1.0_f32).abs() < 1e-6, "expected 1.0, got {v}");
1225 }
1226 }
1227 Ok(())
1228 }
1229
1230 #[test]
1231 fn test_tensor_view_full_range() -> Result<()> {
1232 let tensor = Arc::new(Tensor::ones(&[8])?);
1233 let view = TensorView::slice(tensor, 0, 8)?;
1234 assert_eq!(view.shape(), &[8]);
1235 Ok(())
1236 }
1237
1238 #[test]
1239 fn test_mmap_shape_stored() -> Result<()> {
1240 use std::io::Write;
1241 let tmp_dir = std::env::temp_dir();
1242 let path = tmp_dir.join("trustformers_mmap_shape_test.bin");
1243 let path_str = path.to_string_lossy().to_string();
1244 let data = vec![0u8; 10 * 20 * std::mem::size_of::<f32>()];
1246 {
1247 let mut f = std::fs::File::create(&path).map_err(|e| {
1248 TrustformersError::tensor_op_error(&e.to_string(), "test_mmap_shape_stored")
1249 })?;
1250 f.write_all(&data).map_err(|e| {
1251 TrustformersError::tensor_op_error(&e.to_string(), "test_mmap_shape_stored")
1252 })?;
1253 }
1254 let mmap =
1255 MemoryMappedTensor::new(path_str.clone(), vec![10, 20], crate::tensor::DType::F32)?;
1256 assert_eq!(mmap.shape(), &[10, 20]);
1257 std::fs::remove_file(&path).ok();
1258 Ok(())
1259 }
1260
1261 #[test]
1262 fn test_mmap_file_path_stored() -> Result<()> {
1263 use std::io::Write;
1264 let tmp_dir = std::env::temp_dir();
1265 let path = tmp_dir.join("trustformers_mmap_path_test.bin");
1266 let path_str = path.to_string_lossy().to_string();
1267 let data = vec![0u8; 4 * std::mem::size_of::<f32>()];
1268 {
1269 let mut f = std::fs::File::create(&path).map_err(|e| {
1270 TrustformersError::tensor_op_error(&e.to_string(), "test_mmap_file_path_stored")
1271 })?;
1272 f.write_all(&data).map_err(|e| {
1273 TrustformersError::tensor_op_error(&e.to_string(), "test_mmap_file_path_stored")
1274 })?;
1275 }
1276 let mmap = MemoryMappedTensor::new(path_str.clone(), vec![4], crate::tensor::DType::F32)?;
1277 assert_eq!(mmap.file_path(), path_str);
1278 std::fs::remove_file(&path).ok();
1279 Ok(())
1280 }
1281
1282 #[test]
1283 fn test_global_get_tensor_without_explicit_init() -> Result<()> {
1284 let tensor = get_tensor(&[3, 3], crate::tensor::DType::F32)?;
1288 assert_eq!(tensor.shape(), &[3, 3]);
1289 Ok(())
1290 }
1291}