diff --git a/README.md b/README.md index c6d2c3ae1..bc46ffd9f 100644 --- a/README.md +++ b/README.md @@ -189,11 +189,13 @@ from eval_framework.main import main from eval_framework.tasks.eval_config import EvalConfig from template_formatting.formatter import HFFormatter + # Define your model class MyHuggingFaceModel(HFLLM): LLM_NAME = "microsoft/DialoGPT-medium" DEFAULT_FORMATTER = partial(HFFormatter, "microsoft/DialoGPT-medium") + if __name__ == "__main__": # Initialize your model llm = MyHuggingFaceModel() diff --git a/docs/add_new_benchmark_guide.md b/docs/add_new_benchmark_guide.md index 50a24268f..cc5d0cee3 100644 --- a/docs/add_new_benchmark_guide.md +++ b/docs/add_new_benchmark_guide.md @@ -23,16 +23,16 @@ All benchmarks inherit from `BaseTask[SubjectType]` and must implement several r ```python class YourBenchmark(BaseTask[str]): # or BaseTask[Enum] for multiple subjects # === CORE CONFIGURATION === - NAME: str # Display name for the benchmark - DATASET_PATH: str # HuggingFace dataset path or local path - SAMPLE_SPLIT: str # Dataset split for evaluation samples - FEWSHOT_SPLIT: str # Dataset split for few-shot examples - RESPONSE_TYPE: ResponseType # COMPLETION or LOGLIKELIHOODS - METRICS: list[type[BaseMetric]] # List of metric classes to compute - SUBJECTS: list[SubjectType] # List of subjects/categories to evaluate + NAME: str # Display name for the benchmark + DATASET_PATH: str # HuggingFace dataset path or local path + SAMPLE_SPLIT: str # Dataset split for evaluation samples + FEWSHOT_SPLIT: str # Dataset split for few-shot examples + RESPONSE_TYPE: ResponseType # COMPLETION or LOGLIKELIHOODS + METRICS: list[type[BaseMetric]] # List of metric classes to compute + SUBJECTS: list[SubjectType] # List of subjects/categories to evaluate # === OPTIONAL CONFIGURATION === - HF_REVISION: str | None = None # Git revision for reproducibility + HF_REVISION: str | None = None # Git revision for reproducibility PERTURBATION_UNMODIFIABLE_WORDS: list[str] | None = None # Words to protect from perturbation LANGUAGE: Language | dict[str, Language] | dict[str, tuple[Language, Language]] | None = None # Language(s) tested ``` @@ -44,6 +44,7 @@ def _get_instruction_text(self, item: dict[str, Any]) -> str: """Generate the instruction/question text for a sample.""" pass + def _get_ground_truth(self, item: dict[str, Any]) -> str | None | list[str]: """Extract the correct answer(s) from a dataset item.""" pass @@ -56,38 +57,46 @@ def _get_initial_prompt_text(self, item: dict[str, Any]) -> str: """Text to prepend to the first message.""" return "" + def _get_system_prompt_text(self, item: dict[str, Any]) -> str | None: """System message content.""" return None + def _get_cue_text(self, item: dict[str, Any]) -> str: """Text to append as assistant cue (e.g., 'Answer:').""" return "" + def _get_possible_completions(self, item: dict[str, Any]) -> list[str] | None: """For loglikelihood tasks: list of answer choices.""" return None + def _get_fewshot_target_text(self, item: dict[str, Any]) -> str: """Target text for few-shot examples.""" target = self._get_ground_truth(item) assert target is not None and isinstance(target, str) return target + def _get_context(self, item: dict[str, Any]) -> BaseMetricContext | list[BaseMetricContext] | None: """Additional parameters for evaluation metrics.""" return None + def _sample_fewshot_examples(self, item: dict[str, Any]) -> list[dict]: """Custom few-shot sampling logic.""" # Default implementation samples randomly from FEWSHOT_SPLIT pass + def _create_samples(self, item: dict[str, Any], index: int, subject: str) -> list[Sample]: """Create one or more samples from a dataset item.""" # Default creates single sample - override for multi-sample items pass + def post_process_generated_completion(self, completion_text: str, sample: Sample | None = None) -> str: """Post-process model completions (e.g., extract final answer).""" return completion_text @@ -151,7 +160,6 @@ from eval_framework.metrics.completion.niah_accuracy import NIAHAccuracy from eval_framework.metrics.completion.text_counter import WordCounter from eval_framework.metrics.completion.text_counter import ParagraphCounter from eval_framework.metrics.completion.text_counter import ResponseToOriginalLengthRatio - ``` #### Loglikelihood Metrics @@ -214,7 +222,6 @@ from eval_framework.metrics.llm.llm_judge_sql import LLMJudgeSql from eval_framework.metrics.llm.llm_judge_world_knowledge import LLMJudgeWorldKnowledge # Evaluates whether a summary contains information that goes beyond the reference text (also known as "world knowledge"), returning a boolean classification with detailed reasoning for the assessment. (English, French and German) - ``` ## Implementation Examples and Patterns @@ -231,6 +238,7 @@ from eval_framework.tasks.base import BaseTask from eval_framework.models.sample import ResponseType from eval_framework.metrics.completion.accuracy_completion import AccuracyCompletion + class GeographyQATask(BaseTask[str]): # Required class attributes NAME = "GeographyQA" diff --git a/docs/completion_task_guide.md b/docs/completion_task_guide.md index 226246713..f72cd4a23 100644 --- a/docs/completion_task_guide.md +++ b/docs/completion_task_guide.md @@ -27,7 +27,7 @@ class YourCompletionTask(BaseTask[str]): def _get_ground_truth(self, item: dict[str, Any]) -> str: """Extract the correct answer from the dataset.""" - return item['answer'] + return item["answer"] ``` ## Step-by-Step Implementation @@ -41,6 +41,7 @@ from eval_framework.tasks.base import BaseTask from eval_framework.models.sample import ResponseType from eval_framework.metrics.completion.accuracy_completion import AccuracyCompletion + class MathQATask(BaseTask[str]): NAME = "MathQA" DATASET_PATH = "math_qa_dataset" @@ -76,7 +77,7 @@ This method extracts the correct answer: def _get_ground_truth(self, item: dict[str, Any]) -> str: """Extract the correct answer from the dataset item.""" # Simple case - direct answer - return item['answer'] + return item["answer"] # For numeric answers, you might want to normalize # return str(float(item['answer'])) @@ -102,7 +103,7 @@ class QATask(BaseTask[str]): return f"Question: {item['question']}\nAnswer:" def _get_ground_truth(self, item: dict[str, Any]) -> str: - return item['answer'] + return item["answer"] def _get_cue_text(self, item: dict[str, Any]) -> str: return "Answer:" # Helps model start response correctly @@ -123,13 +124,14 @@ class MathTask(BaseTask[str]): return f"Problem: {item['problem']}\nSolution:" def _get_ground_truth(self, item: dict[str, Any]) -> str: - return item['solution'] + return item["solution"] def post_process_generated_completion(self, completion_text: str, sample: Sample | None = None) -> str: """Extract final numerical answer from solution.""" import re + # Look for "The answer is X" pattern - match = re.search(r'The answer is (\d+(?:\.\d+)?)', completion_text) + match = re.search(r"The answer is (\d+(?:\.\d+)?)", completion_text) if match: return match.group(1) return completion_text.strip() @@ -140,11 +142,14 @@ class MathTask(BaseTask[str]): from eval_framework.metrics.completion.code_execution_pass_at_one import CodeExecutionPassAtOne from eval_framework.shared.types import BaseMetricContext + class CodeTaskMetricContext(BaseMetricContext): """Will be passed to the metric for this task.""" + test_cases: list entry_point: str + class CodeTask(BaseTask[str]): NAME = "Code Generation" DATASET_PATH = "code_dataset" @@ -158,13 +163,13 @@ class CodeTask(BaseTask[str]): return f"Complete this function:\n{item['prompt']}" def _get_ground_truth(self, item: dict[str, Any]) -> str: - return item['canonical_solution'] + return item["canonical_solution"] def _get_context(self, item: dict[str, Any]) -> CodeTaskMetricContext: """Provide test cases for code execution.""" return CodeTaskMetricContext( - test_cases=item['text_cases'], - entry_point=item['entry_point'], + test_cases=item["text_cases"], + entry_point=item["entry_point"], ) ``` @@ -231,6 +236,7 @@ from eval_framework.metrics.completion.csv_format import CSVFormat # Custom metrics using LLM judges from eval_framework.metrics.llm.llm_judge_score import LLMJudgeScore + class YourTask(BaseTask[str]): # Choose metrics appropriate for your task METRICS = [AccuracyCompletion, Rouge1, MathReasoningCompletion] @@ -244,6 +250,7 @@ from eval_framework.tasks.base import BaseTask from eval_framework.models.sample import ResponseType from eval_framework.metrics.completion.accuracy_completion import AccuracyCompletion + class GeographyQuizTask(BaseTask[str]): NAME = "Geography Quiz" DATASET_PATH = "geography_quiz" @@ -259,7 +266,7 @@ class GeographyQuizTask(BaseTask[str]): def _get_ground_truth(self, item: dict[str, Any]) -> str: """Extract the correct capital city.""" - return item['capital'] + return item["capital"] def _get_system_prompt_text(self, item: dict[str, Any]) -> str: """Provide context about the task.""" diff --git a/docs/evaluate_huggingface_model.md b/docs/evaluate_huggingface_model.md index 0e7e810eb..1ae4bdf20 100644 --- a/docs/evaluate_huggingface_model.md +++ b/docs/evaluate_huggingface_model.md @@ -15,11 +15,13 @@ from eval_framework.main import main from eval_framework.tasks.eval_config import EvalConfig from template_formatting.formatter import HFFormatter + # Define your model class MyHuggingFaceModel(HFLLM): LLM_NAME = "context-labs/meta-llama-Llama-3.2-3B-Instruct-FP16" DEFAULT_FORMATTER = partial(HFFormatter, "context-labs/meta-llama-Llama-3.2-3B-Instruct-FP16") + if __name__ == "__main__": # Initialize your model llm = MyHuggingFaceModel() @@ -65,6 +67,7 @@ The formatter determines how prompts are structured for your model. Choose based ```python from template_formatting.formatter import ConcatFormatter + class BaseModel(HFLLM): LLM_NAME = "meta-llama/Llama-3.2-3B" DEFAULT_FORMATTER = ConcatFormatter @@ -75,6 +78,7 @@ class BaseModel(HFLLM): ```python from template_formatting.formatter import Llama3Formatter + class Llama3Model(HFLLM): LLM_NAME = "meta-llama/Meta-Llama-3-8B-Instruct" DEFAULT_FORMATTER = Llama3Formatter @@ -85,6 +89,7 @@ class Llama3Model(HFLLM): ```python from template_formatting.mistral_formatter import MistralFormatter + class MistralModel(HFLLM): LLM_NAME = "mistralai/Mistral-7B-Instruct-v0.1" DEFAULT_FORMATTER = MistralFormatter @@ -95,6 +100,7 @@ class MistralModel(HFLLM): from template_formatting.formatter import HFFormatter from functools import partial + class ChatModel(HFLLM): LLM_NAME = "meta-llama/Llama-3.2-3B-Instruct" DEFAULT_FORMATTER = partial(HFFormatter, "meta-llama/Llama-3.2-3B-Instruct") @@ -113,10 +119,12 @@ class Llama3_8B(HFLLM): LLM_NAME = "meta-llama/Meta-Llama-3-8B-Instruct" DEFAULT_FORMATTER = Llama3Formatter + class Mistral7B(HFLLM): LLM_NAME = "mistralai/Mistral-7B-Instruct-v0.1" DEFAULT_FORMATTER = MistralFormatter + class Qwen2_7B(HFLLM): LLM_NAME = "Qwen/Qwen2-7B-Instruct" DEFAULT_FORMATTER = partial(HFFormatter, "Qwen/Qwen2-7B-Instruct") @@ -128,6 +136,7 @@ class SmolLM(HFLLM): LLM_NAME = "HuggingFaceTB/SmolLM-1.7B-Instruct" DEFAULT_FORMATTER = partial(HFFormatter, "HuggingFaceTB/SmolLM-1.7B-Instruct") + class TinyLlama(HFLLM): LLM_NAME = "TinyLlama/TinyLlama-1.1B-Chat-v1.0" DEFAULT_FORMATTER = partial(HFFormatter, "TinyLlama/TinyLlama-1.1B-Chat-v1.0") @@ -144,15 +153,14 @@ from eval_framework.tasks.eval_config import EvalConfig config = EvalConfig( # Core settings - task_name="MMLU", # Benchmark to run - num_fewshot=5, # Number of examples in prompt - num_samples=100, # How many questions to evaluate - output_dir=Path("./eval_results"), # Where to save results - llm_class=YourModelClass, # Your model class - + task_name="MMLU", # Benchmark to run + num_fewshot=5, # Number of examples in prompt + num_samples=100, # How many questions to evaluate + output_dir=Path("./eval_results"), # Where to save results + llm_class=YourModelClass, # Your model class # Optional settings - task_subjects=["astronomy"], # Specific subjects (if applicable) - batch_size=8, # Batch processing size + task_subjects=["astronomy"], # Specific subjects (if applicable) + batch_size=8, # Batch processing size ) ``` diff --git a/docs/loglikelihood_task_guide.md b/docs/loglikelihood_task_guide.md index 023dc5c75..e391d1020 100644 --- a/docs/loglikelihood_task_guide.md +++ b/docs/loglikelihood_task_guide.md @@ -10,6 +10,7 @@ from eval_framework.tasks.base import BaseTask from eval_framework.models.sample import ResponseType from eval_framework.metrics.loglikelihood.accuracy_loglikelihood import AccuracyLoglikelihood + class YourLoglikelihoodTask(BaseTask[str]): # Required attributes NAME = "YourTaskName" @@ -26,11 +27,11 @@ class YourLoglikelihoodTask(BaseTask[str]): def _get_ground_truth(self, item: dict[str, Any]) -> str: """Return the correct answer choice.""" - return item['correct_answer'] + return item["correct_answer"] def _get_possible_completions(self, item: dict[str, Any]) -> list[str]: """Return all answer choices for ranking.""" - return item['choices'] + return item["choices"] ``` ## Step-by-Step Implementation @@ -44,6 +45,7 @@ from eval_framework.tasks.base import BaseTask from eval_framework.models.sample import ResponseType from eval_framework.metrics.loglikelihood.accuracy_loglikelihood import AccuracyLoglikelihood + class MultipleChoiceTask(BaseTask[str]): NAME = "Multiple Choice" DATASET_PATH = "mcq_dataset" @@ -79,8 +81,8 @@ Return the correct answer choice: def _get_ground_truth(self, item: dict[str, Any]) -> str: """Return the correct answer choice.""" # If dataset has answer index - correct_idx = item['answer_idx'] - return item['choices'][correct_idx] + correct_idx = item["answer_idx"] + return item["choices"][correct_idx] # If dataset has answer directly # return item['correct_answer'] @@ -97,7 +99,7 @@ Return all answer choices: ```python def _get_possible_completions(self, item: dict[str, Any]) -> list[str]: """Return all answer choices for probability ranking.""" - return item['choices'] + return item["choices"] # If choices need formatting # return [f" {choice}" for choice in item['choices']] # Add leading space @@ -115,6 +117,7 @@ from eval_framework.tasks.base import BaseTask from eval_framework.models.sample import ResponseType from eval_framework.metrics.loglikelihood.accuracy_loglikelihood import AccuracyLoglikelihood + class StandardMCQTask(BaseTask[str]): NAME = "Standard MCQ" DATASET_PATH = "mcq_dataset" @@ -128,11 +131,11 @@ class StandardMCQTask(BaseTask[str]): return f"Question: {item['question']}\nAnswer:" def _get_ground_truth(self, item: dict[str, Any]) -> str: - return item['choices'][item['answer_idx']] + return item["choices"][item["answer_idx"]] def _get_possible_completions(self, item: dict[str, Any]) -> list[str]: # Add leading space for better tokenization - return [f" {choice}" for choice in item['choices']] + return [f" {choice}" for choice in item["choices"]] ``` #### Pattern 2: MMLU-style with Labeled Choices @@ -147,12 +150,11 @@ class MMLUStyleTask(BaseTask[str]): SUBJECTS = ["abstract_algebra", "anatomy", "astronomy"] def _get_instruction_text(self, item: dict[str, Any]) -> str: - choices_text = "\n".join([f"{label}. {choice}" - for label, choice in zip(['A', 'B', 'C', 'D'], item['choices'])]) + choices_text = "\n".join([f"{label}. {choice}" for label, choice in zip(["A", "B", "C", "D"], item["choices"])]) return f"Question: {item['question']}\n{choices_text}\nAnswer:" def _get_ground_truth(self, item: dict[str, Any]) -> str: - answer_key = item['answer'] # 'A', 'B', 'C', or 'D' + answer_key = item["answer"] # 'A', 'B', 'C', or 'D' return f" {answer_key}" def _get_possible_completions(self, item: dict[str, Any]) -> list[str]: @@ -174,7 +176,7 @@ class TrueFalseTask(BaseTask[str]): return f"Statement: {item['statement']}\nTrue or False?" def _get_ground_truth(self, item: dict[str, Any]) -> str: - return " True" if item['is_true'] else " False" + return " True" if item["is_true"] else " False" def _get_possible_completions(self, item: dict[str, Any]) -> list[str]: return [" True", " False"] @@ -210,7 +212,7 @@ Format examples consistently: def _get_fewshot_target_text(self, item: dict[str, Any]) -> str: """Format the answer for few-shot examples.""" # For MMLU-style: return the letter - return item['answer'] # 'A', 'B', 'C', or 'D' + return item["answer"] # 'A', 'B', 'C', or 'D' # For true/false: return the word # return "True" if item['is_true'] else "False" @@ -223,6 +225,7 @@ Add helpful context: def _get_system_prompt_text(self, item: dict[str, Any]) -> str: return "You are an expert in multiple choice questions. Choose the best answer." + def _get_initial_prompt_text(self, item: dict[str, Any]) -> str: return "Instructions: Select the most appropriate answer from the given choices." ``` @@ -233,12 +236,14 @@ Handle different subjects: ```python from enum import Enum + class MMLUSubject(Enum): ABSTRACT_ALGEBRA = "abstract_algebra" ANATOMY = "anatomy" ASTRONOMY = "astronomy" # ... more subjects + class MMLUTask(BaseTask[MMLUSubject]): NAME = "MMLU" DATASET_PATH = "mmlu" @@ -249,9 +254,8 @@ class MMLUTask(BaseTask[MMLUSubject]): SUBJECTS = list(MMLUSubject) def _get_instruction_text(self, item: dict[str, Any]) -> str: - subject = item['subject'].replace('_', ' ').title() - choices_text = "\n".join([f"{label}. {choice}" - for label, choice in zip(['A', 'B', 'C', 'D'], item['choices'])]) + subject = item["subject"].replace("_", " ").title() + choices_text = "\n".join([f"{label}. {choice}" for label, choice in zip(["A", "B", "C", "D"], item["choices"])]) return f"The following is a multiple choice question about {subject}.\n\n{item['question']}\n{choices_text}\nAnswer:" ``` @@ -277,6 +281,7 @@ from eval_framework.metrics.loglikelihood.probability_mass_norm import Probabili from eval_framework.metrics.loglikelihood.ternary import TernaryScore from eval_framework.metrics.loglikelihood.dcs import DistributionalCorrectnessScore + class YourTask(BaseTask[str]): # Most common choice METRICS = [AccuracyLoglikelihood] @@ -307,11 +312,13 @@ from eval_framework.tasks.base import BaseTask from eval_framework.models.sample import ResponseType from eval_framework.metrics.loglikelihood.accuracy_loglikelihood import AccuracyLoglikelihood + class ScienceSubject(Enum): PHYSICS = "physics" CHEMISTRY = "chemistry" BIOLOGY = "biology" + class ScienceQuizTask(BaseTask[ScienceSubject]): NAME = "Science Quiz" DATASET_PATH = "science_quiz" @@ -323,9 +330,8 @@ class ScienceQuizTask(BaseTask[ScienceSubject]): def _get_instruction_text(self, item: dict[str, Any]) -> str: """Format science question with choices.""" - subject = item['subject'].replace('_', ' ').title() - choices_text = "\n".join([f"{label}. {choice}" - for label, choice in zip(['A', 'B', 'C', 'D'], item['choices'])]) + subject = item["subject"].replace("_", " ").title() + choices_text = "\n".join([f"{label}. {choice}" for label, choice in zip(["A", "B", "C", "D"], item["choices"])]) return f"Science ({subject}) Question:\n{item['question']}\n\n{choices_text}\n\nAnswer:" def _get_ground_truth(self, item: dict[str, Any]) -> str: @@ -341,7 +347,7 @@ class ScienceQuizTask(BaseTask[ScienceSubject]): def _get_fewshot_target_text(self, item: dict[str, Any]) -> str: """Format target for few-shot examples.""" - return item['answer_key'] # 'A', 'B', 'C', or 'D' (no leading space) + return item["answer_key"] # 'A', 'B', 'C', or 'D' (no leading space) ``` ## Testing Your Completion Task diff --git a/docs/overview_dataloading.md b/docs/overview_dataloading.md index 48e4850d3..464246f69 100644 --- a/docs/overview_dataloading.md +++ b/docs/overview_dataloading.md @@ -10,23 +10,27 @@ The framework uses different data types based on the evaluation approach: from eval_framework.shared.types import Completion, Loglikelihood, RawCompletion, RawLoglikelihood from template_formatting.formatter import Message, Role + # For completion tasks (text generation) class Completion(BaseModel): - completion_text: str # Generated text from the model + completion_text: str # Generated text from the model # Additional fields based on actual implementation + # For loglikelihood tasks (multiple choice) class Loglikelihood(BaseModel): - loglikelihoods: list[float] # Probability scores for each choice + loglikelihoods: list[float] # Probability scores for each choice # Additional fields based on actual implementation + # Raw response types from LLMs class RawCompletion(BaseModel): - text: str # Raw generated text + text: str # Raw generated text # Additional fields based on actual implementation + class RawLoglikelihood(BaseModel): - loglikelihoods: list[float] # Raw probability scores + loglikelihoods: list[float] # Raw probability scores # Additional fields based on actual implementation ``` @@ -37,9 +41,10 @@ Each prompt is structured as a sequence of messages using the template formattin ```python from template_formatting.formatter import Message, Role + class Message(BaseModel): - role: Role # SYSTEM, USER, or ASSISTANT - content: str # Message content + role: Role # SYSTEM, USER, or ASSISTANT + content: str # Message content # Additional fields based on actual formatter implementation ``` @@ -52,6 +57,7 @@ Custom tasks inherit from `BaseTask` and implement specific methods based on the from eval_framework.tasks.base import BaseTask from eval_framework.models.sample import ResponseType + class MyCompletionTask(BaseTask[str]): NAME = "My Task" DATASET_PATH = "dataset_name" @@ -63,7 +69,7 @@ class MyCompletionTask(BaseTask[str]): def _get_ground_truth(self, item: dict) -> str: """Return the expected answer.""" - return item['answer'] + return item["answer"] ``` #### For Loglikelihood Tasks: @@ -79,11 +85,11 @@ class MyLoglikelihoodTask(BaseTask[str]): def _get_ground_truth(self, item: dict) -> str: """Return the correct answer choice.""" - return item['choices'][item['answer_idx']] + return item["choices"][item["answer_idx"]] def _get_possible_completions(self, item: dict) -> list[str]: """Return all answer choices for ranking.""" - return item['choices'] + return item["choices"] ``` ### Few-Shot Example Construction @@ -103,21 +109,12 @@ def construct_prompt(self, item: dict) -> list[Message]: fewshot_examples = self._sample_fewshot_examples(item) for example in fewshot_examples: # User instruction - messages.append(Message( - role=Role.USER, - content=self._get_instruction_text(example) - )) + messages.append(Message(role=Role.USER, content=self._get_instruction_text(example))) # Assistant response - messages.append(Message( - role=Role.ASSISTANT, - content=self._get_fewshot_target_text(example) - )) + messages.append(Message(role=Role.ASSISTANT, content=self._get_fewshot_target_text(example))) # 3. Actual instruction - messages.append(Message( - role=Role.USER, - content=self._get_instruction_text(item) - )) + messages.append(Message(role=Role.USER, content=self._get_instruction_text(item))) # 4. Response cue (optional) if cue := self._get_cue_text(item): @@ -138,7 +135,7 @@ messages = [ Message(Role.USER, "Question: What is the capital of France?"), Message(Role.ASSISTANT, "Answer: Paris"), Message(Role.USER, "Question: What is the capital of Italy?"), - Message(Role.ASSISTANT, "Answer:") + Message(Role.ASSISTANT, "Answer:"), ] ``` diff --git a/src/eval_framework/tasks/registry.py b/src/eval_framework/tasks/registry.py index 5c72dc7fa..89ab96d05 100644 --- a/src/eval_framework/tasks/registry.py +++ b/src/eval_framework/tasks/registry.py @@ -223,20 +223,15 @@ def markdown_doc(self, formatters: Sequence[BaseFormatter]) -> str: class Registry: - """A registry for tasks with support for lazy loading. - - Task names are hashed based on the upper-case name, to avoid issues with - ambiguous naming. - """ + """A registry for Tasks""" def __init__(self) -> None: - # TODO: Lookup only with upper names - self._registry: dict[str, tuple[str, EvalFactory]] = dict() + self._registry: dict[str, EvalFactory] = dict() def __iter__(self) -> Iterator[str]: """Iterate over all task names in the registry.""" - for name, _ in self._registry.values(): - yield name + for factory in self._registry.values(): + yield factory.id() def task_names(self) -> list[str]: """The names of all registered tasks.""" @@ -244,7 +239,8 @@ def task_names(self) -> list[str]: def items(self) -> Iterator[tuple[str, EvalFactory]]: """Iterate over `(task name, EvalFactory)` pairs in the registry.""" - yield from self._registry.values() + for factory in self._registry.values(): + yield factory.id(), factory @staticmethod def _task_key(name: str, /) -> str: @@ -262,26 +258,25 @@ def __contains__(self, name: str) -> bool: def __getitem__(self, name: str, /) -> EvalFactory: task_key = self._task_key(name) try: - _, factory = self._registry[task_key] + return self._registry[task_key] except KeyError: raise KeyError(f"Task not found: {name=} with task_key {task_key=}") - return factory - - def __setitem__(self, name: str, factory: EvalFactory) -> None: - task_key = self._task_key(name) + def add(self, factory: EvalFactory) -> None: + """Register a factory under the key derived from its ``id()``.""" + task_key = self._task_key(factory.id()) if task_key in self._registry: raise ValueError(f"Cannot register duplicate task with key: {task_key}") - self._registry[task_key] = (name, factory) + self._registry[task_key] = factory def register(self, task: type[BaseTask]) -> str: """The class name is used as the task name.""" if not issubclass(task, BaseTask): raise ValueError(f"Can only register subclasses of BaseTask, got {task}") - name = task.__name__ - self[name] = _Eager(task) - return name + factory = _Eager(task) + self.add(factory) + return factory.id() def register_lazy(self, class_path: str, /) -> None: """Register a task by its dotted class path, without importing its module.""" @@ -291,7 +286,7 @@ def register_lazy(self, class_path: str, /) -> None: "`eval_framework.tasks.benchmarks.mmlu.MMLU`): " ) base_module, class_name = class_path.rsplit(".", maxsplit=1) - self[class_name] = _Lazy(class_name=class_name, module=base_module) + self.add(_Lazy(class_name=class_name, module=base_module)) _REGISTRY = Registry() @@ -335,6 +330,7 @@ def register_task(task: type[BaseTask]) -> str: return registry().register(task) -def register_lazy_task(class_path: str, /) -> None: +def register_lazy_task(class_path: str, /, registry: Registry | None = None) -> None: """Register a task by its dotted class path, without importing its module.""" - registry().register_lazy(class_path) + r = registry if registry is not None else _REGISTRY + r.register_lazy(class_path) diff --git a/src/eval_framework/tasks/task_names.py b/src/eval_framework/tasks/task_names.py index 60219ced2..10f193473 100644 --- a/src/eval_framework/tasks/task_names.py +++ b/src/eval_framework/tasks/task_names.py @@ -1,7 +1,8 @@ from enum import Enum from eval_framework.tasks.base import BaseTask -from eval_framework.tasks.registry import register_lazy_task +from eval_framework.tasks.registry import Registry, register_lazy_task +from eval_framework.tasks.registry import registry as global_registry class TaskNameEnum(Enum): @@ -10,73 +11,203 @@ def value(self) -> type[BaseTask]: return super().value -def register_all_tasks() -> None: - """Register all the benchmark tasks with the eval framework.""" - register_lazy_task("eval_framework.tasks.benchmarks.math_reasoning.AIME2024") - register_lazy_task("eval_framework.tasks.benchmarks.math_reasoning.AIME2025") - register_lazy_task("eval_framework.tasks.benchmarks.math_reasoning.AIME2026") - register_lazy_task("eval_framework.tasks.benchmarks.arc.ARC") - register_lazy_task("eval_framework.tasks.benchmarks.arc.ARC_IDK") - register_lazy_task("eval_framework.tasks.benchmarks.arc.ARC_OLMES") - register_lazy_task("eval_framework.tasks.benchmarks.arc_de.ARC_DE") - register_lazy_task("eval_framework.tasks.benchmarks.bigcodebench.BigCodeBench_OLMES") - register_lazy_task("eval_framework.tasks.benchmarks.copa.COPA_OLMES") - register_lazy_task("eval_framework.tasks.benchmarks.goldenswag.GOLDENSWAG") - register_lazy_task("eval_framework.tasks.benchmarks.goldenswag.GOLDENSWAG_IDK") - register_lazy_task("eval_framework.tasks.benchmarks.gpqa.GPQA_OLMES") - register_lazy_task("eval_framework.tasks.benchmarks.gpqa.GPQA_DIAMOND_COT") - register_lazy_task("eval_framework.tasks.benchmarks.gsm8k.GSM8K_OLMES") - register_lazy_task("eval_framework.tasks.benchmarks.gsm8k.GSM8KBPB") - register_lazy_task("eval_framework.tasks.benchmarks.math_reasoning.MATHMinervaBPB") - register_lazy_task("eval_framework.tasks.benchmarks.math_reasoning.GSM8KReasoning") - register_lazy_task("eval_framework.tasks.benchmarks.hellaswag.HELLASWAG") - register_lazy_task("eval_framework.tasks.benchmarks.hellaswag.HELLASWAG_OLMES") - register_lazy_task("eval_framework.tasks.benchmarks.humaneval.HumanEvalBPB") - register_lazy_task("eval_framework.tasks.benchmarks.humaneval.HumanEval_OLMES") - register_lazy_task("eval_framework.tasks.benchmarks.ifeval.IFEval") - register_lazy_task("eval_framework.tasks.benchmarks.ifeval.IFEvalDe") - register_lazy_task("eval_framework.tasks.benchmarks.math_reasoning.MATH500") - register_lazy_task("eval_framework.tasks.benchmarks.math_reasoning.MATHMinerva_OLMES") - register_lazy_task("eval_framework.tasks.benchmarks.math_reasoning.MATHMinerva_OLMES_NONL") - register_lazy_task("eval_framework.tasks.benchmarks.multipl_e.MultiPLEHumanEvalCpp") - register_lazy_task("eval_framework.tasks.benchmarks.multipl_e.MultiPLEHumanEvalJava") - register_lazy_task("eval_framework.tasks.benchmarks.multipl_e.MultiPLEHumanEvalJs") - register_lazy_task("eval_framework.tasks.benchmarks.multipl_e.MultiPLEHumanEvalPhp") - register_lazy_task("eval_framework.tasks.benchmarks.multipl_e.MultiPLEHumanEvalRs") - register_lazy_task("eval_framework.tasks.benchmarks.multipl_e.MultiPLEHumanEvalSh") - register_lazy_task("eval_framework.tasks.benchmarks.multipl_e.MultiPLEMBPPCpp") - register_lazy_task("eval_framework.tasks.benchmarks.multipl_e.MultiPLEMBPPJava") - register_lazy_task("eval_framework.tasks.benchmarks.multipl_e.MultiPLEMBPPJs") - register_lazy_task("eval_framework.tasks.benchmarks.multipl_e.MultiPLEMBPPPhp") - register_lazy_task("eval_framework.tasks.benchmarks.multipl_e.MultiPLEMBPPRs") - register_lazy_task("eval_framework.tasks.benchmarks.multipl_e.MultiPLEMBPPSh") - register_lazy_task("eval_framework.tasks.benchmarks.mbpp.MBPPBPB") - register_lazy_task("eval_framework.tasks.benchmarks.mbpp.MBPP_OLMES") - register_lazy_task("eval_framework.tasks.benchmarks.mbpp.MBPP_EvalPlus") - register_lazy_task("eval_framework.tasks.benchmarks.mbpp.MBPP_BPB_EvalPlus") - register_lazy_task("eval_framework.tasks.benchmarks.mmlu.MMLU") - register_lazy_task("eval_framework.tasks.benchmarks.mmlu.MMLU_IDK") - register_lazy_task("eval_framework.tasks.benchmarks.mmlu.MMLU_OLMES") - register_lazy_task("eval_framework.tasks.benchmarks.mmlu.FullTextMMLU") - register_lazy_task("eval_framework.tasks.benchmarks.mmlu_pro.MMLU_PRO") - register_lazy_task("eval_framework.tasks.benchmarks.mmlu_pro.MMLU_PRO_IDK") - register_lazy_task("eval_framework.tasks.benchmarks.mmlu_pro.MMLU_PRO_OLMES") - register_lazy_task("eval_framework.tasks.benchmarks.mmlu_pro.MMLU_PRO_COT") - register_lazy_task("eval_framework.tasks.benchmarks.mmlu.MMLU_COT") - register_lazy_task("eval_framework.tasks.benchmarks.global_mmlu.GlobalMMLU") - register_lazy_task("eval_framework.tasks.benchmarks.global_mmlu.GlobalMMLU_German") - register_lazy_task("eval_framework.tasks.benchmarks.piqa.PIQA") - register_lazy_task("eval_framework.tasks.benchmarks.piqa.PIQA_IDK") - register_lazy_task("eval_framework.tasks.benchmarks.piqa.PIQA_OLMES") - register_lazy_task("eval_framework.tasks.benchmarks.sciq.SCIQ_OLMES") - register_lazy_task("eval_framework.tasks.benchmarks.squad.SQuAD_OLMES") - register_lazy_task("eval_framework.tasks.benchmarks.squad.SQuAD2_MA") - register_lazy_task("eval_framework.tasks.benchmarks.squad.SQuAD2_MA_NO_SYSPROMPT") - register_lazy_task("eval_framework.tasks.benchmarks.winogrande.WINOGRANDECloze") - register_lazy_task("eval_framework.tasks.benchmarks.csqa.CommonsenseQAMC_OLMES") - register_lazy_task("eval_framework.tasks.benchmarks.drop.DropCompletion_OLMES") - register_lazy_task("eval_framework.tasks.benchmarks.drop.DropMC_OLMES") - register_lazy_task("eval_framework.tasks.benchmarks.naturalqs_open.NaturalQsOpen") - register_lazy_task("eval_framework.tasks.benchmarks.naturalqs_open.NaturalQsOpenMC_OLMES") - register_lazy_task("eval_framework.tasks.benchmarks.social_iqa.SocialIQAMC_OLMES") - register_lazy_task("eval_framework.tasks.benchmarks.medqa.MedQAMC_OLMES") +def register_all_tasks(registry: Registry | None = None) -> None: + """Register all the benchmark tasks with the eval framework + + Uses global registry by default. + """ + registry = registry if registry is not None else global_registry() + + register_math_reasoning_tasks(registry=registry) + register_arc_tasks(registry=registry) + register_arc_de_tasks(registry=registry) + register_bigcodebench_tasks(registry=registry) + register_copa_tasks(registry=registry) + register_goldenswag_tasks(registry=registry) + register_gpqa_tasks(registry=registry) + register_gsm8k_tasks(registry=registry) + register_hellaswag_tasks(registry=registry) + register_humaneval_tasks(registry=registry) + register_ifeval_tasks(registry=registry) + register_multipl_e_tasks(registry=registry) + register_mbpp_tasks(registry=registry) + register_mmlu_tasks(registry=registry) + register_mmlu_pro_tasks(registry=registry) + register_global_mmlu_tasks(registry=registry) + register_piqa_tasks(registry=registry) + register_sciq_tasks(registry=registry) + register_squad_tasks(registry=registry) + register_winogrande_tasks(registry=registry) + register_csqa_tasks(registry=registry) + register_drop_tasks(registry=registry) + register_naturalqs_open_tasks(registry=registry) + register_social_iqa_tasks(registry=registry) + register_medqa_tasks(registry=registry) + + +def register_arc_tasks(registry: Registry) -> None: + """Register arc benchmark tasks.""" + register_lazy_task("eval_framework.tasks.benchmarks.arc.ARC", registry=registry) + register_lazy_task("eval_framework.tasks.benchmarks.arc.ARC_IDK", registry=registry) + register_lazy_task("eval_framework.tasks.benchmarks.arc.ARC_OLMES", registry=registry) + + +def register_hellaswag_tasks(registry: Registry) -> None: + """Register hellaswag benchmark tasks.""" + register_lazy_task("eval_framework.tasks.benchmarks.hellaswag.HELLASWAG", registry=registry) + register_lazy_task("eval_framework.tasks.benchmarks.hellaswag.HELLASWAG_OLMES", registry=registry) + + +def register_piqa_tasks(registry: Registry) -> None: + """Register piqa benchmark tasks.""" + register_lazy_task("eval_framework.tasks.benchmarks.piqa.PIQA", registry=registry) + register_lazy_task("eval_framework.tasks.benchmarks.piqa.PIQA_IDK", registry=registry) + register_lazy_task("eval_framework.tasks.benchmarks.piqa.PIQA_OLMES", registry=registry) + + +def register_gpqa_tasks(registry: Registry) -> None: + """Register gpqa benchmark tasks.""" + register_lazy_task("eval_framework.tasks.benchmarks.gpqa.GPQA_OLMES", registry=registry) + register_lazy_task("eval_framework.tasks.benchmarks.gpqa.GPQA_DIAMOND_COT", registry=registry) + + +def register_gsm8k_tasks(registry: Registry) -> None: + """Register gsm8k benchmark tasks.""" + register_lazy_task("eval_framework.tasks.benchmarks.gsm8k.GSM8K_OLMES", registry=registry) + register_lazy_task("eval_framework.tasks.benchmarks.gsm8k.GSM8KBPB", registry=registry) + + +def register_math_reasoning_tasks(registry: Registry) -> None: + """Register math_reasoning benchmark tasks.""" + register_lazy_task("eval_framework.tasks.benchmarks.math_reasoning.AIME2024", registry=registry) + register_lazy_task("eval_framework.tasks.benchmarks.math_reasoning.AIME2026", registry=registry) + register_lazy_task("eval_framework.tasks.benchmarks.math_reasoning.AIME2025", registry=registry) + register_lazy_task("eval_framework.tasks.benchmarks.math_reasoning.MATHMinervaBPB", registry=registry) + register_lazy_task("eval_framework.tasks.benchmarks.math_reasoning.GSM8KReasoning", registry=registry) + register_lazy_task("eval_framework.tasks.benchmarks.math_reasoning.MATH500", registry=registry) + register_lazy_task("eval_framework.tasks.benchmarks.math_reasoning.MATHMinerva_OLMES", registry=registry) + register_lazy_task("eval_framework.tasks.benchmarks.math_reasoning.MATHMinerva_OLMES_NONL", registry=registry) + + +def register_mmlu_tasks(registry: Registry) -> None: + """Register mmlu benchmark tasks.""" + register_lazy_task("eval_framework.tasks.benchmarks.mmlu.MMLU", registry=registry) + register_lazy_task("eval_framework.tasks.benchmarks.mmlu.MMLU_IDK", registry=registry) + register_lazy_task("eval_framework.tasks.benchmarks.mmlu.MMLU_OLMES", registry=registry) + register_lazy_task("eval_framework.tasks.benchmarks.mmlu.FullTextMMLU", registry=registry) + register_lazy_task("eval_framework.tasks.benchmarks.mmlu.MMLU_COT", registry=registry) + + +def register_humaneval_tasks(registry: Registry) -> None: + """Register humaneval benchmark tasks.""" + register_lazy_task("eval_framework.tasks.benchmarks.humaneval.HumanEvalBPB", registry=registry) + register_lazy_task("eval_framework.tasks.benchmarks.humaneval.HumanEval_OLMES", registry=registry) + + +def register_mbpp_tasks(registry: Registry) -> None: + """Register mbpp benchmark tasks.""" + register_lazy_task("eval_framework.tasks.benchmarks.mbpp.MBPPBPB", registry=registry) + register_lazy_task("eval_framework.tasks.benchmarks.mbpp.MBPP_OLMES", registry=registry) + register_lazy_task("eval_framework.tasks.benchmarks.mbpp.MBPP_EvalPlus", registry=registry) + register_lazy_task("eval_framework.tasks.benchmarks.mbpp.MBPP_BPB_EvalPlus", registry=registry) + + +def register_bigcodebench_tasks(registry: Registry) -> None: + """Register bigcodebench benchmark tasks.""" + register_lazy_task("eval_framework.tasks.benchmarks.bigcodebench.BigCodeBench_OLMES", registry=registry) + + +def register_arc_de_tasks(registry: Registry) -> None: + """Register arc_de benchmark tasks.""" + register_lazy_task("eval_framework.tasks.benchmarks.arc_de.ARC_DE", registry=registry) + + +def register_copa_tasks(registry: Registry) -> None: + """Register copa benchmark tasks.""" + register_lazy_task("eval_framework.tasks.benchmarks.copa.COPA_OLMES", registry=registry) + + +def register_goldenswag_tasks(registry: Registry) -> None: + """Register goldenswag benchmark tasks.""" + register_lazy_task("eval_framework.tasks.benchmarks.goldenswag.GOLDENSWAG", registry=registry) + register_lazy_task("eval_framework.tasks.benchmarks.goldenswag.GOLDENSWAG_IDK", registry=registry) + + +def register_ifeval_tasks(registry: Registry) -> None: + """Register ifeval benchmark tasks.""" + register_lazy_task("eval_framework.tasks.benchmarks.ifeval.IFEval", registry=registry) + register_lazy_task("eval_framework.tasks.benchmarks.ifeval.IFEvalDe", registry=registry) + + +def register_multipl_e_tasks(registry: Registry) -> None: + """Register multipl_e benchmark tasks.""" + register_lazy_task("eval_framework.tasks.benchmarks.multipl_e.MultiPLEHumanEvalCpp", registry=registry) + register_lazy_task("eval_framework.tasks.benchmarks.multipl_e.MultiPLEHumanEvalJava", registry=registry) + register_lazy_task("eval_framework.tasks.benchmarks.multipl_e.MultiPLEHumanEvalJs", registry=registry) + register_lazy_task("eval_framework.tasks.benchmarks.multipl_e.MultiPLEHumanEvalPhp", registry=registry) + register_lazy_task("eval_framework.tasks.benchmarks.multipl_e.MultiPLEHumanEvalRs", registry=registry) + register_lazy_task("eval_framework.tasks.benchmarks.multipl_e.MultiPLEHumanEvalSh", registry=registry) + register_lazy_task("eval_framework.tasks.benchmarks.multipl_e.MultiPLEMBPPCpp", registry=registry) + register_lazy_task("eval_framework.tasks.benchmarks.multipl_e.MultiPLEMBPPJava", registry=registry) + register_lazy_task("eval_framework.tasks.benchmarks.multipl_e.MultiPLEMBPPJs", registry=registry) + register_lazy_task("eval_framework.tasks.benchmarks.multipl_e.MultiPLEMBPPPhp", registry=registry) + register_lazy_task("eval_framework.tasks.benchmarks.multipl_e.MultiPLEMBPPRs", registry=registry) + register_lazy_task("eval_framework.tasks.benchmarks.multipl_e.MultiPLEMBPPSh", registry=registry) + + +def register_mmlu_pro_tasks(registry: Registry) -> None: + """Register mmlu_pro benchmark tasks.""" + register_lazy_task("eval_framework.tasks.benchmarks.mmlu_pro.MMLU_PRO", registry=registry) + register_lazy_task("eval_framework.tasks.benchmarks.mmlu_pro.MMLU_PRO_IDK", registry=registry) + register_lazy_task("eval_framework.tasks.benchmarks.mmlu_pro.MMLU_PRO_OLMES", registry=registry) + register_lazy_task("eval_framework.tasks.benchmarks.mmlu_pro.MMLU_PRO_COT", registry=registry) + + +def register_global_mmlu_tasks(registry: Registry) -> None: + """Register global_mmlu benchmark tasks.""" + register_lazy_task("eval_framework.tasks.benchmarks.global_mmlu.GlobalMMLU", registry=registry) + register_lazy_task("eval_framework.tasks.benchmarks.global_mmlu.GlobalMMLU_German", registry=registry) + + +def register_sciq_tasks(registry: Registry) -> None: + """Register sciq benchmark tasks.""" + register_lazy_task("eval_framework.tasks.benchmarks.sciq.SCIQ_OLMES", registry=registry) + + +def register_squad_tasks(registry: Registry) -> None: + """Register squad benchmark tasks.""" + register_lazy_task("eval_framework.tasks.benchmarks.squad.SQuAD_OLMES", registry=registry) + register_lazy_task("eval_framework.tasks.benchmarks.squad.SQuAD2_MA", registry=registry) + register_lazy_task("eval_framework.tasks.benchmarks.squad.SQuAD2_MA_NO_SYSPROMPT", registry=registry) + + +def register_winogrande_tasks(registry: Registry) -> None: + """Register winogrande benchmark tasks.""" + register_lazy_task("eval_framework.tasks.benchmarks.winogrande.WINOGRANDECloze", registry=registry) + + +def register_csqa_tasks(registry: Registry) -> None: + """Register csqa benchmark tasks.""" + register_lazy_task("eval_framework.tasks.benchmarks.csqa.CommonsenseQAMC_OLMES", registry=registry) + + +def register_drop_tasks(registry: Registry) -> None: + """Register drop benchmark tasks.""" + register_lazy_task("eval_framework.tasks.benchmarks.drop.DropCompletion_OLMES", registry=registry) + register_lazy_task("eval_framework.tasks.benchmarks.drop.DropMC_OLMES", registry=registry) + + +def register_naturalqs_open_tasks(registry: Registry) -> None: + """Register naturalqs_open benchmark tasks.""" + register_lazy_task("eval_framework.tasks.benchmarks.naturalqs_open.NaturalQsOpen", registry=registry) + register_lazy_task("eval_framework.tasks.benchmarks.naturalqs_open.NaturalQsOpenMC_OLMES", registry=registry) + + +def register_social_iqa_tasks(registry: Registry) -> None: + """Register social_iqa benchmark tasks.""" + register_lazy_task("eval_framework.tasks.benchmarks.social_iqa.SocialIQAMC_OLMES", registry=registry) + + +def register_medqa_tasks(registry: Registry) -> None: + """Register medqa benchmark tasks.""" + register_lazy_task("eval_framework.tasks.benchmarks.medqa.MedQAMC_OLMES", registry=registry) diff --git a/tests/tests_eval_framework/tasks/benchmarks/test_arc.py b/tests/tests_eval_framework/tasks/benchmarks/test_arc.py index f97b9ad1b..d03ed6f7b 100644 --- a/tests/tests_eval_framework/tasks/benchmarks/test_arc.py +++ b/tests/tests_eval_framework/tasks/benchmarks/test_arc.py @@ -1,11 +1,17 @@ import pytest +from eval_framework.tasks.registry import Registry +from eval_framework.tasks.task_names import register_arc_tasks from template_formatting.formatter import BaseFormatter, ConcatFormatter, Llama3Formatter -from tests.tests_eval_framework.tasks.benchmarks.utils import get_task_names_for_module, run_formatter_hash_test +from tests.tests_eval_framework.tasks.benchmarks.utils import run_formatter_hash_test + +# Registry for this test suite only holding arc tasks +_arc_registry = Registry() +register_arc_tasks(registry=_arc_registry) @pytest.mark.formatter_hash @pytest.mark.parametrize("formatter_cls", [Llama3Formatter, ConcatFormatter]) -@pytest.mark.parametrize("task_name", get_task_names_for_module("arc")) +@pytest.mark.parametrize("task_name", _arc_registry.task_names()) def test_formatter_hash(task_name: str, formatter_cls: type[BaseFormatter]) -> None: - run_formatter_hash_test(task_name, formatter_cls) + run_formatter_hash_test(task_name, formatter_cls, registry=_arc_registry) diff --git a/tests/tests_eval_framework/tasks/benchmarks/test_arc_de.py b/tests/tests_eval_framework/tasks/benchmarks/test_arc_de.py index c5e3f88da..cebdab2f6 100644 --- a/tests/tests_eval_framework/tasks/benchmarks/test_arc_de.py +++ b/tests/tests_eval_framework/tasks/benchmarks/test_arc_de.py @@ -1,11 +1,17 @@ import pytest +from eval_framework.tasks.registry import Registry +from eval_framework.tasks.task_names import register_arc_de_tasks from template_formatting.formatter import BaseFormatter, ConcatFormatter, Llama3Formatter -from tests.tests_eval_framework.tasks.benchmarks.utils import get_task_names_for_module, run_formatter_hash_test +from tests.tests_eval_framework.tasks.benchmarks.utils import run_formatter_hash_test + +# Registry for this test suite only holding arc_de tasks +_arc_de_registry = Registry() +register_arc_de_tasks(registry=_arc_de_registry) @pytest.mark.formatter_hash @pytest.mark.parametrize("formatter_cls", [Llama3Formatter, ConcatFormatter]) -@pytest.mark.parametrize("task_name", get_task_names_for_module("arc_de")) +@pytest.mark.parametrize("task_name", _arc_de_registry.task_names()) def test_formatter_hash(task_name: str, formatter_cls: type[BaseFormatter]) -> None: - run_formatter_hash_test(task_name, formatter_cls) + run_formatter_hash_test(task_name, formatter_cls, registry=_arc_de_registry) diff --git a/tests/tests_eval_framework/tasks/benchmarks/test_bigcodebench.py b/tests/tests_eval_framework/tasks/benchmarks/test_bigcodebench.py index d021d0401..870962053 100644 --- a/tests/tests_eval_framework/tasks/benchmarks/test_bigcodebench.py +++ b/tests/tests_eval_framework/tasks/benchmarks/test_bigcodebench.py @@ -2,18 +2,21 @@ from datasets import DownloadConfig, load_dataset from eval_framework.tasks.benchmarks.bigcodebench import extract_executable_code +from eval_framework.tasks.registry import Registry +from eval_framework.tasks.task_names import register_bigcodebench_tasks from eval_framework.tasks.utils import BIG_CODE_BENCH_PACKAGE_MAPPING, extract_imports from template_formatting.formatter import BaseFormatter, ConcatFormatter, Llama3Formatter -from tests.tests_eval_framework.tasks.benchmarks.utils import get_task_names_for_module, run_formatter_hash_test +from tests.tests_eval_framework.tasks.benchmarks.utils import run_formatter_hash_test -# BigCodeBenchHard / BigCodeBenchHardInstruct / BigCodeBenchInstruct have non-deterministic -# dataset/sample selection across runs, so their formatter output hashes are not stable. -_SKIPPED_TASKS = ["BigCodeBenchHard", "BigCodeBenchHardInstruct", "BigCodeBenchInstruct"] _NUM_FEWSHOT = { "BigCodeBench": 0, "BigCodeBench_OLMES": 3, } +# Registry for this test suite only holding bigcodebench tasks +_bigcodebench_registry = Registry() +register_bigcodebench_tasks(registry=_bigcodebench_registry) + class TestExtractExecutableCode: def test_python_code_block(self) -> None: @@ -333,6 +336,8 @@ def test_all_imports_in_mapping() -> None: @pytest.mark.formatter_hash @pytest.mark.parametrize("formatter_cls", [Llama3Formatter, ConcatFormatter]) -@pytest.mark.parametrize("task_name", get_task_names_for_module("bigcodebench", skip_tasks=_SKIPPED_TASKS)) +@pytest.mark.parametrize("task_name", _bigcodebench_registry.task_names()) def test_formatter_hash(task_name: str, formatter_cls: type[BaseFormatter]) -> None: - run_formatter_hash_test(task_name, formatter_cls, num_fewshot=_NUM_FEWSHOT.get(task_name, 1)) + run_formatter_hash_test( + task_name, formatter_cls, num_fewshot=_NUM_FEWSHOT.get(task_name, 1), registry=_bigcodebench_registry + ) diff --git a/tests/tests_eval_framework/tasks/benchmarks/test_copa.py b/tests/tests_eval_framework/tasks/benchmarks/test_copa.py index 7985ffc29..7d7b2c9e6 100644 --- a/tests/tests_eval_framework/tasks/benchmarks/test_copa.py +++ b/tests/tests_eval_framework/tasks/benchmarks/test_copa.py @@ -1,11 +1,17 @@ import pytest +from eval_framework.tasks.registry import Registry +from eval_framework.tasks.task_names import register_copa_tasks from template_formatting.formatter import BaseFormatter, ConcatFormatter, Llama3Formatter -from tests.tests_eval_framework.tasks.benchmarks.utils import get_task_names_for_module, run_formatter_hash_test +from tests.tests_eval_framework.tasks.benchmarks.utils import run_formatter_hash_test + +# Registry for this test suite only holding copa tasks +_copa_registry = Registry() +register_copa_tasks(registry=_copa_registry) @pytest.mark.formatter_hash @pytest.mark.parametrize("formatter_cls", [Llama3Formatter, ConcatFormatter]) -@pytest.mark.parametrize("task_name", get_task_names_for_module("copa")) +@pytest.mark.parametrize("task_name", _copa_registry.task_names()) def test_formatter_hash(task_name: str, formatter_cls: type[BaseFormatter]) -> None: - run_formatter_hash_test(task_name, formatter_cls) + run_formatter_hash_test(task_name, formatter_cls, registry=_copa_registry) diff --git a/tests/tests_eval_framework/tasks/benchmarks/test_csqa.py b/tests/tests_eval_framework/tasks/benchmarks/test_csqa.py index 182948b74..29b35071f 100644 --- a/tests/tests_eval_framework/tasks/benchmarks/test_csqa.py +++ b/tests/tests_eval_framework/tasks/benchmarks/test_csqa.py @@ -1,11 +1,17 @@ import pytest +from eval_framework.tasks.registry import Registry +from eval_framework.tasks.task_names import register_csqa_tasks from template_formatting.formatter import BaseFormatter, ConcatFormatter, Llama3Formatter -from tests.tests_eval_framework.tasks.benchmarks.utils import get_task_names_for_module, run_formatter_hash_test +from tests.tests_eval_framework.tasks.benchmarks.utils import run_formatter_hash_test + +# Registry for this test suite only holding csqa tasks +_csqa_registry = Registry() +register_csqa_tasks(registry=_csqa_registry) @pytest.mark.formatter_hash @pytest.mark.parametrize("formatter_cls", [Llama3Formatter, ConcatFormatter]) -@pytest.mark.parametrize("task_name", get_task_names_for_module("csqa")) +@pytest.mark.parametrize("task_name", _csqa_registry.task_names()) def test_formatter_hash(task_name: str, formatter_cls: type[BaseFormatter]) -> None: - run_formatter_hash_test(task_name, formatter_cls) + run_formatter_hash_test(task_name, formatter_cls, registry=_csqa_registry) diff --git a/tests/tests_eval_framework/tasks/benchmarks/test_drop.py b/tests/tests_eval_framework/tasks/benchmarks/test_drop.py index 3231de807..4f0298d9f 100644 --- a/tests/tests_eval_framework/tasks/benchmarks/test_drop.py +++ b/tests/tests_eval_framework/tasks/benchmarks/test_drop.py @@ -1,11 +1,17 @@ import pytest +from eval_framework.tasks.registry import Registry +from eval_framework.tasks.task_names import register_drop_tasks from template_formatting.formatter import BaseFormatter, ConcatFormatter, Llama3Formatter -from tests.tests_eval_framework.tasks.benchmarks.utils import get_task_names_for_module, run_formatter_hash_test +from tests.tests_eval_framework.tasks.benchmarks.utils import run_formatter_hash_test + +# Registry for this test suite only holding drop tasks +_drop_registry = Registry() +register_drop_tasks(registry=_drop_registry) @pytest.mark.formatter_hash @pytest.mark.parametrize("formatter_cls", [Llama3Formatter, ConcatFormatter]) -@pytest.mark.parametrize("task_name", get_task_names_for_module("drop")) +@pytest.mark.parametrize("task_name", _drop_registry.task_names()) def test_formatter_hash(task_name: str, formatter_cls: type[BaseFormatter]) -> None: - run_formatter_hash_test(task_name, formatter_cls) + run_formatter_hash_test(task_name, formatter_cls, registry=_drop_registry) diff --git a/tests/tests_eval_framework/tasks/benchmarks/test_global_mmlu.py b/tests/tests_eval_framework/tasks/benchmarks/test_global_mmlu.py index 75aed0e71..8ba330d38 100644 --- a/tests/tests_eval_framework/tasks/benchmarks/test_global_mmlu.py +++ b/tests/tests_eval_framework/tasks/benchmarks/test_global_mmlu.py @@ -1,11 +1,17 @@ import pytest +from eval_framework.tasks.registry import Registry +from eval_framework.tasks.task_names import register_global_mmlu_tasks from template_formatting.formatter import BaseFormatter, ConcatFormatter, Llama3Formatter -from tests.tests_eval_framework.tasks.benchmarks.utils import get_task_names_for_module, run_formatter_hash_test +from tests.tests_eval_framework.tasks.benchmarks.utils import run_formatter_hash_test + +# Registry for this test suite only holding global_mmlu tasks +_global_mmlu_registry = Registry() +register_global_mmlu_tasks(registry=_global_mmlu_registry) @pytest.mark.formatter_hash @pytest.mark.parametrize("formatter_cls", [Llama3Formatter, ConcatFormatter]) -@pytest.mark.parametrize("task_name", get_task_names_for_module("global_mmlu")) +@pytest.mark.parametrize("task_name", _global_mmlu_registry.task_names()) def test_formatter_hash(task_name: str, formatter_cls: type[BaseFormatter]) -> None: - run_formatter_hash_test(task_name, formatter_cls) + run_formatter_hash_test(task_name, formatter_cls, registry=_global_mmlu_registry) diff --git a/tests/tests_eval_framework/tasks/benchmarks/test_goldenswag.py b/tests/tests_eval_framework/tasks/benchmarks/test_goldenswag.py index 60026bd09..8d7bc92e2 100644 --- a/tests/tests_eval_framework/tasks/benchmarks/test_goldenswag.py +++ b/tests/tests_eval_framework/tasks/benchmarks/test_goldenswag.py @@ -1,11 +1,17 @@ import pytest +from eval_framework.tasks.registry import Registry +from eval_framework.tasks.task_names import register_goldenswag_tasks from template_formatting.formatter import BaseFormatter, ConcatFormatter, Llama3Formatter -from tests.tests_eval_framework.tasks.benchmarks.utils import get_task_names_for_module, run_formatter_hash_test +from tests.tests_eval_framework.tasks.benchmarks.utils import run_formatter_hash_test + +# Registry for this test suite only holding goldenswag tasks +_goldenswag_registry = Registry() +register_goldenswag_tasks(registry=_goldenswag_registry) @pytest.mark.formatter_hash @pytest.mark.parametrize("formatter_cls", [Llama3Formatter, ConcatFormatter]) -@pytest.mark.parametrize("task_name", get_task_names_for_module("goldenswag")) +@pytest.mark.parametrize("task_name", _goldenswag_registry.task_names()) def test_formatter_hash(task_name: str, formatter_cls: type[BaseFormatter]) -> None: - run_formatter_hash_test(task_name, formatter_cls) + run_formatter_hash_test(task_name, formatter_cls, registry=_goldenswag_registry) diff --git a/tests/tests_eval_framework/tasks/benchmarks/test_gpqa.py b/tests/tests_eval_framework/tasks/benchmarks/test_gpqa.py index ac0abe381..8ccc9d897 100644 --- a/tests/tests_eval_framework/tasks/benchmarks/test_gpqa.py +++ b/tests/tests_eval_framework/tasks/benchmarks/test_gpqa.py @@ -3,8 +3,10 @@ import pytest from eval_framework.tasks.benchmarks.gpqa import GPQA, GPQA_COT +from eval_framework.tasks.registry import Registry +from eval_framework.tasks.task_names import register_gpqa_tasks from template_formatting.formatter import BaseFormatter, ConcatFormatter, Llama3Formatter -from tests.tests_eval_framework.tasks.benchmarks.utils import get_task_names_for_module, run_formatter_hash_test +from tests.tests_eval_framework.tasks.benchmarks.utils import run_formatter_hash_test from tests.tests_eval_framework.utils import DatasetPatcher # GPQA_OLMES uses a gated HuggingFace dataset (Idavidrein/gpqa); hashes cannot be computed without auth. @@ -91,8 +93,13 @@ def test_ground_truth_in_completion_cot(self, gpqa_cot_task: GPQA_COT) -> None: assert len(ground_truths) == 1 +# Registry for this test suite only holding gpqa tasks +_gpqa_registry = Registry() +register_gpqa_tasks(registry=_gpqa_registry) + + @pytest.mark.formatter_hash @pytest.mark.parametrize("formatter_cls", [Llama3Formatter, ConcatFormatter]) -@pytest.mark.parametrize("task_name", get_task_names_for_module("gpqa", skip_tasks=_SKIPPED_TASKS)) +@pytest.mark.parametrize("task_name", [name for name in _gpqa_registry.task_names() if name not in _SKIPPED_TASKS]) def test_formatter_hash(task_name: str, formatter_cls: type[BaseFormatter]) -> None: - run_formatter_hash_test(task_name, formatter_cls) + run_formatter_hash_test(task_name, formatter_cls, registry=_gpqa_registry) diff --git a/tests/tests_eval_framework/tasks/benchmarks/test_gsm8k.py b/tests/tests_eval_framework/tasks/benchmarks/test_gsm8k.py index 7d7494b27..7a7804735 100644 --- a/tests/tests_eval_framework/tasks/benchmarks/test_gsm8k.py +++ b/tests/tests_eval_framework/tasks/benchmarks/test_gsm8k.py @@ -4,21 +4,28 @@ from datasets import Dataset, DatasetDict from eval_framework.tasks.benchmarks.gsm8k import GSM8KBPB +from eval_framework.tasks.registry import Registry +from eval_framework.tasks.task_names import register_gsm8k_tasks from template_formatting.formatter import BaseFormatter, ConcatFormatter, Llama3Formatter, Message, Role from tests.tests_eval_framework.tasks.benchmarks.utils import ( ExpectedPrompt, - get_task_names_for_module, run_formatter_hash_test, ) _NUM_FEWSHOT = {"GSM8K_OLMES": 8} +# Registry for this test suite only holding gsm8k tasks +_gsm8k_registry = Registry() +register_gsm8k_tasks(registry=_gsm8k_registry) + @pytest.mark.formatter_hash @pytest.mark.parametrize("formatter_cls", [Llama3Formatter, ConcatFormatter]) -@pytest.mark.parametrize("task_name", get_task_names_for_module("gsm8k")) +@pytest.mark.parametrize("task_name", _gsm8k_registry.task_names()) def test_formatter_hash(task_name: str, formatter_cls: type[BaseFormatter]) -> None: - run_formatter_hash_test(task_name, formatter_cls, num_fewshot=_NUM_FEWSHOT.get(task_name, 1)) + run_formatter_hash_test( + task_name, formatter_cls, num_fewshot=_NUM_FEWSHOT.get(task_name, 1), registry=_gsm8k_registry + ) _SUBJECT = "main" diff --git a/tests/tests_eval_framework/tasks/benchmarks/test_hellaswag.py b/tests/tests_eval_framework/tasks/benchmarks/test_hellaswag.py index 5021aaca8..d66813ca2 100644 --- a/tests/tests_eval_framework/tasks/benchmarks/test_hellaswag.py +++ b/tests/tests_eval_framework/tasks/benchmarks/test_hellaswag.py @@ -1,11 +1,17 @@ import pytest +from eval_framework.tasks.registry import Registry +from eval_framework.tasks.task_names import register_hellaswag_tasks from template_formatting.formatter import BaseFormatter, ConcatFormatter, Llama3Formatter -from tests.tests_eval_framework.tasks.benchmarks.utils import get_task_names_for_module, run_formatter_hash_test +from tests.tests_eval_framework.tasks.benchmarks.utils import run_formatter_hash_test + +# Registry for this test suite only holding hellaswag tasks +_hellaswag_registry = Registry() +register_hellaswag_tasks(registry=_hellaswag_registry) @pytest.mark.formatter_hash @pytest.mark.parametrize("formatter_cls", [Llama3Formatter, ConcatFormatter]) -@pytest.mark.parametrize("task_name", get_task_names_for_module("hellaswag")) +@pytest.mark.parametrize("task_name", _hellaswag_registry.task_names()) def test_formatter_hash(task_name: str, formatter_cls: type[BaseFormatter]) -> None: - run_formatter_hash_test(task_name, formatter_cls) + run_formatter_hash_test(task_name, formatter_cls, registry=_hellaswag_registry) diff --git a/tests/tests_eval_framework/tasks/benchmarks/test_humaneval.py b/tests/tests_eval_framework/tasks/benchmarks/test_humaneval.py index 0d7cce702..cf387151b 100644 --- a/tests/tests_eval_framework/tasks/benchmarks/test_humaneval.py +++ b/tests/tests_eval_framework/tasks/benchmarks/test_humaneval.py @@ -1,9 +1,11 @@ import pytest from eval_framework.tasks.benchmarks.humaneval import HumanEval, HumanEval_OLMES, HumanEvalInstruct +from eval_framework.tasks.registry import Registry +from eval_framework.tasks.task_names import register_humaneval_tasks from eval_framework.tasks.utils import run_python_code from template_formatting.formatter import BaseFormatter, ConcatFormatter, Llama3Formatter -from tests.tests_eval_framework.tasks.benchmarks.utils import get_task_names_for_module, run_formatter_hash_test +from tests.tests_eval_framework.tasks.benchmarks.utils import run_formatter_hash_test from tests.tests_eval_framework.utils import DatasetPatcher _NUM_FEWSHOT = {"HumanEval_OLMES": 3} @@ -89,8 +91,15 @@ def test_code_is_executed(self, human_eval_task_inst: HumanEvalInstruct) -> None assert i == 9 +# Registry for this test suite only holding humaneval tasks +_humaneval_registry = Registry() +register_humaneval_tasks(registry=_humaneval_registry) + + @pytest.mark.formatter_hash @pytest.mark.parametrize("formatter_cls", [Llama3Formatter, ConcatFormatter]) -@pytest.mark.parametrize("task_name", get_task_names_for_module("humaneval")) +@pytest.mark.parametrize("task_name", _humaneval_registry.task_names()) def test_formatter_hash(task_name: str, formatter_cls: type[BaseFormatter]) -> None: - run_formatter_hash_test(task_name, formatter_cls, num_fewshot=_NUM_FEWSHOT.get(task_name, 1)) + run_formatter_hash_test( + task_name, formatter_cls, num_fewshot=_NUM_FEWSHOT.get(task_name, 1), registry=_humaneval_registry + ) diff --git a/tests/tests_eval_framework/tasks/benchmarks/test_ifeval.py b/tests/tests_eval_framework/tasks/benchmarks/test_ifeval.py index 62ea1f8de..2d99581d7 100644 --- a/tests/tests_eval_framework/tasks/benchmarks/test_ifeval.py +++ b/tests/tests_eval_framework/tasks/benchmarks/test_ifeval.py @@ -1,11 +1,17 @@ import pytest +from eval_framework.tasks.registry import Registry +from eval_framework.tasks.task_names import register_ifeval_tasks from template_formatting.formatter import BaseFormatter, ConcatFormatter, Llama3Formatter -from tests.tests_eval_framework.tasks.benchmarks.utils import get_task_names_for_module, run_formatter_hash_test +from tests.tests_eval_framework.tasks.benchmarks.utils import run_formatter_hash_test + +# Registry for this test suite only holding ifeval tasks +_ifeval_registry = Registry() +register_ifeval_tasks(registry=_ifeval_registry) @pytest.mark.formatter_hash @pytest.mark.parametrize("formatter_cls", [Llama3Formatter, ConcatFormatter]) -@pytest.mark.parametrize("task_name", get_task_names_for_module("ifeval")) +@pytest.mark.parametrize("task_name", _ifeval_registry.task_names()) def test_formatter_hash(task_name: str, formatter_cls: type[BaseFormatter]) -> None: - run_formatter_hash_test(task_name, formatter_cls) + run_formatter_hash_test(task_name, formatter_cls, registry=_ifeval_registry) diff --git a/tests/tests_eval_framework/tasks/benchmarks/test_math_reasoning.py b/tests/tests_eval_framework/tasks/benchmarks/test_math_reasoning.py index e5f8f469e..8af60463a 100644 --- a/tests/tests_eval_framework/tasks/benchmarks/test_math_reasoning.py +++ b/tests/tests_eval_framework/tasks/benchmarks/test_math_reasoning.py @@ -4,12 +4,13 @@ from datasets import Dataset, DatasetDict from eval_framework.tasks.benchmarks.math_reasoning import MATH, MATHMinervaBPB +from eval_framework.tasks.registry import Registry +from eval_framework.tasks.task_names import register_math_reasoning_tasks from template_formatting.formatter import BaseFormatter, ConcatFormatter, Llama3Formatter, Message, Role from tests.tests_eval_framework.tasks.benchmarks.utils import ( ExpectedPrompt, assert_offline_oneshot_prompt, assert_offline_zeroshot_prompt, - get_task_names_for_module, run_formatter_hash_test, ) from tests.tests_eval_framework.utils import DatasetPatcher @@ -111,11 +112,18 @@ def test_split_text_command_with_search( assert math_reasoning._split_text_command(string, search) == expected +# Registry for this test suite only holding math_reasoning tasks +_math_reasoning_registry = Registry() +register_math_reasoning_tasks(registry=_math_reasoning_registry) + + @pytest.mark.formatter_hash @pytest.mark.parametrize("formatter_cls", [Llama3Formatter, ConcatFormatter]) -@pytest.mark.parametrize("task_name", get_task_names_for_module("math_reasoning")) +@pytest.mark.parametrize("task_name", _math_reasoning_registry.task_names()) def test_formatter_hash(task_name: str, formatter_cls: type[BaseFormatter]) -> None: - run_formatter_hash_test(task_name, formatter_cls, num_fewshot=_NUM_FEWSHOT.get(task_name, 1)) + run_formatter_hash_test( + task_name, formatter_cls, num_fewshot=_NUM_FEWSHOT.get(task_name, 1), registry=_math_reasoning_registry + ) # --------------------------------------------------------------------------- diff --git a/tests/tests_eval_framework/tasks/benchmarks/test_mbpp.py b/tests/tests_eval_framework/tasks/benchmarks/test_mbpp.py index 8de1e28fa..d598eed80 100644 --- a/tests/tests_eval_framework/tasks/benchmarks/test_mbpp.py +++ b/tests/tests_eval_framework/tasks/benchmarks/test_mbpp.py @@ -1,23 +1,30 @@ import pytest from eval_framework.tasks.benchmarks.mbpp import MBPPBPB +from eval_framework.tasks.registry import Registry +from eval_framework.tasks.task_names import register_mbpp_tasks from template_formatting.formatter import BaseFormatter, ConcatFormatter, Llama3Formatter, Message, Role from tests.tests_eval_framework.tasks.benchmarks.utils import ( ExpectedPrompt, assert_offline_oneshot_prompt, assert_offline_zeroshot_prompt, - get_task_names_for_module, run_formatter_hash_test, ) _NUM_FEWSHOT = {"MBPP_OLMES": 3} +# Registry for this test suite only holding mbpp tasks +_mbpp_registry = Registry() +register_mbpp_tasks(registry=_mbpp_registry) + @pytest.mark.formatter_hash @pytest.mark.parametrize("formatter_cls", [Llama3Formatter, ConcatFormatter]) -@pytest.mark.parametrize("task_name", get_task_names_for_module("mbpp")) +@pytest.mark.parametrize("task_name", _mbpp_registry.task_names()) def test_formatter_hash(task_name: str, formatter_cls: type[BaseFormatter]) -> None: - run_formatter_hash_test(task_name, formatter_cls, num_fewshot=_NUM_FEWSHOT.get(task_name, 1)) + run_formatter_hash_test( + task_name, formatter_cls, num_fewshot=_NUM_FEWSHOT.get(task_name, 1), registry=_mbpp_registry + ) # --------------------------------------------------------------------------- diff --git a/tests/tests_eval_framework/tasks/benchmarks/test_medqa.py b/tests/tests_eval_framework/tasks/benchmarks/test_medqa.py index b30e66afb..1266f78c3 100644 --- a/tests/tests_eval_framework/tasks/benchmarks/test_medqa.py +++ b/tests/tests_eval_framework/tasks/benchmarks/test_medqa.py @@ -1,11 +1,17 @@ import pytest +from eval_framework.tasks.registry import Registry +from eval_framework.tasks.task_names import register_medqa_tasks from template_formatting.formatter import BaseFormatter, ConcatFormatter, Llama3Formatter -from tests.tests_eval_framework.tasks.benchmarks.utils import get_task_names_for_module, run_formatter_hash_test +from tests.tests_eval_framework.tasks.benchmarks.utils import run_formatter_hash_test + +# Registry for this test suite only holding medqa tasks +_medqa_registry = Registry() +register_medqa_tasks(registry=_medqa_registry) @pytest.mark.formatter_hash @pytest.mark.parametrize("formatter_cls", [Llama3Formatter, ConcatFormatter]) -@pytest.mark.parametrize("task_name", get_task_names_for_module("medqa")) +@pytest.mark.parametrize("task_name", _medqa_registry.task_names()) def test_formatter_hash(task_name: str, formatter_cls: type[BaseFormatter]) -> None: - run_formatter_hash_test(task_name, formatter_cls) + run_formatter_hash_test(task_name, formatter_cls, registry=_medqa_registry) diff --git a/tests/tests_eval_framework/tasks/benchmarks/test_mmlu.py b/tests/tests_eval_framework/tasks/benchmarks/test_mmlu.py index 2b7542098..e6bd0db71 100644 --- a/tests/tests_eval_framework/tasks/benchmarks/test_mmlu.py +++ b/tests/tests_eval_framework/tasks/benchmarks/test_mmlu.py @@ -1,11 +1,17 @@ import pytest +from eval_framework.tasks.registry import Registry +from eval_framework.tasks.task_names import register_mmlu_tasks from template_formatting.formatter import BaseFormatter, ConcatFormatter, Llama3Formatter -from tests.tests_eval_framework.tasks.benchmarks.utils import get_task_names_for_module, run_formatter_hash_test +from tests.tests_eval_framework.tasks.benchmarks.utils import run_formatter_hash_test + +# Registry for this test suite only holding mmlu tasks +_mmlu_registry = Registry() +register_mmlu_tasks(registry=_mmlu_registry) @pytest.mark.formatter_hash @pytest.mark.parametrize("formatter_cls", [Llama3Formatter, ConcatFormatter]) -@pytest.mark.parametrize("task_name", get_task_names_for_module("mmlu")) +@pytest.mark.parametrize("task_name", _mmlu_registry.task_names()) def test_formatter_hash(task_name: str, formatter_cls: type[BaseFormatter]) -> None: - run_formatter_hash_test(task_name, formatter_cls) + run_formatter_hash_test(task_name, formatter_cls, registry=_mmlu_registry) diff --git a/tests/tests_eval_framework/tasks/benchmarks/test_mmlu_pro.py b/tests/tests_eval_framework/tasks/benchmarks/test_mmlu_pro.py index d682c7050..a2c2cbdb3 100644 --- a/tests/tests_eval_framework/tasks/benchmarks/test_mmlu_pro.py +++ b/tests/tests_eval_framework/tasks/benchmarks/test_mmlu_pro.py @@ -1,11 +1,17 @@ import pytest +from eval_framework.tasks.registry import Registry +from eval_framework.tasks.task_names import register_mmlu_pro_tasks from template_formatting.formatter import BaseFormatter, ConcatFormatter, Llama3Formatter -from tests.tests_eval_framework.tasks.benchmarks.utils import get_task_names_for_module, run_formatter_hash_test +from tests.tests_eval_framework.tasks.benchmarks.utils import run_formatter_hash_test + +# Registry for this test suite only holding mmlu_pro tasks +_mmlu_pro_registry = Registry() +register_mmlu_pro_tasks(registry=_mmlu_pro_registry) @pytest.mark.formatter_hash @pytest.mark.parametrize("formatter_cls", [Llama3Formatter, ConcatFormatter]) -@pytest.mark.parametrize("task_name", get_task_names_for_module("mmlu_pro")) +@pytest.mark.parametrize("task_name", _mmlu_pro_registry.task_names()) def test_formatter_hash(task_name: str, formatter_cls: type[BaseFormatter]) -> None: - run_formatter_hash_test(task_name, formatter_cls) + run_formatter_hash_test(task_name, formatter_cls, registry=_mmlu_pro_registry) diff --git a/tests/tests_eval_framework/tasks/benchmarks/test_multipl_e.py b/tests/tests_eval_framework/tasks/benchmarks/test_multipl_e.py index e24ebb935..0c046760e 100644 --- a/tests/tests_eval_framework/tasks/benchmarks/test_multipl_e.py +++ b/tests/tests_eval_framework/tasks/benchmarks/test_multipl_e.py @@ -1,11 +1,17 @@ import pytest +from eval_framework.tasks.registry import Registry +from eval_framework.tasks.task_names import register_multipl_e_tasks from template_formatting.formatter import BaseFormatter, ConcatFormatter, Llama3Formatter -from tests.tests_eval_framework.tasks.benchmarks.utils import get_task_names_for_module, run_formatter_hash_test +from tests.tests_eval_framework.tasks.benchmarks.utils import run_formatter_hash_test + +# Registry for this test suite only holding multipl_e tasks +_multipl_e_registry = Registry() +register_multipl_e_tasks(registry=_multipl_e_registry) @pytest.mark.formatter_hash @pytest.mark.parametrize("formatter_cls", [Llama3Formatter, ConcatFormatter]) -@pytest.mark.parametrize("task_name", get_task_names_for_module("multipl_e")) +@pytest.mark.parametrize("task_name", _multipl_e_registry.task_names()) def test_formatter_hash(task_name: str, formatter_cls: type[BaseFormatter]) -> None: - run_formatter_hash_test(task_name, formatter_cls) + run_formatter_hash_test(task_name, formatter_cls, registry=_multipl_e_registry) diff --git a/tests/tests_eval_framework/tasks/benchmarks/test_naturalqs_open.py b/tests/tests_eval_framework/tasks/benchmarks/test_naturalqs_open.py index e3d1564ef..bd07b02ec 100644 --- a/tests/tests_eval_framework/tasks/benchmarks/test_naturalqs_open.py +++ b/tests/tests_eval_framework/tasks/benchmarks/test_naturalqs_open.py @@ -1,11 +1,17 @@ import pytest +from eval_framework.tasks.registry import Registry +from eval_framework.tasks.task_names import register_naturalqs_open_tasks from template_formatting.formatter import BaseFormatter, ConcatFormatter, Llama3Formatter -from tests.tests_eval_framework.tasks.benchmarks.utils import get_task_names_for_module, run_formatter_hash_test +from tests.tests_eval_framework.tasks.benchmarks.utils import run_formatter_hash_test + +# Registry for this test suite only holding naturalqs_open tasks +_naturalqs_open_registry = Registry() +register_naturalqs_open_tasks(registry=_naturalqs_open_registry) @pytest.mark.formatter_hash @pytest.mark.parametrize("formatter_cls", [Llama3Formatter, ConcatFormatter]) -@pytest.mark.parametrize("task_name", get_task_names_for_module("naturalqs_open")) +@pytest.mark.parametrize("task_name", _naturalqs_open_registry.task_names()) def test_formatter_hash(task_name: str, formatter_cls: type[BaseFormatter]) -> None: - run_formatter_hash_test(task_name, formatter_cls) + run_formatter_hash_test(task_name, formatter_cls, registry=_naturalqs_open_registry) diff --git a/tests/tests_eval_framework/tasks/benchmarks/test_piqa.py b/tests/tests_eval_framework/tasks/benchmarks/test_piqa.py index 3e822129b..6034cc8b4 100644 --- a/tests/tests_eval_framework/tasks/benchmarks/test_piqa.py +++ b/tests/tests_eval_framework/tasks/benchmarks/test_piqa.py @@ -1,11 +1,17 @@ import pytest +from eval_framework.tasks.registry import Registry +from eval_framework.tasks.task_names import register_piqa_tasks from template_formatting.formatter import BaseFormatter, ConcatFormatter, Llama3Formatter -from tests.tests_eval_framework.tasks.benchmarks.utils import get_task_names_for_module, run_formatter_hash_test +from tests.tests_eval_framework.tasks.benchmarks.utils import run_formatter_hash_test + +# Registry for this test suite only holding piqa tasks +_piqa_registry = Registry() +register_piqa_tasks(registry=_piqa_registry) @pytest.mark.formatter_hash @pytest.mark.parametrize("formatter_cls", [Llama3Formatter, ConcatFormatter]) -@pytest.mark.parametrize("task_name", get_task_names_for_module("piqa")) +@pytest.mark.parametrize("task_name", _piqa_registry.task_names()) def test_formatter_hash(task_name: str, formatter_cls: type[BaseFormatter]) -> None: - run_formatter_hash_test(task_name, formatter_cls) + run_formatter_hash_test(task_name, formatter_cls, registry=_piqa_registry) diff --git a/tests/tests_eval_framework/tasks/benchmarks/test_sciq.py b/tests/tests_eval_framework/tasks/benchmarks/test_sciq.py index d2d914165..19e789c10 100644 --- a/tests/tests_eval_framework/tasks/benchmarks/test_sciq.py +++ b/tests/tests_eval_framework/tasks/benchmarks/test_sciq.py @@ -1,11 +1,17 @@ import pytest +from eval_framework.tasks.registry import Registry +from eval_framework.tasks.task_names import register_sciq_tasks from template_formatting.formatter import BaseFormatter, ConcatFormatter, Llama3Formatter -from tests.tests_eval_framework.tasks.benchmarks.utils import get_task_names_for_module, run_formatter_hash_test +from tests.tests_eval_framework.tasks.benchmarks.utils import run_formatter_hash_test + +# Registry for this test suite only holding sciq tasks +_sciq_registry = Registry() +register_sciq_tasks(registry=_sciq_registry) @pytest.mark.formatter_hash @pytest.mark.parametrize("formatter_cls", [Llama3Formatter, ConcatFormatter]) -@pytest.mark.parametrize("task_name", get_task_names_for_module("sciq")) +@pytest.mark.parametrize("task_name", _sciq_registry.task_names()) def test_formatter_hash(task_name: str, formatter_cls: type[BaseFormatter]) -> None: - run_formatter_hash_test(task_name, formatter_cls) + run_formatter_hash_test(task_name, formatter_cls, registry=_sciq_registry) diff --git a/tests/tests_eval_framework/tasks/benchmarks/test_social_iqa.py b/tests/tests_eval_framework/tasks/benchmarks/test_social_iqa.py index 1ec4d824e..8e30d89cf 100644 --- a/tests/tests_eval_framework/tasks/benchmarks/test_social_iqa.py +++ b/tests/tests_eval_framework/tasks/benchmarks/test_social_iqa.py @@ -1,11 +1,17 @@ import pytest +from eval_framework.tasks.registry import Registry +from eval_framework.tasks.task_names import register_social_iqa_tasks from template_formatting.formatter import BaseFormatter, ConcatFormatter, Llama3Formatter -from tests.tests_eval_framework.tasks.benchmarks.utils import get_task_names_for_module, run_formatter_hash_test +from tests.tests_eval_framework.tasks.benchmarks.utils import run_formatter_hash_test + +# Registry for this test suite only holding social_iqa tasks +_social_iqa_registry = Registry() +register_social_iqa_tasks(registry=_social_iqa_registry) @pytest.mark.formatter_hash @pytest.mark.parametrize("formatter_cls", [Llama3Formatter, ConcatFormatter]) -@pytest.mark.parametrize("task_name", get_task_names_for_module("social_iqa")) +@pytest.mark.parametrize("task_name", _social_iqa_registry.task_names()) def test_formatter_hash(task_name: str, formatter_cls: type[BaseFormatter]) -> None: - run_formatter_hash_test(task_name, formatter_cls) + run_formatter_hash_test(task_name, formatter_cls, registry=_social_iqa_registry) diff --git a/tests/tests_eval_framework/tasks/benchmarks/test_squad.py b/tests/tests_eval_framework/tasks/benchmarks/test_squad.py index b667f1054..e42ac9c48 100644 --- a/tests/tests_eval_framework/tasks/benchmarks/test_squad.py +++ b/tests/tests_eval_framework/tasks/benchmarks/test_squad.py @@ -1,15 +1,21 @@ import pytest from eval_framework.tasks.benchmarks.squad import SQuAD2_MA +from eval_framework.tasks.registry import Registry +from eval_framework.tasks.task_names import register_squad_tasks from template_formatting.formatter import BaseFormatter, ConcatFormatter, Llama3Formatter -from tests.tests_eval_framework.tasks.benchmarks.utils import get_task_names_for_module, run_formatter_hash_test +from tests.tests_eval_framework.tasks.benchmarks.utils import run_formatter_hash_test + +# Registry for this test suite only holding squad tasks +_squad_registry = Registry() +register_squad_tasks(registry=_squad_registry) @pytest.mark.formatter_hash @pytest.mark.parametrize("formatter_cls", [Llama3Formatter, ConcatFormatter]) -@pytest.mark.parametrize("task_name", get_task_names_for_module("squad")) +@pytest.mark.parametrize("task_name", _squad_registry.task_names()) def test_formatter_hash(task_name: str, formatter_cls: type[BaseFormatter]) -> None: - run_formatter_hash_test(task_name, formatter_cls) + run_formatter_hash_test(task_name, formatter_cls, registry=_squad_registry) _ANSWERABLE = { diff --git a/tests/tests_eval_framework/tasks/benchmarks/test_winogrande.py b/tests/tests_eval_framework/tasks/benchmarks/test_winogrande.py index 6df791c73..eba36eaa3 100644 --- a/tests/tests_eval_framework/tasks/benchmarks/test_winogrande.py +++ b/tests/tests_eval_framework/tasks/benchmarks/test_winogrande.py @@ -1,11 +1,17 @@ import pytest +from eval_framework.tasks.registry import Registry +from eval_framework.tasks.task_names import register_winogrande_tasks from template_formatting.formatter import BaseFormatter, ConcatFormatter, Llama3Formatter -from tests.tests_eval_framework.tasks.benchmarks.utils import get_task_names_for_module, run_formatter_hash_test +from tests.tests_eval_framework.tasks.benchmarks.utils import run_formatter_hash_test + +# Registry for this test suite only holding winogrande tasks +_winogrande_registry = Registry() +register_winogrande_tasks(registry=_winogrande_registry) @pytest.mark.formatter_hash @pytest.mark.parametrize("formatter_cls", [Llama3Formatter, ConcatFormatter]) -@pytest.mark.parametrize("task_name", get_task_names_for_module("winogrande")) +@pytest.mark.parametrize("task_name", _winogrande_registry.task_names()) def test_formatter_hash(task_name: str, formatter_cls: type[BaseFormatter]) -> None: - run_formatter_hash_test(task_name, formatter_cls) + run_formatter_hash_test(task_name, formatter_cls, registry=_winogrande_registry) diff --git a/tests/tests_eval_framework/tasks/benchmarks/utils.py b/tests/tests_eval_framework/tasks/benchmarks/utils.py index f24a263e2..f3acb110e 100644 --- a/tests/tests_eval_framework/tasks/benchmarks/utils.py +++ b/tests/tests_eval_framework/tasks/benchmarks/utils.py @@ -9,36 +9,12 @@ from datasets import Dataset, DatasetDict from eval_framework.tasks.base import BaseTask, Sample -from eval_framework.tasks.registry import ( - _REGISTRY, - registered_task_names, - registry, -) +from eval_framework.tasks.registry import Registry +from eval_framework.tasks.registry import registry as global_registry from template_formatting.formatter import BaseFormatter, ConcatFormatter, Message from tests.tests_eval_framework.utils import assert_hash_string -def _module_of_registered_task(task_name: str) -> str: - task_key = _REGISTRY._task_key(task_name) - _, factory = _REGISTRY._registry[task_key] - return factory.source_module - - -def get_task_names_for_module(module_name: str, skip_tasks: list[str] | None = None) -> list[str]: - """Return registered eval-framework task names declared in a given benchmark module. - - Mirrors `eval_framework_companion.tests.tasks.benchmarks.utils.get_task_names_for_module` - so per-benchmark test files can parametrize over just the tasks they own. - """ - target_module = f"eval_framework.tasks.benchmarks.{module_name}" - skip = set(skip_tasks or []) - return sorted( - name - for name in registered_task_names() - if _module_of_registered_task(name) == target_module and name not in skip - ) - - def _seed_for_determinism() -> None: random.seed(42) try: @@ -55,16 +31,18 @@ def _seed_for_determinism() -> None: pass -def run_formatter_hash_test(task_name: str, formatter_cls: type[BaseFormatter], num_fewshot: int = 1) -> None: +def run_formatter_hash_test( + task_name: str, formatter_cls: type[BaseFormatter], num_fewshot: int = 1, registry: Registry | None = None +) -> None: """Run the formatter hash consistency test for a single task x formatter combination. - Uses the full HuggingFace datasets with seed 42 and a deterministic few-shot sampler, - matching the prior `test_all_formatters.py` behaviour so existing hashes remain valid. + Uses the full HuggingFace datasets with seed 42 and a deterministic few-shot sampler. """ _seed_for_determinism() + registry = registry if registry is not None else global_registry() def _instantiate(num_fewshot_value: int) -> object: - task_instance = registry()[task_name].create( + task_instance = registry[task_name].create( num_fewshot=num_fewshot_value, custom_subjects=None, custom_hf_revision=None,