feat(adr-0250): 2b certified summation turn — full dev holdout 50/50 wrong=0

A total-like unknown (unknown.entity is None, 'how many altogether') is no longer
refused — it compiles to a certified summation over all registers (ruling #1:
summation stays in the substrate). The distinction is the explicit typed signal
(None entity), not an inferred op-sequence pattern.

Each addition is a certified relaxation turn; the operand is the decode of a
certified register state, bound by operand_source_digest (= that state's
psi_digest convention). Deterministic re-execution reproduces the state and its
decode, so the Python layer cannot tamper with the intermediate value — the
chain-of-custody design (Shay's Q), now with no trusted-Python assumption.
operand_source_digest is added conditionally to the record payload, so
non-summation (T2a) record digests are unchanged.

Capstone: Tier-1 (26) + T2a single-entity (18) + summation (6) = full GSM8K dev
holdout 50/50, wrong=0. The reader->Hamiltonian compiler now solves the entire
dev holdout by chained certified relaxation, zero wrong answers.

16/16 pins green (11 multi-register updated + 5 summation incl. capstone).

[Verification]: uv run python -m pytest tests/test_adr_0250_multi_register.py tests/test_adr_0250_summation.py -q
This commit is contained in:
Shay 2026-07-18 14:25:44 -07:00
parent efd5335e86
commit 3d7d583d83
3 changed files with 152 additions and 10 deletions

View file

@ -81,6 +81,19 @@ def _digest(payload: dict) -> str:
).hexdigest()
# Pseudo-entity for certified-summation accumulator records.
_TOTAL = "__total__"
def _state_digest(psi: np.ndarray) -> str:
"""Byte digest of a state (explicit ``<f8`` LE) — the same convention as
``RelaxationCertificate.psi_digest``, so a summation operand binds to the
exact certified state it was decoded from."""
return hashlib.sha256(
np.ascontiguousarray(np.asarray(psi, dtype=np.dtype("<f8"))).tobytes()
).hexdigest()
@dataclass(frozen=True, slots=True)
class RegisterTurn:
"""One program step: a per-register affine op or a coupled transfer."""
@ -97,7 +110,7 @@ class RegisterTurn:
class MultiRegisterProgram:
seeds: tuple[tuple[str, float], ...] # (entity, seed) in graph order
turns: tuple[RegisterTurn, ...]
answer_entity: str
answer_entity: str | None # None ⇒ certified summation over all registers
answer_unit: str
@ -111,9 +124,10 @@ class MultiRegisterRecord:
converged: bool
step: tuple[str, float, float] # (kind, scale, offset)
prev_record_digest: str
operand_source_digest: str = "" # summation only: binds the operand to its source state
def _payload(self) -> dict:
return {
payload = {
"sequence_index": int(self.sequence_index),
"entity": self.entity,
"certificate_id": self.certificate_id,
@ -121,6 +135,10 @@ class MultiRegisterRecord:
"step": [self.step[0], repr(float(self.step[1])), repr(float(self.step[2]))],
"prev_record_digest": self.prev_record_digest,
}
# Present only for summation turns, so non-summation record digests are unchanged.
if self.operand_source_digest:
payload["operand_source_digest"] = self.operand_source_digest
return payload
def record_digest(self) -> str:
return _digest(self._payload())
@ -160,10 +178,10 @@ def compile_multi_register_program(graph: MathProblemGraph) -> MultiRegisterProg
units[entity] = possession.quantity.unit
seeds.append((entity, float(possession.quantity.value)))
# A concrete unknown decodes that register; a None unknown ("altogether") is
# the explicit signal for a certified summation over all registers.
answer_entity = graph.unknown.entity
if answer_entity is None:
raise MultiRegisterError("total_unknown_requires_summation")
if answer_entity not in units:
if answer_entity is not None and answer_entity not in units:
raise MultiRegisterError("unknown_entity_not_a_register", entity=answer_entity)
turns: list[RegisterTurn] = []
@ -284,7 +302,33 @@ def execute_multi_register_program(program: MultiRegisterProgram) -> MultiRegist
records.append(record)
index += 1
answer = decode_quantity(registers[program.answer_entity])
if program.answer_entity is None:
# Certified summation over all registers (ruling #1: summation stays in
# the substrate). Each addition is a certified turn; the operand is the
# decode of a certified register state, bound by ``operand_source_digest``
# (= that state's psi_digest) — deterministic re-execution reproduces it,
# so the Python layer cannot tamper with the intermediate value.
order = [entity for entity, _ in program.seeds]
accumulator = registers[order[0]]
for entity in order[1:]:
source_state = registers[entity]
operand = decode_quantity(source_state)
target = _unit(translate_quantity(accumulator, operand))
accumulator, cert = _relax_to(accumulator, target)
if not cert.converged:
raise MultiRegisterError("summation_nonconverged", entity=entity)
record = MultiRegisterRecord(
index, _TOTAL, cert.certificate_id, cert.converged,
("sum_add", 1.0, operand), prev,
operand_source_digest=_state_digest(source_state),
)
prev = record.record_digest()
records.append(record)
index += 1
answer = decode_quantity(accumulator)
else:
answer = decode_quantity(registers[program.answer_entity])
chain = tuple(records)
return MultiRegisterOutcome(
answer=answer,

View file

@ -145,8 +145,8 @@ def test_deterministic() -> None:
# --- Fail-closed refusals (Tier-2a scope) -----------------------------------
def test_refuses_total_unknown_pending_summation() -> None:
# "how many altogether" needs the certified summation turn (2b), not T2a.
def test_total_unknown_compiles_to_summation() -> None:
# "how many altogether" (None unknown) is no longer refused — 2b sums it.
graph = MathProblemGraph(
entities=("Ann", "Bob"),
initial_state=(
@ -156,8 +156,9 @@ def test_refuses_total_unknown_pending_summation() -> None:
operations=(),
unknown=Unknown(None, "apples"),
)
with pytest.raises(MultiRegisterError):
compile_multi_register_program(graph)
program = compile_multi_register_program(graph)
assert program.answer_entity is None # the summation signal
assert abs(execute_multi_register_program(program).answer - 8.0) < 1e-4
def test_refuses_derived_operand() -> None:

View file

@ -0,0 +1,97 @@
"""ADR-0250 2b — certified summation turn (the 6 "altogether" cases) pins.
A total-like unknown (`unknown.entity is None`) compiles to a certified summation
over all registers: each addition is a certified turn, and the operand is the
decode of a certified register state, bound by ``operand_source_digest`` so the
Python layer cannot tamper with the intermediate value (chain of custody). The
capstone pins the full GSM8K dev holdout at 50/50, wrong=0.
"""
from __future__ import annotations
import dataclasses
import json
from pathlib import Path
from generate.math_problem_graph import (
InitialPossession,
MathProblemGraph,
Operation,
Quantity,
Unknown,
graph_from_dict,
)
from evals.turn_program import (
TurnProgramError,
compile_turn_program,
execute_turn_program,
)
from evals.multi_register_program import (
compile_multi_register_program,
execute_multi_register_program,
verify_multi_register_chain,
)
_DEV = Path(__file__).resolve().parents[1] / "evals" / "gsm8k_math" / "dev" / "cases.jsonl"
def _total(entities_seeds, ops=()):
graph = MathProblemGraph(
entities=tuple(e for e, _ in entities_seeds),
initial_state=tuple(InitialPossession(e, Quantity(v, "x")) for e, v in entities_seeds),
operations=tuple(ops),
unknown=Unknown(None, "x"),
)
return execute_multi_register_program(compile_multi_register_program(graph))
# --- Certified summation ----------------------------------------------------
def test_altogether_bare_registers() -> None:
assert abs(_total([("Ann", 5), ("Bob", 3), ("Cal", 4)]).answer - 12.0) < 1e-4
def test_altogether_after_ops() -> None:
# Ann 5*2=10, Bob 3, total 13 — registers are summed AFTER their turns.
out = _total([("Ann", 5), ("Bob", 3)], ops=(Operation("Ann", "multiply", Quantity(2, "factor")),))
assert abs(out.answer - 13.0) < 1e-4
assert out.certified is True
# --- Chain of custody: operand bound to the source state's digest -----------
def test_summation_records_bind_operand_source_digest() -> None:
out = _total([("Ann", 5), ("Bob", 3)])
sum_records = [r for r in out.records if r.entity == "__total__"]
assert len(sum_records) == 1 # one addition for the second register
assert sum_records[0].operand_source_digest != "" # operand is bound to its source state
assert verify_multi_register_chain(out.records) is True
def test_summation_chain_detects_tamper() -> None:
out = _total([("Ann", 5), ("Bob", 3), ("Cal", 4)])
tampered = list(out.records)
tampered[0] = dataclasses.replace(tampered[0], operand_source_digest="forged")
assert verify_multi_register_chain(tampered) is False
# --- Capstone: the full real GSM8K dev holdout, 50/50 wrong=0 ---------------
def test_full_dev_holdout_solved_wrong_zero() -> None:
cases = [json.loads(line) for line in _DEV.read_text().splitlines() if line.strip()]
solved = wrong = 0
for case in cases:
graph = graph_from_dict(case["ground_truth_graph"])
try: # Tier-1 single-accumulator
answer = execute_turn_program(compile_turn_program(graph)).answer
except TurnProgramError: # Tier-2: multi-register (single-entity or summation)
answer = execute_multi_register_program(compile_multi_register_program(graph)).answer
if abs(answer - float(case["expected_answer"])) < 1e-4:
solved += 1
else:
wrong += 1
assert wrong == 0 # the headline: zero wrong on the entire real holdout
assert solved == 50 # Tier-1 (26) + T2a single-entity (18) + summation (6)