From 9136c1819b8c8a4cc18c636f909c78339edcc182 Mon Sep 17 00:00:00 2001 From: Shay Date: Sat, 18 Jul 2026 16:35:09 -0700 Subject: [PATCH] feat(reader-arc): compare_multiplicative compiler tier (increment 1, compiler half) Cross-register affine definition: actor := factor x reference. Reads the reference register's field STATE, dilates by -ln(factor), writes the actor. Reuses quantity_kernel + the multi-register framework + the 2b summation. Amendments (PR #76): 1. Summation registers-driven: MultiRegisterProgram.register_order (seed order then compare-defined order); execute sums ALL registers in that order, so a compare-DEFINED register (no seed) is included in totals. compile validates the answer target AFTER ops (a concrete unknown may be compare-defined) and propagates the defined register's unit from the reference. 2. Compare record: RegisterTurn.reference; MultiRegisterRecord.source_entity (the reference) + operand_source_digest (reference state digest), both CONDITIONAL in the payload so affine/transfer/summation record digests are byte-identical to before compare. Re-verification reconstructs the source register from records alone. Compare CREATES a quantity (not a transfer) -> exempt from the conservation pin. fraction direction folded in (factor < 1). compare_additive is a later increment (refused). Real official data: the 5 compare parses the reader emits today now solve end-to-end wrong=0 -> corridor real-reach 0/500 -> 5/500 (loop-works proof). 25/25 (9 new + 16 existing regression, digests byte-identical). [Verification]: uv run python -m pytest tests/test_adr_0250_compare.py tests/test_adr_0250_multi_register.py tests/test_adr_0250_summation.py -q --- evals/multi_register_program.py | 84 +++++++++++++--- tests/test_adr_0250_compare.py | 170 ++++++++++++++++++++++++++++++++ 2 files changed, 243 insertions(+), 11 deletions(-) create mode 100644 tests/test_adr_0250_compare.py diff --git a/evals/multi_register_program.py b/evals/multi_register_program.py index 372f3c35..4191bfcc 100644 --- a/evals/multi_register_program.py +++ b/evals/multi_register_program.py @@ -48,7 +48,7 @@ from core.physics.quantity_kernel import ( translate_quantity, ) from core.ports.residual_protocol import GENESIS_DIGEST -from generate.math_problem_graph import MathProblemGraph, Quantity +from generate.math_problem_graph import Comparison, MathProblemGraph, Quantity __all__ = [ "MultiRegisterError", @@ -63,6 +63,7 @@ __all__ = [ _AFFINE = frozenset({"add", "subtract", "multiply", "divide"}) _CONSERVATION_RTOL = 1e-6 +_COMPARE_MULT = "compare_multiply" # internal turn kind for actor := factor × reference class MultiRegisterError(ValueError): @@ -96,7 +97,8 @@ def _state_digest(psi: np.ndarray) -> str: @dataclass(frozen=True, slots=True) class RegisterTurn: - """One program step: a per-register affine op or a coupled transfer.""" + """One program step: a per-register affine op, a coupled transfer, or a + cross-register comparison (``actor := factor × reference``).""" kind: str actor: str @@ -104,6 +106,7 @@ class RegisterTurn: offset: float operand: float target: str | None = None + reference: str | None = None # compare_multiply: the source register read from @dataclass(frozen=True, slots=True) @@ -112,6 +115,8 @@ class MultiRegisterProgram: turns: tuple[RegisterTurn, ...] answer_entity: str | None # None ⇒ certified summation over all registers answer_unit: str + # Deterministic register order for summation: seed order, then compare-defined order. + register_order: tuple[str, ...] = () @dataclass(frozen=True, slots=True) @@ -124,7 +129,8 @@ 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 + operand_source_digest: str = "" # summation/compare: binds the operand to its source state + source_entity: str = "" # compare only: the reference register the dilation read from def _payload(self) -> dict: payload = { @@ -135,9 +141,12 @@ 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. + # Each field is present only when set, so affine/transfer/summation record + # digests are byte-identical to before compare landed. if self.operand_source_digest: payload["operand_source_digest"] = self.operand_source_digest + if self.source_entity: + payload["source_entity"] = self.source_entity return payload def record_digest(self) -> str: @@ -171,21 +180,41 @@ def compile_multi_register_program(graph: MathProblemGraph) -> MultiRegisterProg raise MultiRegisterError("no_initial_state") seeds: list[tuple[str, float]] = [] units: dict[str, str] = {} + order: list[str] = [] # deterministic register order: seeds first, defined later for possession in graph.initial_state: entity = possession.entity if entity in units: raise MultiRegisterError("duplicate_initial_entity", entity=entity) units[entity] = possession.quantity.unit + order.append(entity) 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 not None and answer_entity not in units: - raise MultiRegisterError("unknown_entity_not_a_register", entity=answer_entity) - turns: list[RegisterTurn] = [] for op in graph.operations: + if op.kind == "compare_multiplicative": + # actor := factor × reference — a cross-register definition (AMENDMENT 1/2). + comparison = op.operand + if not isinstance(comparison, Comparison): + raise MultiRegisterError("compare_operand_not_comparison", kind=op.kind) + if comparison.direction not in ("times", "fraction"): + raise MultiRegisterError("compare_additive_out_of_scope", direction=comparison.direction) + reference = comparison.reference_actor + if reference not in units: + raise MultiRegisterError("compare_reference_not_a_register", reference=reference) + if comparison.factor is None: # invariant: times/fraction carry a factor + raise MultiRegisterError("compare_missing_factor", direction=comparison.direction) + factor = float(comparison.factor) + if not math.isfinite(factor) or factor <= 0.0: + raise MultiRegisterError("compare_factor_not_positive", factor=factor) + if op.actor == reference: + raise MultiRegisterError("compare_self_reference", actor=op.actor) + if op.actor not in units: # the comparison DEFINES the actor register + units[op.actor] = units[reference] # unit propagation from the reference + order.append(op.actor) + turns.append( + RegisterTurn(_COMPARE_MULT, op.actor, factor, 0.0, factor, reference=reference) + ) + continue if op.actor not in units: raise MultiRegisterError("actor_not_a_register", actor=op.actor) if op.kind == "transfer": @@ -223,11 +252,18 @@ def compile_multi_register_program(graph: MathProblemGraph) -> MultiRegisterProg raise MultiRegisterError("non_positive_scale", kind="divide", value=value) turns.append(RegisterTurn("divide", op.actor, 1.0 / value, 0.0, value)) + # Validate the answer target AFTER ops, so compare-defined entities count + # (a concrete unknown may be a compare-defined register). None ⇒ summation. + answer_entity = graph.unknown.entity + if answer_entity is not None and answer_entity not in units: + raise MultiRegisterError("unknown_entity_not_a_register", entity=answer_entity) + return MultiRegisterProgram( seeds=tuple(seeds), turns=tuple(turns), answer_entity=answer_entity, answer_unit=graph.unknown.unit, + register_order=tuple(order), ) @@ -284,6 +320,29 @@ def execute_multi_register_program(program: MultiRegisterProgram) -> MultiRegist prev = record.record_digest() records.append(record) index += 1 + elif turn.kind == _COMPARE_MULT: + # actor := factor × reference. Read the reference register's STATE, + # dilate, write the actor (a definition, not a transfer — no + # conservation pin). The record binds the reference entity + its state + # digest so re-verification reconstructs the source from records alone. + reference = turn.reference + if reference is None: # invariant: compile sets a reference for compare + raise MultiRegisterError("compare_missing_reference", actor=turn.actor) + reference_state = registers[reference] + target = _unit(dilate_quantity(reference_state, -math.log(turn.scale))) + new_state, cert = _relax_to(reference_state, target) + if not cert.converged: + raise MultiRegisterError("compare_nonconverged", actor=turn.actor, reference=reference) + registers = {**registers, turn.actor: new_state} + record = MultiRegisterRecord( + index, turn.actor, cert.certificate_id, cert.converged, + (_COMPARE_MULT, turn.scale, 0.0), prev, + operand_source_digest=_state_digest(reference_state), + source_entity=reference, + ) + prev = record.record_digest() + records.append(record) + index += 1 else: target = _unit( translate_quantity( @@ -308,7 +367,10 @@ def execute_multi_register_program(program: MultiRegisterProgram) -> MultiRegist # 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] + # Registers-driven (AMENDMENT 1): sum ALL registers — seeds AND + # compare-defined — in the pinned deterministic order, so a defined + # register is never silently dropped from a total. + order = list(program.register_order) accumulator = registers[order[0]] for entity in order[1:]: source_state = registers[entity] diff --git a/tests/test_adr_0250_compare.py b/tests/test_adr_0250_compare.py new file mode 100644 index 00000000..e960e33d --- /dev/null +++ b/tests/test_adr_0250_compare.py @@ -0,0 +1,170 @@ +"""ADR-0250 reader-arc increment 1 — compare_multiplicative compiler tier. + +`compare_multiplicative` (`actor = factor × reference`) as a cross-register +dilation: read the reference register's field state, dilate, write the actor. +The actor may be *defined* by the comparison (no seed); the registers-driven +summation must include it (amendment 1); the record binds the reference entity ++ its state digest for records-alone re-verification (amendment 2). Compare +*creates* a quantity, so it is exempt from the transfer conservation pin. +""" +from __future__ import annotations + +import dataclasses +import json +from pathlib import Path + +import pytest + +from generate.math_candidate_graph import parse_and_solve +from generate.math_problem_graph import ( + Comparison, + InitialPossession, + MathProblemGraph, + Operation, + Quantity, + Unknown, +) + +from evals.multi_register_program import ( + MultiRegisterError, + compile_multi_register_program, + execute_multi_register_program, + verify_multi_register_chain, +) + +_DEV = Path(__file__).resolve().parents[1] / "evals" / "gsm8k_math" / "holdout_dev" / "v1" / "cases.jsonl" + + +def _cmp_graph(actor, factor, *, direction="times", seed=5, unknown=None, unit="apples"): + """A→seed; actor = factor × A; unknown None (total) or a concrete entity.""" + return MathProblemGraph( + entities=("A", actor), + initial_state=(InitialPossession("A", Quantity(seed, unit)),), + operations=( + Operation(actor, "compare_multiplicative", Comparison("A", None, factor, direction)), + ), + unknown=Unknown(unknown, unit), + ) + + +def _solve(graph): + return execute_multi_register_program(compile_multi_register_program(graph)) + + +# --- The 5 real official compare parses solve end-to-end, wrong=0 ----------- + + +def test_real_compare_cases_solved_wrong_zero() -> None: + cases = [json.loads(line) for line in _DEV.read_text().splitlines() if line.strip()] + solved = wrong = seen = 0 + for case in cases: + try: + result = parse_and_solve(case["problem"]) + except Exception: + continue + graph = result.selected_graph + if graph is None or not any(op.kind == "compare_multiplicative" for op in graph.operations): + continue + seen += 1 + answer = _solve(graph).answer + if abs(answer - float(case["expected_answer"])) < 1e-4: + solved += 1 + else: + wrong += 1 + assert seen == 5 # the compare parses the reader emits today + assert wrong == 0 + assert solved == 5 # corridor real-reach 0/500 → 5/500 (loop-works proof) + + +# --- Amendment 1: compare-defined register in the certified sum + unit prop -- + + +def test_compare_defined_register_included_in_total() -> None: + # A=5, B=3×A=15, total = 20 — B is compare-defined (no seed) yet summed. + assert abs(_solve(_cmp_graph("B", 3.0, unknown=None)).answer - 20.0) < 1e-4 + + +def test_compare_defines_answer_target_with_unit_propagation() -> None: + outcome = _solve(_cmp_graph("B", 3.0, unknown="B", unit="dollars")) + assert abs(outcome.answer - 15.0) < 1e-4 + assert outcome.answer_unit == "dollars" # propagated from the reference + + +def test_fraction_direction() -> None: + assert abs(_solve(_cmp_graph("B", 0.5, direction="fraction", seed=8, unknown="B")).answer - 4.0) < 1e-4 + + +def test_compare_does_not_conserve_and_is_not_rejected() -> None: + # total (20) exceeds the reference (5): compare creates quantity — no pin fires. + outcome = _solve(_cmp_graph("B", 3.0, unknown=None)) + assert outcome.certified is True + + +# --- Amendment 2: the compare record binds the reference (records-alone) ----- + + +def test_compare_record_binds_reference_and_verifies() -> None: + outcome = _solve(_cmp_graph("B", 3.0, unknown="B")) + compare_records = [r for r in outcome.records if r.entity == "B"] + assert len(compare_records) == 1 + rec = compare_records[0] + assert rec.source_entity == "A" # the reference is reconstructable from the record + assert rec.operand_source_digest != "" # bound to the reference state's digest + assert verify_multi_register_chain(outcome.records) is True + + +def test_compare_chain_detects_tamper() -> None: + # 3-record chain (compare B, then a later op) so a non-terminal tamper breaks a link. + graph = MathProblemGraph( + entities=("A", "B"), + initial_state=(InitialPossession("A", Quantity(5, "x")),), + operations=( + Operation("B", "compare_multiplicative", Comparison("A", None, 3.0, "times")), + Operation("A", "add", Quantity(2, "x")), + ), + unknown=Unknown(None, "x"), + ) + outcome = _solve(graph) + tampered = list(outcome.records) + tampered[0] = dataclasses.replace(tampered[0], source_entity="forged") + assert verify_multi_register_chain(tampered) is False + + +# --- Criterion 4: non-compare records carry no source_entity ---------------- + + +def test_summation_record_has_no_source_entity() -> None: + # A summation record keeps only operand_source_digest — source_entity stays + # empty, so pre-compare record digests are byte-identical. + graph = MathProblemGraph( + entities=("A", "B"), + initial_state=( + InitialPossession("A", Quantity(5, "x")), + InitialPossession("B", Quantity(3, "x")), + ), + operations=(), + unknown=Unknown(None, "x"), + ) + outcome = _solve(graph) + for record in outcome.records: + assert record.source_entity == "" + + +# --- Fail-closed refusals ---------------------------------------------------- + + +def test_refuses_compare_additive() -> None: + # compare_additive ("more"/"fewer") is a later increment — out of scope now. + graph = MathProblemGraph( + entities=("A", "B"), + initial_state=( + InitialPossession("A", Quantity(5, "x")), + InitialPossession("B", Quantity(1, "x")), + ), + operations=( + Operation("B", "compare_additive", Comparison("A", Quantity(2, "x"), None, "more")), + ), + unknown=Unknown("B", "x"), + ) + with pytest.raises(MultiRegisterError): + compile_multi_register_program(graph)