BonanDing commited on
Commit
71dedf8
·
1 Parent(s): 9ba5927

Add retry handling to DeMemWM latent dataset

Browse files
datasets/video/minecraft_video_dememwm_latent_dataset.py CHANGED
@@ -90,6 +90,7 @@ class MinecraftVideoDeMemWMLatentDataset(torch.utils.data.Dataset):
90
  self._target_offsets = np.arange(self.n_frames, dtype=np.int64) * self.frame_skip
91
  self._memory_segments = memory_segment_lengths(self.n_frames, self.memory_selection)
92
  self.memory_condition_length = sum(self._memory_segments[key] for key in SEGMENT_KEYS)
 
93
  self._cached_arrays_path = None
94
  self._cached_arrays = None
95
  self.data_paths = self.get_data_paths(split)
@@ -159,7 +160,31 @@ class MinecraftVideoDeMemWMLatentDataset(torch.utils.data.Dataset):
159
  return file_idx, int(idx - prev)
160
 
161
  def __getitem__(self, idx: int):
162
- return self.load_data(self.idx_remap[idx])
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
163
 
164
  def load_data(self, idx: int):
165
  # === 1. Resolve clip and load arrays ===
 
90
  self._target_offsets = np.arange(self.n_frames, dtype=np.int64) * self.frame_skip
91
  self._memory_segments = memory_segment_lengths(self.n_frames, self.memory_selection)
92
  self.memory_condition_length = sum(self._memory_segments[key] for key in SEGMENT_KEYS)
93
+ self.max_getitem_retries = int(cfg_get(cfg, "max_getitem_retries", 32))
94
  self._cached_arrays_path = None
95
  self._cached_arrays = None
96
  self.data_paths = self.get_data_paths(split)
 
160
  return file_idx, int(idx - prev)
161
 
162
  def __getitem__(self, idx: int):
163
+ dataset_len = len(self)
164
+ if dataset_len == 0:
165
+ raise IndexError(
166
+ f"MinecraftVideoDeMemWMLatentDataset split {self.split!r} contains no clips under "
167
+ f"{self.save_dir / self.split}"
168
+ )
169
+
170
+ start_idx = int(idx)
171
+ errors = []
172
+ last_error = None
173
+ for attempt in range(max(1, self.max_getitem_retries)):
174
+ candidate_idx = (start_idx + attempt) % dataset_len
175
+ try:
176
+ return self.load_data(self.idx_remap[candidate_idx])
177
+ except Exception as exc:
178
+ last_error = exc
179
+ if len(errors) < 5:
180
+ errors.append(f"idx={candidate_idx}: {type(exc).__name__}: {exc}")
181
+
182
+ details = "; ".join(errors) if errors else "no error details captured"
183
+ raise RuntimeError(
184
+ f"Failed to load MinecraftVideoDeMemWMLatentDataset item after "
185
+ f"{max(1, self.max_getitem_retries)} attempts starting at index {start_idx} "
186
+ f"for split {self.split!r}. Recent errors: {details}"
187
+ ) from last_error
188
 
189
  def load_data(self, idx: int):
190
  # === 1. Resolve clip and load arrays ===