///|
/// Return whether larger scores are better for a metric.
pub fn higher_score_is_better(metric : DistanceMetric) -> Bool {
match metric {
Cosine | DotProduct => true
Euclidean | Manhattan => false
}
}
///|
/// Compare two scores according to their metric.
pub fn score_is_better(
candidate : Double,
current : Double,
metric : DistanceMetric,
) -> Bool {
if higher_score_is_better(metric) {
candidate > current
} else {
candidate < current
}
}
///|
/// Return ids for a relevance judgement set.
pub fn relevant_ids(results : Array[SearchResult]) -> Map[String, Bool] {
let ids = Map([])
for result in results {
ids.set(result.id, true)
}
ids
}
///|
/// Compute precision@k for an approximate ranking against a relevance set.
pub fn precision_at_k(
results : Array[SearchResult],
relevant : Array[SearchResult],
k : Int,
) -> Double {
if k <= 0 {
return 0.0
}
let ids = relevant_ids(relevant)
let limit = if k < results.length() { k } else { results.length() }
if limit == 0 {
return 0.0
}
let mut hits = 0
for i = 0; i < limit; i = i + 1 {
if ids.contains(results[i].id) {
hits = hits + 1
}
}
hits.to_double() / limit.to_double()
}
///|
/// Compute average precision for a ranked result list.
pub fn average_precision(
results : Array[SearchResult],
relevant : Array[SearchResult],
) -> Double {
if relevant.length() == 0 {
return 0.0
}
let ids = relevant_ids(relevant)
let mut hits = 0
let mut total = 0.0
for i = 0; i < results.length(); i = i + 1 {
if ids.contains(results[i].id) {
hits = hits + 1
total = total + hits.to_double() / (i + 1).to_double()
}
}
total / relevant.length().to_double()
}
///|
/// Compute mean average precision for corresponding query batches.
pub fn mean_average_precision(
approximate : Array[Array[SearchResult]],
exact : Array[Array[SearchResult]],
) -> Double {
let count = if approximate.length() < exact.length() {
approximate.length()
} else {
exact.length()
}
if count == 0 {
return 0.0
}
let mut total = 0.0
for i = 0; i < count; i = i + 1 {
total = total + average_precision(approximate[i], exact[i])
}
total / count.to_double()
}
///|
/// Compute discounted cumulative gain at k.
pub fn dcg_at_k(
results : Array[SearchResult],
relevant : Array[SearchResult],
k : Int,
) -> Double {
if k <= 0 {
return 0.0
}
let ids = relevant_ids(relevant)
let limit = if k < results.length() { k } else { results.length() }
let mut total = 0.0
for i = 0; i < limit; i = i + 1 {
if ids.contains(results[i].id) {
total = total + dcg_discount(i + 1)
}
}
total
}
///|
fn dcg_discount(rank : Int) -> Double {
match rank {
1 => 1.0
2 => 0.6309297535714574
3 => 0.5
4 => 0.43067655807339306
5 => 0.38685280723454163
6 => 0.3562071871080222
7 => 0.3333333333333333
8 => 0.3154648767857287
9 => 0.3010299956639812
10 => 0.2890648263178878
_ => 1.0 / (rank + 1).to_double().sqrt()
}
}
///|
/// Compute normalized discounted cumulative gain at k.
pub fn ndcg_at_k(
results : Array[SearchResult],
relevant : Array[SearchResult],
k : Int,
) -> Double {
if relevant.length() == 0 || k <= 0 {
return 0.0
}
let actual = dcg_at_k(results, relevant, k)
let ideal = dcg_at_k(relevant, relevant, k)
if ideal == 0.0 {
0.0
} else {
actual / ideal
}
}
///|
/// Count exact matches at every requested cutoff.
pub fn recall_curve(
approximate : Array[SearchResult],
exact : Array[SearchResult],
cutoffs : Array[Int],
) -> Array[(Int, Double)] {
let curve = []
for cutoff in cutoffs {
curve.push((cutoff, recall_at_k(approximate, exact, cutoff)))
}
curve
}
///|
/// Return all result scores in ranking order for logging or charting.
pub fn result_scores(results : Array[SearchResult]) -> Array[Double] {
let scores = []
for result in results {
scores.push(result.score)
}
scores
}
///|
/// Compute the largest score gap between adjacent results.
pub fn largest_score_gap(results : Array[SearchResult]) -> Double {
if results.length() < 2 {
return 0.0
}
let mut largest = 0.0
for i = 1; i < results.length(); i = i + 1 {
let gap = (results[i - 1].score - results[i].score).abs()
if gap > largest {
largest = gap
}
}
largest
}
///|
/// Count distinct relevant documents retrieved by an approximate batch.
pub fn retrieved_relevant_count(
approximate : Array[SearchResult],
relevant : Array[SearchResult],
) -> Int {
let ids = relevant_ids(relevant)
let seen = Map([])
for result in approximate {
if ids.contains(result.id) {
seen.set(result.id, true)
}
}
seen.length()
}
///|
/// Measure how much of the relevance set appears anywhere in a result list.
pub fn coverage(
results : Array[SearchResult],
relevant : Array[SearchResult],
) -> Double {
if relevant.length() == 0 {
return 0.0
}
retrieved_relevant_count(results, relevant).to_double() /
relevant.length().to_double()
}
///|
/// Return an aggregate coverage score for corresponding query batches.
pub fn mean_coverage(
approximate : Array[Array[SearchResult]],
exact : Array[Array[SearchResult]],
) -> Double {
let count = if approximate.length() < exact.length() {
approximate.length()
} else {
exact.length()
}
if count == 0 {
return 0.0
}
let mut total = 0.0
for i = 0; i < count; i = i + 1 {
total = total + coverage(approximate[i], exact[i])
}
total / count.to_double()
}