Spaces:
Running
Running
File size: 4,292 Bytes
3b2bae1 | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 | import unittest
from types import SimpleNamespace
from unittest.mock import patch
from chinatravel.data import load_datasets
class QueryLoaderTests(unittest.TestCase):
def test_serialized_oracle_is_parsed_and_preserved_for_evaluation(self):
record = {
"uid": "uid-1",
"hard_logic_py": "['result=True']",
"nature_language": "test",
}
args = SimpleNamespace(
splits="easy",
lang="zh",
oracle_translation=True,
)
with (
patch.object(
load_datasets,
"_load_oracle_snapshot",
return_value=[record],
),
patch.object(
load_datasets,
"_configured_query_ids",
return_value=["uid-1"],
),
):
query_ids, records = load_datasets.load_query(args)
self.assertEqual(query_ids, ["uid-1"])
self.assertEqual(records["uid-1"]["hard_logic_py"], ["result=True"])
def test_agent_facing_load_strips_oracle_fields(self):
record = {
"uid": "uid-1",
"hard_logic_py": ["result=True"],
"nature_language": "test",
}
args = SimpleNamespace(
splits="easy",
lang="zh",
oracle_translation=False,
)
with (
patch.object(
load_datasets,
"_load_huggingface_split",
return_value=[record],
),
patch.object(
load_datasets,
"_configured_query_ids",
return_value=["uid-1"],
),
):
_, records = load_datasets.load_query(args)
self.assertNotIn("hard_logic_py", records["uid-1"])
def test_all_supported_snapshots_have_complete_oracles(self):
supported = {"zh": ("easy", "human"), "en": ("easy", "human")}
for lang, splits in supported.items():
for split in splits:
with self.subTest(lang=lang, split=split):
args = SimpleNamespace(
splits=split,
lang=lang,
oracle_translation=True,
)
query_ids, records = load_datasets.load_query(args)
self.assertEqual(set(query_ids), set(records))
self.assertTrue(
all(
isinstance(record.get("hard_logic_py"), list)
for record in records.values()
)
)
def test_human1000_english_data_is_not_treated_as_an_oracle(self):
args = SimpleNamespace(
splits="human1000",
lang="en",
oracle_translation=True,
)
with self.assertRaisesRegex(ValueError, "No en Oracle source"):
load_datasets.load_query(args)
def test_human1000_oracle_is_loaded_from_runtime_source(self):
args = SimpleNamespace(
splits="human1000",
lang="zh",
oracle_translation=True,
)
record = {
"uid": "uid-1",
"hard_logic_py": "['result=True']",
"nature_language": "test",
}
with (
patch.object(load_datasets, "_load_oracle_snapshot", return_value=None),
patch.object(
load_datasets,
"_load_huggingface_split",
return_value=[record],
) as remote_loader,
patch.object(
load_datasets,
"_configured_query_ids",
return_value=["uid-1"],
),
):
query_ids, records = load_datasets.load_query(args)
remote_loader.assert_called_once_with("human1000")
self.assertEqual(query_ids, ["uid-1"])
self.assertEqual(records["uid-1"]["hard_logic_py"], ["result=True"])
def test_human1000_uses_oracle_benchmark_source(self):
self.assertEqual(
load_datasets.HUGGINGFACE_SPLIT_SOURCES["human1000"],
("LAMDA-NeSy/chinatravel_test", "test"),
)
if __name__ == "__main__":
unittest.main()
|