Skip to main content

torsh_optim/
memory_mapped.rs

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/// Configuration for memory-mapped optimizer states
9#[derive(Debug, Clone)]
10pub struct MemoryMappedConfig {
11    /// Base directory for memory-mapped files
12    pub base_dir: PathBuf,
13    /// Initial size for memory-mapped files (in bytes)
14    pub initial_size: usize,
15    /// Growth factor when expanding files
16    pub growth_factor: f32,
17    /// Whether to sync to disk automatically
18    pub auto_sync: bool,
19    /// Sync frequency (every N operations)
20    pub sync_frequency: usize,
21    /// Whether to use advisory locking
22    pub use_locking: bool,
23    /// Whether to prefault pages
24    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, // 1MB
32            growth_factor: 2.0,
33            auto_sync: true,
34            sync_frequency: 100,
35            use_locking: true,
36            prefault: false,
37        }
38    }
39}
40
41/// Memory-mapped file wrapper
42pub 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    /// Create a new memory-mapped file
53    pub fn new(path: PathBuf, config: MemoryMappedConfig) -> Result<Self, OptimizerError> {
54        // Create parent directory if it doesn't exist
55        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        // Set initial file size
69        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    /// Open an existing memory-mapped file
85    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, // Assume file is fully used when opening
102            capacity,
103            sync_counter: 0,
104            config,
105        })
106    }
107
108    /// Write data to the memory-mapped file
109    pub fn write(&mut self, data: &[u8]) -> Result<usize, OptimizerError> {
110        // Check if we need to expand the file
111        if self.size + data.len() > self.capacity {
112            self.expand(data.len())?;
113        }
114
115        // Seek to the end of the used data
116        self.file
117            .seek(SeekFrom::Start(self.size as u64))
118            .map_err(|e| OptimizerError::MemoryMapError(format!("Failed to seek: {}", e)))?;
119
120        // Write the data
121        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        // Auto-sync if configured
130        if self.config.auto_sync && self.sync_counter >= self.config.sync_frequency {
131            self.sync()?;
132        }
133
134        Ok(bytes_written)
135    }
136
137    /// Read data from the memory-mapped file
138    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    /// Read all data from the file
158    pub fn read_all(&mut self) -> Result<Vec<u8>, OptimizerError> {
159        self.read(0, self.size)
160    }
161
162    /// Overwrite data at a specific offset
163    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        // Update size if we wrote beyond the current end
177        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    /// Sync the file to disk
191    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    /// Truncate the file to a specific size
200    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    /// Get the current file size
212    pub fn size(&self) -> usize {
213        self.size
214    }
215
216    /// Get the current file capacity
217    pub fn capacity(&self) -> usize {
218        self.capacity
219    }
220
221    /// Get the file path
222    pub fn path(&self) -> &Path {
223        &self.path
224    }
225
226    // Private methods
227
228    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
242/// Memory-mapped optimizer state storage
243pub 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    /// Create a new memory-mapped state storage
265    pub fn new(config: MemoryMappedConfig) -> Result<Self, OptimizerError> {
266        // Create base directory
267        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    /// Store optimizer state data
279    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    /// Store raw bytes
293    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        // Get or create the memory-mapped file
303        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        // Write data to the file
313        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        // Update metadata
321        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    /// Load optimizer state data
341    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    /// Load raw bytes
359    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        // Get the memory-mapped file
369        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        // Read data from the file
376        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    /// Update existing state data
383    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    /// Update raw bytes
397    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 new data fits in the existing space, overwrite
414        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            // Otherwise, append new data and update metadata
419            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        // Update timestamp
426        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    /// Remove state data
438    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    /// List all stored keys
444    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    /// Check if a key exists
450    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    /// Get storage statistics
456    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    /// Sync all files to disk
492    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    /// Compact storage by removing unused space
501    pub fn compact(&mut self) -> Result<(), OptimizerError> {
502        // Implementation would involve rewriting files to remove gaps
503        // For now, just sync all files
504        self.sync_all()
505    }
506}
507
508/// Storage statistics
509#[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
538/// Trait for optimizers that support memory-mapped states
539pub trait MemoryMappedSupport {
540    /// Save state to memory-mapped storage
541    fn save_to_mmap(
542        &self,
543        storage: &mut MemoryMappedStateStorage,
544        prefix: &str,
545    ) -> Result<(), OptimizerError>;
546
547    /// Load state from memory-mapped storage
548    fn load_from_mmap(
549        &mut self,
550        storage: &mut MemoryMappedStateStorage,
551        prefix: &str,
552    ) -> Result<(), OptimizerError>;
553
554    /// Get list of state keys for this optimizer
555    fn mmap_state_keys(&self, prefix: &str) -> Vec<String>;
556}
557
558/// Wrapper optimizer that uses memory-mapped storage
559pub 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    /// Create a new memory-mapped optimizer wrapper
573    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    /// Enable or disable auto-save
591    pub fn set_auto_save(&mut self, enabled: bool, frequency: usize) {
592        self.auto_save = enabled;
593        self.save_frequency = frequency;
594    }
595
596    /// Get the inner optimizer
597    pub fn inner(&self) -> &T {
598        &self.inner
599    }
600
601    /// Get the inner optimizer mutably
602    pub fn inner_mut(&mut self) -> &mut T {
603        &mut self.inner
604    }
605
606    /// Get the storage
607    pub fn storage(&self) -> &MemoryMappedStateStorage {
608        &self.storage
609    }
610
611    /// Get the storage mutably
612    pub fn storage_mut(&mut self) -> &mut MemoryMappedStateStorage {
613        &mut self.storage
614    }
615
616    /// Manually save state
617    pub fn save_state(&mut self) -> Result<(), OptimizerError> {
618        self.inner.save_to_mmap(&mut self.storage, &self.prefix)
619    }
620
621    /// Load state
622    pub fn load_state(&mut self) -> Result<(), OptimizerError> {
623        self.inner.load_from_mmap(&mut self.storage, &self.prefix)
624    }
625
626    /// Step the optimizer and handle auto-save
627    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        // Test writing
652        let data = b"Hello, world!";
653        let bytes_written = file.write(data).unwrap();
654        assert_eq!(bytes_written, data.len());
655
656        // Test reading
657        let read_data = file.read(0, data.len()).unwrap();
658        assert_eq!(read_data, data);
659
660        // Test write_at
661        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        // Test storing and loading data
678        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        // Test updating
689        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        // Test statistics
701        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        // Save state
759        mmap_optimizer.save_state().unwrap();
760
761        // Modify optimizer
762        mmap_optimizer.inner_mut().lr = 0.001;
763        mmap_optimizer.inner_mut().momentum = 0.95;
764
765        // Load state (should restore original values)
766        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}