Download tools/encode_t3.py from AllenBi21/tactile-act-screw-baselines: direct link, hf CLI and curl.
- Browser
- Download file 7.12 kB
-
https://huggingface.co/AllenBi21/tactile-act-screw-baselines/resolve/main/tools/encode_t3.py
- Command line
-
hf download hf://AllenBi21/tactile-act-screw-baselines/tools/encode_t3.py
-
curl -L -o encode_t3.py https://huggingface.co/AllenBi21/tactile-act-screw-baselines/resolve/main/tools/encode_t3.py
7.12 kB
| """Pre-encode per-finger tactile camera images through T3 (t3_medium) and write | |
| features back into each episode HDF5 as ``/t3/<f>`` (T, 768) float32. | |
| Mirrors ``encode_sitr.py`` (same input, same batching) but reuses the T3 loader | |
| from tactile_fusion's ``load_baselines.py`` so the feature convention is | |
| IDENTICAL to the T3 probe rows in the paper: encoder tower -> 9-block trunk -> | |
| non-affine LayerNorm -> cls token. (Skipping the trunk was the historical bug | |
| that made features near-random; the LN re-scales the trunk's tiny output std.) | |
| Normalization follows T3's own convention — per-dataset channel stats on raw | |
| [0,1] RGB — computed here over the episodes being encoded (SharpaWave frames), | |
| NOT borrowed from gsmini/9dtact (borrowed stats put inputs several sigma | |
| off-distribution; see load_baselines.py notes). | |
| Source images: ``/images/raw_<f>`` (T, 240, 320) uint8 grayscale -> 3ch -> 224. | |
| Output: ``/t3/<f>`` (T, 768) float32, gzip. | |
| Usage: | |
| python -m tools.encode_t3 --episodes-dir episodes/screwing | |
| python -m tools.encode_t3 --episodes-dir episodes/tofu | |
| python -m tools.encode_t3 --smoke # CPU sanity, no writes | |
| """ | |
| 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 as nn | |
| import torch.nn.functional as F | |
| TACTILE_FUSION = "/home/allenbi/projects_25/tactile_fusion" | |
| N_FINGERS = 5 | |
| T3_FEAT_DIM = 768 | |
| STATS_FRAMES_PER_EP = 40 # subsampled frames per episode for the stats pre-pass | |
| def load_t3(device: torch.device): | |
| # tactile_act's own `data/` package (regular package with __init__) shadows | |
| # tactile_fusion's namespace `data/` package no matter the path order — | |
| # drop tactile_act from sys.path and purge stale modules before importing. | |
| sys.path[:] = [p for p in sys.path | |
| if "tactile_act" not in Path(p or ".").resolve().as_posix()] | |
| for name in list(sys.modules): | |
| if name == "data" or name.startswith("data."): | |
| del sys.modules[name] | |
| if TACTILE_FUSION not in sys.path: | |
| sys.path.insert(0, TACTILE_FUSION) | |
| import os | |
| cwd = os.getcwd() | |
| os.chdir(TACTILE_FUSION) # load_baselines resolves relative third_party paths | |
| try: | |
| from load_baselines import _load_t3_raw # type: ignore | |
| model, domain = _load_t3_raw(modality="9dtact") # SharpaWave treated as | |
| # 9dtact-like, same choice as the HTT (tf9dpe) encoding. | |
| finally: | |
| os.chdir(cwd) | |
| model = model.to(device).eval() | |
| encoder = model.encoders[domain] | |
| trunk = model.trunk | |
| print(f"[t3] encoder domain: {domain}") | |
| def forward(x: torch.Tensor) -> torch.Tensor: | |
| tokens = encoder(x) | |
| tokens = trunk(tokens) | |
| tokens = F.layer_norm(tokens, (tokens.shape[-1],)) | |
| return tokens[:, 0, :] # cls | |
| return forward | |
| def compute_dataset_stats(files, device) -> tuple[torch.Tensor, torch.Tensor]: | |
| """Per-channel mean/std over subsampled raw frames of all fingers/episodes, | |
| on [0,1] grayscale replicated to 3ch (channels are identical -> stats too, | |
| but keep the 3-vector form for parity with the T3 convention).""" | |
| acc_sum, acc_sq, n_px = 0.0, 0.0, 0 | |
| for path in files: | |
| with h5py.File(path, "r") as hf: | |
| T = hf["/qpos"].shape[0] | |
| idx = np.linspace(0, T - 1, min(STATS_FRAMES_PER_EP, T)).astype(int) | |
| for f in range(N_FINGERS): | |
| frames = hf[f"/images/raw_{f}"][idx].astype(np.float64) / 255.0 | |
| acc_sum += frames.sum() | |
| acc_sq += (frames ** 2).sum() | |
| n_px += frames.size | |
| mean = acc_sum / n_px | |
| std = float(np.sqrt(acc_sq / n_px - mean ** 2)) | |
| print(f"[t3] dataset stats over {n_px/1e6:.1f}M px: mean={mean:.5f} std={std:.5f}") | |
| m = torch.full((1, 3, 1, 1), float(mean), device=device) | |
| s = torch.full((1, 3, 1, 1), std, device=device) | |
| return m, s | |
| def preprocess_batch(images_uint8: np.ndarray, mean, std, device) -> torch.Tensor: | |
| 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) | |
| return (x - mean) / std | |
| def encode_episode(hf, forward, mean, std, device, batch_size, overwrite) -> int: | |
| T = hf["/qpos"].shape[0] | |
| n_done = 0 | |
| grp = hf.require_group("t3") | |
| for f in range(N_FINGERS): | |
| key = str(f) | |
| if key in grp and not overwrite: | |
| print(f" finger {f}: /t3/{f} exists, skipping") | |
| continue | |
| if key in grp: | |
| del grp[key] | |
| raw_ds = hf[f"/images/raw_{f}"] | |
| out = np.zeros((T, T3_FEAT_DIM), dtype=np.float32) | |
| for start in range(0, T, batch_size): | |
| end = min(start + batch_size, T) | |
| batch = preprocess_batch(raw_ds[start:end], mean, std, device) | |
| out[start:end] = forward(batch).cpu().numpy() | |
| grp.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=64) | |
| ap.add_argument("--device", type=str, | |
| default="cuda" if torch.cuda.is_available() else "cpu") | |
| ap.add_argument("--overwrite", action="store_true") | |
| ap.add_argument("--smoke", action="store_true", | |
| help="CPU sanity: encode 8 frames of finger 0 of the first " | |
| "episode, print the feature shape/std, write nothing.") | |
| 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("cpu" if args.smoke else args.device) | |
| forward = load_t3(device) | |
| mean, std = compute_dataset_stats(files[: 3 if args.smoke else len(files)], device) | |
| if args.smoke: | |
| with h5py.File(files[0], "r") as hf: | |
| batch = preprocess_batch(hf["/images/raw_0"][:8], mean, std, device) | |
| feats = forward(batch) | |
| print(f"[smoke] feats {tuple(feats.shape)} std={feats.std():.4f} " | |
| f"finite={bool(torch.isfinite(feats).all())}") | |
| return | |
| t0 = time.time() | |
| total = 0 | |
| for i, path in enumerate(files): | |
| with h5py.File(path, "r+") as hf: | |
| t_ep = time.time() | |
| n = encode_episode(hf, forward, mean, std, device, | |
| args.batch_size, args.overwrite) | |
| print(f"[{i+1}/{len(files)}] {path.name} wrote {n} streams " | |
| f"({time.time()-t_ep:.1f}s)", flush=True) | |
| total += n | |
| print(f"\ndone in {time.time()-t0:.1f}s; {total} /t3/<f> datasets total.") | |
| if __name__ == "__main__": | |
| main() | |