summaryrefslogtreecommitdiff
path: root/worldalign/synth_triangle_recovery.py
blob: a9bca61e62598abf2f18b793724e0a6d5ad75387 (plain)
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
168
169
170
171
172
173
"""Blind recovery under the triangle energy.

The triangle gate holds the true assignment at N=512 with no counterfeit
below it. The search question is separate and now worth asking again:
tempering with third-order swap deltas, from random starts, with the
hidden order behind a shuffle. Recovery accuracy is the end-to-end world
matching measurement.
"""

from __future__ import annotations

import argparse
import json
from pathlib import Path

import torch

from .common import read_json, seed_everything, write_json
from .synth_triangle_gate import TriangleEnergy, build_fields, standardized


def parse_args() -> argparse.Namespace:
    parser = argparse.ArgumentParser()
    parser.add_argument("--data-dir", default="artifacts/synth_v0")
    parser.add_argument("--split", choices=["val", "test"], default="test")
    parser.add_argument("--samples", type=int, default=512)
    parser.add_argument("--merge-distance", type=float, default=30.0)
    parser.add_argument("--vision-views", type=int, default=4)
    parser.add_argument("--triples", type=int, default=400000)
    parser.add_argument("--pair-weight", type=float, default=1.0)
    parser.add_argument("--triangle-weight", type=float, default=1.0)
    parser.add_argument("--replicas", type=int, default=6)
    parser.add_argument("--rounds", type=int, default=4000)
    parser.add_argument("--proposals", type=int, default=48)
    parser.add_argument("--temp-high", type=float, default=2e-2)
    parser.add_argument("--temp-low", type=float, default=1e-4)
    parser.add_argument("--exchange-every", type=int, default=20)
    parser.add_argument("--greedy-rounds", type=int, default=3000)
    parser.add_argument("--device", default="cuda:3")
    parser.add_argument("--seed", type=int, default=20260731)
    parser.add_argument(
        "--output", default="artifacts/synth_v0/triangle_recovery.json"
    )
    return parser.parse_args()


def main() -> None:
    args = parse_args()
    seed_everything(args.seed)
    manifest = read_json(Path(args.data_dir, "manifest.json"))
    rows = manifest[args.split][: args.samples]
    visual_field, text_field = build_fields(args, rows, manifest)

    device = torch.device(args.device)
    size = len(rows)
    generator = torch.Generator().manual_seed(args.seed)
    hidden = torch.randperm(size, generator=generator)
    truth = torch.argsort(hidden).to(device)
    text = standardized(text_field[hidden][:, hidden].double().to(device)).float()
    visual = standardized(visual_field.double().to(device)).float()

    triples = torch.randint(0, size, (args.triples, 3), generator=generator)
    triples = triples[
        (triples[:, 0] != triples[:, 1])
        & (triples[:, 1] != triples[:, 2])
        & (triples[:, 0] != triples[:, 2])
    ].to(device)
    energy = TriangleEnergy(
        text, visual, triples, args.pair_weight, args.triangle_weight
    )
    true_energy = energy.total(truth)

    temperatures = torch.logspace(
        torch.log10(torch.tensor(args.temp_low)),
        torch.log10(torch.tensor(args.temp_high)),
        args.replicas,
    )
    states = [
        torch.argsort(torch.rand(size, generator=generator)).to(device)
        for _ in range(args.replicas)
    ]
    energies = [energy.total(state) for state in states]
    history = []
    for round_index in range(args.rounds):
        for replica in range(args.replicas):
            temperature = float(temperatures[replica])
            for _ in range(args.proposals):
                p = int(torch.randint(0, size, (1,), generator=generator))
                q = int(torch.randint(0, size, (1,), generator=generator))
                if p == q:
                    continue
                delta = energy.swap_delta(states[replica], p, q)
                threshold = -temperature * float(
                    torch.rand(1, generator=generator).clamp_min(1e-12).log()
                )
                if delta < threshold:
                    states[replica][[p, q]] = states[replica][[q, p]]
                    energies[replica] += delta
        if round_index % args.exchange_every == 0:
            for replica in range(args.replicas - 1):
                gap = (energies[replica] - energies[replica + 1]) * (
                    1.0 / float(temperatures[replica])
                    - 1.0 / float(temperatures[replica + 1])
                )
                if gap > 0 or float(torch.rand(1, generator=generator)) < min(
                    1.0, float(torch.tensor(gap).exp())
                ):
                    states[replica], states[replica + 1] = (
                        states[replica + 1],
                        states[replica],
                    )
                    energies[replica], energies[replica + 1] = (
                        energies[replica + 1],
                        energies[replica],
                    )
        if round_index % 200 == 0:
            cold = min(range(args.replicas), key=lambda r: energies[r])
            record = {
                "round": round_index,
                "cold_energy": energies[cold],
                "cold_accuracy": float(
                    (states[cold].cpu() == truth.cpu()).float().mean()
                ),
                "energy_over_true": energies[cold] / true_energy - 1.0,
            }
            history.append(record)
            print(json.dumps(record))

    # Greedy polish of the coldest replica.
    cold = min(range(args.replicas), key=lambda r: energies[r])
    current = states[cold].clone()
    for _ in range(args.greedy_rounds):
        best_delta, best_pair = 0.0, None
        for _ in range(96):
            p = int(torch.randint(0, size, (1,), generator=generator))
            q = int(torch.randint(0, size, (1,), generator=generator))
            if p == q:
                continue
            delta = energy.swap_delta(current, p, q)
            if delta < best_delta:
                best_delta, best_pair = delta, (p, q)
        if best_pair is None:
            break
        p, q = best_pair
        current[[p, q]] = current[[q, p]]

    final = {
        "accuracy": float((current.cpu() == truth.cpu()).float().mean()),
        "energy": energy.total(current),
    }
    final["energy_over_true"] = final["energy"] / true_energy - 1.0
    report = {
        "protocol": (
            "Tempering on the pairwise-plus-triangle energy from random "
            "starts; text order hidden behind a shuffle; hidden truth "
            "scores the outcome only."
        ),
        "samples": size,
        "true_energy": true_energy,
        "chance_accuracy": 1.0 / size,
        "history": history,
        "final_polished": final,
        "replica_accuracies": [
            float((state.cpu() == truth.cpu()).float().mean()) for state in states
        ],
    }
    print(json.dumps({"final": final}))
    write_json(args.output, report)
    print(f"Wrote {args.output}")


if __name__ == "__main__":
    main()