///|
fn top_doc_set(
run : Array[RetrievedDoc],
query_id : String,
cutoff : Int,
) -> Map[String, Unit] {
let grouped = group_runs(run)
let docs : Map[String, Unit] = Map([])
let items = grouped.get_or_default(query_id, [])
let limit = Int::min(Int::max(cutoff, 0), items.length())
for index in 0.. Int {
let left = top_doc_set(baseline, query_id, cutoff)
let right = top_doc_set(candidate, query_id, cutoff)
let mut overlap = 0
for doc_id, _ in left {
if right.contains(doc_id) {
overlap += 1
}
}
overlap
}
///|
fn query_metric(
qrels : Array[JudgedDoc],
run : Array[RetrievedDoc],
query_id : String,
cutoff : Int,
) -> Double {
let qrels_by_query = group_qrels(qrels)
let run_by_query = group_runs(run)
let query_qrels = qrels_by_query.get_or_default(query_id, [])
let query_run = run_by_query.get_or_default(query_id, [])
let query = evaluate_query(
query_id,
query_qrels,
query_run,
config=EvalConfig::new(cutoffs=[cutoff]),
)
query.metrics.get_or_default(metric_name("ndcg", cutoff), 0.0)
}
///|
pub fn compare_runs(
qrels : Array[JudgedDoc],
baseline : Array[RetrievedDoc],
candidate : Array[RetrievedDoc],
cutoff~ : Int,
) -> RunComparison {
let qrels_by_query = group_qrels(qrels)
let baseline_by_query = group_runs(baseline)
let candidate_by_query = group_runs(candidate)
let query_ids = sorted_query_ids(qrels_by_query, baseline_by_query)
let all_query_ids = sorted_query_ids(qrels_by_query, candidate_by_query)
for query_id in all_query_ids {
if !query_ids.contains(query_id) {
query_ids.push(query_id)
}
}
query_ids.sort()
let comparisons : Array[QueryComparison] = []
let mut wins = 0
let mut losses = 0
let mut ties = 0
let mut total_delta = 0.0
for query_id in query_ids {
let baseline_value = query_metric(qrels, baseline, query_id, cutoff)
let candidate_value = query_metric(qrels, candidate, query_id, cutoff)
let delta = candidate_value - baseline_value
let outcome = if delta > 0.000000001 {
wins += 1
"win"
} else if delta < -0.000000001 {
losses += 1
"loss"
} else {
ties += 1
"tie"
}
total_delta += delta
comparisons.push({
query_id,
baseline: baseline_value,
candidate: candidate_value,
delta,
outcome,
overlap: overlap_at(baseline, candidate, query_id, cutoff),
})
}
let count = comparisons.length()
{
cutoff,
query_count: count,
wins,
losses,
ties,
mean_delta: if count == 0 {
0.0
} else {
total_delta / Double::from_int(count)
},
queries: comparisons,
}
}
///|
pub fn render_comparison_text(comparison : RunComparison) -> String {
let lines : Array[String] = []
lines.push("MoonRAGBench comparison")
lines.push(
"cutoff=\{comparison.cutoff} queries=\{comparison.query_count} wins=\{comparison.wins} losses=\{comparison.losses} ties=\{comparison.ties}",
)
lines.push("mean_delta=\{format_metric(comparison.mean_delta)}")
for item in comparison.queries {
lines.push(
"\{item.query_id}\toutcome=\{item.outcome}\tdelta=\{format_metric(item.delta)}\toverlap=\{item.overlap}",
)
}
lines.join("\n")
}
///|
pub fn render_comparison_json(comparison : RunComparison) -> String {
ToJson::to_json(comparison).stringify(indent=2)
}