File size: 5,637 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
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
"""Validate the latent dataset against the raw HDF5 and the world model.



The indexing here is easy to get subtly wrong: pixels are strided by frameskip

while actions are not, and the block leaving frame k must be the actions that

actually produced frame k+1. This checks that alignment against ground truth,

then confirms the frozen predictor can reach the sampled goals — if teacher

forcing on real actions doesn't reduce the goal distance, the goals aren't

reachable and the controller has no learnable signal.

"""

import sys
from pathlib import Path

import h5py
import hdf5plugin  # noqa: F401  -- blosc filter
import numpy as np
import torch

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

from lejepa_control.data import (  # noqa: E402
    FRAMESKIP,
    NUM_CONTEXT,
    LatentGoalDataset,
    split_episodes,
)
from lejepa_control.rollout import goal_distance, rollout_plan  # noqa: E402
from lejepa_control.world_model import load_lewm  # noqa: E402

H5 = '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 main():
    device = 'cuda' if torch.cuda.is_available() else 'cpu'
    torch.manual_seed(0)

    ds = LatentGoalDataset('data/latents', max_offset=5)
    print(f'{len(ds)} clips over {len(ds.episodes)} episodes')

    f = h5py.File(H5, 'r')
    model = load_lewm(device=device)

    # --- cached latents match a fresh encode of the same raw frames -------
    sample_idx = [0, 12345, 999999]
    for idx in sample_idx:
        ep = int(ds.clip_episode[idx])
        start = int(ds.clip_start[idx])
        base = int(ds.offsets[ep])
        rows = [base + start + i * FRAMESKIP for i in range(NUM_CONTEXT)]

        px = torch.from_numpy(f['pixels'][rows[0] : rows[-1] + 1 : FRAMESKIP])
        px = px.permute(0, 3, 1, 2).float() / 255.0
        px = ((px - IMAGENET_MEAN) / IMAGENET_STD).unsqueeze(0).to(device)
        with torch.no_grad():
            fresh = model.encode({'pixels': px})['emb'][0].cpu()

        cached = ds[idx]['context']
        rel = (cached - fresh).norm() / fresh.norm()
        assert rel < 0.02, f'latent cache mismatch at {idx}: {rel}'
    print(f'latent cache matches fresh encode (rel err < 2%) at {sample_idx}')

    # --- action blocks are the raw actions between the context frames -----
    idx = 12345
    ep = int(ds.clip_episode[idx])
    start = int(ds.clip_start[idx])
    base = int(ds.offsets[ep])
    item = ds[idx]
    for i in range(NUM_CONTEXT - 1):
        lo = base + start + i * FRAMESKIP
        raw = f['action'][lo : lo + FRAMESKIP]  # (5, 2)
        want = ((raw - ds.action_mean) / ds.action_std).reshape(-1)
        got = item['past_actions'][i].numpy()
        assert np.allclose(got, want, atol=1e-5), f'block {i} misaligned'
    print('past action blocks align with raw actions between context frames')

    # --- goal offsets respect the curriculum ------------------------------
    for max_offset in (2, 3, 5):
        ds.set_max_offset(max_offset)
        offsets = [int(ds[i]['goal_offset']) for i in range(200)]
        assert min(offsets) >= 1 and max(offsets) <= max_offset, offsets
        print(f'  max_offset={max_offset}: sampled {sorted(set(offsets))}')

    # --- reachability: real actions must beat a do-nothing plan -----------
    ds.set_max_offset(5)
    loader = torch.utils.data.DataLoader(ds, batch_size=64, shuffle=True)
    batch = next(iter(loader))

    ctx = batch['context'].to(device)
    past = batch['past_actions'].to(device)
    goal = batch['goal'].to(device)
    real = batch['real_action'].to(device)

    # replay the real first block, then hold still
    H = 5
    plan = torch.zeros(ctx.size(0), H, 10, device=device)
    plan[:, 0] = real
    with torch.no_grad():
        pred = rollout_plan(model, ctx, past, plan)
        d_real = goal_distance(pred, goal)

        zero = torch.zeros_like(plan)
        d_zero = goal_distance(rollout_plan(model, ctx, past, zero), goal)

    # distance from the current frame to the goal, before acting at all
    d_start = (ctx[:, -1] - goal).pow(2).mean(-1)

    print(f'start distance to goal      {d_start.mean():.4f}')
    print(f'after real 1st block (d_1)  {d_real[:, 0].mean():.4f}')
    print(f'after zero 1st block (d_1)  {d_zero[:, 0].mean():.4f}')
    assert d_real[:, 0].mean() < d_start.mean(), (
        'real actions do not move toward the goal — indexing is likely wrong'
    )
    assert d_real[:, 0].mean() < d_zero[:, 0].mean(), (
        'real actions are no better than doing nothing'
    )
    print('reachability ok: real actions reduce latent goal distance')

    # --- one-step goals: teacher forcing should nearly reach --------------
    ds.set_max_offset(1)
    batch = next(iter(torch.utils.data.DataLoader(ds, batch_size=64, shuffle=True)))
    ctx = batch['context'].to(device)
    past = batch['past_actions'].to(device)
    goal = batch['goal'].to(device)
    plan = batch['real_action'].to(device).unsqueeze(1)
    with torch.no_grad():
        d1 = goal_distance(rollout_plan(model, ctx, past, plan), goal)[:, 0]
    d0 = (ctx[:, -1] - goal).pow(2).mean(-1)
    print(f'1-step goal: before {d0.mean():.4f} -> after {d1.mean():.4f} '
          f'({100 * (1 - d1.mean() / d0.mean()):.0f}% closed)')
    assert d1.mean() < 0.5 * d0.mean(), 'one-step goals not being reached'

    print('ALL DATA TESTS PASSED')


if __name__ == '__main__':
    main()