Skip to content

Commit 65a4cd2

Browse files
schroedkclaude
andcommitted
Refactor IronsTuckGrand for cleaner control flow
- Add finalize_output() to consolidate coefficient copy and return - Add project_and_check() to combine projection with convergence check - Remove conv_len parameter from methods; compute via projector internally - Use local bindings for convergence state to avoid repetition - Align method signatures to consistently take projector first - Rename acceleration_step to acceleration_step_check 🤖 Generated with [Claude Code](https://claude.com/claude-code) Co-Authored-By: Claude Opus 4.5 <noreply@anthropic.com>
1 parent 88c4832 commit 65a4cd2

1 file changed

Lines changed: 40 additions & 32 deletions

File tree

src/demean_accelerated/accelerator.rs

Lines changed: 40 additions & 32 deletions
Original file line numberDiff line numberDiff line change
@@ -154,62 +154,66 @@ impl IronsTuckGrand {
154154
projector.coef_len()
155155
);
156156

157-
let conv_len = projector.convergence_len();
158-
159157
// Initial projection and convergence check
160-
projector.project(coef, &mut self.buffers.gx);
161-
if self.check_coef_convergence(coef, conv_len) == ConvergenceState::Converged {
162-
coef.copy_from_slice(&self.buffers.gx);
163-
return (0, ConvergenceState::Converged);
158+
let conv = self.project_and_check(projector, coef);
159+
if conv == ConvergenceState::Converged {
160+
return self.finalize_output(coef, 0, conv);
164161
}
165162

166163
let mut grand_phase = GrandPhase::default();
167164
let mut ssr = 0.0;
168165

169166
for iter in 1..=max_iter {
170167
// Core acceleration step
171-
if self.acceleration_step(projector, coef, conv_len, iter) == ConvergenceState::Converged
172-
{
173-
coef.copy_from_slice(&self.buffers.gx);
174-
return (iter, ConvergenceState::Converged);
168+
let conv = self.acceleration_step_check(projector, coef, iter);
169+
if conv == ConvergenceState::Converged {
170+
return self.finalize_output(coef, iter, conv);
175171
}
176172

177173
// Grand acceleration (every iter_grand_acc iterations)
178174
if iter % self.config.iter_grand_acc == 0 {
179-
if self.grand_acceleration_check(&mut grand_phase, projector, conv_len)
180-
== ConvergenceState::Converged
181-
{
182-
coef.copy_from_slice(&self.buffers.gx);
183-
return (iter, ConvergenceState::Converged);
175+
let conv = self.grand_acceleration_check(projector, &mut grand_phase);
176+
if conv == ConvergenceState::Converged {
177+
return self.finalize_output(coef, iter, conv);
184178
}
185179
}
186180

187181
// SSR convergence check (every ssr_check_interval iterations)
188182
if iter % self.config.ssr_check_interval == 0 {
189-
if self.ssr_convergence_check(projector, iter, &mut ssr)
190-
== ConvergenceState::Converged
191-
{
192-
coef.copy_from_slice(&self.buffers.gx);
193-
return (iter, ConvergenceState::Converged);
183+
let conv = self.ssr_convergence_check(projector, iter, &mut ssr);
184+
if conv == ConvergenceState::Converged {
185+
return self.finalize_output(coef, iter, conv);
194186
}
195187
}
196188
}
189+
self.finalize_output(coef, max_iter, ConvergenceState::NotConverged)
190+
}
197191

192+
/// Copy converged coefficients to the output buffer.
193+
///
194+
/// This method should be called after `run()` has completed to retrieve
195+
/// the final coefficients from the internal `gx` buffer.
196+
#[inline]
197+
fn finalize_output(&self, coef: &mut [f64],
198+
iter: usize,
199+
convergence: ConvergenceState,) -> (usize, ConvergenceState) {
198200
coef.copy_from_slice(&self.buffers.gx);
199-
(max_iter, ConvergenceState::NotConverged)
201+
(iter, convergence)
202+
200203
}
201204

202205
/// Perform the core Irons-Tuck acceleration step.
203206
///
204207
/// Returns `Converged` if convergence detected, `NotConverged` to continue.
205208
#[inline]
206-
fn acceleration_step<P: Projector>(
209+
fn acceleration_step_check<P: Projector>(
207210
&mut self,
208211
projector: &mut P,
209212
coef: &mut [f64],
210-
conv_len: usize,
211213
iter: usize,
212214
) -> ConvergenceState {
215+
let conv_len = projector.convergence_len();
216+
213217
// Double projection for Irons-Tuck: G(G(x))
214218
projector.project(&self.buffers.gx, &mut self.buffers.ggx);
215219

@@ -230,19 +234,17 @@ impl IronsTuckGrand {
230234
}
231235

232236
// Update gx and check coefficient convergence
233-
projector.project(coef, &mut self.buffers.gx);
234-
self.check_coef_convergence(coef, conv_len)
237+
self.project_and_check(projector, coef)
235238
}
236239

237240
/// Perform grand acceleration and check for convergence.
238241
#[inline]
239242
fn grand_acceleration_check<P: Projector>(
240243
&mut self,
241-
grand_phase: &mut GrandPhase,
242244
projector: &mut P,
243-
conv_len: usize,
245+
grand_phase: &mut GrandPhase,
244246
) -> ConvergenceState {
245-
match self.grand_acceleration_step(*grand_phase, projector, conv_len) {
247+
match self.grand_acceleration_step(projector, *grand_phase) {
246248
GrandStepResult::Continue(next) => {
247249
*grand_phase = next;
248250
ConvergenceState::NotConverged
@@ -270,9 +272,15 @@ impl IronsTuckGrand {
270272
}
271273
}
272274

273-
/// Check if coefficients have converged.
275+
/// Project coefficients and check for convergence.
274276
#[inline]
275-
fn check_coef_convergence(&self, coef: &[f64], conv_len: usize) -> ConvergenceState {
277+
fn project_and_check<P: Projector>(
278+
&mut self,
279+
projector: &mut P,
280+
coef: &[f64],
281+
) -> ConvergenceState {
282+
projector.project(coef, &mut self.buffers.gx);
283+
let conv_len = projector.convergence_len();
276284
if Self::should_continue(
277285
&coef[..conv_len],
278286
&self.buffers.gx[..conv_len],
@@ -340,10 +348,10 @@ impl IronsTuckGrand {
340348
#[inline]
341349
fn grand_acceleration_step<P: Projector>(
342350
&mut self,
343-
phase: GrandPhase,
344351
projector: &mut P,
345-
conv_len: usize,
352+
phase: GrandPhase,
346353
) -> GrandStepResult {
354+
let conv_len = projector.convergence_len();
347355
match phase {
348356
GrandPhase::Collect1st => {
349357
self.buffers.y[..conv_len].copy_from_slice(&self.buffers.gx[..conv_len]);

0 commit comments

Comments
 (0)