1use 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#[derive(Debug, Clone)]
26#[cfg_attr(feature = "serde", derive(serde::Serialize, serde::Deserialize))]
27#[non_exhaustive]
28pub struct FeatureHasherConfig {
29 pub dimension: usize,
31 pub seed: u64,
33 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#[derive(Debug, Clone)]
56#[cfg_attr(feature = "serde", derive(serde::Serialize))]
57pub struct FeatureHasher {
58 config: FeatureHasherConfig,
59}
60
61impl FeatureHasher {
62 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 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 pub const fn dimension(&self) -> usize {
83 self.config.dimension
84 }
85
86 pub const fn seed(&self) -> u64 {
88 self.config.seed
89 }
90
91 pub const fn signed(&self) -> bool {
93 self.config.signed
94 }
95
96 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 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 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 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 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 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 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 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 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}