ailake_index/
serialize.rs1use ailake_core::{AilakeError, AilakeResult, VectorMetric};
3use serde::{Deserialize, Serialize};
4
5use crate::hnsw::{HnswConfig, HnswIndex};
6
7#[derive(Serialize, Deserialize)]
8struct HnswSnapshot {
9 m: usize,
10 ef_construction: usize,
11 max_elements: usize,
12 metric: u8,
13 dim: u32,
14 row_ids: Vec<u64>,
16 flat_vecs: Vec<f32>,
18 neighbors: Vec<Vec<Vec<usize>>>,
20 node_levels: Vec<usize>,
21 entry_point: Option<usize>,
22 max_layer: usize,
23}
24
25fn metric_to_u8(m: VectorMetric) -> u8 {
26 match m {
27 VectorMetric::Cosine => 0,
28 VectorMetric::Euclidean => 1,
29 VectorMetric::DotProduct => 2,
30 VectorMetric::NormalizedCosine => 3,
31 }
32}
33
34fn u8_to_metric(v: u8) -> AilakeResult<VectorMetric> {
35 match v {
36 0 => Ok(VectorMetric::Cosine),
37 1 => Ok(VectorMetric::Euclidean),
38 2 => Ok(VectorMetric::DotProduct),
39 3 => Ok(VectorMetric::NormalizedCosine),
40 _ => Err(AilakeError::Catalog(format!(
41 "HNSW index deserialization: unknown metric byte {v} (valid: 0=Cosine, 1=Euclidean, 2=DotProduct, 3=NormalizedCosine)"
42 ))),
43 }
44}
45
46pub struct HnswSerializer;
47
48impl HnswSerializer {
49 pub fn to_bytes(index: &HnswIndex) -> AilakeResult<Vec<u8>> {
50 let snap = HnswSnapshot {
51 m: index.config.m,
52 ef_construction: index.config.ef_construction,
53 max_elements: index.config.max_elements,
54 metric: metric_to_u8(index.metric),
55 dim: index.dim,
56 row_ids: index.row_ids.clone(),
57 flat_vecs: index.flat_vecs.clone(),
58 neighbors: index.neighbors.clone(),
59 node_levels: index.node_levels.clone(),
60 entry_point: index.entry_point,
61 max_layer: index.max_layer,
62 };
63 bincode::serialize(&snap).map_err(|e| AilakeError::Bincode(e.to_string()))
64 }
65
66 pub fn from_bytes(bytes: &[u8]) -> AilakeResult<HnswIndex> {
67 let snap: HnswSnapshot =
68 bincode::deserialize(bytes).map_err(|e| AilakeError::Bincode(e.to_string()))?;
69 let metric = u8_to_metric(snap.metric)?;
70 Ok(HnswIndex {
71 config: HnswConfig {
72 m: snap.m,
73 ef_construction: snap.ef_construction,
74 max_elements: snap.max_elements,
75 },
76 metric,
77 dim: snap.dim,
78 row_ids: snap.row_ids,
79 flat_vecs: snap.flat_vecs,
80 flat_vecs_f16: None, neighbors: snap.neighbors,
82 node_levels: snap.node_levels,
83 entry_point: snap.entry_point,
84 max_layer: snap.max_layer,
85 })
86 }
87}
88
89#[cfg(test)]
90mod tests {
91 use super::*;
92 use crate::hnsw::HnswBuilder;
93 use ailake_core::RowId;
94
95 #[test]
96 fn serialize_roundtrip() {
97 let mut b = HnswBuilder::new(3, VectorMetric::Cosine, Default::default());
98 b.insert(RowId::new(0), vec![1.0, 0.0, 0.0]);
99 b.insert(RowId::new(1), vec![0.0, 1.0, 0.0]);
100 let idx = b.build();
101 let bytes = HnswSerializer::to_bytes(&idx).unwrap();
102 let idx2 = HnswSerializer::from_bytes(&bytes).unwrap();
103 assert_eq!(idx2.node_count(), 2);
104 assert_eq!(idx2.dim(), 3);
105 let r = idx2.search(&[1.0, 0.0, 0.0], 1, 50);
106 assert_eq!(r[0].0, RowId::new(0));
107 }
108
109 #[test]
110 fn serialize_preserves_graph() {
111 use rand::{rngs::StdRng, Rng, SeedableRng};
112 let mut rng = StdRng::seed_from_u64(7);
113 let mut b = HnswBuilder::new(8, VectorMetric::Euclidean, Default::default());
114 for i in 0..50u64 {
115 let v: Vec<f32> = (0..8).map(|_| rng.gen::<f32>()).collect();
116 b.insert(RowId::new(i), v);
117 }
118 let idx = b.build();
119 let query: Vec<f32> = (0..8).map(|_| rng.gen::<f32>()).collect();
120 let r1 = idx.search(&query, 5, 50);
121
122 let bytes = HnswSerializer::to_bytes(&idx).unwrap();
123 let idx2 = HnswSerializer::from_bytes(&bytes).unwrap();
124 let r2 = idx2.search(&query, 5, 50);
125
126 assert_eq!(r1.len(), r2.len());
127 for (a, b) in r1.iter().zip(r2.iter()) {
128 assert_eq!(a.0, b.0);
129 }
130 }
131}