pulsatrix
Loading...
Searching...
No Matches
collate.hpp
Go to the documentation of this file.
1
5#pragma once
6
7#include <functional>
8#include <vector>
9
10#include "pulsatrix/dataset.hpp"
12#include "pulsatrix/tensor.hpp"
13
14namespace pulsatrix {
15
17struct Batch {
18 std::vector<Tensor> fields;
19
21 [[nodiscard]] int64_t size() const { return fields.empty() ? 0 : fields[0].shape().dim(0); }
22};
23
30using CollateFn = std::function<Batch(std::vector<Sample>, DeviceBackend*)>;
31
43[[nodiscard]] Batch DefaultCollate(std::vector<Sample> samples, DeviceBackend* backend);
44
45} // namespace pulsatrix
Vendor-agnostic compute/memory backend. CPUBackend, CUDABackend (Phase 1.5), and HIPBackend (Phase 1....
Definition device_backend.hpp:219
Random-access dataset abstraction – Sample, Dataset (size()/get()).
Abstract interface isolating vendor-specific memory/compute operations from Tensor/ComputationGraph.
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
One collated batch: one stacked Tensor per Sample field position.
Definition collate.hpp:17
int64_t size() const
Number of samples in this batch – fields[0]'s leading dimension.
Definition collate.hpp:21
std::vector< Tensor > fields
Definition collate.hpp:18
N-dimensional tensor – owns a buffer via DeviceBackend*, RAII (Rule of Five).