Skip to main content

haagenti_sparse/
mask.rs

1//! Attention masks for sparse computation
2
3use crate::{CategoryMapping, HeadCategory, Result, SparseError, MIN_ACTIVE_HEADS};
4use serde::{Deserialize, Serialize};
5use std::collections::HashMap;
6
7/// A mask indicating which attention heads to compute
8#[derive(Debug, Clone, Serialize, Deserialize)]
9pub struct AttentionMask {
10    /// Number of heads
11    pub num_heads: usize,
12    /// Number of layers
13    pub num_layers: usize,
14    /// Mask values: true = compute, false = skip
15    /// Indexed as `[layer][head]`
16    pub mask: Vec<Vec<bool>>,
17    /// Per-layer sparsity (fraction of heads skipped)
18    pub layer_sparsity: Vec<f32>,
19    /// Overall sparsity
20    pub overall_sparsity: f32,
21}
22
23impl AttentionMask {
24    /// Create a mask with all heads active
25    pub fn all_active(num_heads: usize, num_layers: usize) -> Self {
26        let mask = vec![vec![true; num_heads]; num_layers];
27        Self {
28            num_heads,
29            num_layers,
30            mask,
31            layer_sparsity: vec![0.0; num_layers],
32            overall_sparsity: 0.0,
33        }
34    }
35
36    /// Create a mask with uniform random sparsity
37    pub fn random(num_heads: usize, num_layers: usize, sparsity: f32) -> Self {
38        use rand::Rng;
39        let mut rng = rand::thread_rng();
40
41        let mask: Vec<Vec<bool>> = (0..num_layers)
42            .map(|_| {
43                (0..num_heads)
44                    .map(|_| rng.gen::<f32>() > sparsity)
45                    .collect()
46            })
47            .collect();
48
49        Self::from_mask(mask)
50    }
51
52    /// Create from raw mask data
53    pub fn from_mask(mask: Vec<Vec<bool>>) -> Self {
54        let num_layers = mask.len();
55        let num_heads = mask.first().map(|l| l.len()).unwrap_or(0);
56
57        let layer_sparsity: Vec<f32> = mask
58            .iter()
59            .map(|layer| {
60                let inactive = layer.iter().filter(|&&active| !active).count();
61                inactive as f32 / layer.len() as f32
62            })
63            .collect();
64
65        let total_heads = num_heads * num_layers;
66        let total_inactive: usize = mask
67            .iter()
68            .flat_map(|layer| layer.iter())
69            .filter(|&&active| !active)
70            .count();
71        let overall_sparsity = total_inactive as f32 / total_heads as f32;
72
73        Self {
74            num_heads,
75            num_layers,
76            mask,
77            layer_sparsity,
78            overall_sparsity,
79        }
80    }
81
82    /// Check if a head is active
83    pub fn is_active(&self, layer: usize, head: usize) -> bool {
84        self.mask
85            .get(layer)
86            .and_then(|l| l.get(head))
87            .copied()
88            .unwrap_or(true)
89    }
90
91    /// Set head activity
92    pub fn set_active(&mut self, layer: usize, head: usize, active: bool) {
93        if let Some(l) = self.mask.get_mut(layer) {
94            if let Some(h) = l.get_mut(head) {
95                *h = active;
96            }
97        }
98        self.update_sparsity();
99    }
100
101    /// Get active head indices for a layer
102    pub fn active_heads(&self, layer: usize) -> Vec<usize> {
103        self.mask
104            .get(layer)
105            .map(|l| {
106                l.iter()
107                    .enumerate()
108                    .filter(|(_, &active)| active)
109                    .map(|(i, _)| i)
110                    .collect()
111            })
112            .unwrap_or_default()
113    }
114
115    /// Get inactive head indices for a layer
116    pub fn inactive_heads(&self, layer: usize) -> Vec<usize> {
117        self.mask
118            .get(layer)
119            .map(|l| {
120                l.iter()
121                    .enumerate()
122                    .filter(|(_, &active)| !active)
123                    .map(|(i, _)| i)
124                    .collect()
125            })
126            .unwrap_or_default()
127    }
128
129    /// Number of active heads in a layer
130    pub fn active_count(&self, layer: usize) -> usize {
131        self.mask
132            .get(layer)
133            .map(|l| l.iter().filter(|&&a| a).count())
134            .unwrap_or(0)
135    }
136
137    /// Update sparsity metrics
138    fn update_sparsity(&mut self) {
139        self.layer_sparsity = self
140            .mask
141            .iter()
142            .map(|layer| {
143                let inactive = layer.iter().filter(|&&active| !active).count();
144                inactive as f32 / layer.len() as f32
145            })
146            .collect();
147
148        let total_heads = self.num_heads * self.num_layers;
149        let total_inactive: usize = self
150            .mask
151            .iter()
152            .flat_map(|layer| layer.iter())
153            .filter(|&&active| !active)
154            .count();
155        self.overall_sparsity = total_inactive as f32 / total_heads as f32;
156    }
157
158    /// Merge with another mask (AND operation - both must be active)
159    pub fn merge_and(&self, other: &AttentionMask) -> Result<Self> {
160        if self.num_heads != other.num_heads || self.num_layers != other.num_layers {
161            return Err(SparseError::InvalidDimensions {
162                expected_heads: self.num_heads,
163                expected_layers: self.num_layers,
164                actual_heads: other.num_heads,
165                actual_layers: other.num_layers,
166            });
167        }
168
169        let mask: Vec<Vec<bool>> = self
170            .mask
171            .iter()
172            .zip(other.mask.iter())
173            .map(|(l1, l2)| l1.iter().zip(l2.iter()).map(|(&a, &b)| a && b).collect())
174            .collect();
175
176        Ok(Self::from_mask(mask))
177    }
178
179    /// Merge with another mask (OR operation - either can be active)
180    pub fn merge_or(&self, other: &AttentionMask) -> Result<Self> {
181        if self.num_heads != other.num_heads || self.num_layers != other.num_layers {
182            return Err(SparseError::InvalidDimensions {
183                expected_heads: self.num_heads,
184                expected_layers: self.num_layers,
185                actual_heads: other.num_heads,
186                actual_layers: other.num_layers,
187            });
188        }
189
190        let mask: Vec<Vec<bool>> = self
191            .mask
192            .iter()
193            .zip(other.mask.iter())
194            .map(|(l1, l2)| l1.iter().zip(l2.iter()).map(|(&a, &b)| a || b).collect())
195            .collect();
196
197        Ok(Self::from_mask(mask))
198    }
199
200    /// Ensure minimum active heads per layer
201    pub fn ensure_minimum(&mut self, min_active: usize) {
202        use rand::seq::SliceRandom;
203        let mut rng = rand::thread_rng();
204
205        for layer_mask in &mut self.mask {
206            let active_count = layer_mask.iter().filter(|&&a| a).count();
207            if active_count < min_active {
208                // Randomly activate more heads
209                let mut inactive: Vec<usize> = layer_mask
210                    .iter()
211                    .enumerate()
212                    .filter(|(_, &a)| !a)
213                    .map(|(i, _)| i)
214                    .collect();
215                inactive.shuffle(&mut rng);
216
217                for &idx in inactive.iter().take(min_active - active_count) {
218                    layer_mask[idx] = true;
219                }
220            }
221        }
222
223        self.update_sparsity();
224    }
225
226    /// Convert to compact byte representation
227    pub fn to_bytes(&self) -> Vec<u8> {
228        let mut bytes = Vec::new();
229
230        // Header: num_heads (u16), num_layers (u16)
231        bytes.extend_from_slice(&(self.num_heads as u16).to_le_bytes());
232        bytes.extend_from_slice(&(self.num_layers as u16).to_le_bytes());
233
234        // Bit-packed mask
235        for layer in &self.mask {
236            for chunk in layer.chunks(8) {
237                let mut byte = 0u8;
238                for (i, &active) in chunk.iter().enumerate() {
239                    if active {
240                        byte |= 1 << i;
241                    }
242                }
243                bytes.push(byte);
244            }
245        }
246
247        bytes
248    }
249
250    /// Deserialize from bytes
251    pub fn from_bytes(bytes: &[u8]) -> Result<Self> {
252        if bytes.len() < 4 {
253            return Err(SparseError::InvalidDimensions {
254                expected_heads: 0,
255                expected_layers: 0,
256                actual_heads: 0,
257                actual_layers: 0,
258            });
259        }
260
261        let num_heads = u16::from_le_bytes([bytes[0], bytes[1]]) as usize;
262        let num_layers = u16::from_le_bytes([bytes[2], bytes[3]]) as usize;
263
264        let bytes_per_layer = num_heads.div_ceil(8);
265        let mut mask = Vec::with_capacity(num_layers);
266
267        let mut offset = 4;
268        for _ in 0..num_layers {
269            let mut layer = Vec::with_capacity(num_heads);
270            for byte_idx in 0..bytes_per_layer {
271                if offset + byte_idx >= bytes.len() {
272                    break;
273                }
274                let byte = bytes[offset + byte_idx];
275                for bit in 0..8 {
276                    if layer.len() < num_heads {
277                        layer.push((byte >> bit) & 1 == 1);
278                    }
279                }
280            }
281            mask.push(layer);
282            offset += bytes_per_layer;
283        }
284
285        Ok(Self::from_mask(mask))
286    }
287}
288
289impl Default for AttentionMask {
290    /// Creates an empty attention mask with no heads or layers.
291    ///
292    /// Use `AttentionMask::all_active(num_heads, num_layers)` to create
293    /// a mask with specific dimensions where all heads are active.
294    fn default() -> Self {
295        Self {
296            num_heads: 0,
297            num_layers: 0,
298            mask: Vec::new(),
299            layer_sparsity: Vec::new(),
300            overall_sparsity: 0.0,
301        }
302    }
303}
304
305/// Pattern for mask generation
306#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize)]
307pub enum MaskPattern {
308    /// Skip every Nth head
309    Strided { stride: usize },
310    /// Skip first N heads per layer
311    SkipFirst { count: usize },
312    /// Skip last N heads per layer
313    SkipLast { count: usize },
314    /// Skip heads below importance threshold
315    Threshold,
316    /// Use category-based pruning
317    CategoryBased,
318}
319
320/// Builder for attention masks
321#[derive(Debug, Clone)]
322pub struct MaskBuilder {
323    num_heads: usize,
324    num_layers: usize,
325    pattern: MaskPattern,
326    target_sparsity: f32,
327    min_active: usize,
328    category_weights: Option<HashMap<HeadCategory, f32>>,
329    category_mapping: Option<CategoryMapping>,
330}
331
332impl MaskBuilder {
333    /// Create a new builder
334    pub fn new(num_heads: usize, num_layers: usize) -> Self {
335        Self {
336            num_heads,
337            num_layers,
338            pattern: MaskPattern::CategoryBased,
339            target_sparsity: 0.5,
340            min_active: MIN_ACTIVE_HEADS,
341            category_weights: None,
342            category_mapping: None,
343        }
344    }
345
346    /// Set target sparsity
347    pub fn sparsity(mut self, sparsity: f32) -> Self {
348        self.target_sparsity = sparsity.clamp(0.0, 0.9);
349        self
350    }
351
352    /// Set mask pattern
353    pub fn pattern(mut self, pattern: MaskPattern) -> Self {
354        self.pattern = pattern;
355        self
356    }
357
358    /// Set minimum active heads
359    pub fn min_active(mut self, min: usize) -> Self {
360        self.min_active = min;
361        self
362    }
363
364    /// Set category weights for category-based masking
365    pub fn category_weights(mut self, weights: HashMap<HeadCategory, f32>) -> Self {
366        self.category_weights = Some(weights);
367        self
368    }
369
370    /// Set category mapping
371    pub fn category_mapping(mut self, mapping: CategoryMapping) -> Self {
372        self.category_mapping = Some(mapping);
373        self
374    }
375
376    /// Build the mask
377    pub fn build(self) -> AttentionMask {
378        let mut mask = match self.pattern {
379            MaskPattern::Strided { stride } => self.build_strided(stride),
380            MaskPattern::SkipFirst { count } => self.build_skip_first(count),
381            MaskPattern::SkipLast { count } => self.build_skip_last(count),
382            MaskPattern::Threshold => self.build_threshold(),
383            MaskPattern::CategoryBased => self.build_category_based(),
384        };
385
386        mask.ensure_minimum(self.min_active);
387        mask
388    }
389
390    fn build_strided(&self, stride: usize) -> AttentionMask {
391        let mask: Vec<Vec<bool>> = (0..self.num_layers)
392            .map(|_| (0..self.num_heads).map(|head| head % stride != 0).collect())
393            .collect();
394        AttentionMask::from_mask(mask)
395    }
396
397    fn build_skip_first(&self, count: usize) -> AttentionMask {
398        let mask: Vec<Vec<bool>> = (0..self.num_layers)
399            .map(|_| (0..self.num_heads).map(|head| head >= count).collect())
400            .collect();
401        AttentionMask::from_mask(mask)
402    }
403
404    fn build_skip_last(&self, count: usize) -> AttentionMask {
405        let mask: Vec<Vec<bool>> = (0..self.num_layers)
406            .map(|_| {
407                (0..self.num_heads)
408                    .map(|head| head < self.num_heads - count)
409                    .collect()
410            })
411            .collect();
412        AttentionMask::from_mask(mask)
413    }
414
415    fn build_threshold(&self) -> AttentionMask {
416        // Use uniform random threshold matching target sparsity
417        AttentionMask::random(self.num_heads, self.num_layers, self.target_sparsity)
418    }
419
420    fn build_category_based(&self) -> AttentionMask {
421        let mapping = self
422            .category_mapping
423            .clone()
424            .unwrap_or_else(CategoryMapping::sdxl_default);
425        let weights = self.category_weights.clone().unwrap_or_default();
426
427        let heads_to_skip = (self.num_heads as f32 * self.target_sparsity) as usize;
428
429        let mask: Vec<Vec<bool>> = (0..self.num_layers)
430            .map(|layer| {
431                // Score each head by its category weight
432                let mut head_scores: Vec<(usize, f32)> = (0..self.num_heads)
433                    .map(|head| {
434                        let category = mapping
435                            .get_category(layer, head)
436                            .unwrap_or(HeadCategory::General);
437                        let weight = weights
438                            .get(&category)
439                            .copied()
440                            .unwrap_or_else(|| category.default_importance());
441                        (head, weight)
442                    })
443                    .collect();
444
445                // Sort by weight (lowest first - these will be skipped)
446                head_scores.sort_by(|a, b| a.1.partial_cmp(&b.1).unwrap());
447
448                // Create mask
449                let skip_heads: std::collections::HashSet<usize> = head_scores
450                    .iter()
451                    .take(heads_to_skip)
452                    .map(|(head, _)| *head)
453                    .collect();
454
455                (0..self.num_heads)
456                    .map(|head| !skip_heads.contains(&head))
457                    .collect()
458            })
459            .collect();
460
461        AttentionMask::from_mask(mask)
462    }
463}
464
465#[cfg(test)]
466mod tests {
467    use super::*;
468
469    #[test]
470    fn test_default() {
471        let mask = AttentionMask::default();
472
473        assert_eq!(mask.num_heads, 0);
474        assert_eq!(mask.num_layers, 0);
475        assert!(mask.mask.is_empty());
476        assert!(mask.layer_sparsity.is_empty());
477        assert_eq!(mask.overall_sparsity, 0.0);
478    }
479
480    #[test]
481    fn test_all_active() {
482        let mask = AttentionMask::all_active(32, 10);
483        assert_eq!(mask.overall_sparsity, 0.0);
484        assert_eq!(mask.active_count(0), 32);
485    }
486
487    #[test]
488    fn test_strided_mask() {
489        let mask = MaskBuilder::new(32, 10)
490            .pattern(MaskPattern::Strided { stride: 2 })
491            .min_active(0)
492            .build();
493
494        // Every other head should be skipped
495        for layer in 0..10 {
496            assert!(!mask.is_active(layer, 0)); // Skipped
497            assert!(mask.is_active(layer, 1)); // Active
498            assert!(!mask.is_active(layer, 2)); // Skipped
499        }
500    }
501
502    #[test]
503    fn test_minimum_active() {
504        let mut mask = AttentionMask::from_mask(vec![vec![false; 32]; 10]);
505        mask.ensure_minimum(8);
506
507        for layer in 0..10 {
508            assert!(mask.active_count(layer) >= 8);
509        }
510    }
511
512    #[test]
513    fn test_byte_serialization() {
514        let original = AttentionMask::random(32, 10, 0.5);
515        let bytes = original.to_bytes();
516        let restored = AttentionMask::from_bytes(&bytes).unwrap();
517
518        assert_eq!(original.num_heads, restored.num_heads);
519        assert_eq!(original.num_layers, restored.num_layers);
520        for layer in 0..10 {
521            for head in 0..32 {
522                assert_eq!(
523                    original.is_active(layer, head),
524                    restored.is_active(layer, head)
525                );
526            }
527        }
528    }
529
530    #[test]
531    fn test_merge() {
532        let mask1 = MaskBuilder::new(32, 10)
533            .pattern(MaskPattern::SkipFirst { count: 8 })
534            .min_active(0)
535            .build();
536        let mask2 = MaskBuilder::new(32, 10)
537            .pattern(MaskPattern::SkipLast { count: 8 })
538            .min_active(0)
539            .build();
540
541        // AND: only middle 16 heads active
542        let merged_and = mask1.merge_and(&mask2).unwrap();
543        assert_eq!(merged_and.active_count(0), 16);
544
545        // OR: all but 0 heads active (first 8 or last 8 cover everything... wait no)
546        // Actually: SkipFirst skips 0-7, SkipLast skips 24-31
547        // OR means: active if either mask says so
548        // mask1 active: 8-31, mask2 active: 0-23
549        // union: 0-31 = all 32
550        let merged_or = mask1.merge_or(&mask2).unwrap();
551        assert_eq!(merged_or.active_count(0), 32);
552    }
553}