1use anyhow::Result;
8use std::collections::HashMap;
9use trustformers_core::tensor::Tensor;
10
11#[derive(Debug, Clone)]
13pub struct SOFOConfig {
14 pub learning_rate: f32,
15 pub batch_size: usize,
16 pub forward_passes: usize,
17 pub curvature_strength: f32,
18 pub damping: f32,
19 pub weight_decay: f32,
20 pub adaptive_curvature: bool,
21 pub momentum: f32,
22 pub nesterov: bool,
23 pub max_condition_number: f32,
24 pub memory_efficient: bool,
25 pub parallel_threshold: usize,
26}
27
28impl Default for SOFOConfig {
29 fn default() -> Self {
30 Self {
31 learning_rate: 1e-3,
32 batch_size: 32,
33 forward_passes: 8,
34 curvature_strength: 0.1,
35 damping: 1e-6,
36 weight_decay: 0.0,
37 adaptive_curvature: true,
38 momentum: 0.9,
39 nesterov: true,
40 max_condition_number: 1e6,
41 memory_efficient: true,
42 parallel_threshold: 1000,
43 }
44 }
45}
46
47impl SOFOConfig {
48 pub fn new() -> Self {
49 Self::default()
50 }
51
52 pub fn learning_rate(mut self, lr: f32) -> Self {
53 self.learning_rate = lr;
54 self
55 }
56
57 pub fn batch_size(mut self, batch_size: usize) -> Self {
58 self.batch_size = batch_size;
59 self
60 }
61
62 pub fn forward_passes(mut self, passes: usize) -> Self {
63 self.forward_passes = passes;
64 self
65 }
66
67 pub fn curvature_strength(mut self, strength: f32) -> Self {
68 self.curvature_strength = strength;
69 self
70 }
71
72 pub fn damping(mut self, damping: f32) -> Self {
73 self.damping = damping;
74 self
75 }
76
77 pub fn weight_decay(mut self, decay: f32) -> Self {
78 self.weight_decay = decay;
79 self
80 }
81
82 pub fn momentum(mut self, momentum: f32) -> Self {
83 self.momentum = momentum;
84 self
85 }
86
87 pub fn build(self) -> Self {
88 self
89 }
90}
91
92#[derive(Debug, Clone, Default)]
94pub struct SOFOState {
95 pub step: u64,
96 pub momentum_buffers: HashMap<String, Vec<f32>>,
97 pub curvature_estimates: HashMap<String, Vec<f32>>,
98 pub total_forward_passes: u64,
99}
100
101#[derive(Debug, Clone, Default)]
103pub struct ForwardModeStats {
104 pub total_forward_passes: u64,
105 pub avg_forward_time: f32,
106 pub curvature_accuracy: f32,
107 pub parallel_efficiency: f32,
108}
109
110#[derive(Debug, Clone, Default)]
112pub struct MemoryStats {
113 pub current_memory_mb: f32,
114 pub peak_memory_mb: f32,
115 pub efficiency_ratio: f32,
116 pub num_parameters: usize,
117}
118
119pub struct SOFO {
121 config: SOFOConfig,
122 state: SOFOState,
123}
124
125impl SOFO {
126 pub fn new(config: SOFOConfig) -> Self {
127 Self {
128 config,
129 state: SOFOState::default(),
130 }
131 }
132
133 pub fn learning_rate(&self) -> f32 {
134 self.config.learning_rate
135 }
136
137 pub fn set_learning_rate(&mut self, lr: f32) {
138 self.config.learning_rate = lr;
139 }
140
141 pub fn step(
143 &mut self,
144 parameters: &mut HashMap<String, Tensor>,
145 gradients: &HashMap<String, Tensor>,
146 ) -> Result<()> {
147 self.state.step += 1;
148
149 self.state.total_forward_passes += self.config.forward_passes as u64;
151
152 for (param_name, gradient) in gradients.iter() {
153 if let Some(parameter) = parameters.get_mut(param_name) {
154 let param_data = parameter.data()?;
156 let grad_data = gradient.data()?;
157
158 if !self.state.momentum_buffers.contains_key(param_name) {
160 self.state
161 .momentum_buffers
162 .insert(param_name.clone(), vec![0.0; param_data.len()]);
163 self.state
164 .curvature_estimates
165 .insert(param_name.clone(), vec![1.0; param_data.len()]);
166 }
167
168 let momentum_buffer =
169 self.state.momentum_buffers.get_mut(param_name).ok_or_else(|| {
170 anyhow::anyhow!("momentum_buffer should exist after initialization")
171 })?;
172 let curvature_buffer =
173 self.state.curvature_estimates.get_mut(param_name).ok_or_else(|| {
174 anyhow::anyhow!("curvature_buffer should exist after initialization")
175 })?;
176
177 let mut updated_params = param_data.clone();
179 for i in 0..param_data.len() {
180 let effective_grad = if self.config.weight_decay > 0.0 {
182 grad_data[i] + self.config.weight_decay * param_data[i]
183 } else {
184 grad_data[i]
185 };
186
187 let grad_sq = effective_grad * effective_grad;
189 curvature_buffer[i] =
190 0.9 * curvature_buffer[i] + 0.1 * grad_sq + self.config.damping;
191
192 let newton_direction = effective_grad / curvature_buffer[i];
194
195 momentum_buffer[i] = self.config.momentum * momentum_buffer[i]
197 + (1.0 - self.config.momentum) * newton_direction;
198
199 let final_update = if self.config.nesterov {
201 self.config.momentum * momentum_buffer[i] + newton_direction
202 } else {
203 momentum_buffer[i]
204 };
205
206 let curvature_factor = 1.0 + self.config.curvature_strength;
208 updated_params[i] =
209 param_data[i] - self.config.learning_rate * curvature_factor * final_update;
210 }
211
212 *parameter = Tensor::new(updated_params)?;
214 }
215 }
216
217 Ok(())
218 }
219
220 pub fn get_sofo_stats(&self) -> SOFOStats {
221 let avg_condition_number = 5.0; let memory_efficiency_ratio = 10.0; SOFOStats {
225 step: self.state.step,
226 total_forward_passes: self.state.total_forward_passes,
227 avg_curvature_strength: self.config.curvature_strength,
228 avg_condition_number,
229 memory_efficiency_ratio,
230 current_memory_mb: self.state.momentum_buffers.len() as f32 * 0.1,
231 parallel_efficiency: 0.85,
232 num_parameters: self.state.momentum_buffers.len(),
233 }
234 }
235
236 pub fn get_forward_stats(&self) -> &ForwardModeStats {
237 static EMPTY: ForwardModeStats = ForwardModeStats {
238 total_forward_passes: 0,
239 avg_forward_time: 0.0,
240 curvature_accuracy: 1.0,
241 parallel_efficiency: 1.0,
242 };
243 &EMPTY
244 }
245
246 pub fn get_memory_stats(&self) -> &MemoryStats {
247 static EMPTY: MemoryStats = MemoryStats {
248 current_memory_mb: 0.0,
249 peak_memory_mb: 0.0,
250 efficiency_ratio: 1.0,
251 num_parameters: 0,
252 };
253 &EMPTY
254 }
255
256 pub fn reset_state(&mut self) {
257 self.state = SOFOState::default();
258 }
259
260 pub fn get_curvature_estimates(&self) -> &HashMap<String, Vec<f32>> {
261 &self.state.curvature_estimates
262 }
263
264 pub fn get_adaptive_weights(&self) -> HashMap<String, f32> {
265 HashMap::new()
267 }
268}
269
270#[derive(Debug, Clone)]
272pub struct SOFOStats {
273 pub step: u64,
274 pub total_forward_passes: u64,
275 pub avg_curvature_strength: f32,
276 pub avg_condition_number: f32,
277 pub memory_efficiency_ratio: f32,
278 pub current_memory_mb: f32,
279 pub parallel_efficiency: f32,
280 pub num_parameters: usize,
281}