File size: 8,502 Bytes
9ede8c0
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
203
204
205
206
"""
Stage 1: self-supervised training of encoder + dynamics + decoder on random
(state, action, next_state, reward) transitions and multi-step rollouts --
no labels, no oracle, works for any domain implementing `environment.py`'s
`Environment` interface.

(Value-head training -- "stage 2" -- lives in `verifier.py`, since ConnectX
uses on-policy Monte Carlo value learning rather than an oracle-labeled
regression: there's no tractable exact solver at the real 7x6 board scale.)
"""
import torch
import torch.nn as nn

DEVICE = torch.device("cuda" if torch.cuda.is_available() else "cpu")


def sample_legal_action(env, rng, state):
    legal = [i for i in range(env.num_actions) if env.is_legal(state, i)]
    return rng.choice(legal)


def generate_transitions(env, rng, n_problems=4000, walk_len=6, **problem_kwargs):
    """Random-walk from random problem starts, collecting (s, a, s', r)
    along the way -- covers realistic reachable states, not just problem
    starts."""
    transitions = []
    for _ in range(n_problems):
        state, _ = env.random_problem(rng, **problem_kwargs)
        for _ in range(walk_len):
            if env.is_solved(state):
                break
            a_idx = sample_legal_action(env, rng, state)
            next_state, reward, done = env.step(state, a_idx)
            transitions.append((state, a_idx, next_state, reward))
            state = next_state
            if done:
                break
    return transitions


def generate_rollout_sequences(env, rng, n_problems=4000, k=3, **problem_kwargs):
    """K-step (states, actions) sequences for the unrolled/open-loop
    training objective below."""
    sequences = []
    for _ in range(n_problems):
        state, _ = env.random_problem(rng, **problem_kwargs)
        states = [state]
        actions = []
        cur = state
        for _ in range(k):
            a_idx = sample_legal_action(env, rng, cur)
            cur, _, _ = env.step(cur, a_idx)
            actions.append(a_idx)
            states.append(cur)
        sequences.append((states, actions))
    return sequences


class StateNormalizer:
    def __init__(self, states):
        t = torch.tensor(states, dtype=torch.float32)
        self.mean = t.mean(dim=0)
        self.std = t.std(dim=0).clamp_min(1e-3)

    def to(self, device):
        self.mean = self.mean.to(device)
        self.std = self.std.to(device)
        return self

    def normalize(self, state_tensor):
        return (state_tensor - self.mean) / self.std

    def denormalize(self, norm_tensor):
        return norm_tensor * self.std + self.mean


def states_to_tensor(states):
    return torch.tensor([list(s) for s in states], dtype=torch.float32)


def train_stage1(model, normalizer, transitions, sequences=None, k=3,
                  epochs=400, batch_size=256, lr=1e-3):
    """`sequences` (see generate_rollout_sequences) adds a k-step UNROLLED
    loss: the dynamics model is chained k times, feeding its own predicted
    latent back in as input at each step (never re-encoding the real
    intermediate state). This matches how the model is actually used at
    inference (chained multi-step search) -- training on 1-step transitions
    alone leaves compounding rollout error unaddressed."""
    state_dim = normalizer.mean.shape[0]
    states = states_to_tensor([t[0] for t in transitions]).to(DEVICE)
    actions = torch.tensor([t[1] for t in transitions], dtype=torch.long).to(DEVICE)
    next_states = states_to_tensor([t[2] for t in transitions]).to(DEVICE)
    rewards = torch.tensor([t[3] for t in transitions], dtype=torch.float32).to(DEVICE)

    norm_states = normalizer.normalize(states)
    norm_next_states = normalizer.normalize(next_states)

    if sequences is not None:
        seq_states_raw = torch.stack([states_to_tensor(s) for s, _ in sequences]).to(DEVICE)  # [N, k+1, D]
        seq_actions = torch.tensor([a for _, a in sequences], dtype=torch.long).to(DEVICE)  # [N, k]
        norm_seq_states = normalizer.normalize(seq_states_raw.view(-1, state_dim)).view(seq_states_raw.shape)

    n = states.shape[0]
    opt = torch.optim.Adam(
        list(model.encoder.parameters()) + list(model.dynamics.parameters()) + list(model.decoder.parameters()),
        lr=lr,
    )
    sched = torch.optim.lr_scheduler.CosineAnnealingLR(opt, T_max=epochs)
    mse = nn.MSELoss()

    for epoch in range(epochs):
        perm = torch.randperm(n, device=DEVICE)
        total_loss = 0.0
        for i in range(0, n, batch_size):
            idx = perm[i:i + batch_size]
            s, a, s_next, r = norm_states[idx], actions[idx], norm_next_states[idx], rewards[idx]

            z = model.encode(s)
            z_next_target = model.encode(s_next)
            pred_next_z, pred_reward = model.imagine_step(z, a)

            recon = model.reconstruct(z)
            recon_next_from_dynamics = model.reconstruct(pred_next_z)

            loss = (
                mse(recon, s)
                + mse(pred_next_z, z_next_target)
                + mse(recon_next_from_dynamics, s_next)
                + mse(pred_reward, r)
            )

            opt.zero_grad()
            loss.backward()
            opt.step()
            total_loss += loss.item() * idx.shape[0]

        if sequences is not None:
            n_seq = norm_seq_states.shape[0]
            seq_perm = torch.randperm(n_seq, device=DEVICE)
            unrolled_total = 0.0
            for i in range(0, n_seq, batch_size):
                sidx = seq_perm[i:i + batch_size]
                seq_s = norm_seq_states[sidx]  # [B, k+1, D]
                seq_a = seq_actions[sidx]  # [B, k]

                z = model.encode(seq_s[:, 0, :])
                loss_unrolled = 0.0
                for step in range(k):
                    z, _pred_r = model.imagine_step(z, seq_a[:, step])
                    target = seq_s[:, step + 1, :]
                    target_z = model.encode(target)
                    loss_unrolled = loss_unrolled + mse(z, target_z) + mse(model.reconstruct(z), target)
                loss_unrolled = loss_unrolled / k

                opt.zero_grad()
                loss_unrolled.backward()
                opt.step()
                unrolled_total += loss_unrolled.item() * sidx.shape[0]

        sched.step()
        if (epoch + 1) % 20 == 0 or epoch == 0:
            msg = f"  [stage1] epoch {epoch+1:3d}/{epochs}  loss={total_loss/n:.4f}"
            if sequences is not None:
                msg += f"  unrolled_loss={unrolled_total/n_seq:.4f}"
            print(msg)


@torch.no_grad()
def eval_stage1(model, normalizer, transitions):
    """Decoder reconstruction / dynamics-rollout exact-match accuracy (after
    rounding to the nearest integer) -- the diagnostic that checks the
    latent isn't collapsing to something the decoder can't read back out."""
    states = states_to_tensor([t[0] for t in transitions]).to(DEVICE)
    actions = torch.tensor([t[1] for t in transitions], dtype=torch.long).to(DEVICE)
    next_states = states_to_tensor([t[2] for t in transitions]).to(DEVICE)

    norm_states = normalizer.normalize(states)
    z = model.encode(norm_states)
    recon = normalizer.denormalize(model.reconstruct(z))
    pred_next_z, _ = model.imagine_step(z, actions)
    recon_next = normalizer.denormalize(model.reconstruct(pred_next_z))

    recon_acc = (recon.round() == states).all(dim=1).float().mean().item()
    dyn_acc = (recon_next.round() == next_states).all(dim=1).float().mean().item()
    return {"decoder_recon_exact_acc": recon_acc, "dynamics_rollout_exact_acc": dyn_acc}


@torch.no_grad()
def eval_multistep_rollout(model, normalizer, sequences, k):
    """Chained (open-loop) rollout exact-match accuracy at each step 1..k --
    exposes compounding error the way a single-step eval can't."""
    state_dim = normalizer.mean.shape[0]
    seq_states_raw = torch.stack([states_to_tensor(s) for s, _ in sequences]).to(DEVICE)
    seq_actions = torch.tensor([a for _, a in sequences], dtype=torch.long).to(DEVICE)
    norm_seq_states = normalizer.normalize(seq_states_raw.view(-1, state_dim)).view(seq_states_raw.shape)

    z = model.encode(norm_seq_states[:, 0, :])
    results = {}
    for step in range(k):
        z, _ = model.imagine_step(z, seq_actions[:, step])
        recon = normalizer.denormalize(model.reconstruct(z))
        real = seq_states_raw[:, step + 1, :]
        acc = (recon.round() == real).all(dim=1).float().mean().item()
        results[f"k={step+1}"] = acc
    return results