Skip to content

zorch.sumcheck.eq.stage

Equality-factored sumcheck roles.

EqSumClaim dataclass

Public sum claim for eq(equality_point, x) * summand(x).

Source code in zorch/sumcheck/eq/stage.py
24
25
26
27
28
29
30
@dataclass(frozen=True)
class EqSumClaim:
    """Public sum claim for ``eq(equality_point, x) * summand(x)``."""

    equality_point: Array
    value: Array
    rounds: int

EqPolyWitness dataclass

Factor tables witnessing an EqSumClaim.

Source code in zorch/sumcheck/eq/stage.py
33
34
35
36
37
@dataclass(frozen=True)
class EqPolyWitness:
    """Factor tables witnessing an ``EqSumClaim``."""

    factors: Array

EqPolyProver

Bases: ProverStage[EqSumClaim, EqPolyWitness, EvaluationClaim, Array]

Prove equality-weighted sumcheck with the equality factor separate.

Source code in zorch/sumcheck/eq/stage.py
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
class EqPolyProver(ProverStage[EqSumClaim, EqPolyWitness, EvaluationClaim, Array]):
    """Prove equality-weighted sumcheck with the equality factor separate."""

    def __init__(
        self,
        summand: SumcheckSummand,
        *,
        challenges: ChallengePolicy,
    ) -> None:
        self.summand = summand
        self.challenges = challenges
        self.degree = summand.degree + 1
        self.verifier_round = SumcheckRound(self.degree, challenges)

    def prove(
        self,
        claim: EqSumClaim,
        witness: EqPolyWitness,
        transcript: Transcript,
    ) -> ProveResult[EvaluationClaim, Array]:
        _check_claim(claim)
        pre = transcript
        domain = natural_domain(self.degree, witness.factors.dtype)
        prover_round = EqPolyRound(
            self.summand,
            claim.equality_point,
            domain,
            challenges=self.challenges,
        )
        state: EqPolyState = (
            witness.factors,
            fnp.ones(1, dtype=witness.factors.dtype),
        )
        carry, transcript, messages = fold_rounds(
            prover_round,
            initial_claim(state, claim.value, claim.rounds),
            transcript,
            claim.rounds,
        )
        reduced = carry.claim
        return ProveResult(
            EvaluationClaim(reduced.point, reduced.value),
            fnp.stack(messages),
            transcript,
        )

EqPolyVerifier

Bases: VerifierStage[EqSumClaim, EvaluationClaim, Array]

Verify equality-weighted sumcheck.

Source code in zorch/sumcheck/eq/stage.py
 95
 96
 97
 98
 99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
class EqPolyVerifier(VerifierStage[EqSumClaim, EvaluationClaim, Array]):
    """Verify equality-weighted sumcheck."""

    def __init__(
        self,
        summand: SumcheckSummand,
        *,
        challenges: ChallengePolicy,
    ) -> None:
        self.verifier_round = SumcheckRound(summand.degree + 1, challenges)

    def verify(
        self,
        claim: EqSumClaim,
        reduction_proof: Array,
        transcript: Transcript,
    ) -> VerifyResult[EvaluationClaim]:
        _check_claim(claim)
        if reduction_proof.shape[0] != claim.rounds:
            raise ValueError(
                f"expected {claim.rounds} sumcheck rounds, "
                f"got {reduction_proof.shape[0]}"
            )
        point, value, transcript, ok = verify(
            self.verifier_round, claim.value, reduction_proof, transcript
        )
        return VerifyResult(EvaluationClaim(point, value), transcript, ok)