lawkit_core/generate/
zipf.rs1use 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 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 for prob in &mut probabilities {
39 *prob /= total_weight;
40 }
41
42 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); break;
52 }
53 }
54 }
55
56 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 for _ in 0..fraud_count {
80 let index = rng.gen_range(0..numbers.len());
81 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 let mut frequencies = HashMap::new();
101 for &rank in &result {
102 *frequencies.entry(rank).or_insert(0) += 1;
103 }
104
105 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 for &rank in &result {
113 assert!((1..=1000).contains(&rank));
114 }
115 }
116}