Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
18 changes: 17 additions & 1 deletion src/lighteval/tasks/lighteval_task.py
Original file line number Diff line number Diff line change
Expand Up @@ -249,13 +249,29 @@ def __init__(
self.generation_grammar = config.generation_grammar
self.stop_sequence = config.stop_sequence

# We assume num_samples always contains 1 (for base generative evals)
# num_samples always contains 1 (for base generative evals). On top of
# that, honor an explicit task-level override from the config and the
# requirements of any sampling metric. The request builder later takes
# the max, so the largest requested value wins.
self.num_samples = [1]
if config.num_samples is not None:
requested = config.num_samples if isinstance(config.num_samples, (list, tuple)) else [config.num_samples]
self.num_samples.extend(int(n) for n in requested)
for metric in self.metrics:
if isinstance(metric.sample_level_fn, SamplingMetric):
# Update the number of samples to generate using the information in the metric name
self.num_samples.append(metric.sample_level_fn.num_samples())

if max(self.num_samples) > 1 and not any(
isinstance(metric.sample_level_fn, SamplingMetric) for metric in self.metrics
):
logger.warning(
f"Task {self.name}: num_samples > 1 is set but no metric consumes "
"multiple samples, so the extra generations will not affect any "
"score. Attach a sampling metric (for example pass@k or maj@k) or "
"set num_samples to 1."
)

def get_first_possible_fewshot_splits(self, available_splits: ListLike[str]) -> str | None:
"""Parses the possible fewshot split keys in order: train, then validation
keys and matches them with the available keys. Returns the first
Expand Down
43 changes: 43 additions & 0 deletions tests/unit/tasks/test_lighteval_task.py
Original file line number Diff line number Diff line change
Expand Up @@ -21,6 +21,8 @@
# SOFTWARE.


import logging

from lighteval.tasks.lighteval_task import LightevalTask, LightevalTaskConfig
from lighteval.tasks.requests import Doc

Expand Down Expand Up @@ -84,3 +86,44 @@ def test_hf_data_files(tmp_path):

eval_docs = task.eval_docs()
assert [doc.query for doc in eval_docs] == src_docs


def _num_samples_config(num_samples, metrics=None):
return LightevalTaskConfig(
name="test_num_samples",
prompt_function=dummy_prompt_function,
hf_repo="lighteval-tests-datasets/dataset-test-1",
hf_subset="default",
evaluation_splits=["train"],
metrics=metrics or [],
num_samples=num_samples,
)


def test_num_samples_defaults_to_one_when_unset():
cfg = _num_samples_config(num_samples=None)
task = LightevalTask(cfg)
assert task.num_samples == [1]


def test_num_samples_config_is_respected():
# Regression test for the config-level `num_samples` being silently dropped:
# the request builder takes max(num_samples), so the override must survive.
cfg = _num_samples_config(num_samples=[16])
task = LightevalTask(cfg)
assert max(task.num_samples) == 16
# 1 stays present so base generative scoring still works.
assert 1 in task.num_samples


def test_num_samples_config_accepts_bare_int():
cfg = _num_samples_config(num_samples=8)
task = LightevalTask(cfg)
assert max(task.num_samples) == 8


def test_num_samples_without_sampling_metric_warns(caplog):
cfg = _num_samples_config(num_samples=[16])
with caplog.at_level(logging.WARNING):
LightevalTask(cfg)
assert any("num_samples" in record.message for record in caplog.records)