Skip to main content

uqa_graph/
embedding.rs

1//
2// Unified Query Algebra
3//
4// Copyright (c) 2023-2026 Cognica, Inc.
5//
6
7//! Structural per-vertex graph embeddings.
8//!
9//! For each vertex we assemble a feature vector from
10//!
11//! * out-degree, in-degree
12//! * out-edge label distribution (capped at `dims / 2`)
13//! * k-hop frontier counts (one per layer)
14//!
15//! padded or truncated to `dims` and L2-normalized. The result lives
16//! on the payload's `_embedding` field as a `Value::List(Float)` so
17//! downstream vector-similarity scorers can consume it directly.
18
19use std::collections::{BTreeMap, BTreeSet};
20
21use uqa_core::{DocId, EdgeId, Payload, PostingEntry, PostingList, Value, VertexId};
22
23use crate::posting_list::{GraphPayload, GraphPostingList};
24use crate::store::{GraphStore, GraphStoreError, GraphStoreResult};
25
26pub const MAX_GRAPH_EMBEDDING_DIMENSIONS: usize = 4_096;
27pub const MAX_GRAPH_EMBEDDING_LAYERS: u32 = 256;
28const MAX_EXACT_F64_INTEGER: u64 = 9_007_199_254_740_992;
29
30pub struct GraphEmbedding<'a> {
31    pub graph: &'a str,
32    pub dimensions: usize,
33    pub k_layers: u32,
34}
35
36impl<'a> GraphEmbedding<'a> {
37    pub fn new(graph: &'a str) -> Self {
38        Self {
39            graph,
40            dimensions: 32,
41            k_layers: 2,
42        }
43    }
44
45    pub fn dimensions(mut self, dims: usize) -> Self {
46        self.dimensions = dims;
47        self
48    }
49
50    pub fn k_layers(mut self, k: u32) -> Self {
51        self.k_layers = k;
52        self
53    }
54
55    pub fn execute<G: GraphStore>(&self, store: &G) -> GraphStoreResult<GraphPostingList> {
56        if self.dimensions == 0 || self.dimensions > MAX_GRAPH_EMBEDDING_DIMENSIONS {
57            return Err(GraphStoreError::InvalidQuery(format!(
58                "graph embedding dimensions must be in 1..={MAX_GRAPH_EMBEDDING_DIMENSIONS}, got {}",
59                self.dimensions
60            )));
61        }
62        if self.k_layers > MAX_GRAPH_EMBEDDING_LAYERS {
63            return Err(GraphStoreError::InvalidQuery(format!(
64                "graph embedding layer count {} exceeds limit {MAX_GRAPH_EMBEDDING_LAYERS}",
65                self.k_layers
66            )));
67        }
68        let vertex_ids = store.vertex_ids_in_graph(self.graph)?;
69        let mut vertices: Vec<VertexId> = Vec::new();
70        vertices
71            .try_reserve_exact(vertex_ids.len())
72            .map_err(|error| allocation_error("vertex id list", vertex_ids.len(), &error))?;
73        vertices.extend(vertex_ids.iter().copied());
74        if vertices.is_empty() {
75            return Ok(GraphPostingList::new());
76        }
77        for vid in &vertices {
78            if store.get_vertex(*vid).is_none() {
79                return Err(GraphStoreError::CorruptGraph(format!(
80                    "graph embedding graph {:?} references missing vertex {vid}",
81                    self.graph
82                )));
83            }
84        }
85
86        // Collect alphabet for label-distribution one-hot.
87        let mut all_labels: BTreeSet<String> = BTreeSet::new();
88        for vid in &vertices {
89            for eid in store.out_edge_ids(*vid, self.graph)? {
90                let edge = store.get_edge(eid).ok_or_else(|| {
91                    GraphStoreError::CorruptGraph(format!("missing embedding edge {eid}"))
92                })?;
93                all_labels.insert(edge.label.clone());
94            }
95        }
96        let label_to_idx: BTreeMap<String, usize> = all_labels
97            .iter()
98            .enumerate()
99            .map(|(i, label)| (label.clone(), i))
100            .collect();
101        let n_labels = all_labels.len();
102
103        let mut entries: Vec<PostingEntry> = Vec::new();
104        entries
105            .try_reserve_exact(vertices.len())
106            .map_err(|error| allocation_error("posting entries", vertices.len(), &error))?;
107        let mut graph_payloads: BTreeMap<DocId, GraphPayload> = BTreeMap::new();
108        for vid in &vertices {
109            let embedding =
110                self.compute_embedding(store, *vid, &vertex_ids, &label_to_idx, n_labels)?;
111            let mut fields: BTreeMap<String, Value> = BTreeMap::new();
112            let mut embedding_values = Vec::new();
113            embedding_values
114                .try_reserve_exact(embedding.len())
115                .map_err(|error| allocation_error("embedding payload", embedding.len(), &error))?;
116            embedding_values.extend(embedding.iter().map(|x| Value::Float(*x)));
117            fields.insert("_embedding".into(), Value::List(embedding_values));
118            let payload = Payload {
119                positions: Vec::new(),
120                score: 0.0,
121                fields,
122            };
123            entries.push(PostingEntry::new(*vid, payload));
124            graph_payloads.insert(
125                *vid,
126                GraphPayload {
127                    subgraph_vertices: vec![*vid],
128                    subgraph_edges: Vec::new(),
129                    graph_name: self.graph.to_string(),
130                    score_override: None,
131                },
132            );
133        }
134        GraphPostingList::try_from_parts(
135            PostingList::from_sorted_unchecked(entries),
136            graph_payloads,
137        )
138        .map_err(Into::into)
139    }
140
141    fn compute_embedding<G: GraphStore>(
142        &self,
143        store: &G,
144        vid: VertexId,
145        graph_vertices: &BTreeSet<VertexId>,
146        label_to_idx: &BTreeMap<String, usize>,
147        n_labels: usize,
148    ) -> GraphStoreResult<Vec<f64>> {
149        let out_edges: BTreeSet<EdgeId> = store.out_edge_ids(vid, self.graph)?;
150        let in_edges: BTreeSet<EdgeId> = store.in_edge_ids(vid, self.graph)?;
151        let out_degree = out_edges.len();
152        let in_degree = in_edges.len();
153
154        let label_dims = n_labels.min(self.dimensions / 2);
155        let mut label_dist = Vec::new();
156        label_dist
157            .try_reserve_exact(label_dims)
158            .map_err(|error| allocation_error("edge-label distribution", label_dims, &error))?;
159        label_dist.resize(label_dims, 0.0_f64);
160        for eid in &out_edges {
161            let edge = store.get_edge(*eid).ok_or_else(|| {
162                GraphStoreError::CorruptGraph(format!("missing embedding edge {eid}"))
163            })?;
164            if let Some(&idx) = label_to_idx.get(&edge.label) {
165                if idx < label_dims {
166                    label_dist[idx] += 1.0;
167                }
168            }
169        }
170        let total: f64 = label_dist.iter().sum();
171        if total > 0.0 {
172            for v in &mut label_dist {
173                *v /= total;
174            }
175        }
176
177        let layer_capacity = usize::try_from(self.k_layers).map_err(|_| {
178            GraphStoreError::InvalidQuery(format!(
179                "graph embedding layer count {} does not fit usize",
180                self.k_layers
181            ))
182        })?;
183        let mut hop_counts: Vec<f64> = Vec::new();
184        hop_counts
185            .try_reserve_exact(layer_capacity)
186            .map_err(|error| allocation_error("hop counts", layer_capacity, &error))?;
187        let mut visited: BTreeSet<VertexId> = BTreeSet::from([vid]);
188        let mut frontier: BTreeSet<VertexId> = BTreeSet::from([vid]);
189        for _ in 0..self.k_layers {
190            let mut next_frontier: BTreeSet<VertexId> = BTreeSet::new();
191            for v in &frontier {
192                for eid in store.out_edge_ids(*v, self.graph)? {
193                    let edge = store.get_edge(eid).ok_or_else(|| {
194                        GraphStoreError::CorruptGraph(format!("missing embedding edge {eid}"))
195                    })?;
196                    if !graph_vertices.contains(&edge.target_id) {
197                        return Err(GraphStoreError::CorruptGraph(format!(
198                            "embedding edge {eid} targets vertex {} outside graph {:?}",
199                            edge.target_id, self.graph
200                        )));
201                    }
202                    if !visited.contains(&edge.target_id) {
203                        next_frontier.insert(edge.target_id);
204                        visited.insert(edge.target_id);
205                    }
206                }
207            }
208            hop_counts.push(usize_to_f64_exact(
209                next_frontier.len(),
210                "graph embedding hop frontier",
211            )?);
212            frontier = next_frontier;
213        }
214
215        let generated_features = 2_usize
216            .checked_add(label_dims)
217            .and_then(|value| value.checked_add(layer_capacity))
218            .ok_or_else(|| {
219                GraphStoreError::InvalidQuery(
220                    "graph embedding feature count overflows usize".to_string(),
221                )
222            })?;
223        let raw_capacity = self.dimensions.max(generated_features);
224        let mut raw: Vec<f64> = Vec::new();
225        raw.try_reserve_exact(raw_capacity)
226            .map_err(|error| allocation_error("raw embedding", raw_capacity, &error))?;
227        raw.push(usize_to_f64_exact(
228            out_degree,
229            "graph embedding out-degree",
230        )?);
231        raw.push(usize_to_f64_exact(in_degree, "graph embedding in-degree")?);
232        raw.extend(label_dist);
233        raw.extend(hop_counts);
234        normalize_embedding(raw, self.dimensions, vid)
235    }
236}
237
238fn usize_to_f64_exact(value: usize, context: &str) -> GraphStoreResult<f64> {
239    if !u64::try_from(value).is_ok_and(|value| value <= MAX_EXACT_F64_INTEGER) {
240        return Err(GraphStoreError::InvalidQuery(format!(
241            "{context} {value} exceeds the exact f64 integer range"
242        )));
243    }
244    Ok(value as f64)
245}
246
247fn allocation_error(
248    context: &str,
249    elements: usize,
250    error: &std::collections::TryReserveError,
251) -> GraphStoreError {
252    GraphStoreError::InvalidQuery(format!(
253        "cannot allocate {elements} elements for graph embedding {context}: {error}"
254    ))
255}
256
257fn normalize_embedding(
258    mut raw: Vec<f64>,
259    dimensions: usize,
260    vertex_id: VertexId,
261) -> GraphStoreResult<Vec<f64>> {
262    if raw.len() < dimensions {
263        raw.resize(dimensions, 0.0);
264    } else {
265        raw.truncate(dimensions);
266    }
267    let norm = raw.iter().map(|value| value * value).sum::<f64>().sqrt();
268    if !norm.is_finite() {
269        return Err(GraphStoreError::InvalidQuery(format!(
270            "graph embedding norm is not finite for vertex {vertex_id}"
271        )));
272    }
273    if norm > 0.0 {
274        for value in &mut raw {
275            *value /= norm;
276        }
277    }
278    Ok(raw)
279}
280
281#[cfg(test)]
282mod tests {
283    use uqa_core::{Edge, Vertex};
284
285    use super::*;
286    use crate::memory_store::MemoryGraphStore;
287
288    #[test]
289    fn missing_edge_record_is_corruption_not_an_empty_label_bucket() {
290        let mut store = MemoryGraphStore::new();
291        store.create_graph("g");
292        store.add_vertex(Vertex::new(1, "n"), "g").unwrap();
293        store.add_vertex(Vertex::new(2, "n"), "g").unwrap();
294        store.add_edge(Edge::new(10, 1, 2, "edge"), "g").unwrap();
295        store.remove_edge_record_for_corruption_test(10);
296
297        let error = GraphEmbedding::new("g").execute(&store).unwrap_err();
298        assert!(matches!(error, GraphStoreError::CorruptGraph(_)));
299    }
300}