Download Transfer_and_Test/src/two_condition_dataset.py from DancingNow/swag-train-bundle: direct link, hf CLI and curl.
- Browser
- Download file 2.66 kB
-
https://huggingface.co/DancingNow/swag-train-bundle/resolve/main/Transfer_and_Test/src/two_condition_dataset.py
- Command line
-
hf download hf://DancingNow/swag-train-bundle/Transfer_and_Test/src/two_condition_dataset.py
-
curl -L -o two_condition_dataset.py https://huggingface.co/DancingNow/swag-train-bundle/resolve/main/Transfer_and_Test/src/two_condition_dataset.py
2.66 kB
| from __future__ import annotations | |
| from pathlib import Path | |
| import h5py | |
| import numpy as np | |
| import torch | |
| from torch.utils.data import Dataset | |
| class TwoConditionDataset(Dataset): | |
| def __init__(self, path: Path, cache_in_memory: bool = False): | |
| self.path = Path(path) | |
| self.cache_in_memory = bool(cache_in_memory) | |
| with h5py.File(self.path, 'r') as f: | |
| self.data_shape = tuple(f['data'].shape) | |
| self.label_shape = tuple(f['labels'].shape) | |
| self.dtype = f['data'].dtype | |
| if self.data_shape != (3000, 6000, 3) or self.label_shape != (3000, 2): | |
| raise ValueError(f'unexpected shapes: data={self.data_shape}, labels={self.label_shape}') | |
| self._file = None | |
| self._data = None | |
| self._labels = None | |
| self._waveforms = None | |
| self._cached_labels = None | |
| if self.cache_in_memory: | |
| with h5py.File(self.path, 'r') as f: | |
| waveforms = np.asarray(f['data'], dtype=np.float32) | |
| labels = np.asarray(f['labels'], dtype=np.float32) | |
| if not np.isfinite(waveforms).all() or not np.isfinite(labels).all(): | |
| raise ValueError('training HDF5 contains non-finite values') | |
| self._waveforms = np.ascontiguousarray(waveforms.transpose(0, 2, 1)) | |
| self._cached_labels = np.ascontiguousarray(labels) | |
| print( | |
| f'Cached training data in RAM: waveforms={self._waveforms.nbytes / 2**20:.1f} MiB, ' | |
| f'labels={self._cached_labels.nbytes / 2**20:.1f} MiB' | |
| ) | |
| def __len__(self): | |
| return self.data_shape[0] | |
| def _open(self): | |
| if self._file is None: | |
| self._file = h5py.File(self.path, 'r') | |
| self._data, self._labels = self._file['data'], self._file['labels'] | |
| def __getitem__(self, i): | |
| if self._waveforms is not None: | |
| x = self._waveforms[i] | |
| raw = self._cached_labels[i] | |
| else: | |
| self._open() | |
| x = np.asarray(self._data[i], dtype=np.float32) | |
| raw = np.asarray(self._labels[i], dtype=np.float32) | |
| if not np.isfinite(x).all() or not np.isfinite(raw).all(): | |
| raise ValueError(f'non-finite row {i}') | |
| label = raw | |
| if not (0 <= label[0] < label[1] <= 6000): | |
| raise ValueError(f'invalid label row {i}: {label.tolist()}') | |
| waveform = x if self._waveforms is not None else x.transpose(1, 0).copy() | |
| return torch.from_numpy(waveform), torch.from_numpy(label) | |
| def close(self): | |
| if self._file is not None: | |
| self._file.close() | |
| self._file = None | |