File size: 4,855 Bytes
dc9f917
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
"""Verify the frozen-LeWM contract the controller depends on.



Checks the action-block layout, the normalization the checkpoint was trained

with, and that gradients reach candidate actions through a multi-step rollout.

"""

import sys
import time
from pathlib import Path

import h5py
import numpy as np
import torch

sys.path.insert(0, str(Path(__file__).resolve().parents[1]))

from lejepa_control.world_model import load_lewm  # noqa: E402

H5 = Path('data/swm_home/datasets/pusht_expert_train.h5')
IMAGENET_MEAN = torch.tensor([0.485, 0.456, 0.406]).view(1, 3, 1, 1)
IMAGENET_STD = torch.tensor([0.229, 0.224, 0.225]).view(1, 3, 1, 1)


def preprocess(frames_hwc: np.ndarray) -> torch.Tensor:
    x = torch.from_numpy(frames_hwc).permute(0, 3, 1, 2).float() / 255.0
    return (x - IMAGENET_MEAN) / IMAGENET_STD


def main():
    device = 'cuda' if torch.cuda.is_available() else 'cpu'
    model = load_lewm(device=device)
    D = model.predictor.input_dim
    print(f'num_frames={model.predictor.num_frames} '
          f'action_input_dim={model.action_encoder.input_dim} D={D}')

    f = h5py.File(H5, 'r')
    action = f['action']
    offs, lens = f['ep_offset'][:], f['ep_len'][:]

    # --- action normalization the checkpoint was trained with -------------
    stats = action[:]
    mean, std = stats.mean(0), stats.std(0)
    print(f'action mean={mean} std={std}')
    print(f'raw [-1,1] -> normalized +-{(1.0 / std)}')

    # --- frameskip / block layout ----------------------------------------
    # A clip loads pixels strided by 5 and actions unstrided, reshaped (T, 10).
    # So block k must be the 5 raw actions between frame k and frame k+1.
    ep, start, fs, T = 0, 0, 5, 4
    o = offs[ep]
    raw = action[o + start : o + start + T * fs]  # (20, 2)
    block = raw.reshape(T, fs * 2)  # (4, 10)
    assert np.array_equal(block[1], raw[5:10].reshape(-1)), 'block layout'
    print('block layout ok: block[k] = raw actions [5k, 5k+5)')

    # --- encode a real 3-frame history + goal -----------------------------
    N = model.predictor.num_frames
    idx = [o + i * fs for i in range(N + 1)]  # 3 history + 1 target
    pixels = preprocess(f['pixels'][idx[0] : idx[-1] + 1 : fs])
    pixels = pixels.unsqueeze(0).to(device)  # (1, 4, 3, 224, 224)
    with torch.no_grad():
        emb = model.encode({'pixels': pixels})['emb']
    print(f'emb {tuple(emb.shape)} mean={emb.mean():.4f} std={emb.std():.4f} '
          f'per-coord std={emb.std(dim=(0, 1)).mean():.4f}')

    # --- teacher forcing: does the frozen predictor actually predict? -----
    ctx = emb[:, :N]
    a_norm = (raw.reshape(T, fs, 2) - mean) / std
    act = torch.from_numpy(a_norm.reshape(1, T, fs * 2)).float().to(device)
    with torch.no_grad():
        pred = model.predict(ctx, model.action_encoder(act[:, :N]))
    tgt = emb[:, 1 : N + 1]
    err = (pred - tgt).pow(2).mean().item()
    # baseline: predicting "no change" (copy the last context frame)
    copy_err = (ctx - tgt).pow(2).mean().item()
    print(f'1-step pred MSE={err:.5f}  copy-last-frame MSE={copy_err:.5f}')

    # --- multi-step differentiable rollout, gradient to candidate actions -
    Hz = 5
    cand = torch.zeros(1, Hz, fs * 2, device=device, requires_grad=True)
    hist = list(emb[:, :N].detach().unbind(dim=1))
    past = act[:, : N - 1]  # blocks between the context frames
    all_act = torch.cat([past, cand], dim=1)  # (1, N-1+Hz, 10)
    all_emb = model.action_encoder(all_act)
    for t in range(Hz):
        lo = max(0, N + t - N)
        e = torch.stack(hist[lo:], dim=1)
        a = all_emb[:, lo : N + t]
        hist.append(model.predict(e, a)[:, -1])
    x_H = hist[-1]
    goal = emb[:, -1].detach()
    loss = (x_H - goal).pow(2).sum() / D
    loss.backward()
    g = cand.grad
    print(f'rollout ok: terminal d={loss.item():.5f}')
    print(f'grad norm={g.norm():.5f} per-step={g.norm(dim=-1).squeeze().tolist()}')
    assert g.norm() > 0, 'no gradient reached candidate actions'
    assert g[0, -1].norm() > 0, 'last block got no gradient'

    # --- encoder throughput (decides whether caching latents is viable) ---
    batch = preprocess(f['pixels'][:256]).to(device)
    with torch.no_grad():
        for _ in range(2):
            model.encode({'pixels': batch.unsqueeze(0)})
        torch.cuda.synchronize()
        t0 = time.perf_counter()
        model.encode({'pixels': batch.unsqueeze(0)})
        torch.cuda.synchronize()
    dt = time.perf_counter() - t0
    ips = 256 / dt
    total = int(lens.sum())
    print(f'encoder {ips:.0f} img/s -> {total} frames in {total / ips / 60:.1f} min')
    print(f'latent cache fp16 = {total * D * 2 / 1e9:.2f} GB')
    print('ALL CONTRACT CHECKS PASSED')


if __name__ == '__main__':
    main()