File size: 3,032 Bytes
31dc8dc
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
"""
Benchmark Datasets - Dataset loaders for benchmark evaluation
Supports common evaluation datasets such as GSM8K, HumanEval, etc.
"""

from typing import List, Dict, Any, Optional, Callable
from datasets import load_dataset


def load_gsm8k(
    split: str = "test",
    limit: Optional[int] = None,
    prompt_template: Optional[Callable[[str], str]] = None,
) -> List[Dict[str, Any]]:
    """
    Load GSM8K dataset

    Args:
        split: Dataset split, default "test"
        limit: Limit number of samples, None means all
        prompt_template: Prompt template function that takes question string and returns full prompt

    Returns:
        List of dataset items, each containing 'prompt' and 'answer' fields
    """
    dataset = load_dataset("gsm8k", "main", split=split)

    if limit:
        dataset = dataset[:limit]

    results = []
    for item in dataset:
        question = item["question"]
        answer = item["answer"]

        if prompt_template:
            prompt = prompt_template(question)
        else:
            # Default template
            prompt = f"Question: {question}\nAnswer:"

        results.append(
            {
                "prompt": prompt,
                "answer": answer,
                "question": question,
            }
        )

    return results


def load_humaneval(
    limit: Optional[int] = None,
    prompt_template: Optional[Callable[[str], str]] = None,
) -> List[Dict[str, Any]]:
    """
    Load HumanEval dataset

    Args:
        limit: Limit number of samples, None means all
        prompt_template: Prompt template function that takes prompt string and returns full prompt

    Returns:
        List of dataset items, each containing 'prompt', 'test', 'entry_point' fields
    """
    dataset = load_dataset("openai/humaneval", split="test")

    if limit:
        dataset = dataset[:limit]

    results = []
    for item in dataset:
        prompt = item["prompt"]
        test = item["test"]
        entry_point = item["entry_point"]

        if prompt_template:
            full_prompt = prompt_template(prompt)
        else:
            full_prompt = prompt

        results.append(
            {
                "prompt": full_prompt,
                "original_prompt": prompt,
                "test": test,
                "entry_point": entry_point,
                "task_id": item.get("task_id", ""),
            }
        )

    return results


def load_benchmark_dataset(dataset_name: str, **kwargs) -> List[Dict[str, Any]]:
    """
    Unified dataset loading interface

    Args:
        dataset_name: Dataset name, supports "gsm8k", "humaneval"
        **kwargs: Arguments passed to the specific dataset loader

    Returns:
        List of dataset items
    """
    loaders = {
        "gsm8k": load_gsm8k,
        "humaneval": load_humaneval,
    }

    if dataset_name not in loaders:
        raise ValueError(f"Unknown dataset: {dataset_name}. Supported datasets: {list(loaders.keys())}")

    return loaders[dataset_name](**kwargs)