pulsatrix
Loading...
Searching...
No Matches
sampler.hpp
Go to the documentation of this file.
1
5#pragma once
6
7#include <cstdint>
8#include <optional>
9#include <random>
10#include <vector>
11
12namespace pulsatrix {
13
20class Sampler {
21public:
22 virtual ~Sampler() = default;
23
31 virtual void reset(int64_t dataset_size) = 0;
32
37 [[nodiscard]] virtual std::optional<int64_t> next() = 0;
38};
39
41class SequentialSampler : public Sampler {
42public:
43 void reset(int64_t dataset_size) override;
44 [[nodiscard]] std::optional<int64_t> next() override;
45
46private:
47 int64_t size_ = 0;
48 int64_t position_ = 0;
49};
50
56class ShuffleSampler : public Sampler {
57public:
58 explicit ShuffleSampler(unsigned seed);
59 void reset(int64_t dataset_size) override;
60 [[nodiscard]] std::optional<int64_t> next() override;
61
62private:
63 unsigned seed_;
64 std::mt19937 rng_;
65 std::vector<int64_t> indices_;
66 size_t position_ = 0;
67};
68
69} // namespace pulsatrix
Produces the order in which a DataLoader visits a Dataset's indices for one epoch.
Definition sampler.hpp:20
virtual std::optional< int64_t > next()=0
Fetches the next index in this epoch's order.
virtual void reset(int64_t dataset_size)=0
Begins a new epoch over a dataset of the given size.
virtual ~Sampler()=default
Visits indices [0, dataset_size) in ascending order.
Definition sampler.hpp:41
void reset(int64_t dataset_size) override
Begins a new epoch over a dataset of the given size.
std::optional< int64_t > next() override
Fetches the next index in this epoch's order.
Visits indices [0, dataset_size) in a seeded pseudo-random permutation – reproducible across runs giv...
Definition sampler.hpp:56
ShuffleSampler(unsigned seed)
void reset(int64_t dataset_size) override
Begins a new epoch over a dataset of the given size.
std::optional< int64_t > next() override
Fetches the next index in this epoch's order.
Definition acquisition_functions.hpp:16