1use crate::{CategoryMapping, HeadCategory, Result, SparseError, MIN_ACTIVE_HEADS};
4use serde::{Deserialize, Serialize};
5use std::collections::HashMap;
6
7#[derive(Debug, Clone, Serialize, Deserialize)]
9pub struct AttentionMask {
10 pub num_heads: usize,
12 pub num_layers: usize,
14 pub mask: Vec<Vec<bool>>,
17 pub layer_sparsity: Vec<f32>,
19 pub overall_sparsity: f32,
21}
22
23impl AttentionMask {
24 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 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 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 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 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 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 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 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 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 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 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 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 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 pub fn to_bytes(&self) -> Vec<u8> {
228 let mut bytes = Vec::new();
229
230 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 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 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 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#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize)]
307pub enum MaskPattern {
308 Strided { stride: usize },
310 SkipFirst { count: usize },
312 SkipLast { count: usize },
314 Threshold,
316 CategoryBased,
318}
319
320#[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 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 pub fn sparsity(mut self, sparsity: f32) -> Self {
348 self.target_sparsity = sparsity.clamp(0.0, 0.9);
349 self
350 }
351
352 pub fn pattern(mut self, pattern: MaskPattern) -> Self {
354 self.pattern = pattern;
355 self
356 }
357
358 pub fn min_active(mut self, min: usize) -> Self {
360 self.min_active = min;
361 self
362 }
363
364 pub fn category_weights(mut self, weights: HashMap<HeadCategory, f32>) -> Self {
366 self.category_weights = Some(weights);
367 self
368 }
369
370 pub fn category_mapping(mut self, mapping: CategoryMapping) -> Self {
372 self.category_mapping = Some(mapping);
373 self
374 }
375
376 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 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 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 head_scores.sort_by(|a, b| a.1.partial_cmp(&b.1).unwrap());
447
448 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 for layer in 0..10 {
496 assert!(!mask.is_active(layer, 0)); assert!(mask.is_active(layer, 1)); assert!(!mask.is_active(layer, 2)); }
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 let merged_and = mask1.merge_and(&mask2).unwrap();
543 assert_eq!(merged_and.active_count(0), 16);
544
545 let merged_or = mask1.merge_or(&mask2).unwrap();
551 assert_eq!(merged_or.active_count(0), 32);
552 }
553}