Add context-cutoff dataset materialization

This commit is contained in:
jiachun
2026-08-18 19:28:20 +08:00
parent fea85047f1
commit c7943843cc
6 changed files with 414 additions and 6 deletions
+72
View File
@@ -0,0 +1,72 @@
"""Tests for publishable context-cutoff split materialization."""
from __future__ import annotations
import json
from pathlib import Path
import pyarrow as pa
import pyarrow.parquet as pq
from swe_data_processing.materialize import materialize_dataset
def _decision(
sample_id: str,
source: Path,
bucket: str,
resolved: int,
*,
review: bool = False,
) -> dict:
return {
"sample_id": sample_id,
"source_parquet": str(source),
"resolved": resolved,
"length_bucket": bucket,
"hard_reject_reasons": [],
"review_flags": ["EARLY_FAILURE_CLUSTER"] if review else [],
"metrics": {
"token_count": 50_000 if bucket == "LE_81920" else 100_000,
"failed_tool_call_count": 1,
"failed_tool_call_rate": 0.1,
"longest_consecutive_failure_run": 1,
},
}
def test_materialize_creates_nested_context_splits(tmp_path: Path) -> None:
source = tmp_path / "source" / "family" / "part.parquet"
source.parent.mkdir(parents=True)
pq.write_table(
pa.Table.from_pylist(
[
{"trajectory_id": "short", "resolved": 1},
{"trajectory_id": "medium", "resolved": 0},
{"trajectory_id": "review", "resolved": -1},
]
),
source,
)
decisions = tmp_path / "decisions.jsonl"
rows = [
_decision("short", source, "LE_81920", 1),
_decision("medium", source, "81921_TO_131072", 0),
_decision("review", source, "LE_81920", -1, review=True),
]
decisions.write_text(
"".join(json.dumps(row) + "\n" for row in rows), encoding="utf-8"
)
output = tmp_path / "output"
manifest = materialize_dataset(decisions, output, workers=1, resume=False)
assert manifest["splits"]["context_131072"]["samples"] == 2
assert manifest["splits"]["context_81920"]["samples"] == 1
long_table = pq.read_table(output / "data/context_131072")
short_table = pq.read_table(output / "data/context_81920")
assert long_table.column("trajectory_id").to_pylist() == ["short", "medium"]
assert short_table.column("trajectory_id").to_pylist() == ["short"]
assert set(long_table.column("qc_filter_version").to_pylist()) == {
"deterministic-v3.1"
}