Skip to main content

rill_ml/
feature_hasher.rs

1//! Feature hashing for dimensionality reduction.
2//!
3//! Maps high-dimensional sparse features into a fixed-dimensional dense
4//! vector using a deterministic hash function. Supports signed hashing
5//! to reduce collision bias.
6//!
7//! # Examples
8//!
9//! ```
10//! use rill_ml::feature_hasher::FeatureHasher;
11//! use rill_ml::sparse::SparseFeatures;
12//!
13//! let hasher = FeatureHasher::new(8, 42).unwrap();
14//! let sf = SparseFeatures::from_sorted(vec![(1, 3.0), (5, -2.0)]).unwrap();
15//! let dense = hasher.transform(&sf).unwrap();
16//! assert_eq!(dense.len(), 8);
17//! ```
18
19use crate::error::{RillError, checked_finite_add, ensure_finite};
20use crate::sparse::{FeatureId, SparseFeatures};
21use std::collections::hash_map::DefaultHasher;
22use std::hash::{Hash, Hasher};
23
24/// Configuration for [`FeatureHasher`].
25#[derive(Debug, Clone)]
26#[cfg_attr(feature = "serde", derive(serde::Serialize, serde::Deserialize))]
27#[non_exhaustive]
28pub struct FeatureHasherConfig {
29    /// Output dimension. Must be > 0.
30    pub dimension: usize,
31    /// Random seed for reproducible hashing.
32    pub seed: u64,
33    /// Whether to use signed hashing (alternate sign based on hash bit).
34    pub signed: bool,
35}
36
37impl Default for FeatureHasherConfig {
38    fn default() -> Self {
39        Self {
40            dimension: 1024,
41            seed: 0,
42            signed: true,
43        }
44    }
45}
46
47/// Fixed-dimension feature hasher.
48///
49/// Uses two independent hash functions:
50/// - The first determines the target bucket (`hash1 % dimension`).
51/// - The second determines the sign (`hash2 & 1`) when `signed = true`.
52///
53/// The hash is deterministic given the same `seed`, ensuring reproducible
54/// output across runs.
55#[derive(Debug, Clone)]
56#[cfg_attr(feature = "serde", derive(serde::Serialize))]
57pub struct FeatureHasher {
58    config: FeatureHasherConfig,
59}
60
61impl FeatureHasher {
62    /// Create a new hasher with the given dimension and seed.
63    ///
64    /// Uses signed hashing by default.
65    pub fn new(dimension: usize, seed: u64) -> Result<Self, RillError> {
66        Self::with_config(FeatureHasherConfig {
67            dimension,
68            seed,
69            signed: true,
70        })
71    }
72
73    /// Create a new hasher with a custom configuration.
74    pub fn with_config(config: FeatureHasherConfig) -> Result<Self, RillError> {
75        if config.dimension == 0 {
76            return Err(RillError::InvalidHashDimension(config.dimension));
77        }
78        Ok(Self { config })
79    }
80
81    /// The output dimension.
82    pub const fn dimension(&self) -> usize {
83        self.config.dimension
84    }
85
86    /// The random seed.
87    pub const fn seed(&self) -> u64 {
88        self.config.seed
89    }
90
91    /// Whether signed hashing is enabled.
92    pub const fn signed(&self) -> bool {
93        self.config.signed
94    }
95
96    /// Hash a `FeatureId` to a `(bucket, sign)` pair.
97    fn hash_id(&self, id: FeatureId) -> (usize, f64) {
98        let bucket = self.hash_bucket(id);
99        let sign = if self.config.signed {
100            self.hash_sign(id)
101        } else {
102            1.0
103        };
104        (bucket, sign)
105    }
106
107    /// Compute the bucket index for a feature id.
108    fn hash_bucket(&self, id: FeatureId) -> usize {
109        let mut hasher = DefaultHasher::new();
110        self.config.seed.hash(&mut hasher);
111        id.hash(&mut hasher);
112        (hasher.finish() as usize) % self.config.dimension
113    }
114
115    /// Compute the sign for a feature id (signed hashing).
116    fn hash_sign(&self, id: FeatureId) -> f64 {
117        let mut hasher = DefaultHasher::new();
118        (self.config.seed.wrapping_mul(0x517cc1b727220a95)).hash(&mut hasher);
119        id.hash(&mut hasher);
120        if hasher.finish() & 1 == 1 { -1.0 } else { 1.0 }
121    }
122
123    /// Hash a string feature name to a `FeatureId`.
124    pub fn hash_string(&self, name: &str) -> FeatureId {
125        let mut hasher = DefaultHasher::new();
126        self.config.seed.hash(&mut hasher);
127        name.hash(&mut hasher);
128        hasher.finish()
129    }
130
131    /// Create `SparseFeatures` from string name/value pairs.
132    ///
133    /// Each string is hashed to a `FeatureId`, then the pairs are sorted
134    /// and duplicates are merged by summing values.
135    pub fn hash_strings(&self, pairs: &[(&str, f64)]) -> Result<SparseFeatures, RillError> {
136        let mut ids: Vec<(FeatureId, f64)> = Vec::with_capacity(pairs.len());
137        for (name, value) in pairs {
138            ensure_finite("hash_value", *value)?;
139            ids.push((self.hash_string(name), *value));
140        }
141        SparseFeatures::from_unsorted(ids)
142    }
143
144    /// Transform `SparseFeatures` into a dense `Vec<f64>`.
145    ///
146    /// Each feature's value is added to its target bucket (multiplied by
147    /// the sign if signed hashing is enabled). Collisions cause values to
148    /// accumulate.
149    pub fn transform(&self, features: &SparseFeatures) -> Result<Vec<f64>, RillError> {
150        if self.config.dimension == 0 {
151            return Err(RillError::InvalidHashDimension(0));
152        }
153        features.validate()?;
154        let mut output = vec![0.0; self.config.dimension];
155        for &(id, value) in features.values() {
156            ensure_finite("sparse_value", value)?;
157            let (bucket, sign) = self.hash_id(id);
158            output[bucket] = checked_finite_add(output[bucket], sign * value, "hashed feature")?;
159        }
160        Ok(output)
161    }
162}
163
164#[cfg(feature = "serde")]
165impl<'de> serde::Deserialize<'de> for FeatureHasher {
166    fn deserialize<D>(deserializer: D) -> Result<Self, D::Error>
167    where
168        D: serde::Deserializer<'de>,
169    {
170        #[derive(serde::Deserialize)]
171        struct FeatureHasherState {
172            config: FeatureHasherConfig,
173        }
174
175        let state = FeatureHasherState::deserialize(deserializer)?;
176        Self::with_config(state.config).map_err(serde::de::Error::custom)
177    }
178}
179
180#[cfg(test)]
181mod tests {
182    use super::*;
183
184    #[test]
185    fn reproducible_output() {
186        let h = FeatureHasher::new(16, 42).unwrap();
187        let sf = SparseFeatures::from_sorted(vec![(1, 3.0), (5, -2.0), (10, 1.0)]).unwrap();
188        let out1 = h.transform(&sf).unwrap();
189        let out2 = h.transform(&sf).unwrap();
190        assert_eq!(out1, out2);
191    }
192
193    #[test]
194    fn collision_overflow_is_rejected() {
195        let hasher = FeatureHasher::with_config(FeatureHasherConfig {
196            dimension: 1,
197            seed: 42,
198            signed: false,
199        })
200        .unwrap();
201        let features = SparseFeatures::from_sorted(vec![(1, f64::MAX), (2, f64::MAX)]).unwrap();
202        assert!(hasher.transform(&features).is_err());
203    }
204
205    #[cfg(feature = "serde")]
206    #[test]
207    fn serde_rejects_zero_dimension() {
208        let malformed = r#"{"config":{"dimension":0,"seed":0,"signed":true}}"#;
209        assert!(serde_json::from_str::<FeatureHasher>(malformed).is_err());
210    }
211
212    #[test]
213    fn different_seeds_produce_different_output() {
214        let h1 = FeatureHasher::new(16, 1).unwrap();
215        let h2 = FeatureHasher::new(16, 2).unwrap();
216        let sf = SparseFeatures::from_sorted(vec![(1, 1.0), (2, 2.0)]).unwrap();
217        let out1 = h1.transform(&sf).unwrap();
218        let out2 = h2.transform(&sf).unwrap();
219        assert_ne!(out1, out2);
220    }
221
222    #[test]
223    fn signed_hashing_produces_negatives() {
224        let h = FeatureHasher::with_config(FeatureHasherConfig {
225            dimension: 256,
226            seed: 42,
227            signed: true,
228        })
229        .unwrap();
230        let sf =
231            SparseFeatures::from_sorted((0..100).map(|i| (i, 1.0)).collect::<Vec<_>>()).unwrap();
232        let out = h.transform(&sf).unwrap();
233        // With 100 features and signed hashing, at least some should be negative
234        assert!(out.iter().any(|&v| v < 0.0));
235    }
236
237    #[test]
238    fn unsigned_hashing_all_positive() {
239        let h = FeatureHasher::with_config(FeatureHasherConfig {
240            dimension: 256,
241            seed: 42,
242            signed: false,
243        })
244        .unwrap();
245        let sf =
246            SparseFeatures::from_sorted((0..100).map(|i| (i, 1.0)).collect::<Vec<_>>()).unwrap();
247        let out = h.transform(&sf).unwrap();
248        assert!(out.iter().all(|&v| v >= 0.0));
249    }
250
251    #[test]
252    fn dimension_one_all_same_bucket() {
253        let h = FeatureHasher::new(1, 42).unwrap();
254        let sf = SparseFeatures::from_sorted(vec![(1, 3.0), (2, 5.0)]).unwrap();
255        let out = h.transform(&sf).unwrap();
256        assert_eq!(out.len(), 1);
257        // Both values land in bucket 0, with signs
258        assert!(out[0].abs() > 0.0);
259    }
260
261    #[test]
262    fn empty_features_returns_zeros() {
263        let h = FeatureHasher::new(8, 42).unwrap();
264        let sf = SparseFeatures::new();
265        let out = h.transform(&sf).unwrap();
266        assert_eq!(out, vec![0.0; 8]);
267    }
268
269    #[test]
270    fn string_hash_reproducible() {
271        let h = FeatureHasher::new(8, 42).unwrap();
272        let id1 = h.hash_string("user_id");
273        let id2 = h.hash_string("user_id");
274        assert_eq!(id1, id2);
275    }
276
277    #[test]
278    fn different_strings_different_ids() {
279        let h = FeatureHasher::new(8, 42).unwrap();
280        let id1 = h.hash_string("user_id");
281        let id2 = h.hash_string("device_id");
282        assert_ne!(id1, id2);
283    }
284
285    #[test]
286    fn hash_strings_creates_sorted_features() {
287        let h = FeatureHasher::new(8, 42).unwrap();
288        let sf = h
289            .hash_strings(&[("alpha", 1.0), ("beta", 2.0), ("gamma", 3.0)])
290            .unwrap();
291        // Should be valid sorted sparse features
292        assert!(sf.validate().is_ok());
293        assert_eq!(sf.len(), 3);
294    }
295
296    #[test]
297    fn invalid_dimension_rejected() {
298        assert!(matches!(
299            FeatureHasher::new(0, 42),
300            Err(RillError::InvalidHashDimension(0))
301        ));
302    }
303
304    #[test]
305    #[cfg(feature = "serde")]
306    fn serde_roundtrip() {
307        let h = FeatureHasher::new(16, 42).unwrap();
308        let json = serde_json::to_string(&h).unwrap();
309        let restored: FeatureHasher = serde_json::from_str(&json).unwrap();
310        assert_eq!(restored.dimension(), 16);
311        assert_eq!(restored.seed(), 42);
312        assert!(restored.signed());
313    }
314}