Skip to main content

lawkit_core/generate/
zipf.rs

1use super::{DataGenerator, GenerateConfig};
2use crate::error::Result;
3use rand::prelude::*;
4
5#[derive(Debug, Clone)]
6pub struct ZipfGenerator {
7    pub exponent: f64,
8    pub vocabulary_size: usize,
9}
10
11impl ZipfGenerator {
12    pub fn new(exponent: f64, vocabulary_size: usize) -> Self {
13        Self {
14            exponent,
15            vocabulary_size,
16        }
17    }
18}
19
20impl DataGenerator for ZipfGenerator {
21    type Output = Vec<usize>;
22
23    fn generate(&self, config: &GenerateConfig) -> Result<Self::Output> {
24        let mut rng = config.create_rng();
25        let mut numbers = Vec::with_capacity(config.samples);
26
27        // Pre-calculate probabilities for each rank
28        let mut probabilities = Vec::with_capacity(self.vocabulary_size);
29        let mut total_weight = 0.0;
30
31        for rank in 1..=self.vocabulary_size {
32            let weight = 1.0 / (rank as f64).powf(self.exponent);
33            probabilities.push(weight);
34            total_weight += weight;
35        }
36
37        // Normalize probabilities
38        for prob in &mut probabilities {
39            *prob /= total_weight;
40        }
41
42        // Generate samples using inverse transform sampling
43        for _ in 0..config.samples {
44            let u: f64 = rng.gen();
45            let mut cumulative = 0.0;
46
47            for (rank, &prob) in probabilities.iter().enumerate() {
48                cumulative += prob;
49                if u <= cumulative {
50                    numbers.push(rank + 1); // ranks are 1-indexed
51                    break;
52                }
53            }
54        }
55
56        // Inject fraud if specified (flatten the distribution)
57        if config.fraud_rate > 0.0 {
58            inject_zipf_fraud(
59                &mut numbers,
60                config.fraud_rate,
61                self.vocabulary_size,
62                &mut rng,
63            );
64        }
65
66        Ok(numbers)
67    }
68}
69
70fn inject_zipf_fraud(
71    numbers: &mut [usize],
72    fraud_rate: f64,
73    vocab_size: usize,
74    rng: &mut impl Rng,
75) {
76    let fraud_count = (numbers.len() as f64 * fraud_rate) as usize;
77
78    // Fraud: inject more uniform distribution (less Zipf-like)
79    for _ in 0..fraud_count {
80        let index = rng.gen_range(0..numbers.len());
81        // Replace with a more uniformly distributed rank
82        numbers[index] = rng.gen_range(1..=vocab_size);
83    }
84}
85
86#[cfg(test)]
87mod tests {
88    use super::*;
89    use std::collections::HashMap;
90
91    #[test]
92    fn test_zipf_generator() {
93        let generator = ZipfGenerator::new(1.0, 1000);
94        let config = GenerateConfig::new(10000).with_seed(42);
95
96        let result = generator.generate(&config).unwrap();
97        assert_eq!(result.len(), 10000);
98
99        // Count frequencies
100        let mut frequencies = HashMap::new();
101        for &rank in &result {
102            *frequencies.entry(rank).or_insert(0) += 1;
103        }
104
105        // Check that rank 1 appears most frequently
106        let rank1_freq = frequencies.get(&1).unwrap_or(&0);
107        let rank2_freq = frequencies.get(&2).unwrap_or(&0);
108
109        assert!(rank1_freq > rank2_freq);
110
111        // All ranks should be within vocabulary size
112        for &rank in &result {
113            assert!((1..=1000).contains(&rank));
114        }
115    }
116}