File size: 6,814 Bytes
8815366
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
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
93
94
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
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
#!/usr/bin/env python3
"""Exact finite-sample audit of the Appendix-B printed bit threshold.

For an encoded one-bit, Appendix B assigns probability p1=3/(5r) and then
decodes one when the empirical mass is at least the same number.  Therefore a
single bit is K~Binomial(m,p1) and its exact recovery probability is
P[K >= ceil(m p1)], which converges to 1/2 rather than 1.  Since all-bit
recovery is a subset of recovery of any particular one-bit, this marginal
calculation is also a rigorous upper bound for every all-ones target.
"""

from __future__ import annotations

import argparse
import csv
import hashlib
import json
import math
import platform
import time
from datetime import datetime, timezone
from importlib.metadata import version
from pathlib import Path

import matplotlib

matplotlib.use("Agg")
import matplotlib.pyplot as plt
import numpy as np
from scipy.stats import binom


def threshold_count(m: int, threshold: float) -> int:
    # Guard against a binary floating representation just above an integer.
    return int(math.ceil(m * threshold - 1e-12))


def success_one_bit(m: int, r: int, threshold: float) -> float:
    p_one = 3.0 / (5.0 * r)
    k = threshold_count(m, threshold)
    return float(binom.sf(k - 1, m, p_one))


def success_zero_bit(m: int, r: int, threshold: float) -> float:
    p_zero = 2.0 / (5.0 * r)
    k = threshold_count(m, threshold)
    return float(binom.cdf(k - 1, m, p_zero))


def paper_order_samples(r: int, epsilon: float, constant: float) -> int:
    return int(math.ceil(constant * r * r * (math.log(2.0 * r) + math.log(1.0 / epsilon))))


def main() -> None:
    parser = argparse.ArgumentParser()
    parser.add_argument("--output", type=Path, default=Path("outputs/exact_decoder_threshold"))
    parser.add_argument("--constant", type=float, default=40.0)
    args = parser.parse_args()
    args.output.mkdir(parents=True, exist_ok=True)
    started = datetime.now(timezone.utc)
    t0 = time.perf_counter()

    rows: list[dict] = []
    epsilons = [1e-1, 1e-2, 1e-4, 1e-8, 1e-12]
    for r in [1, 4, 8, 16, 32]:
        p_one = 3.0 / (5.0 * r)
        printed = p_one
        midpoint = 2.5 / (5.0 * r)
        for epsilon in epsilons:
            m = paper_order_samples(r, epsilon, args.constant)
            printed_one = success_one_bit(m, r, printed)
            midpoint_one = success_one_bit(m, r, midpoint)
            printed_zero = success_zero_bit(m, r, printed)
            midpoint_zero = success_zero_bit(m, r, midpoint)
            rows.append({
                "r": r,
                "epsilon": epsilon,
                "samples_m": m,
                "p_one": p_one,
                "p_zero": 2.0 / (5.0 * r),
                "printed_threshold": printed,
                "midpoint_threshold": midpoint,
                "printed_one_bit_success_exact": printed_one,
                "midpoint_one_bit_success_exact": midpoint_one,
                "printed_zero_bit_success_exact": printed_zero,
                "midpoint_zero_bit_success_exact": midpoint_zero,
                "all_ones_success_upper_bound_printed": printed_one,
                "printed_exceeds_two_thirds": int(printed_one >= 2.0 / 3.0),
            })

    csv_path = args.output / "exact_binomial_audit.csv"
    with csv_path.open("w", newline="") as handle:
        writer = csv.DictWriter(handle, fieldnames=list(rows[0]))
        writer.writeheader()
        writer.writerows(rows)

    r1 = [row for row in rows if row["r"] == 1]
    summary = {
        "paper": "Positive Distribution Shift as a Framework for Understanding Tractable Learning",
        "openreview_id": "DkLQ40hTlt",
        "claim": "Appendix-B printed decoder threshold for Theorem 3.2",
        "exact_identity": "K~Binomial(m,3/(5r)); printed success=P[K>=ceil(3m/(5r))]",
        "logical_consequence": (
            "All-bit success for an all-ones code is at most this one-bit marginal. "
            "For r=1 the bound is exact, so the printed decoder cannot attain "
            "the usual >=2/3 success probability as epsilon decreases."
        ),
        "rows": len(rows),
        "all_printed_probabilities_below_two_thirds": bool(
            all(row["printed_one_bit_success_exact"] < 2.0 / 3.0 for row in rows)
        ),
        "r1_smallest_epsilon": r1[-1],
        "maximum_printed_one_bit_success": float(
            max(row["printed_one_bit_success_exact"] for row in rows)
        ),
        "minimum_midpoint_one_bit_success": float(
            min(row["midpoint_one_bit_success_exact"] for row in rows)
        ),
        "asymptotic_limit_printed_one_bit": 0.5,
        "scope": (
            "Exact audit of the printed decoder, not a disproof of a repaired "
            "midpoint decoder or of every possible f-PDS universality theorem."
        ),
        "execution": {
            "started_at_utc": started.isoformat(),
            "completed_at_utc": datetime.now(timezone.utc).isoformat(),
            "elapsed_seconds": time.perf_counter() - t0,
            "python": platform.python_version(),
            "platform": platform.platform(),
            "packages": {name: version(name) for name in ["numpy", "scipy", "matplotlib"]},
        },
    }
    (args.output / "summary.json").write_text(json.dumps(summary, indent=2) + "\n")

    fig, axes = plt.subplots(1, 2, figsize=(10.5, 4.0))
    for r in [1, 4, 8, 16, 32]:
        rr = [row for row in rows if row["r"] == r]
        axes[0].semilogx(
            [row["samples_m"] for row in rr],
            [row["printed_one_bit_success_exact"] for row in rr],
            marker="o", label=f"r={r}",
        )
        axes[1].semilogx(
            [row["samples_m"] for row in rr],
            [row["midpoint_one_bit_success_exact"] for row in rr],
            marker="o", label=f"r={r}",
        )
    axes[0].axhline(0.5, color="black", linestyle="--", linewidth=1, label="limit 1/2")
    axes[0].axhline(2.0 / 3.0, color="red", linestyle=":", linewidth=1, label="2/3 target")
    axes[0].set(title="Printed threshold = one-bit mean", xlabel="sample count m", ylabel="exact one-bit recovery", ylim=(0.45, 0.70))
    axes[1].set(title="Correct midpoint threshold", xlabel="sample count m", ylabel="exact one-bit recovery", ylim=(0.45, 1.02))
    axes[0].legend(ncol=2, fontsize=8)
    axes[1].legend(ncol=2, fontsize=8)
    fig.tight_layout()
    fig.savefig(args.output / "exact_decoder_threshold.png", dpi=180)
    plt.close(fig)

    manifest = {}
    for name in ["exact_binomial_audit.csv", "summary.json", "exact_decoder_threshold.png"]:
        manifest[name] = hashlib.sha256((args.output / name).read_bytes()).hexdigest()
    (args.output / "manifest.json").write_text(json.dumps(manifest, indent=2) + "\n")
    print(json.dumps(summary, indent=2))


if __name__ == "__main__":
    main()