1use super::{FusedKerSpec, FusedSpec, MatMatMulKer, OutputStoreKer};
2use crate::{BinOp, LADatum};
3use downcast_rs::{Downcast, impl_downcast};
4use std::cell::RefCell;
5use std::fmt::Debug;
6use std::sync::atomic::AtomicUsize;
7use tract_data::internal::num_integer::Integer;
8use tract_data::internal::*;
9
10static GENERATION: AtomicUsize = AtomicUsize::new(1);
11
12thread_local! {
13 static TLS: RefCell<TLSScratch> = Default::default();
14}
15
16#[derive(Default, Debug)]
17pub(crate) struct TLSScratch {
18 generation: usize,
19 blob: Blob,
20 ker_specs_16: Vec<FusedKerSpec<f16>>,
21 ker_specs_32: Vec<FusedKerSpec<f32>>,
22 ker_specs_64: Vec<FusedKerSpec<f64>>,
23}
24
25impl TLSScratch {
26 #[allow(unknown_lints, clippy::missing_transmute_annotations)]
27 fn ker_specs<TI: LADatum>(&mut self) -> &mut Vec<FusedKerSpec<TI>> {
28 unsafe {
29 if TI::datum_type() == f32::datum_type() || TI::datum_type() == i32::datum_type() {
30 std::mem::transmute(&mut self.ker_specs_32)
31 } else if TI::datum_type() == f16::datum_type() {
32 std::mem::transmute(&mut self.ker_specs_16)
33 } else if TI::datum_type() == f64::datum_type() {
34 std::mem::transmute(&mut self.ker_specs_64)
35 } else {
36 todo!();
37 }
38 }
39 }
40
41 fn sync<TI: LADatum>(&mut self, scratch: &ScratchSpaceImpl<TI>) {
42 if self.generation == scratch.generation {
43 return;
44 }
45 let ker_specs = self.ker_specs::<TI>();
46 ker_specs.clear();
47 ker_specs.extend_from_slice(&scratch.ker_specs);
48
49 unsafe {
50 self.blob.ensure_size_and_align(scratch.blob_size, scratch.blob_align);
51
52 for LocDependent { loc, ker_spec, .. } in &scratch.loc_dependent {
53 #[allow(clippy::single_match)]
54 if matches!(scratch.ker_specs[*ker_spec], FusedKerSpec::AddMatMul { .. }) {
55 let scratch = &mut *(self.blob.as_ptr().add(*loc) as *mut AddMatMulTemp);
56 scratch.panel_a_id = usize::MAX;
57 scratch.panel_b_id = usize::MAX;
58 };
59 }
60 }
61 self.generation = scratch.generation;
62 }
63}
64
65pub trait ScratchSpace: Downcast + Send {}
66impl_downcast!(ScratchSpace);
67
68#[derive(Debug, Default)]
69pub struct ScratchSpaceImpl<TI: LADatum> {
70 generation: usize,
71 blob_size: usize,
72 blob_align: usize,
73 ker_specs: Vec<FusedKerSpec<TI>>,
74 loc_dependent: TVec<LocDependent>,
75 valid_down_tiles: usize,
76 remnant_down: usize,
77 valid_right_tiles: usize,
78 remnant_right: usize,
79}
80
81#[derive(Debug, new)]
82struct LocDependent {
83 spec: usize,
84 ker_spec: usize,
85 loc: usize,
87 buffer_a: Option<usize>,
89 buffer_b: Option<usize>,
90}
91
92impl<TI: LADatum> ScratchSpace for ScratchSpaceImpl<TI> {}
93unsafe impl<TI: LADatum> Send for ScratchSpaceImpl<TI> {}
94
95#[derive(Debug)]
96struct AddMatMulTemp {
97 ptr_a: *const u8,
98 panel_a_id: usize,
99 ptr_b: *const u8,
100 panel_b_id: usize,
101}
102
103impl<TI: LADatum> ScratchSpaceImpl<TI> {
104 pub unsafe fn prepare(
105 &mut self,
106 ker: &impl MatMatMulKer<Acc = TI>,
107 m: usize,
108 n: usize,
109 specs: &[FusedSpec],
110 ) -> TractResult<()> {
111 use FusedKerSpec as FKS;
112 use FusedSpec as FS;
113 self.ker_specs.clear();
114 self.loc_dependent.clear();
115 self.ker_specs.reserve(specs.len() + 2);
116 self.ker_specs.push(FusedKerSpec::Clear);
117 self.valid_down_tiles = m / ker.mr();
118 self.remnant_down = m % ker.mr();
119 self.valid_right_tiles = n / ker.nr();
120 self.remnant_right = n % ker.nr();
121 let mut offset = 0;
122 let mut align = std::mem::size_of::<*const ()>();
123 fn ld(spec: usize, uspec: usize, loc: usize) -> LocDependent {
124 LocDependent { spec, ker_spec: uspec, loc, buffer_a: None, buffer_b: None }
125 }
126 for (ix, spec) in specs.iter().enumerate() {
127 offset = offset.next_multiple_of(&align);
128 let ker_spec = match spec {
129 FS::BinScalar(t, op) => match op {
130 BinOp::Min => FKS::ScalarMin(*t.try_as_plain()?.to_scalar()?),
131 BinOp::Max => FKS::ScalarMax(*t.try_as_plain()?.to_scalar()?),
132 BinOp::Mul => FKS::ScalarMul(*t.try_as_plain()?.to_scalar()?),
133 BinOp::Add => FKS::ScalarAdd(*t.try_as_plain()?.to_scalar()?),
134 BinOp::Sub => FKS::ScalarSub(*t.try_as_plain()?.to_scalar()?),
135 BinOp::SubF => FKS::ScalarSubF(*t.try_as_plain()?.to_scalar()?),
136 },
137 FS::ShiftLeft(s) => FKS::ShiftLeft(*s),
138 FS::RoundingShiftRight(s, rp) => FKS::RoundingShiftRight(*s, *rp),
139 FS::QScale(s, rp, m) => FKS::QScale(*s, *rp, *m),
140 FS::BinPerRow(_, _) => {
141 self.loc_dependent.push(ld(ix, self.ker_specs.len(), offset));
142 offset += TI::datum_type().size_of() * ker.mr();
143 FusedKerSpec::Done
144 }
145 FS::BinPerCol(_, _) => {
146 self.loc_dependent.push(ld(ix, self.ker_specs.len(), offset));
147 offset += TI::datum_type().size_of() * ker.nr();
148 FusedKerSpec::Done
149 }
150 FS::AddRowColProducts(_, _) => {
151 self.loc_dependent.push(ld(ix, self.ker_specs.len(), offset));
152 offset += TI::datum_type().size_of() * (ker.mr() + ker.nr());
153 FusedKerSpec::Done
154 }
155 FS::AddUnicast(_) => {
156 self.loc_dependent.push(ld(ix, self.ker_specs.len(), offset));
157 offset += TI::datum_type().size_of() * ker.mr() * ker.nr();
158 FusedKerSpec::Done
159 }
160 FS::Store(store) => {
161 let row_major = ker.stores_row_major_tile()
165 && store.col_byte_stride == store.item_size as isize;
166 let tile_bytes = if row_major {
167 align = align.lcm(&128);
171 offset = Integer::next_multiple_of(&offset, &128);
172 Integer::next_multiple_of(&(store.item_size * ker.nr()), &128) * ker.mr()
173 } else {
174 store.item_size * ker.mr() * ker.nr()
175 };
176 self.loc_dependent.push(ld(ix, self.ker_specs.len(), offset));
177 offset += tile_bytes;
178 FusedKerSpec::Done
179 }
180 FS::LeakyRelu(t) => FKS::LeakyRelu(*t.try_as_plain()?.to_scalar()?),
181 FS::AddMatMul { a, b, packing } => {
182 let mut ld = ld(ix, self.ker_specs.len(), offset);
183 offset += std::mem::size_of::<AddMatMulTemp>();
184 if let Some(tmp) = a.scratch_panel_buffer_layout() {
185 align = tmp.align().lcm(&align);
186 offset = Integer::next_multiple_of(&offset, &tmp.align());
187 ld.buffer_a = Some(offset);
188 offset += tmp.size();
189 }
190 if let Some(tmp) = b.scratch_panel_buffer_layout() {
191 align = tmp.align().lcm(&align);
192 offset = Integer::next_multiple_of(&offset, &tmp.align());
193 ld.buffer_b = Some(offset);
194 offset += tmp.size();
195 }
196 self.loc_dependent.push(ld);
197 FusedKerSpec::AddMatMul {
198 k: 0,
199 pa: std::ptr::null(),
200 pb: std::ptr::null(),
201 packing: *packing,
202 }
203 }
204 };
205 self.ker_specs.push(ker_spec);
206 }
207 self.ker_specs.push(FKS::Done);
208 self.blob_size = offset;
209 self.blob_align = align;
210
211 self.generation = GENERATION.fetch_add(1, std::sync::atomic::Ordering::Relaxed);
212 Ok(())
213 }
214
215 pub unsafe fn run(
216 &self,
217 ker: &impl MatMatMulKer<Acc = TI>,
218 specs: &[FusedSpec],
219 down: usize,
220 right: usize,
221 ) -> TractResult<()> {
222 unsafe {
226 self.run_in_tls_scope(|this, tls| this.run_one_tile(ker, specs, tls, down, right))
227 }
228 }
229
230 pub(crate) unsafe fn run_in_tls_scope<F, R>(&self, f: F) -> R
235 where
236 F: FnOnce(&Self, &mut TLSScratch) -> R,
237 {
238 TLS.with_borrow_mut(|tls| {
239 tls.sync(self);
240 f(self, tls)
241 })
242 }
243
244 #[inline(always)]
247 pub(crate) unsafe fn run_one_tile(
248 &self,
249 ker: &impl MatMatMulKer<Acc = TI>,
250 specs: &[FusedSpec],
251 tls: &mut TLSScratch,
252 down: usize,
253 right: usize,
254 ) -> TractResult<()> {
255 unsafe {
256 if down < self.valid_down_tiles && right < self.valid_right_tiles {
257 self.for_valid_tile(ker, specs, tls, down, right)?;
258 let err = ker.kernel(tls.ker_specs());
259 debug_assert_eq!(err, 0, "Kernel return error {err}");
260 } else {
261 let remnant_down =
262 if down < self.valid_down_tiles { ker.mr() } else { self.remnant_down };
263 let remnant_right =
264 if right < self.valid_right_tiles { ker.nr() } else { self.remnant_right };
265 self.for_border_tile(ker, specs, tls, down, right, remnant_down, remnant_right)?;
266 let err = ker.kernel(tls.ker_specs());
267 debug_assert_eq!(err, 0, "Kernel return error {err}");
268 self.postprocess_tile(specs, tls, down, right, remnant_down, remnant_right)?;
269 }
270 Ok(())
271 }
272 }
273
274 #[inline(always)]
275 unsafe fn for_valid_tile(
276 &self,
277 ker: &impl MatMatMulKer<Acc = TI>,
278 specs: &[FusedSpec],
279 tls: &mut TLSScratch,
280 down: usize,
281 right: usize,
282 ) -> TractResult<()> {
283 unsafe {
284 use FusedKerSpec as FKS;
285 use FusedSpec as FS;
286 let ScratchSpaceImpl { ker_specs, loc_dependent, .. } = self;
287 debug_assert!(specs.len() + 2 == ker_specs.len());
288 for LocDependent { spec, ker_spec, loc, buffer_a, buffer_b } in loc_dependent {
289 let spec = specs.get_unchecked(*spec);
290 let it = match spec {
291 FS::BinPerRow(v, op) => {
292 let v = v.as_ptr_unchecked::<TI>().add(down * ker.mr());
293 match op {
294 BinOp::Min => FKS::PerRowMin(v),
295 BinOp::Max => FKS::PerRowMax(v),
296 BinOp::Add => FKS::PerRowAdd(v),
297 BinOp::Mul => FKS::PerRowMul(v),
298 BinOp::Sub => FKS::PerRowSub(v),
299 BinOp::SubF => FKS::PerRowSubF(v),
300 }
301 }
302 FS::BinPerCol(v, op) => {
303 let v = v.as_ptr_unchecked::<TI>().add(right * ker.nr());
304 match op {
305 BinOp::Min => FKS::PerColMin(v),
306 BinOp::Max => FKS::PerColMax(v),
307 BinOp::Add => FKS::PerColAdd(v),
308 BinOp::Mul => FKS::PerColMul(v),
309 BinOp::Sub => FKS::PerColSub(v),
310 BinOp::SubF => FKS::PerColSubF(v),
311 }
312 }
313 FS::AddRowColProducts(rows, cols) => {
314 let row_ptr = rows.as_ptr_unchecked::<TI>().add(down * ker.mr());
315 let col_ptr = cols.as_ptr_unchecked::<TI>().add(right * ker.nr());
316 FKS::AddRowColProducts(row_ptr, col_ptr)
317 }
318 FS::AddUnicast(store) => FKS::AddUnicast(store.tile_c(down, right)),
319 FS::Store(c_store) => FKS::Store(c_store.tile_c(down, right)),
320 FS::AddMatMul { a, b, packing } => {
321 let scratch = (tls.blob.as_mut_ptr().add(*loc) as *mut AddMatMulTemp)
322 .as_mut()
323 .unwrap();
324 if scratch.panel_a_id != down {
325 scratch.ptr_a = a.panel_bytes(
326 down,
327 buffer_a.map(|o| tls.blob.as_mut_ptr().add(o)),
328 )?;
329 scratch.panel_a_id = down;
330 }
331 if scratch.panel_b_id != right {
332 scratch.ptr_b = b.panel_bytes(
333 right,
334 buffer_b.map(|o| tls.blob.as_mut_ptr().add(o)),
335 )?;
336 scratch.panel_b_id = right;
337 }
338 FKS::AddMatMul {
339 k: b.k(),
340 pa: scratch.ptr_a,
341 pb: scratch.ptr_b,
342 packing: *packing,
343 }
344 }
345 _ => std::hint::unreachable_unchecked(),
346 };
347 *tls.ker_specs().get_unchecked_mut(*ker_spec) = it;
348 }
349 Ok(())
350 }
351 }
352
353 #[inline(never)]
354 #[allow(clippy::too_many_arguments)]
355 unsafe fn for_border_tile(
356 &self,
357 ker: &impl MatMatMulKer<Acc = TI>,
358 specs: &[FusedSpec],
359 tls: &mut TLSScratch,
360 down: usize,
361 right: usize,
362 m_remnant: usize,
363 n_remnant: usize,
364 ) -> TractResult<()> {
365 unsafe {
366 use FusedKerSpec as FKS;
367 use FusedSpec as FS;
368 for LocDependent { spec, ker_spec: uspec, loc, buffer_a, buffer_b } in
369 &self.loc_dependent
370 {
371 let loc = tls.blob.as_mut_ptr().add(*loc);
372 let spec = specs.get_unchecked(*spec);
373 let it = match spec {
374 FS::BinPerRow(v, op) => {
375 let buf = std::slice::from_raw_parts_mut(loc as *mut TI, ker.mr());
376 let ptr = if m_remnant < ker.mr() {
377 if m_remnant > 0 {
378 buf.get_unchecked_mut(..m_remnant).copy_from_slice(
379 v.as_slice_unchecked()
380 .get_unchecked(down * ker.mr()..)
381 .get_unchecked(..m_remnant),
382 );
383 }
384 buf.get_unchecked_mut(m_remnant..)
388 .iter_mut()
389 .for_each(|x| *x = TI::zero());
390 buf.as_ptr()
391 } else {
392 v.as_ptr_unchecked::<TI>().add(down * ker.mr())
393 };
394 match op {
395 BinOp::Min => FKS::PerRowMin(ptr),
396 BinOp::Max => FKS::PerRowMax(ptr),
397 BinOp::Add => FKS::PerRowAdd(ptr),
398 BinOp::Mul => FKS::PerRowMul(ptr),
399 BinOp::Sub => FKS::PerRowSub(ptr),
400 BinOp::SubF => FKS::PerRowSubF(ptr),
401 }
402 }
403 FS::BinPerCol(v, op) => {
404 let buf = std::slice::from_raw_parts_mut(loc as *mut TI, ker.nr());
405 let ptr = if n_remnant < ker.nr() {
406 if n_remnant > 0 {
407 buf.get_unchecked_mut(..n_remnant).copy_from_slice(
408 v.as_slice_unchecked()
409 .get_unchecked(right * ker.nr()..)
410 .get_unchecked(..n_remnant),
411 );
412 }
413 buf.get_unchecked_mut(n_remnant..)
414 .iter_mut()
415 .for_each(|x| *x = TI::zero());
416 buf.as_ptr()
417 } else {
418 v.as_ptr_unchecked::<TI>().add(right * ker.nr())
419 };
420 match op {
421 BinOp::Min => FKS::PerColMin(ptr),
422 BinOp::Max => FKS::PerColMax(ptr),
423 BinOp::Add => FKS::PerColAdd(ptr),
424 BinOp::Mul => FKS::PerColMul(ptr),
425 BinOp::Sub => FKS::PerColSub(ptr),
426 BinOp::SubF => FKS::PerColSubF(ptr),
427 }
428 }
429 FS::AddRowColProducts(rows, cols) => {
430 let r = std::slice::from_raw_parts_mut(loc as *mut TI, ker.mr());
431 let row_ptr = if m_remnant < ker.mr() {
432 r.get_unchecked_mut(..m_remnant).copy_from_slice(
433 rows.as_slice_unchecked()
434 .get_unchecked(down * ker.mr()..)
435 .get_unchecked(..m_remnant),
436 );
437 r.get_unchecked_mut(m_remnant..)
438 .iter_mut()
439 .for_each(|x| *x = TI::zero());
440 r.as_ptr()
441 } else {
442 rows.as_ptr_unchecked::<TI>().add(down * ker.mr())
443 };
444 let c = std::slice::from_raw_parts_mut(
445 (loc as *mut TI).add(ker.mr()),
446 ker.nr(),
447 );
448 let col_ptr = if n_remnant < ker.nr() {
449 c.get_unchecked_mut(..n_remnant).copy_from_slice(
450 cols.as_slice_unchecked()
451 .get_unchecked(right * ker.nr()..)
452 .get_unchecked(..n_remnant),
453 );
454 c.get_unchecked_mut(n_remnant..)
455 .iter_mut()
456 .for_each(|x| *x = TI::zero());
457 c.as_ptr()
458 } else {
459 cols.as_ptr_unchecked::<TI>().add(right * ker.nr())
460 };
461 FKS::AddRowColProducts(row_ptr, col_ptr)
462 }
463 FS::AddUnicast(store) => {
464 let row_byte_stride = store.row_byte_stride;
465 let col_byte_stride = store.col_byte_stride;
466 let tile_offset = row_byte_stride * down as isize * ker.mr() as isize
467 + col_byte_stride * right as isize * ker.nr() as isize;
468 let tile_ptr = store.ptr.offset(tile_offset);
469 let tmp_d_tile =
470 std::slice::from_raw_parts_mut(loc as *mut TI, ker.mr() * ker.nr());
471 tmp_d_tile.iter_mut().for_each(|t| *t = TI::zero());
472 for r in 0..m_remnant as isize {
473 for c in 0..n_remnant as isize {
474 let inner_offset = c * col_byte_stride + r * row_byte_stride;
475 if inner_offset + tile_offset
476 < (store.item_size * store.item_count) as isize
477 {
478 *tmp_d_tile
479 .get_unchecked_mut(r as usize + c as usize * ker.mr()) =
480 *(tile_ptr.offset(inner_offset) as *const TI);
481 }
482 }
483 }
484 FKS::AddUnicast(OutputStoreKer {
485 ptr: tmp_d_tile.as_ptr() as _,
486 row_byte_stride: std::mem::size_of::<TI>() as isize,
487 col_byte_stride: (std::mem::size_of::<TI>() * ker.mr()) as isize,
488 item_size: std::mem::size_of::<TI>(),
489 })
490 }
491 FS::Store(c_store) => {
492 let row_major = ker.stores_row_major_tile()
493 && c_store.col_byte_stride == c_store.item_size as isize;
494 let (row_byte_stride, col_byte_stride) = if row_major {
495 let row =
498 Integer::next_multiple_of(&(c_store.item_size * ker.nr()), &128);
499 (row as isize, c_store.item_size as isize)
500 } else {
501 (c_store.item_size as isize, (c_store.item_size * ker.mr()) as isize)
502 };
503 let tmpc = OutputStoreKer {
504 ptr: loc as _,
505 item_size: c_store.item_size,
506 row_byte_stride,
507 col_byte_stride,
508 };
509 FKS::Store(tmpc)
510 }
511 FS::AddMatMul { a, b, packing } => {
512 let scratch = (loc as *mut AddMatMulTemp).as_mut().unwrap();
513 if scratch.panel_a_id != down {
514 scratch.ptr_a = a.panel_bytes(
515 down,
516 buffer_a.map(|o| tls.blob.as_mut_ptr().add(o)),
517 )?;
518 scratch.panel_a_id = down;
519 }
520 if scratch.panel_b_id != right {
521 scratch.ptr_b = b.panel_bytes(
522 right,
523 buffer_b.map(|o| tls.blob.as_mut_ptr().add(o)),
524 )?;
525 scratch.panel_b_id = right;
526 }
527 FKS::AddMatMul {
528 k: b.k(),
529 pa: scratch.ptr_a,
530 pb: scratch.ptr_b,
531 packing: *packing,
532 }
533 }
534 _ => std::hint::unreachable_unchecked(),
535 };
536 *tls.ker_specs().get_unchecked_mut(*uspec) = it;
537 }
538 Ok(())
539 }
540 }
541
542 #[inline]
543 pub fn uspecs(&self) -> &[FusedKerSpec<TI>] {
544 &self.ker_specs
545 }
546
547 unsafe fn postprocess_tile(
548 &self,
549 specs: &[FusedSpec],
550 tls: &mut TLSScratch,
551 down: usize,
552 right: usize,
553 m_remnant: usize,
554 n_remnant: usize,
555 ) -> TractResult<()>
556 where
557 TI: LADatum,
558 {
559 unsafe {
560 for LocDependent { spec, ker_spec: uspec, .. } in self.loc_dependent.iter() {
561 let spec = specs.get_unchecked(*spec);
562 let ker_spec = tls.ker_specs::<TI>().get_unchecked(*uspec);
563 if let (FusedSpec::Store(c_store), FusedKerSpec::Store(tmp)) = (spec, ker_spec) {
564 c_store.set_from_tile(down, right, m_remnant, n_remnant, tmp)
565 }
566 }
567 Ok(())
568 }
569 }
570}