use futures_util::stream::StreamExt;
use lamellar::array::prelude::*;
use matrixmultiply::sgemm;
#[lamellar::main]
fn main() {
let args: Vec<String> = std::env::args().collect();
let elem_per_pe = args
.get(1)
.and_then(|s| s.parse::<usize>().ok())
.unwrap_or_else(|| 2000);
let world = lamellar::LamellarWorldBuilder::new().build();
let my_pe = world.my_pe();
let num_pes = world.num_pes();
let dim = elem_per_pe * num_pes;
let m = dim; let n = dim; let p = dim;
let a = LocalLockArray::<f32>::new(&world, m * n, Distribution::Block).block(); let b = LocalLockArray::<f32>::new(&world, n * p, Distribution::Block).block(); let c = LocalLockArray::<f32>::new(&world, m * p, Distribution::Block).block(); let a_init = a
.dist_iter_mut()
.enumerate()
.for_each(move |(i, x)| *x = i as f32);
let b_init = b.dist_iter_mut().enumerate().for_each(move |(i, x)| {
let row = i / dim;
let col = i % dim;
if row == col {
*x = 1 as f32
} else {
*x = 0 as f32;
}
});
let c_init = c.dist_iter_mut().for_each(move |x| *x = 0.0);
world.block_on_all([a_init, b_init, c_init]);
let a = a.into_read_only().block();
let b = b.into_read_only().block();
world.barrier();
let num_gops = ((2 * dim * dim * dim) - dim * dim) as f64 / 1_000_000_000.0; let blocksize = std::cmp::min(1000000, dim / num_pes);
let m_blks = m / blocksize; let m_blks_pe = m_blks / num_pes;
let n_blks = n / blocksize; let p_blks = p / blocksize; let p_blks_pe = p_blks / num_pes;
println!("{blocksize} {m_blks} {m_blks_pe} {n_blks} {p_blks} {p_blks_pe}");
let nblks_array = LocalLockArray::new(&world, n_blks * num_pes, Distribution::Block).block();
nblks_array
.dist_iter_mut()
.enumerate()
.for_each(move |(i, x)| *x = i % n_blks)
.block();
let m_blks_pe_array =
LocalLockArray::new(&world, m_blks_pe * num_pes, Distribution::Block).block();
m_blks_pe_array
.dist_iter_mut()
.enumerate()
.for_each(move |(i, x)| *x = i % m_blks_pe)
.block();
world.barrier();
let nblks_array = nblks_array.into_read_only().block();
let m_blks_pe_array = m_blks_pe_array.into_read_only().block();
println!("{blocksize} {m_blks} {m_blks_pe} {n_blks} {p_blks}");
let start = std::time::Instant::now();
let a = a.clone();
let b = b.clone();
let c_clone = c.clone();
let gemm = nblks_array.dist_iter().for_each_async(move |k_blk| {
let a = a.clone();
let b = b.clone();
let c_clone = c_clone.clone();
let m_blks_pe_array = m_blks_pe_array.clone();
async move {
let my_p_blks = (p_blks_pe * my_pe..p_blks).chain(0..p_blks_pe * my_pe); for j_blk in my_p_blks {
let b_block: Vec<f32> = b
.onesided_iter() .chunks(blocksize) .skip(*k_blk * n_blks * blocksize + j_blk) .step_by(n_blks) .into_stream() .take(blocksize) .fold(Vec::new(), |mut vec, x| {
vec.extend_from_slice(x.as_slice());
async move { vec }
})
.await;
let a = a.clone();
let c_clone = c_clone.clone();
m_blks_pe_array
.local_iter()
.for_each_async_with_schedule(
Schedule::Chunk(m_blks_pe_array.len()),
move |i_blk| {
let c = c_clone.clone();
let b_block_vec = b_block.clone();
let a_vec: Vec<f32> = a
.local_as_slice()
.chunks(blocksize) .skip(i_blk * m_blks * blocksize + *k_blk) .step_by(m_blks) .take(blocksize) .fold(Vec::new(), |mut vec, x| {
vec.extend(x);
vec
});
let mut c_vec = vec![0.0; blocksize * blocksize]; unsafe {
sgemm(
blocksize,
blocksize,
blocksize,
1.0,
a_vec.as_ptr(),
blocksize as isize,
1,
b_block_vec.as_ptr(),
1,
blocksize as isize,
0.0,
c_vec.as_mut_ptr(),
blocksize as isize,
1,
);
}
async move {
let mut c_slice = c.write_local_data().await;
for row in 0..blocksize {
let row_offset = (i_blk * blocksize + row) * n;
for col in 0..blocksize {
let col_offset = j_blk * blocksize + col;
c_slice[row_offset + col_offset] +=
c_vec[row * blocksize + col];
}
}
}
},
)
.await;
}
}
});
gemm.block();
println!(
"[{:?}] block_on done {:?}",
my_pe,
start.elapsed().as_secs_f64()
);
world.wait_all();
println!(
"[{:?}] wait_all done {:?}",
my_pe,
start.elapsed().as_secs_f64()
);
world.barrier();
let elapsed = start.elapsed().as_secs_f64();
if my_pe == 0 {
println!("Elapsed: {:?}", elapsed);
println!(
"blksize: {:?} elapsed {:?} Gflops: {:?}",
blocksize,
elapsed,
num_gops / elapsed,
);
}
}