File size: 6,644 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
"""
On-policy Monte Carlo value-head training -- no oracle, no bootstrapping,
no target network. Walk the environment with the current value head's own
epsilon-greedy policy, and label every visited state along a walk that
actually reached a solved state with its REALIZED return (what happened,
not a network's own possibly-wrong estimate of what happens next). This is
the only value-training method used for ConnectX, since the real 7x6 board
has no tractable exact solver to regress against.

`unsolved_penalty` is what makes this work for an ADVERSARIAL domain
specifically: `env.step` can return `done=True` on a LOSS (the opponent
won) or a draw, not just our own win. Discarding every one of those walks
(the natural default for a single-agent puzzle, where "unsolved" just means
"further away") would mean the value head never sees a single labeled
example of "this leads to losing" -- exactly the signal an adversarial
domain needs to learn to avoid bad moves.
"""
import collections

import torch
import torch.nn as nn

from .train_utils import states_to_tensor, DEVICE


def _observed_states_to_tensor(env, states):
    return states_to_tensor([env.observe(s) for s in states])


def _greedy_walk_action(env, model, normalizer, state, rng, epsilon):
    """Epsilon-greedy action choice using the value head's own 1-step
    lookahead, scored the same way `search.plan_action`'s depth=1 does
    (`-reward + value`, using the REAL reward/next-state from `env.step`,
    not an imagined one) -- so training optimizes for the exact decision
    rule actually used at inference."""
    legal = [a for a in range(env.num_actions) if env.is_legal(state, a)]
    if rng.random() < epsilon:
        return rng.choice(legal)
    next_states, rewards = [], []
    for a in legal:
        ns, r, _done = env.step(state, a)
        next_states.append(ns)
        rewards.append(r)
    z = model.encode(normalizer.normalize(_observed_states_to_tensor(env, next_states).to(DEVICE)))
    values = model.evaluate(z)
    rewards_t = torch.tensor(rewards, dtype=torch.float32, device=DEVICE)
    scores = -rewards_t + values
    return legal[torch.argmin(scores).item()]


def generate_mc_walks(env, model, normalizer, rng, n_problems, max_steps, epsilon,
                       unsolved_penalty=None, **problem_kwargs):
    """Complete walks (up to max_steps) using the current value head's
    epsilon-greedy policy. A walk that reaches `is_solved` labels every
    visited state with its real steps-remaining. A walk that doesn't
    (a loss, a draw, or simply running out of steps) is discarded UNLESS
    `unsolved_penalty` is set, in which case every state along it gets that
    fixed, uniform label instead -- deliberately not scaled by how early or
    late the walk went wrong; a state that's actually fine keeps
    reappearing in OTHER (solved) walks too, so its label self-corrects
    over many rounds rather than staying pinned to one bad walk's worst
    case."""
    labeled = []
    for _ in range(n_problems):
        state, _ = env.random_problem(rng, **problem_kwargs)
        path_states = [state]
        for _ in range(max_steps):
            if env.is_solved(state):
                break
            a = _greedy_walk_action(env, model, normalizer, state, rng, epsilon)
            state, _reward, done = env.step(state, a)
            path_states.append(state)
            if done:
                break
        if env.is_solved(path_states[-1]):
            T = len(path_states) - 1
            for t, s in enumerate(path_states):
                labeled.append((s, float(T - t)))
        elif unsolved_penalty is not None:
            for s in path_states:
                labeled.append((s, float(unsolved_penalty)))
    return labeled


def train_mc_value_onpolicy(env, model, normalizer, rng, n_rounds=15, n_problems_per_round=400,
                             max_steps=10, epochs_per_round=40, lr=1e-3,
                             epsilon_start=1.0, epsilon_end=0.05, warmup_rounds=3,
                             replay_capacity=20000, min_replay_before_train=30,
                             unsolved_penalty=None, verbose_every=5, **problem_kwargs):
    """`epsilon_start=1.0` + `warmup_rounds`: a freshly-initialized value
    head's "greedy" choice is pure noise, which can be WORSE than uniform
    random at stumbling into a solved state by chance. Holding epsilon at
    1.0 (pure random walk) for the first `warmup_rounds` guarantees a
    baseline solve rate to bootstrap training from, before annealing toward
    exploitation."""
    opt = torch.optim.Adam(model.value.parameters(), lr=lr)
    buffer = collections.deque(maxlen=replay_capacity)

    for round_idx in range(n_rounds):
        if round_idx < warmup_rounds:
            epsilon = 1.0
        else:
            progress = (round_idx - warmup_rounds) / max(1, n_rounds - 1 - warmup_rounds)
            epsilon = epsilon_start + (epsilon_end - epsilon_start) * progress
        labeled = generate_mc_walks(env, model, normalizer, rng, n_problems_per_round, max_steps, epsilon,
                                     unsolved_penalty=unsolved_penalty, **problem_kwargs)
        buffer.extend(labeled)
        if len(buffer) < min_replay_before_train:
            if verbose_every:
                print(f"  [mc-onpolicy] round {round_idx+1}/{n_rounds}  epsilon={epsilon:.2f}  "
                      f"only {len(buffer)} labeled states so far (need {min_replay_before_train}) -- skipping fit")
            continue

        all_data = list(buffer)
        states_t = _observed_states_to_tensor(env, [s for s, _ in all_data]).to(DEVICE)
        returns_t = torch.tensor([r for _, r in all_data], dtype=torch.float32, device=DEVICE)

        model.value_target_mean.copy_(returns_t.mean())
        model.value_target_std.copy_(returns_t.std().clamp(min=1e-3))
        returns_norm = (returns_t - model.value_target_mean) / model.value_target_std

        with torch.no_grad():
            z_states = model.encode(normalizer.normalize(states_t))

        n = len(all_data)
        for _epoch in range(epochs_per_round):
            perm = torch.randperm(n, device=DEVICE)
            pred_norm = model.value(z_states[perm])
            loss = nn.functional.mse_loss(pred_norm, returns_norm[perm])
            opt.zero_grad()
            loss.backward()
            opt.step()

        if verbose_every and (round_idx + 1) % verbose_every == 0:
            print(f"  [mc-onpolicy] round {round_idx+1}/{n_rounds}  epsilon={epsilon:.2f}  "
                  f"buffer_size={n}  loss={loss.item():.4f}  return_mean={model.value_target_mean.item():.2f}")