1use ndarray::Array1;
11
12use crate::gpu_kernels::arrow_schur::{
13 ArrowSchurGpuFailure, solve_arrow_newton_step, solve_arrow_newton_step_dense_reference,
14};
15use gam_problem::ExecutionPath;
16
17#[derive(Clone, Copy, Debug, Eq, PartialEq)]
35pub enum InnerSolveMode {
36 DeviceResident,
37 DeviceReupload,
38 CpuReference,
39}
40
41impl InnerSolveMode {
42 #[inline]
47 const fn execution_path(self) -> ExecutionPath {
48 match self {
49 Self::DeviceResident => ExecutionPath::GpuResidentFull,
50 Self::DeviceReupload => ExecutionPath::GpuReupload,
51 Self::CpuReference => ExecutionPath::Cpu,
52 }
53 }
54}
55use crate::arrow_schur::{ArrowSchurError, ArrowSchurSystem};
56
57#[derive(Clone, Copy, Debug, Eq, PartialEq)]
64pub struct DeviceResidentArrowShape {
65 pub n: usize,
66 pub p: usize,
67 pub basis_cols: usize,
68 pub d: usize,
69}
70
71impl DeviceResidentArrowShape {
72 #[inline]
73 pub const fn qwen_non_gating() -> Self {
74 Self {
75 n: 2_000,
76 p: 2_048,
77 basis_cols: 8,
78 d: 2,
79 }
80 }
81
82 #[inline]
86 pub const fn color_arm() -> Self {
87 Self {
88 n: 180,
89 p: 5_120,
90 basis_cols: 9,
91 d: 2,
92 }
93 }
94
95 #[inline]
96 pub const fn target_len(self) -> usize {
97 self.n * self.p
98 }
99
100 #[inline]
101 pub const fn basis_len(self) -> usize {
102 self.n * self.basis_cols
103 }
104
105 #[inline]
106 pub const fn row_hessian_len(self) -> usize {
107 self.n * self.d * self.d
108 }
109
110 #[inline]
111 pub const fn row_cross_len(self) -> usize {
112 self.n * self.d * self.p
113 }
114
115 #[inline]
116 pub const fn row_gradient_len(self) -> usize {
117 self.n * self.d
118 }
119
120 #[inline]
121 pub const fn border_hessian_len(self) -> usize {
122 self.p * self.p
123 }
124}
125
126#[derive(Clone, Debug)]
133pub struct DeviceResidentArrowSlabs {
134 pub row_hessian_slabs: Vec<f64>,
135 pub row_cross_slabs: Vec<f64>,
136 pub row_gradient_slabs: Vec<f64>,
137 pub border_hessian: Vec<f64>,
138 pub border_gradient: Vec<f64>,
139}
140
141#[derive(Clone, Debug)]
143pub struct DeviceResidentArrowStep {
144 pub delta_t: Array1<f64>,
145 pub delta_beta: Array1<f64>,
146 pub objective: f64,
147 pub gradient_norm: f64,
148 pub log_det_hessian: f64,
149 pub execution_path: ExecutionPath,
150}
151
152#[derive(Debug, Clone)]
153pub enum DeviceResidentArrowError {
154 Shape { reason: String },
155 Unavailable { reason: String },
156 Solve { reason: String },
157}
158
159impl std::fmt::Display for DeviceResidentArrowError {
160 fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
161 match self {
162 Self::Shape { reason } | Self::Unavailable { reason } | Self::Solve { reason } => {
163 f.write_str(reason)
164 }
165 }
166 }
167}
168
169impl std::error::Error for DeviceResidentArrowError {}
170
171#[cfg(target_os = "linux")]
172pub struct DeviceResidentArrowBuffers {
173 pub stream: std::sync::Arc<cudarc::driver::CudaStream>,
174 pub target_x_dev: cudarc::driver::CudaSlice<f64>,
175 pub basis_values_dev: cudarc::driver::CudaSlice<f64>,
176 pub gate_activations_dev: cudarc::driver::CudaSlice<f64>,
177 pub row_hessian_dev: cudarc::driver::CudaSlice<f64>,
178 pub row_cross_dev: cudarc::driver::CudaSlice<f64>,
179 pub row_gradient_dev: cudarc::driver::CudaSlice<f64>,
180 pub border_hessian_dev: cudarc::driver::CudaSlice<f64>,
181 pub border_gradient_dev: cudarc::driver::CudaSlice<f64>,
182 pub bytes: usize,
183}
184
185pub struct DeviceResidentArrowWorkspace {
187 shape: DeviceResidentArrowShape,
188 target_x: Vec<f64>,
189 basis_values: Vec<f64>,
190 gate_activations: Vec<f64>,
191 slabs: DeviceResidentArrowSlabs,
192 #[cfg(target_os = "linux")]
193 device: Option<DeviceResidentArrowBuffers>,
194}
195
196impl DeviceResidentArrowWorkspace {
197 pub fn new(
198 shape: DeviceResidentArrowShape,
199 target_x: Vec<f64>,
200 basis_values: Vec<f64>,
201 gate_activations: Vec<f64>,
202 slabs: DeviceResidentArrowSlabs,
203 ) -> Result<Self, DeviceResidentArrowError> {
204 validate_shape(shape, &target_x, &basis_values, &gate_activations, &slabs)?;
205 #[cfg(target_os = "linux")]
206 let device =
207 upload_resident_buffers(shape, &target_x, &basis_values, &gate_activations, &slabs);
208 Ok(Self {
209 shape,
210 target_x,
211 basis_values,
212 gate_activations,
213 slabs,
214 #[cfg(target_os = "linux")]
215 device,
216 })
217 }
218
219 #[inline]
220 pub const fn shape(&self) -> DeviceResidentArrowShape {
221 self.shape
222 }
223
224 #[must_use]
225 pub fn device_resident(&self) -> bool {
226 #[cfg(target_os = "linux")]
227 {
228 self.device.is_some()
229 }
230 #[cfg(not(target_os = "linux"))]
231 {
232 false
233 }
234 }
235
236 #[must_use]
237 pub fn resident_device_bytes(&self) -> usize {
238 #[cfg(target_os = "linux")]
239 {
240 self.device.as_ref().map_or(0, |device| device.bytes)
241 }
242 #[cfg(not(target_os = "linux"))]
243 {
244 0
245 }
246 }
247
248 #[must_use]
253 fn context_id(&self) -> usize {
254 usize::from(self.device_resident())
255 }
256
257 #[must_use]
260 fn frame_upload_bytes(&self) -> usize {
261 [
262 self.slabs.row_hessian_slabs.len(),
263 self.slabs.row_cross_slabs.len(),
264 self.slabs.row_gradient_slabs.len(),
265 self.slabs.border_hessian.len(),
266 self.slabs.border_gradient.len(),
267 ]
268 .into_iter()
269 .sum::<usize>()
270 * std::mem::size_of::<f64>()
271 }
272
273 #[must_use]
274 pub fn host_shadow_bytes(&self) -> usize {
275 [
276 self.target_x.len(),
277 self.basis_values.len(),
278 self.gate_activations.len(),
279 self.slabs.row_hessian_slabs.len(),
280 self.slabs.row_cross_slabs.len(),
281 self.slabs.row_gradient_slabs.len(),
282 self.slabs.border_hessian.len(),
283 self.slabs.border_gradient.len(),
284 ]
285 .into_iter()
286 .sum::<usize>()
287 * std::mem::size_of::<f64>()
288 }
289
290 pub fn one_inner_iteration(
293 &self,
294 ridge_t: f64,
295 ridge_beta: f64,
296 ) -> Result<DeviceResidentArrowStep, DeviceResidentArrowError> {
297 if !self.device_resident() {
298 return Err(DeviceResidentArrowError::Unavailable {
299 reason: "SAE resident inner iteration unavailable: CUDA runtime did not admit the qwen-scale row-block workload".to_string(),
300 });
301 }
302 let sys = self.to_arrow_system();
303 let frame = crate::gpu_kernels::arrow_schur::ResidentArrowFrameHandle::new(
304 &sys, ridge_t, ridge_beta,
305 )
306 .map_err(map_gpu_error)?;
307 let g_t: Vec<f64> = sys
308 .rows
309 .iter()
310 .flat_map(|row| row.gt.iter().copied())
311 .collect();
312 let g_beta: Vec<f64> = sys.gb.iter().copied().collect();
313 frame
319 .solve_gradient(&g_t, &g_beta)
320 .map(|solution| self.finish_step(solution, ExecutionPath::GpuResidentLinearization))
321 .map_err(map_gpu_error)
322 }
323
324 pub fn cpu_reference_step(
327 &self,
328 ridge_t: f64,
329 ridge_beta: f64,
330 ) -> Result<DeviceResidentArrowStep, DeviceResidentArrowError> {
331 let sys = self.to_arrow_system();
332 solve_arrow_newton_step_dense_reference(&sys, ridge_t, ridge_beta)
333 .map(|solution| self.finish_step(solution, ExecutionPath::Cpu))
334 .map_err(|reason| DeviceResidentArrowError::Solve { reason })
335 }
336
337 pub fn inner_iteration_for_production(
362 &self,
363 mode: gam_gpu::GpuPolicy,
364 ridge_t: f64,
365 ridge_beta: f64,
366 ) -> Result<DeviceResidentArrowStep, DeviceResidentArrowError> {
367 match mode {
368 gam_gpu::GpuPolicy::Off => {
369 note_resident_engagement(false, "GpuPolicy::Off — CPU reference step");
370 self.cpu_reference_step(ridge_t, ridge_beta)
371 }
372 gam_gpu::GpuPolicy::Required => {
373 if !self.device_resident() {
374 return Err(DeviceResidentArrowError::Unavailable {
375 reason: format!(
376 "SAE resident inner step GpuPolicy::Required: workspace is not \
377 device-resident (the CUDA runtime did not admit shape n={} p={} d={} \
378 at break-even); refusing to run on the CPU",
379 self.shape.n, self.shape.p, self.shape.d
380 ),
381 });
382 }
383 note_resident_engagement(true, "GpuPolicy::Required — resident device step");
384 self.one_inner_iteration(ridge_t, ridge_beta)
385 }
386 gam_gpu::GpuPolicy::Auto => {
387 if !self.device_resident() {
388 note_resident_engagement(
389 false,
390 "GpuPolicy::Auto — workspace not device-resident; CPU reference step",
391 );
392 return self.cpu_reference_step(ridge_t, ridge_beta);
393 }
394 match self.one_inner_iteration(ridge_t, ridge_beta) {
395 Ok(step) => {
396 note_resident_engagement(true, "GpuPolicy::Auto — resident device step");
397 Ok(step)
398 }
399 Err(err) => {
400 note_resident_engagement(
401 false,
402 &format!(
403 "GpuPolicy::Auto — device solve fault, CPU reference fallback: {err}"
404 ),
405 );
406 self.cpu_reference_step(ridge_t, ridge_beta)
407 }
408 }
409 }
410 }
411 }
412
413 pub fn to_arrow_system(&self) -> ArrowSchurSystem {
414 let shape = self.shape;
415 let mut sys = ArrowSchurSystem::new(shape.n, shape.d, shape.p);
416 for i in 0..shape.n {
417 let h_base = i * shape.d * shape.d;
418 let b_base = i * shape.d * shape.p;
419 let g_base = i * shape.d;
420 for r in 0..shape.d {
421 for c in 0..shape.d {
422 sys.rows[i].htt[[r, c]] =
423 self.slabs.row_hessian_slabs[h_base + r * shape.d + c];
424 }
425 sys.rows[i].gt[r] = self.slabs.row_gradient_slabs[g_base + r];
426 for c in 0..shape.p {
427 sys.rows[i].htbeta[[r, c]] =
428 self.slabs.row_cross_slabs[b_base + r * shape.p + c];
429 }
430 }
431 }
432 for r in 0..shape.p {
433 sys.gb[r] = self.slabs.border_gradient[r];
434 for c in 0..shape.p {
435 sys.hbb[[r, c]] = self.slabs.border_hessian[r * shape.p + c];
436 }
437 }
438 sys.refresh_row_hessian_fingerprint();
439 sys
440 }
441
442 fn finish_step(
443 &self,
444 solution: crate::gpu_kernels::arrow_schur::ArrowSchurGpuSolution,
445 execution_path: ExecutionPath,
446 ) -> DeviceResidentArrowStep {
447 DeviceResidentArrowStep {
448 delta_t: solution.delta_t,
449 delta_beta: solution.delta_beta,
450 objective: 0.5 * squared_norm(&self.target_x),
451 gradient_norm: self.gradient_norm(),
452 log_det_hessian: solution.log_det_hessian,
453 execution_path,
454 }
455 }
456
457 fn gradient_norm(&self) -> f64 {
458 let row = squared_norm(&self.slabs.row_gradient_slabs);
459 let border = squared_norm(&self.slabs.border_gradient);
460 (row + border).sqrt()
461 }
462
463 pub fn device_fit(
492 &self,
493 opts: &DeviceResidentInnerOptions,
494 ) -> Result<DeviceResidentInnerOutcome, DeviceResidentArrowError> {
495 if !self.device_resident() {
496 return Err(DeviceResidentArrowError::Unavailable {
497 reason: "SAE resident inner loop unavailable: CUDA runtime did not admit the qwen-scale row-block workload".to_string(),
498 });
499 }
500 self.run_inner_loop(opts, InnerSolveMode::DeviceResident)
501 }
502
503 pub fn device_reupload_fit(
511 &self,
512 opts: &DeviceResidentInnerOptions,
513 ) -> Result<DeviceResidentInnerOutcome, DeviceResidentArrowError> {
514 if !self.device_resident() {
515 return Err(DeviceResidentArrowError::Unavailable {
516 reason: "SAE re-uploading inner loop unavailable: CUDA runtime did not admit the row-block workload".to_string(),
517 });
518 }
519 self.run_inner_loop(opts, InnerSolveMode::DeviceReupload)
520 }
521
522 pub fn cpu_reference_fit(
526 &self,
527 opts: &DeviceResidentInnerOptions,
528 ) -> Result<DeviceResidentInnerOutcome, DeviceResidentArrowError> {
529 self.run_inner_loop(opts, InnerSolveMode::CpuReference)
530 }
531
532 fn run_inner_loop(
533 &self,
534 opts: &DeviceResidentInnerOptions,
535 mode: InnerSolveMode,
536 ) -> Result<DeviceResidentInnerOutcome, DeviceResidentArrowError> {
537 let execution_path = mode.execution_path();
538 let n = self.shape.n;
539 let d = self.shape.d;
540 let p = self.shape.p;
541 let t_len = n * d;
542
543 let mut t = vec![0.0_f64; t_len];
547 let mut beta = vec![0.0_f64; p];
548
549 let base = self.to_arrow_system();
550 let half_target_energy = 0.5 * squared_norm(&self.target_x);
551
552 let mut ridge_t = opts.initial_ridge_t.max(0.0);
553 let mut ridge_beta = opts.initial_ridge_beta.max(0.0);
554 let mut resident_frame: Option<(
563 f64,
564 f64,
565 crate::gpu_kernels::arrow_schur::ResidentArrowFrameHandle,
566 )> = None;
567 let mut current_objective = self.objective_at(&base, half_target_energy, &t, &beta);
568 let mut accepted_iters = 0_usize;
569 let mut total_iters = 0_usize;
570 let mut converged = false;
571 let mut last_step = DeviceResidentArrowStep {
572 delta_t: Array1::zeros(t_len),
573 delta_beta: Array1::zeros(p),
574 objective: current_objective,
575 gradient_norm: 0.0,
576 log_det_hessian: 0.0,
577 execution_path,
578 };
579
580 while total_iters < opts.max_iterations {
581 let residual = self.residual_system(&base, &t, &beta);
583 let g_norm = arrow_system_gradient_norm(&residual);
584 let scale = 1.0 + iterate_norm(&t, &beta);
585 if g_norm / scale < opts.convergence_tolerance {
586 converged = true;
587 break;
588 }
589
590 let solution = match mode {
591 InnerSolveMode::DeviceResident => {
592 let frame_matches = resident_frame
597 .as_ref()
598 .is_some_and(|(rt, rb, _)| *rt == ridge_t && *rb == ridge_beta);
599 let mut frame_build_error: Option<DeviceResidentArrowError> = None;
600 if !frame_matches {
601 resident_frame = None;
602 match crate::gpu_kernels::arrow_schur::ResidentArrowFrameHandle::new(
603 &residual, ridge_t, ridge_beta,
604 ) {
605 Ok(frame) => {
606 gam_gpu::profile::telemetry_record_handle_creation(
612 self.context_id(),
613 );
614 gam_gpu::profile::telemetry_record_factorization();
615 gam_gpu::profile::telemetry_record_h2d(self.frame_upload_bytes());
616 resident_frame = Some((ridge_t, ridge_beta, frame));
617 }
618 Err(err) => frame_build_error = Some(map_gpu_error(err)),
619 }
620 }
621 match resident_frame.as_ref() {
622 Some((_, _, frame)) => {
623 let mut g_t = Vec::with_capacity(n * d);
626 for row in &residual.rows {
627 for &v in row.gt.iter() {
628 g_t.push(v);
629 }
630 }
631 let g_beta: Vec<f64> = residual.gb.iter().copied().collect();
632 let grad_bytes =
636 (g_t.len() + g_beta.len()) * std::mem::size_of::<f64>();
637 gam_gpu::profile::telemetry_record_h2d(grad_bytes);
638 gam_gpu::profile::telemetry_record_kernel_launch();
639 gam_gpu::profile::telemetry_record_d2h(
640 (n * d + p) * std::mem::size_of::<f64>(),
641 );
642 frame.solve_gradient(&g_t, &g_beta).map_err(map_gpu_error)
643 }
644 None => Err(frame_build_error.unwrap_or_else(|| {
645 DeviceResidentArrowError::Solve {
646 reason: "SAE resident frame build declined".to_string(),
647 }
648 })),
649 }
650 }
651 InnerSolveMode::DeviceReupload => {
652 gam_gpu::profile::telemetry_record_handle_creation(self.context_id());
658 gam_gpu::profile::telemetry_record_factorization();
659 gam_gpu::profile::telemetry_record_h2d(self.frame_upload_bytes());
660 gam_gpu::profile::telemetry_record_kernel_launch();
661 gam_gpu::profile::telemetry_record_d2h(
662 (n * d + p) * std::mem::size_of::<f64>(),
663 );
664 solve_arrow_newton_step(&residual, ridge_t, ridge_beta).map_err(map_gpu_error)
665 }
666 InnerSolveMode::CpuReference => {
667 solve_arrow_newton_step_dense_reference(&residual, ridge_t, ridge_beta)
668 .map_err(|reason| DeviceResidentArrowError::Solve { reason })
669 }
670 };
671
672 let solution = match solution {
673 Ok(sol) => sol,
674 Err(DeviceResidentArrowError::Solve { .. })
675 | Err(DeviceResidentArrowError::Unavailable { .. }) => {
676 ridge_t = grow_ridge(ridge_t, opts.lm_grow);
680 ridge_beta = grow_ridge(ridge_beta, opts.lm_grow);
681 if ridge_t > opts.max_ridge || ridge_beta > opts.max_ridge {
682 return Err(DeviceResidentArrowError::Solve {
683 reason: format!(
684 "SAE resident inner loop: LM ridge exceeded max ({:e}) at iter {total_iters}",
685 opts.max_ridge
686 ),
687 });
688 }
689 total_iters += 1;
690 continue;
691 }
692 Err(other) => return Err(other),
693 };
694
695 let predicted_reduction = crate::arrow_schur::arrow_bare_quadratic_model_reduction(
698 &residual,
699 solution.delta_t.view(),
700 solution.delta_beta.view(),
701 ridge_t,
702 ridge_beta,
703 )
704 .map_err(|err| DeviceResidentArrowError::Solve {
705 reason: format!("SAE resident inner loop predicted-reduction failed: {err}"),
706 })?;
707
708 let mut trial_t = t.clone();
710 let mut trial_beta = beta.clone();
711 for (slot, dv) in trial_t.iter_mut().zip(solution.delta_t.iter()) {
712 *slot += *dv;
713 }
714 for (slot, dv) in trial_beta.iter_mut().zip(solution.delta_beta.iter()) {
715 *slot += *dv;
716 }
717 let trial_objective =
718 self.objective_at(&base, half_target_energy, &trial_t, &trial_beta);
719
720 let objective_scale = current_objective.abs();
732 let noise_floor = objective_scale * 1e-14;
733 let actual_reduction = current_objective - trial_objective;
734 let rho = if predicted_reduction > noise_floor {
735 actual_reduction / predicted_reduction
736 } else if actual_reduction >= -noise_floor {
737 1.0
738 } else {
739 -1.0
740 };
741
742 if rho > 0.0 && trial_objective.is_finite() {
743 t = trial_t;
744 beta = trial_beta;
745 current_objective = trial_objective;
746 ridge_t = (ridge_t * opts.lm_shrink).max(0.0);
747 ridge_beta = (ridge_beta * opts.lm_shrink).max(0.0);
748 last_step = DeviceResidentArrowStep {
749 delta_t: solution.delta_t,
750 delta_beta: solution.delta_beta,
751 objective: current_objective,
752 gradient_norm: g_norm,
753 log_det_hessian: solution.log_det_hessian,
754 execution_path,
755 };
756 accepted_iters += 1;
757 total_iters += 1;
758 } else {
759 ridge_t = grow_ridge(ridge_t, opts.lm_grow);
760 ridge_beta = grow_ridge(ridge_beta, opts.lm_grow);
761 if ridge_t > opts.max_ridge || ridge_beta > opts.max_ridge {
762 return Err(DeviceResidentArrowError::Solve {
763 reason: format!(
764 "SAE resident inner loop: LM rejected step until ridge exceeded max ({:e}) at iter {total_iters} (rho={rho:.3e})",
765 opts.max_ridge
766 ),
767 });
768 }
769 total_iters += 1;
770 }
771 }
772
773 Ok(DeviceResidentInnerOutcome {
774 t: Array1::from_vec(t),
775 beta: Array1::from_vec(beta),
776 objective: current_objective,
777 gradient_norm: last_step.gradient_norm,
778 log_det_hessian: last_step.log_det_hessian,
779 iterations: total_iters,
780 accepted_iterations: accepted_iters,
781 converged,
782 execution_path,
783 })
784 }
785
786 pub fn device_fit_outer_sequence(
822 &self,
823 base_gradient_overrides: &[(Vec<f64>, Vec<f64>)],
824 opts: &DeviceResidentInnerOptions,
825 ) -> Result<OuterSequenceOutcome, DeviceResidentArrowError> {
826 if !self.device_resident() {
827 return Err(DeviceResidentArrowError::Unavailable {
828 reason: "SAE outer-sequence residency unavailable: CUDA runtime did not admit the row-block workload".to_string(),
829 });
830 }
831 self.run_outer_sequence(
832 base_gradient_overrides,
833 opts,
834 InnerSolveMode::DeviceResident,
835 )
836 }
837
838 pub fn cpu_reference_outer_sequence(
843 &self,
844 base_gradient_overrides: &[(Vec<f64>, Vec<f64>)],
845 opts: &DeviceResidentInnerOptions,
846 ) -> Result<OuterSequenceOutcome, DeviceResidentArrowError> {
847 self.run_outer_sequence(base_gradient_overrides, opts, InnerSolveMode::CpuReference)
848 }
849
850 fn run_outer_sequence(
851 &self,
852 base_gradient_overrides: &[(Vec<f64>, Vec<f64>)],
853 opts: &DeviceResidentInnerOptions,
854 mode: InnerSolveMode,
855 ) -> Result<OuterSequenceOutcome, DeviceResidentArrowError> {
856 let n = self.shape.n;
857 let d = self.shape.d;
858 let p = self.shape.p;
859 let t_len = n * d;
860 let half_target_energy = 0.5 * squared_norm(&self.target_x);
861
862 let mut shared = SharedFrameState::default();
869 let mut outcomes = Vec::with_capacity(base_gradient_overrides.len());
870
871 for (g_t_override, g_beta_override) in base_gradient_overrides {
872 if g_t_override.len() != t_len || g_beta_override.len() != p {
873 return Err(DeviceResidentArrowError::Shape {
874 reason: format!(
875 "outer-sequence gradient shape mismatch: g_t={} (want {t_len}), g_beta={} (want {p})",
876 g_t_override.len(),
877 g_beta_override.len()
878 ),
879 });
880 }
881 let mut base = self.to_arrow_system();
884 for (i, row) in base.rows.iter_mut().enumerate() {
885 for r in 0..d {
886 row.gt[r] = g_t_override[i * d + r];
887 }
888 }
889 for (j, gb) in base.gb.iter_mut().enumerate() {
890 *gb = g_beta_override[j];
891 }
892 base.refresh_row_hessian_fingerprint();
893
894 let outcome = self.run_one_outer(&base, half_target_energy, opts, mode, &mut shared)?;
895 outcomes.push(outcome);
896 }
897
898 Ok(OuterSequenceOutcome {
899 outers: outcomes,
900 frame_builds: shared.frame_builds,
901 })
902 }
903
904 fn run_one_outer(
911 &self,
912 base: &ArrowSchurSystem,
913 half_target_energy: f64,
914 opts: &DeviceResidentInnerOptions,
915 mode: InnerSolveMode,
916 shared: &mut SharedFrameState,
917 ) -> Result<DeviceResidentInnerOutcome, DeviceResidentArrowError> {
918 let execution_path = mode.execution_path();
919 let n = self.shape.n;
920 let d = self.shape.d;
921 let p = self.shape.p;
922 let t_len = n * d;
923
924 let mut t = vec![0.0_f64; t_len];
925 let mut beta = vec![0.0_f64; p];
926 let mut ridge_t = opts.initial_ridge_t.max(0.0);
927 let mut ridge_beta = opts.initial_ridge_beta.max(0.0);
928 let mut current_objective = self.objective_at(base, half_target_energy, &t, &beta);
929 let mut accepted_iters = 0_usize;
930 let mut total_iters = 0_usize;
931 let mut converged = false;
932 let mut last_gradient_norm = 0.0_f64;
933 let mut last_log_det = 0.0_f64;
934
935 while total_iters < opts.max_iterations {
936 let residual = self.residual_system(base, &t, &beta);
937 let g_norm = arrow_system_gradient_norm(&residual);
938 let scale = 1.0 + iterate_norm(&t, &beta);
939 if g_norm / scale < opts.convergence_tolerance {
940 converged = true;
941 break;
942 }
943
944 let solution = match mode {
945 InnerSolveMode::DeviceResident => {
946 let frame_matches = shared
947 .frame
948 .as_ref()
949 .is_some_and(|(rt, rb, _)| *rt == ridge_t && *rb == ridge_beta);
950 let mut frame_build_error: Option<DeviceResidentArrowError> = None;
951 if !frame_matches {
952 shared.frame = None;
953 match crate::gpu_kernels::arrow_schur::ResidentArrowFrameHandle::new(
954 &residual, ridge_t, ridge_beta,
955 ) {
956 Ok(frame) => {
957 shared.frame_builds += 1;
958 gam_gpu::profile::telemetry_record_handle_creation(
959 self.context_id(),
960 );
961 gam_gpu::profile::telemetry_record_factorization();
962 gam_gpu::profile::telemetry_record_h2d(self.frame_upload_bytes());
963 shared.frame = Some((ridge_t, ridge_beta, frame));
964 }
965 Err(err) => frame_build_error = Some(map_gpu_error(err)),
966 }
967 }
968 match shared.frame.as_ref() {
969 Some((_, _, frame)) => {
970 let mut g_t = Vec::with_capacity(n * d);
971 for row in &residual.rows {
972 for &v in row.gt.iter() {
973 g_t.push(v);
974 }
975 }
976 let g_beta: Vec<f64> = residual.gb.iter().copied().collect();
977 let grad_bytes =
978 (g_t.len() + g_beta.len()) * std::mem::size_of::<f64>();
979 gam_gpu::profile::telemetry_record_h2d(grad_bytes);
980 gam_gpu::profile::telemetry_record_kernel_launch();
981 gam_gpu::profile::telemetry_record_d2h(
982 (n * d + p) * std::mem::size_of::<f64>(),
983 );
984 frame.solve_gradient(&g_t, &g_beta).map_err(map_gpu_error)
985 }
986 None => Err(frame_build_error.unwrap_or_else(|| {
987 DeviceResidentArrowError::Solve {
988 reason: "SAE resident frame build declined".to_string(),
989 }
990 })),
991 }
992 }
993 InnerSolveMode::DeviceReupload => {
994 solve_arrow_newton_step(&residual, ridge_t, ridge_beta).map_err(map_gpu_error)
995 }
996 InnerSolveMode::CpuReference => {
997 solve_arrow_newton_step_dense_reference(&residual, ridge_t, ridge_beta)
998 .map_err(|reason| DeviceResidentArrowError::Solve { reason })
999 }
1000 };
1001
1002 let solution = match solution {
1003 Ok(sol) => sol,
1004 Err(DeviceResidentArrowError::Solve { .. })
1005 | Err(DeviceResidentArrowError::Unavailable { .. }) => {
1006 ridge_t = grow_ridge(ridge_t, opts.lm_grow);
1007 ridge_beta = grow_ridge(ridge_beta, opts.lm_grow);
1008 if ridge_t > opts.max_ridge || ridge_beta > opts.max_ridge {
1009 return Err(DeviceResidentArrowError::Solve {
1010 reason: format!(
1011 "SAE outer-sequence inner loop: LM ridge exceeded max ({:e}) at iter {total_iters}",
1012 opts.max_ridge
1013 ),
1014 });
1015 }
1016 total_iters += 1;
1017 continue;
1018 }
1019 Err(other) => return Err(other),
1020 };
1021
1022 let predicted_reduction = crate::arrow_schur::arrow_bare_quadratic_model_reduction(
1023 &residual,
1024 solution.delta_t.view(),
1025 solution.delta_beta.view(),
1026 ridge_t,
1027 ridge_beta,
1028 )
1029 .map_err(|err| DeviceResidentArrowError::Solve {
1030 reason: format!("SAE outer-sequence predicted-reduction failed: {err}"),
1031 })?;
1032
1033 let mut trial_t = t.clone();
1034 let mut trial_beta = beta.clone();
1035 for (slot, dv) in trial_t.iter_mut().zip(solution.delta_t.iter()) {
1036 *slot += *dv;
1037 }
1038 for (slot, dv) in trial_beta.iter_mut().zip(solution.delta_beta.iter()) {
1039 *slot += *dv;
1040 }
1041 let trial_objective =
1042 self.objective_at(base, half_target_energy, &trial_t, &trial_beta);
1043
1044 let objective_scale = current_objective.abs();
1045 let noise_floor = objective_scale * 1e-14;
1046 let actual_reduction = current_objective - trial_objective;
1047 let rho = if predicted_reduction > noise_floor {
1048 actual_reduction / predicted_reduction
1049 } else if actual_reduction >= -noise_floor {
1050 1.0
1051 } else {
1052 -1.0
1053 };
1054
1055 if rho > 0.0 && trial_objective.is_finite() {
1056 t = trial_t;
1057 beta = trial_beta;
1058 current_objective = trial_objective;
1059 ridge_t = (ridge_t * opts.lm_shrink).max(0.0);
1060 ridge_beta = (ridge_beta * opts.lm_shrink).max(0.0);
1061 last_gradient_norm = g_norm;
1062 last_log_det = solution.log_det_hessian;
1063 accepted_iters += 1;
1064 total_iters += 1;
1065 } else {
1066 ridge_t = grow_ridge(ridge_t, opts.lm_grow);
1067 ridge_beta = grow_ridge(ridge_beta, opts.lm_grow);
1068 if ridge_t > opts.max_ridge || ridge_beta > opts.max_ridge {
1069 return Err(DeviceResidentArrowError::Solve {
1070 reason: format!(
1071 "SAE outer-sequence inner loop: LM rejected step until ridge exceeded max ({:e}) at iter {total_iters} (rho={rho:.3e})",
1072 opts.max_ridge
1073 ),
1074 });
1075 }
1076 total_iters += 1;
1077 }
1078 }
1079
1080 Ok(DeviceResidentInnerOutcome {
1081 t: Array1::from_vec(t),
1082 beta: Array1::from_vec(beta),
1083 objective: current_objective,
1084 gradient_norm: last_gradient_norm,
1085 log_det_hessian: last_log_det,
1086 iterations: total_iters,
1087 accepted_iterations: accepted_iters,
1088 converged,
1089 execution_path,
1090 })
1091 }
1092
1093 fn objective_at(
1100 &self,
1101 base: &ArrowSchurSystem,
1102 half_target_energy: f64,
1103 t: &[f64],
1104 beta: &[f64],
1105 ) -> f64 {
1106 let n = self.shape.n;
1107 let d = self.shape.d;
1108 let p = self.shape.p;
1109 let mut quad = 0.0_f64;
1111 let mut lin = 0.0_f64;
1112 for i in 0..n {
1115 let t_base = i * d;
1116 for r in 0..d {
1117 let mut htt_t = 0.0_f64;
1119 for c in 0..d {
1120 htt_t += base.rows[i].htt[[r, c]] * t[t_base + c];
1121 }
1122 let mut htb_b = 0.0_f64;
1124 for c in 0..p {
1125 htb_b += base.rows[i].htbeta[[r, c]] * beta[c];
1126 }
1127 quad += t[t_base + r] * (htt_t + 2.0 * htb_b);
1128 lin += base.rows[i].gt[r] * t[t_base + r];
1129 }
1130 }
1131 for r in 0..p {
1133 let mut hbb_b = 0.0_f64;
1134 for c in 0..p {
1135 hbb_b += base.hbb[[r, c]] * beta[c];
1136 }
1137 quad += beta[r] * hbb_b;
1138 lin += base.gb[r] * beta[r];
1139 }
1140 half_target_energy + 0.5 * quad - lin
1141 }
1142
1143 fn residual_system(
1148 &self,
1149 base: &ArrowSchurSystem,
1150 t: &[f64],
1151 beta: &[f64],
1152 ) -> ArrowSchurSystem {
1153 let n = self.shape.n;
1154 let d = self.shape.d;
1155 let p = self.shape.p;
1156 let mut sys = self.to_arrow_system();
1164 for i in 0..n {
1165 let t_base = i * d;
1166 for r in 0..d {
1167 let mut hz = 0.0_f64;
1168 for c in 0..d {
1169 hz += base.rows[i].htt[[r, c]] * t[t_base + c];
1170 }
1171 for c in 0..p {
1172 hz += base.rows[i].htbeta[[r, c]] * beta[c];
1173 }
1174 sys.rows[i].gt[r] = hz - base.rows[i].gt[r];
1175 }
1176 }
1177 for r in 0..p {
1178 let mut hz = 0.0_f64;
1179 for c in 0..p {
1181 hz += base.hbb[[r, c]] * beta[c];
1182 }
1183 for i in 0..n {
1185 let t_base = i * d;
1186 for rr in 0..d {
1187 hz += base.rows[i].htbeta[[rr, r]] * t[t_base + rr];
1188 }
1189 }
1190 sys.gb[r] = hz - base.gb[r];
1191 }
1192 sys.refresh_row_hessian_fingerprint();
1193 sys
1194 }
1195}
1196
1197#[derive(Clone, Copy, Debug)]
1201pub struct DeviceResidentInnerOptions {
1202 pub max_iterations: usize,
1203 pub convergence_tolerance: f64,
1204 pub initial_ridge_t: f64,
1205 pub initial_ridge_beta: f64,
1206 pub lm_grow: f64,
1207 pub lm_shrink: f64,
1208 pub max_ridge: f64,
1209}
1210
1211impl Default for DeviceResidentInnerOptions {
1212 fn default() -> Self {
1213 Self {
1214 max_iterations: 16,
1215 convergence_tolerance: 1e-9,
1216 initial_ridge_t: 0.0,
1217 initial_ridge_beta: 0.0,
1218 lm_grow: 4.0,
1219 lm_shrink: 0.5,
1220 max_ridge: 1e9,
1221 }
1222 }
1223}
1224
1225#[derive(Clone, Debug)]
1227pub struct DeviceResidentInnerOutcome {
1228 pub t: Array1<f64>,
1229 pub beta: Array1<f64>,
1230 pub objective: f64,
1231 pub gradient_norm: f64,
1232 pub log_det_hessian: f64,
1233 pub iterations: usize,
1234 pub accepted_iterations: usize,
1235 pub converged: bool,
1236 pub execution_path: ExecutionPath,
1237}
1238
1239#[derive(Clone, Debug)]
1250pub struct OuterSequenceOutcome {
1251 pub outers: Vec<DeviceResidentInnerOutcome>,
1252 pub frame_builds: usize,
1253}
1254
1255#[derive(Default)]
1260struct SharedFrameState {
1261 frame: Option<(
1262 f64,
1263 f64,
1264 crate::gpu_kernels::arrow_schur::ResidentArrowFrameHandle,
1265 )>,
1266 frame_builds: usize,
1267}
1268
1269fn note_resident_engagement(engaged: bool, detail: &str) {
1276 use std::sync::Once;
1277 static ENGAGED_ONCE: Once = Once::new();
1278 static DECLINED_ONCE: Once = Once::new();
1279 let once = if engaged {
1280 &ENGAGED_ONCE
1281 } else {
1282 &DECLINED_ONCE
1283 };
1284 once.call_once(|| {
1285 let verdict = if engaged {
1286 "device ENGAGED"
1287 } else {
1288 "device DECLINED - CPU reference"
1289 };
1290 log::warn!("[gam-solve sae_resident inner step] {verdict}: {detail}");
1291 });
1292}
1293
1294fn grow_ridge(current: f64, grow: f64) -> f64 {
1295 if current == 0.0 { 1e-6 } else { current * grow }
1296}
1297
1298fn arrow_system_gradient_norm(sys: &ArrowSchurSystem) -> f64 {
1299 let mut acc = 0.0_f64;
1300 for row in &sys.rows {
1301 for &v in row.gt.iter() {
1302 acc += v * v;
1303 }
1304 }
1305 for &v in sys.gb.iter() {
1306 acc += v * v;
1307 }
1308 acc.sqrt()
1309}
1310
1311fn iterate_norm(t: &[f64], beta: &[f64]) -> f64 {
1312 (squared_norm(t) + squared_norm(beta)).sqrt()
1313}
1314
1315fn validate_shape(
1316 shape: DeviceResidentArrowShape,
1317 target_x: &[f64],
1318 basis_values: &[f64],
1319 gate_activations: &[f64],
1320 slabs: &DeviceResidentArrowSlabs,
1321) -> Result<(), DeviceResidentArrowError> {
1322 let checks = [
1323 ("target_x", target_x.len(), shape.target_len()),
1324 ("basis_values", basis_values.len(), shape.basis_len()),
1325 (
1326 "gate_activations",
1327 gate_activations.len(),
1328 shape.basis_len(),
1329 ),
1330 (
1331 "row_hessian_slabs",
1332 slabs.row_hessian_slabs.len(),
1333 shape.row_hessian_len(),
1334 ),
1335 (
1336 "row_cross_slabs",
1337 slabs.row_cross_slabs.len(),
1338 shape.row_cross_len(),
1339 ),
1340 (
1341 "row_gradient_slabs",
1342 slabs.row_gradient_slabs.len(),
1343 shape.row_gradient_len(),
1344 ),
1345 (
1346 "border_hessian",
1347 slabs.border_hessian.len(),
1348 shape.border_hessian_len(),
1349 ),
1350 ("border_gradient", slabs.border_gradient.len(), shape.p),
1351 ];
1352 for (label, got, want) in checks {
1353 if got != want {
1354 return Err(DeviceResidentArrowError::Shape {
1355 reason: format!(
1356 "SAE resident workspace shape mismatch for {label}: got {got}, expected {want}"
1357 ),
1358 });
1359 }
1360 }
1361 if shape.n == 0 || shape.p == 0 || shape.d == 0 || shape.basis_cols == 0 {
1362 return Err(DeviceResidentArrowError::Shape {
1363 reason: "SAE resident workspace requires nonzero n, p, basis_cols, and d".to_string(),
1364 });
1365 }
1366 Ok(())
1367}
1368
1369#[cfg(target_os = "linux")]
1370fn upload_resident_buffers(
1371 shape: DeviceResidentArrowShape,
1372 target_x: &[f64],
1373 basis_values: &[f64],
1374 gate_activations: &[f64],
1375 slabs: &DeviceResidentArrowSlabs,
1376) -> Option<DeviceResidentArrowBuffers> {
1377 use gam_gpu::linalg_dispatch::{DispatchOp, route_through_gpu};
1378
1379 let runtime = route_through_gpu(DispatchOp::SmallDenseBatchedPotrf {
1380 p: shape.d,
1381 batch: shape.n,
1382 })
1383 .or_else(|| {
1384 route_through_gpu(DispatchOp::Gemm {
1385 m: shape.p,
1386 n: shape.p,
1387 k: shape.n * shape.basis_cols,
1388 })
1389 })?;
1390 let ctx = gam_gpu::device_runtime::cuda_context_for(runtime.device.ordinal)?;
1391 let stream = ctx.new_stream().ok()?;
1392 let target_x_dev = stream.clone_htod(target_x).ok()?;
1393 let basis_values_dev = stream.clone_htod(basis_values).ok()?;
1394 let gate_activations_dev = stream.clone_htod(gate_activations).ok()?;
1395 let row_hessian_dev = stream.clone_htod(&slabs.row_hessian_slabs).ok()?;
1396 let row_cross_dev = stream.clone_htod(&slabs.row_cross_slabs).ok()?;
1397 let row_gradient_dev = stream.clone_htod(&slabs.row_gradient_slabs).ok()?;
1398 let border_hessian_dev = stream.clone_htod(&slabs.border_hessian).ok()?;
1399 let border_gradient_dev = stream.clone_htod(&slabs.border_gradient).ok()?;
1400 let bytes = [
1401 target_x.len(),
1402 basis_values.len(),
1403 gate_activations.len(),
1404 slabs.row_hessian_slabs.len(),
1405 slabs.row_cross_slabs.len(),
1406 slabs.row_gradient_slabs.len(),
1407 slabs.border_hessian.len(),
1408 slabs.border_gradient.len(),
1409 ]
1410 .into_iter()
1411 .sum::<usize>()
1412 * std::mem::size_of::<f64>();
1413 Some(DeviceResidentArrowBuffers {
1414 stream,
1415 target_x_dev,
1416 basis_values_dev,
1417 gate_activations_dev,
1418 row_hessian_dev,
1419 row_cross_dev,
1420 row_gradient_dev,
1421 border_hessian_dev,
1422 border_gradient_dev,
1423 bytes,
1424 })
1425}
1426
1427fn map_gpu_error(err: ArrowSchurGpuFailure) -> DeviceResidentArrowError {
1428 match err {
1429 ArrowSchurGpuFailure::Unavailable => DeviceResidentArrowError::Unavailable {
1430 reason: "SAE resident inner iteration unavailable after GPU admission".to_string(),
1431 },
1432 ArrowSchurGpuFailure::RidgeBumpRequired { row, bump } => DeviceResidentArrowError::Solve {
1433 reason: format!("SAE resident inner iteration row {row} requires ridge bump {bump:e}"),
1434 },
1435 ArrowSchurGpuFailure::SchurFactorFailed { reason } => {
1436 DeviceResidentArrowError::Solve { reason }
1437 }
1438 ArrowSchurGpuFailure::GpuRequiresDenseSystem {
1439 had_hbb_matvec,
1440 had_htbeta_matvec,
1441 } => DeviceResidentArrowError::Solve {
1442 reason: format!(
1443 "SAE resident inner iteration requires dense slabs; hbb_matvec={had_hbb_matvec} htbeta_matvec={had_htbeta_matvec}"
1444 ),
1445 },
1446 }
1447}
1448
1449fn squared_norm(values: &[f64]) -> f64 {
1450 values.iter().map(|v| v * v).sum()
1451}
1452
1453impl From<ArrowSchurError> for DeviceResidentArrowError {
1454 fn from(err: ArrowSchurError) -> Self {
1455 Self::Solve {
1456 reason: err.to_string(),
1457 }
1458 }
1459}
1460
1461pub fn qwen_non_gating_fixture() -> Result<DeviceResidentArrowWorkspace, DeviceResidentArrowError> {
1463 qwen_non_gating_fixture_seeded(0x1017_0003_D3A1_5EED)
1464}
1465
1466pub fn qwen_non_gating_fixture_seeded(
1470 seed: u64,
1471) -> Result<DeviceResidentArrowWorkspace, DeviceResidentArrowError> {
1472 fixture_for_shape_seeded(DeviceResidentArrowShape::qwen_non_gating(), seed)
1473}
1474
1475pub fn color_arm_fixture() -> Result<DeviceResidentArrowWorkspace, DeviceResidentArrowError> {
1480 fixture_for_shape_seeded(DeviceResidentArrowShape::color_arm(), 0x1017_C010_2A12_5EED)
1481}
1482
1483fn fixture_for_shape_seeded(
1488 shape: DeviceResidentArrowShape,
1489 seed: u64,
1490) -> Result<DeviceResidentArrowWorkspace, DeviceResidentArrowError> {
1491 if shape.d == 0 {
1492 return Err(DeviceResidentArrowError::Shape {
1493 reason: "fixture_for_shape_seeded requires d >= 1".to_string(),
1494 });
1495 }
1496 let d = shape.d;
1497 let mut rng = SplitMix64::new(seed);
1498 let mut target_x = vec![0.0_f64; shape.target_len()];
1499 for i in 0..shape.n {
1500 for j in 0..shape.p {
1501 let phase = ((i % 97) as f64) * 0.013 + ((j % 131) as f64) * 0.007;
1502 target_x[i * shape.p + j] = 0.02 * phase.sin() + 0.001 * rng.sample_signed();
1503 }
1504 }
1505 let mut basis_values = vec![0.0_f64; shape.basis_len()];
1506 let mut gate_activations = vec![1.0_f64; shape.basis_len()];
1507 for i in 0..shape.n {
1508 for a in 0..shape.basis_cols {
1509 let phase = ((i + 1) as f64) * ((a + 1) as f64) * 0.003;
1510 basis_values[i * shape.basis_cols + a] = phase.cos();
1511 gate_activations[i * shape.basis_cols + a] = 1.0;
1512 }
1513 }
1514 let mut row_hessian_slabs = vec![0.0_f64; shape.row_hessian_len()];
1515 let mut row_cross_slabs = vec![0.0_f64; shape.row_cross_len()];
1516 let mut row_gradient_slabs = vec![0.0_f64; shape.row_gradient_len()];
1517 for i in 0..shape.n {
1518 let mut basis_sum = 0.0_f64;
1519 for a in 0..shape.basis_cols {
1520 basis_sum +=
1521 basis_values[i * shape.basis_cols + a] * gate_activations[i * shape.basis_cols + a];
1522 }
1523 let h_base = i * d * d;
1526 for r in 0..d {
1527 for c in 0..d {
1528 let v = if r == c {
1529 3.0 + 0.01 * basis_sum.abs() + 0.1 * (r as f64)
1530 } else {
1531 0.02 * (basis_sum + (r + c) as f64).sin() / (d as f64)
1532 };
1533 row_hessian_slabs[h_base + r * d + c] = v;
1534 }
1535 }
1536 for r in 0..d {
1538 for c in 0..r {
1539 let avg = 0.5
1540 * (row_hessian_slabs[h_base + r * d + c]
1541 + row_hessian_slabs[h_base + c * d + r]);
1542 row_hessian_slabs[h_base + r * d + c] = avg;
1543 row_hessian_slabs[h_base + c * d + r] = avg;
1544 }
1545 }
1546 let b_base = i * d * shape.p;
1548 let g_base = i * d;
1549 for r in 0..d {
1550 for j in 0..shape.p {
1551 let feature = ((j % 257) as f64) * 0.011;
1552 row_cross_slabs[b_base + r * shape.p + j] =
1553 1.0e-4 * (basis_sum + r as f64).sin() * feature.cos();
1554 }
1555 row_gradient_slabs[g_base + r] = 0.01 * (basis_sum + r as f64).sin();
1556 }
1557 }
1558 let mut border_hessian = vec![0.0_f64; shape.border_hessian_len()];
1559 for r in 0..shape.p {
1560 border_hessian[r * shape.p + r] = 4.0;
1561 if r + 1 < shape.p {
1562 border_hessian[r * shape.p + r + 1] = 0.01;
1563 border_hessian[(r + 1) * shape.p + r] = 0.01;
1564 }
1565 }
1566 let mut border_gradient = vec![0.0_f64; shape.p];
1567 for j in 0..shape.p {
1568 border_gradient[j] = 0.001 * ((j % 193) as f64 * 0.017).sin();
1569 }
1570 DeviceResidentArrowWorkspace::new(
1571 shape,
1572 target_x,
1573 basis_values,
1574 gate_activations,
1575 DeviceResidentArrowSlabs {
1576 row_hessian_slabs,
1577 row_cross_slabs,
1578 row_gradient_slabs,
1579 border_hessian,
1580 border_gradient,
1581 },
1582 )
1583}
1584
1585pub struct MultiplexedFit {
1587 pub outcome: DeviceResidentInnerOutcome,
1588}
1589
1590pub fn run_resident_fits_multiplexed(
1620 workspaces: Vec<DeviceResidentArrowWorkspace>,
1621 opts: DeviceResidentInnerOptions,
1622) -> Result<Vec<Result<MultiplexedFit, DeviceResidentArrowError>>, String> {
1623 run_resident_fits_multiplexed_with(workspaces, opts, |workspace, opts| {
1624 workspace.device_fit(opts)
1625 })
1626}
1627
1628fn run_resident_fits_multiplexed_with<Run>(
1632 workspaces: Vec<DeviceResidentArrowWorkspace>,
1633 opts: DeviceResidentInnerOptions,
1634 run_one: Run,
1635) -> Result<Vec<Result<MultiplexedFit, DeviceResidentArrowError>>, String>
1636where
1637 Run: Fn(
1638 &DeviceResidentArrowWorkspace,
1639 &DeviceResidentInnerOptions,
1640 ) -> Result<DeviceResidentInnerOutcome, DeviceResidentArrowError>
1641 + Sync,
1642{
1643 let rows = crate::topology_selector::run_topology_race_parallel(
1644 workspaces,
1645 move |workspace: DeviceResidentArrowWorkspace| {
1646 run_one(&workspace, &opts).map(|outcome| MultiplexedFit { outcome })
1647 },
1648 )?;
1649 Ok(rows.into_iter().map(|row| row.result).collect())
1650}
1651
1652pub fn run_resident_fits_sequential(
1657 workspaces: &[DeviceResidentArrowWorkspace],
1658 opts: &DeviceResidentInnerOptions,
1659) -> Vec<Result<MultiplexedFit, DeviceResidentArrowError>> {
1660 workspaces
1661 .iter()
1662 .map(|workspace| {
1663 workspace
1664 .device_fit(opts)
1665 .map(|outcome| MultiplexedFit { outcome })
1666 })
1667 .collect()
1668}
1669
1670#[derive(Clone, Copy, Debug)]
1689pub struct SweepVariant {
1690 pub dim: DeviceResidentArrowShape,
1692 pub seed: u64,
1694}
1695
1696#[derive(Clone, Copy, Debug)]
1698pub struct SweepThroughput {
1699 pub fits: usize,
1700 pub succeeded: usize,
1701 pub wall_seconds: f64,
1702 pub fits_per_second: f64,
1704}
1705
1706pub fn build_sweep_workspaces(
1711 variants: &[SweepVariant],
1712) -> Result<Vec<DeviceResidentArrowWorkspace>, DeviceResidentArrowError> {
1713 variants
1714 .iter()
1715 .map(|v| fixture_for_shape_seeded(v.dim, v.seed))
1716 .collect()
1717}
1718
1719pub fn run_variant_sweep_multiplexed(
1724 variants: &[SweepVariant],
1725 opts: DeviceResidentInnerOptions,
1726) -> Result<
1727 (
1728 Vec<Result<MultiplexedFit, DeviceResidentArrowError>>,
1729 SweepThroughput,
1730 ),
1731 String,
1732> {
1733 let workspaces = build_sweep_workspaces(variants).map_err(|e| e.to_string())?;
1734 run_battery_sweep_multiplexed(workspaces, opts)
1735}
1736
1737pub fn run_battery_sweep_multiplexed(
1748 workspaces: Vec<DeviceResidentArrowWorkspace>,
1749 opts: DeviceResidentInnerOptions,
1750) -> Result<
1751 (
1752 Vec<Result<MultiplexedFit, DeviceResidentArrowError>>,
1753 SweepThroughput,
1754 ),
1755 String,
1756> {
1757 let fits = workspaces.len();
1758 let start = std::time::Instant::now();
1759 let results = run_resident_fits_multiplexed(workspaces, opts)?;
1760 let wall_seconds = start.elapsed().as_secs_f64();
1761 let succeeded = results.iter().filter(|r| r.is_ok()).count();
1762 let throughput = SweepThroughput {
1763 fits,
1764 succeeded,
1765 wall_seconds,
1766 fits_per_second: (fits as f64) / wall_seconds.max(1e-9),
1767 };
1768 Ok((results, throughput))
1769}
1770
1771#[must_use]
1778pub fn color_arm_variant_matrix() -> Vec<SweepVariant> {
1779 let topologies = ["euclidean", "circle", "torus", "sphere"];
1780 let mut variants = Vec::with_capacity(4 * topologies.len() * 2);
1781 for k in 1..=4u64 {
1782 for (t_idx, _topology) in topologies.iter().enumerate() {
1783 for &(d, basis_cols, basis_tag) in &[(2usize, 8usize, 0u64), (1usize, 2usize, 1u64)] {
1785 let mut dim = DeviceResidentArrowShape::color_arm();
1786 dim.d = d;
1787 dim.basis_cols = basis_cols;
1788 let seed = 0x1017_C010_0000_0000 ^ (k << 16) ^ ((t_idx as u64) << 8) ^ basis_tag;
1789 variants.push(SweepVariant { dim, seed });
1790 }
1791 }
1792 }
1793 variants
1794}
1795
1796pub fn assert_sweep_parity_vs_sequential(
1803 variants: &[SweepVariant],
1804 opts: &DeviceResidentInnerOptions,
1805 multiplexed: &[Result<MultiplexedFit, DeviceResidentArrowError>],
1806) -> Result<SweepThroughput, String> {
1807 let workspaces = build_sweep_workspaces(variants).map_err(|e| e.to_string())?;
1808 let start = std::time::Instant::now();
1809 let sequential = run_resident_fits_sequential(&workspaces, opts);
1810 let wall_seconds = start.elapsed().as_secs_f64();
1811 if sequential.len() != multiplexed.len() {
1812 return Err(format!(
1813 "sweep parity: length mismatch seq={} mux={}",
1814 sequential.len(),
1815 multiplexed.len()
1816 ));
1817 }
1818 for (idx, (seq, mux)) in sequential.iter().zip(multiplexed.iter()).enumerate() {
1819 match (seq, mux) {
1820 (Ok(s), Ok(m)) => {
1821 if s.outcome.t.as_slice() != m.outcome.t.as_slice()
1822 || s.outcome.beta.as_slice() != m.outcome.beta.as_slice()
1823 || s.outcome.objective.to_bits() != m.outcome.objective.to_bits()
1824 {
1825 return Err(format!(
1826 "sweep parity: fit {idx} multiplexed result differs from sequential"
1827 ));
1828 }
1829 }
1830 (Err(_), Err(_)) => {}
1831 _ => {
1832 return Err(format!(
1833 "sweep parity: fit {idx} success/failure disagrees seq-vs-mux"
1834 ));
1835 }
1836 }
1837 }
1838 let fits = variants.len();
1839 let succeeded = sequential.iter().filter(|r| r.is_ok()).count();
1840 Ok(SweepThroughput {
1841 fits,
1842 succeeded,
1843 wall_seconds,
1844 fits_per_second: (fits as f64) / wall_seconds.max(1e-9),
1845 })
1846}
1847
1848struct SplitMix64 {
1849 state: u64,
1850}
1851
1852impl SplitMix64 {
1853 const fn new(seed: u64) -> Self {
1854 Self { state: seed }
1855 }
1856
1857 fn next_u64(&mut self) -> u64 {
1858 gam_linalg::utils::splitmix64(&mut self.state)
1859 }
1860
1861 fn sample_signed(&mut self) -> f64 {
1862 let unit = (self.next_u64() >> 11) as f64 / ((1_u64 << 53) as f64);
1863 2.0 * unit - 1.0
1864 }
1865}
1866
1867#[cfg(test)]
1868mod tests {
1869 use super::*;
1870 use ndarray::Array2;
1871
1872 fn small_fixture(seed: u64) -> DeviceResidentArrowWorkspace {
1876 let shape = DeviceResidentArrowShape {
1884 n: 8,
1885 p: 4,
1886 basis_cols: 2,
1887 d: 2,
1888 };
1889 let mut rng = SplitMix64::new(seed);
1890 let target_x = vec![0.0_f64; shape.target_len()];
1891 let basis_values = vec![0.5_f64; shape.basis_len()];
1892 let gate_activations = vec![1.0_f64; shape.basis_len()];
1893
1894 let mut row_hessian_slabs = vec![0.0_f64; shape.row_hessian_len()];
1895 let mut row_cross_slabs = vec![0.0_f64; shape.row_cross_len()];
1896 let mut row_gradient_slabs = vec![0.0_f64; shape.row_gradient_len()];
1897 for i in 0..shape.n {
1898 let h = i * shape.d * shape.d;
1899 row_hessian_slabs[h] = 5.0 + 0.1 * rng.sample_signed();
1900 row_hessian_slabs[h + 1] = 0.05 * rng.sample_signed();
1901 row_hessian_slabs[h + 2] = row_hessian_slabs[h + 1];
1902 row_hessian_slabs[h + 3] = 4.0 + 0.1 * rng.sample_signed();
1903 let b = i * shape.d * shape.p;
1904 for j in 0..shape.p {
1905 row_cross_slabs[b + j] = 0.01 * rng.sample_signed();
1906 row_cross_slabs[b + shape.p + j] = 0.01 * rng.sample_signed();
1907 }
1908 let g = i * shape.d;
1909 row_gradient_slabs[g] = rng.sample_signed();
1910 row_gradient_slabs[g + 1] = rng.sample_signed();
1911 }
1912 let mut border_hessian = vec![0.0_f64; shape.border_hessian_len()];
1913 for r in 0..shape.p {
1914 border_hessian[r * shape.p + r] = 6.0 + 0.1 * rng.sample_signed();
1915 }
1916 let border_gradient: Vec<f64> = (0..shape.p).map(|_| rng.sample_signed()).collect();
1917
1918 DeviceResidentArrowWorkspace::new(
1919 shape,
1920 target_x,
1921 basis_values,
1922 gate_activations,
1923 DeviceResidentArrowSlabs {
1924 row_hessian_slabs,
1925 row_cross_slabs,
1926 row_gradient_slabs,
1927 border_hessian,
1928 border_gradient,
1929 },
1930 )
1931 .expect("small resident fixture must validate")
1932 }
1933
1934 fn dense_hz(
1937 ws: &DeviceResidentArrowWorkspace,
1938 sys: &ArrowSchurSystem,
1939 ) -> (Array2<f64>, Array1<f64>) {
1940 let shape = ws.shape;
1941 let total = shape.n * shape.d + shape.p;
1942 let mut h = Array2::<f64>::zeros((total, total));
1943 let mut g0 = Array1::<f64>::zeros(total);
1944 for i in 0..shape.n {
1945 let base = i * shape.d;
1946 for r in 0..shape.d {
1947 for c in 0..shape.d {
1948 h[[base + r, base + c]] = sys.rows[i].htt[[r, c]];
1949 }
1950 for c in 0..shape.p {
1951 let v = sys.rows[i].htbeta[[r, c]];
1952 h[[base + r, shape.n * shape.d + c]] = v;
1953 h[[shape.n * shape.d + c, base + r]] = v;
1954 }
1955 g0[base + r] = sys.rows[i].gt[r];
1956 }
1957 }
1958 for r in 0..shape.p {
1959 for c in 0..shape.p {
1960 h[[shape.n * shape.d + r, shape.n * shape.d + c]] = sys.hbb[[r, c]];
1961 }
1962 g0[shape.n * shape.d + r] = sys.gb[r];
1963 }
1964 (h, g0)
1965 }
1966
1967 #[test]
1968 fn cpu_inner_loop_reaches_quadratic_minimiser() {
1969 let ws = small_fixture(0xABCD_0001);
1970 let opts = DeviceResidentInnerOptions::default();
1971 let outcome = ws.cpu_reference_fit(&opts).expect("cpu fit");
1972 assert!(
1973 outcome.converged,
1974 "inner loop must converge on a PD quadratic"
1975 );
1976
1977 let base = ws.to_arrow_system();
1979 let (h, g0) = dense_hz(&ws, &base);
1980 let total = ws.shape.n * ws.shape.d + ws.shape.p;
1981 let mut z = Array1::<f64>::zeros(total);
1982 for r in 0..ws.shape.n * ws.shape.d {
1983 z[r] = outcome.t[r];
1984 }
1985 for c in 0..ws.shape.p {
1986 z[ws.shape.n * ws.shape.d + c] = outcome.beta[c];
1987 }
1988 let hz = h.dot(&z);
1989 let mut max_resid = 0.0_f64;
1990 for r in 0..total {
1991 max_resid = max_resid.max((hz[r] - g0[r]).abs());
1992 }
1993 assert!(
1994 max_resid < 1e-9,
1995 "inner loop fixed point must solve H z = g0; residual {max_resid:e}"
1996 );
1997 }
1998
1999 #[test]
2000 fn cpu_multiplex_matches_sequential_bit_identical() {
2001 let seeds = [0x11, 0x22, 0x33, 0x44, 0x55, 0x66];
2002 let opts = DeviceResidentInnerOptions::default();
2003
2004 let seq_workspaces: Vec<_> = seeds.iter().map(|&s| small_fixture(s)).collect();
2005 let sequential: Vec<_> = seq_workspaces
2006 .iter()
2007 .map(|ws| ws.cpu_reference_fit(&opts).expect("seq cpu fit"))
2008 .collect();
2009
2010 let mux_workspaces: Vec<_> = seeds.iter().map(|&s| small_fixture(s)).collect();
2011 let multiplexed = run_resident_fits_multiplexed_with(mux_workspaces, opts, |ws, opts| {
2012 ws.cpu_reference_fit(opts)
2013 })
2014 .expect("multiplexed cpu fits");
2015
2016 assert_eq!(sequential.len(), multiplexed.len());
2017 for (seq, mux) in sequential.iter().zip(multiplexed.iter()) {
2018 let mux = mux.as_ref().expect("mux fit ok");
2019 assert_eq!(seq.t.as_slice(), mux.outcome.t.as_slice());
2022 assert_eq!(seq.beta.as_slice(), mux.outcome.beta.as_slice());
2023 assert_eq!(seq.objective.to_bits(), mux.outcome.objective.to_bits());
2024 }
2025 }
2026
2027 #[test]
2036 fn device_resident_fit_matches_cpu_reference() {
2037 let ws = small_fixture(0x5AE_1017);
2038 let opts = DeviceResidentInnerOptions::default();
2039
2040 let cpu = ws.cpu_reference_fit(&opts).expect("cpu reference fit");
2042 assert!(cpu.converged, "cpu reference must converge on PD quadratic");
2043
2044 let base = ws.to_arrow_system();
2045
2046 println!(
2047 "DIAG_RESIDENT device_resident={} shape=({},{},{})",
2048 ws.device_resident(),
2049 ws.shape.n,
2050 ws.shape.d,
2051 ws.shape.p
2052 );
2053 if ws.device_resident() {
2054 let dev = ws.device_fit(&opts).expect("device resident fit");
2056 assert_eq!(
2057 dev.execution_path,
2058 ExecutionPath::GpuResidentFull,
2059 "device_fit must report the full device-resident execution path"
2060 );
2061 assert!(dev.converged, "device resident loop must converge");
2062
2063 let t_scale = cpu.t.iter().fold(1.0_f64, |m, &v| m.max(v.abs()));
2069 let b_scale = cpu.beta.iter().fold(1.0_f64, |m, &v| m.max(v.abs()));
2070 let mut max_rel = 0.0_f64;
2071 for (a, b) in dev.t.iter().zip(cpu.t.iter()) {
2072 max_rel = max_rel.max((a - b).abs() / t_scale);
2073 }
2074 for (a, b) in dev.beta.iter().zip(cpu.beta.iter()) {
2075 max_rel = max_rel.max((a - b).abs() / b_scale);
2076 }
2077 assert!(
2078 max_rel < 1e-9,
2079 "resident device fit must match CPU reference (rel {max_rel:e})"
2080 );
2081
2082 let one = ws
2088 .one_inner_iteration(opts.initial_ridge_t, opts.initial_ridge_beta)
2089 .expect("resident one_inner_iteration");
2090 assert_eq!(
2091 one.execution_path,
2092 ExecutionPath::GpuResidentLinearization,
2093 "one_inner_iteration must report resident single-linearization residency"
2094 );
2095
2096 match crate::gpu_kernels::arrow_schur::ResidentArrowFrameHandle::new(
2104 &base,
2105 opts.initial_ridge_t,
2106 opts.initial_ridge_beta,
2107 ) {
2108 Err(err) => panic!("resident frame must build on CUDA host: {err:?}"),
2109 Ok(frame) => {
2110 let g_t: Vec<f64> = base
2111 .rows
2112 .iter()
2113 .flat_map(|r| r.gt.iter().copied())
2114 .collect();
2115 let g_beta: Vec<f64> = base.gb.iter().copied().collect();
2116 let resident_sol = frame
2117 .solve_gradient(&g_t, &g_beta)
2118 .expect("resident single-gradient solve");
2119 let full =
2120 crate::gpu_kernels::arrow_schur::solve_arrow_newton_step_dense_reference(
2121 &base,
2122 opts.initial_ridge_t,
2123 opts.initial_ridge_beta,
2124 )
2125 .expect("dense reference single solve");
2126 let mut max_step_rel = 0.0_f64;
2127 let step_scale = full
2128 .delta_t
2129 .iter()
2130 .chain(full.delta_beta.iter())
2131 .fold(1.0_f64, |m, &v| m.max(v.abs()));
2132 for (a, b) in resident_sol.delta_t.iter().zip(full.delta_t.iter()) {
2133 max_step_rel = max_step_rel.max((a - b).abs() / step_scale);
2134 }
2135 for (a, b) in resident_sol.delta_beta.iter().zip(full.delta_beta.iter()) {
2136 max_step_rel = max_step_rel.max((a - b).abs() / step_scale);
2137 }
2138 assert!(
2139 max_step_rel < 1e-9,
2140 "resident solve_gradient must match full dense reference step \
2141 (rel {max_step_rel:e})"
2142 );
2143 }
2144 }
2145
2146 let reup = ws
2149 .device_reupload_fit(&opts)
2150 .expect("device re-uploading fit");
2151 assert_eq!(
2152 reup.execution_path,
2153 ExecutionPath::GpuReupload,
2154 "device_reupload_fit must report the re-uploading device path"
2155 );
2156 assert!(reup.converged, "re-uploading loop must converge");
2157 let mut max_reup_rel = 0.0_f64;
2158 for (a, b) in reup.t.iter().zip(cpu.t.iter()) {
2159 max_reup_rel = max_reup_rel.max((a - b).abs() / t_scale);
2160 }
2161 for (a, b) in reup.beta.iter().zip(cpu.beta.iter()) {
2162 max_reup_rel = max_reup_rel.max((a - b).abs() / b_scale);
2163 }
2164 assert!(
2165 max_reup_rel < 1e-9,
2166 "re-uploading GPU fit must match CPU reference (rel {max_reup_rel:e})"
2167 );
2168 } else {
2169 assert!(
2177 gam_gpu::device_runtime::GpuRuntime::resolve(gam_gpu::GpuPolicy::Auto)
2178 .unwrap_or_else(|error| {
2179 panic!("GPU probe fault in resident SAE engagement test: {error}")
2180 })
2181 .is_none(),
2182 "device_resident() is false on a host WITH a CUDA runtime present, \
2183 despite a floor-clearing fixture (batch=8): the resident device \
2184 buffers failed to bind — a real device fault, not a CPU-only skip."
2185 );
2186 let dev = ws.device_fit(&opts);
2188 assert!(
2189 matches!(dev, Err(DeviceResidentArrowError::Unavailable { .. })),
2190 "device_fit must report Unavailable on a CPU-only host, got {dev:?}"
2191 );
2192 let reup = ws.device_reupload_fit(&opts);
2193 assert!(
2194 matches!(reup, Err(DeviceResidentArrowError::Unavailable { .. })),
2195 "device_reupload_fit must report Unavailable on a CPU-only host, got {reup:?}"
2196 );
2197 let frame = crate::gpu_kernels::arrow_schur::ResidentArrowFrameHandle::new(
2198 &base,
2199 opts.initial_ridge_t,
2200 opts.initial_ridge_beta,
2201 );
2202 assert!(
2203 frame.is_err(),
2204 "resident frame construction must decline on a CPU-only host"
2205 );
2206 }
2207 }
2208
2209 #[test]
2227 fn resident_inner_solve_matches_production_arrow_core() {
2228 use crate::arrow_schur::{ArrowSolveOptions, solve_arrow_newton_step_core};
2229
2230 let ws = small_fixture(0x1017_F17);
2231 let opts = DeviceResidentInnerOptions::default();
2232
2233 let resident = ws.cpu_reference_fit(&opts).expect("resident cpu fit");
2235 assert!(
2236 resident.converged,
2237 "resident reference must converge on the PD quadratic"
2238 );
2239
2240 let sys = ws.to_arrow_system();
2244 let (delta_t, delta_beta, _diag) = solve_arrow_newton_step_core(
2245 &sys,
2246 opts.initial_ridge_t,
2247 opts.initial_ridge_beta,
2248 &ArrowSolveOptions::direct(),
2249 )
2250 .expect("production arrow-core solve");
2251
2252 let t_scale = resident.t.iter().fold(1.0_f64, |m, &v| m.max(v.abs()));
2257 let b_scale = resident.beta.iter().fold(1.0_f64, |m, &v| m.max(v.abs()));
2258 let mut max_rel_t = 0.0_f64;
2264 let mut worst_t: Option<(usize, f64, f64)> = None;
2265 for (i, (prod, res)) in delta_t.iter().zip(resident.t.iter()).enumerate() {
2266 let rel = (prod + res).abs() / t_scale;
2267 if rel > max_rel_t {
2268 max_rel_t = rel;
2269 worst_t = Some((i, *prod, *res));
2270 }
2271 }
2272 let mut max_rel_b = 0.0_f64;
2273 let mut worst_b: Option<(usize, f64, f64)> = None;
2274 for (i, (prod, res)) in delta_beta.iter().zip(resident.beta.iter()).enumerate() {
2275 let rel = (prod + res).abs() / b_scale;
2276 if rel > max_rel_b {
2277 max_rel_b = rel;
2278 worst_b = Some((i, *prod, *res));
2279 }
2280 }
2281 let max_rel = max_rel_t.max(max_rel_b);
2282 assert!(
2283 max_rel < 1e-9,
2284 "production arrow-core Newton step must be −(resident converged fit) on \
2285 the same quadratic; wiring the device seam into the SAE inner loop must \
2286 not change the system being solved. rel_t={max_rel_t:e} (worst {worst_t:?}: \
2287 Δt+t* must be 0), rel_beta={max_rel_b:e} (worst {worst_b:?}: Δβ+β* must \
2288 be 0). A t-only gap implicates the per-row factor / row-gradient \
2289 assembly; a β-only gap the border Schur path."
2290 );
2291 }
2292
2293 #[test]
2301 fn outer_sequence_reuses_frame_and_matches_independent() {
2302 let ws = super::color_arm_fixture().expect("color_arm fixture");
2303 let opts = DeviceResidentInnerOptions::default();
2304 let n = ws.shape.n;
2305 let d = ws.shape.d;
2306 let p = ws.shape.p;
2307
2308 let outers: Vec<(Vec<f64>, Vec<f64>)> = (0..3)
2312 .map(|s| {
2313 let g_t: Vec<f64> = (0..n * d)
2314 .map(|i| 0.01 * (((i + 3 * s) as f64) * 0.002).sin())
2315 .collect();
2316 let g_beta: Vec<f64> = (0..p)
2317 .map(|j| 0.001 * (((j + 11 * s) as f64) * 0.0009).cos())
2318 .collect();
2319 (g_t, g_beta)
2320 })
2321 .collect();
2322
2323 let independent = ws
2326 .cpu_reference_outer_sequence(&outers, &opts)
2327 .expect("cpu reference outer sequence");
2328 assert_eq!(independent.outers.len(), outers.len());
2329
2330 if ws.device_resident() {
2331 let shared = ws
2333 .device_fit_outer_sequence(&outers, &opts)
2334 .expect("device outer sequence");
2335 assert_eq!(
2336 shared.frame_builds,
2337 1,
2338 "across-outer residency must build the resident frame exactly once \
2339 for an unchanged operator (got {} builds over {} outers) — a count \
2340 > 1 means the frame was needlessly re-factored per outer",
2341 shared.frame_builds,
2342 outers.len()
2343 );
2344 for (idx, (sh, ind)) in shared
2347 .outers
2348 .iter()
2349 .zip(independent.outers.iter())
2350 .enumerate()
2351 {
2352 let scale = ind
2353 .t
2354 .iter()
2355 .chain(ind.beta.iter())
2356 .fold(1.0_f64, |m, &v| m.max(v.abs()));
2357 let mut max_rel = 0.0_f64;
2358 for (a, b) in sh.t.iter().zip(ind.t.iter()) {
2359 max_rel = max_rel.max((a - b).abs() / scale);
2360 }
2361 for (a, b) in sh.beta.iter().zip(ind.beta.iter()) {
2362 max_rel = max_rel.max((a - b).abs() / scale);
2363 }
2364 assert!(
2365 max_rel < 1e-9,
2366 "outer {idx}: across-outer-shared frame must match independent fit \
2367 (rel {max_rel:e})"
2368 );
2369 }
2370 println!(
2371 "[#1017 outer-seq color_arm] outers={} frame_builds={} (across-outer factor \
2372 amortized) parity<1e-9 OK",
2373 outers.len(),
2374 shared.frame_builds
2375 );
2376 } else {
2377 println!(
2378 "[#1017 outer-seq color_arm] no CUDA device — across-outer residency skipped; \
2379 run on the GPU node to assert frame_builds==1 + device parity"
2380 );
2381 }
2382 }
2383
2384 #[test]
2405 fn gpu_residency_per_solve_bench() {
2406 use std::time::Instant;
2407 const N_SOLVES: usize = 24;
2408 for (label, ws) in [
2409 ("color_arm", super::color_arm_fixture()),
2410 ("qwen_non_gating", super::qwen_non_gating_fixture()),
2411 ] {
2412 let ws = ws.expect("bench fixture must validate");
2413 let base = ws.to_arrow_system();
2414 let n = ws.shape.n;
2419 let d = ws.shape.d;
2420 let p = ws.shape.p;
2421 let gradients: Vec<(Vec<f64>, Vec<f64>)> = (0..N_SOLVES)
2422 .map(|s| {
2423 let g_t: Vec<f64> =
2424 (0..n * d).map(|i| ((i + s) as f64 * 0.001).sin()).collect();
2425 let g_beta: Vec<f64> = (0..p)
2426 .map(|j| ((j + 7 * s) as f64 * 0.0007).cos())
2427 .collect();
2428 (g_t, g_beta)
2429 })
2430 .collect();
2431
2432 if !ws.device_resident() {
2433 println!(
2434 "[#1017 per-solve {label}] no CUDA device — {N_SOLVES} solves skipped; \
2435 run on the GPU node for the across-iteration residency speedup"
2436 );
2437 continue;
2438 }
2439
2440 let t_build = Instant::now();
2444 match crate::gpu_kernels::arrow_schur::ResidentArrowFrameHandle::new(&base, 0.0, 0.0) {
2449 Err(err) => panic!("resident frame must build on CUDA host: {err:?}"),
2450 Ok(frame) => {
2451 let frame_build_ms = t_build.elapsed().as_secs_f64() * 1e3;
2452
2453 frame
2460 .solve_gradient(&gradients[0].0, &gradients[0].1)
2461 .expect("resident warm-up solve");
2462 {
2463 let mut sys = ws.to_arrow_system();
2464 for (i, row) in sys.rows.iter_mut().enumerate() {
2465 for r in 0..d {
2466 row.gt[r] = gradients[0].0[i * d + r];
2467 }
2468 }
2469 for (j, gb) in sys.gb.iter_mut().enumerate() {
2470 *gb = gradients[0].1[j];
2471 }
2472 sys.refresh_row_hessian_fingerprint();
2473 crate::gpu_kernels::arrow_schur::solve_arrow_newton_step(&sys, 0.0, 0.0)
2474 .expect("reupload warm-up solve");
2475 }
2476
2477 let t_res = Instant::now();
2482 let mut resident_steps = Vec::with_capacity(N_SOLVES);
2483 for (g_t, g_beta) in &gradients {
2484 resident_steps.push(
2485 frame
2486 .solve_gradient(g_t, g_beta)
2487 .expect("resident solve_gradient"),
2488 );
2489 }
2490 let resident_ms = t_res.elapsed().as_secs_f64() * 1e3;
2491
2492 let t_reup = Instant::now();
2494 let mut reupload_steps = Vec::with_capacity(N_SOLVES);
2495 for (g_t, g_beta) in &gradients {
2496 let mut sys = ws.to_arrow_system();
2497 for (i, row) in sys.rows.iter_mut().enumerate() {
2498 for r in 0..d {
2499 row.gt[r] = g_t[i * d + r];
2500 }
2501 }
2502 for (j, gb) in sys.gb.iter_mut().enumerate() {
2503 *gb = g_beta[j];
2504 }
2505 sys.refresh_row_hessian_fingerprint();
2506 reupload_steps.push(
2507 crate::gpu_kernels::arrow_schur::solve_arrow_newton_step(
2508 &sys, 0.0, 0.0,
2509 )
2510 .expect("reupload solve_arrow_newton_step"),
2511 );
2512 }
2513 let reupload_ms = t_reup.elapsed().as_secs_f64() * 1e3;
2514
2515 let mut max_rel = 0.0_f64;
2518 for (rs, us) in resident_steps.iter().zip(reupload_steps.iter()) {
2519 let scale = us
2520 .delta_t
2521 .iter()
2522 .chain(us.delta_beta.iter())
2523 .fold(1.0_f64, |m, &v| m.max(v.abs()));
2524 for (a, b) in rs.delta_t.iter().zip(us.delta_t.iter()) {
2525 max_rel = max_rel.max((a - b).abs() / scale);
2526 }
2527 for (a, b) in rs.delta_beta.iter().zip(us.delta_beta.iter()) {
2528 max_rel = max_rel.max((a - b).abs() / scale);
2529 }
2530 }
2531
2532 let resident_per_solve = resident_ms / N_SOLVES as f64;
2533 let reupload_per_solve = reupload_ms / N_SOLVES as f64;
2534 let residency_speedup = reupload_ms / resident_ms.max(1e-9);
2535 println!(
2536 "[#1017 per-solve {label}] N={N_SOLVES} frame_build={frame_build_ms:.2}ms \
2537 resident={resident_ms:.2}ms ({resident_per_solve:.3}ms/solve, \
2538 grad-upload + warm factors) reupload={reupload_ms:.2}ms \
2539 ({reupload_per_solve:.3}ms/solve, N factors + N D/B uploads) \
2540 residency_speedup={residency_speedup:.2}x parity_rel={max_rel:e}"
2541 );
2542 assert!(
2543 max_rel < 1e-9,
2544 "{label}: resident per-solve steps must match reupload (rel {max_rel:e})"
2545 );
2546
2547 let min_speedup = if label == "color_arm" { 1.5 } else { 1.0 };
2563 assert!(
2564 residency_speedup > min_speedup,
2565 "{label}: across-iteration residency must beat per-solve re-upload \
2566 (residency_speedup={residency_speedup:.3}x, required >{min_speedup}x; \
2567 resident {resident_per_solve:.3}ms/solve vs reupload \
2568 {reupload_per_solve:.3}ms/solve over N={N_SOLVES} solves) — the resident \
2569 frame either silently re-uploaded D/B or the dispatch dropped the \
2570 amortized factor path"
2571 );
2572 }
2573 }
2574 }
2575 }
2576
2577 fn battery_variant_matrix() -> Vec<super::SweepVariant> {
2582 let mut variants = Vec::new();
2583 for k in 1..=4u64 {
2586 for basis_cols in [4usize, 8, 12] {
2587 let mut dim = DeviceResidentArrowShape::color_arm();
2588 dim.basis_cols = basis_cols;
2589 variants.push(super::SweepVariant {
2590 dim,
2591 seed: 0x1017_0040_0000_0000 ^ (k << 8) ^ (basis_cols as u64),
2592 });
2593 }
2594 }
2595 variants
2596 }
2597
2598 fn battery_variant_matrix_cpu_gate() -> Vec<super::SweepVariant> {
2608 battery_variant_matrix()
2609 .into_iter()
2610 .map(|mut variant| {
2611 variant.dim.p = 256;
2612 variant
2613 })
2614 .collect()
2615 }
2616
2617 #[test]
2621 fn variant_sweep_multiplex_matches_sequential() {
2622 let variants = battery_variant_matrix_cpu_gate();
2623 let opts = DeviceResidentInnerOptions::default();
2624
2625 let workspaces =
2628 super::build_sweep_workspaces(&variants).expect("sweep workspaces must build");
2629 let multiplexed =
2630 super::run_resident_fits_multiplexed_with(workspaces, opts, |ws, opts| {
2631 ws.cpu_reference_fit(opts)
2632 })
2633 .expect("multiplexed cpu sweep");
2634
2635 let seq_workspaces =
2636 super::build_sweep_workspaces(&variants).expect("sweep workspaces must build");
2637 let sequential: Vec<_> = seq_workspaces
2638 .iter()
2639 .map(|ws| ws.cpu_reference_fit(&opts))
2640 .collect();
2641
2642 assert_eq!(multiplexed.len(), sequential.len());
2643 for (idx, (mux, seq)) in multiplexed.iter().zip(sequential.iter()).enumerate() {
2644 let mux = &mux.as_ref().unwrap().outcome;
2645 let seq = seq.as_ref().unwrap();
2646 assert_eq!(
2647 mux.t.as_slice(),
2648 seq.t.as_slice(),
2649 "variant {idx}: multiplexed t differs from sequential"
2650 );
2651 assert_eq!(
2652 mux.beta.as_slice(),
2653 seq.beta.as_slice(),
2654 "variant {idx}: multiplexed beta differs from sequential"
2655 );
2656 assert_eq!(
2657 mux.objective.to_bits(),
2658 seq.objective.to_bits(),
2659 "variant {idx}: multiplexed objective differs from sequential"
2660 );
2661 }
2662 }
2663
2664 #[test]
2670 fn gpu_multiplex_throughput_bench() {
2671 let variants = battery_variant_matrix();
2672 let opts = DeviceResidentInnerOptions::default();
2673
2674 let probe = super::build_sweep_workspaces(&variants).expect("sweep workspaces");
2675 let any_device = probe.iter().any(|w| w.device_resident());
2676 if !any_device {
2677 println!(
2678 "[#1017 mux-bench] no CUDA device — {} variants (K1..4 x 3 basis) \
2679 skipped; run on the GPU node for cross-fit throughput",
2680 variants.len()
2681 );
2682 return;
2683 }
2684
2685 let (results, mux_tp) =
2686 super::run_variant_sweep_multiplexed(&variants, opts).expect("multiplexed sweep");
2687 let seq_tp = super::assert_sweep_parity_vs_sequential(&variants, &opts, &results)
2688 .expect("sweep parity vs sequential must hold");
2689 println!(
2690 "[#1017 mux-bench] fits={} succeeded={} multiplexed={:.3}s ({:.1} fits/s) \
2691 sequential={:.3}s ({:.1} fits/s) cross-fit-speedup={:.2}x",
2692 mux_tp.fits,
2693 mux_tp.succeeded,
2694 mux_tp.wall_seconds,
2695 mux_tp.fits_per_second,
2696 seq_tp.wall_seconds,
2697 seq_tp.fits_per_second,
2698 mux_tp.fits_per_second / seq_tp.fits_per_second.max(1e-9),
2699 );
2700 assert_eq!(
2701 mux_tp.succeeded, mux_tp.fits,
2702 "all battery variants must fit successfully on device"
2703 );
2704 }
2705}