1use crate::tuning::{L1_CACHE_SIZE, L2_CACHE_SIZE};
23use core::mem::size_of;
24
25pub const BASE_CASE_THRESHOLD: usize = 64;
30
31pub const MIN_BLOCK_SIZE: usize = 16;
33
34pub const MAX_BLOCK_SIZE: usize = 512;
36
37pub fn gemm_block_sizes<T>(m: usize, n: usize, k: usize) -> (usize, usize, usize) {
51 let elem_size = size_of::<T>();
52
53 let target_bytes = L2_CACHE_SIZE / 2;
58
59 let max_block = ((target_bytes / elem_size / 3) as f64).sqrt() as usize;
61 let mut block = max_block.clamp(MIN_BLOCK_SIZE, MAX_BLOCK_SIZE);
62
63 block = (block / 8) * 8;
65 if block < MIN_BLOCK_SIZE {
66 block = MIN_BLOCK_SIZE;
67 }
68
69 let block_m = block.min(m);
71 let block_n = block.min(n);
72 let block_k = block.min(k);
73
74 (block_m, block_n, block_k)
75}
76
77pub fn trsm_block_size<T>(n: usize, nrhs: usize) -> usize {
83 let elem_size = size_of::<T>();
84
85 let max_block = ((2 * L1_CACHE_SIZE / elem_size) as f64).sqrt() as usize;
88 let block = max_block.clamp(MIN_BLOCK_SIZE, MAX_BLOCK_SIZE / 2);
89
90 let block = (block / 8) * 8;
93
94 block.min(n).min(nrhs)
99}
100
101pub fn factorization_panel_width<T>(n: usize) -> usize {
106 let elem_size = size_of::<T>();
107
108 let max_panel = L2_CACHE_SIZE / (elem_size * n.max(1));
111 let panel = max_panel.clamp(16, 128);
112
113 let panel = (panel / 4) * 4;
116
117 panel.min(n)
121}
122
123#[derive(Debug, Clone, Copy)]
128pub struct BlockRange {
129 pub start: usize,
131 pub end: usize,
133}
134
135impl BlockRange {
136 #[inline]
138 pub const fn new(start: usize, end: usize) -> Self {
139 BlockRange { start, end }
140 }
141
142 #[inline]
144 pub const fn from_len(n: usize) -> Self {
145 BlockRange { start: 0, end: n }
146 }
147
148 #[inline]
150 pub const fn len(&self) -> usize {
151 self.end.saturating_sub(self.start)
152 }
153
154 #[inline]
156 pub const fn is_empty(&self) -> bool {
157 self.start >= self.end
158 }
159
160 #[inline]
162 pub fn is_base_case(&self, threshold: usize) -> bool {
163 self.len() <= threshold
164 }
165
166 #[inline]
170 pub fn split(&self) -> (Self, Self) {
171 let mid = self.start + self.len() / 2;
172 (
173 BlockRange::new(self.start, mid),
174 BlockRange::new(mid, self.end),
175 )
176 }
177
178 #[inline]
180 pub fn split_at(&self, point: usize) -> (Self, Self) {
181 let split = (self.start + point).min(self.end);
182 (
183 BlockRange::new(self.start, split),
184 BlockRange::new(split, self.end),
185 )
186 }
187}
188
189#[derive(Debug, Clone, Copy)]
193pub struct RecursiveTask {
194 pub rows: BlockRange,
196 pub cols: BlockRange,
198}
199
200impl RecursiveTask {
201 #[inline]
203 pub const fn new(rows: BlockRange, cols: BlockRange) -> Self {
204 RecursiveTask { rows, cols }
205 }
206
207 #[inline]
209 pub const fn from_dims(m: usize, n: usize) -> Self {
210 RecursiveTask {
211 rows: BlockRange::from_len(m),
212 cols: BlockRange::from_len(n),
213 }
214 }
215
216 #[inline]
218 pub fn size(&self) -> usize {
219 self.rows.len() * self.cols.len()
220 }
221
222 #[inline]
224 pub fn is_base_case(&self, threshold: usize) -> bool {
225 self.rows.len() <= threshold && self.cols.len() <= threshold
226 }
227
228 pub fn split(&self) -> (Self, Self) {
232 if self.rows.len() >= self.cols.len() {
233 let (r1, r2) = self.rows.split();
235 (
236 RecursiveTask::new(r1, self.cols),
237 RecursiveTask::new(r2, self.cols),
238 )
239 } else {
240 let (c1, c2) = self.cols.split();
242 (
243 RecursiveTask::new(self.rows, c1),
244 RecursiveTask::new(self.rows, c2),
245 )
246 }
247 }
248
249 pub fn quadrants(&self) -> (Self, Self, Self, Self) {
253 let (r1, r2) = self.rows.split();
254 let (c1, c2) = self.cols.split();
255
256 (
257 RecursiveTask::new(r1, c1), RecursiveTask::new(r1, c2), RecursiveTask::new(r2, c1), RecursiveTask::new(r2, c2), )
262 }
263}
264
265pub trait BlockVisitor {
269 type Error;
271
272 fn visit_block(
278 &mut self,
279 row_start: usize,
280 row_end: usize,
281 col_start: usize,
282 col_end: usize,
283 ) -> Result<(), Self::Error>;
284}
285
286pub fn cache_oblivious_traverse<V: BlockVisitor>(
291 visitor: &mut V,
292 task: RecursiveTask,
293 threshold: usize,
294) -> Result<(), V::Error> {
295 if task.is_base_case(threshold) {
296 visitor.visit_block(
298 task.rows.start,
299 task.rows.end,
300 task.cols.start,
301 task.cols.end,
302 )
303 } else {
304 let (t1, t2) = task.split();
306 cache_oblivious_traverse(visitor, t1, threshold)?;
307 cache_oblivious_traverse(visitor, t2, threshold)
308 }
309}
310
311#[inline]
316pub fn morton_index(x: u32, y: u32) -> u64 {
317 fn expand_bits(v: u32) -> u64 {
318 let mut v = v as u64;
319 v = (v | (v << 16)) & 0x0000_FFFF_0000_FFFF;
320 v = (v | (v << 8)) & 0x00FF_00FF_00FF_00FF;
321 v = (v | (v << 4)) & 0x0F0F_0F0F_0F0F_0F0F;
322 v = (v | (v << 2)) & 0x3333_3333_3333_3333;
323 v = (v | (v << 1)) & 0x5555_5555_5555_5555;
324 v
325 }
326 expand_bits(x) | (expand_bits(y) << 1)
327}
328
329#[inline]
331pub fn morton_decode(z: u64) -> (u32, u32) {
332 fn compact_bits(mut v: u64) -> u32 {
333 v &= 0x5555_5555_5555_5555;
334 v = (v | (v >> 1)) & 0x3333_3333_3333_3333;
335 v = (v | (v >> 2)) & 0x0F0F_0F0F_0F0F_0F0F;
336 v = (v | (v >> 4)) & 0x00FF_00FF_00FF_00FF;
337 v = (v | (v >> 8)) & 0x0000_FFFF_0000_FFFF;
338 v = (v | (v >> 16)) & 0x0000_0000_FFFF_FFFF;
339 v as u32
340 }
341 (compact_bits(z), compact_bits(z >> 1))
342}
343
344#[cfg(test)]
345mod tests {
346 use super::*;
347
348 #[test]
349 fn test_gemm_block_sizes() {
350 let (bm, bn, bk) = gemm_block_sizes::<f64>(1024, 1024, 1024);
351
352 assert!(bm >= MIN_BLOCK_SIZE);
354 assert!(bn >= MIN_BLOCK_SIZE);
355 assert!(bk >= MIN_BLOCK_SIZE);
356 assert!(bm <= MAX_BLOCK_SIZE);
357 assert!(bn <= MAX_BLOCK_SIZE);
358 assert!(bk <= MAX_BLOCK_SIZE);
359
360 assert_eq!(bm % 8, 0);
362 }
363
364 #[test]
365 fn test_trsm_block_size_clamped_to_matrix_extent() {
366 for &n in &[0usize, 1, 4, 8, 15, 16] {
372 for &nrhs in &[0usize, 1, 4, 8] {
373 let block = trsm_block_size::<f64>(n, nrhs);
374 assert!(block <= n, "block {block} exceeds n={n} (nrhs={nrhs})");
375 assert!(block <= nrhs, "block {block} exceeds nrhs={nrhs} (n={n})");
376 }
377 }
378
379 assert_eq!(trsm_block_size::<f64>(0, 64), 0);
382 assert_eq!(trsm_block_size::<f64>(64, 0), 0);
383
384 assert_eq!(trsm_block_size::<f64>(87, 4096), 87);
390 assert_eq!(trsm_block_size::<f64>(88, 4096), 88);
391 assert_eq!(trsm_block_size::<f64>(89, 4096), 88);
392
393 let unclamped = trsm_block_size::<f64>(4096, 4096);
396 assert!(unclamped >= MIN_BLOCK_SIZE);
397 assert!(unclamped <= MAX_BLOCK_SIZE / 2);
398 }
399
400 #[test]
401 fn test_factorization_panel_width_clamped_to_matrix_extent() {
402 for &n in &[0usize, 1, 4, 8, 15, 16] {
405 let panel = factorization_panel_width::<f64>(n);
406 assert!(panel <= n, "panel {panel} exceeds n={n}");
407 }
408
409 assert_eq!(factorization_panel_width::<f64>(0), 0);
410 assert_eq!(factorization_panel_width::<f64>(1), 1);
411 assert_eq!(factorization_panel_width::<f64>(4), 4);
412
413 assert_eq!(factorization_panel_width::<f64>(256), 128);
417 assert_eq!(factorization_panel_width::<f64>(257), 124);
418 }
419
420 #[test]
421 fn test_block_range() {
422 let range = BlockRange::new(0, 100);
423 assert_eq!(range.len(), 100);
424
425 let (left, right) = range.split();
426 assert_eq!(left.start, 0);
427 assert_eq!(left.end, 50);
428 assert_eq!(right.start, 50);
429 assert_eq!(right.end, 100);
430
431 assert!(BlockRange::new(0, 32).is_base_case(64));
432 assert!(!BlockRange::new(0, 100).is_base_case(64));
433 }
434
435 #[test]
436 fn test_recursive_task() {
437 let task = RecursiveTask::from_dims(100, 200);
438 assert_eq!(task.size(), 20000);
439
440 let (t1, t2) = task.split();
442 assert_eq!(t1.cols.len(), 100);
443 assert_eq!(t2.cols.len(), 100);
444 assert_eq!(t1.rows.len(), 100);
445 assert_eq!(t2.rows.len(), 100);
446 }
447
448 #[test]
449 fn test_quadrants() {
450 let task = RecursiveTask::from_dims(100, 100);
451 let (tl, _tr, _bl, br) = task.quadrants();
452
453 assert_eq!(tl.rows.start, 0);
454 assert_eq!(tl.rows.end, 50);
455 assert_eq!(tl.cols.start, 0);
456 assert_eq!(tl.cols.end, 50);
457
458 assert_eq!(br.rows.start, 50);
459 assert_eq!(br.rows.end, 100);
460 assert_eq!(br.cols.start, 50);
461 assert_eq!(br.cols.end, 100);
462 }
463
464 #[test]
465 fn test_morton_index() {
466 assert_eq!(morton_index(0, 0), 0);
468 assert_eq!(morton_index(1, 0), 1);
469 assert_eq!(morton_index(0, 1), 2);
470 assert_eq!(morton_index(1, 1), 3);
471 assert_eq!(morton_index(2, 0), 4);
472
473 for x in 0..100 {
475 for y in 0..100 {
476 let z = morton_index(x, y);
477 let (dx, dy) = morton_decode(z);
478 assert_eq!((dx, dy), (x, y));
479 }
480 }
481 }
482
483 struct CountingVisitor {
484 count: usize,
485 total_elements: usize,
486 }
487
488 impl BlockVisitor for CountingVisitor {
489 type Error = ();
490
491 fn visit_block(
492 &mut self,
493 row_start: usize,
494 row_end: usize,
495 col_start: usize,
496 col_end: usize,
497 ) -> Result<(), ()> {
498 self.count += 1;
499 self.total_elements += (row_end - row_start) * (col_end - col_start);
500 Ok(())
501 }
502 }
503
504 #[test]
505 fn test_cache_oblivious_traverse() {
506 let task = RecursiveTask::from_dims(128, 128);
507 let mut visitor = CountingVisitor {
508 count: 0,
509 total_elements: 0,
510 };
511
512 cache_oblivious_traverse(&mut visitor, task, 32).unwrap();
513
514 assert!(visitor.count > 1);
516 assert_eq!(visitor.total_elements, 128 * 128);
518 }
519}