1use crate::OptimizerError;
2use std::collections::HashMap;
3use std::fs::{File, OpenOptions};
4use std::io::{self, Read, Seek, SeekFrom, Write};
5use std::path::{Path, PathBuf};
6use std::sync::{Arc, Mutex, RwLock};
7
8#[derive(Debug, Clone)]
10pub struct MemoryMappedConfig {
11 pub base_dir: PathBuf,
13 pub initial_size: usize,
15 pub growth_factor: f32,
17 pub auto_sync: bool,
19 pub sync_frequency: usize,
21 pub use_locking: bool,
23 pub prefault: bool,
25}
26
27impl Default for MemoryMappedConfig {
28 fn default() -> Self {
29 Self {
30 base_dir: PathBuf::from("mmap_states"),
31 initial_size: 1024 * 1024, growth_factor: 2.0,
33 auto_sync: true,
34 sync_frequency: 100,
35 use_locking: true,
36 prefault: false,
37 }
38 }
39}
40
41pub struct MemoryMappedFile {
43 file: File,
44 path: PathBuf,
45 size: usize,
46 capacity: usize,
47 sync_counter: usize,
48 config: MemoryMappedConfig,
49}
50
51impl MemoryMappedFile {
52 pub fn new(path: PathBuf, config: MemoryMappedConfig) -> Result<Self, OptimizerError> {
54 if let Some(parent) = path.parent() {
56 std::fs::create_dir_all(parent).map_err(|e| {
57 OptimizerError::MemoryMapError(format!("Failed to create directory: {}", e))
58 })?;
59 }
60
61 let file = OpenOptions::new()
62 .read(true)
63 .write(true)
64 .create(true)
65 .open(&path)
66 .map_err(|e| OptimizerError::MemoryMapError(format!("Failed to open file: {}", e)))?;
67
68 let initial_size = config.initial_size;
70 file.set_len(initial_size as u64).map_err(|e| {
71 OptimizerError::MemoryMapError(format!("Failed to set file size: {}", e))
72 })?;
73
74 Ok(Self {
75 file,
76 path,
77 size: 0,
78 capacity: initial_size,
79 sync_counter: 0,
80 config,
81 })
82 }
83
84 pub fn open(path: PathBuf, config: MemoryMappedConfig) -> Result<Self, OptimizerError> {
86 let file = OpenOptions::new()
87 .read(true)
88 .write(true)
89 .open(&path)
90 .map_err(|e| OptimizerError::MemoryMapError(format!("Failed to open file: {}", e)))?;
91
92 let metadata = file.metadata().map_err(|e| {
93 OptimizerError::MemoryMapError(format!("Failed to get file metadata: {}", e))
94 })?;
95
96 let capacity = metadata.len() as usize;
97
98 Ok(Self {
99 file,
100 path,
101 size: capacity, capacity,
103 sync_counter: 0,
104 config,
105 })
106 }
107
108 pub fn write(&mut self, data: &[u8]) -> Result<usize, OptimizerError> {
110 if self.size + data.len() > self.capacity {
112 self.expand(data.len())?;
113 }
114
115 self.file
117 .seek(SeekFrom::Start(self.size as u64))
118 .map_err(|e| OptimizerError::MemoryMapError(format!("Failed to seek: {}", e)))?;
119
120 let bytes_written = self
122 .file
123 .write(data)
124 .map_err(|e| OptimizerError::MemoryMapError(format!("Failed to write: {}", e)))?;
125
126 self.size += bytes_written;
127 self.sync_counter += 1;
128
129 if self.config.auto_sync && self.sync_counter >= self.config.sync_frequency {
131 self.sync()?;
132 }
133
134 Ok(bytes_written)
135 }
136
137 pub fn read(&mut self, offset: usize, length: usize) -> Result<Vec<u8>, OptimizerError> {
139 if offset + length > self.size {
140 return Err(OptimizerError::MemoryMapError(
141 "Read extends beyond file size".to_string(),
142 ));
143 }
144
145 self.file
146 .seek(SeekFrom::Start(offset as u64))
147 .map_err(|e| OptimizerError::MemoryMapError(format!("Failed to seek: {}", e)))?;
148
149 let mut buffer = vec![0u8; length];
150 self.file
151 .read_exact(&mut buffer)
152 .map_err(|e| OptimizerError::MemoryMapError(format!("Failed to read: {}", e)))?;
153
154 Ok(buffer)
155 }
156
157 pub fn read_all(&mut self) -> Result<Vec<u8>, OptimizerError> {
159 self.read(0, self.size)
160 }
161
162 pub fn write_at(&mut self, offset: usize, data: &[u8]) -> Result<(), OptimizerError> {
164 if offset + data.len() > self.capacity {
165 self.expand(offset + data.len() - self.capacity)?;
166 }
167
168 self.file
169 .seek(SeekFrom::Start(offset as u64))
170 .map_err(|e| OptimizerError::MemoryMapError(format!("Failed to seek: {}", e)))?;
171
172 self.file
173 .write_all(data)
174 .map_err(|e| OptimizerError::MemoryMapError(format!("Failed to write: {}", e)))?;
175
176 if offset + data.len() > self.size {
178 self.size = offset + data.len();
179 }
180
181 self.sync_counter += 1;
182
183 if self.config.auto_sync && self.sync_counter >= self.config.sync_frequency {
184 self.sync()?;
185 }
186
187 Ok(())
188 }
189
190 pub fn sync(&mut self) -> Result<(), OptimizerError> {
192 self.file
193 .sync_all()
194 .map_err(|e| OptimizerError::MemoryMapError(format!("Failed to sync: {}", e)))?;
195 self.sync_counter = 0;
196 Ok(())
197 }
198
199 pub fn truncate(&mut self, size: usize) -> Result<(), OptimizerError> {
201 self.file
202 .set_len(size as u64)
203 .map_err(|e| OptimizerError::MemoryMapError(format!("Failed to truncate: {}", e)))?;
204
205 self.size = size.min(self.size);
206 self.capacity = size;
207
208 Ok(())
209 }
210
211 pub fn size(&self) -> usize {
213 self.size
214 }
215
216 pub fn capacity(&self) -> usize {
218 self.capacity
219 }
220
221 pub fn path(&self) -> &Path {
223 &self.path
224 }
225
226 fn expand(&mut self, min_additional: usize) -> Result<(), OptimizerError> {
229 let new_capacity =
230 ((self.capacity + min_additional) as f32 * self.config.growth_factor) as usize;
231
232 self.file
233 .set_len(new_capacity as u64)
234 .map_err(|e| OptimizerError::MemoryMapError(format!("Failed to expand file: {}", e)))?;
235
236 self.capacity = new_capacity;
237
238 Ok(())
239 }
240}
241
242pub struct MemoryMappedStateStorage {
244 config: MemoryMappedConfig,
245 files: HashMap<String, Arc<Mutex<MemoryMappedFile>>>,
246 metadata: Arc<RwLock<StateMetadata>>,
247}
248
249#[derive(Debug, Clone, Default)]
250struct StateMetadata {
251 entries: HashMap<String, StateEntry>,
252}
253
254#[derive(Debug, Clone)]
255struct StateEntry {
256 file_key: String,
257 offset: usize,
258 length: usize,
259 data_type: String,
260 timestamp: u64,
261}
262
263impl MemoryMappedStateStorage {
264 pub fn new(config: MemoryMappedConfig) -> Result<Self, OptimizerError> {
266 std::fs::create_dir_all(&config.base_dir).map_err(|e| {
268 OptimizerError::MemoryMapError(format!("Failed to create base directory: {}", e))
269 })?;
270
271 Ok(Self {
272 config,
273 files: HashMap::new(),
274 metadata: Arc::new(RwLock::new(StateMetadata::default())),
275 })
276 }
277
278 pub fn store<T: serde::Serialize>(
280 &mut self,
281 key: &str,
282 data: &T,
283 ) -> Result<(), OptimizerError> {
284 let serialized =
285 oxicode::serde::encode_to_vec(data, oxicode::config::standard()).map_err(|e| {
286 OptimizerError::MemoryMapError(format!("Failed to serialize data: {}", e))
287 })?;
288
289 self.store_raw(key, &serialized, std::any::type_name::<T>())
290 }
291
292 pub fn store_raw(
294 &mut self,
295 key: &str,
296 data: &[u8],
297 data_type: &str,
298 ) -> Result<(), OptimizerError> {
299 let file_key = format!("state_{}", key.replace('/', "_"));
300 let file_path = self.config.base_dir.join(format!("{}.mmap", file_key));
301
302 let file_mutex = if let Some(existing) = self.files.get(&file_key) {
304 existing.clone()
305 } else {
306 let mmap_file = MemoryMappedFile::new(file_path, self.config.clone())?;
307 let file_mutex = Arc::new(Mutex::new(mmap_file));
308 self.files.insert(file_key.clone(), file_mutex.clone());
309 file_mutex
310 };
311
312 let offset = {
314 let mut file = file_mutex.lock().expect("lock should not be poisoned");
315 let offset = file.size();
316 file.write(data)?;
317 offset
318 };
319
320 let timestamp = std::time::SystemTime::now()
322 .duration_since(std::time::UNIX_EPOCH)
323 .expect("system time should be after UNIX epoch")
324 .as_secs();
325
326 let entry = StateEntry {
327 file_key: file_key.clone(),
328 offset,
329 length: data.len(),
330 data_type: data_type.to_string(),
331 timestamp,
332 };
333
334 let mut metadata = self.metadata.write().expect("lock should not be poisoned");
335 metadata.entries.insert(key.to_string(), entry);
336
337 Ok(())
338 }
339
340 pub fn load<T: serde::de::DeserializeOwned>(
342 &mut self,
343 key: &str,
344 ) -> Result<Option<T>, OptimizerError> {
345 if let Some(data) = self.load_raw(key)? {
346 let (deserialized, _): (T, usize) =
347 oxicode::serde::decode_from_slice(&data, oxicode::config::standard()).map_err(
348 |e| {
349 OptimizerError::MemoryMapError(format!("Failed to deserialize data: {}", e))
350 },
351 )?;
352 Ok(Some(deserialized))
353 } else {
354 Ok(None)
355 }
356 }
357
358 pub fn load_raw(&mut self, key: &str) -> Result<Option<Vec<u8>>, OptimizerError> {
360 let metadata = self.metadata.read().expect("lock should not be poisoned");
361 let entry = if let Some(entry) = metadata.entries.get(key) {
362 entry.clone()
363 } else {
364 return Ok(None);
365 };
366 drop(metadata);
367
368 let file_mutex = self
370 .files
371 .get(&entry.file_key)
372 .ok_or_else(|| OptimizerError::MemoryMapError("File not found".to_string()))?
373 .clone();
374
375 let mut file = file_mutex.lock().expect("lock should not be poisoned");
377 let data = file.read(entry.offset, entry.length)?;
378
379 Ok(Some(data))
380 }
381
382 pub fn update<T: serde::Serialize>(
384 &mut self,
385 key: &str,
386 data: &T,
387 ) -> Result<(), OptimizerError> {
388 let serialized =
389 oxicode::serde::encode_to_vec(data, oxicode::config::standard()).map_err(|e| {
390 OptimizerError::MemoryMapError(format!("Failed to serialize data: {}", e))
391 })?;
392
393 self.update_raw(key, &serialized)
394 }
395
396 pub fn update_raw(&mut self, key: &str, data: &[u8]) -> Result<(), OptimizerError> {
398 let metadata = self.metadata.read().expect("lock should not be poisoned");
399 let mut entry = if let Some(entry) = metadata.entries.get(key) {
400 entry.clone()
401 } else {
402 drop(metadata);
403 return self.store_raw(key, data, "unknown");
404 };
405 drop(metadata);
406
407 let file_mutex = self
408 .files
409 .get(&entry.file_key)
410 .ok_or_else(|| OptimizerError::MemoryMapError("File not found".to_string()))?
411 .clone();
412
413 if data.len() <= entry.length {
415 let mut file = file_mutex.lock().expect("lock should not be poisoned");
416 file.write_at(entry.offset, data)?;
417 } else {
418 let mut file = file_mutex.lock().expect("lock should not be poisoned");
420 entry.offset = file.size();
421 entry.length = data.len();
422 file.write(data)?;
423 }
424
425 entry.timestamp = std::time::SystemTime::now()
427 .duration_since(std::time::UNIX_EPOCH)
428 .expect("system time should be after UNIX epoch")
429 .as_secs();
430
431 let mut metadata = self.metadata.write().expect("lock should not be poisoned");
432 metadata.entries.insert(key.to_string(), entry);
433
434 Ok(())
435 }
436
437 pub fn remove(&mut self, key: &str) -> Result<bool, OptimizerError> {
439 let mut metadata = self.metadata.write().expect("lock should not be poisoned");
440 Ok(metadata.entries.remove(key).is_some())
441 }
442
443 pub fn keys(&self) -> Vec<String> {
445 let metadata = self.metadata.read().expect("lock should not be poisoned");
446 metadata.entries.keys().cloned().collect()
447 }
448
449 pub fn contains_key(&self, key: &str) -> bool {
451 let metadata = self.metadata.read().expect("lock should not be poisoned");
452 metadata.entries.contains_key(key)
453 }
454
455 pub fn statistics(&self) -> StorageStatistics {
457 let metadata = self.metadata.read().expect("lock should not be poisoned");
458 let total_entries = metadata.entries.len();
459
460 let total_size: usize = self
461 .files
462 .values()
463 .map(|file_mutex| {
464 let file = file_mutex.lock().expect("lock should not be poisoned");
465 file.size()
466 })
467 .sum();
468
469 let total_capacity: usize = self
470 .files
471 .values()
472 .map(|file_mutex| {
473 let file = file_mutex.lock().expect("lock should not be poisoned");
474 file.capacity()
475 })
476 .sum();
477
478 StorageStatistics {
479 total_entries,
480 total_size,
481 total_capacity,
482 utilization: if total_capacity > 0 {
483 total_size as f32 / total_capacity as f32
484 } else {
485 0.0
486 },
487 num_files: self.files.len(),
488 }
489 }
490
491 pub fn sync_all(&mut self) -> Result<(), OptimizerError> {
493 for file_mutex in self.files.values() {
494 let mut file = file_mutex.lock().expect("lock should not be poisoned");
495 file.sync()?;
496 }
497 Ok(())
498 }
499
500 pub fn compact(&mut self) -> Result<(), OptimizerError> {
502 self.sync_all()
505 }
506}
507
508#[derive(Debug, Clone)]
510pub struct StorageStatistics {
511 pub total_entries: usize,
512 pub total_size: usize,
513 pub total_capacity: usize,
514 pub utilization: f32,
515 pub num_files: usize,
516}
517
518impl std::fmt::Display for StorageStatistics {
519 fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
520 writeln!(f, "Memory-Mapped Storage Statistics:")?;
521 writeln!(f, " Total Entries: {}", self.total_entries)?;
522 writeln!(
523 f,
524 " Total Size: {:.2} MB",
525 self.total_size as f64 / 1024.0 / 1024.0
526 )?;
527 writeln!(
528 f,
529 " Total Capacity: {:.2} MB",
530 self.total_capacity as f64 / 1024.0 / 1024.0
531 )?;
532 writeln!(f, " Utilization: {:.1}%", self.utilization * 100.0)?;
533 writeln!(f, " Number of Files: {}", self.num_files)?;
534 Ok(())
535 }
536}
537
538pub trait MemoryMappedSupport {
540 fn save_to_mmap(
542 &self,
543 storage: &mut MemoryMappedStateStorage,
544 prefix: &str,
545 ) -> Result<(), OptimizerError>;
546
547 fn load_from_mmap(
549 &mut self,
550 storage: &mut MemoryMappedStateStorage,
551 prefix: &str,
552 ) -> Result<(), OptimizerError>;
553
554 fn mmap_state_keys(&self, prefix: &str) -> Vec<String>;
556}
557
558pub struct MemoryMappedOptimizer<T> {
560 inner: T,
561 storage: MemoryMappedStateStorage,
562 prefix: String,
563 auto_save: bool,
564 save_frequency: usize,
565 step_count: usize,
566}
567
568impl<T> MemoryMappedOptimizer<T>
569where
570 T: MemoryMappedSupport,
571{
572 pub fn new(
574 inner: T,
575 config: MemoryMappedConfig,
576 prefix: String,
577 ) -> Result<Self, OptimizerError> {
578 let storage = MemoryMappedStateStorage::new(config)?;
579
580 Ok(Self {
581 inner,
582 storage,
583 prefix,
584 auto_save: true,
585 save_frequency: 100,
586 step_count: 0,
587 })
588 }
589
590 pub fn set_auto_save(&mut self, enabled: bool, frequency: usize) {
592 self.auto_save = enabled;
593 self.save_frequency = frequency;
594 }
595
596 pub fn inner(&self) -> &T {
598 &self.inner
599 }
600
601 pub fn inner_mut(&mut self) -> &mut T {
603 &mut self.inner
604 }
605
606 pub fn storage(&self) -> &MemoryMappedStateStorage {
608 &self.storage
609 }
610
611 pub fn storage_mut(&mut self) -> &mut MemoryMappedStateStorage {
613 &mut self.storage
614 }
615
616 pub fn save_state(&mut self) -> Result<(), OptimizerError> {
618 self.inner.save_to_mmap(&mut self.storage, &self.prefix)
619 }
620
621 pub fn load_state(&mut self) -> Result<(), OptimizerError> {
623 self.inner.load_from_mmap(&mut self.storage, &self.prefix)
624 }
625
626 pub fn step_with_mmap(&mut self) -> Result<(), OptimizerError> {
628 self.step_count += 1;
629
630 if self.auto_save && (self.step_count % self.save_frequency == 0) {
631 self.save_state()?;
632 }
633
634 Ok(())
635 }
636}
637
638#[cfg(test)]
639mod tests {
640 use super::*;
641 use std::collections::HashMap;
642
643 #[test]
644 fn test_memory_mapped_file() {
645 let temp_dir = tempfile::tempdir().unwrap();
646 let config = MemoryMappedConfig::default();
647 let path = temp_dir.path().join("test.mmap");
648
649 let mut file = MemoryMappedFile::new(path.clone(), config).unwrap();
650
651 let data = b"Hello, world!";
653 let bytes_written = file.write(data).unwrap();
654 assert_eq!(bytes_written, data.len());
655
656 let read_data = file.read(0, data.len()).unwrap();
658 assert_eq!(read_data, data);
659
660 let new_data = b"Hi";
662 file.write_at(0, new_data).unwrap();
663 let read_data = file.read(0, new_data.len()).unwrap();
664 assert_eq!(read_data, new_data);
665 }
666
667 #[test]
668 fn test_memory_mapped_storage() -> Result<(), OptimizerError> {
669 let temp_dir = tempfile::tempdir().unwrap();
670 let config = MemoryMappedConfig {
671 base_dir: temp_dir.path().to_path_buf(),
672 ..Default::default()
673 };
674
675 let mut storage = MemoryMappedStateStorage::new(config).unwrap();
676
677 let test_data = HashMap::from([
679 ("lr".to_string(), 0.01f32),
680 ("momentum".to_string(), 0.9f32),
681 ]);
682
683 storage.store("optimizer_params", &test_data).unwrap();
684 let loaded_data: HashMap<String, f32> = storage.load("optimizer_params").unwrap().unwrap();
685
686 assert_eq!(test_data, loaded_data);
687
688 let updated_data = HashMap::from([
690 ("lr".to_string(), 0.001f32),
691 ("momentum".to_string(), 0.95f32),
692 ]);
693
694 storage.update("optimizer_params", &updated_data).unwrap();
695 let loaded_updated: HashMap<String, f32> =
696 storage.load("optimizer_params").unwrap().unwrap();
697
698 assert_eq!(updated_data, loaded_updated);
699
700 let stats = storage.statistics();
702 assert_eq!(stats.total_entries, 1);
703 assert!(stats.total_size > 0);
704 Ok(())
705 }
706
707 #[derive(Debug)]
708 struct MockOptimizer {
709 lr: f32,
710 momentum: f32,
711 }
712
713 impl MemoryMappedSupport for MockOptimizer {
714 fn save_to_mmap(
715 &self,
716 storage: &mut MemoryMappedStateStorage,
717 prefix: &str,
718 ) -> Result<(), OptimizerError> {
719 storage.store(&format!("{}_lr", prefix), &self.lr)?;
720 storage.store(&format!("{}_momentum", prefix), &self.momentum)?;
721 Ok(())
722 }
723
724 fn load_from_mmap(
725 &mut self,
726 storage: &mut MemoryMappedStateStorage,
727 prefix: &str,
728 ) -> Result<(), OptimizerError> {
729 if let Some(lr) = storage.load(&format!("{}_lr", prefix))? {
730 self.lr = lr;
731 }
732 if let Some(momentum) = storage.load(&format!("{}_momentum", prefix))? {
733 self.momentum = momentum;
734 }
735 Ok(())
736 }
737
738 fn mmap_state_keys(&self, prefix: &str) -> Vec<String> {
739 vec![format!("{}_lr", prefix), format!("{}_momentum", prefix)]
740 }
741 }
742
743 #[test]
744 fn test_memory_mapped_optimizer() {
745 let temp_dir = tempfile::tempdir().unwrap();
746 let config = MemoryMappedConfig {
747 base_dir: temp_dir.path().to_path_buf(),
748 ..Default::default()
749 };
750
751 let optimizer = MockOptimizer {
752 lr: 0.01,
753 momentum: 0.9,
754 };
755 let mut mmap_optimizer =
756 MemoryMappedOptimizer::new(optimizer, config, "test_opt".to_string()).unwrap();
757
758 mmap_optimizer.save_state().unwrap();
760
761 mmap_optimizer.inner_mut().lr = 0.001;
763 mmap_optimizer.inner_mut().momentum = 0.95;
764
765 mmap_optimizer.load_state().unwrap();
767
768 assert_eq!(mmap_optimizer.inner().lr, 0.01);
769 assert_eq!(mmap_optimizer.inner().momentum, 0.9);
770 }
771}