pulsatrix
Loading...
Searching...
No Matches
data_loader.hpp
Go to the documentation of this file.
1
5#pragma once
6
7#include <memory>
8#include <optional>
9
10#include "pulsatrix/collate.hpp"
11#include "pulsatrix/dataset.hpp"
14#include "pulsatrix/sampler.hpp"
15
16namespace pulsatrix {
17
20 int64_t batch_size = 1;
21 bool shuffle = false;
24 std::optional<unsigned> shuffle_seed;
26 int num_workers = 0;
27 int64_t prefetch_batches = 2;
28 bool drop_last = false;
30};
31
42public:
47 DataLoader(std::shared_ptr<Dataset> dataset, DeviceBackend* backend, DataLoaderOptions options = {});
48
55 DataLoader(std::shared_ptr<IterableDataset> dataset, DeviceBackend* backend, DataLoaderOptions options = {});
56
59
65 [[nodiscard]] std::optional<Batch> next_batch();
66
72 [[nodiscard]] int64_t num_batches() const;
73
74private:
75 [[nodiscard]] std::optional<Sample> fetch_next_sample();
76
77 std::shared_ptr<Dataset> dataset_; // non-null iff constructed from a Dataset
78 std::shared_ptr<IterableDataset> iterable_dataset_; // non-null iff constructed from an IterableDataset
79 DeviceBackend* backend_;
80 DataLoaderOptions options_;
81 std::unique_ptr<Sampler> sampler_; // only used when dataset_ is set
82};
83
84} // namespace pulsatrix
Orchestrates sampling and collation into batches – pulsatrix's DataLoader (PyTorch DataLoader / torch...
Definition data_loader.hpp:41
void reset_epoch()
Resets to the start of a new epoch (re-seeds/reshuffles the sampler, or resets the stream).
int64_t num_batches() const
Number of batches per epoch.
DataLoader(std::shared_ptr< Dataset > dataset, DeviceBackend *backend, DataLoaderOptions options={})
Constructs a DataLoader over a random-access Dataset.
std::optional< Batch > next_batch()
Fetches the next batch.
DataLoader(std::shared_ptr< IterableDataset > dataset, DeviceBackend *backend, DataLoaderOptions options={})
Constructs a DataLoader over a streaming IterableDataset.
Vendor-agnostic compute/memory backend. CPUBackend, CUDABackend (Phase 1.5), and HIPBackend (Phase 1....
Definition device_backend.hpp:219
Batch assembly – Batch, CollateFn, DefaultCollate.
Random-access dataset abstraction – Sample, Dataset (size()/get()).
Abstract interface isolating vendor-specific memory/compute operations from Tensor/ComputationGraph.
Streaming dataset abstraction – reset()/next() for sources with no random access.
Definition acquisition_functions.hpp:16
Batch DefaultCollate(std::vector< Sample > samples, DeviceBackend *backend)
Stacks a list of samples into one Batch, field-by-field, via Tensor::Stack – pulsatrix's default Coll...
std::function< Batch(std::vector< Sample >, DeviceBackend *)> CollateFn
A function assembling a list of Samples into one Batch – pulsatrix's analogue of PyTorch's collate_fn...
Definition collate.hpp:30
Index-order abstraction for DataLoader – SequentialSampler, ShuffleSampler.
Configuration for a DataLoader.
Definition data_loader.hpp:19
int num_workers
0 = fully synchronous, no threads spawned (this phase's only exercised path).
Definition data_loader.hpp:26
std::optional< unsigned > shuffle_seed
Shuffle seed. Unset draws one from the global seed stream (next_seed(), FND-7) when the loader is bui...
Definition data_loader.hpp:24
bool shuffle
Definition data_loader.hpp:21
int64_t prefetch_batches
Definition data_loader.hpp:27
CollateFn collate_fn
Definition data_loader.hpp:29
bool drop_last
Definition data_loader.hpp:28
int64_t batch_size
Definition data_loader.hpp:20