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 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 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 ===
|