Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension


Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
8 changes: 8 additions & 0 deletions linalg/Cargo.toml
Original file line number Diff line number Diff line change
Expand Up @@ -52,3 +52,11 @@ required-features = ["blas"]
[[example]]
name = "jit_bench"
required-features = ["jit"]

[[example]]
name = "gnn_million"
required-features = ["jit"]

[[example]]
name = "gnn_stress"
required-features = ["jit"]
58 changes: 51 additions & 7 deletions linalg/benches/perf.rs
Original file line number Diff line number Diff line change
Expand Up @@ -15,9 +15,11 @@

use std::time::Instant;

#[cfg(feature = "jit")]
use linalg::jit::{EinsumF32Plan, JitInput};
use linalg::{
any::Tensor,
blocked::{attention, Blocked, Blocked16, Blocked8},
blocked::{Blocked, Blocked8, Blocked16, attention},
csr::Csr,
dense::Dense,
einsum::{einsum, einsum_homogenous},
Expand Down Expand Up @@ -170,6 +172,22 @@ fn section_dense_matmul() {
});
println!(" ratio (dyn / homo): {:.2}×", dyn_t / homo);
println!(" ratio (enum / homo): {:.2}×", enum_t / homo);
#[cfg(feature = "jit")]
{
let plan = EinsumF32Plan::compile(
"ab,bc->ac",
&[JitInput::Dense(&a), JitInput::Dense(&b)],
&[vec![n, n]],
)
.unwrap();
let plan_name = format!("EinsumF32Plan ({:?})", plan.backend());
let plan_t = bench(&plan_name, iters, || {
let mut c = Dense::<f32>::zeros(vec![n, n]);
plan.run(&[JitInput::Dense(&a), JitInput::Dense(&b)], &mut [&mut c]);
std::hint::black_box(&c);
});
println!(" ratio (plan / homo): {:.4}×", plan_t / homo);
}
}
}

Expand All @@ -181,11 +199,14 @@ fn section_csr_matmul() {
for &s in &[5usize, 10, 20] {
let a = lattice_csr(s, 3.0, 42);
let n = (s * s * s) as u32;
println!(
"\n--- side={s} (n={n}, nnz={}) ---",
a.nnz()
);
let iters: u32 = if s == 5 { 2000 } else if s == 10 { 200 } else { 30 };
println!("\n--- side={s} (n={n}, nnz={}) ---", a.nnz());
let iters: u32 = if s == 5 {
2000
} else if s == 10 {
200
} else {
30
};

let nat = bench("Csr::matmul (native)", iters, || {
let r = a.matmul(&a);
Expand Down Expand Up @@ -230,7 +251,13 @@ fn section_csr_par() {
let a = lattice_csr(s, 3.0, 42);
let n = (s * s * s) as u32;
println!("\n--- side={s} (n={n}, nnz={}) ---", a.nnz());
let iters: u32 = if s == 10 { 100 } else if s == 20 { 20 } else { 5 };
let iters: u32 = if s == 10 {
100
} else if s == 20 {
20
} else {
5
};

let seq = bench("Csr::matmul", iters, || {
let r = a.matmul(&a);
Expand Down Expand Up @@ -286,6 +313,23 @@ fn section_csr_times_dense() {
std::hint::black_box(&y);
});
println!(" ratio (enum / dyn): {:.2}×", enum_t / dyn_t);
#[cfg(feature = "jit")]
{
let plan = EinsumF32Plan::compile(
"ab,bc->ac",
&[JitInput::Csr(&af), JitInput::Dense(&x)],
&[vec![n, d]],
)
.unwrap();
let plan_name = format!("EinsumF32Plan ({:?})", plan.backend());
let plan_t = bench(&plan_name, iters, || {
let mut y = Dense::<f32>::zeros(vec![n, d]);
plan.run(&[JitInput::Csr(&af), JitInput::Dense(&x)], &mut [&mut y]);
std::hint::black_box(&y);
});
println!(" ratio (plan / dyn): {:.4}×", plan_t / dyn_t);
println!(" ratio (plan / enum): {:.4}×", plan_t / enum_t);
}
}
}

Expand Down
112 changes: 112 additions & 0 deletions linalg/examples/gnn_million.rs
Original file line number Diff line number Diff line change
@@ -0,0 +1,112 @@
//! Million-node graph propagation with the sparse einsum JIT.
//!
//! Run with:
//! `cargo run --release --features jit --example gnn_million`

use std::time::Instant;

use linalg::csr::Csr;
use linalg::dense::Dense;
use linalg::jit::{EinsumF32Jit, JitInput};
use linalg::tensor::NDIndex;

fn fill_rand(t: &mut Dense<f32>, mut state: u64) {
for v in t.data.iter_mut() {
state = state
.wrapping_mul(6364136223846793005)
.wrapping_add(1442695040888963407);
*v = ((state >> 32) as f32 / u32::MAX as f32) * 2.0 - 1.0;
}
}

fn rand_graph(n: usize, per_row: usize, mut state: u64) -> Csr<u32, f32> {
let mut triples = Vec::with_capacity(n * per_row);
let scale = 1.0 / per_row as f32;
for r in 0..n {
for _ in 0..per_row {
state = state
.wrapping_mul(6364136223846793005)
.wrapping_add(1442695040888963407);
let c = (state >> 33) as usize % n;
triples.push((r as u32, c as u32, scale));
}
}
Csr::<u32, f32>::from_coo(n as u32, &mut triples)
}

fn manual_row_feature(a: &Csr<u32, f32>, x: &Dense<f32>, row: usize, feature: usize) -> f32 {
let mut acc = 0.0;
for (col, value) in a.row(row as u32) {
acc += value * x.get(&[col as usize, feature]);
}
acc
}

fn main() {
const N: usize = 1_048_576;
const PER_ROW: usize = 16;
const FEATURES: usize = 32;
const LAYERS: usize = 4;

println!("million-node sparse propagation: A({N}x{N}, {PER_ROW}/row) · X({N}x{FEATURES})");

let build_start = Instant::now();
let a = rand_graph(N, PER_ROW, 1);
let mut x = Dense::<f32>::zeros(vec![N, FEATURES]);
fill_rand(&mut x, 2);
println!(
" graph nnz={} build+feature-init {:.3}s",
a.nnz(),
build_start.elapsed().as_secs_f64()
);

let compile_start = Instant::now();
let jit = EinsumF32Jit::compile(
"ab,bc->ac",
&[JitInput::Csr(&a), JitInput::Dense(&x)],
&[vec![N, FEATURES]],
)
.unwrap();
println!(
" one-time JIT compile {:.3} ms",
compile_start.elapsed().as_secs_f64() * 1_000.0
);

let mut y = Dense::<f32>::zeros(vec![N, FEATURES]);
let start = Instant::now();
jit.run(&[JitInput::Csr(&a), JitInput::Dense(&x)], &mut [&mut y]);
let one_layer_s = start.elapsed().as_secs_f64();

let row = 123_456;
let feature = 17;
let expected = manual_row_feature(&a, &x, row, feature);
let got = y.get(&[row, feature]);
println!(
" one layer JIT {:8.3} ms spot_check row={row} feature={feature} diff={:.3e}",
one_layer_s * 1_000.0,
(expected - got).abs()
);

let mut ping = x;
let mut pong = y;
let start = Instant::now();
for _ in 1..LAYERS {
pong.clear();
jit.run(
&[JitInput::Csr(&a), JitInput::Dense(&ping)],
&mut [&mut pong],
);
std::mem::swap(&mut ping, &mut pong);
}
let remaining_s = start.elapsed().as_secs_f64();
let total_s = one_layer_s + remaining_s;
let edge_feature_updates = a.nnz() as f64 * FEATURES as f64 * LAYERS as f64;
let flops = 2.0 * edge_feature_updates;

println!(
" {LAYERS} layers total {:8.3} ms {:.2} billion edge-feature updates/s {:.2} GFLOP/s",
total_s * 1_000.0,
edge_feature_updates / total_s / 1.0e9,
flops / total_s / 1.0e9
);
}
128 changes: 128 additions & 0 deletions linalg/examples/gnn_stress.rs
Original file line number Diff line number Diff line change
@@ -0,0 +1,128 @@
//! Stress test: graph-neural-network style sparse propagation.
//!
//! Run with:
//! `cargo run --release --features jit --example gnn_stress`

use std::time::Instant;

use linalg::csr::Csr;
use linalg::dense::Dense;
use linalg::einsum::einsum;
use linalg::jit::{EinsumF32Jit, JitInput};
use linalg::tensor::NDIndex;

fn fill_rand(t: &mut Dense<f32>, mut state: u64) {
for v in t.data.iter_mut() {
state = state
.wrapping_mul(6364136223846793005)
.wrapping_add(1442695040888963407);
*v = ((state >> 32) as f32 / u32::MAX as f32) * 2.0 - 1.0;
}
}

fn rand_graph(n: usize, per_row: usize, mut state: u64) -> Csr<u32, f32> {
let mut triples = Vec::with_capacity(n * per_row);
let scale = 1.0 / per_row as f32;
for r in 0..n {
for _ in 0..per_row {
state = state
.wrapping_mul(6364136223846793005)
.wrapping_add(1442695040888963407);
let c = (state >> 33) as usize % n;
triples.push((r as u32, c as u32, scale));
}
}
Csr::<u32, f32>::from_coo(n as u32, &mut triples)
}

fn max_abs_diff(a: &Dense<f32>, b: &Dense<f32>) -> f32 {
a.data
.iter()
.zip(b.data.iter())
.map(|(x, y)| (x - y).abs())
.fold(0.0, f32::max)
}

fn run_vm(a: &Csr<u32, f32>, x: &Dense<f32>, y: &mut Dense<f32>) -> f64 {
y.clear();
let start = Instant::now();
einsum::<f32>(
"ab,bc->ac",
&[a as &dyn NDIndex<f32>, x as &dyn NDIndex<f32>],
&mut [y as &mut dyn NDIndex<f32>],
)
.unwrap();
start.elapsed().as_secs_f64()
}

fn run_jit(jit: &EinsumF32Jit, a: &Csr<u32, f32>, x: &Dense<f32>, y: &mut Dense<f32>) -> f64 {
y.clear();
let start = Instant::now();
jit.run(&[JitInput::Csr(a), JitInput::Dense(x)], &mut [y]);
start.elapsed().as_secs_f64()
}

fn main() {
const N: usize = 65_536;
const PER_ROW: usize = 16;
const FEATURES: usize = 64;
const LAYERS: usize = 4;

println!("GNN-style sparse propagation: A({N}x{N}, {PER_ROW}/row) · X({N}x{FEATURES})");

let build_start = Instant::now();
let a = rand_graph(N, PER_ROW, 1);
let mut x = Dense::<f32>::zeros(vec![N, FEATURES]);
fill_rand(&mut x, 2);
println!(
" graph nnz={} build+feature-init {:.3}s",
a.nnz(),
build_start.elapsed().as_secs_f64()
);

let compile_start = Instant::now();
let jit = EinsumF32Jit::compile(
"ab,bc->ac",
&[JitInput::Csr(&a), JitInput::Dense(&x)],
&[vec![N, FEATURES]],
)
.unwrap();
println!(
" one-time JIT compile {:.3} ms",
compile_start.elapsed().as_secs_f64() * 1_000.0
);

let mut vm_out = Dense::<f32>::zeros(vec![N, FEATURES]);
let mut jit_out = Dense::<f32>::zeros(vec![N, FEATURES]);

let vm_s = run_vm(&a, &x, &mut vm_out);
let jit_s = run_jit(&jit, &a, &x, &mut jit_out);
let diff = max_abs_diff(&vm_out, &jit_out);
let flops = 2.0 * a.nnz() as f64 * FEATURES as f64;

println!(
" one layer VM {:8.3} ms JIT {:8.3} ms {:5.1}x faster max_abs_diff {:.3e}",
vm_s * 1_000.0,
jit_s * 1_000.0,
vm_s / jit_s,
diff
);
println!(
" JIT effective throughput {:.2} GFLOP/s",
flops / jit_s / 1.0e9
);

let mut ping = x;
let mut pong = Dense::<f32>::zeros(vec![N, FEATURES]);
let layers_start = Instant::now();
for _ in 0..LAYERS {
run_jit(&jit, &a, &ping, &mut pong);
std::mem::swap(&mut ping, &mut pong);
}
let layers_s = layers_start.elapsed().as_secs_f64();
println!(
" {LAYERS} JIT layers {:8.3} ms {:.2} billion edge-feature updates/s",
layers_s * 1_000.0,
(a.nnz() as f64 * FEATURES as f64 * LAYERS as f64) / layers_s / 1.0e9
);
}
Loading