"""Pre-encode per-finger tactile camera images through SITR and write the features back into each episode HDF5 as ``/sitr/`` (T, 768) float32. The reasoning: SITR is a ~86M-param ViT-base. Running it once per frame is cheap (~5–10 ms) but re-running it inside every ACT training step is wasteful. This script does one forward pass per (episode × finger × frame), stores the cls-token features in the HDF5, and lets ACT consume them via the existing ``--tactile-keys sitr_0 … sitr_4 --tactile-dim 3840`` machinery. Source images: ``/images/raw_`` (T, 240, 320) uint8 grayscale. Replicated to 3 channels, resized to 224×224, normalized with the stats SITR was trained under (``gsrl/dataloaders.py``: sample_mu / sample_std). Output: ``/sitr/`` (T, 768) float32, gzip-compressed. """ from __future__ import annotations import argparse import sys import time from pathlib import Path import h5py import numpy as np import torch import torch.nn.functional as F GSRL_ROOT = "/home/allenbi/projects_25/gsrl" SITR_CKPT = f"{GSRL_ROOT}/checkpoints/SITR_B18.pth" # SITR training normalization (gsrl/dataloaders.py:14-15). These are applied # *after* ToTensor() (i.e. on inputs already scaled to [0, 1]). SITR_MEAN = (-1.2223, -1.8114, -1.7090) SITR_STD = (11.7932, 12.7956, 13.6452) N_FINGERS = 5 SITR_FEAT_DIM = 768 def load_sitr(device: torch.device) -> torch.nn.Module: sys.path.insert(0, GSRL_ROOT) from models.networks import SITR_base # type: ignore model = SITR_base(num_calibration=0) # inference without calibration frames ckpt = torch.load(SITR_CKPT, map_location="cpu") state_dict = ckpt["state_dict"] if isinstance(ckpt, dict) and "state_dict" in ckpt else ckpt missing, unexpected = model.load_state_dict(state_dict, strict=False) print(f"[sitr] loaded {SITR_CKPT} missing={len(missing)} unexpected={len(unexpected)}") return model.to(device).eval() def preprocess_batch(images_uint8: np.ndarray, device: torch.device) -> torch.Tensor: """images_uint8: (B, H, W) uint8 grayscale → (B, 3, 224, 224) float32 normalized.""" x = torch.from_numpy(images_uint8).to(device).float() / 255.0 # (B, H, W) x = x.unsqueeze(1).expand(-1, 3, -1, -1) # (B, 3, H, W) x = F.interpolate(x, size=(224, 224), mode="bilinear", align_corners=False) mean = torch.tensor(SITR_MEAN, device=device).view(1, 3, 1, 1) std = torch.tensor(SITR_STD, device=device).view(1, 3, 1, 1) return (x - mean) / std @torch.no_grad() def encode_episode(hf: h5py.File, model: torch.nn.Module, device: torch.device, batch_size: int, overwrite: bool) -> int: """Encode all 5 finger streams and write /sitr/. Returns number of features written.""" T = hf["/qpos"].shape[0] n_done = 0 sitr_group = hf.require_group("sitr") for f in range(N_FINGERS): key = str(f) if key in sitr_group and not overwrite: print(f" finger {f}: /sitr/{f} exists, skipping") continue if key in sitr_group: del sitr_group[key] raw_ds = hf[f"/images/raw_{f}"] out = np.zeros((T, SITR_FEAT_DIM), dtype=np.float32) for start in range(0, T, batch_size): end = min(start + batch_size, T) batch_np = raw_ds[start:end] # (b, 240, 320) uint8 batch = preprocess_batch(batch_np, device) # (b, 3, 224, 224) feats = model.forward_encoder(batch, c=None) # (b, N+1, 768) out[start:end] = feats[:, 0, :].cpu().numpy() # cls token sitr_group.create_dataset(key, data=out, compression="gzip", compression_opts=4) n_done += 1 return n_done def main() -> None: ap = argparse.ArgumentParser() ap.add_argument("--episodes-dir", type=Path, default=Path("episodes/screwing")) ap.add_argument("--episode-glob", type=str, default="episode_*.hdf5") ap.add_argument("--batch-size", type=int, default=32) ap.add_argument("--device", type=str, default="cuda" if torch.cuda.is_available() else "cpu") ap.add_argument("--overwrite", action="store_true", help="Re-encode and overwrite existing /sitr/ datasets.") args = ap.parse_args() files = sorted(args.episodes_dir.glob(args.episode_glob)) if not files: print(f"No HDF5 under {args.episodes_dir} matching {args.episode_glob}", file=sys.stderr) sys.exit(1) device = torch.device(args.device) model = load_sitr(device) t0 = time.time() total_written = 0 for i, path in enumerate(files): with h5py.File(path, "r+") as hf: T = hf["/qpos"].shape[0] t_ep = time.time() n = encode_episode(hf, model, device, args.batch_size, args.overwrite) dt = time.time() - t_ep print(f"[{i+1}/{len(files)}] {path.name} T={T} wrote {n} streams ({dt:.1f}s)", flush=True) total_written += n print(f"\ndone in {time.time() - t0:.1f}s; {total_written} /sitr/ datasets total.") if __name__ == "__main__": main()