Skip to main content

gam_gpu/
calibration.rs

1use crate::device::GpuDeviceInfo;
2use crate::gpu_error::GpuError;
3use crate::policy::GpuDispatchPolicy;
4use faer::Side;
5use gam_linalg::faer_ndarray::FaerCholesky;
6use gam_runtime::warm_start::{Fingerprint, Fingerprinter};
7use ndarray::{Array1, Array2};
8use serde::{Deserialize, Serialize};
9use std::fs;
10use std::path::PathBuf;
11use std::time::Instant;
12
13const SCHEMA_VERSION: u32 = 1;
14const CACHE_ROOT_COMPONENTS: [&str; 4] = ["gam", "gpu", "policy", "v1"];
15const GEMM_DIMS: [usize; 3] = [64, 128, 256];
16const POTRF_DIMS: [usize; 3] = [64, 128, 256];
17const XTWX_DIMS: [(usize, usize); 3] = [(2048, 32), (4096, 64), (8192, 96)];
18const GPU_WIN_RATIO: f64 = 0.95;
19
20// Pin the pre-probe admission floors to the smallest calibration measurements:
21// `crossover_flops` / `crossover_rows` can lower `gemm_min_flops` /
22// `potrf_min_p` at most to the smallest measured GEMM's flop count / smallest
23// POTRF dimension. `DispatchOp::admissible_under_any_policy` (and every
24// pre-probe size gate built on these constants) relies on that bound, so a
25// change to the measurement grid must consciously update the constants.
26const _: () = assert!(
27    2 * (GEMM_DIMS[0] as u128) * (GEMM_DIMS[0] as u128) * (GEMM_DIMS[0] as u128)
28        == GpuDispatchPolicy::MIN_CALIBRATABLE_GEMM_FLOPS
29);
30const _: () = assert!(POTRF_DIMS[0] == GpuDispatchPolicy::MIN_CALIBRATABLE_POTRF_P);
31const _: () = assert!(XTWX_DIMS[0].0 == GpuDispatchPolicy::MIN_CALIBRATABLE_ROW_KERNEL_N);
32
33#[derive(Clone, Debug, Serialize, Deserialize)]
34struct CachedCalibration {
35    schema_version: u32,
36    device_fingerprint: String,
37    policy: GpuDispatchPolicy,
38    measurements: Vec<MeasurementRecord>,
39}
40
41#[derive(Clone, Debug, Serialize, Deserialize)]
42struct MeasurementRecord {
43    operation: String,
44    rows: usize,
45    cols: usize,
46    inner: usize,
47    flops: usize,
48    cpu_seconds: f64,
49    gpu_seconds: f64,
50}
51
52#[derive(Clone, Debug)]
53struct Measurement {
54    operation: &'static str,
55    rows: usize,
56    cols: usize,
57    inner: usize,
58    flops: usize,
59    cpu_seconds: f64,
60    gpu_seconds: f64,
61}
62
63pub(crate) fn calibrated_policy_for_device(device: &GpuDeviceInfo) -> GpuDispatchPolicy {
64    let fingerprint = device_fingerprint(device);
65    if let Some(cached) = load_cached_policy(fingerprint) {
66        log::info!(
67            "[GPU] loaded calibrated dispatch policy for {} ({fingerprint})",
68            device.name
69        );
70        return cached;
71    }
72
73    match calibrate_device(device, fingerprint) {
74        Ok(record) => {
75            let policy = record.policy.clone();
76            store_cached_policy(fingerprint, &record);
77            policy
78        }
79        Err(err) => {
80            log::warn!(
81                "[GPU] dispatch calibration unavailable for {}: {}; using default policy",
82                device.name,
83                err
84            );
85            GpuDispatchPolicy::default()
86        }
87    }
88}
89
90fn calibrate_device(
91    device: &GpuDeviceInfo,
92    fingerprint: Fingerprint,
93) -> Result<CachedCalibration, GpuError> {
94    let mut measurements = Vec::new();
95    measurements.extend(measure_gemm(device.ordinal)?);
96    measurements.extend(measure_potrf(device.ordinal)?);
97    measurements.extend(measure_xtwx(device.ordinal)?);
98    if measurements.is_empty() {
99        return Err(GpuError::CalibrationFailed {
100            reason: "no GPU calibration measurements completed".to_string(),
101        });
102    }
103
104    let mut policy = GpuDispatchPolicy::default();
105    if let Some(flops) = crossover_flops(&measurements, "gemm", policy.gemm_min_flops) {
106        policy.gemm_min_flops = flops;
107    }
108    if let Some(flops) = crossover_flops(&measurements, "xtwx", policy.xtwx_flops_min) {
109        policy.xtwx_flops_min = flops;
110    }
111    if let Some(rows) = crossover_rows(&measurements, "xtwx", policy.xtwx_n_min) {
112        policy.xtwx_n_min = rows;
113        policy.row_kernel_min_n = rows;
114        policy.fused_kernel_min_n = rows.saturating_mul(2);
115    }
116    if let Some(p) = crossover_rows(&measurements, "potrf", policy.potrf_min_p) {
117        policy.potrf_min_p = p;
118        policy.prefer_gpu_factorization_min_p = p;
119    }
120
121    log::info!(
122        "[GPU] calibrated dispatch policy for {} ({fingerprint}) from {} measurements",
123        device.name,
124        measurements.len()
125    );
126
127    Ok(CachedCalibration {
128        schema_version: SCHEMA_VERSION,
129        device_fingerprint: fingerprint.to_hex(),
130        policy,
131        measurements: measurements
132            .into_iter()
133            .map(Measurement::into_record)
134            .collect(),
135    })
136}
137
138fn measure_gemm(ordinal: usize) -> Result<Vec<Measurement>, GpuError> {
139    let mut out = Vec::with_capacity(GEMM_DIMS.len());
140    for dim in GEMM_DIMS {
141        let a = deterministic_matrix(dim, dim, 0.13);
142        let b = deterministic_matrix(dim, dim, 0.37);
143        let cpu_seconds = time_cpu(|| a.dot(&b))?;
144        let gpu_seconds = time_gpu(|| {
145            crate::blas::gemm_on_ordinal_cuda(ordinal, a.view(), b.view(), false, false)
146        })?;
147        out.push(Measurement {
148            operation: "gemm",
149            rows: dim,
150            cols: dim,
151            inner: dim,
152            flops: 2usize
153                .saturating_mul(dim)
154                .saturating_mul(dim)
155                .saturating_mul(dim),
156            cpu_seconds,
157            gpu_seconds,
158        });
159    }
160    Ok(out)
161}
162
163fn measure_potrf(ordinal: usize) -> Result<Vec<Measurement>, GpuError> {
164    let mut out = Vec::with_capacity(POTRF_DIMS.len());
165    for dim in POTRF_DIMS {
166        let a = deterministic_spd_matrix(dim);
167        let cpu_seconds = time_gpu_result(|| {
168            a.cholesky(Side::Lower)
169                .map(|factor| factor.lower_triangular())
170                .map_err(|err| format!("cpu POTRF failed: {err}"))
171        })?;
172        let gpu_seconds =
173            time_gpu_result(|| crate::solver::cholesky_lower_on_ordinal_gpu(ordinal, a.view()))?;
174        out.push(Measurement {
175            operation: "potrf",
176            rows: dim,
177            cols: dim,
178            inner: dim,
179            flops: dim.saturating_mul(dim).saturating_mul(dim) / 3,
180            cpu_seconds,
181            gpu_seconds,
182        });
183    }
184    Ok(out)
185}
186
187fn measure_xtwx(ordinal: usize) -> Result<Vec<Measurement>, GpuError> {
188    let mut out = Vec::with_capacity(XTWX_DIMS.len());
189    for (n, p) in XTWX_DIMS {
190        let x = deterministic_matrix(n, p, 0.61);
191        let w = deterministic_weights(n);
192        let cpu_seconds = time_cpu(|| cpu_xtwx(&x, &w))?;
193        let gpu_seconds =
194            time_gpu(|| crate::blas::xt_diag_x_on_ordinal_cuda(ordinal, x.view(), w.view()))?;
195        out.push(Measurement {
196            operation: "xtwx",
197            rows: n,
198            cols: p,
199            inner: p,
200            flops: 2usize.saturating_mul(n).saturating_mul(p).saturating_mul(p),
201            cpu_seconds,
202            gpu_seconds,
203        });
204    }
205    Ok(out)
206}
207
208fn time_cpu<F>(mut f: F) -> Result<f64, GpuError>
209where
210    F: FnMut() -> Array2<f64>,
211{
212    time_gpu_result(|| Result::<Array2<f64>, GpuError>::Ok(f()))
213}
214
215fn time_gpu<F>(mut f: F) -> Result<f64, GpuError>
216where
217    F: FnMut() -> Option<Array2<f64>>,
218{
219    time_gpu_result(|| {
220        f().ok_or_else(|| GpuError::CalibrationFailed {
221            reason: "GPU calibration kernel returned no result".to_string(),
222        })
223    })
224}
225
226fn time_gpu_result<F, E>(mut f: F) -> Result<f64, GpuError>
227where
228    F: FnMut() -> Result<Array2<f64>, E>,
229    E: std::fmt::Display,
230{
231    let start = Instant::now();
232    let out = f().map_err(|err| GpuError::CalibrationFailed {
233        reason: err.to_string(),
234    })?;
235    let elapsed = start.elapsed().as_secs_f64();
236    let checksum = out.iter().fold(0.0, |acc, value| acc + value.abs());
237    if elapsed.is_finite() && elapsed > 0.0 && checksum.is_finite() {
238        Ok(elapsed)
239    } else {
240        Err(GpuError::CalibrationFailed {
241            reason: format!(
242                "invalid calibration timing/checksum: elapsed={elapsed}, checksum={checksum}"
243            ),
244        })
245    }
246}
247
248fn crossover_flops(
249    measurements: &[Measurement],
250    operation: &'static str,
251    fallback: usize,
252) -> Option<usize> {
253    crossover_measurement(measurements, operation)
254        .map(|measurement| measurement.flops.max(1))
255        .or_else(|| {
256            measurements
257                .iter()
258                .filter(|measurement| measurement.operation == operation)
259                .map(|measurement| measurement.flops)
260                .max()
261                .map(|max_seen| fallback.max(max_seen.saturating_mul(2)))
262        })
263}
264
265fn crossover_rows(
266    measurements: &[Measurement],
267    operation: &'static str,
268    fallback: usize,
269) -> Option<usize> {
270    crossover_measurement(measurements, operation)
271        .map(|measurement| measurement.rows.max(1))
272        .or_else(|| {
273            measurements
274                .iter()
275                .filter(|measurement| measurement.operation == operation)
276                .map(|measurement| measurement.rows)
277                .max()
278                .map(|max_seen| fallback.max(max_seen.saturating_mul(2)))
279        })
280}
281
282fn crossover_measurement<'a>(
283    measurements: &'a [Measurement],
284    operation: &'static str,
285) -> Option<&'a Measurement> {
286    measurements
287        .iter()
288        .filter(|measurement| measurement.operation == operation)
289        .find(|measurement| measurement.gpu_seconds <= measurement.cpu_seconds * GPU_WIN_RATIO)
290}
291
292fn deterministic_matrix(rows: usize, cols: usize, phase: f64) -> Array2<f64> {
293    Array2::from_shape_fn((rows, cols), |(row, col)| {
294        let x = (row as f64 + 1.0) * 0.017 + (col as f64 + 1.0) * 0.031 + phase;
295        x.sin() + 0.25 * (2.0 * x).cos()
296    })
297}
298
299fn deterministic_spd_matrix(dim: usize) -> Array2<f64> {
300    let a = deterministic_matrix(dim, dim, 0.89);
301    let mut spd = a.t().dot(&a);
302    for idx in 0..dim {
303        spd[[idx, idx]] += dim as f64;
304    }
305    spd
306}
307
308fn deterministic_weights(n: usize) -> Array1<f64> {
309    Array1::from_shape_fn(n, |idx| 0.5 + ((idx as f64 + 1.0) * 0.019).sin().abs())
310}
311
312fn cpu_xtwx(x: &Array2<f64>, w: &Array1<f64>) -> Array2<f64> {
313    let mut weighted = x.clone();
314    for (mut row, weight) in weighted.outer_iter_mut().zip(w.iter()) {
315        row *= *weight;
316    }
317    x.t().dot(&weighted)
318}
319
320fn load_cached_policy(fingerprint: Fingerprint) -> Option<GpuDispatchPolicy> {
321    let path = cache_path(fingerprint);
322    let bytes = fs::read(path).ok()?;
323    let record: CachedCalibration = serde_json::from_slice(&bytes).ok()?;
324    if record.schema_version == SCHEMA_VERSION && record.device_fingerprint == fingerprint.to_hex()
325    {
326        Some(record.policy)
327    } else {
328        None
329    }
330}
331
332fn store_cached_policy(fingerprint: Fingerprint, record: &CachedCalibration) {
333    let path = cache_path(fingerprint);
334    if let Some(parent) = path.parent() {
335        if let Err(err) = fs::create_dir_all(parent) {
336            log::warn!("[GPU] unable to create calibration cache dir: {err}");
337            return;
338        }
339    }
340    let tmp = path.with_extension("json.tmp");
341    let bytes = match serde_json::to_vec_pretty(record) {
342        Ok(bytes) => bytes,
343        Err(err) => {
344            log::warn!("[GPU] unable to serialize calibration cache: {err}");
345            return;
346        }
347    };
348    if let Err(err) = fs::write(&tmp, bytes).and_then(|_| fs::rename(&tmp, &path)) {
349        log::warn!("[GPU] unable to write calibration cache: {err}");
350    }
351}
352
353fn cache_path(fingerprint: Fingerprint) -> PathBuf {
354    let mut root = std::env::temp_dir();
355    for component in CACHE_ROOT_COMPONENTS {
356        root.push(component);
357    }
358    root.push(format!("{fingerprint}.json"));
359    root
360}
361
362fn device_fingerprint(device: &GpuDeviceInfo) -> Fingerprint {
363    let mut fp = Fingerprinter::new();
364    fp.absorb_tag(b"gpu-dispatch-calibration");
365    fp.absorb_u64(b"schema-version", u64::from(SCHEMA_VERSION));
366    fp.absorb_str(b"name", &device.name);
367    fp.absorb_u64(
368        b"compute-major",
369        u64::try_from(device.capability.compute_major).unwrap_or(0),
370    );
371    fp.absorb_u64(
372        b"compute-minor",
373        u64::try_from(device.capability.compute_minor).unwrap_or(0),
374    );
375    fp.absorb_u64(b"sm-count", u64::try_from(device.sm_count).unwrap_or(0));
376    fp.absorb_u64(
377        b"max-threads-per-sm",
378        u64::try_from(device.max_threads_per_sm).unwrap_or(0),
379    );
380    fp.absorb_u64(
381        b"max-shared-mem-per-block",
382        device.max_shared_mem_per_block as u64,
383    );
384    fp.absorb_u64(b"l2-cache-bytes", device.l2_cache_bytes as u64);
385    fp.absorb_u64(b"total-mem-bytes", device.total_mem_bytes as u64);
386    fp.absorb_u64(b"ecc-enabled", bool_fingerprint_value(device.ecc_enabled));
387    fp.absorb_u64(b"integrated", bool_fingerprint_value(device.integrated));
388    fp.absorb_u64(b"mig-mode", bool_fingerprint_value(device.mig_mode));
389    fp.finalize()
390}
391
392const fn bool_fingerprint_value(value: bool) -> u64 {
393    if value { 1 } else { 0 }
394}
395
396impl Measurement {
397    fn into_record(self) -> MeasurementRecord {
398        MeasurementRecord {
399            operation: self.operation.to_string(),
400            rows: self.rows,
401            cols: self.cols,
402            inner: self.inner,
403            flops: self.flops,
404            cpu_seconds: self.cpu_seconds,
405            gpu_seconds: self.gpu_seconds,
406        }
407    }
408}
409
410#[cfg(test)]
411mod tests {
412    use super::*;
413    use crate::device::GpuCapability;
414
415    fn measurement(
416        operation: &'static str,
417        rows: usize,
418        cols: usize,
419        flops: usize,
420        cpu_seconds: f64,
421        gpu_seconds: f64,
422    ) -> Measurement {
423        Measurement {
424            operation,
425            rows,
426            cols,
427            inner: cols,
428            flops,
429            cpu_seconds,
430            gpu_seconds,
431        }
432    }
433
434    #[test]
435    fn calibration_crossover_uses_first_measured_gpu_win() {
436        let measurements = vec![
437            measurement("gemm", 64, 64, 524_288, 0.001, 0.004),
438            measurement("gemm", 128, 128, 4_194_304, 0.010, 0.009),
439            measurement("gemm", 256, 256, 33_554_432, 0.080, 0.010),
440        ];
441
442        assert_eq!(
443            crossover_flops(&measurements, "gemm", 100_000_000),
444            Some(4_194_304)
445        );
446    }
447
448    #[test]
449    fn calibration_crossover_raises_threshold_when_gpu_never_wins() {
450        let measurements = vec![
451            measurement("xtwx", 2_048, 32, 4_194_304, 0.001, 0.004),
452            measurement("xtwx", 4_096, 64, 33_554_432, 0.010, 0.040),
453            measurement("xtwx", 8_192, 96, 150_994_944, 0.080, 0.400),
454        ];
455
456        assert_eq!(
457            crossover_flops(&measurements, "xtwx", 100_000_000),
458            Some(301_989_888)
459        );
460        assert_eq!(crossover_rows(&measurements, "xtwx", 50_000), Some(50_000));
461    }
462
463    #[test]
464    fn calibration_cache_key_tracks_device_fingerprint() {
465        let device = GpuDeviceInfo {
466            ordinal: 0,
467            name: "unit-test GPU".to_string(),
468            capability: GpuCapability::from_compute_capability(8, 0),
469            sm_count: 108,
470            max_threads_per_sm: 2048,
471            max_shared_mem_per_block: 99_328,
472            l2_cache_bytes: 40 * 1024 * 1024,
473            total_mem_bytes: 80 * 1024 * 1024 * 1024,
474            free_mem_bytes: 70 * 1024 * 1024 * 1024,
475            ecc_enabled: true,
476            integrated: false,
477            mig_mode: false,
478        };
479
480        let fingerprint = device_fingerprint(&device);
481        let path = cache_path(fingerprint);
482        assert!(path.ends_with(format!("{}.json", fingerprint.to_hex())));
483        assert!(
484            path.components()
485                .map(|component| component.as_os_str().to_string_lossy().into_owned())
486                .collect::<Vec<_>>()
487                .windows(CACHE_ROOT_COMPONENTS.len())
488                .any(|window| window == CACHE_ROOT_COMPONENTS)
489        );
490    }
491}