metal-linalg

Eigensolver routing on Apple M5 Pro (20 GPU cores)

What every section and number below means: reading-reports.md.

Cost model scaled to this device from one probe point: block x0.14, cpu x0.94, whole-matrix x1.00 (1.00 is an M1).

Machine state: load 3.6/18 at the start, load 1.5/18 at the end; power mains.

Probe point after the sweep relative to before it: block x1.01, cpu x1.01, tg x1.01 (stable).

Generated by tuning/tune_eigh.py from 205 (N, batch) points, N in [2, 4, 8, 12, 16, 24, 32, 48, 64, 96, 128, 192, 256, 384, 512, 768, 1024, 1536, 2048, 3072, 4096], batch in [1, 2, 4, 8, 16, 32, 64, 128, 256, 512, 1024, 2048, 4096], four backends, two or more passes, min-of-repeats.

Answer

Row for kTuned[] in src/eigh.mm:

// device, GPU cores,   simd_max_n, block_min_n, block_min_n_batched, block_min_batch,   gpu_max_n, gpu_min_batch_times_n, gpu_min_batch,   values_gpu_max_n, values_gpu_min_batch_times_n, values_gpu_min_batch,   tridiag_min_n, values_tridiag_min_n
{"Apple M5 Pro", 20,   0, 96, 0, 0,   1024, 512, 16,   256, 2048, 32,   1024, 0},

To try it without rebuilding:

EIGH_SIMD_MAX_N=0 EIGH_BLOCK_MIN_N=96 EIGH_BLOCK_MIN_N_BATCHED=0 EIGH_BLOCK_MIN_BATCH=0 EIGH_GPU_MAX_N=1024 EIGH_GPU_MIN_BATCH_TIMES_N=512 EIGH_GPU_MIN_BATCH=16 EIGH_TRIDIAG_MIN_N=1024 EIGH_VALUES_TRIDIAG_MIN_N=0 EIGH_VALUES_GPU_MAX_N=256 EIGH_VALUES_GPU_MIN_BATCH_TIMES_N=2048 EIGH_VALUES_GPU_MIN_BATCH=32

The policy in effect on this device came from tuned:Apple M5 Pro. It differs from the fitted one; see the warnings.

Against the best measured backend at every point the whole rule scores 1.0275 geometric-mean regret, worst 1.67x, 15 of 205 points losing more than 10%, and 1.319x the oracle’s total time. The decision is fitted in two stages, below, because the CPU routing would otherwise hide the GPU backend crossover.

Warnings

Stage 1: which GPU backend

Scored against the best GPU backend at each of the 205 points, as if there were no CPU: this is the rule a forced-GPU call (EIGH_DEVICE=gpu) and the detail entry points follow, and it is what a GPU with more cores will lean on.

rule geomean regret worst >10% total time / oracle est. picks
policy in effect (‘0’, ‘96’, ‘none’, ‘none’) 1.0180 1.55x 10 1.003 0
fitted (‘0’, ‘96’, ‘none’, ‘none’) 1.0180 1.55x 10 1.003 0

2 of 128 (simd_max_n, block_min_n) pairs are within 0.5% of the best geomean: simd_max_n 0 .. 2, block_min_n 96 .. 96.

xychart-beta
    title "Regret by block_min_n"
    x-axis "block_min_n" [32, 48, 64, 96, 128, 192, 256, 384, 512, 768, 1024, 1536, 2048, 3072, 4096, none]
    y-axis "geometric-mean regret" 1.0 --> 2.53
    line [1.1934, 1.0992, 1.0359, 1.0180, 1.0425, 1.1121, 1.2376, 1.3783, 1.5490, 1.7218, 1.9334, 2.1172, 2.2551, 2.4010, 2.4591, 2.5187]
block_min_n 32 48 64 96 128 192 256 384 512 768 1024 1536 2048 3072 4096 none
geomean 1.1934 1.0992 1.0359 1.0180 1.0425 1.1121 1.2376 1.3783 1.5490 1.7218 1.9334 2.1172 2.2551 2.4010 2.4591 2.5187
worst 8.70x 4.49x 2.58x 1.55x 1.92x 2.24x 4.59x 5.66x 10.11x 15.91x 15.91x 15.91x 15.91x 15.91x 15.91x 15.91x
xychart-beta
    title "Regret by simd_max_n"
    x-axis "simd_max_n" [0, 2, 4, 8, 12, 16, 24, 32]
    y-axis "geometric-mean regret" 1.0 --> 1.20
    line [1.0180, 1.0200, 1.0255, 1.0345, 1.0491, 1.0726, 1.1149, 1.1824]
simd_max_n 0 2 4 8 12 16 24 32
geomean 1.0180 1.0200 1.0255 1.0345 1.0491 1.0726 1.1149 1.1824
worst 1.55x 1.55x 1.90x 2.42x 3.18x 3.66x 3.66x 4.98x

Held-out check of a batch-dependent crossover (block from a lower N once the batch is large enough). Fitted on 112 points, scored on the other 93; the verdict is a bootstrap over the test points.

rule fitted on train train geomean test geomean test worst verdict
two constants [0, 96, 1000000000, 1000000000] 1.0225 1.0127 1.55x baseline
batch-dependent block crossover {“block_lo”: 64, “batch_hi”: 128} 1.0020 1.0064 1.51x rejected (better in 78% of resamples, median gain 0.6%)

Best GPU backend per point (s simd, t threadgroup, B block), then what the split picks:

  N \ batch     1     2     4     8    16    32    64   128   256   512  1024  2048  4096
          2     s     t     t     s     t     t     s     t     t     t     t     t     s
          4     s     t     t     t     t     t     s     t     t     s     s     s     t
          8     t     t     s     t     t     s     t     s     t     t     t     t     s
         12     t     t     t     t     t     t     t     t     t     t     t     s     t
         16     t     t     t     t     t     t     t     t     t     t     t     t     t
         24     t     t     t     t     t     t     t     t     t     t     t     t     t
         32     t     t     t     t     t     t     t     t     t     t     t     t     t
         48     t     t     t     t     t     t     t     t     t     t     t     t     t
         64     t     t     t     t     t     t     t     t     B     B     B     B     B
         96     t     t     t     t     t     B     B     B     B     B     B     B     B
        128     B     B     B     B     B     B     B     B     B     B     B     B     B
        192     B     B     B     B     B     B     B     B     B     B     B     B     B
        256     B     B     B     B     B     B     B     B     B     B     B     .     .
        384     B     B     B     B     B     B     B     B     B     B     .     .     .
        512     B     B     B     B     B     B     B     B     .     .     .     .     .
        768     B     B     B     B     B     B     B     .     .     .     .     .     .
       1024     B     B     B     B     B     .     .     .     .     .     .     .     .
       1536     B     B     B     .     .     .     .     .     .     .     .     .     .
       2048     B     B     B     .     .     .     .     .     .     .     .     .     .
       3072     B     .     .     .     .     .     .     .     .     .     .     .     .
       4096     B     .     .     .     .     .     .     .     .     .     .     .     .
  N \ batch     1     2     4     8    16    32    64   128   256   512  1024  2048  4096
          2     t     t     t     t     t     t     t     t     t     t     t     t     t
          4     t     t     t     t     t     t     t     t     t     t     t     t     t
          8     t     t     t     t     t     t     t     t     t     t     t     t     t
         12     t     t     t     t     t     t     t     t     t     t     t     t     t
         16     t     t     t     t     t     t     t     t     t     t     t     t     t
         24     t     t     t     t     t     t     t     t     t     t     t     t     t
         32     t     t     t     t     t     t     t     t     t     t     t     t     t
         48     t     t     t     t     t     t     t     t     t     t     t     t     t
         64     t     t     t     t     t     t     t     t     t     t     t     t     t
         96     B     B     B     B     B     B     B     B     B     B     B     B     B
        128     B     B     B     B     B     B     B     B     B     B     B     B     B
        192     B     B     B     B     B     B     B     B     B     B     B     B     B
        256     B     B     B     B     B     B     B     B     B     B     B     .     .
        384     B     B     B     B     B     B     B     B     B     B     .     .     .
        512     B     B     B     B     B     B     B     B     .     .     .     .     .
        768     B     B     B     B     B     B     B     .     .     .     .     .     .
       1024     B     B     B     B     B     .     .     .     .     .     .     .     .
       1536     B     B     B     .     .     .     .     .     .     .     .     .     .
       2048     B     B     B     .     .     .     .     .     .     .     .     .     .
       3072     B     .     .     .     .     .     .     .     .     .     .     .     .
       4096     B     .     .     .     .     .     .     .     .     .     .     .     .

Stage 2: GPU or CPU

Given the split above, GPU iff N <= gpu_max_n, batch * N >= gpu_min_batch_times_n and batch >= gpu_min_batch, scored against the best of all four backends. worst is over the points where the chosen backend was timed; a pick the cost model had to guess is listed in the warnings instead.

rule geomean regret worst >10% total time / oracle est. picks
oracle (best per point) 1.0000 1.00x 0 1.000 0
policy in effect (‘1024’, ‘512’, ‘16’) 1.0275 1.67x 15 1.319 1
fitted (‘1024’, ‘512’, ‘16’) 1.0275 1.67x 15 1.319 1

6 of 1188 combinations are within 0.5% of the best geomean: gpu_max_n 1024 .. none, gpu_min_batch_times_n 512 .. 512, gpu_min_batch 16 .. 16.

xychart-beta
    title "Regret by gpu_min_batch_times_n"
    x-axis "gpu_min_batch_times_n" [0, 64, 128, 256, 512, 1024, 2048, 4096, 8192, 16384, none]
    y-axis "geometric-mean regret" 1.0 --> 1.97
    line [1.1217, 1.1053, 1.0800, 1.0483, 1.0275, 1.0333, 1.0648, 1.1267, 1.2227, 1.3501, 1.9557]
gpu_min_batch_times_n 0 64 128 256 512 1024 2048 4096 8192 16384 none
geomean 1.1217 1.1053 1.0800 1.0483 1.0275 1.0333 1.0648 1.1267 1.2227 1.3501 1.9557
worst 20.66x 11.52x 6.68x 3.35x 1.67x 2.18x 3.49x 5.26x 7.03x 9.75x 16.90x
xychart-beta
    title "Regret by gpu_min_batch"
    x-axis "gpu_min_batch" [1, 2, 4, 8, 16, 32]
    y-axis "geometric-mean regret" 1.0 --> 1.11
    line [1.0929, 1.0753, 1.0557, 1.0376, 1.0275, 1.0446]
gpu_min_batch 1 2 4 8 16 32
geomean 1.0929 1.0753 1.0557 1.0376 1.0275 1.0446
worst 3.32x 3.00x 3.00x 2.11x 1.67x 1.76x
xychart-beta
    title "Regret by gpu_max_n"
    x-axis "gpu_max_n" [16, 24, 32, 48, 64, 96, 128, 192, 256, 384, 512, 768, 1024, 1536, 2048, 3072, 4096, none]
    y-axis "geometric-mean regret" 1.0 --> 1.57
    line [1.5578, 1.4512, 1.3574, 1.2796, 1.2197, 1.1678, 1.1222, 1.0838, 1.0600, 1.0461, 1.0410, 1.0330, 1.0275, 1.0275, 1.0275, 1.0275, 1.0275, 1.0275]
gpu_max_n 16 24 32 48 64 96 128 192 256 384 512 768 1024 1536 2048 3072 4096 none
geomean 1.5578 1.4512 1.3574 1.2796 1.2197 1.1678 1.1222 1.0838 1.0600 1.0461 1.0410 1.0330 1.0275 1.0275 1.0275 1.0275 1.0275 1.0275
worst 10.64x 8.09x 5.68x 5.68x 3.80x 3.00x 2.34x 2.00x 1.60x 1.60x 1.60x 1.67x 1.67x 1.67x 1.67x 1.67x 1.67x 1.67x

Held-out check of a per-N boundary (a lookup table of the smallest batch at which the GPU wins, per N) against the product rule. Fitted on 112 points, scored on the other 93.

rule fitted on train train geomean test geomean test worst verdict
product rule [1024, 512, 16] 1.0394 1.0134 1.55x baseline
per-N table {“min_batch_by_n”: {“2”: 1024, “4”: 256, “8”: 128, “12”: 64, “16”: 32, “24”: 64, “32”: 64, “48”: 32, “64”: 32, “96”: 32, “128”: 16, “192”: 16, “256”: 8, “384”: 8, “512”: 16, “768”: null, “1024”: 16, “1536”: null, “2048”: null, “3072”: null, “4096”: 1}} 1.0147 1.0707 2.52x rejected (better in 0% of resamples, median gain -5.2%)

Best backend per point (c CPU, s simd, t threadgroup, B block, . not measured), what the whole rule picks, and the speedup of the best GPU backend over the CPU:

  N \ batch     1     2     4     8    16    32    64   128   256   512  1024  2048  4096
          2     c     c     c     c     c     c     c     c     c     t     t     t     s
          4     c     c     c     c     c     c     c     c     t     s     s     s     t
          8     c     c     c     c     c     c     c     s     t     t     t     t     s
         12     c     c     c     c     c     c     t     t     t     t     t     s     t
         16     c     c     c     c     c     t     t     t     t     t     t     t     t
         24     c     c     c     c     t     t     t     t     t     t     t     t     t
         32     c     c     c     c     t     t     t     t     t     t     t     t     t
         48     c     c     c     c     t     t     t     t     t     t     t     t     t
         64     c     c     c     c     t     t     t     t     B     B     B     B     B
         96     c     c     c     c     t     B     B     B     B     B     B     B     B
        128     c     c     c     c     B     B     B     B     B     B     B     B     B
        192     c     c     c     c     B     B     B     B     B     B     B     B     B
        256     c     c     c     B     B     B     B     B     B     B     B     .     .
        384     c     c     c     B     B     B     B     B     B     B     .     .     .
        512     c     c     c     B     B     c     c     B     .     .     .     .     .
        768     c     c     c     c     c     B     B     .     .     .     .     .     .
       1024     c     c     c     c     B     .     .     .     .     .     .     .     .
       1536     c     c     c     .     .     .     .     .     .     .     .     .     .
       2048     c     c     c     .     .     .     .     .     .     .     .     .     .
       3072     c     .     .     .     .     .     .     .     .     .     .     .     .
       4096     B     .     .     .     .     .     .     .     .     .     .     .     .
  N \ batch     1     2     4     8    16    32    64   128   256   512  1024  2048  4096
          2     c     c     c     c     c     c     c     c     t     t     t     t     t
          4     c     c     c     c     c     c     c     t     t     t     t     t     t
          8     c     c     c     c     c     c     t     t     t     t     t     t     t
         12     c     c     c     c     c     c     t     t     t     t     t     t     t
         16     c     c     c     c     c     t     t     t     t     t     t     t     t
         24     c     c     c     c     c     t     t     t     t     t     t     t     t
         32     c     c     c     c     t     t     t     t     t     t     t     t     t
         48     c     c     c     c     t     t     t     t     t     t     t     t     t
         64     c     c     c     c     t     t     t     t     t     t     t     t     t
         96     c     c     c     c     B     B     B     B     B     B     B     B     B
        128     c     c     c     c     B     B     B     B     B     B     B     B     B
        192     c     c     c     c     B     B     B     B     B     B     B     B     B
        256     c     c     c     c     B     B     B     B     B     B     B     .     .
        384     c     c     c     c     B     B     B     B     B     B     .     .     .
        512     c     c     c     c     B     B     B     B     .     .     .     .     .
        768     c     c     c     c     B     B     B     .     .     .     .     .     .
       1024     c     c     c     c     B     .     .     .     .     .     .     .     .
       1536     c     c     c     .     .     .     .     .     .     .     .     .     .
       2048     c     c     c     .     .     .     .     .     .     .     .     .     .
       3072     c     .     .     .     .     .     .     .     .     .     .     .     .
       4096     c     .     .     .     .     .     .     .     .     .     .     .     .
  N \ batch     1     2     4     8    16    32    64   128   256   512  1024  2048  4096
          2  0.01  0.01  0.02  0.02  0.05  0.09  0.15  0.30  0.63  1.18  1.92  3.77  7.03
          4  0.01  0.02  0.03  0.05  0.10  0.16  0.31  0.69  1.14  2.32  4.35  7.30 11.17
          8  0.02  0.04  0.06  0.12  0.23  0.38  0.81  1.63  3.16  5.06  8.61 12.72 16.90
         12  0.03  0.06  0.10  0.19  0.39  0.72  1.44  2.77  4.59  7.03  9.75 11.65 14.03
         16  0.04  0.09  0.15  0.29  0.62  1.15  2.30  4.01  5.91  8.67  9.84 12.70 14.57
         24  0.08  0.17  0.30  0.57  1.13  2.18  3.49  5.26  6.33  8.26  9.19 10.06 10.64
         32  0.11  0.19  0.37  0.74  1.44  2.52  3.64  4.23  5.29  6.70  7.35  7.53  8.09
         48  0.10  0.22  0.43  0.83  1.76  2.53  3.07  3.96  4.69  5.03  5.14  5.41  5.08
         64  0.12  0.23  0.44  0.89  1.72  2.38  2.63  3.28  4.28  5.13  5.42  5.44  5.68
         96  0.08  0.16  0.31  0.62  1.23  1.65  2.36  3.09  3.68  3.66  3.70  3.69  3.80
        128  0.09  0.17  0.33  0.64  1.13  1.91  2.51  3.00  2.99  2.92  2.86  2.94  2.93
        192  0.13  0.26  0.49  0.91  1.59  1.98  2.34  2.28  2.12  2.10  2.13   gpu   gpu
        256  0.19  0.37  0.69  1.20  1.70  2.00  1.87  1.75  1.76  1.76   gpu     .     .
        384  0.23  0.41  0.70  1.13  1.32  1.19  1.17  1.07   gpu   gpu     .     .     .
        512  0.30  0.50  0.81  1.11  1.04  0.97  0.93   gpu     .     .     .     .     .
        768  0.31  0.50  0.70  0.63  0.60   gpu   gpu     .     .     .     .     .     .
       1024  0.39  0.60  0.64  0.59   gpu     .     .     .     .     .     .     .     .
       1536  0.40  0.44  0.40     .     .     .     .     .     .     .     .     .     .
       2048  0.39  0.39  0.39     .     .     .     .     .     .     .     .     .     .
       3072  0.33     .     .     .     .     .     .     .     .     .     .     .     .
       4096   gpu     .     .     .     .     .     .     .     .     .     .     .     .

Stage 3: GPU or CPU, eigenvalues alone

The same rule for eigvalsh, with its own thresholds (values_gpu_max_n, values_gpu_min_batch_times_n, values_gpu_min_batch), fitted on the _vals timings of 195 points given the split above. The CPU computes eigenvalues alone by LAPACK’s two-stage reduction from N = 128, so the boundary need not be eigh’s.

rule geomean regret worst >10% total time / oracle est. picks
policy in effect 1.0148 1.55x 10 1.013 0
eigh’s fitted boundary 1.0680 3.95x 29 1.096 0
fitted (‘256’, ‘2048’, ‘32’) 1.0148 1.55x 10 1.013 0

7 combinations are within 0.5% of the best geomean: values_gpu_max_n 192 .. 384, values_gpu_min_batch_times_n 1024 .. 2048, values_gpu_min_batch 16 .. 32.

Held out: fitted on 106 points (‘256’, ‘2048’, ‘32’), scored on the other 89: geomean 1.0204x, worst 1.53x, against 1.0566x, worst 2.19x for eigh’s boundary on the same points.

Stage 4: the tridiag backend instead of the CPU

Where the rule above chooses the CPU, the tridiag backend from a threshold N on (0: never), fitted over the measured N against the best of all backends, tridiag included, on the points where tridiag was timed (N >= 128, within the cost cap): the region the threshold decides.

  threshold geomean regret worst without tridiag: geomean worst held out (fitted on half)
with eigenvectors 1024 1.0393 2.09x 1.1750 2.73x from 1024: 1.0691 vs 1.1136
eigenvalues alone never 1.0021 1.13x 1.0021 1.13x from never: 1.0000 vs 1.0000

tridiag over the CPU (with eigenvectors), N x batch: 128x1 0.21x, 128x2 0.21x, 128x4 0.21x, 128x8 0.21x, 128x16 0.20x, 128x32 0.21x, 128x64 0.21x, 128x128 0.21x, 128x256 0.21x, 128x512 0.21x, 192x1 0.31x, 192x2 0.31x, 192x4 0.29x, 192x8 0.30x, 192x16 0.30x, 192x32 0.31x, 192x64 0.30x, 192x128 0.30x, 192x256 0.30x, 256x1 0.42x, 256x2 0.42x, 256x4 0.41x, 256x8 0.42x, 256x16 0.42x, 256x32 0.41x, 256x64 0.42x, 256x128 0.41x, 256x256 0.42x, 384x1 0.51x, 384x2 0.50x, 384x4 0.48x, 384x8 0.50x, 384x16 0.50x, 384x32 0.50x, 384x64 0.50x, 384x128 0.49x, 512x1 0.69x, 512x2 0.67x, 512x4 0.67x, 512x8 0.67x, 512x16 0.66x, 512x32 0.67x, 512x64 0.67x, 768x1 0.81x, 768x2 0.80x, 768x4 0.80x, 768x8 0.79x, 768x16 0.80x, 1024x1 1.13x, 1024x2 1.13x, 1024x4 1.13x, 1024x8 1.13x, 1536x1 1.47x, 1536x2 1.43x, 1536x4 1.44x, 2048x1 1.93x, 2048x2 1.95x, 2048x4 1.93x, 3072x1 2.73x

tridiag over the CPU (eigenvalues alone), N x batch: 128x1 0.14x, 128x2 0.14x, 128x4 0.13x, 128x8 0.13x, 128x16 0.13x, 128x32 0.13x, 128x64 0.13x, 128x128 0.13x, 128x256 0.13x, 128x512 0.13x, 192x1 0.19x, 192x2 0.19x, 192x4 0.18x, 192x8 0.19x, 192x16 0.19x, 192x32 0.18x, 192x64 0.19x, 192x128 0.19x, 192x256 0.19x, 256x1 0.24x, 256x2 0.24x, 256x4 0.23x, 256x8 0.24x, 256x16 0.23x, 256x32 0.23x, 256x64 0.24x, 256x128 0.24x, 256x256 0.24x, 384x1 0.33x, 384x2 0.33x, 384x4 0.32x, 384x8 0.33x, 384x16 0.32x, 384x32 0.33x, 384x64 0.32x, 384x128 0.33x, 512x1 0.40x, 512x2 0.40x, 512x4 0.40x, 512x8 0.39x, 512x16 0.40x, 512x32 0.40x, 512x64 0.40x, 768x1 0.54x, 768x2 0.55x, 768x4 0.54x, 768x8 0.53x, 768x16 0.54x, 1024x1 0.68x, 1024x2 0.67x, 1024x4 0.68x, 1024x8 0.67x, 1536x1 0.85x, 1536x2 0.85x, 1536x4 0.84x, 2048x1 0.93x, 2048x2 0.94x, 2048x4 0.94x, 3072x1 1.13x

Noise floor

Pass-to-pass ratio (max/min of the same measurement across passes), 1284 measurements: median 1.013, p90 1.091, max 2.98. The held-out verdicts use a bootstrap rather than this figure, since a mean over many points is far less noisy than one measurement.

runtime n median p90 max
<1 ms 465 1.034 1.298 2.98
1-3 ms 119 1.008 1.031 1.16
3-10 ms 187 1.008 1.032 1.11
10-30 ms 123 1.008 1.035 1.09
30-100 ms 141 1.012 1.066 1.13
>100 ms 249 1.010 1.050 1.14