Skip to main content

tract_linalg/frame/mmm/
scratch.rs

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    // offset for the location dependent structure
86    loc: usize,
87    // offset of its associated dynamic-size buffers
88    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                    // Only worth a row-major tile when the destination is itself
162                    // n-contiguous: set_from_tile then copies without transposing
163                    // and the kernel's aligned bulk-store path applies.
164                    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                        // 128-align the row-major store tile, with rows padded to
168                        // 128 bytes, so the kernel can hit its aligned bulk-store
169                        // path (e.g. Apple AMX stz-direct).
170                        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        // Per-tile entry: enter the TLS scope (does sync once) then run a single
223        // tile. Single-threaded callers should prefer `run_in_tls_scope`+
224        // `run_one_tile` to amortise the TLS borrow + sync over many tiles.
225        unsafe {
226            self.run_in_tls_scope(|this, tls| this.run_one_tile(ker, specs, tls, down, right))
227        }
228    }
229
230    /// Borrow the per-thread scratch blob for a single MMM call and `sync` it
231    /// once. The closure is invoked once with a mutable reference to the TLS
232    /// scratch and to `self`. Used by single-threaded matmul drivers to avoid
233    /// re-entering TLS / re-running `sync` per tile.
234    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    /// Run a single tile against an already-borrowed TLS scratch. Caller is
245    /// responsible for entering `run_in_tls_scope` first (so `sync` has run).
246    #[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                            // The kernel computes on the tail lanes before their
385                            // results are discarded; garbage there decodes to
386                            // denormals and stalls the fp pipeline. Zero them.
387                            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                            // Pad the row stride to 128 bytes so it stays aligned
496                            // for the kernel's bulk-store path.
497                            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}