///|
pub struct SearchResult {
id : String
score : Double
metadata : Array[(String, String)]
}
///|
pub impl Show for SearchResult with fn output(self, logger) {
logger.write_string(
"SearchResult{id: " + self.id + ", score: " + self.score.to_string() + "}",
)
}
///|
fn sort_results(results : Array[SearchResult], ascending : Bool) -> Unit {
let n = results.length()
for i = 0; i < n; i = i + 1 {
for j = i + 1; j < n; j = j + 1 {
let swap = if ascending {
results[i].score > results[j].score
} else {
results[i].score < results[j].score
}
if swap {
let temp = results[i]
results[i] = results[j]
results[j] = temp
}
}
}
}
///|
fn matches_filters(
metadata : Array[(String, String)],
filters : Array[(String, String)],
) -> Bool {
if filters.length() == 0 {
return true
}
for f in filters {
let mut found = false
for m in metadata {
if m.0 == f.0 && m.1 == f.1 {
found = true
break
}
}
if !found {
return false
}
}
true
}
// Flat baseline index
///|
pub struct FlatIndex {
documents : Array[Document]
}
///|
pub fn FlatIndex::new() -> FlatIndex {
{ documents: [] }
}
///|
pub fn FlatIndex::add(self : FlatIndex, doc : Document) -> Unit {
self.documents.push(doc)
}
///|
pub fn FlatIndex::search(
self : FlatIndex,
query : Array[Double],
top_k : Int,
metric : DistanceMetric,
filters : Array[(String, String)],
) -> Array[SearchResult] raise VectorError {
let results = []
for doc in self.documents {
if matches_filters(doc.metadata, filters) {
let score = calculate_distance(query, doc.vector, metric)
results.push({ id: doc.id, score, metadata: doc.metadata })
}
}
let ascending = match metric {
Euclidean | Manhattan => true
Cosine | DotProduct => false
}
sort_results(results, ascending)
let limit = if top_k < results.length() { top_k } else { results.length() }
let top = []
for i = 0; i < limit; i = i + 1 {
top.push(results[i])
}
top
}
// IVF-Flat Index (Inverted File)
///|
pub struct IvfIndex {
mut centroids : Array[Array[Double]]
inverted_lists : Array[Array[Document]]
k : Int
metric : DistanceMetric
}
///|
pub fn IvfIndex::new(k : Int, metric : DistanceMetric) -> IvfIndex {
let lists = Array::make(k, [])
for i = 0; i < k; i = i + 1 {
lists[i] = []
}
{ centroids: [], inverted_lists: lists, k, metric }
}
///|
pub fn IvfIndex::build(
self : IvfIndex,
docs : Array[Document],
) -> Unit raise VectorError {
if docs.length() == 0 {
return
}
let vectors = []
for d in docs {
vectors.push(d.vector)
}
self.centroids = kmeans(vectors, self.k, 20)
for i = 0; i < self.k; i = i + 1 {
self.inverted_lists[i].clear()
}
for doc in docs {
let mut min_dist = -1.0
let mut best_centroid = 0
for c = 0; c < self.k; c = c + 1 {
let dist = euclidean_distance(doc.vector, self.centroids[c])
if min_dist < 0.0 || dist < min_dist {
min_dist = dist
best_centroid = c
}
}
self.inverted_lists[best_centroid].push(doc)
}
}
///|
pub fn IvfIndex::search(
self : IvfIndex,
query : Array[Double],
top_k : Int,
nprobe : Int,
filters : Array[(String, String)],
) -> Array[SearchResult] raise VectorError {
if self.centroids.length() == 0 {
raise IndexError("IVF index is empty or not built")
}
let np = if nprobe > self.k { self.k } else { nprobe }
let centroid_dists = []
for c = 0; c < self.k; c = c + 1 {
let dist = euclidean_distance(query, self.centroids[c])
centroid_dists.push((c, dist))
}
let n_centroids = centroid_dists.length()
for i = 0; i < n_centroids; i = i + 1 {
for j = i + 1; j < n_centroids; j = j + 1 {
if centroid_dists[i].1 > centroid_dists[j].1 {
let temp = centroid_dists[i]
centroid_dists[i] = centroid_dists[j]
centroid_dists[j] = temp
}
}
}
let results = []
for i = 0; i < np; i = i + 1 {
let centroid_idx = centroid_dists[i].0
let list = self.inverted_lists[centroid_idx]
for doc in list {
if matches_filters(doc.metadata, filters) {
let score = calculate_distance(query, doc.vector, self.metric)
results.push({ id: doc.id, score, metadata: doc.metadata })
}
}
}
let ascending = match self.metric {
Euclidean | Manhattan => true
Cosine | DotProduct => false
}
sort_results(results, ascending)
let limit = if top_k < results.length() { top_k } else { results.length() }
let top = []
for i = 0; i < limit; i = i + 1 {
top.push(results[i])
}
top
}
// KD-Tree Index
///|
pub enum KdNode {
Empty
Node(Int, Document, KdNode, KdNode)
}
///|
pub struct KdTreeIndex {
mut root : KdNode
metric : DistanceMetric
}
///|
pub fn KdTreeIndex::new(metric : DistanceMetric) -> KdTreeIndex {
{ root: Empty, metric }
}
///|
fn build_kd_node(docs : Array[Document], depth : Int, dim : Int) -> KdNode {
if docs.length() == 0 {
return Empty
}
let axis = depth % dim
let n = docs.length()
for i = 0; i < n; i = i + 1 {
for j = i + 1; j < n; j = j + 1 {
if docs[i].vector[axis] > docs[j].vector[axis] {
let temp = docs[i]
docs[i] = docs[j]
docs[j] = temp
}
}
}
let median = n / 2
let doc = docs[median]
let left_docs = []
for i = 0; i < median; i = i + 1 {
left_docs.push(docs[i])
}
let right_docs = []
for i = median + 1; i < n; i = i + 1 {
right_docs.push(docs[i])
}
let left = build_kd_node(left_docs, depth + 1, dim)
let right = build_kd_node(right_docs, depth + 1, dim)
Node(axis, doc, left, right)
}
///|
pub fn KdTreeIndex::build(self : KdTreeIndex, docs : Array[Document]) -> Unit {
if docs.length() == 0 {
self.root = Empty
return
}
let dim = docs[0].vector.length()
self.root = build_kd_node(docs, 0, dim)
}
///|
fn add_to_best(
best : Array[SearchResult],
doc : Document,
score : Double,
top_k : Int,
metric : DistanceMetric,
) -> Unit {
let is_better = if best.length() < top_k {
true
} else {
let worst_idx = best.length() - 1
let worst_score = best[worst_idx].score
match metric {
Euclidean | Manhattan => score < worst_score
Cosine | DotProduct => score > worst_score
}
}
if is_better {
let result = { id: doc.id, score, metadata: doc.metadata }
if best.length() < top_k {
best.push(result)
} else {
let worst_idx = best.length() - 1
best[worst_idx] = result
}
let ascending = match metric {
Euclidean | Manhattan => true
Cosine | DotProduct => false
}
sort_results(best, ascending)
}
}
///|
fn search_kd_node(
node : KdNode,
query : Array[Double],
top_k : Int,
metric : DistanceMetric,
filters : Array[(String, String)],
best : Array[SearchResult],
) -> Unit raise VectorError {
match node {
Empty => ()
Node(axis, doc, left, right) => {
let dist = calculate_distance(query, doc.vector, metric)
if matches_filters(doc.metadata, filters) {
add_to_best(best, doc, dist, top_k, metric)
}
let query_val = query[axis]
let node_val = doc.vector[axis]
let (closer_branch, further_branch) = if query_val < node_val {
(left, right)
} else {
(right, left)
}
search_kd_node(closer_branch, query, top_k, metric, filters, best)
let plane_dist = if query_val < node_val {
node_val - query_val
} else {
query_val - node_val
}
let search_further = if best.length() < top_k {
true
} else {
let worst_idx = best.length() - 1
let worst_score = best[worst_idx].score
match metric {
Euclidean | Manhattan => plane_dist < worst_score
Cosine | DotProduct => true
}
}
if search_further {
search_kd_node(further_branch, query, top_k, metric, filters, best)
}
}
}
}
///|
pub fn KdTreeIndex::search(
self : KdTreeIndex,
query : Array[Double],
top_k : Int,
filters : Array[(String, String)],
) -> Array[SearchResult] raise VectorError {
if top_k <= 0 {
return []
}
let best = []
search_kd_node(self.root, query, top_k, self.metric, filters, best)
let ascending = match self.metric {
Euclidean | Manhattan => true
Cosine | DotProduct => false
}
sort_results(best, ascending)
best
}