diff --git a/plugins/data-designer-retrieval-sdg/src/data_designer_retrieval_sdg/postprocess.py b/plugins/data-designer-retrieval-sdg/src/data_designer_retrieval_sdg/postprocess.py index 1bbff7f..e31085f 100644 --- a/plugins/data-designer-retrieval-sdg/src/data_designer_retrieval_sdg/postprocess.py +++ b/plugins/data-designer-retrieval-sdg/src/data_designer_retrieval_sdg/postprocess.py @@ -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) @@ -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"] if query_id in query_to_modality: doc_to_modality[corpus_id].add(query_to_modality[query_id]) @@ -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 diff --git a/plugins/data-designer-retrieval-sdg/tests/test_postprocess.py b/plugins/data-designer-retrieval-sdg/tests/test_postprocess.py index f73b24d..8cc9a25 100644 --- a/plugins/data-designer-retrieval-sdg/tests/test_postprocess.py +++ b/plugins/data-designer-retrieval-sdg/tests/test_postprocess.py @@ -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 @@ -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"}