pulsatrix
Loading...
Searching...
No Matches
mnist_dataset_adapter.hpp
Go to the documentation of this file.
1
5#pragma once
6
7#include <stdexcept>
8#include <utility>
9
10#include "pulsatrix/dataset.hpp"
12
13namespace pulsatrix {
14
24public:
26 : dataset_(std::move(dataset)), backend_(backend) {}
27
28 [[nodiscard]] int64_t size() const override { return static_cast<int64_t>(dataset_.images.size()); }
29
37 [[nodiscard]] Sample get(int64_t index) const override {
38 if (index < 0 || index >= size()) {
39 throw std::out_of_range("MnistDatasetAdapter::get: index out of range");
40 }
41 size_t i = static_cast<size_t>(index);
42 Tensor label(Shape({1}), backend_, {static_cast<float>(dataset_.labels[i])});
43 return Sample{{dataset_.images[i], std::move(label)}};
44 }
45
46private:
47 MnistDataset dataset_;
48 DeviceBackend* backend_;
49};
50
51} // namespace pulsatrix
Random-access dataset abstraction – pulsatrix's analogue of PyTorch's torch.utils....
Definition dataset.hpp:33
Vendor-agnostic compute/memory backend. CPUBackend, CUDABackend (Phase 1.5), and HIPBackend (Phase 1....
Definition device_backend.hpp:219
Adapts a pre-loaded MnistDataset (MnistIdxLoader::Load's output) onto the generic Dataset interface –...
Definition mnist_dataset_adapter.hpp:23
int64_t size() const override
Number of samples in this dataset.
Definition mnist_dataset_adapter.hpp:28
Sample get(int64_t index) const override
Definition mnist_dataset_adapter.hpp:37
MnistDatasetAdapter(MnistDataset dataset, DeviceBackend *backend)
Definition mnist_dataset_adapter.hpp:25
An N-dimensional shape. A plain aggregate of dimensions with no invariant beyond "non-negative dimens...
Definition shape.hpp:24
N-dimensional tensor. Owns its data buffer exclusively; a DeviceBackend* is injected (not owned) – th...
Definition tensor.hpp:29
Random-access dataset abstraction – Sample, Dataset (size()/get()).
Parses real MNIST IDX/ubyte files (fetched by tools/fetch_mnist.py) into Tensor images and integer la...
Definition acquisition_functions.hpp:16
One IDX file pair's contents: parallel images/labels, same length.
Definition mnist_loader.hpp:18
std::vector< Tensor > images
Shape (1, 28, 28), pixel values normalized to [0,1].
Definition mnist_loader.hpp:20
std::vector< int64_t > labels
0-9 class index, one per image, same order/length as images.
Definition mnist_loader.hpp:22
One dataset sample: an ordered list of Tensor fields (e.g. {features, label} or {image,...
Definition dataset.hpp:19