Download dataset/base_dataset.py from ducido/diffusion_policy_gbc: direct link, hf CLI and curl.
- Browser
- Download file 1.38 kB
-
https://huggingface.co/ducido/diffusion_policy_gbc/resolve/main/dataset/base_dataset.py
- Command line
-
hf download hf://ducido/diffusion_policy_gbc/dataset/base_dataset.py
-
curl -L -o base_dataset.py https://huggingface.co/ducido/diffusion_policy_gbc/resolve/main/dataset/base_dataset.py
1.38 kB
| from typing import Dict | |
| import torch | |
| import torch.nn | |
| from diffusion_policy.model.common.normalizer import LinearNormalizer | |
| class BaseLowdimDataset(torch.utils.data.Dataset): | |
| def get_validation_dataset(self) -> 'BaseLowdimDataset': | |
| # return an empty dataset by default | |
| return BaseLowdimDataset() | |
| def get_normalizer(self, **kwargs) -> LinearNormalizer: | |
| raise NotImplementedError() | |
| def get_all_actions(self) -> torch.Tensor: | |
| raise NotImplementedError() | |
| def __len__(self) -> int: | |
| return 0 | |
| def __getitem__(self, idx: int) -> Dict[str, torch.Tensor]: | |
| """ | |
| output: | |
| obs: T, Do | |
| action: T, Da | |
| """ | |
| raise NotImplementedError() | |
| class BaseImageDataset(torch.utils.data.Dataset): | |
| def get_validation_dataset(self) -> 'BaseLowdimDataset': | |
| # return an empty dataset by default | |
| return BaseImageDataset() | |
| def get_normalizer(self, **kwargs) -> LinearNormalizer: | |
| raise NotImplementedError() | |
| def get_all_actions(self) -> torch.Tensor: | |
| raise NotImplementedError() | |
| def __len__(self) -> int: | |
| return 0 | |
| def __getitem__(self, idx: int) -> Dict[str, torch.Tensor]: | |
| """ | |
| output: | |
| obs: | |
| key: T, * | |
| action: T, Da | |
| """ | |
| raise NotImplementedError() | |