Skip to content

Commit 65c1341

Browse files
committed
fix registry mockup
1 parent 7037815 commit 65c1341

5 files changed

Lines changed: 28 additions & 26 deletions

File tree

src/lighteval/pipeline.py

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -245,7 +245,7 @@ def _init_tasks_and_requests(self, tasks: str):
245245
# The registry contains all the potential tasks
246246
registry = Registry(tasks=tasks, custom_tasks=self.pipeline_parameters.custom_tasks_directory)
247247

248-
# load the tasks fro the configs and their datasets
248+
# load the tasks from the configs and their datasets
249249
self.tasks_dict: dict[str, LightevalTask] = registry.load_tasks()
250250
LightevalTask.load_datasets(self.tasks_dict, self.pipeline_parameters.dataset_loading_processes)
251251
self.documents_dict = {

src/lighteval/tasks/lighteval_task.py

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -113,7 +113,7 @@ def __post_init__(self):
113113
self.evaluation_splits = tuple(self.evaluation_splits)
114114
self.suite = tuple(self.suite)
115115
self.stop_sequence = self.stop_sequence if self.stop_sequence is not None else ()
116-
self.full_name = f"{self.name}|{self.num_fewshots}"
116+
self.full_name = f"{self.name}|{self.num_fewshots}" # todo clefourrier: this is likely incorrect
117117

118118
def print(self):
119119
md_writer = MarkdownTableWriter()

tests/pipeline/test_reasoning_tags.py

Lines changed: 15 additions & 14 deletions
Original file line numberDiff line numberDiff line change
@@ -61,11 +61,13 @@ def setUp(self):
6161
stop_sequence=["\n"],
6262
num_fewshots=0,
6363
)
64+
self.input_task_name = "test|test_reasoning_task|0"
65+
self.task_config_name = self.task_config.full_name
6466

6567
# Create test documents with reasoning tags in expected responses
6668
self.test_docs = [
6769
Doc(
68-
task_name="test|test_reasoning_task|0",
70+
task_name=self.input_task_name,
6971
query="What is 2+2?",
7072
choices=["4"],
7173
gold_index=[0],
@@ -77,7 +79,7 @@ def setUp(self):
7779
# Mock dataset
7880
self.mock_dataset = {"test": self.test_docs}
7981

80-
def _mock_task_registry(self, task_config, task_docs, responses_with_reasoning_tags):
82+
def _mock_task_registry(self, input_task_name, task_config, task_docs, responses_with_reasoning_tags):
8183
"""Create a fake registry for testing."""
8284

8385
class FakeTask(LightevalTask):
@@ -96,13 +98,12 @@ class FakeRegistry(Registry):
9698
def __init__(
9799
self, tasks: Optional[str] = None, custom_tasks: Optional[Union[str, Path, ModuleType]] = None
98100
):
99-
super().__init__(tasks=tasks, custom_tasks=custom_tasks)
101+
self.tasks_list = [input_task_name]
102+
# suite_name, task_name, few_shot = input_task_name.split("|")
103+
self.task_to_configs = {input_task_name: [task_config]}
100104

101-
def get_tasks_configs(self, task: str):
102-
return [task_config]
103-
104-
def get_tasks_from_configs(self, tasks_configs):
105-
return {f"{task_config.suite[0]}|{task_config.full_name}": FakeTask(task_config)}
105+
def load_tasks(self):
106+
return {input_task_name: FakeTask(config=task_config)}
106107

107108
# Create a DummyModel that returns responses with reasoning tags
108109
class TestDummyModel(DummyModel):
@@ -124,7 +125,7 @@ def test_remove_reasoning_tags_enabled(self):
124125
]
125126

126127
FakeRegistry, TestDummyModel = self._mock_task_registry(
127-
self.task_config, self.test_docs, responses_with_reasoning
128+
self.input_task_name, self.task_config, self.test_docs, responses_with_reasoning
128129
)
129130

130131
# Initialize accelerator if available
@@ -170,7 +171,7 @@ def test_remove_reasoning_tags_enabled_tags_as_string(self):
170171
]
171172

172173
FakeRegistry, TestDummyModel = self._mock_task_registry(
173-
self.task_config, self.test_docs, responses_with_reasoning
174+
self.input_task_name, self.task_config, self.test_docs, responses_with_reasoning
174175
)
175176

176177
# Initialize accelerator if available
@@ -216,7 +217,7 @@ def test_remove_reasoning_tags_enabled_default_tags(self):
216217
]
217218

218219
FakeRegistry, TestDummyModel = self._mock_task_registry(
219-
self.task_config, self.test_docs, responses_with_reasoning
220+
self.input_task_name, self.task_config, self.test_docs, responses_with_reasoning
220221
)
221222

222223
# Initialize accelerator if available
@@ -259,7 +260,7 @@ def test_remove_reasoning_tags_disabled(self):
259260
]
260261

261262
FakeRegistry, TestDummyModel = self._mock_task_registry(
262-
self.task_config, self.test_docs, responses_with_reasoning
263+
self.input_task_name, self.task_config, self.test_docs, responses_with_reasoning
263264
)
264265

265266
# Initialize accelerator if available
@@ -305,7 +306,7 @@ def test_custom_reasoning_tags(self):
305306
]
306307

307308
FakeRegistry, TestDummyModel = self._mock_task_registry(
308-
self.task_config, self.test_docs, responses_with_reasoning
309+
self.input_task_name, self.task_config, self.test_docs, responses_with_reasoning
309310
)
310311

311312
# Initialize accelerator if available
@@ -351,7 +352,7 @@ def test_multiple_reasoning_tags(self):
351352
]
352353

353354
FakeRegistry, TestDummyModel = self._mock_task_registry(
354-
self.task_config, self.test_docs, responses_with_reasoning
355+
self.input_task_name, self.task_config, self.test_docs, responses_with_reasoning
355356
)
356357

357358
# Initialize accelerator if available

tests/tasks/test_registry.py

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -95,7 +95,7 @@ def test_superset_with_subset_task():
9595
registry = Registry(tasks="original|mmlu|3,original|mmlu:abstract_algebra|5")
9696

9797
# We have all mmlu tasks
98-
assert registry.tasks_list == ["original|mmlu|3", "original|mmlu:abstract_algebra|5"]
98+
assert set(registry.tasks_list) == {"original|mmlu|3", "original|mmlu:abstract_algebra|5"}
9999
assert len(registry.task_to_configs.keys()) == 57
100100

101101
task_info: list[LightevalTaskConfig] = registry.task_to_configs["original|mmlu:abstract_algebra"]

tests/utils.py

Lines changed: 10 additions & 9 deletions
Original file line numberDiff line numberDiff line change
@@ -98,6 +98,7 @@ def fake_evaluate_task(
9898
# Mock the Registry.get_task_dict method
9999

100100
task_name = f"{lighteval_task.suite[0]}|{lighteval_task.name}"
101+
task_name_fs = f"{lighteval_task.suite[0]}|{lighteval_task.name}|{n_fewshot}"
101102

102103
task_dict = {task_name: lighteval_task}
103104
evaluation_tracker = EvaluationTracker(output_dir="outputs")
@@ -106,17 +107,17 @@ def fake_evaluate_task(
106107

107108
class FakeRegistry(Registry):
108109
def __init__(self, tasks: Optional[str], custom_tasks: Optional[Union[str, Path, ModuleType]] = None):
109-
super().__init__(tasks=tasks, custom_tasks=custom_tasks)
110+
self.tasks_list = [task_name_fs]
111+
self.task_to_configs = {task_name: [lighteval_task.config]}
110112

111-
def get_task_dict(self, task_names: list[str]):
112-
return task_dict
113+
def load_tasks(self):
114+
return {lighteval_task.config.full_name: lighteval_task}
113115

114-
def get_tasks_configs(self, task: str):
115-
config = lighteval_task.config
116-
config.num_fewshots = n_fewshot
117-
config.truncate_fewshots = False
118-
config.full_name = f"{task_name}|{config.num_fewshots}"
119-
return [config]
116+
# def get_tasks_configs(self, task: str):
117+
# config = lighteval_task.config
118+
# config.num_fewshots = n_fewshot
119+
# config.full_name = f"{task_name}|{config.num_fewshots}"
120+
# return [config]
120121

121122
# This is due to logger complaining we have no initialised the accelerator
122123
# It's hard to mock as it's global singleton

0 commit comments

Comments
 (0)