621 lines
21 KiB
Python
621 lines
21 KiB
Python
#!/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()
|