mirror of
https://github.com/huggingface/lerobot.git
synced 2026-07-23 17:56:07 +00:00
fix(dataset): fix data access bottleneck for faster training (#2408)
This commit is contained in:
@@ -940,11 +940,26 @@ class LeRobotDataset(torch.utils.data.Dataset):
|
|||||||
return query_timestamps
|
return query_timestamps
|
||||||
|
|
||||||
def _query_hf_dataset(self, query_indices: dict[str, list[int]]) -> dict:
|
def _query_hf_dataset(self, query_indices: dict[str, list[int]]) -> dict:
|
||||||
return {
|
"""
|
||||||
key: torch.stack(self.hf_dataset[q_idx][key])
|
Query dataset for indices across keys, skipping video keys.
|
||||||
for key, q_idx in query_indices.items()
|
|
||||||
if key not in self.meta.video_keys
|
Tries column-first [key][indices] for speed, falls back to row-first.
|
||||||
}
|
|
||||||
|
Args:
|
||||||
|
query_indices: Dict mapping keys to index lists to retrieve
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
Dict with stacked tensors of queried data (video keys excluded)
|
||||||
|
"""
|
||||||
|
result: dict = {}
|
||||||
|
for key, q_idx in query_indices.items():
|
||||||
|
if key in self.meta.video_keys:
|
||||||
|
continue
|
||||||
|
try:
|
||||||
|
result[key] = torch.stack(self.hf_dataset[key][q_idx])
|
||||||
|
except (KeyError, TypeError, IndexError):
|
||||||
|
result[key] = torch.stack(self.hf_dataset[q_idx][key])
|
||||||
|
return result
|
||||||
|
|
||||||
def _query_videos(self, query_timestamps: dict[str, list[float]], ep_idx: int) -> dict[str, torch.Tensor]:
|
def _query_videos(self, query_timestamps: dict[str, list[float]], ep_idx: int) -> dict[str, torch.Tensor]:
|
||||||
"""Note: When using data workers (e.g. DataLoader with num_workers>0), do not call this function
|
"""Note: When using data workers (e.g. DataLoader with num_workers>0), do not call this function
|
||||||
|
|||||||
Reference in New Issue
Block a user