research: add task bootstrap evaluation pipeline

This commit is contained in:
wuyang
2026-07-30 04:45:40 +08:00
parent a09e6b14ce
commit 65a91ae178
3 changed files with 1208 additions and 0 deletions
@@ -0,0 +1,620 @@
#!/usr/bin/env python3
"""Build task-paired bootstrap and multi-tape diagnostics for Round 08."""
from __future__ import annotations
import argparse
import hashlib
import json
from collections import Counter
from pathlib import Path
from statistics import mean, median
from typing import Any
import numpy as np
PROTOCOL_ID = "llm-atlas-deepseek-chat-task-bootstrap-crn-v1"
CONDITIONS = (
"s0_eos",
"s1_eos",
"s0_period",
"s1_period",
)
DOMAINS = ("code", "math")
BOOTSTRAP_SEED = 1364512825
BOOTSTRAP_RESAMPLES = 10000
CONTRASTS = {
"period_at_s0": ("s0_eos", "s0_period"),
"period_at_s1": ("s1_eos", "s1_period"),
"system_at_eos": ("s0_eos", "s1_eos"),
"system_at_period": ("s0_period", "s1_period"),
}
METRICS = (
"fixed_budget_success",
"strict_complete_success",
"natural_eos",
"generated_tokens",
)
def parse_args() -> argparse.Namespace:
parser = argparse.ArgumentParser()
parser.add_argument("--sampling-json", type=Path, required=True)
parser.add_argument("--evaluation-json", type=Path, required=True)
parser.add_argument("--reproduction-json", type=Path)
parser.add_argument("--output", type=Path, required=True)
return parser.parse_args()
def sha256_file(path: Path) -> str:
digest = hashlib.sha256()
with path.open("rb") as handle:
while chunk := handle.read(16 * 1024 * 1024):
digest.update(chunk)
return digest.hexdigest()
def canonical_hash(value: Any) -> str:
return hashlib.sha256(
json.dumps(
value,
ensure_ascii=False,
sort_keys=True,
separators=(",", ":"),
).encode()
).hexdigest()
def metric_value(row: dict[str, Any], metric: str) -> float:
evaluation = row["task_evaluation"]
if metric == "fixed_budget_success":
key = (
"fixed_budget_tests_pass"
if row["domain"] == "code"
else "fixed_budget_numeric_exact"
)
return float(evaluation[key])
if metric == "strict_complete_success":
key = (
"strict_complete_tests_pass"
if row["domain"] == "code"
else "strict_complete_numeric_exact"
)
return float(evaluation[key])
if metric == "natural_eos":
return float(row["hit_eos"])
if metric == "generated_tokens":
return float(row["generated_tokens"])
raise KeyError(metric)
def summarize(values: list[float]) -> dict[str, Any]:
return {
"count": len(values),
"mean": mean(values) if values else None,
"median": median(values) if values else None,
"min": min(values) if values else None,
"max": max(values) if values else None,
}
def percentile_band(values: np.ndarray) -> dict[str, float]:
quantiles = np.percentile(values, [2.5, 50.0, 97.5])
return {
"p2_5": float(quantiles[0]),
"p50": float(quantiles[1]),
"p97_5": float(quantiles[2]),
}
def direction_counts(values: list[float]) -> dict[str, int]:
epsilon = 1e-12
return {
"positive": sum(value > epsilon for value in values),
"zero": sum(abs(value) <= epsilon for value in values),
"negative": sum(value < -epsilon for value in values),
}
def success_transition(
left: list[float],
right: list[float],
) -> dict[str, int]:
pairs = [(int(a), int(b)) for a, b in zip(left, right, strict=True)]
return {
"fail_to_fail": sum(a == 0 and b == 0 for a, b in pairs),
"fail_to_pass": sum(a == 0 and b == 1 for a, b in pairs),
"pass_to_fail": sum(a == 1 and b == 0 for a, b in pairs),
"pass_to_pass": sum(a == 1 and b == 1 for a, b in pairs),
}
def common_prefix(left: list[int], right: list[int]) -> int:
count = 0
for left_id, right_id in zip(left, right):
if left_id != right_id:
break
count += 1
return count
def tape_uint64(tape: str, source_id: str, step: int) -> int:
payload = (
f"{PROTOCOL_ID}\0uniform\0"
f"{tape}\0{source_id}\0{step}"
).encode()
return int.from_bytes(hashlib.sha256(payload).digest()[:8], "big")
def build_indices(
sampling: dict[str, Any],
evaluation: dict[str, Any],
) -> tuple[
dict[tuple[str, str, str], dict[str, Any]],
dict[tuple[str, str, str], dict[str, Any]],
]:
sample_index = {
(
source["id"],
run["tape_label"],
output["condition"],
): output
for source in sampling["sources"]
for run in source["runs"]
for output in run["outputs"]
}
eval_index = {
(
row["source_id"],
row["tape_label"],
row["condition"],
): row
for row in evaluation["rows"]
}
if set(sample_index) != set(eval_index):
raise RuntimeError("sampling/evaluation cell keys differ")
return sample_index, eval_index
def main_analysis(
*,
sampling: dict[str, Any],
sample_index: dict[tuple[str, str, str], dict[str, Any]],
eval_index: dict[tuple[str, str, str], dict[str, Any]],
) -> dict[str, Any]:
result = {}
for domain in DOMAINS:
sources = [
source
for source in sampling["sources"]
if source["domain"] == domain
]
sources.sort(key=lambda row: row["within_domain_index"])
if len(sources) != 32:
raise RuntimeError(f"{domain}: main frame must have 32 tasks")
source_ids = [source["id"] for source in sources]
rng = np.random.default_rng(BOOTSTRAP_SEED)
sampled_indices = rng.integers(
0,
len(source_ids),
size=(BOOTSTRAP_RESAMPLES, len(source_ids)),
endpoint=False,
)
conditions = {}
for condition in CONDITIONS:
rows = [
eval_index[(source_id, "T0", condition)]
for source_id in source_ids
]
conditions[condition] = {
metric: summarize(
[metric_value(row, metric) for row in rows]
)
for metric in METRICS
} | {
"task_outcomes": dict(
sorted(
Counter(row["task_outcome"] for row in rows).items()
)
)
}
contrasts = {}
for name, (left_condition, right_condition) in CONTRASTS.items():
metrics = {}
for metric in METRICS:
left = [
metric_value(
eval_index[(source_id, "T0", left_condition)],
metric,
)
for source_id in source_ids
]
right = [
metric_value(
eval_index[(source_id, "T0", right_condition)],
metric,
)
for source_id in source_ids
]
differences = np.asarray(right) - np.asarray(left)
bootstrap_means = differences[
sampled_indices
].mean(axis=1)
metric_result = {
"right_minus_left_point": float(differences.mean()),
"selected_task_resampling_band": percentile_band(
bootstrap_means
),
"task_differences": summarize(
differences.tolist()
),
"direction_counts": direction_counts(
differences.tolist()
),
"by_source": [
{
"source_id": source_id,
"domain_index": index,
"left": float(left[index]),
"right": float(right[index]),
"right_minus_left": float(
differences[index]
),
}
for index, source_id in enumerate(source_ids)
],
}
if metric in {
"fixed_budget_success",
"strict_complete_success",
}:
metric_result["transition"] = success_transition(
left,
right,
)
metrics[metric] = metric_result
trajectory_rows = []
crn_exact = 0
for source_id in source_ids:
left_output = sample_index[
(source_id, "T0", left_condition)
]
right_output = sample_index[
(source_id, "T0", right_condition)
]
left_ids = left_output["generated_token_ids"]
right_ids = right_output["generated_token_ids"]
shared_steps = min(len(left_ids), len(right_ids))
expected_left_uniforms = [
tape_uint64("T0", source_id, step)
for step in range(
left_output["uniform_steps_consumed"]
)
]
expected_right_uniforms = [
tape_uint64("T0", source_id, step)
for step in range(
right_output["uniform_steps_consumed"]
)
]
left_uniform_exact = (
canonical_hash(expected_left_uniforms)
== left_output["uniform_uint64_prefix_sha256"]
)
right_uniform_exact = (
canonical_hash(expected_right_uniforms)
== right_output["uniform_uint64_prefix_sha256"]
)
crn_exact += left_uniform_exact and right_uniform_exact
trajectory_rows.append(
{
"source_id": source_id,
"shared_active_steps": shared_steps,
"common_prefix_tokens": common_prefix(
left_ids,
right_ids,
),
"token_ids_exact": left_ids == right_ids,
"shared_uniform_prefix_exact": (
left_uniform_exact and right_uniform_exact
),
}
)
contrasts[name] = {
"left": left_condition,
"right": right_condition,
"metrics": metrics,
"trajectory": {
"sources": len(trajectory_rows),
"shared_uniform_prefix_exact": crn_exact,
"exact_trajectories": sum(
row["token_ids_exact"]
for row in trajectory_rows
),
"common_prefix_tokens": summarize(
[
float(row["common_prefix_tokens"])
for row in trajectory_rows
]
),
"rows": trajectory_rows,
},
}
result[domain] = {
"tasks": len(source_ids),
"source_ids": source_ids,
"tape": "T0",
"conditions": conditions,
"contrasts": contrasts,
}
return result
def tape_diagnostic(
*,
sampling: dict[str, Any],
eval_index: dict[tuple[str, str, str], dict[str, Any]],
) -> dict[str, Any]:
result = {}
for domain in DOMAINS:
sources = [
source
for source in sampling["sources"]
if source["domain"] == domain and len(source["runs"]) == 4
]
sources.sort(key=lambda row: row["within_domain_index"])
if [row["within_domain_index"] for row in sources] != [
0,
8,
16,
24,
]:
raise RuntimeError(
f"{domain}: diagnostic task indices differ"
)
source_ids = [row["id"] for row in sources]
tapes = ["T0", "T1", "T2", "T3"]
by_tape_condition = {
tape: {
condition: {
metric: summarize(
[
metric_value(
eval_index[
(source_id, tape, condition)
],
metric,
)
for source_id in source_ids
]
)
for metric in METRICS
}
for condition in CONDITIONS
}
for tape in tapes
}
contrasts = {}
for name, (left_condition, right_condition) in CONTRASTS.items():
metric_results = {}
for metric in METRICS:
matrix = np.asarray(
[
[
metric_value(
eval_index[
(
source_id,
tape,
right_condition,
)
],
metric,
)
- metric_value(
eval_index[
(
source_id,
tape,
left_condition,
)
],
metric,
)
for tape in tapes
]
for source_id in source_ids
],
dtype=float,
)
tape_means = matrix.mean(axis=0)
task_means = matrix.mean(axis=1)
task_ranges_within_tape = (
matrix.max(axis=0) - matrix.min(axis=0)
)
tape_ranges_within_task = (
matrix.max(axis=1) - matrix.min(axis=1)
)
metric_results[metric] = {
"matrix_task_by_tape": matrix.tolist(),
"tape_means": {
tape: float(tape_means[index])
for index, tape in enumerate(tapes)
},
"task_means": {
source_id: float(task_means[index])
for index, source_id in enumerate(source_ids)
},
"direction_by_tape": {
tape: direction_counts(
matrix[:, index].tolist()
)
for index, tape in enumerate(tapes)
},
"task_range_within_tape": summarize(
task_ranges_within_tape.tolist()
),
"tape_range_within_task": summarize(
tape_ranges_within_task.tolist()
),
"grand_mean_descriptive": float(matrix.mean()),
}
contrasts[name] = {
"left": left_condition,
"right": right_condition,
"metrics": metric_results,
}
result[domain] = {
"tasks": len(source_ids),
"source_ids": source_ids,
"tapes": tapes,
"by_tape_condition": by_tape_condition,
"contrasts": contrasts,
"independence_warning": (
"The 4 tasks x 4 tapes are crossed repeated measures, "
"not 16 independent tasks."
),
}
return result
def main() -> None:
args = parse_args()
for path in (args.sampling_json, args.evaluation_json):
if not path.is_file():
raise FileNotFoundError(path)
sampling = json.loads(
args.sampling_json.read_text(encoding="utf-8")
)
evaluation = json.loads(
args.evaluation_json.read_text(encoding="utf-8")
)
if (
sampling["protocol_id"] != PROTOCOL_ID
or evaluation["protocol_id"] != PROTOCOL_ID
):
raise RuntimeError("protocol ID differs")
if sampling["summary"]["outputs"] != 352:
raise RuntimeError("sampling formal output count differs")
if len(evaluation["rows"]) != 352:
raise RuntimeError("evaluation row count differs")
sample_index, eval_index = build_indices(sampling, evaluation)
reproduction = None
if args.reproduction_json is not None:
if not args.reproduction_json.is_file():
raise FileNotFoundError(args.reproduction_json)
reproduction_payload = json.loads(
args.reproduction_json.read_text(encoding="utf-8")
)
if reproduction_payload["protocol_id"] != PROTOCOL_ID:
raise RuntimeError("reproduction protocol ID differs")
reproduction = {
"path": str(args.reproduction_json),
"sha256": sha256_file(args.reproduction_json),
"summary": reproduction_payload["summary"],
}
result = {
"schema_version": 1,
"protocol_id": PROTOCOL_ID,
"inputs": {
"sampling": {
"path": str(args.sampling_json),
"sha256": sha256_file(args.sampling_json),
"content_hash": sampling["content_hash"],
},
"evaluation": {
"path": str(args.evaluation_json),
"sha256": sha256_file(args.evaluation_json),
"content_hash": evaluation["content_hash"],
},
"reproduction": reproduction,
},
"bootstrap_contract": {
"resamples": BOOTSTRAP_RESAMPLES,
"seed": BOOTSTRAP_SEED,
"rng": "numpy.random.default_rng reset per domain",
"unit": "selected task",
"paired_conditions": True,
"tape": "T0",
"interval_label": (
"selected-task resampling band for the fixed "
"32-task frame and T0"
),
"not": [
"benchmark-population confidence interval",
"model-ability confidence interval",
"generation-seed uncertainty interval",
"causal-effect confidence interval",
],
},
"main_t0_selected_task_analysis": main_analysis(
sampling=sampling,
sample_index=sample_index,
eval_index=eval_index,
),
"multi_tape_diagnostic": tape_diagnostic(
sampling=sampling,
eval_index=eval_index,
),
"claim_boundary": [
"HumanEval and GSM8K are analyzed separately.",
(
"T0 selected-task bands describe only this frozen "
"32-task frame."
),
(
"T1-T3 are sensitivity diagnostics and are not pooled "
"into the T0 primary success rates."
),
(
"Common random numbers align probability quantiles; "
"they do not force identical sampled tokens."
),
"Period prompts are counterfactual, not official-valid chats.",
],
}
result["content_hash"] = canonical_hash(
{
"bootstrap_contract": result["bootstrap_contract"],
"main": result["main_t0_selected_task_analysis"],
"diagnostic": result["multi_tape_diagnostic"],
}
)
args.output.parent.mkdir(parents=True, exist_ok=True)
args.output.write_text(
json.dumps(result, ensure_ascii=False, indent=2) + "\n",
encoding="utf-8",
)
payload = args.output.read_bytes()
compact = {
domain: {
name: {
metric: result[
"main_t0_selected_task_analysis"
][domain]["contrasts"][name]["metrics"][metric][
"right_minus_left_point"
]
for metric in METRICS
}
for name in CONTRASTS
}
for domain in DOMAINS
}
print(
json.dumps(
{
"output": str(args.output),
"sha256": hashlib.sha256(payload).hexdigest(),
"bytes": len(payload),
"main_contrast_points": compact,
},
ensure_ascii=False,
indent=2,
)
)
if __name__ == "__main__":
main()
@@ -0,0 +1,364 @@
#!/usr/bin/env python3
"""Evaluate Round 08 while keeping T0 and multi-tape diagnostics separate."""
from __future__ import annotations
import argparse
import json
from collections import Counter
from pathlib import Path
from statistics import mean
from typing import Any
import v2_lite_chat_completion_evaluator as completion
import v2_lite_chat_sampling_evaluator as shared
PROTOCOL_ID = "llm-atlas-deepseek-chat-task-bootstrap-crn-v1"
CONDITIONS = (
"s0_eos",
"s1_eos",
"s0_period",
"s1_period",
)
DOMAINS = ("code", "math")
def parse_args() -> argparse.Namespace:
parser = argparse.ArgumentParser()
parser.add_argument("--sampling-json", type=Path, required=True)
parser.add_argument("--human-eval", type=Path, required=True)
parser.add_argument("--gsm8k", type=Path, required=True)
parser.add_argument("--output", type=Path, required=True)
parser.add_argument(
"--sandbox-image",
required=True,
help="Pinned image reference including @sha256 digest.",
)
parser.add_argument("--timeout-seconds", type=float, default=5.0)
parser.add_argument("--skip-code-execution", action="store_true")
return parser.parse_args()
def task_outcome(
domain: str,
evaluation: dict[str, Any],
) -> str:
if domain == "math":
if evaluation["fixed_budget_numeric_exact"]:
return "passed"
if evaluation["predicted_final"] is None:
return "no_numeric_answer"
return "wrong_numeric_answer"
if not evaluation["contains_entry_point_definition"]:
return "extract_failed"
if not evaluation["python_ast_parse"]:
return "syntax_error"
status = evaluation["execution"]["status"]
return {
"passed": "passed",
"timeout": "timeout",
"runtime_error": "runtime_error",
"assertion_failed": "assertion_failed",
"not_run": "not_run",
}.get(status, f"execution_{status}")
def success(row: dict[str, Any], *, strict: bool = False) -> int:
evaluation = row["task_evaluation"]
if row["domain"] == "code":
key = (
"strict_complete_tests_pass"
if strict
else "fixed_budget_tests_pass"
)
else:
key = (
"strict_complete_numeric_exact"
if strict
else "fixed_budget_numeric_exact"
)
return int(evaluation[key])
def compact_summary(rows: list[dict[str, Any]]) -> dict[str, Any]:
result = {
"outputs": len(rows),
"sources": len({row["source_id"] for row in rows}),
"tapes": sorted({row["tape_label"] for row in rows}),
"natural_eos": sum(row["hit_eos"] for row in rows),
"budget_truncated": sum(
row["stopped_at_max_new_tokens"] for row in rows
),
"mean_generated_tokens": (
mean(row["generated_tokens"] for row in rows)
if rows
else None
),
"fixed_budget_success": sum(success(row) for row in rows),
"strict_complete_success": sum(
success(row, strict=True) for row in rows
),
"task_outcomes": dict(
sorted(Counter(row["task_outcome"] for row in rows).items())
),
"completion_classes": dict(
sorted(
Counter(
row["completion_class"] for row in rows
).items()
)
),
}
result["by_domain"] = {
domain: compact_summary_no_recursion(
[row for row in rows if row["domain"] == domain]
)
for domain in DOMAINS
}
return result
def compact_summary_no_recursion(
rows: list[dict[str, Any]],
) -> dict[str, Any]:
return {
"outputs": len(rows),
"sources": len({row["source_id"] for row in rows}),
"tapes": sorted({row["tape_label"] for row in rows}),
"natural_eos": sum(row["hit_eos"] for row in rows),
"budget_truncated": sum(
row["stopped_at_max_new_tokens"] for row in rows
),
"mean_generated_tokens": (
mean(row["generated_tokens"] for row in rows)
if rows
else None
),
"fixed_budget_success": sum(success(row) for row in rows),
"strict_complete_success": sum(
success(row, strict=True) for row in rows
),
"task_outcomes": dict(
sorted(Counter(row["task_outcome"] for row in rows).items())
),
}
def main() -> None:
args = parse_args()
for path in (
args.sampling_json,
args.human_eval,
args.gsm8k,
):
if not path.is_file():
raise FileNotFoundError(path)
if "@sha256:" not in args.sandbox_image:
raise ValueError("--sandbox-image must include immutable digest")
if args.timeout_seconds <= 0:
raise ValueError("--timeout-seconds must be positive")
sampling = json.loads(
args.sampling_json.read_text(encoding="utf-8")
)
if sampling["protocol_id"] != PROTOCOL_ID:
raise RuntimeError("sampling protocol ID differs")
if tuple(
sampling["seed_contract"]["condition_row_order"]
) != CONDITIONS:
raise RuntimeError("sampling condition order differs")
if sampling["summary"]["outputs"] != 352:
raise RuntimeError("formal grid must contain 352 outputs")
human_eval, gsm8k = completion.load_tasks(
args.human_eval,
args.gsm8k,
)
code_cache: dict[
tuple[str, str, str],
dict[str, Any],
] = {}
rows = []
for source in sampling["sources"]:
for run in source["runs"]:
for output in run["outputs"]:
if source["domain"] == "math":
evaluation = completion.evaluate_math(
output["text"],
output["hit_eos"],
gsm8k[source["id"]]["answer"],
)
cache_hit = False
else:
evaluation, cache_hit = (
shared.evaluate_code_cached(
text=output["text"],
hit_eos=output["hit_eos"],
task=human_eval[source["id"]],
image=args.sandbox_image,
timeout_seconds=args.timeout_seconds,
skip_execution=args.skip_code_execution,
cache=code_cache,
)
)
row = {
"source_id": source["id"],
"domain": source["domain"],
"domain_index": source["within_domain_index"],
"replicate_index": run["replicate_index"],
"replicate_label": run["replicate_label"],
"tape_label": run["tape_label"],
"base_seed": run["base_seed"],
"run_seed": run["run_seed"],
"condition": output["condition"],
"generated_tokens": output["generated_tokens"],
"hit_eos": output["hit_eos"],
"stopped_at_max_new_tokens": output[
"stopped_at_max_new_tokens"
],
"prompt_token_ids_sha256": output[
"prompt_token_ids_sha256"
],
"generated_token_ids_sha256": output[
"generated_token_ids_sha256"
],
"text_sha256": output["text_sha256"],
"uniform_uint64_prefix_sha256": output[
"uniform_uint64_prefix_sha256"
],
"task_evaluation": evaluation,
"completion_class": completion.completion_class(
output,
evaluation,
),
"code_execution_cache_hit": cache_hit,
}
row["task_outcome"] = task_outcome(
source["domain"],
evaluation,
)
rows.append(row)
main_rows = [
row for row in rows if row["tape_label"] == "T0"
]
diagnostic_sources = {
source["id"]
for source in sampling["sources"]
if len(source["runs"]) == 4
}
diagnostic_rows = [
row
for row in rows
if row["source_id"] in diagnostic_sources
]
additional_rows = [
row for row in rows if row["tape_label"] != "T0"
]
if (
len(main_rows) != 256
or len(diagnostic_rows) != 128
or len(additional_rows) != 96
):
raise RuntimeError(
"T0 / diagnostic grid counts differ from protocol"
)
result = {
"schema_version": 1,
"protocol_id": PROTOCOL_ID,
"input": {
"sampling_path": str(args.sampling_json),
"sampling_sha256": completion.sha256_file(
args.sampling_json
),
"sampling_content_hash": sampling["content_hash"],
"human_eval_sha256": completion.sha256_file(
args.human_eval
),
"gsm8k_sha256": completion.sha256_file(args.gsm8k),
"model_revision": sampling["model"]["revision"],
"max_new_tokens": sampling["generation_contract"][
"max_new_tokens"
],
},
"sandbox": {
"image": args.sandbox_image,
"timeout_seconds": args.timeout_seconds,
"code_execution_skipped": args.skip_code_execution,
"network": "none",
"filesystem": "read-only",
"user": "65534:65534",
"capabilities": "ALL dropped",
"no_new_privileges": True,
"memory": "256m",
"memory_swap": "256m",
"pids_limit": 64,
"cpus": 0.5,
"tmpfs": "/tmp:rw,noexec,nosuid,size=16m",
"host_mounts": 0,
"unique_code_cache_entries": len(code_cache),
},
"rows": rows,
"summary": {
"formal_all_outputs": compact_summary(rows),
"main_t0": compact_summary(main_rows),
"diagnostic_all_four_tapes": compact_summary(
diagnostic_rows
),
"diagnostic_additional_t1_t3": compact_summary(
additional_rows
),
"main_t0_by_domain_condition": {
domain: {
condition: compact_summary_no_recursion(
[
row
for row in main_rows
if row["domain"] == domain
and row["condition"] == condition
]
)
for condition in CONDITIONS
}
for domain in DOMAINS
},
},
"claim_boundary": [
(
"Primary summaries use T0 only; T1-T3 are isolated "
"multi-tape sensitivity diagnostics."
),
"HumanEval and GSM8K are never pooled into one ability rate.",
(
"The frozen 32 tasks per domain are a selected task "
"frame, not a full benchmark sample."
),
"Passing HumanEval tests is functional, not safety, evidence.",
"Counterfactual period prompts are not official-valid chats.",
],
}
result["content_hash"] = completion.canonical_hash(rows)
args.output.parent.mkdir(parents=True, exist_ok=True)
args.output.write_text(
json.dumps(result, ensure_ascii=False, indent=2) + "\n",
encoding="utf-8",
)
payload = args.output.read_bytes()
print(
json.dumps(
{
"output": str(args.output),
"sha256": completion.sha256_bytes(payload),
"bytes": len(payload),
"rows": len(rows),
"main_t0": result["summary"]["main_t0"],
"unique_code_cache_entries": len(code_cache),
},
ensure_ascii=False,
indent=2,
)
)
if __name__ == "__main__":
main()
@@ -0,0 +1,224 @@
#!/usr/bin/env python3
"""Compare the preregistered 64-cell Round 08 replay with the formal grid."""
from __future__ import annotations
import argparse
import hashlib
import json
from pathlib import Path
from typing import Any
PROTOCOL_ID = "llm-atlas-deepseek-chat-task-bootstrap-crn-v1"
def parse_args() -> argparse.Namespace:
parser = argparse.ArgumentParser()
parser.add_argument("--formal-json", type=Path, required=True)
parser.add_argument("--rerun-json", type=Path, required=True)
parser.add_argument("--output", type=Path, required=True)
return parser.parse_args()
def sha256_file(path: Path) -> str:
digest = hashlib.sha256()
with path.open("rb") as handle:
while chunk := handle.read(16 * 1024 * 1024):
digest.update(chunk)
return digest.hexdigest()
def canonical_hash(value: Any) -> str:
return hashlib.sha256(
json.dumps(
value,
ensure_ascii=False,
sort_keys=True,
separators=(",", ":"),
).encode()
).hexdigest()
def run_index(
payload: dict[str, Any],
) -> dict[
tuple[str, str, str],
tuple[dict[str, Any], dict[str, Any]],
]:
return {
(
source["id"],
run["tape_label"],
output["condition"],
): (run, output)
for source in payload["sources"]
for run in source["runs"]
for output in run["outputs"]
}
def main() -> None:
args = parse_args()
for path in (args.formal_json, args.rerun_json):
if not path.is_file():
raise FileNotFoundError(path)
formal = json.loads(
args.formal_json.read_text(encoding="utf-8")
)
rerun = json.loads(
args.rerun_json.read_text(encoding="utf-8")
)
if (
formal["protocol_id"] != PROTOCOL_ID
or rerun["protocol_id"] != PROTOCOL_ID
):
raise RuntimeError("protocol ID differs")
if formal["model"]["revision"] != rerun["model"]["revision"]:
raise RuntimeError("model revision differs")
if formal["generation_contract"] != rerun["generation_contract"]:
raise RuntimeError("generation contract differs")
if rerun["summary"]["outputs"] != 64:
raise RuntimeError("replay grid must contain 64 outputs")
formal_rows = run_index(formal)
rerun_rows = run_index(rerun)
if any(key[1] != "T0" for key in rerun_rows):
raise RuntimeError("replay grid must use T0 only")
missing = sorted(set(rerun_rows) - set(formal_rows))
if missing:
raise RuntimeError(
f"{len(missing)} replay cells are absent from formal grid"
)
rows = []
for key in sorted(rerun_rows):
formal_run, formal_output = formal_rows[key]
rerun_run, rerun_output = rerun_rows[key]
checks = {
"run_seed_exact": (
formal_run["run_seed"] == rerun_run["run_seed"]
),
"prompt_hash_exact": (
formal_output["prompt_token_ids_sha256"]
== rerun_output["prompt_token_ids_sha256"]
),
"uniform_run_hash_exact": (
formal_run["uniform_uint64_sha256"]
== rerun_run["uniform_uint64_sha256"]
),
"uniform_output_prefix_hash_exact": (
formal_output["uniform_uint64_prefix_sha256"]
== rerun_output["uniform_uint64_prefix_sha256"]
),
"uniform_steps_exact": (
formal_output["uniform_steps_consumed"]
== rerun_output["uniform_steps_consumed"]
),
"generated_token_ids_exact": (
formal_output["generated_token_ids"]
== rerun_output["generated_token_ids"]
),
"decoded_text_exact": (
formal_output["text"] == rerun_output["text"]
),
"eos_state_exact": (
formal_output["hit_eos"]
== rerun_output["hit_eos"]
),
"truncation_state_exact": (
formal_output["stopped_at_max_new_tokens"]
== rerun_output["stopped_at_max_new_tokens"]
),
"cpu_rng_pre_state_exact": (
formal_run["rng_state_before"]["cpu_sha256"]
== rerun_run["rng_state_before"]["cpu_sha256"]
),
"cuda_rng_pre_state_exact": (
formal_run["rng_state_before"][
"cuda_combined_sha256"
]
== rerun_run["rng_state_before"][
"cuda_combined_sha256"
]
),
"torch_rng_unchanged_exact": (
formal_run["torch_rng_unchanged"]
and rerun_run["torch_rng_unchanged"]
),
}
rows.append(
{
"source_id": key[0],
"tape_label": key[1],
"condition": key[2],
**checks,
"all_preregistered_fields_exact": all(
checks.values()
),
}
)
check_names = [
key
for key in rows[0]
if key.endswith("_exact")
and key != "all_preregistered_fields_exact"
]
result = {
"schema_version": 1,
"protocol_id": PROTOCOL_ID,
"formal": {
"path": str(args.formal_json),
"sha256": sha256_file(args.formal_json),
"content_hash": formal["content_hash"],
},
"rerun": {
"path": str(args.rerun_json),
"sha256": sha256_file(args.rerun_json),
"content_hash": rerun["content_hash"],
},
"rows": rows,
"summary": {
"cells": len(rows),
"all_preregistered_fields_exact": sum(
row["all_preregistered_fields_exact"]
for row in rows
),
"by_field": {
name: sum(row[name] for row in rows)
for name in check_names
},
},
"claim_boundary": [
"Only the preregistered 16-source T0 subset is replayed.",
(
"Exact replay is scoped to the pinned checkpoint, "
"software, sampler, precision, and hardware contract."
),
"Reproduction does not imply tape-invariant trajectories.",
],
}
if result["summary"]["cells"] != 64:
raise RuntimeError("reproduction comparison must contain 64 cells")
result["content_hash"] = canonical_hash(rows)
args.output.parent.mkdir(parents=True, exist_ok=True)
args.output.write_text(
json.dumps(result, ensure_ascii=False, indent=2) + "\n",
encoding="utf-8",
)
payload = args.output.read_bytes()
print(
json.dumps(
{
"output": str(args.output),
"sha256": hashlib.sha256(payload).hexdigest(),
"bytes": len(payload),
"summary": result["summary"],
},
ensure_ascii=False,
indent=2,
)
)
if __name__ == "__main__":
main()