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()
|