414 lines
15 KiB
Python
414 lines
15 KiB
Python
|
|
"""Dataset, judging, and sandbox helpers for the PKPO coding run.
|
||
|
|
|
||
|
|
All generated code is executed inside Modal containers. This module still applies
|
||
|
|
per-process time and memory limits and uses a guarded Python runner that blocks
|
||
|
|
network-oriented imports.
|
||
|
|
"""
|
||
|
|
from __future__ import annotations
|
||
|
|
|
||
|
|
import base64
|
||
|
|
import json
|
||
|
|
import os
|
||
|
|
import pickle
|
||
|
|
import random
|
||
|
|
import re
|
||
|
|
import resource
|
||
|
|
import subprocess
|
||
|
|
import sys
|
||
|
|
import tempfile
|
||
|
|
import textwrap
|
||
|
|
import time
|
||
|
|
import zlib
|
||
|
|
from dataclasses import dataclass, asdict
|
||
|
|
from pathlib import Path
|
||
|
|
from typing import Any
|
||
|
|
|
||
|
|
from agent_core import coding_task_message, extract_code
|
||
|
|
|
||
|
|
|
||
|
|
@dataclass
|
||
|
|
class Problem:
|
||
|
|
source: str
|
||
|
|
problem_id: str
|
||
|
|
title: str
|
||
|
|
statement: str
|
||
|
|
public_tests: list[dict[str, str]]
|
||
|
|
hidden_tests: list[dict[str, str]]
|
||
|
|
starter_code: str = ""
|
||
|
|
difficulty: str = ""
|
||
|
|
contest_date: str = ""
|
||
|
|
|
||
|
|
def first_user_message(self, max_turns: int = 1) -> str:
|
||
|
|
return coding_task_message(
|
||
|
|
f"{self.title}\n\n{self.statement}",
|
||
|
|
starter_code=self.starter_code,
|
||
|
|
max_turns=max_turns,
|
||
|
|
)
|
||
|
|
|
||
|
|
def to_public_dict(self) -> dict[str, Any]:
|
||
|
|
d = asdict(self)
|
||
|
|
d["hidden_tests"] = [{"input": t.get("input", "")[:80], "output_len": len(t.get("output", ""))}
|
||
|
|
for t in self.hidden_tests[:3]]
|
||
|
|
d["statement"] = self.statement[:500]
|
||
|
|
return d
|
||
|
|
|
||
|
|
|
||
|
|
def _json_or_encoded_tests(value: Any) -> list[dict[str, str]]:
|
||
|
|
if value is None or value == "":
|
||
|
|
return []
|
||
|
|
if isinstance(value, list):
|
||
|
|
return value
|
||
|
|
if not isinstance(value, str):
|
||
|
|
return []
|
||
|
|
s = value.strip()
|
||
|
|
if not s:
|
||
|
|
return []
|
||
|
|
try:
|
||
|
|
obj = json.loads(s)
|
||
|
|
return obj if isinstance(obj, list) else []
|
||
|
|
except Exception:
|
||
|
|
pass
|
||
|
|
raw = zlib.decompress(base64.b64decode(s))
|
||
|
|
try:
|
||
|
|
obj = pickle.loads(raw)
|
||
|
|
except Exception:
|
||
|
|
obj = raw.decode("utf-8")
|
||
|
|
if isinstance(obj, str):
|
||
|
|
obj = json.loads(obj)
|
||
|
|
return obj if isinstance(obj, list) else []
|
||
|
|
|
||
|
|
|
||
|
|
def _stdin_tests(tests: list[dict[str, Any]], limit: int | None = None) -> list[dict[str, str]]:
|
||
|
|
out = []
|
||
|
|
for t in tests:
|
||
|
|
if t.get("testtype", "stdin") != "stdin":
|
||
|
|
continue
|
||
|
|
if "input" not in t or "output" not in t:
|
||
|
|
continue
|
||
|
|
out.append({"input": str(t["input"]), "output": str(t["output"]), "testtype": "stdin"})
|
||
|
|
if limit and len(out) >= limit:
|
||
|
|
break
|
||
|
|
return out
|
||
|
|
|
||
|
|
|
||
|
|
def load_lcb_v6_subset(limit: int = 12, seed: int = 7341) -> list[Problem]:
|
||
|
|
"""Latest code_generation_lite release file is v6/test6.jsonl."""
|
||
|
|
from datasets import load_dataset
|
||
|
|
|
||
|
|
ds = load_dataset("livecodebench/code_generation_lite", "v6", split="test", trust_remote_code=True)
|
||
|
|
candidates: list[Problem] = []
|
||
|
|
for row in ds:
|
||
|
|
if str(row.get("starter_code") or "").strip():
|
||
|
|
continue
|
||
|
|
public = _stdin_tests(_json_or_encoded_tests(row.get("public_test_cases")))
|
||
|
|
hidden = _stdin_tests(_json_or_encoded_tests(row.get("private_test_cases")))
|
||
|
|
if not public or not hidden:
|
||
|
|
continue
|
||
|
|
candidates.append(Problem(
|
||
|
|
source="livecodebench/code_generation_lite:v6",
|
||
|
|
problem_id=str(row["question_id"]),
|
||
|
|
title=str(row.get("question_title") or row["question_id"]),
|
||
|
|
statement=str(row.get("question_content") or ""),
|
||
|
|
public_tests=public,
|
||
|
|
hidden_tests=hidden,
|
||
|
|
starter_code=str(row.get("starter_code") or ""),
|
||
|
|
difficulty=str(row.get("difficulty") or ""),
|
||
|
|
contest_date=str(row.get("contest_date") or ""),
|
||
|
|
))
|
||
|
|
candidates.sort(key=lambda p: (p.contest_date, p.problem_id), reverse=True)
|
||
|
|
rng = random.Random(seed)
|
||
|
|
latest_pool = candidates[: max(limit * 4, limit)]
|
||
|
|
rng.shuffle(latest_pool)
|
||
|
|
return latest_pool[:limit]
|
||
|
|
|
||
|
|
|
||
|
|
def _tests_from_io_dict(value: Any, limit: int | None = None) -> list[dict[str, str]]:
|
||
|
|
if not isinstance(value, dict):
|
||
|
|
return []
|
||
|
|
ins = value.get("input") or []
|
||
|
|
outs = value.get("output") or []
|
||
|
|
if len(ins) != len(outs):
|
||
|
|
return []
|
||
|
|
out = []
|
||
|
|
for i, o in zip(ins, outs):
|
||
|
|
out.append({"input": str(i), "output": str(o), "testtype": "stdin"})
|
||
|
|
if limit and len(out) >= limit:
|
||
|
|
break
|
||
|
|
return out
|
||
|
|
|
||
|
|
|
||
|
|
def load_codecontest_train(
|
||
|
|
limit: int = 12,
|
||
|
|
seed: int = 20260709,
|
||
|
|
exclude_problem_ids: set[str] | None = None,
|
||
|
|
max_rows: int | None = None,
|
||
|
|
min_cf_rating: int = 800,
|
||
|
|
max_cf_rating: int = 1300,
|
||
|
|
) -> list[Problem]:
|
||
|
|
"""Sample valid old CodeContests train problems without an early-prefix bias.
|
||
|
|
|
||
|
|
``streaming=True`` yields a stable dataset prefix; shuffling only after
|
||
|
|
prematurely stopping at ``limit`` is not a sample. Reservoir sampling lets a
|
||
|
|
bounded scan supply a deterministic, broader curriculum while keeping SFT and
|
||
|
|
RL problem ids disjoint.
|
||
|
|
|
||
|
|
A Qwen3-8B base policy earns ~0 reward on unrated/hard Codeforces problems,
|
||
|
|
which starves PKPO of nonzero groups, so the RL pool is restricted to
|
||
|
|
problems with a KNOWN cf_rating inside [min_cf_rating, max_cf_rating].
|
||
|
|
"""
|
||
|
|
from datasets import load_dataset
|
||
|
|
|
||
|
|
stream = load_dataset("deepmind/code_contests", split="train", streaming=True)
|
||
|
|
excluded = exclude_problem_ids or set()
|
||
|
|
rng = random.Random(seed)
|
||
|
|
rows: list[Problem] = []
|
||
|
|
seen = 0
|
||
|
|
max_rows = max_rows or max(3000, limit * 45)
|
||
|
|
for row_idx, row in enumerate(stream, start=1):
|
||
|
|
if row_idx > max_rows:
|
||
|
|
break
|
||
|
|
rating = int(row.get("cf_rating") or 0)
|
||
|
|
if rating < min_cf_rating or rating > max_cf_rating:
|
||
|
|
continue
|
||
|
|
desc = str(row.get("description") or "")
|
||
|
|
if not desc or len(desc) > 4500:
|
||
|
|
continue
|
||
|
|
low = desc.lower()
|
||
|
|
if "interactive" in low or "output-only" in low:
|
||
|
|
continue
|
||
|
|
if str(row.get("input_file") or "").strip() or str(row.get("output_file") or "").strip():
|
||
|
|
continue
|
||
|
|
public = _tests_from_io_dict(row.get("public_tests"), limit=3)
|
||
|
|
generated = _tests_from_io_dict(row.get("generated_tests"), limit=18)
|
||
|
|
private = _tests_from_io_dict(row.get("private_tests"), limit=18)
|
||
|
|
hidden = generated or private
|
||
|
|
problem_id = str(row.get("name") or f"codecontest-{row_idx}")
|
||
|
|
if problem_id in excluded or not public or len(hidden) < 6:
|
||
|
|
continue
|
||
|
|
problem = Problem(
|
||
|
|
source="deepmind/code_contests:train",
|
||
|
|
problem_id=problem_id,
|
||
|
|
title=str(row.get("name") or "CodeContests problem"),
|
||
|
|
statement=desc,
|
||
|
|
public_tests=public,
|
||
|
|
hidden_tests=hidden,
|
||
|
|
difficulty=str(row.get("difficulty") or ""),
|
||
|
|
)
|
||
|
|
seen += 1
|
||
|
|
if len(rows) < limit:
|
||
|
|
rows.append(problem)
|
||
|
|
else:
|
||
|
|
replace_idx = rng.randrange(seen)
|
||
|
|
if replace_idx < limit:
|
||
|
|
rows[replace_idx] = problem
|
||
|
|
rng.shuffle(rows)
|
||
|
|
return rows
|
||
|
|
|
||
|
|
|
||
|
|
def python3_verified_solutions(
|
||
|
|
limit: int = 8,
|
||
|
|
seed: int = 20260710,
|
||
|
|
candidate_multiplier: int = 2,
|
||
|
|
max_rows: int | None = None,
|
||
|
|
min_cf_rating: int = 800,
|
||
|
|
max_cf_rating: int = 1500,
|
||
|
|
) -> list[tuple[Problem, str]]:
|
||
|
|
"""Return runnable CodeContests Python-3 reference solutions only.
|
||
|
|
|
||
|
|
``solutions.language`` is a ClassLabel integer in CodeContests: 1 is Python
|
||
|
|
(Python 2) while 3 is Python 3. The earlier loader ignored that paired field,
|
||
|
|
so it could SFT on ``raw_input``/bare-``print`` programs that fail our Python 3
|
||
|
|
judge. Each candidate here is also compiled and run on held-out generated
|
||
|
|
tests before it is admitted to the warm-up set.
|
||
|
|
"""
|
||
|
|
from datasets import load_dataset
|
||
|
|
|
||
|
|
stream = load_dataset("deepmind/code_contests", split="train", streaming=True)
|
||
|
|
pairs: list[tuple[Problem, str]] = []
|
||
|
|
wanted = max(limit, limit * max(1, candidate_multiplier))
|
||
|
|
max_rows = max_rows or min(6000, max(500, wanted * 40))
|
||
|
|
|
||
|
|
for row_idx, row in enumerate(stream, start=1):
|
||
|
|
if row_idx > max_rows:
|
||
|
|
break
|
||
|
|
rating = int(row.get("cf_rating") or 0)
|
||
|
|
if rating < min_cf_rating or rating > max_cf_rating:
|
||
|
|
continue
|
||
|
|
desc = str(row.get("description") or "")
|
||
|
|
if not desc or len(desc) > 4500:
|
||
|
|
continue
|
||
|
|
low = desc.lower()
|
||
|
|
if "interactive" in low or "output-only" in low:
|
||
|
|
continue
|
||
|
|
if str(row.get("input_file") or "").strip() or str(row.get("output_file") or "").strip():
|
||
|
|
continue
|
||
|
|
sols = row.get("solutions") or {}
|
||
|
|
languages = sols.get("language") or []
|
||
|
|
solutions = sols.get("solution") or []
|
||
|
|
public = _tests_from_io_dict(row.get("public_tests"), limit=2)
|
||
|
|
generated = _tests_from_io_dict(row.get("generated_tests"), limit=3)
|
||
|
|
private = _tests_from_io_dict(row.get("private_tests"), limit=3)
|
||
|
|
held_out = generated or private
|
||
|
|
if not public or not held_out:
|
||
|
|
continue
|
||
|
|
|
||
|
|
for language, raw_solution in zip(languages, solutions):
|
||
|
|
# CodeContests' ClassLabel id 3 is PYTHON3. Do not accept PYTHON (id
|
||
|
|
# 1), even if it happens to compile under the local interpreter.
|
||
|
|
if language != 3:
|
||
|
|
continue
|
||
|
|
solution = textwrap.dedent(str(raw_solution)).strip()
|
||
|
|
if not 40 <= len(solution) <= 7000:
|
||
|
|
continue
|
||
|
|
try:
|
||
|
|
compile(solution, "<codecontests-python3>", "exec")
|
||
|
|
except (SyntaxError, ValueError, TypeError):
|
||
|
|
continue
|
||
|
|
verdict = judge_code(solution, held_out, timeout_s=3, memory_mb=768, max_tests=3)
|
||
|
|
if not verdict.get("passed"):
|
||
|
|
continue
|
||
|
|
prob = Problem(
|
||
|
|
source="deepmind/code_contests:train",
|
||
|
|
problem_id=str(row.get("name") or f"sft-{len(pairs)}"),
|
||
|
|
title=str(row.get("name") or "CodeContests problem"),
|
||
|
|
statement=desc,
|
||
|
|
public_tests=public,
|
||
|
|
hidden_tests=held_out,
|
||
|
|
difficulty=str(row.get("difficulty") or ""),
|
||
|
|
)
|
||
|
|
pairs.append((prob, solution))
|
||
|
|
break
|
||
|
|
if len(pairs) >= wanted:
|
||
|
|
break
|
||
|
|
|
||
|
|
rng = random.Random(seed)
|
||
|
|
rng.shuffle(pairs)
|
||
|
|
return pairs[:limit]
|
||
|
|
|
||
|
|
|
||
|
|
GUARD = r'''
|
||
|
|
import builtins
|
||
|
|
import sys
|
||
|
|
|
||
|
|
blocked = {
|
||
|
|
"_socket", "socket", "ssl", "urllib", "http", "ftplib", "requests",
|
||
|
|
"subprocess", "multiprocessing", "ctypes",
|
||
|
|
}
|
||
|
|
real_import = builtins.__import__
|
||
|
|
|
||
|
|
def guarded_import(name, globals=None, locals=None, fromlist=(), level=0):
|
||
|
|
root = name.split(".", 1)[0]
|
||
|
|
if root in blocked:
|
||
|
|
raise ImportError(f"blocked import: {name}")
|
||
|
|
return real_import(name, globals, locals, fromlist, level)
|
||
|
|
|
||
|
|
builtins.__import__ = guarded_import
|
||
|
|
path = sys.argv[1]
|
||
|
|
with open(path, "r", encoding="utf-8", errors="replace") as f:
|
||
|
|
src = f.read()
|
||
|
|
ns = {"__name__": "__main__", "__file__": path}
|
||
|
|
exec(compile(src, path, "exec"), ns, ns)
|
||
|
|
'''
|
||
|
|
|
||
|
|
|
||
|
|
def _limit_child(timeout_s: int, memory_mb: int, cwd: str):
|
||
|
|
def inner():
|
||
|
|
os.chdir(cwd)
|
||
|
|
cpu = max(1, int(timeout_s) + 1)
|
||
|
|
for res, lim in [
|
||
|
|
(resource.RLIMIT_CPU, (cpu, cpu + 1)),
|
||
|
|
(resource.RLIMIT_AS, (memory_mb * 1024 * 1024, memory_mb * 1024 * 1024)),
|
||
|
|
(resource.RLIMIT_FSIZE, (32 * 1024 * 1024, 32 * 1024 * 1024)),
|
||
|
|
(resource.RLIMIT_NOFILE, (64, 64)),
|
||
|
|
]:
|
||
|
|
try:
|
||
|
|
resource.setrlimit(res, lim)
|
||
|
|
except Exception:
|
||
|
|
pass
|
||
|
|
return inner
|
||
|
|
|
||
|
|
|
||
|
|
def normalize_output(text: str) -> str:
|
||
|
|
text = text.replace("\r\n", "\n").replace("\r", "\n")
|
||
|
|
return "\n".join(line.rstrip() for line in text.strip().split("\n")).strip()
|
||
|
|
|
||
|
|
|
||
|
|
def judge_code(code: str, tests: list[dict[str, str]], timeout_s: int = 3,
|
||
|
|
memory_mb: int = 768, max_tests: int | None = None) -> dict[str, Any]:
|
||
|
|
tests = _stdin_tests(tests, limit=max_tests)
|
||
|
|
started = time.time()
|
||
|
|
if not code.strip():
|
||
|
|
return {
|
||
|
|
"passed": False,
|
||
|
|
"passed_tests": 0,
|
||
|
|
"total_tests": len(tests),
|
||
|
|
"error": "empty code",
|
||
|
|
"seconds": time.time() - started,
|
||
|
|
}
|
||
|
|
with tempfile.TemporaryDirectory(prefix="pkpo_exec_") as td:
|
||
|
|
sol = Path(td) / "solution.py"
|
||
|
|
runner = Path(td) / "guarded_runner.py"
|
||
|
|
sol.write_text(code, encoding="utf-8")
|
||
|
|
runner.write_text(GUARD, encoding="utf-8")
|
||
|
|
for idx, test in enumerate(tests):
|
||
|
|
try:
|
||
|
|
cp = subprocess.run(
|
||
|
|
[sys.executable, "-I", str(runner), str(sol)],
|
||
|
|
input=test["input"],
|
||
|
|
text=True,
|
||
|
|
stdout=subprocess.PIPE,
|
||
|
|
stderr=subprocess.PIPE,
|
||
|
|
timeout=timeout_s + 1,
|
||
|
|
preexec_fn=_limit_child(timeout_s, memory_mb, td),
|
||
|
|
)
|
||
|
|
except subprocess.TimeoutExpired:
|
||
|
|
return {
|
||
|
|
"passed": False,
|
||
|
|
"passed_tests": idx,
|
||
|
|
"total_tests": len(tests),
|
||
|
|
"error": "timeout",
|
||
|
|
"seconds": time.time() - started,
|
||
|
|
}
|
||
|
|
if cp.returncode != 0:
|
||
|
|
return {
|
||
|
|
"passed": False,
|
||
|
|
"passed_tests": idx,
|
||
|
|
"total_tests": len(tests),
|
||
|
|
"error": "runtime_error",
|
||
|
|
"exit_code": cp.returncode,
|
||
|
|
"stderr": cp.stderr[-500:],
|
||
|
|
"seconds": time.time() - started,
|
||
|
|
}
|
||
|
|
got = normalize_output(cp.stdout)
|
||
|
|
want = normalize_output(test["output"])
|
||
|
|
if got != want:
|
||
|
|
return {
|
||
|
|
"passed": False,
|
||
|
|
"passed_tests": idx,
|
||
|
|
"total_tests": len(tests),
|
||
|
|
"error": "wrong_answer",
|
||
|
|
"got": got[:300],
|
||
|
|
"want": want[:300],
|
||
|
|
"seconds": time.time() - started,
|
||
|
|
}
|
||
|
|
return {"passed": True, "passed_tests": len(tests), "total_tests": len(tests), "seconds": time.time() - started}
|
||
|
|
|
||
|
|
|
||
|
|
def judge_final_answer(final_answer: str, tests: list[dict[str, str]], **kwargs) -> dict[str, Any]:
|
||
|
|
return judge_code(extract_code(final_answer), tests, **kwargs)
|
||
|
|
|
||
|
|
|
||
|
|
def compact_json_dump(path: str | Path, obj: Any):
|
||
|
|
Path(path).parent.mkdir(parents=True, exist_ok=True)
|
||
|
|
Path(path).write_text(json.dumps(obj, indent=2, sort_keys=True) + "\n", encoding="utf-8")
|
||
|
|
|
||
|
|
|
||
|
|
def short_completion_for_sft(solution: str) -> str:
|
||
|
|
solution = textwrap.dedent(solution).strip()
|
||
|
|
return (
|
||
|
|
"I will provide a direct Python solution.\n</think>\n"
|
||
|
|
"<answer>Tool type: final\nTool query: ```python\n"
|
||
|
|
+ solution
|
||
|
|
+ "\n```</answer>"
|
||
|
|
)
|