AllenBi21's picture
T3 + SITR ACT screw baselines (act_best, args, encode tools, deploy doc)
07c6d07 verified
Raw History Blame Contribute Delete
5.17 kB
"""Pre-encode per-finger tactile camera images through SITR and write the
features back into each episode HDF5 as ``/sitr/<f>`` (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_<f>`` (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/<f>`` (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/<f>. 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/<f> 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/<f> datasets total.")
if __name__ == "__main__":
main()