Skip to main content

hybit_matrix/
abtm.rs

1use std::collections::BTreeMap;
2
3use crate::{Csr32Matrix, DofMask};
4use hybit_core::{HybitError, LinearOperator};
5
6pub const TILE_WIDTH: usize = 64;
7const KIND_SHIFT: u32 = 30;
8const BASE_WORD_MASK: u32 = (1u32 << KIND_SHIFT) - 1;
9
10#[repr(u8)]
11#[derive(Clone, Copy, Debug, PartialEq, Eq)]
12pub enum TileKind {
13    Sparse = 0,
14    Bitmap = 1,
15    Dense = 2,
16}
17
18#[repr(C)]
19#[derive(Clone, Copy, Debug)]
20pub struct TileDesc {
21    pub mask: u64,
22    pub value_offset: u32,
23    meta: u32,
24}
25
26impl TileDesc {
27    fn new(
28        base_col: usize,
29        value_offset: usize,
30        mask: u64,
31        kind: TileKind,
32    ) -> Result<Self, HybitError> {
33        let word = base_col / TILE_WIDTH;
34        if word > BASE_WORD_MASK as usize || value_offset > u32::MAX as usize {
35            return Err(HybitError::SizeOverflow);
36        }
37        Ok(Self {
38            mask,
39            value_offset: value_offset as u32,
40            meta: word as u32 | ((kind as u32) << KIND_SHIFT),
41        })
42    }
43
44    #[inline(always)]
45    pub fn kind(&self) -> TileKind {
46        match self.meta >> KIND_SHIFT {
47            0 => TileKind::Sparse,
48            1 => TileKind::Bitmap,
49            2 => TileKind::Dense,
50            _ => unreachable!("invalid tile kind"),
51        }
52    }
53
54    #[inline(always)]
55    pub fn base_col(&self) -> usize {
56        ((self.meta & BASE_WORD_MASK) as usize) * TILE_WIDTH
57    }
58}
59
60#[derive(Clone, Copy, Debug)]
61pub struct AbtmConfig {
62    pub sparse_max_nnz: u8,
63    pub dense_min_nnz: u8,
64}
65
66impl Default for AbtmConfig {
67    fn default() -> Self {
68        Self {
69            sparse_max_nnz: 8,
70            dense_min_nnz: 40,
71        }
72    }
73}
74
75impl AbtmConfig {
76    fn validate(self) -> Result<Self, HybitError> {
77        if self.sparse_max_nnz == 0
78            || self.dense_min_nnz as usize > TILE_WIDTH
79            || self.sparse_max_nnz >= self.dense_min_nnz
80        {
81            return Err(HybitError::InvalidArgument(
82                "invalid ABTM density thresholds",
83            ));
84        }
85        Ok(self)
86    }
87}
88
89#[derive(Clone, Debug, Default)]
90pub struct AbtmStats {
91    pub nrows: usize,
92    pub ncols: usize,
93    pub matrix_nnz: usize,
94    pub tiles: usize,
95    pub sparse_tiles: usize,
96    pub bitmap_tiles: usize,
97    pub dense_tiles: usize,
98    pub value_slots: usize,
99    pub metadata_bytes: usize,
100}
101
102impl AbtmStats {
103    pub fn metadata_bytes_per_nnz(&self) -> f64 {
104        if self.matrix_nnz == 0 {
105            0.0
106        } else {
107            self.metadata_bytes as f64 / self.matrix_nnz as f64
108        }
109    }
110    pub fn total_estimated_bytes(&self) -> usize {
111        self.metadata_bytes + self.value_slots * std::mem::size_of::<f64>()
112    }
113}
114
115#[derive(Clone, Debug)]
116pub struct AbtmMatrix {
117    nrows: usize,
118    ncols: usize,
119    row_tile_ptr: Vec<u32>,
120    tiles: Vec<TileDesc>,
121    values: Vec<f64>,
122    matrix_nnz: usize,
123}
124
125impl AbtmMatrix {
126    pub fn from_csr32(csr: &Csr32Matrix, config: AbtmConfig) -> Result<Self, HybitError> {
127        csr.validate()?;
128        let config = config.validate()?;
129        let mut row_tile_ptr = Vec::with_capacity(csr.nrows() + 1);
130        let mut tiles = Vec::new();
131        let mut values = Vec::new();
132        let mut matrix_nnz = 0usize;
133        row_tile_ptr.push(0);
134
135        for row in 0..csr.nrows() {
136            let start = csr.row_ptr()[row] as usize;
137            let end = csr.row_ptr()[row + 1] as usize;
138            let mut blocks: BTreeMap<usize, BTreeMap<u8, f64>> = BTreeMap::new();
139            for p in start..end {
140                let col = csr.col_idx()[p] as usize;
141                let value = csr.values()[p];
142                if value == 0.0 {
143                    continue;
144                }
145                let block = col / TILE_WIDTH;
146                let offset = (col % TILE_WIDTH) as u8;
147                *blocks.entry(block).or_default().entry(offset).or_default() += value;
148            }
149
150            for (block, entries) in blocks {
151                let canonical: Vec<(u8, f64)> =
152                    entries.into_iter().filter(|(_, v)| *v != 0.0).collect();
153                if canonical.is_empty() {
154                    continue;
155                }
156                let nnz = canonical.len();
157                matrix_nnz += nnz;
158                let base_col = block
159                    .checked_mul(TILE_WIDTH)
160                    .ok_or(HybitError::SizeOverflow)?;
161                let value_offset = values.len();
162                let mut mask = 0u64;
163                for &(offset, _) in &canonical {
164                    mask |= 1u64 << offset;
165                }
166                let kind = if nnz <= config.sparse_max_nnz as usize {
167                    TileKind::Sparse
168                } else if nnz >= config.dense_min_nnz as usize {
169                    TileKind::Dense
170                } else {
171                    TileKind::Bitmap
172                };
173                match kind {
174                    TileKind::Sparse | TileKind::Bitmap => {
175                        values.extend(canonical.iter().map(|&(_, value)| value));
176                    }
177                    TileKind::Dense => {
178                        let mut dense = [0.0f64; TILE_WIDTH];
179                        for &(offset, value) in &canonical {
180                            dense[offset as usize] = value;
181                        }
182                        values.extend_from_slice(&dense);
183                    }
184                }
185                tiles.push(TileDesc::new(base_col, value_offset, mask, kind)?);
186            }
187            if tiles.len() > u32::MAX as usize {
188                return Err(HybitError::SizeOverflow);
189            }
190            row_tile_ptr.push(tiles.len() as u32);
191        }
192
193        Ok(Self {
194            nrows: csr.nrows(),
195            ncols: csr.ncols(),
196            row_tile_ptr,
197            tiles,
198            values,
199            matrix_nnz,
200        })
201    }
202
203    pub fn expand_mask_one_hop(&self, mask: &DofMask) -> Result<DofMask, HybitError> {
204        if self.nrows != self.ncols {
205            return Err(HybitError::InvalidMatrix(
206                "mask expansion requires a square ABTM matrix",
207            ));
208        }
209        if mask.len() != self.nrows {
210            return Err(HybitError::DimensionMismatch {
211                expected: self.nrows,
212                actual: mask.len(),
213            });
214        }
215        let mut expanded = mask.clone();
216        for row in mask.indices() {
217            let start = self.row_tile_ptr[row] as usize;
218            let end = self.row_tile_ptr[row + 1] as usize;
219            for tile in &self.tiles[start..end] {
220                expanded.or_word(tile.base_col() / TILE_WIDTH, tile.mask);
221            }
222        }
223        Ok(expanded)
224    }
225
226    pub fn stats(&self) -> AbtmStats {
227        let mut stats = AbtmStats {
228            nrows: self.nrows,
229            ncols: self.ncols,
230            matrix_nnz: self.matrix_nnz,
231            tiles: self.tiles.len(),
232            value_slots: self.values.len(),
233            metadata_bytes: self.row_tile_ptr.len() * std::mem::size_of::<u32>()
234                + self.tiles.len() * std::mem::size_of::<TileDesc>(),
235            ..AbtmStats::default()
236        };
237        for tile in &self.tiles {
238            match tile.kind() {
239                TileKind::Sparse => stats.sparse_tiles += 1,
240                TileKind::Bitmap => stats.bitmap_tiles += 1,
241                TileKind::Dense => stats.dense_tiles += 1,
242            }
243        }
244        stats
245    }
246
247    #[inline(always)]
248    fn apply_compact_tile(&self, tile: &TileDesc, x: &[f64]) -> f64 {
249        let mut bits = tile.mask;
250        let mut p = tile.value_offset as usize;
251        let base = tile.base_col();
252        let mut sum = 0.0;
253        while bits != 0 {
254            let bit = bits.trailing_zeros() as usize;
255            sum += self.values[p] * x[base + bit];
256            p += 1;
257            bits &= bits - 1;
258        }
259        sum
260    }
261}
262
263impl LinearOperator for AbtmMatrix {
264    fn rows(&self) -> usize {
265        self.nrows
266    }
267    fn cols(&self) -> usize {
268        self.ncols
269    }
270
271    fn apply(&self, x: &[f64], y: &mut [f64]) -> Result<(), HybitError> {
272        if x.len() != self.ncols {
273            return Err(HybitError::DimensionMismatch {
274                expected: self.ncols,
275                actual: x.len(),
276            });
277        }
278        if y.len() != self.nrows {
279            return Err(HybitError::DimensionMismatch {
280                expected: self.nrows,
281                actual: y.len(),
282            });
283        }
284        for (row, yi) in y.iter_mut().enumerate() {
285            let start = self.row_tile_ptr[row] as usize;
286            let end = self.row_tile_ptr[row + 1] as usize;
287            let mut sum = 0.0;
288            for tile in &self.tiles[start..end] {
289                match tile.kind() {
290                    TileKind::Sparse | TileKind::Bitmap => sum += self.apply_compact_tile(tile, x),
291                    TileKind::Dense => {
292                        let base = tile.base_col();
293                        let vo = tile.value_offset as usize;
294                        let limit = (self.ncols - base).min(TILE_WIDTH);
295                        for j in 0..limit {
296                            sum += self.values[vo + j] * x[base + j];
297                        }
298                    }
299                }
300            }
301            *yi = sum;
302        }
303        Ok(())
304    }
305}
306
307#[cfg(test)]
308mod tests {
309    use super::*;
310
311    #[test]
312    fn tile_descriptor_is_16_bytes() {
313        assert_eq!(std::mem::size_of::<TileDesc>(), 16);
314    }
315
316    #[test]
317    fn bitmap_topology_expands_one_hop() {
318        let csr = Csr32Matrix::new(
319            4,
320            4,
321            vec![0, 2, 5, 8, 10],
322            vec![0, 1, 0, 1, 2, 1, 2, 3, 2, 3],
323            vec![4.0, -1.0, -1.0, 4.0, -1.0, -1.0, 4.0, -1.0, -1.0, 3.0],
324        )
325        .unwrap();
326        let abtm = AbtmMatrix::from_csr32(&csr, AbtmConfig::default()).unwrap();
327        let seed = DofMask::from_indices(4, &[1]).unwrap();
328        let expanded = abtm.expand_mask_one_hop(&seed).unwrap();
329        assert_eq!(expanded.indices(), vec![0, 1, 2]);
330    }
331
332    #[test]
333    fn abtm_matches_csr() {
334        let csr = Csr32Matrix::new(
335            4,
336            4,
337            vec![0, 2, 5, 8, 10],
338            vec![0, 1, 0, 1, 2, 1, 2, 3, 2, 3],
339            vec![4.0, -1.0, -1.0, 4.0, -1.0, -1.0, 4.0, -1.0, -1.0, 3.0],
340        )
341        .unwrap();
342        let abtm = AbtmMatrix::from_csr32(&csr, AbtmConfig::default()).unwrap();
343        let x = [1.0, 2.0, 3.0, 4.0];
344        let mut yc = vec![0.0; 4];
345        let mut ya = vec![0.0; 4];
346        csr.apply(&x, &mut yc).unwrap();
347        abtm.apply(&x, &mut ya).unwrap();
348        for (a, b) in yc.iter().zip(&ya) {
349            assert!((a - b).abs() < 1.0e-12);
350        }
351    }
352}