Spaces:
Running on Zero
Running on Zero
File size: 14,335 Bytes
23a59ea | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 142 143 144 145 146 147 148 149 150 151 152 153 154 155 156 157 158 159 160 161 162 163 164 165 166 167 168 169 170 171 172 173 174 175 176 177 178 179 180 181 182 183 184 185 186 187 188 189 190 191 192 193 194 195 196 197 198 199 200 201 202 203 204 205 206 207 208 209 210 211 212 213 214 215 216 217 218 219 220 221 222 223 224 225 226 227 228 229 230 231 232 233 234 235 236 237 238 239 240 241 242 243 244 245 246 247 248 249 250 251 252 253 254 255 256 257 258 259 260 261 262 263 264 265 266 267 268 269 270 271 272 273 274 275 276 277 278 279 280 281 282 283 284 285 286 287 288 289 290 291 292 293 294 295 296 297 298 299 300 301 302 303 304 305 306 307 308 309 310 311 312 313 314 315 | # sharded_frame_dataset.py
import json
import os
import bisect
import random
from collections import OrderedDict
from pathlib import Path
from typing import Sequence, List, Dict, Union, Mapping, Optional
import torch
import torch.distributed as dist
from torch.utils.data import Dataset
class ShardedFrameDataset(Dataset):
"""
Samples contiguous sequences from preprocessed shards across multiple roots:
root/<task>/<task>_index.json with {"shard_name": num_frames, ...}
root/<task>/*.pt with {"frames": (N, 3, H, W) uint8}
Returns: (T, 3, H, W) float32 in [0,1], where T = seq_len.
If iid_sampling=True, ignores idx and samples a random sequence from this
rank's slice of shards. Each worker holds a "current" shard for
samples_per_shard draws before picking a new one — this amortizes the cost
of loading a shard from disk over many sequences and is the main lever for
avoiding I/O-bound training on large datasets.
cache_size controls how many shards are kept in memory (LRU eviction) per
worker process. With samples_per_shard > 1, a small cache (4-8) suffices
since most accesses go to the current shard.
ddp_partition controls whether iid_sampling shards are partitioned across
DDP ranks. Set to False for rank-0-only validation loaders so rank 0 can
still see the full validation set.
"""
def __init__(
self,
outdirs: Union[str, Sequence[str]],
tasks: Sequence[str] = (),
seq_len: int = 16,
iid_sampling: bool = True,
cache_size: int = 8,
samples_per_shard: int = 1,
ddp_partition: bool = True,
task_weights: Optional[Mapping[str, float]] = None,
verbose: bool = True,
return_task_idx: bool = False,
):
super().__init__()
assert outdirs is not None, "outdirs must be specified"
if isinstance(outdirs, (str, Path)):
self.outdirs = [str(outdirs)]
else:
self.outdirs = [str(p) for p in outdirs]
self.tasks = list(tasks)
self.seq_len = int(seq_len)
self.iid_sampling = bool(iid_sampling)
self._cache_size = max(1, int(cache_size))
self.samples_per_shard = max(1, int(samples_per_shard))
self.verbose = bool(verbose)
self.return_task_idx = bool(return_task_idx)
# Resolve optional per-task sampling weights. Mirrors WMDataset: a dict
# {task_name: weight} sets relative task draw probabilities; tasks not
# present get weight 0 (excluded from sampling). If None, fall back to
# the legacy behavior where shard weight = num_starts, which makes
# P(task) ∝ total valid_starts (short-trajectory domains get starved).
if task_weights is not None:
tw_list: List[float] = []
missing: List[str] = []
for t in self.tasks:
if t in task_weights:
w = float(task_weights[t])
if w < 0:
raise ValueError(f"task_weights[{t!r}] is negative: {w}")
tw_list.append(w)
else:
missing.append(t)
tw_list.append(0.0)
if missing and self.verbose:
print(f"[ShardedFrameDataset] Warning: {len(missing)} tasks have no "
f"task_weights entry and will be excluded from sampling "
f"(e.g. {missing[:5]})")
self.task_weights: Optional[List[float]] = tw_list
else:
self.task_weights = None
self.shards: List[Dict] = []
self.cum_starts: List[int] = []
total_starts = 0
for root in self.outdirs:
root = Path(root)
for ti, task in enumerate(self.tasks):
task_dir = root / task
if not task_dir.exists():
continue
index_path = task_dir / f"{task}_index.json"
if index_path.exists():
# Fast path: read frame counts from the metadata index written by
# preprocess_dataset.py — no tensor data is loaded at init time.
with open(index_path) as f:
index = json.load(f)
for shard_name, num_frames in sorted(index.items()):
path = task_dir / shard_name
if not path.exists():
print(f"[ShardedFrameDataset] Shard {path} listed in index but missing, skipping")
continue
N = int(num_frames)
if N < self.seq_len:
print(f"[ShardedFrameDataset] Skipping {path} (N={N} < seq_len={self.seq_len})")
continue
num_starts = N - self.seq_len + 1
self.shards.append({"path": str(path), "num_frames": N, "num_starts": num_starts, "task_idx": ti})
total_starts += num_starts
self.cum_starts.append(total_starts)
else:
# Slow fallback: load every shard to inspect its shape.
# Run preprocess_dataset.py to generate index files and avoid this.
print(f"[ShardedFrameDataset] No index for task={task} in {root}, scanning shards (slow)")
for fname in sorted(os.listdir(task_dir)):
if not fname.endswith(".pt"):
continue
path = task_dir / fname
try:
td = torch.load(path, map_location="cpu", weights_only=True)
except Exception as e:
print(f"[ShardedFrameDataset] Skipping shard {path} (load error): {e}")
continue
frames = td.get("frames", None)
if not isinstance(frames, torch.Tensor):
print(f"[ShardedFrameDataset] Skipping shard {path} (no 'frames' tensor)")
continue
if frames.ndim != 4 or frames.shape[1] != 3:
print(f"[ShardedFrameDataset] Skipping shard {path} (unexpected shape {frames.shape})")
continue
N = int(frames.shape[0])
if N < self.seq_len:
print(f"[ShardedFrameDataset] Skipping shard {path} (N={N} < seq_len={self.seq_len})")
continue
num_starts = N - self.seq_len + 1
self.shards.append({"path": str(path), "num_frames": N, "num_starts": num_starts, "task_idx": ti})
total_starts += num_starts
self.cum_starts.append(total_starts)
self.total_starts = total_starts
# ---- DDP-aware sharding for iid_sampling ----
# When running under DDP with iid_sampling=True, partition shards across
# ranks so each rank only samples from its own slice. This (a) makes the
# "epoch" concept meaningful by removing cross-rank overlap, and (b)
# shrinks each worker's shard pool by world_size, dramatically improving
# per-worker LRU cache effectiveness.
rank, world_size = 0, 1
if (
ddp_partition
and self.iid_sampling
and dist.is_available()
and dist.is_initialized()
):
rank = dist.get_rank()
world_size = dist.get_world_size()
if world_size > 1 and len(self.shards) >= world_size:
# Strided (round-robin) partitioning rather than contiguous: shards
# are appended in task order, so a contiguous slice would lock each
# rank into a small set of tasks. Striding interleaves tasks across
# ranks, so every rank (and therefore the rank-0 viz) sees a
# task-diverse sample.
self._iid_shard_indices = list(range(rank, len(self.shards), world_size))
else:
self._iid_shard_indices = list(range(len(self.shards)))
# Precompute selection weights. Default: proportional to num_starts so
# the marginal distribution over sequences stays uniform across shards
# of varying length — but this makes P(task) ∝ total valid_starts.
# When task_weights is set, keep the intra-task distribution ∝ num_starts
# (for cache locality) but rescale per-task totals to match task_weights.
# Shards of tasks with weight 0 (or tasks absent from the dict) are
# dropped from the sampling pool entirely.
if self._iid_shard_indices:
if self.task_weights is None:
weights = [float(self.shards[i]["num_starts"]) for i in self._iid_shard_indices]
else:
# Group this rank's shards by task, compute per-task scale.
task_to_total_ns: Dict[int, float] = {}
for i in self._iid_shard_indices:
ti = self.shards[i]["task_idx"]
task_to_total_ns[ti] = task_to_total_ns.get(ti, 0.0) + float(self.shards[i]["num_starts"])
task_scale: Dict[int, float] = {}
for ti, total_ns in task_to_total_ns.items():
w_task = float(self.task_weights[ti])
if w_task <= 0 or total_ns <= 0:
task_scale[ti] = 0.0
else:
task_scale[ti] = w_task / total_ns
kept_indices: List[int] = []
weights = []
for i in self._iid_shard_indices:
ti = self.shards[i]["task_idx"]
s = task_scale.get(ti, 0.0)
if s <= 0:
continue
kept_indices.append(i)
weights.append(s * float(self.shards[i]["num_starts"]))
self._iid_shard_indices = kept_indices
total_w = float(sum(weights))
if total_w > 0:
self._iid_shard_weights = [w / total_w for w in weights]
else:
self._iid_shard_weights = []
else:
self._iid_shard_weights = []
if self.total_starts == 0:
print("[ShardedFrameDataset] WARNING: no usable sequences found in outdirs")
else:
extra = ""
if world_size > 1 and self.iid_sampling and ddp_partition:
extra = (
f", ddp_rank={rank}/{world_size}, "
f"local_shards={len(self._iid_shard_indices):,}"
)
if self.task_weights is not None:
nz = sum(1 for w in self.task_weights if w > 0)
extra += f", task_weights active ({nz}/{len(self.tasks)} tasks with w>0)"
print(
f"[ShardedFrameDataset] roots={len(self.outdirs)}, "
f"shards={len(self.shards):,}, seq_starts={self.total_starts:,}, "
f"samples_per_shard={self.samples_per_shard}{extra}"
)
# LRU shard cache: most-recently-used shards stay resident in memory.
# With samples_per_shard > 1, the bulk of accesses hit the current
# shard, so a small cache_size suffices.
self._cache: OrderedDict[str, torch.Tensor] = OrderedDict()
# Per-worker state for samples_per_shard. These attributes are inherited
# by forked workers but mutated independently in each worker process,
# so each worker maintains its own "currently active" shard.
self._cur_shard_idx: Union[int, None] = None
self._cur_shard_draws: int = 0
def __len__(self) -> int:
return self.total_starts
def _load_shard(self, path: str) -> torch.Tensor:
if path in self._cache:
self._cache.move_to_end(path) # mark as most-recently-used
return self._cache[path]
td = torch.load(path, map_location="cpu", weights_only=True)
frames = td["frames"]
if len(self._cache) >= self._cache_size:
self._cache.popitem(last=False) # evict least-recently-used
self._cache[path] = frames
return frames
def _map_global_start_to_shard(self, global_start: int) -> tuple[int, int]:
# global_start in [0, total_starts)
shard_idx = bisect.bisect_right(self.cum_starts, global_start)
prev_cum = 0 if shard_idx == 0 else self.cum_starts[shard_idx - 1]
start_idx_in_shard = global_start - prev_cum
return shard_idx, start_idx_in_shard
def __getitem__(self, idx: int) -> torch.Tensor:
if self.total_starts == 0:
raise IndexError("Empty dataset")
if self.iid_sampling:
# Reuse the current shard for samples_per_shard draws before picking
# a new one. This is the key locality optimization that lets large
# datasets train without saturating disk bandwidth.
if (
self._cur_shard_idx is None
or self._cur_shard_draws >= self.samples_per_shard
):
self._cur_shard_idx = random.choices(
self._iid_shard_indices,
weights=self._iid_shard_weights,
k=1,
)[0]
self._cur_shard_draws = 0
shard_idx = self._cur_shard_idx
self._cur_shard_draws += 1
meta = self.shards[shard_idx]
start = random.randrange(meta["num_starts"])
else:
if idx < 0 or idx >= self.total_starts:
raise IndexError(idx)
shard_idx, start = self._map_global_start_to_shard(int(idx))
meta = self.shards[shard_idx]
frames = self._load_shard(meta["path"]) # (N, 3, H, W)
end = start + self.seq_len
seq_u8 = frames[start:end] # (T, 3, H, W), guaranteed valid by construction
seq = seq_u8.to(torch.float32) / 255.0
if self.return_task_idx:
return {"frames": seq, "task_idx": int(meta["task_idx"])}
return seq
|