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}