Source code for netloader.utils.types

"""
Type definitions.
"""
from __future__ import annotations

from warnings import warn
from typing import Union, TypeVar, Any, Sequence, Protocol, TypeAlias, ParamSpec, TYPE_CHECKING

from numpy import ndarray
from torch import nn, Tensor

if TYPE_CHECKING:
    from netloader.data import Data, DataList


ShapeLike: TypeAlias = list[int] | list[list[int]]
Param: TypeAlias = dict[str, Union[Tensor, 'Param']]
DataLike: TypeAlias = Union[ndarray, Tensor, 'Data[ArrayCT]']
ArrayLike: TypeAlias = ndarray | Tensor
TensorLike: TypeAlias = Union[Tensor, 'Data[Tensor]', 'DataList[Tensor | Data[Tensor]]']
NDArrayLike: TypeAlias = Union[ndarray, 'Data[ndarray]', 'DataList[ndarray | Data[ndarray]]']
DataListLike: TypeAlias = Union[DataLike, 'DataList[DataLike]']
ShapeListLike: TypeAlias = list[int] | Sequence['ShapeListLike']
TensorListLike: TypeAlias = Union[Tensor, 'DataList[Tensor]']
NDArrayListLike: TypeAlias = Union[ndarray, 'DataList[ndarray]']

T = TypeVar('T')
P = ParamSpec('P')
DataT = TypeVar('DataT', bound=DataLike)
ArrayT = TypeVar('ArrayT', bound=ArrayLike)
ModuleT = TypeVar('ModuleT', bound=nn.Module)
TensorT = TypeVar('TensorT', bound=TensorLike)
NDArrayT = TypeVar('NDArrayT', bound=NDArrayLike)
DataListT = TypeVar('DataListT', bound=DataListLike)
TensorListT = TypeVar('TensorListT', bound=TensorListLike)
ModuleListT = TypeVar('ModuleListT', bound=nn.Module | nn.ModuleList)

LossCT = TypeVar('LossCT', float, dict[str, float])
DataCT = TypeVar('DataCT', ndarray, Tensor, 'Data[ndarray]', 'Data[Tensor]')
ArrayCT = TypeVar('ArrayCT', ndarray, Tensor)
TensorLossCT = TypeVar('TensorLossCT', Tensor, dict[str, Tensor])


[docs] class DatasetProtocol(Protocol[DataT]): """ Protocol for datasets used in NetLoader Attributes ---------- extra : list[Any] | ndarray | None Additional data for each sample in the dataset of length N with shape (N,...) and type Any idxs : ndarray Index for each sample in the dataset with shape (N) and type int low_dim : ndarray | Tensor | None Low dimensional data for each sample in the dataset with shape (N) high_dim : ndarray | Tensor | object High dimensional data for each sample in the dataset with shape (N), this is required """ extra: list[Any] | ndarray | Tensor | None idxs: ndarray low_dim: DataT | None high_dim: DataT | None
[docs] def __len__(self) -> int: """ Returns the number of samples in the dataset Returns ------- int Number of samples in the dataset """
[docs] def __getitem__(self, idx: int) -> tuple[int, DataListT, DataListT, Any]: """ Parameters ---------- idx : int Sample index Returns ------- tuple[int, DataListT, DataListT, Any] Sample index, low dimensional data, high dimensional data, and extra data """
[docs] def step(self, epoch: float) -> None: """ Step the dataset for each training iteration. Parameters ---------- epoch : float Current epoch number """
DatasetT = TypeVar('DatasetT', bound=DatasetProtocol) __all__ = [ 'DatasetProtocol', 'Param', 'DataLike', 'ShapeLike', 'ArrayLike', 'TensorLike', 'NDArrayLike', 'DataListLike', 'ShapeListLike', 'TensorListLike', 'NDArrayListLike', 'T', 'P', 'DataT', 'ArrayT', 'ModuleT', 'TensorT', 'NDArrayT', 'DatasetT', 'DataListT', 'TensorListT', 'ModuleListT', 'LossCT', 'DataCT', 'ArrayCT', 'TensorLossCT', ] def __getattr__(name: str) -> Any: if name == 'ArrayTC': warn( 'ArrayTC is deprecated, use ArrayCT instead', DeprecationWarning, stacklevel=2, ) return ArrayCT raise AttributeError(f"module {__name__} has no attribute {name}")