pulsatrix
Loading...
Searching...
No Matches
transform.hpp
Go to the documentation of this file.
1
5#pragma once
6
7#include <memory>
8#include <utility>
9#include <vector>
10
11#include "pulsatrix/dataset.hpp"
12
13namespace pulsatrix {
14
20class Transform {
21public:
22 virtual ~Transform() = default;
23
29 [[nodiscard]] virtual Sample apply(Sample sample) const = 0;
30};
31
38class Compose : public Transform {
39public:
40 explicit Compose(std::vector<std::shared_ptr<Transform>> steps) : steps_(std::move(steps)) {}
41
42 [[nodiscard]] Sample apply(Sample sample) const override {
43 for (const auto& step : steps_) {
44 sample = step->apply(std::move(sample));
45 }
46 return sample;
47 }
48
49private:
50 std::vector<std::shared_ptr<Transform>> steps_;
51};
52
57class TransformDataset : public Dataset {
58public:
59 TransformDataset(std::shared_ptr<Dataset> base, std::shared_ptr<Transform> transform)
60 : base_(std::move(base)), transform_(std::move(transform)) {}
61
62 [[nodiscard]] int64_t size() const override { return base_->size(); }
63
64 [[nodiscard]] Sample get(int64_t index) const override { return transform_->apply(base_->get(index)); }
65
66private:
67 std::shared_ptr<Dataset> base_;
68 std::shared_ptr<Transform> transform_;
69};
70
71} // namespace pulsatrix
Eager, ordered list of Transforms applied in sequence – pulsatrix's analogue of torchvision....
Definition transform.hpp:38
Sample apply(Sample sample) const override
Applies this transform to a sample.
Definition transform.hpp:42
Compose(std::vector< std::shared_ptr< Transform > > steps)
Definition transform.hpp:40
Random-access dataset abstraction – pulsatrix's analogue of PyTorch's torch.utils....
Definition dataset.hpp:33
Decorates a Dataset with a Transform, applied to every sample get() returns – lets Transform/Compose ...
Definition transform.hpp:57
int64_t size() const override
Number of samples in this dataset.
Definition transform.hpp:62
Sample get(int64_t index) const override
Loads one sample by index.
Definition transform.hpp:64
TransformDataset(std::shared_ptr< Dataset > base, std::shared_ptr< Transform > transform)
Definition transform.hpp:59
A single sample-level preprocessing step (normalize, augment, tokenize, ...).
Definition transform.hpp:20
virtual Sample apply(Sample sample) const =0
Applies this transform to a sample.
virtual ~Transform()=default
Random-access dataset abstraction – Sample, Dataset (size()/get()).
Definition acquisition_functions.hpp:16
One dataset sample: an ordered list of Tensor fields (e.g. {features, label} or {image,...
Definition dataset.hpp:19