perf(backend-cpu): parallelize blocked gemm over row blocks with rayon
Each task owns a disjoint BLOCK_SIZE-row slice of the result; the inner blocked kernel is unchanged. Needed for dg-gnn HetGAT training throughput (node-level [M,64]x[64,64] matmuls dominated single-threaded step time). Co-Authored-By: Claude Fable 5 <[email protected]>
This commit is contained in:
co-authored by
Claude Fable 5
parent
35b2b2cdf4
commit
67c47898fa
@@ -20,26 +20,33 @@ pub fn matmul<E: CpuFloat>(
|
||||
// Block size for cache efficiency
|
||||
const BLOCK_SIZE: usize = 64;
|
||||
|
||||
// Blocked matrix multiplication
|
||||
for i_block in (0..m).step_by(BLOCK_SIZE) {
|
||||
for j_block in (0..n).step_by(BLOCK_SIZE) {
|
||||
for k_block in (0..k).step_by(BLOCK_SIZE) {
|
||||
let i_end = (i_block + BLOCK_SIZE).min(m);
|
||||
let j_end = (j_block + BLOCK_SIZE).min(n);
|
||||
let k_end = (k_block + BLOCK_SIZE).min(k);
|
||||
// Blocked matrix multiplication, parallelized over row blocks: each rayon
|
||||
// task owns a disjoint `BLOCK_SIZE`-row slice of the result, so the inner
|
||||
// blocked kernel is unchanged and no synchronization is needed.
|
||||
result
|
||||
.par_chunks_mut(BLOCK_SIZE * n)
|
||||
.enumerate()
|
||||
.for_each(|(bi, res_rows)| {
|
||||
let i_block = bi * BLOCK_SIZE;
|
||||
let i_end = (i_block + BLOCK_SIZE).min(m);
|
||||
for j_block in (0..n).step_by(BLOCK_SIZE) {
|
||||
for k_block in (0..k).step_by(BLOCK_SIZE) {
|
||||
let j_end = (j_block + BLOCK_SIZE).min(n);
|
||||
let k_end = (k_block + BLOCK_SIZE).min(k);
|
||||
|
||||
for i in i_block..i_end {
|
||||
for j in j_block..j_end {
|
||||
let mut sum = result[i * n + j];
|
||||
for kk in k_block..k_end {
|
||||
sum = sum + lhs.data[i * k + kk] * rhs.data[kk * n + j];
|
||||
for i in i_block..i_end {
|
||||
let row = &mut res_rows[(i - i_block) * n..(i - i_block) * n + n];
|
||||
for j in j_block..j_end {
|
||||
let mut sum = row[j];
|
||||
for kk in k_block..k_end {
|
||||
sum = sum + lhs.data[i * k + kk] * rhs.data[kk * n + j];
|
||||
}
|
||||
row[j] = sum;
|
||||
}
|
||||
result[i * n + j] = sum;
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
});
|
||||
|
||||
CpuTensorPrimitive::new(result, [m, n], lhs.device.clone())
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user