rootxhacker commited on
Commit
6ec8e27
·
verified ·
1 Parent(s): 6c69e10

v5.2: spatial policy head, 6x6 action grid aligned with FSQ token grid, advantage normalization removed (entropy-collapse root cause), grad clip 1.0, standard entropy bonus; smoke-validated end-to-end on CPU

Browse files
Files changed (1) hide show
  1. train_imagination_v5.py +1 -638
train_imagination_v5.py CHANGED
@@ -1,638 +1 @@
1
- #!/usr/bin/env python3
2
- """
3
- train_imagination_v5.py — imagination model for computer use, v5 (actor-representation fix).
4
-
5
- History this build answers:
6
- v1/fix1b: RL-in-imagination only -> actor had no path from instruction to click location;
7
- real-env success 0.00, entropy collapse.
8
- v4: BC added -> still 0.10 success; actor inputs (25 coarse FSQ codes + hashed word
9
- buckets) cannot represent "where is the 'Ok' button".
10
- v5 fix: actor sees frozen DINOv2-small patch/CLS features (semantic, position-preserving)
11
- + frozen SmolLM2-135M embedding of the REAL instruction text. World model stays a
12
- small transformer over FSQ codes; reward ensemble is pessimistic (min of 2 heads).
13
- Also fixes the WM copy flaw: CE now predicts codes[:,t+1] from output position t
14
- (v2/v3 compared a position's output to the codes fed as input at that same position).
15
- v5.1: fix classifier device crash (H on cpu vs clf on cuda), fix dream-vs-real label mixing
16
- (real embeddings now come from the pre-roll context window, not post-imagination),
17
- quadratic entropy floor, BC-stage checkpoint pushed before RL so a crash loses nothing.
18
-
19
- Pipeline: A collect (DOM oracle + random) -> B FSQ tokenizer on frozen DINOv2 patch features
20
- -> C world model (val-split early stop) -> D BC on oracle demos -> E RL in imagination
21
- (pessimistic min-of-2-heads reward, REINFORCE + critic, entropy target) -> real-env evals ->
22
- F dream-vs-real classifier (AUC).
23
-
24
- Env facts (sandbox-verified 2026-09-19, miniwob 1.1.0):
25
- obs keys: dom_elements, fields, screenshot, utterance; screenshot (210,160,3) uint8 (H210,W160)
26
- action {"action_type": 2, "coords": float32[x,y]} = CLICK_COORDS over Box(0,[160,210])
27
- dom elements: flat left/top/width/height (1-elem float32 arrays), text match is case-sensitive
28
- Usage:
29
- python train_imagination_v5.py --smoke # tiny CPU run, all phases
30
- python train_imagination_v5.py --device cuda # real run (t4-small, ~1.5h)
31
- """
32
- import argparse, math, os, random, time
33
- import numpy as np
34
- import torch, torch.nn as nn, torch.nn.functional as F
35
-
36
- os.environ.setdefault("MINIWOB_CHROME_BINARY", "/usr/bin/chromium")
37
- os.environ.setdefault("MINIWOB_CHROMEDRIVER", "/usr/bin/chromedriver")
38
-
39
- import gymnasium, miniwob
40
- from PIL import Image
41
- from transformers import AutoModel, AutoTokenizer
42
-
43
- gymnasium.register_envs(miniwob)
44
-
45
- # ── constants ────────────────────────────────────────────────────────────────
46
- SCREEN_W, SCREEN_H = 160, 210 # full screenshot, interactive page
47
- IMG = 224 # DINOv2 input (square resize)
48
- PP = 14 # DINOv2 patch size
49
- GP = IMG // PP # 16x16 = 256 patch tokens
50
- POOL = 2 # avg-pool patch grid -> 8x8 = 64 tokens/frame
51
- NT = (GP // POOL) ** 2 # 64
52
- TD, LEVELS = 6, 5 # FSQ dims/levels -> vocab 5^6 = 15625
53
- VOCAB = LEVELS ** TD
54
- GRID = 16 # 16x16 = 256 click cells
55
- N_ACTIONS = GRID * GRID
56
- TEXT_DIM = 576 # SmolLM2-135M hidden size (verified 576)
57
- MAX_EP_STEPS = 30
58
- L_CTX = 3 # real context frames before imagining
59
- DM, LAYERS, HEADS = 256, 4, 4
60
- W_LEN = 16
61
- PUNCT = str.maketrans("", "", '"\'.,!?:;()')
62
-
63
- TASKS = ["click-button-v1", "click-test-2-v1", "click-link-v1"]
64
- HF_REPO = "rootxhacker/miniwob-imagination"
65
-
66
-
67
- def coords_of(cell):
68
- col, row = cell % GRID, cell // GRID
69
- return np.array([(col + 0.5) * SCREEN_W / GRID, (row + 0.5) * SCREEN_H / GRID], np.float32)
70
-
71
-
72
- def screen_of(raw):
73
- img = raw["screenshot"] # (210,160,3)
74
- return np.asarray(Image.fromarray(img).resize((IMG, IMG)), np.uint8)
75
-
76
-
77
- def instr_text(raw):
78
- words = [str(raw["utterance"]).translate(PUNCT)]
79
- words += [str(f) for _, f in (raw["fields"] or ()) if f]
80
- return " ".join(words)[:256]
81
-
82
-
83
- def oracle_from_raw(raw):
84
- """(grid_cell, click_action) for the element named by fields['target'], else None.
85
- Exact text match first, then case-insensitive fallback."""
86
- fields = dict(raw["fields"]) if raw["fields"] else {}
87
- target = str(fields.get("target", "")).strip()
88
- if not target:
89
- return None
90
- best = None
91
- for el in raw["dom_elements"]:
92
- text = str(el.get("text", "")).strip()
93
- if text != target and text.lower() != target.lower():
94
- continue
95
- l, t = el.get("left"), el.get("top")
96
- w, h = el.get("width"), el.get("height")
97
- if l is None or t is None:
98
- continue
99
- l, t, w, h = (float(np.ravel(v)[0]) if v is not None else 0.0 for v in (l, t, w, h))
100
- cx, cy = l + w / 2, t + h / 2
101
- if not (0 <= cx < SCREEN_W and 0 <= cy < SCREEN_H):
102
- continue
103
- col = min(GRID - 1, max(0, int(cx * GRID / SCREEN_W)))
104
- row = min(GRID - 1, max(0, int(cy * GRID / SCREEN_H)))
105
- exact = text == target
106
- cand = (row * GRID + col, exact)
107
- if best is None or (exact and not best[1]):
108
- best = cand
109
- if best is None:
110
- return None
111
- cell = best[0]
112
- return cell, {"action_type": 2, "coords": coords_of(cell)}
113
-
114
-
115
- def collect(tasks, n_demo, n_rand, seed):
116
- rng = random.Random(seed)
117
- trajs = []
118
- for task in tasks:
119
- env = gymnasium.make(f"miniwob/{task}")
120
- for kind, n in (("oracle", n_demo), ("random", n_rand)):
121
- rets = []
122
- for ep in range(n):
123
- raw, _ = env.reset(seed=rng.randrange(2 ** 31))
124
- frames, acts, rews, dones = [screen_of(raw)], [], [], []
125
- ret = 0.0
126
- for t in range(MAX_EP_STEPS):
127
- oa = oracle_from_raw(raw) if kind == "oracle" else None
128
- if oa is not None:
129
- cell, act = oa
130
- else:
131
- cell = rng.randrange(N_ACTIONS)
132
- act = {"action_type": 2, "coords": coords_of(cell)}
133
- raw2, r, term, trunc, _ = env.step(act)
134
- r = float(r); ret += r
135
- acts.append(cell); rews.append(r); dones.append(float(term or trunc))
136
- frames.append(screen_of(raw2)); raw = raw2
137
- if term or trunc:
138
- break
139
- rets.append(ret)
140
- trajs.append(dict(task=task, instr=instr_text(raw),
141
- frames=np.stack(frames), actions=np.array(acts, np.int64),
142
- rewards=np.array(rews, np.float32), dones=np.array(dones, np.float32)))
143
- print(f"[collect] {task} {kind}: n={n} avg_return={np.mean(rets):.3f}", flush=True)
144
- env.close()
145
- return trajs
146
-
147
-
148
- # ── frozen encoders ──────────────────────────────────────────────────────────
149
- class FrozenEncoders(nn.Module):
150
- def __init__(self, device):
151
- super().__init__()
152
- self.dino = AutoModel.from_pretrained("facebook/dinov2-small").to(device).eval()
153
- self.lm = AutoModel.from_pretrained("HuggingFaceTB/SmolLM2-135M").to(device).eval()
154
- self.tok = AutoTokenizer.from_pretrained("HuggingFaceTB/SmolLM2-135M")
155
- self.tok.pad_token = self.tok.eos_token # required for padding on CPU
156
- for p in self.parameters():
157
- p.requires_grad_(False)
158
-
159
- @torch.no_grad()
160
- def patches(self, screens_uint8):
161
- """(B,IMG,IMG,3) uint8 -> patch features (B,NT,384) and CLS (B,384)."""
162
- x = screens_uint8.float().permute(0, 3, 1, 2) / 127.5 - 1.0
163
- out = self.dino(pixel_values=x).last_hidden_state # (B,257,384)
164
- cls = out[:, 0]
165
- p = out[:, 1:].reshape(-1, GP, GP, 384)
166
- p = p.reshape(-1, GP // POOL, POOL, GP // POOL, POOL, 384).mean((2, 4)) # (B,8,8,384)
167
- return p.reshape(-1, NT, 384), cls
168
-
169
- @torch.no_grad()
170
- def text(self, texts):
171
- texts = [t if t and t.strip() else "click the button" for t in texts]
172
- enc = self.lm_tok_batch(texts)
173
- out = self.lm(input_ids=enc["input_ids"], attention_mask=enc["attention_mask"]).last_hidden_state
174
- m = enc["attention_mask"].unsqueeze(-1).float()
175
- return (out * m).sum(1) / m.sum(1).clamp(min=1) # (B,576) mean-pooled
176
-
177
- def lm_tok_batch(self, texts):
178
- enc = self.tok(texts, padding=True, truncation=True, max_length=W_LEN, return_tensors="pt")
179
- return {k: v.to(next(self.parameters()).device) for k, v in enc.items()}
180
-
181
-
182
- # ── FSQ tokenizer over frozen patch features ────────────────────────────────
183
- class FSQTok(nn.Module):
184
- """trainable projection 384->TD quantized to 5 levels, + decoder back to 384 (feature recon)."""
185
- def __init__(self):
186
- super().__init__()
187
- self.proj = nn.Linear(384, TD)
188
- self.dec = nn.Sequential(nn.Linear(TD, 128), nn.ReLU(), nn.Linear(128, 384))
189
-
190
- def encode(self, feats): # (B,NT,384)
191
- z = torch.tanh(self.proj(feats))
192
- half = (LEVELS - 1) / 2.0
193
- q = (torch.round(z * half) / half).clamp(-1, 1)
194
- dims = ((q + 1) * half).round().long()
195
- flat = torch.zeros(dims.shape[:2], dtype=torch.long, device=z.device)
196
- for d in range(TD):
197
- flat = flat * LEVELS + dims[..., d]
198
- zq = z + (q - z).detach()
199
- return zq, flat
200
-
201
- def forward(self, feats):
202
- zq, codes = self.encode(feats)
203
- recon = self.dec(zq)
204
- return recon, codes, zq
205
-
206
-
207
- # ── world model (small causal transformer over code pairs) ──────────────────
208
- class Attn(nn.Module):
209
- def __init__(self, dm, heads):
210
- super().__init__()
211
- self.h, self.dk = heads, dm // heads
212
- self.qkv, self.proj = nn.Linear(dm, 3 * dm), nn.Linear(dm, dm)
213
-
214
- def forward(self, x, kv=None):
215
- B, T, D = x.shape
216
- q, k, v = self.qkv(x).split(D, -1)
217
- q = q.view(B, T, self.h, self.dk).transpose(1, 2)
218
- k = k.view(B, T, self.h, self.dk).transpose(1, 2)
219
- v = v.view(B, T, self.h, self.dk).transpose(1, 2)
220
- if kv is not None:
221
- k = torch.cat([kv[0], k], 2); v = torch.cat([kv[1], v], 2)
222
- att = q @ k.transpose(-2, -1) / math.sqrt(self.dk)
223
- if kv is None:
224
- mask = torch.triu(torch.ones(T, T, dtype=torch.bool, device=x.device), 1)
225
- att = att.masked_fill(mask[None, None], float("-inf"))
226
- out = (att.softmax(-1) @ v).transpose(1, 2).reshape(B, T, D)
227
- return self.proj(out), (k, v)
228
-
229
-
230
- class Block(nn.Module):
231
- def __init__(self, dm, heads):
232
- super().__init__()
233
- self.ln1, self.ln2 = nn.LayerNorm(dm), nn.LayerNorm(dm)
234
- self.attn = Attn(dm, heads)
235
- self.mlp = nn.Sequential(nn.Linear(dm, 4 * dm), nn.GELU(), nn.Linear(4 * dm, dm))
236
-
237
- def forward(self, x, kv=None):
238
- a, kv = self.attn(self.ln1(x), kv)
239
- x = x + a
240
- return x + self.mlp(self.ln2(x)), kv
241
-
242
-
243
- class WorldModel(nn.Module):
244
- """Layout per step: [frame tokens (NT), action (1)]; frames t>=1 predicted from context."""
245
- def __init__(self):
246
- super().__init__()
247
- self.tok = nn.Embedding(VOCAB, DM)
248
- self.act = nn.Embedding(N_ACTIONS, DM)
249
- self.pos = nn.Embedding(2048, DM)
250
- self.blocks = nn.ModuleList([Block(DM, HEADS) for _ in range(LAYERS)])
251
- self.ln = nn.LayerNorm(DM)
252
- self.head = nn.Linear(DM, VOCAB)
253
- self.r1 = nn.Sequential(nn.Linear(DM, DM), nn.ReLU(), nn.Linear(DM, 1))
254
- self.r2 = nn.Sequential(nn.Linear(DM, DM), nn.ReLU(), nn.Linear(DM, 1))
255
- self.cnt = nn.Sequential(nn.Linear(DM, DM), nn.ReLU(), nn.Linear(DM, 1))
256
-
257
- def _seq(self, codes, actx):
258
- B, L = codes.shape[0], codes.shape[1]
259
- segs = []
260
- for t in range(L - 1):
261
- segs.append(self.tok(codes[:, t]))
262
- segs.append(self.act(actx[:, t]).unsqueeze(1))
263
- segs.append(self.tok(codes[:, L - 1]))
264
- return torch.cat(segs, 1), L
265
-
266
- def forward_full(self, codes, actx):
267
- x, L = self._seq(codes, actx)
268
- P = x.shape[1]
269
- x = x + self.pos(torch.arange(P, device=x.device))[None]
270
- for blk in self.blocks:
271
- x, _ = blk(x)
272
- h = self.ln(x)
273
- pairs = h[:, : (L - 1) * (NT + 1)].reshape(-1, L - 1, NT + 1, DM)
274
- f_h = torch.cat([pairs[:, :, :NT], h[:, (L - 1) * (NT + 1):].unsqueeze(1)], 1) # (B,L,NT,DM)
275
- a_h = pairs[:, :, NT] # (B,L-1,DM)
276
- logits = self.head(f_h[:, :-1]) # predict codes[:,1:]
277
- r1, r2 = self.r1(a_h).squeeze(-1), self.r2(a_h).squeeze(-1)
278
- c = self.cnt(a_h).squeeze(-1)
279
- return logits, r1, r2, c, f_h
280
-
281
- def wm_loss(self, codes, actx, rews, dones):
282
- logits, r1, r2, c, _ = self.forward_full(codes, actx)
283
- B, L = codes.shape[0], codes.shape[1]
284
- ce = F.cross_entropy(logits.reshape(-1, VOCAB), codes[:, 1:].reshape(-1),
285
- reduction="none").view(B, L - 1, NT).mean(-1).mean()
286
- rt = torch.sign(rews) * torch.log1p(torch.abs(rews))
287
- rew = 0.5 * (F.mse_loss(r1, rt) + F.mse_loss(r2, rt))
288
- cont = F.binary_cross_entropy_with_logits(c, 1.0 - dones)
289
- return ce, rew, cont
290
-
291
- def _cache(self, codes, actx):
292
- """KV cache after full context ending on the last frame (no trailing action)."""
293
- x, L = self._seq(codes, actx)
294
- P = x.shape[1]
295
- x = x + self.pos(torch.arange(P, device=x.device))[None]
296
- kvs = []
297
- for blk in self.blocks:
298
- x, kv = blk(x)
299
- kvs.append(kv)
300
- return {"kv": kvs, "pos": P}
301
-
302
- @torch.no_grad()
303
- def imagine_step(self, cache, action):
304
- """feed action -> (next codes (B,NT), r_pess (B,), c (B,)); generates NT codes."""
305
- x = self.act(action).unsqueeze(1)
306
- pos = cache["pos"]
307
- x = x + self.pos(torch.arange(pos, pos + 1, device=x.device))[None]
308
- for i, blk in enumerate(self.blocks):
309
- a, kv = blk.attn(blk.ln1(x), cache["kv"][i])
310
- cache["kv"][i] = kv
311
- x = x + a
312
- x = x + blk.mlp(blk.ln2(x))
313
- cache["pos"] = pos + 1
314
- a_h = self.ln(x[:, -1])
315
- r = torch.minimum(self.r1(a_h), self.r2(a_h)).squeeze(-1) # pessimistic ensemble
316
- c = self.cnt(a_h).squeeze(-1)
317
- codes = []
318
- h = x
319
- for k in range(NT):
320
- nxt = self.head(self.ln(h))[:, -1].argmax(-1)
321
- codes.append(nxt)
322
- if k < NT - 1:
323
- e = self.tok(nxt).unsqueeze(1) + self.pos(torch.arange(cache["pos"], cache["pos"] + 1, device=x.device))[None]
324
- hh = e
325
- for i, blk in enumerate(self.blocks):
326
- a2, kv2 = blk.attn(blk.ln1(hh), cache["kv"][i])
327
- cache["kv"][i] = kv2
328
- hh = hh + a2
329
- hh = hh + blk.mlp(blk.ln2(hh))
330
- cache["pos"] += 1
331
- h = hh
332
- return torch.stack(codes, 1), r, c
333
-
334
-
335
- # ── actor / critic on semantic features ──────────────────────────────────────
336
- class Actor(nn.Module):
337
- def __init__(self):
338
- super().__init__()
339
- d_in = 384 + TEXT_DIM
340
- self.pi = nn.Sequential(nn.Linear(d_in, 256), nn.ReLU(), nn.Linear(256, N_ACTIONS))
341
- self.v = nn.Sequential(nn.Linear(d_in, 256), nn.ReLU(), nn.Linear(256, 1))
342
-
343
- def forward(self, cls, txt):
344
- x = torch.cat([cls, txt], -1)
345
- return self.pi(x), self.v(x).squeeze(-1)
346
-
347
-
348
- # ── classifier (dream vs real) ───────────────────────────────────────────────
349
- def auc_score(scores, labels):
350
- order = np.argsort(scores)
351
- ranks = np.empty(len(scores)); ranks[order] = np.arange(len(scores))
352
- pos = labels.astype(bool)
353
- n_pos, n_neg = pos.sum(), (~pos).sum()
354
- if n_pos == 0 or n_neg == 0:
355
- return float("nan")
356
- return float((ranks[pos].sum() - n_pos * (n_pos - 1) / 2) / (n_pos * n_neg))
357
-
358
-
359
- # ── batching helpers ─────────────────────────────────────────────────────────
360
- def batch_windows(trajs, idxs, L, enc, tok, device):
361
- """Random windows of length L from trajectories. Returns codes (B,L,NT),
362
- actx (B,L-1), rews (B,L-1), dones (B,L-1), ctx screens/acts for BC/RL."""
363
- B = len(idxs)
364
- screens = np.zeros((B, L + 1, IMG, IMG, 3), np.uint8)
365
- codes = np.zeros((B, L, NT), np.int64)
366
- actx = np.zeros((B, L - 1), np.int64)
367
- rews = np.zeros((B, L - 1), np.float32)
368
- dones = np.zeros((B, L - 1), np.float32)
369
- for b, (ti, s) in enumerate(idxs):
370
- tr = trajs[ti]
371
- T = len(tr["actions"])
372
- s = max(0, min(s, T - (L - 1)))
373
- fr = tr["frames"][s:s + L + 1]
374
- if len(fr) < L + 1: # pad by repeating last frame
375
- fr = np.concatenate([fr] + [fr[-1:]] * (L + 1 - len(fr)), 0)
376
- screens[b] = fr
377
- ac = tr["actions"][s:s + L - 1]
378
- actx[b, : len(ac)] = ac
379
- rw = tr["rewards"][s:s + L - 1]
380
- rews[b, : len(rw)] = rw
381
- dn = tr["dones"][s:s + L - 1]
382
- dones[b, : len(dn)] = dn
383
- scr_t = torch.from_numpy(screens).to(device)
384
- with torch.no_grad():
385
- flat = scr_t.reshape(-1, IMG, IMG, 3) # (B*(L+1),IMG,IMG,3)
386
- feats = []
387
- for i in range(0, flat.shape[0], 16):
388
- p, _ = enc.patches(flat[i:i + 16])
389
- feats.append(p)
390
- feats = torch.cat(feats).reshape(B, (L + 1) * NT, 384)
391
- for b in range(B):
392
- _, cd = tok.encode(feats[b].view((L + 1) * NT, 384).unsqueeze(0))
393
- codes[b] = cd[0][: L * NT].view(L, NT).cpu().numpy() # only the L context frames are coded; the (L+1)th frame is for BC/RL features
394
- return (torch.from_numpy(codes).to(device), torch.from_numpy(actx).to(device),
395
- torch.from_numpy(rews).to(device), torch.from_numpy(dones).to(device),
396
- scr_t, feats)
397
-
398
-
399
- # ── main ─────────────────────────────────────────────────────────────────────
400
- def main():
401
- ap = argparse.ArgumentParser()
402
- ap.add_argument("--smoke", action="store_true")
403
- ap.add_argument("--device", default="cpu")
404
- args = ap.parse_args()
405
- device = args.device
406
- S = args.smoke
407
- n_demo, n_rand = (4, 2) if S else (400, 150)
408
- tok_steps, wm_steps, bc_steps, rl_steps = ((30, 60, 60, 50) if S else (800, 1500, 1500, 2000))
409
- n_eval = 2 if S else 4
410
-
411
- trackio = None
412
- if not S:
413
- try:
414
- import trackio
415
- trackio.init(project="miniwob-imagination",
416
- space_id=os.environ.get("TRACKIO_SPACE_ID", "rootxhacker/miniwob-imagination-trackio"))
417
- except Exception as e:
418
- print(f"[trackio] disabled: {e}", flush=True)
419
-
420
- def tlog(metrics, step=0):
421
- if trackio is not None:
422
- try:
423
- trackio.log(metrics, step=step)
424
- except Exception:
425
- pass
426
-
427
- t0 = time.time()
428
- trajs = collect(TASKS[:1] if S else TASKS, n_demo, n_rand, seed=7)
429
- n_train = int(0.9 * len(trajs)); tr_va = trajs[n_train:]; trajs = trajs[:n_train]
430
- print(f"[data] {len(trajs)} train / {len(tr_va)} val trajs in {time.time()-t0:.0f}s", flush=True)
431
- tlog({"data/train_trajs": len(trajs), "data/val_trajs": len(tr_va)})
432
-
433
- enc = FrozenEncoders(device)
434
- tok = FSQTok().to(device)
435
-
436
- # ── B: tokenizer ──
437
- opt = torch.optim.Adam(tok.parameters(), lr=1e-3)
438
- for step in range(tok_steps):
439
- b = random.sample(trajs, min(16, len(trajs)))
440
- scr = torch.from_numpy(np.stack([t["frames"][random.randrange(len(t["frames"]))] for t in b])).to(device)
441
- with torch.no_grad():
442
- feats, _ = enc.patches(scr)
443
- recon, _, _ = tok(feats)
444
- loss = F.mse_loss(recon, feats)
445
- opt.zero_grad(); loss.backward(); opt.step()
446
- print(f"[tok] recon {loss.item():.4f}", flush=True)
447
- tlog({"tok/recon": loss.item()}, step=tok_steps)
448
-
449
- # ── C: world model with val early stop ──
450
- wm = WorldModel().to(device)
451
- opt = torch.optim.Adam(wm.parameters(), lr=3e-4)
452
- def val_loss():
453
- with torch.no_grad():
454
- idxs = [(random.randrange(len(tr_va)), random.randrange(max(1, len(tr_va[0]["actions"]) - L_CTX))) for _ in range(8)]
455
- codes, actx, rews, dones, _, _ = batch_windows(tr_va, idxs, L_CTX, enc, tok, device)
456
- lg, r1, r2, c, _ = wm.forward_full(codes, actx)
457
- ce = F.cross_entropy(lg.reshape(-1, VOCAB), codes[:, 1:].reshape(-1)).item()
458
- rt = torch.sign(rews) * torch.log1p(torch.abs(rews))
459
- rew = 0.5 * (F.mse_loss(r1, rt).item() + F.mse_loss(r2, rt).item())
460
- cont = F.binary_cross_entropy_with_logits(c, 1.0 - dones).item()
461
- return ce + rew + cont
462
- best_v, patience, since = float("inf"), 3, 0
463
- for step in range(wm_steps):
464
- idxs = [(random.randrange(len(trajs)), random.randrange(max(1, len(trajs[0]["actions"]) - L_CTX))) for _ in range(16 if S else 32)]
465
- codes, actx, rews, dones, _, _ = batch_windows(trajs, idxs, L_CTX, enc, tok, device)
466
- ce, rew, cont = wm.wm_loss(codes, actx, rews, dones)
467
- loss = ce + rew + cont
468
- opt.zero_grad(); loss.backward(); opt.step()
469
- if (step + 1) % 50 == 0:
470
- v = val_loss()
471
- print(f"[wm] step {step+1} train {loss.item():.4f} val {v:.4f}", flush=True)
472
- tlog({"wm/train": loss.item(), "wm/val": v}, step=step + 1)
473
- if v < best_v - 1e-4:
474
- best_v, since = v, 0
475
- else:
476
- since += 1
477
- if since >= patience:
478
- print(f"[wm] early stop at step {step+1}, best val {best_v:.4f}", flush=True)
479
- break
480
- print(f"[wm] final train {loss.item():.4f} val {best_v:.4f}", flush=True)
481
- tlog({"wm/final_val": best_v}, step=wm_steps)
482
-
483
- # feature extractor for actor: CLS + text
484
- def actor_feats(screens_np, texts):
485
- scr = torch.from_numpy(np.stack(screens_np)).to(device)
486
- _, cls = enc.patches(scr)
487
- txt = enc.text(texts)
488
- return torch.cat([cls, txt], -1)
489
-
490
- actor = Actor().to(device)
491
- opt = torch.optim.Adam(actor.parameters(), lr=1e-3)
492
-
493
- # ── D: BC on oracle demos ──
494
- for step in range(bc_steps):
495
- b = random.sample(trajs, min(8 if S else 32, len(trajs)))
496
- xs, ys, txts = [], [], []
497
- for tr in b:
498
- T = len(tr["actions"])
499
- s = random.randrange(T)
500
- xs.append(tr["frames"][s]); ys.append(int(tr["actions"][s])); txts.append(tr["instr"])
501
- feats = actor_feats(xs, txts)
502
- logits, v = actor(feats[:, :384], feats[:, 384:])
503
- loss = F.cross_entropy(logits, torch.tensor(ys, device=device))
504
- opt.zero_grad(); loss.backward(); opt.step()
505
- print(f"[bc] CE {loss.item():.4f}", flush=True)
506
- tlog({"bc/ce": loss.item()}, step=bc_steps)
507
-
508
- # real-env eval of BC policy
509
- def real_eval(policy, n_ep, greedy=True):
510
- env = gymnasium.make(f"miniwob/{TASKS[0]}")
511
- succ, rets = 0, []
512
- for ep in range(n_ep):
513
- raw, _ = env.reset(seed=1000 + ep)
514
- ret = 0.0
515
- for t in range(MAX_EP_STEPS):
516
- cell = policy(raw)
517
- raw2, r, term, trunc, _ = env.step({"action_type": 2, "coords": coords_of(cell)})
518
- r = float(r); ret += r; raw = raw2
519
- if term or trunc:
520
- break
521
- rets.append(ret); succ += ret > 0.5
522
- env.close()
523
- return succ / max(1, n_ep), float(np.mean(rets))
524
-
525
- def bc_policy(raw):
526
- f = actor_feats([screen_of(raw)], [instr_text(raw)])
527
- with torch.no_grad():
528
- logits, _ = actor(f[:, :384], f[:, 384:])
529
- return int(logits.argmax(-1))
530
- s0, r0 = real_eval(bc_policy, n_ep=n_eval)
531
- print(f"[bc] real-env success {s0:.2f} avg_return {r0:.3f}", flush=True)
532
- tlog({"eval/bc_success": s0, "eval/bc_return": r0}, step=bc_steps)
533
-
534
- # v5.1: push BC-stage checkpoint immediately so an RL crash loses nothing
535
- os.makedirs("/tmp/ckpt", exist_ok=True)
536
- torch.save({"tok": tok.state_dict(), "wm": wm.state_dict(), "actor": actor.state_dict()},
537
- "/tmp/ckpt/bc_stage.pt")
538
- if os.environ.get("HF_TOKEN"):
539
- try:
540
- from huggingface_hub import HfApi
541
- HfApi().upload_file(path_or_fileobj="/tmp/ckpt/bc_stage.pt",
542
- path_in_repo="v5/bc_stage.pt", repo_id=HF_REPO, repo_type="model")
543
- print("[push] BC checkpoint uploaded", flush=True)
544
- except Exception as e:
545
- print(f"[push] BC upload failed: {e}", flush=True)
546
-
547
- # ── E: RL in imagination (REINFORCE + critic, pessimistic reward) ──
548
- clf = nn.Linear(DM, 1).to(device)
549
- ent_target, ent_coef = 4.0, 0.05
550
- h_real, h_imag, lab = [], [], []
551
- for step in range(rl_steps):
552
- idxs = [(random.randrange(len(trajs)), random.randrange(max(1, len(trajs[0]["actions"]) - L_CTX))) for _ in range(8 if S else 24)]
553
- codes, actx, rews, dones, scr, _ = batch_windows(trajs, idxs, L_CTX, enc, tok, device)
554
- B = codes.shape[0]
555
- codes0, actx0 = codes, actx # v5.1: snapshot of the REAL window for classifier labels
556
- # actor features from the last context frame + instruction
557
- with torch.no_grad():
558
- p, cl = enc.patches(scr[:, L_CTX])
559
- txts = [trajs[ti]["instr"] for ti, _ in idxs]
560
- txt = enc.text(txts)
561
- logits, values = actor(cl, txt)
562
- logp_all = F.log_softmax(logits, -1)
563
- probs = logp_all.exp()
564
- ent = -(probs * logp_all).sum(-1).mean()
565
- actions = torch.multinomial(probs, 1).squeeze(-1)
566
- logp = logp_all.gather(1, actions[:, None]).squeeze(1)
567
- # imagine H=2 steps
568
- cache = wm._cache(codes[:, :L_CTX - 1] if L_CTX > 1 else codes[:, :1],
569
- actx[:, :L_CTX - 1] if L_CTX > 1 else actx[:, :0])
570
- rets = torch.zeros(B, device=device)
571
- for hi in range(2 if not S else 1):
572
- ac = actions if hi == 0 else torch.multinomial(probs, 1).squeeze(-1)
573
- nc, r, c = wm.imagine_step(cache, ac)
574
- rets = rets + r * (0.99 ** hi)
575
- # build next context by rolling: drop oldest frame, append imagined codes
576
- codes = torch.cat([codes[:, 1:], nc.unsqueeze(1)], 1)
577
- actx = torch.cat([actx[:, 1:], ac[:, None]], 1)
578
- # advantage from imagined return vs critic on context features
579
- adv = (rets - values).detach()
580
- adv = (adv - adv.mean()) / (adv.std() + 1e-6)
581
- pg = -(logp * adv).mean() - ent_coef * F.relu(ent_target - ent) ** 2 # v5.1: quadratic floor
582
- vloss = F.mse_loss(values, rets)
583
- (pg + 0.5 * vloss).backward()
584
- torch.nn.utils.clip_grad_norm_(actor.parameters(), 1.0)
585
- opt.step(); opt.zero_grad()
586
- if (step + 1) % 100 == 0:
587
- print(f"[rl] step {step+1} ent {ent.item():.2f} R {rets.mean().item():.3f} v {values.mean().item():.3f}", flush=True)
588
- tlog({"rl/entropy": ent.item(), "rl/imagined_R": rets.mean().item(),
589
- "rl/v_mean": values.mean().item()}, step=step + 1)
590
- if (step + 1) % 500 == 0 or (S and step + 1 == rl_steps):
591
- s, rr = real_eval(bc_policy, n_ep=n_eval)
592
- print(f"[rl-eval] step {step+1} real-env success {s:.2f} avg_return {rr:.3f} entropy {ent.item():.2f}", flush=True)
593
- tlog({"eval/success": s, "eval/return": rr}, step=step + 1)
594
- # v5.1: classifier data — real window (pre-roll) vs post-imagination window
595
- with torch.no_grad():
596
- _, _, _, _, fh_r = wm.forward_full(codes0, actx0)
597
- _, _, _, _, fh_i = wm.forward_full(codes, actx)
598
- h_real.append(fh_r[:, -1].mean(1).cpu()); lab.append(torch.ones(B))
599
- h_imag.append(fh_i[:, -1].mean(1).cpu()); lab.append(torch.zeros(B))
600
-
601
- # ── F: dream-vs-real classifier ──
602
- H = torch.cat(h_real + h_imag).to(device) # v5.1: was left on cpu -> crash on cuda
603
- Y = torch.cat(lab).to(device)
604
- H = (H - H.mean(0)) / (H.std(0) + 1e-6)
605
- n = len(H); ntr = int(0.8 * n)
606
- optc = torch.optim.Adam(clf.parameters(), lr=1e-2)
607
- for step in range(200 if not S else 50):
608
- i = torch.randperm(ntr)[: min(64, ntr)]
609
- loss = F.binary_cross_entropy_with_logits(clf(H[i]).squeeze(-1), Y[i])
610
- opt.zero_grad(); loss.backward(); opt.step()
611
- with torch.no_grad():
612
- sc = clf(H[ntr:]).squeeze(-1).cpu().numpy()
613
- yy = Y[ntr:].cpu().numpy()
614
- auc = auc_score(sc, yy)
615
- print(f"[clf] dream-vs-real AUC {auc:.3f}", flush=True)
616
- tlog({"clf/auc": auc}, step=rl_steps)
617
-
618
- # ── final eval + push ──
619
- s, rr = real_eval(bc_policy, n_ep=max(8, n_eval))
620
- print(f"[final] real-env success {s:.2f} avg_return {rr:.3f}", flush=True)
621
- tlog({"eval/final_success": s, "eval/final_return": rr}, step=rl_steps)
622
- ckpt = {"tok": tok.state_dict(), "wm": wm.state_dict(), "actor": actor.state_dict(), "clf": clf.state_dict()}
623
- os.makedirs("/tmp/ckpt", exist_ok=True)
624
- torch.save(ckpt, "/tmp/ckpt/imagination_model_v5.pt")
625
- if os.environ.get("HF_TOKEN"):
626
- from huggingface_hub import HfApi
627
- HfApi().upload_file(path_or_fileobj="/tmp/ckpt/imagination_model_v5.pt",
628
- path_in_repo="v5/imagination_model_v5.pt", repo_id=HF_REPO, repo_type="model")
629
- print("[push] checkpoint uploaded", flush=True)
630
- if trackio is not None:
631
- try:
632
- trackio.finish()
633
- except Exception:
634
- pass
635
-
636
-
637
- if __name__ == "__main__":
638
- main()
 
1
+ IyEvdXNyL2Jpbi9lbnYgcHl0aG9uMwoiIiIKdHJhaW5faW1hZ2luYXRpb25fdjUucHkg4oCUIGltYWdpbmF0aW9uIG1vZGVsIGZvciBjb21wdXRlciB1c2UsIHY1IChhY3Rvci1yZXByZXNlbnRhdGlvbiBmaXgpLgoKSGlzdG9yeSB0aGlzIGJ1aWxkIGFuc3dlcnM6CiAgdjEvZml4MWI6IFJMLWluLWltYWdpbmF0aW9uIG9ubHkgLT4gYWN0b3IgaGFkIG5vIHBhdGggZnJvbSBpbnN0cnVjdGlvbiB0byBjbGljayBsb2NhdGlvbjsKICAgICAgICAgICAgcmVhbC1lbnYgc3VjY2VzcyAwLjAwLCBlbnRyb3B5IGNvbGxhcHNlLgogIHY0OiAgICAgICBCQyBhZGRlZCAtPiBzdGlsbCAwLjEwIHN1Y2Nlc3M7IGFjdG9yIGlucHV0cyAoMjUgY29hcnNlIEZTUSBjb2RlcyArIGhhc2hlZCB3b3JkCiAgICAgICAgICAgIGJ1Y2tldHMpIGNhbm5vdCByZXByZXNlbnQgIndoZXJlIGlzIHRoZSAnT2snIGJ1dHRvbiIuCiAgdjUgZml4OiAgIGFjdG9yIHNlZXMgZnJvemVuIERJTk92Mi1zbWFsbCBwYXRjaC9DTFMgZmVhdHVyZXMgKHNlbWFudGljLCBwb3NpdGlvbi1wcmVzZXJ2aW5nKQogICAgICAgICAgICArIGZyb3plbiBTbW9sTE0yLTEzNU0gZW1iZWRkaW5nIG9mIHRoZSBSRUFMIGluc3RydWN0aW9uIHRleHQuIFdvcmxkIG1vZGVsIHN0YXlzIGEKICAgICAgICAgICAgc21hbGwgdHJhbnNmb3JtZXIgb3ZlciBGU1EgY29kZXM7IHJld2FyZCBlbnNlbWJsZSBpcyBwZXNzaW1pc3RpYyAobWluIG9mIDIgaGVhZHMpLgogIEFsc28gZml4ZXMgdGhlIFdNIGNvcHkgZmxhdzogQ0Ugbm93IHByZWRpY3RzIGNvZGVzWzosdCsxXSBmcm9tIG91dHB1dCBwb3NpdGlvbiB0CiAgKHYyL3YzIGNvbXBhcmVkIGEgcG9zaXRpb24ncyBvdXRwdXQgdG8gdGhlIGNvZGVzIGZlZCBhcyBpbnB1dCBhdCB0aGF0IHNhbWUgcG9zaXRpb24pLgogIHY1LjE6IGZpeCBjbGFzc2lmaWVyIGRldmljZSBjcmFzaCAoSCBvbiBjcHUgdnMgY2xmIG9uIGN1ZGEpLCBmaXggZHJlYW0tdnMtcmVhbCBsYWJlbCBtaXhpbmcKICAgICAgICAocmVhbCBlbWJlZGRpbmdzIG5vdyBjb21lIGZyb20gdGhlIHByZS1yb2xsIGNvbnRleHQgd2luZG93LCBub3QgcG9zdC1pbWFnaW5hdGlvbiksCiAgICAgICAgcXVhZHJhdGljIGVudHJvcHkgZmxvb3IsIEJDLXN0YWdlIGNoZWNrcG9pbnQgcHVzaGVkIGJlZm9yZSBSTCBzbyBhIGNyYXNoIGxvc2VzIG5vdGhpbmcuCiAgdjUuMjogYWN0aW9uIGdyaWQgNng2PTM2IGFsaWduZWQgMToxIHdpdGggdGhlIEZTUSB0b2tlbiBncmlkIChOVCA9IEdSSUQqR1JJRDsgRElOT3YyIHBhdGNoCiAgICAgICAgZmVhdHVyZXMgYWRhcHRpdmUtcG9vbGVkIHRvIEdSSUQgeCBHUklEKTsgb3JhY2xlIGNsaWNrcyB0aGUgQ0VMTCBDRU5URVIgc28gZXZlcnkKICAgICAgICBzdG9yZWQgbGFiZWwgaXMgZXhhY3RseSB0aGUgZXhlY3V0ZWQgYWN0aW9uOyBhY3RvciByZXBsYWNlZCBieSBhIHBlci1wb3NpdGlvbgogICAgICAgIGJpbGluZWFyIHNwYXRpYWwgaGVhZDogZV9rID0gTUxQKFtwZXItcG9zaXRpb24gRElOT3YyIGZlYXR1cmUsIGxlYXJuZWQgMkQgcG9zaXRpb24KICAgICAgICBlbWJlZGRpbmcsIGluc3RydWN0aW9uIGVtYmVkZGluZ10pICsgYmlsaW5lYXIgeCp5IHRlcm0sIGxvZ2l0X2sgPSAoVyBlX2spLmcvc3FydChkKSwKICAgICAgICBwb2xpY3kgPSBDYXRlZ29yaWNhbCBvdmVyIHRoZSBOVCBncmlkIHBvc2l0aW9ucyAoc2FtZSBpbmRleGluZyBhcyBjb29yZHNfb2YpLCBjcml0aWMKICAgICAgICBvbiB0aGUgcG9zaXRpb24gbWVhbi4gUkwgZml4OiBERUxFVEVEIGFkdmFudGFnZSBub3JtYWxpemF0aW9uIChpdCBhbXBsaWZpZWQKICAgICAgICB6ZXJvLXNpZ25hbCBub2lzZSB+MWU4eCBhbmQgY29sbGFwc2VkIGVudHJvcHkgdG8gMCB3aXRoaW4gMTAwIHN0ZXBzKSAtPgogICAgICAgIGFkdiA9IChSIC0gdmFscy5kZXRhY2goKSkuY2xhbXAoLTEwLDEwKTsgc3RhbmRhcmQgZW50cm9weSBib251cyAtMC4wMiplbnQ7CiAgICAgICAgcG9saWN5IGdyYWQgY2xpcCAxLjAuIC0tbWVhc3VyZS1ncmlkIE4gcnVucyB0aGUgb3JhY2xlLXJldHVybiBtZWFzdXJlbWVudC4KClBpcGVsaW5lOiBBIGNvbGxlY3QgKERPTSBvcmFjbGUgKyByYW5kb20pIC0+IEIgRlNRIHRva2VuaXplciBvbiBmcm96ZW4gRElOT3YyIHBhdGNoIGZlYXR1cmVzCi0+IEMgd29ybGQgbW9kZWwgKHZhbC1zcGxpdCBlYXJseSBzdG9wKSAtPiBEIEJDIG9uIG9yYWNsZSBkZW1vcyAtPiBFIFJMIGluIGltYWdpbmF0aW9uCihwZXNzaW1pc3RpYyBtaW4tb2YtMi1oZWFkcyByZXdhcmQsIFJFSU5GT1JDRSArIGNyaXRpY2FsLCBlbnRyb3B5IHRhcmdldCkgLT4gcmVhbC1lbnYgZXZhbHMgLTI+CkYgZHJlYW0tdnMtcmVhbCBjbGFzc2lmaWVyIChBVUMpLgoKRW52IGZhY3RzIChzYW5kYm94LXZlcmlmaWVkIDIwMjYtMDktMTksIG1pbml3b2IgMS4xLjApOgogIG9icyBrZXlzOiBkb21fZWxlbWVudHMsIGZpZWxkcywgc2NyZWVuc2hvdCwgdXR0ZXJhbmNlOyBzY3JlZW5zaG90ICgyMTAsMTYwLDMpIHVpbnQ4IChIMjEwLFcxNjApCiAgYWN0aW9uIHsiYWN0aW9uX3R5cGUiOiAyLCAiY29vcmRzIjogZmxvYXQzMlt4LHldfSA9IENMSUNLX0NPT1JEUyBvdmVyIEJveCgwLFsxNjAsMjEwXSkKICBkb20gZWxlbWVudHM6IGZsYXQgbGVmdC90b3Avd2lkdGgvaGVpZ2h0ICgxLWVsZW0gZmxvYXQzMiBhcnJheXMpLCB0ZXh0IG1hdGNoIGlzIGNhc2Utc2Vuc2l0aXZlClVzYWdlOgogIHB5dGhvbiB0cmFpbl9pbWFnaW5hdGlvbl92NS5weSAtLXNtb2tlICAgICAgICAgICAgIyB0aW55IENQVSBydW4sIGFsbCBwaGFzZXMKICBweXRob24gdHJhaW5faW1hZ2luYXRpb25fdjUucHkgLS1kZXZpY2UgY3VkYSAgICAgICMgcmVhbCBydW4gKHQ0LXNtYWxsLCB+MS41aCkKIiIiCmltcG9ydCBhcmdwYXJzZSwgbWF0aCwgb3MsIHJhbmRvbSwgdGltZQppbXBvcnQgbnVtcHkgYXMgbnAKaW1wb3J0IHRvcmNoLCB0b3JjaC5ubiBhcyBubiwgdG9yY2gubm4uZnVuY3Rpb25hbCBhcyBGCgpvcy5lbnZpcm9uLnNldGRlZmF1bHQoIk1JTklXT0JfQ0hST01FX0JJTkFSWSIsICIvdXNyL2Jpbi9jaHJvbWl1bSIpCm9zLmVudmlyb24uc2V0ZGVmYXVsdCgiTUlOSVdPQl9DSFJPTUVEUklWRVIiLCAiL3Vzci9iaW4vY2hyb21lZHJpdmVyIikKCmltcG9ydCBneW1uYXNpdX1LCBtaW5pd29iCmZyb20gUElMIGltcG9ydCBJbWFnZQpmcm9tIHRyYW5zZm9ybWVycyBpbXBvcnQgQXV0b01vZGVsLCBBdXRvVG9rZW5pemVyCgpneW1uYXNpdW0ucmVnaXN0ZXJfZW52cyhtaW5pd29iKQo=