1use std::time::Instant;
15
16use rlx_optim::{Adam, AdamW, Lion, Optimizer, Sgd};
17
18fn 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 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 (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}