AllenBi21's picture
T3 + SITR ACT screw baselines (act_best, args, encode tools, deploy doc)
07c6d07 verified
Raw History Blame Contribute Delete
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}")
@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/<f> datasets total.")
if __name__ == "__main__":
main()