Files
llm-atlas/experiments/deepseek/v2_lite_chat_cross_source_sampling_analysis.py
T
2026-07-30 03:57:17 +08:00

679 lines
22 KiB
Python

#!/usr/bin/env python3
"""Build source-blocked summaries for the Round 07 sampling grid.
The primary unit is the preregistered source. Seed-level outputs remain
within-source repeats and are never counted as independent benchmark tasks.
No p-values or population confidence intervals are produced from four sources
per domain.
"""
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, Callable
PROTOCOL_ID = "llm-atlas-deepseek-chat-cross-source-sampling-v1"
CONDITIONS = (
"s0_eos",
"s1_eos",
"s0_period",
"s1_period",
)
DOMAINS = ("english", "chinese", "code", "math")
TASK_DOMAINS = ("code", "math")
METRICS = (
"natural_eos_rate",
"mean_generated_tokens",
"fixed_budget_success_rate",
"strict_complete_success_rate",
)
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 summarize(values: list[float | int]) -> 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 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 row_success(row: dict[str, Any], strict: bool) -> int | None:
evaluation = row["task_evaluation"]
if row["domain"] == "math":
key = (
"strict_complete_numeric_exact"
if strict
else "fixed_budget_numeric_exact"
)
return int(evaluation[key])
if row["domain"] == "code":
key = (
"strict_complete_tests_pass"
if strict
else "fixed_budget_tests_pass"
)
return int(evaluation[key])
return None
def cell_summary(rows: list[dict[str, Any]]) -> dict[str, Any]:
if not rows:
raise ValueError("cell has no rows")
fixed = [
value
for row in rows
if (value := row_success(row, strict=False)) is not None
]
strict = [
value
for row in rows
if (value := row_success(row, strict=True)) is not None
]
lengths = [row["generated_tokens"] for row in rows]
hashes = {
row["generated_token_ids_sha256"] for row in rows
}
result = {
"outputs": len(rows),
"natural_eos": sum(row["hit_eos"] for row in rows),
"natural_eos_rate": mean(
int(row["hit_eos"]) for row in rows
),
"budget_truncated": sum(
row["stopped_at_max_new_tokens"] for row in rows
),
"generated_tokens": summarize(lengths),
"generated_token_seed_range": max(lengths) - min(lengths),
"unique_trajectory_hashes": len(hashes),
"completion_classes": dict(
Counter(row["completion_class"] for row in rows)
),
"fixed_budget_success": sum(fixed) if fixed else None,
"fixed_budget_success_rate": mean(fixed) if fixed else None,
"strict_complete_success": sum(strict) if strict else None,
"strict_complete_success_rate": (
mean(strict) if strict else None
),
"observed_any_fixed_budget_pass": (
any(fixed) if fixed else None
),
"observed_any_strict_complete_pass": (
any(strict) if strict else None
),
"seed_success_range": (
max(fixed) - min(fixed) if fixed else None
),
}
if rows[0]["domain"] == "math":
answers = Counter(
row["task_evaluation"]["predicted_final"]
for row in rows
if row["task_evaluation"]["predicted_final"] is not None
)
largest = max(answers.values(), default=0)
modes = sorted(
answer
for answer, count in answers.items()
if count == largest
)
result["math_answers"] = {
"frequencies": dict(answers),
"modes": modes,
"unique_absolute_majority": (
modes[0]
if len(modes) == 1 and largest > len(rows) / 2
else None
),
"gold": rows[0]["task_evaluation"]["gold_final"],
}
if rows[0]["domain"] == "code":
result["code"] = {
"ast_parse": sum(
row["task_evaluation"]["python_ast_parse"]
for row in rows
),
"executed": sum(
row["task_evaluation"]["execution"]["status"]
!= "not_run"
for row in rows
),
"execution_statuses": dict(
Counter(
row["task_evaluation"]["execution"]["status"]
for row in rows
)
),
}
return result
def contrast(
cells: dict[str, dict[str, Any]],
metric: str,
) -> dict[str, float] | None:
def metric_value(cell: dict[str, Any]) -> float | None:
if metric == "mean_generated_tokens":
return cell["generated_tokens"]["mean"]
return cell[metric]
values = {
condition: metric_value(cells[condition])
for condition in CONDITIONS
}
if any(value is None for value in values.values()):
return None
system = (
(values["s1_eos"] + values["s1_period"]) / 2
- (values["s0_eos"] + values["s0_period"]) / 2
)
boundary = (
(values["s0_period"] + values["s1_period"]) / 2
- (values["s0_eos"] + values["s1_eos"]) / 2
)
interaction = (
values["s1_period"]
- values["s1_eos"]
- values["s0_period"]
+ values["s0_eos"]
)
return {
"system_main": system,
"boundary_main_period_minus_eos": boundary,
"interaction": interaction,
"period_minus_eos_s0": (
values["s0_period"] - values["s0_eos"]
),
"period_minus_eos_s1": (
values["s1_period"] - values["s1_eos"]
),
"system_on_minus_off_eos": (
values["s1_eos"] - values["s0_eos"]
),
"system_on_minus_off_period": (
values["s1_period"] - values["s0_period"]
),
}
def domain_condition_summary(
*,
domain: str,
condition: str,
source_ids: list[str],
source_cells: dict[str, dict[str, dict[str, Any]]],
rows: list[dict[str, Any]],
) -> dict[str, Any]:
subset = [
row
for row in rows
if row["domain"] == domain
and row["condition"] == condition
]
cells = [source_cells[source_id][condition] for source_id in source_ids]
fixed_rates = [
cell["fixed_budget_success_rate"]
for cell in cells
if cell["fixed_budget_success_rate"] is not None
]
strict_rates = [
cell["strict_complete_success_rate"]
for cell in cells
if cell["strict_complete_success_rate"] is not None
]
source_mean_lengths = [
cell["generated_tokens"]["mean"] for cell in cells
]
return {
"sources": len(source_ids),
"outputs": len(subset),
"micro": {
"natural_eos": sum(row["hit_eos"] for row in subset),
"natural_eos_rate": mean(
int(row["hit_eos"]) for row in subset
),
"budget_truncated": sum(
row["stopped_at_max_new_tokens"]
for row in subset
),
"mean_generated_tokens": mean(
row["generated_tokens"] for row in subset
),
"fixed_budget_success": (
sum(
value
for row in subset
if (
value := row_success(row, strict=False)
)
is not None
)
if fixed_rates
else None
),
"strict_complete_success": (
sum(
value
for row in subset
if (
value := row_success(row, strict=True)
)
is not None
)
if strict_rates
else None
),
},
"source_macro": {
"natural_eos_rate": summarize(
[cell["natural_eos_rate"] for cell in cells]
),
"mean_generated_tokens": summarize(source_mean_lengths),
"fixed_budget_success_rate": summarize(fixed_rates),
"strict_complete_success_rate": summarize(strict_rates),
},
"variability": {
"within_source_seed_length_range": summarize(
[
cell["generated_token_seed_range"]
for cell in cells
]
),
"between_source_mean_length_range": (
max(source_mean_lengths) - min(source_mean_lengths)
),
"within_source_seed_success_range": summarize(
[
cell["seed_success_range"]
for cell in cells
if cell["seed_success_range"] is not None
]
),
"between_source_success_rate_range": (
max(fixed_rates) - min(fixed_rates)
if fixed_rates
else None
),
},
}
def build_source_cells(
rows: list[dict[str, Any]],
source_order: list[str],
) -> dict[str, dict[str, dict[str, Any]]]:
result: dict[str, dict[str, dict[str, Any]]] = {}
for source_id in source_order:
result[source_id] = {}
for condition in CONDITIONS:
subset = [
row
for row in rows
if row["source_id"] == source_id
and row["condition"] == condition
]
if len(subset) != 4:
raise RuntimeError(
f"{source_id}/{condition} has {len(subset)} rows; "
"expected 4"
)
result[source_id][condition] = cell_summary(subset)
return result
def load_json(path: Path) -> dict[str, Any]:
if not path.is_file():
raise FileNotFoundError(path)
return json.loads(path.read_text(encoding="utf-8"))
def main() -> None:
args = parse_args()
sampling = load_json(args.sampling_json)
evaluation = load_json(args.evaluation_json)
if sampling["protocol_id"] != PROTOCOL_ID:
raise RuntimeError("sampling protocol ID differs")
if evaluation["protocol_id"] != PROTOCOL_ID:
raise RuntimeError("evaluation protocol ID differs")
if evaluation["input"]["sampling_sha256"] != sha256_file(
args.sampling_json
):
raise RuntimeError("evaluation does not reference this sampling file")
if tuple(
sampling["seed_contract"]["condition_row_order"]
) != CONDITIONS:
raise RuntimeError("condition order differs from preregistration")
if len(
sampling["seed_contract"]["executed_base_seeds"]
) != 4:
raise RuntimeError("formal grid must contain four seeds")
source_metadata = {
source["id"]: {
"domain": source["domain"],
"within_domain_index": source["within_domain_index"],
"selection_rank": source["selection_rank"],
}
for source in sampling["sources"]
}
source_order = [source["id"] for source in sampling["sources"]]
if len(source_order) != 16 or len(set(source_order)) != 16:
raise RuntimeError("formal grid must contain 16 unique sources")
by_domain = {
domain: [
source_id
for source_id in source_order
if source_metadata[source_id]["domain"] == domain
]
for domain in DOMAINS
}
if any(len(source_ids) != 4 for source_ids in by_domain.values()):
raise RuntimeError("each domain must contain four sources")
rows = evaluation["rows"]
if len(rows) != 256:
raise RuntimeError(f"expected 256 evaluated rows; got {len(rows)}")
grid_keys = {
(
row["source_id"],
row["base_seed"],
row["condition"],
)
for row in rows
}
if len(grid_keys) != 256:
raise RuntimeError("evaluation grid has duplicate or missing cells")
source_cells = build_source_cells(rows, source_order)
source_contrasts = {
source_id: {
metric: contrast(source_cells[source_id], metric)
for metric in METRICS
}
for source_id in source_order
}
domain_conditions = {
domain: {
condition: domain_condition_summary(
domain=domain,
condition=condition,
source_ids=by_domain[domain],
source_cells=source_cells,
rows=rows,
)
for condition in CONDITIONS
}
for domain in DOMAINS
}
domain_contrasts: dict[str, Any] = {}
for domain in DOMAINS:
domain_contrasts[domain] = {}
for metric in METRICS:
available = {
source_id: source_contrasts[source_id][metric]
for source_id in by_domain[domain]
if source_contrasts[source_id][metric] is not None
}
domain_contrasts[domain][metric] = {}
for contrast_name in (
"system_main",
"boundary_main_period_minus_eos",
"interaction",
"period_minus_eos_s0",
"period_minus_eos_s1",
"system_on_minus_off_eos",
"system_on_minus_off_period",
):
values = [
value[contrast_name]
for value in available.values()
]
domain_contrasts[domain][metric][contrast_name] = {
**summarize(values),
"directions": direction_counts(values),
"by_source": {
source_id: value[contrast_name]
for source_id, value in available.items()
},
}
task_matrix = {
domain: [
{
"source_id": source_id,
"within_domain_index": source_metadata[source_id][
"within_domain_index"
],
"conditions": {
condition: {
key: source_cells[source_id][condition][key]
for key in (
"fixed_budget_success",
"strict_complete_success",
"observed_any_fixed_budget_pass",
"observed_any_strict_complete_pass",
"natural_eos",
"generated_tokens",
)
}
for condition in CONDITIONS
},
}
for source_id in by_domain[domain]
]
for domain in TASK_DOMAINS
}
reproduction = None
if args.reproduction_json is not None:
reproduction_payload = load_json(args.reproduction_json)
if reproduction_payload["protocol_id"] != PROTOCOL_ID:
raise RuntimeError("reproduction protocol ID differs")
if reproduction_payload["formal"]["sha256"] != sha256_file(
args.sampling_json
):
raise RuntimeError(
"reproduction does not reference this formal file"
)
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,
},
"contract": {
"source_is_primary_coverage_unit": True,
"seed_is_within_source_repeat": True,
"sources": len(source_order),
"sources_per_domain": 4,
"seeds_per_source_condition": 4,
"conditions": list(CONDITIONS),
"outputs": len(rows),
"source_order": source_order,
"sources_by_domain": by_domain,
"no_population_confidence_intervals": True,
"no_p_values": True,
},
"overall": {
"outputs": len(rows),
"natural_eos": sum(row["hit_eos"] for row in rows),
"budget_truncated": sum(
row["stopped_at_max_new_tokens"] for row in rows
),
"unique_trajectory_hashes": len(
{
row["generated_token_ids_sha256"]
for row in rows
}
),
"math": {
"outputs": sum(
row["domain"] == "math" for row in rows
),
"fixed_budget_success": sum(
row_success(row, strict=False) or 0
for row in rows
if row["domain"] == "math"
),
"strict_complete_success": sum(
row_success(row, strict=True) or 0
for row in rows
if row["domain"] == "math"
),
},
"code": {
"outputs": sum(
row["domain"] == "code" for row in rows
),
"fixed_budget_success": sum(
row_success(row, strict=False) or 0
for row in rows
if row["domain"] == "code"
),
"strict_complete_success": sum(
row_success(row, strict=True) or 0
for row in rows
if row["domain"] == "code"
),
},
},
"source_metadata": source_metadata,
"source_condition_cells": source_cells,
"domain_condition_summary": domain_conditions,
"source_contrasts": source_contrasts,
"domain_contrasts": domain_contrasts,
"task_matrix": task_matrix,
"prior_direction_check": {
domain: {
"period_shortens_mean_tokens": {
"by_source": {
source_id: (
source_contrasts[source_id][
"mean_generated_tokens"
][
"boundary_main_period_minus_eos"
]
< 0
)
for source_id in by_domain[domain]
},
},
"system_on_raises_natural_eos_rate": {
"by_source": {
source_id: (
source_contrasts[source_id][
"natural_eos_rate"
]["system_main"]
> 0
)
for source_id in by_domain[domain]
},
},
}
for domain in DOMAINS
},
"claim_boundary": [
"Four sources per domain are not full benchmark estimates.",
"Seed-level repeats are not independent task observations.",
"Source-blocked contrasts are descriptive and have no p-values.",
"Observed any-pass is not standard HumanEval pass@4.",
"Natural EOS, evaluator coverage, and correctness remain separate.",
"The period condition is a counterfactual, not an official-valid chat.",
"Round 06 motivated the period contrast, so this is a directed follow-up.",
],
}
result["content_hash"] = canonical_hash(
{
"protocol_id": result["protocol_id"],
"contract": result["contract"],
"source_condition_cells": result[
"source_condition_cells"
],
"source_contrasts": result["source_contrasts"],
"task_matrix": result["task_matrix"],
}
)
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",
)
print(
json.dumps(
{
"output": str(args.output),
"sha256": sha256_file(args.output),
"bytes": args.output.stat().st_size,
"overall": result["overall"],
"reproduction": reproduction,
},
ensure_ascii=False,
indent=2,
)
)
if __name__ == "__main__":
main()