Skip to main content

gam_gpu/
policy.rs

1use serde::{Deserialize, Serialize};
2
3#[derive(Clone, Copy, Debug, Eq, PartialEq, Serialize, Deserialize)]
4pub enum GpuMixedPrecisionPolicy {
5    /// Always use fp64 factorization; no refinement attempted.
6    Off,
7    /// Attempt fp32 Cholesky factorization followed by up to
8    /// `REFINEMENT_MAX_STEPS` fp64-residual refinement steps. Policy admits
9    /// the attempt only when `p ≥ REFINEMENT_MIN_P` (so that the fp64 GEMV
10    /// overhead is amortized) and the measured residual drops monotonically.
11    /// Falls back to fp64 factorization automatically when the residual does
12    /// not decrease (κ(A)·u ≥ 1 regime) or when the fp32 POTRF itself fails.
13    Refinement,
14    /// Always use fp64 factorization; equivalent to `Off` but signals that
15    /// an explicit policy decision was taken.
16    Never,
17}
18
19#[derive(Clone, Debug, Eq, PartialEq, Serialize, Deserialize)]
20pub struct GpuDispatchPolicy {
21    pub xtwx_n_min: usize,
22    pub xtwx_flops_min: usize,
23    pub xtwx_use_fused_below_p: usize,
24    pub gemm_min_flops: usize,
25    pub potrf_min_p: usize,
26    pub small_dense_batched_potrf_max_p: usize,
27    pub small_dense_batched_potrf_min_batch: usize,
28    pub syevd_min_p: usize,
29    pub sparse_min_nnz: usize,
30    pub fused_kernel_min_n: usize,
31    pub keep_design_resident_min_bytes: usize,
32    pub prefer_gpu_factorization_min_p: usize,
33    pub row_kernel_min_n: usize,
34    pub mixed_precision: GpuMixedPrecisionPolicy,
35}
36
37impl Default for GpuDispatchPolicy {
38    /// Conservative seed thresholds used before device calibration and when
39    /// calibration cannot run on the current host.
40    ///
41    /// The production runtime replaces these with
42    /// [`crate::calibration::calibrated_policy_for_device`] after the CUDA
43    /// probe selects a concrete device. Keep these values conservative: they
44    /// are the typed baseline for CPU-only builds, failed calibration, and unit
45    /// tests that exercise policy predicates without initializing CUDA.
46    fn default() -> Self {
47        Self {
48            xtwx_n_min: 50_000,
49            xtwx_flops_min: 100_000_000,
50            xtwx_use_fused_below_p: 256,
51            gemm_min_flops: 100_000_000,
52            potrf_min_p: 512,
53            small_dense_batched_potrf_max_p: 32,
54            small_dense_batched_potrf_min_batch: 8,
55            syevd_min_p: 256,
56            sparse_min_nnz: 1_000_000,
57            fused_kernel_min_n: 100_000,
58            keep_design_resident_min_bytes: 32 * 1024 * 1024,
59            prefer_gpu_factorization_min_p: 512,
60            row_kernel_min_n: 50_000,
61            mixed_precision: GpuMixedPrecisionPolicy::Refinement,
62        }
63    }
64}
65
66impl GpuDispatchPolicy {
67    /// The smallest `gemm_min_flops` ANY production dispatch policy can carry.
68    ///
69    /// Production policies are exactly two: [`Self::default`] (seed,
70    /// `gemm_min_flops = 1e8`) and the device-calibrated policy, whose
71    /// `crossover_flops` can lower the floor at most to the flop count of the
72    /// smallest calibration measurement — the 64×64×64 GEMM in
73    /// `calibration::GEMM_DIMS`, i.e. `2·64³ = 524_288` (a compile-time assert
74    /// in `calibration.rs` pins the correspondence). Work below this floor is
75    /// therefore inadmissible for GPU dispatch under EVERY reachable policy, so
76    /// a caller may refuse it BEFORE probing the device — this is the pre-probe
77    /// size gate that lets CPU-sized problems skip CUDA context creation
78    /// entirely (the startup-tax ordering fix). Work at or above it must fall
79    /// through to the probed runtime's real (possibly calibrated) policy gate,
80    /// so genuinely GPU-sized problems behave exactly as before.
81    pub const MIN_CALIBRATABLE_GEMM_FLOPS: u128 = 524_288;
82
83    /// The smallest `potrf_min_p` ANY production dispatch policy can carry:
84    /// the smallest POTRF calibration dimension (`calibration::POTRF_DIMS[0]`,
85    /// pinned by a compile-time assert there). A single (batch ≤ 1) POTRF with
86    /// `p` below this is inadmissible under every reachable policy.
87    pub const MIN_CALIBRATABLE_POTRF_P: usize = 64;
88
89    /// The smallest `row_kernel_min_n` / `xtwx_n_min` ANY production dispatch
90    /// policy can carry: the smallest XtWX calibration row count
91    /// (`calibration::XTWX_DIMS[0].0`, pinned by a compile-time assert there).
92    /// A row-kernel workload with fewer rows is inadmissible under every
93    /// reachable policy, so per-fit GPU-eligibility deciders may refuse it
94    /// BEFORE probing the device.
95    pub const MIN_CALIBRATABLE_ROW_KERNEL_N: usize = 2_048;
96
97    /// Minimum problem dimension for the fp32+refinement path.
98    ///
99    /// Below this threshold the fp64 GEMV needed for the residual check costs
100    /// more than the savings from fp32 factorization. The threshold is set so
101    /// that a single `p × p` DGEMV (2p² flops) is at least 10× cheaper than
102    /// the `p³/3` POTRF (i.e. p ≥ 64) while still leaving margin for the
103    /// POTRF/POTRS launches. In practice `p ≥ 64` matches the existing
104    /// `potrf_min_p = 512` floor for GPU dispatch, so the refinement path only
105    /// activates when the GPU factorization path is already chosen.
106    pub const REFINEMENT_MIN_P: usize = 64;
107
108    /// Maximum number of fp32-correction steps per solve.
109    ///
110    /// Two steps suffice for κ(A) ≤ 10⁵ at fp32 (u ≈ 6 × 10⁻⁸): after step
111    /// 1 the error is O(κ u)² ≈ 10⁻⁶, after step 2 it is O(κ u)⁴ ≈ 10⁻¹²,
112    /// which is well within the fp64 unit roundoff of 10⁻¹⁶ × κ. A cap of 3
113    /// is used defensively.
114    pub const REFINEMENT_MAX_STEPS: usize = 3;
115
116    /// Relative residual tolerance for declaring convergence.
117    ///
118    /// `‖r‖ / ‖b‖ ≤ tol` is considered a converged solve. 10⁻¹² is two
119    /// orders of magnitude above the fp64 machine epsilon times a moderate
120    /// condition number, leaving the policy conservative.
121    pub const REFINEMENT_TOL: f64 = 1e-12;
122
123    /// Return `true` when the policy and problem size together suggest that
124    /// attempting fp32 factorization + iterative refinement will be profitable.
125    ///
126    /// The predicate is conservative:
127    ///   * `GpuMixedPrecisionPolicy::Off` or `Never` → always `false`.
128    ///   * `Refinement` with `p < REFINEMENT_MIN_P` → `false` (GEMV overhead
129    ///     not amortised by fp32 POTRF savings below this threshold).
130    ///   * Otherwise `true`; the caller still falls back to fp64 factorization
131    ///     when the runtime fp32 POTRF fails or when the measured residual is
132    ///     non-monotone.
133    #[inline]
134    pub const fn iterative_refinement_should_attempt(&self, p: usize) -> bool {
135        match self.mixed_precision {
136            GpuMixedPrecisionPolicy::Off | GpuMixedPrecisionPolicy::Never => false,
137            GpuMixedPrecisionPolicy::Refinement => p >= Self::REFINEMENT_MIN_P,
138        }
139    }
140
141    pub const fn dense_gemv_target_is_gpu(&self, n: usize, p: usize, resident: bool) -> bool {
142        resident || n.saturating_mul(p).saturating_mul(2) >= self.gemm_min_flops
143    }
144
145    pub const fn xtwx_target_is_gpu(&self, n: usize, p: usize, materialized: bool) -> bool {
146        materialized && n > 0 && p > 0 && self.xtwx_flops(n, p) >= self.dense_reduction_flops_min()
147    }
148
149    pub const fn xtwy_target_is_gpu(
150        &self,
151        n: usize,
152        px: usize,
153        q: usize,
154        materialized: bool,
155    ) -> bool {
156        materialized
157            && n > 0
158            && px > 0
159            && q > 0
160            && self.xtwy_flops(n, px, q) >= self.dense_reduction_flops_min()
161    }
162
163    pub const fn potrf_target_is_gpu(&self, p: usize, h_resident: bool) -> bool {
164        h_resident && p >= self.potrf_min_p
165    }
166
167    pub const fn dense_hessian_work_target_is_gpu(&self, n: usize, p: usize) -> bool {
168        n > 0
169            && p >= Self::DEVICE_LOOP_MIN_P
170            && self.xtwx_flops(n, p) >= self.dense_reduction_flops_min()
171    }
172
173    const fn dense_reduction_flops_min(&self) -> u128 {
174        if self.xtwx_flops_min < self.gemm_min_flops {
175            self.xtwx_flops_min as u128
176        } else {
177            self.gemm_min_flops as u128
178        }
179    }
180
181    const fn xtwx_flops(&self, n: usize, p: usize) -> u128 {
182        2u128 * (n as u128) * (p as u128) * (p as u128)
183    }
184
185    const fn xtwy_flops(&self, n: usize, px: usize, q: usize) -> u128 {
186        2u128 * (n as u128) * (px as u128) * (q as u128)
187    }
188
189    /// Minimum total CG-amortised matvec flops below which the host↔device
190    /// transfer of the row frames + CG vectors is not repaid by the device
191    /// matvec, so the reduced-Schur PCG hot loop stays on the CPU.
192    ///
193    /// The dense-Direct path keys on `dense_reduction_flops_min` (a single big
194    /// factorization). The matrix-free SAE matvec is different: no single apply
195    /// trips that floor (each is a stack of `n` tiny `d×d` solves + sparse
196    /// `m·k` gather/scatter), but the *whole CG solve* runs the apply
197    /// `O(cg_iters)` times over the same resident frames. The device wins when
198    /// the **summed** matvec work over the solve exceeds the one-time staging
199    /// cost — so the gate keys on `cg_iters · per_apply_flops`, not one apply.
200    ///
201    /// Set one order of magnitude below the dense floor: the matvec frames stay
202    /// resident across CG iterations (uploaded once), so the per-flop transfer
203    /// amortization is `1/cg_iters` of a cold dense launch, and the breakeven
204    /// drops accordingly.
205    pub const MATVEC_OFFLOAD_FLOPS_MIN: u128 = 10_000_000;
206
207    /// Thin-curve (`d_atom = 1`) SAE dictionaries are the common manifold-SAE
208    /// production shape: each per-row frame is a scalar, so the staged device
209    /// payload is much smaller than the general `d > 1` row-frame bundle, while
210    /// the work is still a large batched gather/scatter over `K` atoms and `n`
211    /// rows.  Use a lower admission floor for this scalar-frame regime so a
212    /// realistic token block with a moderately wide curve dictionary is not kept
213    /// on the CPU solely because the conservative general-frame lower-bound
214    /// undercounts the transpose cross term.
215    pub const THIN_CURVE_MATVEC_OFFLOAD_FLOPS_MIN: u128 = 1_000_000;
216
217    /// Conservative seed for the reduced-Schur PCG iteration count when the
218    /// caller cannot supply a measured budget. InexactPCG on an SAE β-block of
219    /// width `k` converges in `O(√κ)` iterations; this floor keeps the work
220    /// estimate honest (≥ this many applies) without over-claiming a tight
221    /// solve. Used only to amortise the staging cost in the work estimate.
222    pub const MATVEC_OFFLOAD_MIN_CG_ITERS: usize = 8;
223
224    /// Per-apply flop estimate for one reduced-Schur matvec `S·x` of a
225    /// matrix-free SAE Kronecker system, as a pure function of the system shape.
226    ///
227    /// Per row block `i` the apply does: a forward cross-block GEMV
228    /// `v_i = H_tβ^(i)·x` (`≈ 2·d·k` multiply-adds, with the per-row latent
229    /// depth `d` as the M-frame width and `k` the border), a `d×d` triangular
230    /// solve through the cached Cholesky factor (`≈ d²`), and a transpose
231    /// cross-block GEMV `H_βt^(i)·w_i` (`≈ 2·d·k`). The two `2·d·k` GEMVs would
232    /// sum to `4·d·k`; this estimate deliberately undercounts to a single
233    /// `2·d·k` cross term as a conservative (lower-bound) admission floor, so
234    /// the apply is modelled as `≈ n·(2·d·k + d²)`. This is a deliberate
235    /// lower bound on the true `≈ n·(4·d·k + d²)` arithmetic — admitting a
236    /// shape under the smaller figure can only be more conservative, never
237    /// over-eager. It is keyed on the *frame depth* `d` (M) and border width
238    /// `k` (p), not row count alone, so LLM shapes (few rows, wide `k`, modest
239    /// `d`) register arithmetic the row-count gate misses.
240    ///
241    /// USE FOR DISPATCH GATING ONLY. This is **not** a flop count: it omits the
242    /// transpose cross-block GEMV (`2·d·k`), so it is a strict lower bound on the
243    /// true per-apply work `n·(4·d·k + d²)`. The gate can therefore only
244    /// under-admit, never over-admit. Do not reuse it for benchmark / speedup
245    /// accounting.
246    const fn admission_work_lower_bound(n: usize, k: usize, d: usize) -> u128 {
247        let n = n as u128;
248        let k = k as u128;
249        let d = d as u128;
250        // 2·d·k cross-block apply (forward only) + d² per-row solve — the
251        // transpose GEMV is intentionally dropped so this stays a lower bound.
252        n.saturating_mul(
253            2u128
254                .saturating_mul(d)
255                .saturating_mul(k)
256                .saturating_add(d * d),
257        )
258    }
259
260    /// Work-based admission for offloading the **reduced-Schur PCG matvec**
261    /// (the InexactPCG hot loop for matrix-free SAE β-blocks) to the device.
262    ///
263    /// This is the Phase-1 (#1017) re-keying: the dense gates key on row count
264    /// (`xtwx_n_min`, `row_kernel_min_n` at 50k) or a single big-factorization
265    /// flop floor, neither of which the SAE LLM shape trips — `(n≈2000) ×
266    /// (k≈2048) × (d≈8)` is *thousands of small dense ops*, no single op large,
267    /// so the row-count gate keeps the whole fit on one CPU core. Here the gate
268    /// is the **total batched work over the CG solve**:
269    ///
270    /// ```text
271    /// estimated_device_flops = cg_iters · per_apply_flops(n, k, d)
272    /// should_offload = estimated_device_flops ≥ T_breakeven
273    /// ```
274    ///
275    /// where `T_breakeven = MATVEC_OFFLOAD_FLOPS_MIN` accounts for the
276    /// host↔device staging of the row frames + CG vectors amortised over the
277    /// `cg_iters` applies that reuse the resident frames (so the per-flop
278    /// transfer cost is `1/cg_iters` of a cold launch, an order of magnitude
279    /// below the dense-Direct floor).
280    ///
281    /// Pure function of the shape: no device needed to evaluate, so it is unit-
282    /// testable. The caller still falls back to the bit-identical CPU matvec
283    /// whenever the backend build declines, so admitting a shape never changes
284    /// the numerics — only where the `Σ_i Y_iᵀ(Y_i x)` flops execute.
285    ///
286    /// * `n`        — number of row blocks (SAE observations / latent rows).
287    /// * `k`        — border β width (the SAE decoder atom count `K`).
288    /// * `d`        — per-row latent / active-frame depth (the M dimension).
289    /// * `cg_iters` — expected PCG iteration budget; the per-apply work is
290    ///   multiplied by this because the frames stay resident across iterations.
291    ///   Pass [`Self::MATVEC_OFFLOAD_MIN_CG_ITERS`] when no measured budget is
292    ///   available; a tighter (smaller) value only makes the gate stricter.
293    ///
294    /// ## Live arrow-Schur call site
295    ///
296    /// `crate::solver::arrow_schur::maybe_inject_gpu_schur_matvec` gates the
297    /// InexactPCG reduced-Schur matvec injection on this predicate:
298    /// `reduced_schur_matvec_should_offload(sys.rows.len(), sys.k, sys.d,
299    /// options.pcg.max_iterations.min(options.trust_region.max_iterations))`,
300    /// where `sys.d` is the system's max per-row latent depth and the iteration
301    /// budget is the same `max_iterations` the PCG loop launches with.
302    /// `try_device_arrow_direct` (the **dense** Direct point solve) correctly
303    /// keeps `dense_hessian_work_target_is_gpu`: that path is a single large
304    /// factorization, not the amortised matvec.
305    pub const fn reduced_schur_matvec_should_offload(
306        &self,
307        n: usize,
308        k: usize,
309        d: usize,
310        cg_iters: usize,
311    ) -> bool {
312        if n == 0 || k == 0 || d == 0 || cg_iters == 0 {
313            return false;
314        }
315        // The border width must clear the device-loop floor: below it the per-
316        // apply launch latency (one kernel sequence per matvec) dominates any
317        // arithmetic regardless of how many CG iterations run.
318        if k < Self::DEVICE_LOOP_MIN_P {
319            return false;
320        }
321        let per_apply = Self::admission_work_lower_bound(n, k, d);
322        let total = per_apply.saturating_mul(cg_iters as u128);
323        let floor = if d == 1 {
324            Self::THIN_CURVE_MATVEC_OFFLOAD_FLOPS_MIN
325        } else {
326            Self::MATVEC_OFFLOAD_FLOPS_MIN
327        };
328        total >= floor
329    }
330}
331
332/// Factorization strategy for the arrow-Schur border (shared `β`) solve, chosen
333/// from the *shape* of the joint system rather than a single fixed border-width
334/// cut (`ArrowSolverMode::automatic`'s `DIRECT_SOLVE_MAX_K = 2000`).
335///
336/// The border width alone is a blunt selector: it cannot see that the data-fit
337/// contribution to the `k × k` border is only rank `Σ_i d_i ≈ n·d`. For the
338/// #1017 color arm (`n = 180`, per-row depth `d = 2`, border `k = 15360`) the
339/// data information is rank `360` yet a dense Direct solve pays a full `k³/3 ≈
340/// 1.2e12`-flop Cholesky — the measured 26-min-class fit. This maps cleanly onto
341/// the two `ArrowSolverMode` variants the solver already implements.
342#[derive(Clone, Copy, Debug, Eq, PartialEq, Serialize, Deserialize)]
343pub enum ArrowBorderStrategy {
344    /// Eliminate the per-row blocks, form the dense `k × k` reduced Schur, and
345    /// Cholesky-factor it (`ArrowSolverMode::Direct`). Appropriate for modest,
346    /// near-square borders where the `k³/3` factorization is cheap and the
347    /// data-fit rank is comparable to `k`.
348    DenseDirect,
349    /// Solve the reduced Schur iteratively by matrix-free PCG
350    /// (`ArrowSolverMode::InexactPCG`), never materialising the `k × k` factor.
351    /// Appropriate when the dense `k³` factorization dominates and/or the
352    /// data-fit contribution to the border is rank-deficient (`n·d < k`).
353    ReducedIterative,
354}
355
356/// Cost model + recommendation for the arrow-Schur border solve, a pure function
357/// of the joint-system shape (unit-testable, no device required).
358///
359/// This operationalises the measured #1017 finding that the full arrow-Schur
360/// Newton solve is dominated by the dense `k × k` border Cholesky (the on-device
361/// dense Direct solve was measured at ~0.94× — a slowdown — because the `k³/3`
362/// factorization, not the GPU-favourable batched per-row work, is the bottleneck
363/// at LLM/SAE border widths). The lever the issue calls for is to *shrink or
364/// factor the dense border* so the batched `n`-row work dominates; the plan
365/// makes that decision inspectable and honest.
366///
367/// ## Flop model (deliberate, documented approximations)
368///
369/// * **Dense Direct** ≈ `2·n·d·k²` (assemble the reduced Schur: per row a
370///   rank-`d` symmetric update `H_βt (H_tt)⁻¹ H_tβ` to the `k × k` border,
371///   `≈ 2·d·k²` flops) `+ k³/3` (Cholesky of the dense `k × k` Schur).
372/// * **Reduced iterative** ≈ `cg_iters · n·(4·d·k + d²)` (matrix-free PCG:
373///   per matvec a forward + transpose cross-block GEMV `4·d·k` plus the per-row
374///   `d × d` solve `d²`, summed over `n` row blocks, over `cg_iters` applies).
375///
376/// Both are dispatch-grade estimates, not exact operation counts; they omit
377/// preconditioner setup and lower-order terms symmetrically, so their ratio (the
378/// only thing the recommendation consumes) is meaningful while neither figure
379/// should be reused for speedup accounting.
380///
381/// ## Status
382///
383/// Advisory / diagnostic. It is **not** wired into the live
384/// `ArrowSolverMode::automatic` selector: replacing the fixed `DIRECT_SOLVE_MAX_K`
385/// cut with this shape-driven crossover changes which production fits take the
386/// Direct vs PCG path and must be validated on GPU hardware (#1017 Phase 2–4)
387/// before it can change numerics. Today it is consumed by the honest
388/// `examples/full_color_fit_1017.rs` measurement harness (modeled-vs-measured)
389/// and by the unit tests below.
390#[derive(Clone, Copy, Debug, Eq, PartialEq)]
391pub struct ArrowBorderSolvePlan {
392    /// Number of per-row blocks (SAE observations / latent rows).
393    pub n: usize,
394    /// Border `β` width (the SAE decoder atom count `K` × basis width).
395    pub k: usize,
396    /// Per-row latent / active-frame depth (the `M` dimension).
397    pub d: usize,
398    /// CG iteration budget assumed for the iterative estimate.
399    pub cg_iters: usize,
400    /// Effective rank of the data-fit contribution to the `k × k` border,
401    /// bounded by `Σ_i d_i ≈ n·d` and never more than `k`.
402    pub data_fit_rank: usize,
403    /// True when `n·d < k`: the dense `k × k` Cholesky spends `O(k³)` factorising
404    /// a border whose data information is only rank `n·d` — the pathological
405    /// wide-sparse-border regime (color arm: `n·d = 360 ≪ k = 15360`).
406    pub dense_border_rank_deficient: bool,
407    /// `≈ 2·n·d·k² + k³/3` — reduced-Schur assembly plus dense border Cholesky.
408    pub dense_direct_flops: u128,
409    /// `≈ cg_iters · n·(4·d·k + d²)` — matrix-free PCG matvecs.
410    pub reduced_iterative_flops: u128,
411    /// The recommended strategy. `DenseDirect` is chosen only for a full-rank
412    /// border (`n·d ≥ k`, so the `k × k` reduced Schur is non-singular and its
413    /// Cholesky exists) whose `k³/3` border factorization is no costlier than
414    /// the matrix-free CG solve at `cg_iters`; otherwise `ReducedIterative`. A
415    /// rank-deficient border is always `ReducedIterative` — a dense Cholesky of
416    /// a singular border does not exist.
417    pub recommended: ArrowBorderStrategy,
418    /// Whether running the *recommended* strategy on the device is expected to
419    /// pay off. For `ReducedIterative` this is `reduced_schur_matvec_should_offload`;
420    /// for `DenseDirect` the device wins only when the batched per-row assembly
421    /// work (`2·n·d·k²`, GPU-favourable batched GEMM/POTRF) at least matches the
422    /// border Cholesky (`k³/3`) *and* clears the dense flop floor — the honest
423    /// encoding of the measured 0.94× dense-Direct-on-device slowdown.
424    pub device_favorable: bool,
425}
426
427impl GpuDispatchPolicy {
428    /// Assembly flops for the dense reduced Schur: per row a rank-`d` update to
429    /// the `k × k` border (`≈ 2·d·k²`), summed over `n` rows.
430    const fn dense_schur_assembly_flops(n: usize, k: usize, d: usize) -> u128 {
431        2u128
432            .saturating_mul(n as u128)
433            .saturating_mul(d as u128)
434            .saturating_mul((k as u128).saturating_mul(k as u128))
435    }
436
437    /// Cholesky flops for the dense `k × k` reduced Schur: `≈ k³/3`.
438    const fn dense_border_cholesky_flops(k: usize) -> u128 {
439        let k = k as u128;
440        k.saturating_mul(k).saturating_mul(k) / 3
441    }
442
443    /// Total matrix-free PCG flops: `cg_iters · n·(4·d·k + d²)`.
444    const fn reduced_iterative_flops(n: usize, k: usize, d: usize, cg_iters: usize) -> u128 {
445        let n = n as u128;
446        let k = k as u128;
447        let d = d as u128;
448        let per_apply = n.saturating_mul(
449            4u128
450                .saturating_mul(d)
451                .saturating_mul(k)
452                .saturating_add(d.saturating_mul(d)),
453        );
454        per_apply.saturating_mul(cg_iters as u128)
455    }
456
457    /// Build the shape-driven [`ArrowBorderSolvePlan`] for a joint arrow-Schur
458    /// system with `n` row blocks, border width `k`, per-row depth `d`, and an
459    /// assumed CG budget `cg_iters` (pass
460    /// [`Self::MATVEC_OFFLOAD_MIN_CG_ITERS`] when none is measured; a smaller
461    /// value only biases the recommendation toward `DenseDirect`, never the
462    /// reverse).
463    ///
464    /// Degenerate shapes (`n`, `k`, or `d` zero) return an all-zero plan
465    /// recommending `DenseDirect` (the trivial/empty solve stays on the simple
466    /// path) with `device_favorable = false`.
467    pub fn arrow_border_solve_plan(
468        &self,
469        n: usize,
470        k: usize,
471        d: usize,
472        cg_iters: usize,
473    ) -> ArrowBorderSolvePlan {
474        if n == 0 || k == 0 || d == 0 {
475            return ArrowBorderSolvePlan {
476                n,
477                k,
478                d,
479                cg_iters,
480                data_fit_rank: 0,
481                dense_border_rank_deficient: false,
482                dense_direct_flops: 0,
483                reduced_iterative_flops: 0,
484                recommended: ArrowBorderStrategy::DenseDirect,
485                device_favorable: false,
486            };
487        }
488
489        let assembly = Self::dense_schur_assembly_flops(n, k, d);
490        let border_chol = Self::dense_border_cholesky_flops(k);
491        let dense_direct_flops = assembly.saturating_add(border_chol);
492        let iters = if cg_iters == 0 { 1 } else { cg_iters };
493        let reduced_iterative_flops = Self::reduced_iterative_flops(n, k, d, iters);
494
495        let data_fit_rank = (n.saturating_mul(d)).min(k);
496        let dense_border_rank_deficient = n.saturating_mul(d) < k;
497
498        // Recommend the exact dense factorization only when it is both VALID and
499        // not the bottleneck:
500        //   * Validity — a rank-deficient border (`n·d < k`) has a singular
501        //     `k × k` reduced Schur, so its Cholesky does not exist. DenseDirect
502        //     is inadmissible there and we must solve matrix-free. (The pure
503        //     assembly+Cholesky-vs-iterative flop rule this replaced ignored
504        //     rank and could recommend factorizing a provably-singular border
505        //     whenever the small-`k` dense flops happened to be the cheaper
506        //     count — e.g. `n=1, d=1, k=2`.)
507        //   * Cost — the reduced-Schur reduction is an embarrassingly-parallel
508        //     batched GEMM whichever path runs; the term that scales badly with
509        //     border width is the `k³/3` Cholesky. Prefer the exact,
510        //     RHS-reusable, convergence-free dense solve while that Cholesky is
511        //     no costlier than the full matrix-free CG solve, and fall to
512        //     ReducedIterative once the `k³` factorization overtakes it.
513        let recommended =
514            if !dense_border_rank_deficient && border_chol <= reduced_iterative_flops {
515                ArrowBorderStrategy::DenseDirect
516            } else {
517                ArrowBorderStrategy::ReducedIterative
518            };
519
520        let device_favorable = match recommended {
521            ArrowBorderStrategy::ReducedIterative => {
522                self.reduced_schur_matvec_should_offload(n, k, d, iters)
523            }
524            ArrowBorderStrategy::DenseDirect => {
525                // Dense Direct wins on device only when the batched per-row
526                // assembly work dominates the (poorly GPU-scaling, and here
527                // rank-deficient) border Cholesky, and the total clears the
528                // dense reduction floor. This is the honest encoding of the
529                // measured 0.94× on-device dense-Direct slowdown: when the k³
530                // Cholesky dominates, stay on the CPU.
531                assembly >= border_chol && dense_direct_flops >= self.dense_reduction_flops_min()
532            }
533        };
534
535        ArrowBorderSolvePlan {
536            n,
537            k,
538            d,
539            cg_iters: iters,
540            data_fit_rank,
541            dense_border_rank_deficient,
542            dense_direct_flops,
543            reduced_iterative_flops,
544            recommended,
545            device_favorable,
546        }
547    }
548}
549
550/// The aspirational single-GPU design-row throughput the #1412 decision gate is
551/// supposed to establish for the LLM-shape batched-Cholesky + tile-GEMM fit
552/// pipeline: 100 000 design rows processed per wall-clock second per device.
553///
554/// The original gate *claimed* this number without ever measuring it. The
555/// honest contract is the other way around: a benchmark
556/// (`examples/throughput_1412.rs`) measures the true rows/sec on a real device,
557/// and [`GpuThroughputVerdict::from_measurement`] reports whether the measured
558/// value meets the target — the verdict is a *function of the measurement*, not
559/// a hardcoded assertion. See `tests/owed_1412.rs`.
560pub const GPU_THROUGHPUT_TARGET_ROWS_PER_SEC: f64 = 100_000.0;
561
562/// Outcome of comparing a *measured* GPU throughput against the target. The
563/// only way to construct one is [`Self::from_measurement`], so a verdict can
564/// never assert a target that was not actually established by a measurement.
565#[derive(Clone, Copy, Debug, PartialEq)]
566pub struct GpuThroughputVerdict {
567    /// The measured design-rows-per-second on the device under test.
568    pub measured_rows_per_sec: f64,
569    /// The target the measurement is compared against.
570    pub target_rows_per_sec: f64,
571    /// `measured / target`. ≥ 1.0 means the target was established.
572    pub fraction_of_target: f64,
573    /// True iff `measured_rows_per_sec >= target_rows_per_sec`.
574    pub meets_target: bool,
575}
576
577impl GpuThroughputVerdict {
578    /// Build a verdict from a measured throughput against
579    /// [`GPU_THROUGHPUT_TARGET_ROWS_PER_SEC`]. A non-finite or non-positive
580    /// measurement can never meet the target (it is not a usable measurement).
581    #[inline]
582    pub fn from_measurement(measured_rows_per_sec: f64) -> Self {
583        Self::from_measurement_against(measured_rows_per_sec, GPU_THROUGHPUT_TARGET_ROWS_PER_SEC)
584    }
585
586    /// Build a verdict against an explicit target (used by tests that probe the
587    /// comparison logic without depending on the global target constant).
588    #[inline]
589    pub fn from_measurement_against(measured_rows_per_sec: f64, target_rows_per_sec: f64) -> Self {
590        let usable = measured_rows_per_sec.is_finite() && measured_rows_per_sec > 0.0;
591        let fraction_of_target = if usable && target_rows_per_sec > 0.0 {
592            measured_rows_per_sec / target_rows_per_sec
593        } else {
594            0.0
595        };
596        Self {
597            measured_rows_per_sec,
598            target_rows_per_sec,
599            fraction_of_target,
600            meets_target: usable && measured_rows_per_sec >= target_rows_per_sec,
601        }
602    }
603}
604
605/// Why a Stage-3 encode deployment decision could not be made from a real device
606/// measurement (#988, #1412). Each variant is a state in which the
607/// `100_000` rows/sec/GPU target was neither established NOR refuted on a
608/// device — the decision is blocked on hardware, not green-washed from a CPU
609/// proxy.
610#[derive(Clone, Copy, Debug, PartialEq, Eq)]
611pub enum EncodeDecisionBlocked {
612    /// No CUDA device on this host: the exact encode could not be measured on a
613    /// device at all (a CPU rate cannot substitute — that was the #1412 defect).
614    NoDevice,
615    /// A device is present but there is no device-resident *exact-encode* kernel,
616    /// so the FULL per-row encode cannot be measured on the device. (The resident
617    /// normal-equations solve in [`crate::encode_throughput`] is only ONE
618    /// component of the encode, not the encode; a component measurement cannot
619    /// decide the encode surrogate question — #988.)
620    NoDeviceEncodeKernel,
621    /// A device is present and a measurement was attempted, but the device path
622    /// did not engage (false routing) — refused rather than reported as a pass.
623    DeviceNotEngaged,
624}
625
626/// Tri-state Stage-3 encode deployment / amortized-surrogate decision
627/// (#988, #1412).
628///
629/// The decision the throughput gate exists to make is empirical: does the EXACT
630/// per-row encode clear the `100_000` rows/sec/GPU deployment target on a real
631/// device? Only a real device measurement can answer it:
632///   * [`Self::Met`] — a device measurement CLEARED the target: ship the exact
633///     encode; the certified amortized surrogate is NOT needed.
634///   * [`Self::Unmet`] — a device measurement MISSED the target: the certified
635///     amortized surrogate becomes justified.
636///   * [`Self::Undetermined`] — no device measurement is available. The decision
637///     is BLOCKED on hardware; it is neither "surrogate unneeded" nor "surrogate
638///     justified".
639///
640/// The critical anti-green-wash property (#1412): there is NO constructor that
641/// takes a CPU rate. A CPU measurement, however fast, can never move the decision
642/// out of [`Self::Undetermined`]. Projecting a CPU rate through an assumed
643/// CPU→GPU factor to declare the target met was the exact #1412 defect and is
644/// structurally impossible here — [`Self::Met`] / [`Self::Unmet`] come only from
645/// [`Self::from_device_measurement`] with `engaged == true`.
646#[derive(Clone, Copy, Debug, PartialEq)]
647pub enum EncodeDeploymentDecision {
648    /// A device measurement established the deployment target.
649    Met {
650        /// The measured device rows/sec that cleared the target.
651        measured_rows_per_sec: f64,
652        /// The target it was compared against.
653        target_rows_per_sec: f64,
654    },
655    /// A device measurement fell short of the deployment target.
656    Unmet {
657        /// The measured device rows/sec that missed the target.
658        measured_rows_per_sec: f64,
659        /// The target it was compared against.
660        target_rows_per_sec: f64,
661    },
662    /// No device measurement is available; the decision is blocked on hardware.
663    Undetermined {
664        /// Why no device measurement could be made.
665        reason: EncodeDecisionBlocked,
666    },
667}
668
669impl EncodeDeploymentDecision {
670    /// The ONLY path to a `Met`/`Unmet` decision: a device measurement that
671    /// actually engaged the device and produced a usable rate. `engaged == false`
672    /// (false routing / CPU decline) or a non-finite / non-positive rate yields
673    /// [`Self::Undetermined`] — never a fabricated pass or fail.
674    #[must_use]
675    pub fn from_device_measurement(engaged: bool, measured_rows_per_sec: f64) -> Self {
676        Self::from_device_measurement_against(
677            engaged,
678            measured_rows_per_sec,
679            GPU_THROUGHPUT_TARGET_ROWS_PER_SEC,
680        )
681    }
682
683    /// [`Self::from_device_measurement`] against an explicit target (for tests
684    /// that probe the decision logic without the global target constant).
685    #[must_use]
686    pub fn from_device_measurement_against(
687        engaged: bool,
688        measured_rows_per_sec: f64,
689        target_rows_per_sec: f64,
690    ) -> Self {
691        let usable = measured_rows_per_sec.is_finite() && measured_rows_per_sec > 0.0;
692        if !engaged || !usable {
693            return Self::Undetermined {
694                reason: EncodeDecisionBlocked::DeviceNotEngaged,
695            };
696        }
697        if measured_rows_per_sec >= target_rows_per_sec {
698            Self::Met {
699                measured_rows_per_sec,
700                target_rows_per_sec,
701            }
702        } else {
703            Self::Unmet {
704                measured_rows_per_sec,
705                target_rows_per_sec,
706            }
707        }
708    }
709
710    /// Construct the blocked decision for a host that cannot measure the exact
711    /// encode on a device. This is the honest CPU-only / no-device-kernel outcome
712    /// — the deployment target is left undetermined rather than projected.
713    #[must_use]
714    pub fn blocked(reason: EncodeDecisionBlocked) -> Self {
715        Self::Undetermined { reason }
716    }
717
718    /// True ONLY when a device measurement cleared the target: the exact encode
719    /// ships and no surrogate is built. Never true from a CPU proxy.
720    #[must_use]
721    pub fn surrogate_unneeded(&self) -> bool {
722        matches!(self, Self::Met { .. })
723    }
724
725    /// True ONLY when a device measurement missed the target: the certified
726    /// amortized surrogate becomes justified. Never true without a measurement.
727    #[must_use]
728    pub fn surrogate_justified(&self) -> bool {
729        matches!(self, Self::Unmet { .. })
730    }
731
732    /// True when no device measurement is available and the decision is blocked
733    /// on hardware (neither [`Self::surrogate_unneeded`] nor
734    /// [`Self::surrogate_justified`]).
735    #[must_use]
736    pub fn is_undetermined(&self) -> bool {
737        matches!(self, Self::Undetermined { .. })
738    }
739}
740
741/// Which `(response, link)` family the Stage 3.3 device-resident PIRLS loop
742/// can evaluate without going through the Level-B raw-body NVRTC path.
743///
744/// Mirrors `PirlsRowFamily::ALL` at the policy layer so the predicate stays
745/// linkable from the CPU PIRLS entry without dragging a Linux-only enum into
746/// every host compilation unit.
747#[derive(Clone, Copy, Debug, Eq, PartialEq)]
748pub enum PirlsLoopFamilyKind {
749    BernoulliLogit,
750    BernoulliProbit,
751    BernoulliCLogLog,
752    PoissonLog,
753    GaussianIdentity,
754    GammaLog,
755}
756
757#[derive(Clone, Copy, Debug, Eq, PartialEq)]
758pub enum PirlsLoopCurvatureKind {
759    Fisher,
760    Observed,
761}
762
763/// Inputs to [`should_run_reml_outer_on_device`]. The admission predicate
764/// for routing the *outer* REML BFGS-over-ρ loop onto a fully device-resident
765/// driver (rather than the host orchestrator that hops out per step).
766///
767/// Fields are intentionally lifted from data the CPU REML entry has on hand
768/// before it touches the seed generator or the inner P-IRLS loop, so the
769/// admission check is allocation-free and can short-circuit before any
770/// device call.
771#[derive(Clone, Copy, Debug)]
772pub struct RemlOuterAdmission {
773    /// Active design rows (post-transform).
774    pub n: usize,
775    /// Active design columns / penalised-Hessian dimension.
776    pub p: usize,
777    /// Number of smoothing parameters ρ the outer BFGS optimises over.
778    pub num_rho: usize,
779    /// Inner family / link pair the device-resident PIRLS loop can evaluate.
780    /// `None` means the family does not map onto the six JIT-cached row
781    /// kernels — the outer loop must stay on the host orchestrator because
782    /// the inner step would already hop out anyway.
783    pub family: Option<PirlsLoopFamilyKind>,
784    /// Curvature surface the inner loop will use; tied to `family` via
785    /// `pirls_loop_curvature_for`.
786    pub curvature: PirlsLoopCurvatureKind,
787    /// True when the CUDA runtime is initialised on this host.
788    pub gpu_available: bool,
789}
790
791/// Inputs to [`should_use_gpu_pirls_loop`]. Each field comes from data the
792/// CPU PIRLS entry has on hand before it touches the eigendecomposition
793/// engine, so the admission check itself is allocation-free and can short-
794/// circuit before any heavy work happens.
795#[derive(Clone, Copy, Debug)]
796pub struct PirlsLoopAdmission {
797    /// Number of rows in the active (post-transform) design matrix.
798    pub n: usize,
799    /// Number of columns in the active design (i.e. `p` of `Xᵀ X`).
800    pub p: usize,
801    /// `Some(_)` when the inner family maps onto one of the six JIT-cached
802    /// `PirlsRowFamily` variants; `None` for custom families that still
803    /// require Stage 6 Level B and have not yet been admitted here.
804    pub family: Option<PirlsLoopFamilyKind>,
805    /// Curvature surface the inner loop will use; the GPU loop has Fisher +
806    /// Observed kernels, anything else (e.g. expected-projection surrogates)
807    /// is not admitted.
808    pub curvature: PirlsLoopCurvatureKind,
809    /// True when the CUDA runtime is initialised on this host (i.e.
810    /// lossless Auto resolution returned an available runtime).
811    pub gpu_available: bool,
812}
813
814impl GpuDispatchPolicy {
815    /// Minimum design column count for the device-resident inner/outer loops.
816    ///
817    /// Below this width the per-iteration `XᵀWX + Cholesky` is dominated by
818    /// launch latency and PCIe staging rather than arithmetic, so the host LM
819    /// loop (which populates the full `PirlsResult` surface as a free
820    /// side-effect) is strictly cheaper. Shared by both the inner PIRLS and
821    /// outer REML admission predicates so they cannot drift apart.
822    pub const DEVICE_LOOP_MIN_P: usize = 32;
823
824    /// Conservative admission predicate for routing
825    /// `fit_model_for_fixed_rho_with_adaptive_kkt` through the Stage 3.3
826    /// device-resident PIRLS loop instead of the CPU LM loop.
827    ///
828    /// The threshold is the dense `XᵀWX` work estimate, not row count alone:
829    /// LLM/SAE fits can have only a few thousand rows but thousands of columns,
830    /// so `2*n*p^2` already dwarfs launch/staging overhead. Smaller fits stay on
831    /// the CPU LM loop where the full `PirlsResult` surface (firth, EDF,
832    /// per-row weights, …) is already populated as a free side-effect of the
833    /// iteration.
834    pub const fn should_use_gpu_pirls_loop(&self, adm: PirlsLoopAdmission) -> bool {
835        if !adm.gpu_available {
836            return false;
837        }
838        if !self.dense_hessian_work_target_is_gpu(adm.n, adm.p) {
839            return false;
840        }
841        match adm.family {
842            Some(_) => true,
843            None => false,
844        }
845    }
846
847    /// Admission predicate for routing the outer REML BFGS-over-ρ loop onto
848    /// a device-resident driver that keeps the BFGS state (ρ, gradient,
849    /// Hessian approx) on-device and only downloads the per-step scalar
850    /// metrics (objective value, gradient norm, convergence flag).
851    ///
852    /// The dense-work threshold piggybacks on the existing inner-PIRLS admission
853    /// predicate because the device-resident outer loop calls
854    /// `pirls_loop_on_stream` per step and must not pay the host hop for small
855    /// fits the inner loop would have rejected anyway. The
856    /// `num_rho ≥ 2` floor rules out the trivial single-smoother case where
857    /// host orchestration is already negligible and the device BFGS state
858    /// (one length-`num_rho` gradient + a `num_rho × num_rho` Hessian
859    /// approx) collapses to a couple of scalars not worth keeping on device.
860    pub const fn should_run_reml_outer_on_device(&self, adm: RemlOuterAdmission) -> bool {
861        if !adm.gpu_available {
862            return false;
863        }
864        if !self.dense_hessian_work_target_is_gpu(adm.n, adm.p) {
865            return false;
866        }
867        if adm.num_rho < 2 {
868            return false;
869        }
870        match adm.family {
871            Some(_) => true,
872            None => false,
873        }
874    }
875}
876
877#[cfg(test)]
878mod refinement_policy_tests {
879    use super::*;
880
881    #[test]
882    fn refinement_policy_admits_large_p() {
883        let pol = GpuDispatchPolicy::default();
884        // Default policy is Refinement; large p should be admitted.
885        assert!(pol.iterative_refinement_should_attempt(512));
886        assert!(pol.iterative_refinement_should_attempt(GpuDispatchPolicy::REFINEMENT_MIN_P));
887    }
888
889    #[test]
890    fn refinement_policy_rejects_small_p() {
891        let pol = GpuDispatchPolicy::default();
892        assert!(!pol.iterative_refinement_should_attempt(GpuDispatchPolicy::REFINEMENT_MIN_P - 1));
893        assert!(!pol.iterative_refinement_should_attempt(0));
894    }
895
896    #[test]
897    fn off_policy_never_attempts_refinement() {
898        let pol = GpuDispatchPolicy {
899            mixed_precision: GpuMixedPrecisionPolicy::Off,
900            ..Default::default()
901        };
902        assert!(!pol.iterative_refinement_should_attempt(1024));
903    }
904
905    #[test]
906    fn never_policy_never_attempts_refinement() {
907        let pol = GpuDispatchPolicy {
908            mixed_precision: GpuMixedPrecisionPolicy::Never,
909            ..Default::default()
910        };
911        assert!(!pol.iterative_refinement_should_attempt(1024));
912    }
913}
914
915#[cfg(test)]
916mod reduced_schur_matvec_offload_tests {
917    use super::*;
918
919    /// The LLM/SAE shape the whole #1017 Phase-1 re-keying targets: a few
920    /// thousand row blocks, a *wide* border (decoder atom count in the
921    /// thousands), a modest per-row frame depth, and a realistic CG budget.
922    /// The row-count gate (50k) and the dense-Direct flop floor both miss this
923    /// "thousands of tiny dense ops" shape; the work-amortised matvec gate must
924    /// fire on it.
925    #[test]
926    fn admits_llm_sae_matvec_shape() {
927        let pol = GpuDispatchPolicy::default();
928        // n≈2000 rows, k≈2048 atoms, M≈8 frame depth — n is far below the 50k
929        // row gate, yet the summed CG matvec work is large.
930        assert!(pol.reduced_schur_matvec_should_offload(
931            2_000,
932            2_048,
933            8,
934            GpuDispatchPolicy::MATVEC_OFFLOAD_MIN_CG_ITERS,
935        ));
936        // The same shape would be rejected by the row-count-style dense gate,
937        // confirming the re-keying is what admits it.
938        assert!(!pol.dense_hessian_work_target_is_gpu(2_000, 8));
939    }
940
941    /// Even with only a single conservative CG iteration the wide LLM border
942    /// clears the breakeven (the per-apply work alone is `2_000·(2·8·2_048 +
943    /// 8²) ≈ 6.6e7` flops > 1e7 by the conservative `n·(2·d·k + d²)` model;
944    /// the true `n·(4·d·k + d²)` arithmetic is ≈1.3e8),
945    /// so the gate is not relying on an inflated iteration count.
946    #[test]
947    fn admits_llm_shape_with_one_cg_iter() {
948        let pol = GpuDispatchPolicy::default();
949        assert!(pol.reduced_schur_matvec_should_offload(2_000, 2_048, 8, 1));
950    }
951
952    /// #1783: the primary manifold-SAE regime is a `d_atom = 1` curve
953    /// dictionary.  Its scalar row frames have much lower staging cost than the
954    /// general framed matvec, so realistic token blocks must not be stranded on
955    /// the CPU merely because the conservative admission lower bound is thin in
956    /// `d`.
957    #[test]
958    fn admits_thin_curve_atoms_at_realistic_scale() {
959        let pol = GpuDispatchPolicy::default();
960        assert!(pol.reduced_schur_matvec_should_offload(24_576, 64, 1, 1));
961        assert!(pol.reduced_schur_matvec_should_offload(40_456, 256, 1, 1));
962        assert!(!pol.reduced_schur_matvec_should_offload(300, 6, 1, 8));
963    }
964
965    /// Tiny shapes where the host↔device transfer dominates must stay on the
966    /// CPU: a handful of rows, a narrow border, shallow frames. The summed
967    /// matvec work is orders of magnitude below the staging breakeven.
968    #[test]
969    fn rejects_tiny_shape_where_transfer_dominates() {
970        let pol = GpuDispatchPolicy::default();
971        assert!(!pol.reduced_schur_matvec_should_offload(
972            30,
973            8,
974            2,
975            GpuDispatchPolicy::MATVEC_OFFLOAD_MIN_CG_ITERS,
976        ));
977        // The 300×8 shape the production seam tests use as the "stay CPU"
978        // canary is rejected here too.
979        assert!(!pol.reduced_schur_matvec_should_offload(300, 8, 4, 16));
980    }
981
982    /// A narrow border (k below the device-loop floor) is rejected regardless
983    /// of how much row/iteration work is piled on: per-apply launch latency
984    /// dominates a sub-`DEVICE_LOOP_MIN_P` border.
985    #[test]
986    fn rejects_narrow_border_even_with_huge_row_count() {
987        let pol = GpuDispatchPolicy::default();
988        let narrow = GpuDispatchPolicy::DEVICE_LOOP_MIN_P - 1;
989        assert!(!pol.reduced_schur_matvec_should_offload(1_000_000, narrow, 64, 64));
990    }
991
992    /// Degenerate dimensions are never offloaded (no work, or no solve).
993    #[test]
994    fn rejects_degenerate_dimensions() {
995        let pol = GpuDispatchPolicy::default();
996        assert!(!pol.reduced_schur_matvec_should_offload(0, 2_048, 8, 8));
997        assert!(!pol.reduced_schur_matvec_should_offload(2_000, 0, 8, 8));
998        assert!(!pol.reduced_schur_matvec_should_offload(2_000, 2_048, 0, 8));
999        assert!(!pol.reduced_schur_matvec_should_offload(2_000, 2_048, 8, 0));
1000    }
1001
1002    /// The gate is monotone in the CG budget: once a shape is admitted at a
1003    /// given iteration count it stays admitted for any larger count (more
1004    /// applies over the same resident frames only improves amortization), and
1005    /// a borderline shape crosses the breakeven as iterations grow.
1006    #[test]
1007    fn monotone_in_cg_iters() {
1008        let pol = GpuDispatchPolicy::default();
1009        // A border at the floor with shallow frames and few rows: per-apply
1010        // work ~ n·(2·d·k + d²). Choose a shape that is below breakeven at 1
1011        // iter but above it once enough iterations accumulate.
1012        let (n, k, d) = (200usize, GpuDispatchPolicy::DEVICE_LOOP_MIN_P, 4usize);
1013        // per_apply ≈ 200·(2·4·32 + 16) = 200·272 = 54_400 flops.
1014        assert!(!pol.reduced_schur_matvec_should_offload(n, k, d, 1));
1015        // Once the summed work clears 1e7 the gate fires; ~184 iters here.
1016        assert!(pol.reduced_schur_matvec_should_offload(n, k, d, 1_000));
1017        // Monotonicity: admitted at 1_000 ⇒ admitted at every larger budget.
1018        assert!(pol.reduced_schur_matvec_should_offload(n, k, d, 5_000));
1019    }
1020
1021    /// The admission lower bound must stay strictly below the true per-apply
1022    /// work `n·(4·d·k + d²)` for any non-degenerate cross-block shape (it drops
1023    /// the transpose GEMV). Treating the lower bound as a flop count would
1024    /// over-report device speedups, so this asserts the gap is real.
1025    #[test]
1026    fn admission_lower_bound_undercounts_actual_work() {
1027        for &(n, k, d) in &[
1028            (2_000usize, 2_048usize, 8usize),
1029            (200, GpuDispatchPolicy::DEVICE_LOOP_MIN_P, 4),
1030            (1, 1, 1),
1031        ] {
1032            let lower = GpuDispatchPolicy::admission_work_lower_bound(n, k, d);
1033            // True per-apply work models the full forward+transpose GEMV pair
1034            // plus the d×d solve: n·(4·d·k + d²).
1035            let actual = (n as u128) * (4 * (d as u128) * (k as u128) + (d as u128) * (d as u128));
1036            assert!(
1037                lower < actual,
1038                "admission lower bound {lower} must undercount actual work {actual} for ({n},{k},{d})"
1039            );
1040        }
1041    }
1042}
1043
1044#[cfg(test)]
1045mod arrow_border_solve_plan_tests {
1046    use super::*;
1047
1048    /// The #1017 color arm — few rows, shallow per-row depth, a very wide border
1049    /// (`k = 15360 = 3 × 5120`). The dense `k³/3` Cholesky (`≈ 1.2e12` flops)
1050    /// dwarfs a matrix-free PCG solve at any realistic CG budget, and the border
1051    /// is grossly rank-deficient (`n·d = 360 ≪ k`). The plan must recommend
1052    /// `ReducedIterative` and flag the rank deficiency.
1053    #[test]
1054    fn color_arm_recommends_reduced_iterative_and_flags_rank_deficiency() {
1055        let pol = GpuDispatchPolicy::default();
1056        let plan = pol.arrow_border_solve_plan(180, 15_360, 2, 30);
1057        assert_eq!(plan.recommended, ArrowBorderStrategy::ReducedIterative);
1058        assert!(plan.dense_border_rank_deficient);
1059        assert_eq!(plan.data_fit_rank, 360);
1060        // The dense path is orders of magnitude more expensive here.
1061        assert!(plan.dense_direct_flops > plan.reduced_iterative_flops * 100);
1062        // The recommended (iterative) path is device-favorable at this shape:
1063        // the wide border × summed CG work clears the matvec offload floor.
1064        assert!(plan.device_favorable);
1065    }
1066
1067    /// A modest, near-square border where the data-fit rank is comparable to `k`
1068    /// and the `k³/3` Cholesky is cheap: dense Direct is the right call.
1069    #[test]
1070    fn small_square_border_recommends_dense_direct() {
1071        let pol = GpuDispatchPolicy::default();
1072        // n·d = 400 > k = 64: not rank-deficient; a 64³/3 Cholesky is trivial.
1073        let plan = pol.arrow_border_solve_plan(200, 64, 2, 8);
1074        assert_eq!(plan.recommended, ArrowBorderStrategy::DenseDirect);
1075        assert!(!plan.dense_border_rank_deficient);
1076        assert_eq!(plan.data_fit_rank, 64);
1077    }
1078
1079    /// The rank-deficiency flag is exactly `n·d < k`, and `data_fit_rank` is
1080    /// clamped at `k` (the border can carry no more than `k` data directions).
1081    #[test]
1082    fn rank_flag_and_clamp_track_n_d_versus_k() {
1083        let pol = GpuDispatchPolicy::default();
1084        // n·d == k exactly: full-rank border, not deficient.
1085        let exact = pol.arrow_border_solve_plan(50, 100, 2, 8);
1086        assert!(!exact.dense_border_rank_deficient);
1087        assert_eq!(exact.data_fit_rank, 100);
1088        // n·d one below k: deficient.
1089        let deficient = pol.arrow_border_solve_plan(49, 100, 2, 8);
1090        assert!(deficient.dense_border_rank_deficient);
1091        assert_eq!(deficient.data_fit_rank, 98);
1092    }
1093
1094    /// The recommendation is monotone toward `ReducedIterative` as the border
1095    /// widens at fixed row work: once the dense `k³` term overtakes the linear-
1096    /// in-`k` iterative cost, growing `k` keeps it recommending iterative.
1097    #[test]
1098    fn wider_border_only_moves_toward_iterative() {
1099        let pol = GpuDispatchPolicy::default();
1100        let narrow = pol.arrow_border_solve_plan(200, 128, 4, 16);
1101        let wide = pol.arrow_border_solve_plan(200, 8_192, 4, 16);
1102        // The wide border must recommend iterative.
1103        assert_eq!(wide.recommended, ArrowBorderStrategy::ReducedIterative);
1104        // If the narrow one already recommends iterative, the wide one still
1105        // does (monotone); if not, the wide one is a strict switch. Either way
1106        // the wide border's dense/iterative flop ratio exceeds the narrow one's.
1107        let narrow_ratio = narrow.dense_direct_flops as f64 / narrow.reduced_iterative_flops as f64;
1108        let wide_ratio = wide.dense_direct_flops as f64 / wide.reduced_iterative_flops as f64;
1109        assert!(wide_ratio > narrow_ratio);
1110    }
1111
1112    /// A larger CG budget makes the iterative path more expensive, so the
1113    /// crossover can only move toward `DenseDirect`, never away from it. If a
1114    /// shape is `DenseDirect` at a small budget it stays `DenseDirect` at a
1115    /// larger one.
1116    #[test]
1117    fn larger_cg_budget_never_switches_away_from_dense() {
1118        let pol = GpuDispatchPolicy::default();
1119        let shape = (200usize, 96usize, 3usize);
1120        let small = pol.arrow_border_solve_plan(shape.0, shape.1, shape.2, 4);
1121        let large = pol.arrow_border_solve_plan(shape.0, shape.1, shape.2, 400);
1122        if small.recommended == ArrowBorderStrategy::DenseDirect {
1123            assert_eq!(large.recommended, ArrowBorderStrategy::DenseDirect);
1124        }
1125        assert!(large.reduced_iterative_flops >= small.reduced_iterative_flops);
1126    }
1127
1128    /// Degenerate shapes yield an all-zero plan on the trivial `DenseDirect`
1129    /// path and are never device-favorable.
1130    #[test]
1131    fn degenerate_shapes_are_trivial_dense_and_not_device_favorable() {
1132        let pol = GpuDispatchPolicy::default();
1133        for shape in [(0usize, 100usize, 2usize), (100, 0, 2), (100, 100, 0)] {
1134            let plan = pol.arrow_border_solve_plan(shape.0, shape.1, shape.2, 8);
1135            assert_eq!(plan.recommended, ArrowBorderStrategy::DenseDirect);
1136            assert!(!plan.device_favorable);
1137            assert_eq!(plan.dense_direct_flops, 0);
1138            assert_eq!(plan.reduced_iterative_flops, 0);
1139        }
1140    }
1141
1142    /// A zero CG budget is treated as one apply (a plan must still be
1143    /// comparable), never a divide-by-zero or an all-free iterative path.
1144    #[test]
1145    fn zero_cg_budget_is_treated_as_one_apply() {
1146        let pol = GpuDispatchPolicy::default();
1147        let plan = pol.arrow_border_solve_plan(180, 15_360, 2, 0);
1148        assert_eq!(plan.cg_iters, 1);
1149        assert!(plan.reduced_iterative_flops > 0);
1150    }
1151}
1152
1153#[cfg(test)]
1154mod encode_deployment_decision_tests {
1155    use super::*;
1156
1157    /// #1412 anti-green-wash core: a CPU rate can NEVER produce a `Met`/`Unmet`
1158    /// decision. The only Met/Unmet constructor requires `engaged == true`; a
1159    /// CPU-only host has no device measurement, so it can only ever be
1160    /// `Undetermined`, no matter how fast the CPU is.
1161    #[test]
1162    fn cpu_rate_can_never_meet_or_refute_the_target() {
1163        // Even a CPU rate a thousand times the target cannot certify the gate:
1164        // there is simply no `from_cpu_measurement` — the type has no such door.
1165        // The blocked constructor is the only CPU-side option.
1166        let cpu_only = EncodeDeploymentDecision::blocked(EncodeDecisionBlocked::NoDevice);
1167        assert!(cpu_only.is_undetermined());
1168        assert!(!cpu_only.surrogate_unneeded());
1169        assert!(!cpu_only.surrogate_justified());
1170
1171        // A "device" measurement that did not engage (false routing) is refused —
1172        // it becomes Undetermined even with a huge rate.
1173        let false_routed = EncodeDeploymentDecision::from_device_measurement(false, 1.0e9);
1174        assert!(false_routed.is_undetermined());
1175        assert!(!false_routed.surrogate_unneeded());
1176    }
1177
1178    #[test]
1179    fn engaged_measurement_decides_by_the_number() {
1180        let target = GPU_THROUGHPUT_TARGET_ROWS_PER_SEC;
1181        // Clears the target => Met => surrogate unneeded.
1182        let met = EncodeDeploymentDecision::from_device_measurement(true, target * 2.0);
1183        assert!(matches!(met, EncodeDeploymentDecision::Met { .. }));
1184        assert!(met.surrogate_unneeded());
1185        assert!(!met.surrogate_justified());
1186        assert!(!met.is_undetermined());
1187
1188        // Misses the target => Unmet => surrogate justified.
1189        let unmet = EncodeDeploymentDecision::from_device_measurement(true, target * 0.25);
1190        assert!(matches!(unmet, EncodeDeploymentDecision::Unmet { .. }));
1191        assert!(unmet.surrogate_justified());
1192        assert!(!unmet.surrogate_unneeded());
1193
1194        // Exact boundary meets the target.
1195        let boundary = EncodeDeploymentDecision::from_device_measurement(true, target);
1196        assert!(boundary.surrogate_unneeded());
1197    }
1198
1199    #[test]
1200    fn engaged_but_non_usable_rate_is_undetermined_not_a_pass() {
1201        for bad in [0.0, -1.0, f64::NAN, f64::INFINITY] {
1202            let d = EncodeDeploymentDecision::from_device_measurement(true, bad);
1203            assert!(
1204                d.is_undetermined(),
1205                "an engaged-but-unusable rate {bad} must be Undetermined, not a decision"
1206            );
1207            assert!(!d.surrogate_unneeded());
1208            assert!(!d.surrogate_justified());
1209        }
1210    }
1211
1212    #[test]
1213    fn blocked_reasons_are_all_undetermined() {
1214        for reason in [
1215            EncodeDecisionBlocked::NoDevice,
1216            EncodeDecisionBlocked::NoDeviceEncodeKernel,
1217            EncodeDecisionBlocked::DeviceNotEngaged,
1218        ] {
1219            let d = EncodeDeploymentDecision::blocked(reason);
1220            assert!(d.is_undetermined());
1221            assert!(!d.surrogate_unneeded());
1222            assert!(!d.surrogate_justified());
1223        }
1224    }
1225}