Skip to main content

step_bench/
step_bench.rs

1// RLX — versatile ML compiler + runtime.
2// Copyright (C) 2026 Eugene Hauptmann, Nataliya Kosmyna.
3// SPDX-License-Identifier: MIT OR Apache-2.0
4
5//! Per-element cost of one optimizer step.
6//!
7//! The host optimizer is 40–55% of a training step at MNIST-MLP scale and does
8//! not vary with the device, so this is the number to move. Run with
9//! `--release`; a debug build measures the debug build.
10//!
11//!     cargo run --release -p rlx-optim --example step_bench
12//!     cargo run --release -p rlx-optim --example step_bench --features parallel
13
14use std::time::Instant;
15
16use rlx_optim::{Adam, AdamW, Lion, Optimizer, Sgd};
17
18/// A deterministic spread of magnitudes, so the timing is not measuring
19/// denormal or all-equal fast paths.
20fn fill(n: usize, seed: u32) -> Vec<f32> {
21    let mut s = seed;
22    (0..n)
23        .map(|_| {
24            s = s.wrapping_mul(1_664_525).wrapping_add(1_013_904_223);
25            ((s >> 8) as f32 / 8_388_608.0 - 1.0) * 0.1
26        })
27        .collect()
28}
29
30fn bench(label: &str, n: usize, mut opt: Box<dyn Optimizer>) -> f64 {
31    let shape = [n];
32    let mut param = fill(n, 7);
33    let grad = fill(n, 11);
34    // Warm: the first call allocates the moment buffers.
35    for _ in 0..5 {
36        opt.step("w", &shape, &mut param, &grad);
37        opt.end_iteration();
38    }
39    let iters = (200_000_000 / n).clamp(20, 2000);
40    let start = Instant::now();
41    for _ in 0..iters {
42        opt.step("w", &shape, &mut param, &grad);
43        opt.end_iteration();
44    }
45    let per_step = start.elapsed().as_secs_f64() / iters as f64;
46    let per_elem_ns = per_step * 1e9 / n as f64;
47    println!(
48        "  {label:22} {:>9.1} µs/step  {per_elem_ns:>6.2} ns/element  {:>6.2} GB/s",
49        per_step * 1e6,
50        // Adam touches p, m, v (read+write) and g (read): 7 × 4 bytes per element.
51        (n as f64 * 7.0 * 4.0) / per_step / 1e9
52    );
53    per_elem_ns
54}
55
56fn main() {
57    for n in [64 * 1024usize, 1024 * 1024, 8 * 1024 * 1024] {
58        println!(
59            "\n{} elements ({:.1} MiB per buffer)",
60            n,
61            (n * 4) as f64 / 1048576.0
62        );
63        bench("adamw (f64 default)", n, Box::new(AdamW::new(1e-3)));
64        bench(
65            "adamw f32Math",
66            n,
67            Box::new(AdamW::new(1e-3).with_f32_math(true)),
68        );
69        bench("adam (f64 default)", n, Box::new(Adam::new(1e-3)));
70        bench(
71            "adam f32Math",
72            n,
73            Box::new(Adam::new(1e-3).with_f32_math(true)),
74        );
75        bench("sgd+momentum", n, {
76            let mut o = Sgd::new(1e-3);
77            o.momentum = 0.9;
78            Box::new(o)
79        });
80        bench("lion", n, Box::new(Lion::new(1e-3)));
81    }
82}