///|
fn strategy_name(strategy : NegativeStrategy) -> String {
match strategy {
Tail(size) => "tail-\{size}"
HardWindow(size) => "hard-\{size}"
Stride(step) => "stride-\{step}"
}
}
///|
fn pool_map(pools : Array[CandidatePool]) -> Map[String, Array[String]] {
let mapped : Map[String, Array[String]] = Map([])
for pool in pools {
mapped[pool.query_id] = pool.doc_ids.copy()
}
mapped
}
///|
fn sampled_doc_ids(
doc_ids : Array[String],
strategy : NegativeStrategy,
limit : Int,
) -> Array[(String, Int)] {
let picked : Array[(String, Int)] = []
if limit <= 0 {
return picked
}
match strategy {
Tail(window) => {
let start = Int::max(0, doc_ids.length() - Int::max(1, window))
for index in start..= limit {
break
}
picked.push((doc_ids[index], index + 1))
}
}
HardWindow(window) => {
let top = Int::min(doc_ids.length(), Int::max(1, window))
for index in 0..= limit {
break
}
picked.push((doc_ids[index], index + 1))
}
}
Stride(step) => {
let stride = Int::max(step, 1)
for index in 0..= limit {
break
}
if index % stride == 0 {
picked.push((doc_ids[index], index + 1))
}
}
}
}
picked
}
///|
pub fn sample_negatives(
qrels : Array[JudgedDoc],
run : Array[RetrievedDoc],
pools? : Array[CandidatePool] = [],
config? : NegativeSampleConfig = NegativeSampleConfig::default(),
) -> Array[NegativeSample] {
let qrels_by_query = group_qrels(qrels)
let run_pools = build_candidate_pools(run)
let external_pools = pool_map(pools)
let fallback_pools = pool_map(run_pools)
let query_ids = sorted_query_ids(qrels_by_query, group_runs(run))
let samples : Array[NegativeSample] = []
for query_id in query_ids {
let judgments = qrels_by_query.get_or_default(query_id, [])
let blocked : Map[String, Unit] = Map([])
for item in judgments {
if item.relevance >= config.relevant_threshold || config.skip_judged {
blocked[item.doc_id] = ()
}
}
let doc_ids = match external_pools.get(query_id) {
Some(ids) => ids
None => fallback_pools.get_or_default(query_id, [])
}
let selected = sampled_doc_ids(
doc_ids,
config.strategy,
config.per_query * 3,
)
let seen : Map[String, Unit] = Map([])
for pair in selected {
let (doc_id, rank) = pair
if blocked.contains(doc_id) || seen.contains(doc_id) {
continue
}
seen[doc_id] = ()
samples.push({
query_id,
doc_id,
source_rank: rank,
strategy: strategy_name(config.strategy),
})
if seen.length() >= config.per_query {
break
}
}
}
samples.sort_by(fn(a, b) {
let by_query = a.query_id.compare(b.query_id)
if by_query == 0 {
a.source_rank - b.source_rank
} else {
by_query
}
})
samples
}
///|
pub fn render_negative_samples_tsv(samples : Array[NegativeSample]) -> String {
let rows : Array[String] = ["query_id\tdoc_id\tsource_rank\tstrategy"]
for sample in samples {
rows.push(
"\{sample.query_id}\t\{sample.doc_id}\t\{sample.source_rank}\t\{sample.strategy}",
)
}
rows.join("\n")
}