File size: 5,866 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
140
141
142
143
144
145
146
147
148
149
150
151
"""Encode every DMC Reacher frame with the frozen LeWM encoder into a latent cache.



Reacher port of ``scripts/encode_latents.py``. Controller training never needs

pixels: the encoder is frozen and no image augmentation is used, so latents can

be computed once. Also stores the action z-score statistics — LeWM was trained

on z-scored actions (``column_normalizer`` defaults to ``method='zscore'``),

so the controller must emit actions in that same normalized space.



The only reacher-specific differences from the PushT script are the default

dataset path and the episode-index column names: the DMC collect pipeline

writes ``ep_idx`` where PushT wrote ``episode_idx``. Both spellings are

accepted, and ``ep_len``/``ep_offset`` are derived when absent.

"""

import argparse
import json
import time
from pathlib import Path

import h5py
import hdf5plugin  # noqa: F401  -- registers the blosc filter used by the h5
import numpy as np
import torch

from lejepa_control.world_model import load_lewm

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 episode_layout(f):
    """``(lengths, offsets)`` and the dataset row count, per spelling."""
    keys = set(f.keys())
    if 'ep_len' in keys and 'ep_offset' in keys:
        lengths = f['ep_len'][:].astype(np.int64)
        offsets = f['ep_offset'][:].astype(np.int64)
        return lengths, offsets, int(offsets[-1] + lengths[-1])
    ep_col = 'episode_idx' if 'episode_idx' in keys else 'ep_idx'
    ep_idx = f[ep_col][:].astype(np.int64)
    step_idx = f['step_idx'][:].astype(np.int64)
    n_eps = int(ep_idx.max()) + 1
    lengths = np.zeros(n_eps, dtype=np.int64)
    np.maximum.at(lengths, ep_idx, step_idx + 1)
    keep = lengths > 0
    lengths = lengths[keep]
    offsets = np.concatenate([[0], np.cumsum(lengths)[:-1]])
    return lengths, offsets, len(ep_idx)


def main():
    parser = argparse.ArgumentParser()
    parser.add_argument(
        '--h5', default='data/swm_home/datasets/dmc/reacher.h5'
    )
    parser.add_argument('--out', default='data/latents_reacher')
    parser.add_argument('--batch-size', type=int, default=512)
    parser.add_argument('--limit-episodes', type=int, default=None)
    parser.add_argument('--wm-name', default='quentinll/lewm-reacher')
    args = parser.parse_args()

    device = 'cuda' if torch.cuda.is_available() else 'cpu'
    model = load_lewm(name=args.wm_name, device=device)
    encoder, projector = model.encoder, model.projector
    D = model.predictor.input_dim

    out_dir = Path(args.out)
    out_dir.mkdir(parents=True, exist_ok=True)

    mean = IMAGENET_MEAN.to(device)
    std = IMAGENET_STD.to(device)

    with h5py.File(args.h5, 'r') as f:
        lengths, offsets, end = episode_layout(f)

        if args.limit_episodes is not None:
            lengths = lengths[: args.limit_episodes]
            offsets = offsets[: args.limit_episodes]
            end = int(offsets[-1] + lengths[-1])

        n_frames = int(lengths.sum())
        print(f'{len(lengths)} episodes, {n_frames} frames -> {D}-dim latents')

        actions = f['action'][:end].astype(np.float32)
        # the frame that ends an episode has no outgoing action (there is no
        # next state to leave it for) and is NaN-padded; ignore those rows for
        # the stats, then fill them with the mean so a block that straddles an
        # episode boundary never hands training a NaN.
        nan_mask = np.isnan(actions).any(axis=1)
        a_mean = np.nanmean(actions, axis=0)
        a_std = np.nanstd(actions, axis=0)
        if nan_mask.any():
            print(
                f'{int(nan_mask.sum())} NaN (terminal-step) action rows'
                ' -> filled with mean'
            )
            actions[nan_mask] = a_mean
        print(f'action mean={a_mean} std={a_std}')

        latents = np.lib.format.open_memmap(
            out_dir / 'latents.npy',
            mode='w+',
            dtype=np.float16,
            shape=(n_frames, D),
        )

        t0 = time.perf_counter()
        for start in range(0, end, args.batch_size):
            stop = min(start + args.batch_size, end)
            frames = f['pixels'][start:stop]  # (B, 224, 224, 3) uint8

            x = torch.from_numpy(frames).to(device, non_blocking=True)
            x = x.permute(0, 3, 1, 2).float().div_(255.0).sub_(mean).div_(std)

            with torch.no_grad(), torch.autocast('cuda', dtype=torch.bfloat16):
                out = encoder(x, interpolate_pos_encoding=True)
                emb = projector(out.last_hidden_state[:, 0].float())

            latents[start:stop] = emb.float().cpu().numpy().astype(np.float16)

            if start % (args.batch_size * 200) == 0:
                done = stop / end
                rate = stop / (time.perf_counter() - t0)
                eta = (end - stop) / rate / 60
                print(
                    f'  {done:6.1%}  {rate:6.0f} img/s  eta {eta:5.1f} min',
                    flush=True,
                )

        latents.flush()

    np.save(out_dir / 'actions.npy', actions)
    np.save(out_dir / 'ep_len.npy', lengths)
    np.save(out_dir / 'ep_offset.npy', offsets)

    stats = {
        'action_mean': a_mean.tolist(),
        'action_std': a_std.tolist(),
        'latent_dim': int(D),
        'n_frames': n_frames,
        'n_episodes': len(lengths),
    }
    (out_dir / 'stats.json').write_text(json.dumps(stats, indent=2))

    z = np.asarray(latents[:10000], dtype=np.float32)
    print(f'latent per-coord std (mean) = {z.std(0).mean():.4f}')
    print(f'done in {(time.perf_counter() - t0) / 60:.1f} min -> {out_dir}')


if __name__ == '__main__':
    main()