"""Pre-encode per-finger tactile camera images through T3 (t3_medium) and write features back into each episode HDF5 as ``/t3/`` (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_`` (T, 240, 320) uint8 grayscale -> 3ch -> 224. Output: ``/t3/`` (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}") @torch.no_grad() 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 @torch.no_grad() 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/ datasets total.") if __name__ == "__main__": main()