import abc
import pathlib
import random
import typing
import numpy
class Dataset(abc.ABC):
@property
@abc.abstractmethod
def name(self) -> str:
pass
@property
@abc.abstractmethod
def data(self):
pass
@property
@abc.abstractmethod
def indices(self) -> numpy.ndarray:
pass
@abc.abstractmethod
def __eq__(self, other: "Dataset") -> bool:
pass
@property
def cardinality(self) -> int:
return self.indices.shape[0]
@property
@abc.abstractmethod
def max_instance_size(self) -> int:
pass
@property
@abc.abstractmethod
def approx_memory_size(self) -> int:
pass
@abc.abstractmethod
def __getitem__(self, item: typing.Union[int, typing.Iterable[int]]):
pass
@abc.abstractmethod
def subset(self, indices: list[int], subset_name: str) -> "Dataset":
pass
def max_batch_size(self, available_memory: int) -> int:
return available_memory // self.max_instance_size
def complement_indices(self, indices: list[int]) -> list[int]:
return list(set(range(self.cardinality)) - set(indices))
def subsample_indices(self, n: int) -> tuple[list[int], list[int]]:
if not (0 < n < self.cardinality):
msg = f"`n` must be a positive integer smaller than the cardinality of the dataset ({self.cardinality}). Got {n} instead."
raise ValueError(
msg,
)
indices = list(range(n))
random.shuffle(indices)
return indices[:n], indices[n:]
class TabularDataset(Dataset):
def __init__(self, data: numpy.ndarray, name: str) -> None:
self.__data = data
self.__name = name
self.__indices = numpy.asarray(list(range(data.shape[0])), dtype=numpy.uint)
@property
def name(self) -> str:
return self.__name
@property
def data(self) -> numpy.ndarray:
return self.__data
@property
def indices(self) -> numpy.ndarray:
return self.__indices
def __eq__(self, other: "TabularDataset") -> bool:
return self.__name == other.__name
@property
def max_instance_size(self) -> int:
item_size = self.__data.itemsize
num_items = numpy.prod(self.__data.shape[1:])
return int(item_size * num_items)
@property
def approx_memory_size(self) -> int:
return self.cardinality * self.max_instance_size
def __getitem__(
self,
item: typing.Union[int, typing.Iterable[int]],
) -> numpy.ndarray:
return self.__data[item]
def subset(self, indices: list[int], subset_name: str) -> "TabularDataset":
return TabularDataset(self.__data[indices, :], subset_name)
class TabularMMap(TabularDataset):
def __init__(
self,
file_path: typing.Union[pathlib.Path, str],
name: str,
indices: list[int] = None,
) -> None:
self.file_path = file_path
data = numpy.load(str(file_path), mmap_mode="r")
indices = list(range(data.shape[0])) if indices is None else indices
self.__indices = numpy.asarray(indices, dtype=numpy.uint)
if self.__indices.max(initial=0) >= data.shape[0]:
msg = "Invalid indices provided."
raise IndexError(msg)
super().__init__(data, name)
@property
def indices(self) -> numpy.ndarray:
return self.__indices
def __getitem__(self, item: typing.Union[int, typing.Iterable[int]]):
indices = self.__indices[item]
return numpy.asarray(self.data[indices, :])
def subset(self, indices: list[int], subset_name: str) -> "TabularDataset":
return TabularMMap(self.file_path, subset_name, indices)
__all__ = [
"Dataset",
"TabularDataset",
"TabularMMap",
]