swag-train-bundle / Transfer_and_Test /src /two_condition_dataset.py
DancingNow's picture
Add files using upload-large-folder tool
275b5a1 verified
Raw History Blame Contribute Delete
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