Download tools/encode_sitr.py from AllenBi21/tactile-act-screw-baselines: direct link, hf CLI and curl.
- Browser
- Download file 5.17 kB
-
https://huggingface.co/AllenBi21/tactile-act-screw-baselines/resolve/main/tools/encode_sitr.py
- Command line
-
hf download hf://AllenBi21/tactile-act-screw-baselines/tools/encode_sitr.py
-
curl -L -o encode_sitr.py https://huggingface.co/AllenBi21/tactile-act-screw-baselines/resolve/main/tools/encode_sitr.py
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 | |
| 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() | |