jaxsnn.base.dataset

Modules

jaxsnn.base.dataset.circle

jaxsnn.base.dataset.constant

jaxsnn.base.dataset.dataloader

jaxsnn.base.dataset.linear

jaxsnn.base.dataset.shd

SHD Dataset class

jaxsnn.base.dataset.yinyang

Functions

jaxsnn.base.dataset.circle_dataset(rng: jax.Array, size: int, mirror: bool = True, bias_spike: Optional[float, None] = 0.0)Tuple[jax.Array, jax.Array]
jaxsnn.base.dataset.constant_dataset(t_max: float, size: int)Tuple[jax.Array, jax.Array]
jaxsnn.base.dataset.data_loader(dataset: Tuple[Any, Any], batch_size: int, num_batches: Optional[int, None] = None, rng: Optional[jax.Array, None] = None)
jaxsnn.base.dataset.linear_dataset(rng: jax.Array, size: int, mirror: bool, bias_spike: Optional[float, None])Tuple[jax.Array, jax.Array]
jaxsnn.base.dataset.shd_dataset(train: bool, download: bool, t_max: float, root: str, time_scale: float = 1.0, channel_start: int = 0, channel_modulo: int = 1, num_samples: Optional[int, None] = None, dtype: numpy.dtype = <class 'numpy.float32'>)Tuple[jax.Array, jax.Array]

SHD dataset

jaxsnn.base.dataset.yinyang_dataset(rng: jax.Array, size: int, mirror: bool, bias_spike: Optional[float, None])Tuple[jax.Array, jax.Array]