tactile-vae / test /visualize_color_jitter.py
WitneyWW's picture
Initial upload of tactile_vae (code, model, config, inference)
3770c94 verified
Raw
History Blame Contribute Delete
2.98 kB
"""Visual sanity check: 10 samples with vs. without color jitter.
Writes `color_jitter_compare.png` next to this script. Layout:
row 0 — base image (no augmentation)
row 1 — color-jittered image (training-time augmentation)
row 2 — |jittered − base| amplified for visibility
"""
from __future__ import annotations
import sys
from pathlib import Path
import matplotlib.pyplot as plt
import numpy as np
import torch
_REPO_ROOT = Path(__file__).resolve().parents[2]
if str(_REPO_ROOT) not in sys.path:
sys.path.insert(0, str(_REPO_ROOT))
from tactile_vae.dataset import ColorJitterConfig, TactileParquetDataset
N_SAMPLES = 10
SEED = 0
JITTER = ColorJitterConfig(brightness=0.3, contrast=0.3, saturation=0.3, hue=0.05)
OUT_PNG = Path(__file__).with_name("color_jitter_compare.png")
def _to_hwc(t: torch.Tensor) -> np.ndarray:
return t.detach().cpu().clamp(0, 1).permute(1, 2, 0).numpy()
def main() -> None:
ds_plain = TactileParquetDataset(image_size=128, color_jitter=None)
ds_jit = TactileParquetDataset(image_size=128, color_jitter=JITTER)
# Sample 10 distinct indices across the dataset for variety.
n = len(ds_plain)
rng = np.random.default_rng(SEED)
indices = sorted(rng.choice(n, size=N_SAMPLES, replace=False).tolist())
print(f"Comparing {N_SAMPLES} samples (indices: {indices[:3]}... of {n:,})")
print(f"Jitter config: {JITTER}")
torch.manual_seed(SEED) # fix jitter RNG so the PNG is reproducible
fig, axes = plt.subplots(3, N_SAMPLES, figsize=(2 * N_SAMPLES, 6.6))
for col, idx in enumerate(indices):
base = ds_plain[idx] # (3,128,128) float in [0,1]
jittered = ds_jit[idx]
base_hwc = _to_hwc(base)
jit_hwc = _to_hwc(jittered)
# Amplify residual so visible per-pixel changes are visible at a glance.
diff = np.abs(jit_hwc - base_hwc)
scale = max(diff.max(), 1e-6)
diff_vis = np.clip(diff / scale, 0, 1)
axes[0, col].imshow(base_hwc)
axes[0, col].set_title(f"idx={idx}", fontsize=8)
axes[1, col].imshow(jit_hwc)
axes[1, col].set_title(f"Δmean={float(jit_hwc.mean()-base_hwc.mean()):+.3f}",
fontsize=8)
axes[2, col].imshow(diff_vis)
axes[2, col].set_title(f"|Δ|max={float(diff.max()):.3f}", fontsize=8)
for r in range(3):
axes[r, col].set_xticks([])
axes[r, col].set_yticks([])
axes[0, 0].set_ylabel("no jitter", fontsize=10)
axes[1, 0].set_ylabel("with jitter", fontsize=10)
axes[2, 0].set_ylabel("|Δ| (norm)", fontsize=10)
fig.suptitle(
f"Color jitter comparison (b={JITTER.brightness}, c={JITTER.contrast}, "
f"s={JITTER.saturation}, h={JITTER.hue})",
fontsize=11,
)
fig.tight_layout(rect=(0, 0, 1, 0.96))
fig.savefig(OUT_PNG, dpi=120)
print(f"wrote {OUT_PNG} ({OUT_PNG.stat().st_size/1024:.1f} KB)")
if __name__ == "__main__":
main()