Skip to content

Commit 6379678

Browse files
committed
address code review feedback
1 parent e4dd3ac commit 6379678

4 files changed

Lines changed: 69 additions & 11 deletions

File tree

‎src/inspect_ai/_cli/eval.py‎

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -256,7 +256,7 @@ def eval_options(func: Callable[..., Any]) -> Callable[..., click.Context]:
256256
)
257257
@click.option(
258258
"--max-dataset-memory",
259-
type=int,
259+
type=click.IntRange(min=0),
260260
help="Maximum MB of dataset sample data to hold in memory per task. When exceeded, samples are paged to disk.",
261261
envvar="INSPECT_EVAL_MAX_DATASET_MEMORY",
262262
)

‎src/inspect_ai/_eval/task/run.py‎

Lines changed: 7 additions & 5 deletions
Original file line numberDiff line numberDiff line change
@@ -243,6 +243,10 @@ async def task_run(options: TaskRunOptions) -> EvalLog:
243243
# optionally page dataset to disk if it exceeds the memory budget
244244
sample_store = maybe_page_to_disk(dataset, config.max_dataset_memory)
245245

246+
# release in-memory samples now that they're paged to disk
247+
if sample_store is not dataset:
248+
del dataset
249+
246250
# resolve the plan (unroll chains)
247251
solver = solver or task.solver
248252
plan = resolve_plan(task, solver)
@@ -298,14 +302,12 @@ async def task_run(options: TaskRunOptions) -> EvalLog:
298302
if options.task.early_stopping is not None:
299303
stopping_manager = await options.task.early_stopping.start_task(
300304
logger.eval,
301-
samples=[deepcopy(s) for s in dataset],
305+
samples=[
306+
deepcopy(sample_store[i]) for i in range(len(sample_store))
307+
],
302308
epochs=epochs,
303309
)
304310

305-
# now safe to release in-memory samples
306-
if sample_store is not dataset:
307-
del dataset
308-
309311
with td.progress() as p:
310312
# forward progress
311313
def progress(number: int) -> None:

‎src/inspect_ai/_eval/task/store.py‎

Lines changed: 12 additions & 5 deletions
Original file line numberDiff line numberDiff line change
@@ -15,11 +15,18 @@ class DiskSampleStore:
1515
def __init__(self, samples: Sequence[Sample]) -> None:
1616
self._len = len(samples)
1717
fd, self._path = tempfile.mkstemp(suffix=".pkl")
18-
with os.fdopen(fd, "wb") as f:
19-
self._offsets: list[int] = []
20-
for sample in samples:
21-
self._offsets.append(f.tell())
22-
pickle.dump(sample, f)
18+
try:
19+
with os.fdopen(fd, "wb") as f:
20+
self._offsets: list[int] = []
21+
for sample in samples:
22+
self._offsets.append(f.tell())
23+
pickle.dump(sample, f)
24+
except Exception:
25+
try:
26+
os.unlink(self._path)
27+
except OSError:
28+
pass
29+
raise
2330
self._reader: BinaryIO | None = None
2431

2532
def __len__(self) -> int:

‎tests/test_disk_sample_store.py‎

Lines changed: 49 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -1,4 +1,7 @@
11
import os
2+
from typing import TYPE_CHECKING
3+
4+
from pydantic import JsonValue
25

36
from inspect_ai import Task, eval
47
from inspect_ai._eval.task.store import (
@@ -10,6 +13,11 @@
1013
from inspect_ai.dataset._dataset import MemoryDataset
1114
from inspect_ai.scorer import match
1215
from inspect_ai.solver._solver import generate
16+
from inspect_ai.util._early_stopping import EarlyStop
17+
18+
if TYPE_CHECKING:
19+
from inspect_ai.log._log import EvalSpec
20+
from inspect_ai.scorer._metric import SampleScore
1321

1422

1523
def _make_samples(n: int = 5) -> list[Sample]:
@@ -109,3 +117,44 @@ def test_eval_with_max_dataset_memory() -> None:
109117
assert log.status == "success"
110118
assert log.samples is not None
111119
assert len(log.samples) == 3
120+
121+
122+
# -- Part 5: Early stopping + disk paging integration test --
123+
124+
125+
class _NoopEarlyStopping:
126+
"""Minimal early stopping that never stops anything."""
127+
128+
async def start_task(
129+
self, task: "EvalSpec", samples: list[Sample], epochs: int
130+
) -> str:
131+
return "noop"
132+
133+
async def schedule_sample(self, id: str | int, epoch: int) -> EarlyStop | None:
134+
return None
135+
136+
async def complete_sample(
137+
self,
138+
id: str | int,
139+
epoch: int,
140+
scores: dict[str, "SampleScore"],
141+
) -> None:
142+
pass
143+
144+
async def complete_task(self) -> dict[str, JsonValue]:
145+
return {}
146+
147+
148+
def test_eval_with_max_dataset_memory_and_early_stopping() -> None:
149+
samples = [Sample(input=f"Say {i}", target=str(i)) for i in range(3)]
150+
task = Task(
151+
dataset=samples,
152+
solver=[generate()],
153+
scorer=match(),
154+
early_stopping=_NoopEarlyStopping(),
155+
)
156+
log = eval(task, model="mockllm/model", max_dataset_memory=0)[0]
157+
158+
assert log.status == "success"
159+
assert log.samples is not None
160+
assert len(log.samples) == 3

0 commit comments

Comments
 (0)