Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
Original file line number Diff line number Diff line change
Expand Up @@ -190,6 +190,8 @@ def load_positive_docs_with_modality(
Tuple of ``(positive_docs_df, doc_to_modality_final)``.
"""
qrels_df = pd.read_csv(test_tsv_path, sep="\t")
if "corpus-id" not in qrels_df.columns and "corpus-id " in qrels_df.columns:
qrels_df = qrels_df.rename(columns={"corpus-id ": "corpus-id"})

with open(split_json_path, encoding="utf-8") as f:
splits = json.load(f)
Expand All @@ -202,7 +204,7 @@ def load_positive_docs_with_modality(
doc_to_modality: dict[str, set[str]] = defaultdict(set)
for _, row in qrels_df.iterrows():
query_id = row["query-id"]
corpus_id = row["corpus-id "] # trailing space in column name
corpus_id = row["corpus-id"]
Comment thread
oliverholworthy marked this conversation as resolved.
if query_id in query_to_modality:
doc_to_modality[corpus_id].add(query_to_modality[query_id])

Expand All @@ -212,7 +214,7 @@ def load_positive_docs_with_modality(
doc_to_modality_final[doc_id] = next(iter(modalities))
else:
modality_counts: dict[str, int] = defaultdict(int)
for _, r in qrels_df[qrels_df["corpus-id "] == doc_id].iterrows():
for _, r in qrels_df[qrels_df["corpus-id"] == doc_id].iterrows():
qid = r["query-id"]
if qid in query_to_modality:
modality_counts[query_to_modality[qid]] += 1
Expand Down
52 changes: 51 additions & 1 deletion plugins/data-designer-retrieval-sdg/tests/test_postprocess.py
Original file line number Diff line number Diff line change
@@ -1,9 +1,17 @@
# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
# SPDX-License-Identifier: Apache-2.0

import json
from pathlib import Path

import pandas as pd
import pytest

from data_designer_retrieval_sdg.postprocess import filter_qa_pairs_by_quality, postprocess_retriever_data
from data_designer_retrieval_sdg.postprocess import (
filter_qa_pairs_by_quality,
load_positive_docs_with_modality,
postprocess_retriever_data,
)

# ---------------------------------------------------------------------------
# postprocess_retriever_data
Expand Down Expand Up @@ -85,3 +93,45 @@ def test_filter_skips_mismatched() -> None:
filtered_df, skipped = filter_qa_pairs_by_quality(df, quality_threshold=5.0)
assert len(filtered_df) == 0
assert len(skipped) == 1


# ---------------------------------------------------------------------------
# load_positive_docs_with_modality
# ---------------------------------------------------------------------------


@pytest.mark.parametrize("corpus_id_header", ["corpus-id", "corpus-id "])
def test_load_positive_docs_with_modality_accepts_corpus_id_headers(tmp_path: Path, corpus_id_header: str) -> None:
"""Canonical and legacy corpus ID headers produce the same result."""
qrels_path = tmp_path / "test.tsv"
qrels_path.write_text(f"query-id\t{corpus_id_header}\tscore\nq1\td1\t1\n", encoding="utf-8")

corpus_path = tmp_path / "corpus.jsonl"
corpus_path.write_text(json.dumps({"_id": "d1", "text": "hello world", "title": "T"}) + "\n", encoding="utf-8")

split_path = tmp_path / "split.json"
split_path.write_text(json.dumps({"text": ["q1"]}), encoding="utf-8")

docs_df, modality_map = load_positive_docs_with_modality(qrels_path, corpus_path, split_path)

assert len(docs_df) == 1
assert docs_df.iloc[0]["doc_id"] == "d1"
assert docs_df.iloc[0]["modality"] == "text"
assert modality_map == {"d1": "text"}


def test_load_positive_docs_with_modality_uses_most_common_modality(tmp_path: Path) -> None:
"""A document shared across modalities uses its most frequent modality."""
qrels_path = tmp_path / "test.tsv"
qrels_path.write_text("query-id\tcorpus-id\tscore\nq1\td1\t1\nq2\td1\t1\nq3\td1\t1\n", encoding="utf-8")

corpus_path = tmp_path / "corpus.jsonl"
corpus_path.write_text(json.dumps({"_id": "d1", "text": "hello world", "title": "T"}) + "\n", encoding="utf-8")

split_path = tmp_path / "split.json"
split_path.write_text(json.dumps({"text": ["q1", "q2"], "image": ["q3"]}), encoding="utf-8")

docs_df, modality_map = load_positive_docs_with_modality(qrels_path, corpus_path, split_path)

assert docs_df.iloc[0]["modality"] == "text"
assert modality_map == {"d1": "text"}
Loading