From 6b69d7769d326cdef75448a70b91c3f6ec23b35a Mon Sep 17 00:00:00 2001 From: Abelo9996 Date: Thu, 27 Aug 2026 16:36:01 -0400 Subject: [PATCH] fix: honor task-level num_samples and warn when it has no effect LightevalTaskConfig exposes a num_samples field, but LightevalTask.__init__ reset self.num_samples to [1] and only extended it from sampling metrics, so a task-level num_samples was silently dropped and generation stayed at one sample. Seed self.num_samples from config.num_samples (accepting a list or a bare int) so the documented field is respected, and warn when num_samples > 1 is requested but no attached metric consumes multiple samples, instead of silently generating extra samples that never reach a score. Refs #618 --- src/lighteval/tasks/lighteval_task.py | 18 ++++++++++- tests/unit/tasks/test_lighteval_task.py | 43 +++++++++++++++++++++++++ 2 files changed, 60 insertions(+), 1 deletion(-) diff --git a/src/lighteval/tasks/lighteval_task.py b/src/lighteval/tasks/lighteval_task.py index 5e9bac215..598e28b78 100644 --- a/src/lighteval/tasks/lighteval_task.py +++ b/src/lighteval/tasks/lighteval_task.py @@ -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 diff --git a/tests/unit/tasks/test_lighteval_task.py b/tests/unit/tasks/test_lighteval_task.py index 7cdb7b6f5..e50c92968 100644 --- a/tests/unit/tasks/test_lighteval_task.py +++ b/tests/unit/tasks/test_lighteval_task.py @@ -21,6 +21,8 @@ # SOFTWARE. +import logging + from lighteval.tasks.lighteval_task import LightevalTask, LightevalTaskConfig from lighteval.tasks.requests import Doc @@ -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)