syCen commited on
Commit
1bdcf8a
·
verified ·
1 Parent(s): 94d70bf

Create physical_dataset.py

Browse files
Files changed (1) hide show
  1. physical_dataset.py +112 -0
physical_dataset.py ADDED
@@ -0,0 +1,112 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ """
2
+ physical_dataset.py
3
+
4
+ Dataset for Contact / Force VAE finetuning.
5
+
6
+ Reads clips.json (from build_clip_index.py) and statistics.json, loads each
7
+ 17-frame clip of per-finger contact (2,H,W) or force (6,H,W) maps, normalizes,
8
+ and returns a tensor plus an active mask for weighted reconstruction loss.
9
+
10
+ Normalization (background stays 0 in both cases):
11
+ contact: x[ch] / contact_ch_max[ch] -> (0, 1], background 0
12
+ force: x[ch] / force_ch_std[ch] -> ~O(1) signed, background 0
13
+ Rationale: contact/force are sparse; subtracting a mean would turn the 0
14
+ background into a nonzero value and destroy sparsity, so we only scale.
15
+
16
+ Returns per item:
17
+ {
18
+ "data": (C, T, H, W) float32, normalized
19
+ "active_mask": (1, T, H, W) float32, 1 where any channel != 0
20
+ "episode": str
21
+ }
22
+ """
23
+
24
+ import json
25
+ import os
26
+ import numpy as np
27
+ import torch
28
+ from torch.utils.data import Dataset
29
+
30
+
31
+ class PhysicalClipDataset(Dataset):
32
+ def __init__(self, clips_json, statistics_json, source_root,
33
+ modality="contact", eps=1e-6):
34
+ assert modality in ("contact", "force")
35
+ self.modality = modality
36
+ self.source_root = source_root
37
+ self.eps = eps
38
+
39
+ with open(clips_json) as f:
40
+ blob = json.load(f)
41
+ self.clips = blob["clips"] if isinstance(blob, dict) and "clips" in blob else blob
42
+ self.config = blob.get("config", {}) if isinstance(blob, dict) else {}
43
+
44
+ with open(statistics_json) as f:
45
+ self.stats = json.load(f)
46
+
47
+ # build per-channel scale vector, shaped (C,1,1,1) for broadcasting
48
+ if modality == "contact":
49
+ scale = np.asarray(self.stats["contact_ch_max"], dtype=np.float32) # (2,)
50
+ else:
51
+ scale = np.asarray(self.stats["force_ch_std"], dtype=np.float32) # (6,)
52
+ scale = np.maximum(scale, eps) # avoid div-by-zero
53
+ self.scale = scale.reshape(-1, 1, 1, 1) # (C,1,1,1)
54
+ self.n_ch = self.scale.shape[0]
55
+
56
+ self.path_key = "contact_paths" if modality == "contact" else "force_paths"
57
+
58
+ def __len__(self):
59
+ return len(self.clips)
60
+
61
+ def _load_clip(self, clip):
62
+ """Load 17 frames -> (C, T, H, W) raw float32."""
63
+ ep = clip["episode"]
64
+ frames = []
65
+ for rel in clip[self.path_key]:
66
+ arr = np.load(os.path.join(self.source_root, ep, rel)) # (C,H,W)
67
+ frames.append(arr.astype(np.float32))
68
+ # stack along time: list of (C,H,W) -> (T,C,H,W) -> (C,T,H,W)
69
+ x = np.stack(frames, axis=0).transpose(1, 0, 2, 3)
70
+ return x
71
+
72
+ def __getitem__(self, idx):
73
+ clip = self.clips[idx]
74
+ x = self._load_clip(clip) # (C,T,H,W) raw
75
+
76
+ # active mask BEFORE normalization (any channel nonzero)
77
+ active = (np.abs(x).sum(axis=0, keepdims=True) > 0).astype(np.float32) # (1,T,H,W)
78
+
79
+ # normalize: scale only, background 0 stays 0
80
+ x = x / self.scale # (C,T,H,W)
81
+
82
+ return {
83
+ "data": torch.from_numpy(x), # (C,T,H,W)
84
+ "active_mask": torch.from_numpy(active), # (1,T,H,W)
85
+ "episode": clip["episode"],
86
+ }
87
+
88
+
89
+ def collate_physical(batch):
90
+ """Stack into (B,C,T,H,W). Assumes uniform shape (it is: fixed 17 frames)."""
91
+ data = torch.stack([b["data"] for b in batch], dim=0)
92
+ mask = torch.stack([b["active_mask"] for b in batch], dim=0)
93
+ episodes = [b["episode"] for b in batch]
94
+ return {"data": data, "active_mask": mask, "episodes": episodes}
95
+
96
+
97
+ if __name__ == "__main__":
98
+ import argparse
99
+ ap = argparse.ArgumentParser()
100
+ ap.add_argument("--clips", required=True)
101
+ ap.add_argument("--stats", required=True)
102
+ ap.add_argument("--source_root", required=True)
103
+ ap.add_argument("--modality", choices=["contact", "force"], default="contact")
104
+ args = ap.parse_args()
105
+
106
+ ds = PhysicalClipDataset(args.clips, args.stats, args.source_root, args.modality)
107
+ print(f"{args.modality} dataset: {len(ds)} clips, scale={ds.scale.ravel()}")
108
+ item = ds[0]
109
+ print("data:", tuple(item["data"].shape), item["data"].dtype,
110
+ "min/max:", float(item["data"].min()), float(item["data"].max()))
111
+ print("active_mask:", tuple(item["active_mask"].shape),
112
+ "active frac:", float(item["active_mask"].mean()))