pulsatrix
Loading...
Searching...
No Matches
tensor.hpp
Go to the documentation of this file.
1
5#pragma once
6
7#include <string>
8#include <stdexcept>
9#include <initializer_list>
10#include <vector>
11
12#include "pulsatrix/assert.hpp"
15#include "pulsatrix/shape.hpp"
16
17namespace pulsatrix {
18
29class Tensor {
30public:
37
47
58 Tensor(Shape shape, DeviceBackend* backend, std::initializer_list<float> values, DeviceType device);
59
61 Tensor(Shape shape, DeviceBackend* backend, std::initializer_list<float> values);
62
76 Tensor(Shape shape, DeviceBackend* backend, const std::vector<float>& values, DeviceType device);
77
79 Tensor(Shape shape, DeviceBackend* backend, const std::vector<float>& values);
80
82
99 [[nodiscard]] static Tensor Stack(const std::vector<Tensor>& tensors, DeviceBackend* backend);
100
107 Tensor(const Tensor& other);
108 Tensor& operator=(const Tensor& other);
109 Tensor(Tensor&& other) noexcept;
110 Tensor& operator=(Tensor&& other) noexcept;
111
113 [[nodiscard]] const Shape& shape() const { return shape_; }
114
116 [[nodiscard]] int64_t numel() const { return shape_.numel(); }
117
119 [[nodiscard]] int64_t rank() const { return shape_.rank(); }
120
122 [[nodiscard]] DeviceType device() const { return device_; }
123
130 [[nodiscard]] DeviceBackend* backend() const { return backend_; }
131
141 [[nodiscard]] bool requires_grad() const { return requires_grad_; }
142
144 void set_requires_grad(bool requires_grad) { requires_grad_ = requires_grad; }
145
151 [[nodiscard]] float read_element(int64_t flat_index) const;
152
161 [[nodiscard]] std::vector<float> to_host_vector() const;
162
164 void write_element(int64_t flat_index, float value);
165
167 [[nodiscard]] const float* data() const { return data_; }
168
170 [[nodiscard]] float* data() { return data_; }
171
182 [[nodiscard]] float& at(std::initializer_list<int64_t> index);
183
185 [[nodiscard]] const float& at(std::initializer_list<int64_t> index) const;
186
195 [[nodiscard]] float& operator[](int64_t flat_index) {
196 PULSATRIX_ASSERT(flat_index >= 0 && flat_index < numel());
197 return data_[flat_index];
198 }
199
201 [[nodiscard]] const float& operator[](int64_t flat_index) const {
202 PULSATRIX_ASSERT(flat_index >= 0 && flat_index < numel());
203 return data_[flat_index];
204 }
205
211 Tensor& fill(float value);
212
223 Tensor& accumulate(const Tensor& other);
224
231 Tensor& reshape(Shape new_shape);
232
244
265 Tensor& to(DeviceType target, DeviceBackend* target_backend);
266
267private:
268 [[nodiscard]] int64_t flat_index_of(std::initializer_list<int64_t> index) const;
269
270 float* data_;
271 Shape shape_;
272 DeviceBackend* backend_;
273 DeviceType device_;
274 bool requires_grad_ = true;
275};
276
285inline void require_device(const Tensor& t, DeviceType expected, const char* where) {
286 if (t.device() != expected) {
287 auto name = [](DeviceType d) {
288 switch (d) {
289 case DeviceType::Cpu:
290 return "Cpu";
291 case DeviceType::Cuda:
292 return "Cuda";
293 case DeviceType::Hip:
294 return "Hip";
295 }
296 return "unknown";
297 };
298 throw std::invalid_argument(std::string(where) + ": tensor is on " + name(t.device()) +
299 " but this computes on " + name(expected) + "; move it first with Tensor::to()");
300 }
301}
302
303} // namespace pulsatrix
PULSATRIX_ASSERT – debug-only invariant check for programmer errors, distinct from throw (used for ca...
#define PULSATRIX_ASSERT(cond)
Aborts with a diagnostic message if cond is false. Debug-only – use for conditions that indicate a bu...
Definition assert.hpp:22
Vendor-agnostic compute/memory backend. CPUBackend, CUDABackend (Phase 1.5), and HIPBackend (Phase 1....
Definition device_backend.hpp:219
An N-dimensional shape. A plain aggregate of dimensions with no invariant beyond "non-negative dimens...
Definition shape.hpp:24
int64_t numel() const
Total element count – the product of all dimensions.
Definition shape.hpp:74
int64_t rank() const
Number of dimensions. 0 for a scalar.
Definition shape.hpp:62
N-dimensional tensor. Owns its data buffer exclusively; a DeviceBackend* is injected (not owned) – th...
Definition tensor.hpp:29
Tensor & reshape(Shape new_shape)
Reinterprets this tensor's dimensions in place – same buffer, new shape.
Tensor(const Tensor &other)
Deep-copies another tensor's buffer.
Tensor & to(DeviceType target, DeviceBackend *target_backend)
Moves this tensor's buffer to another device, owned by target_backend.
static Tensor Stack(const std::vector< Tensor > &tensors, DeviceBackend *backend)
Concatenates N tensors along their leading dimension into one batch Tensor – pulsatrix's collate-time...
const float & at(std::initializer_list< int64_t > index) const
Const overload of at().
Tensor(Shape shape, DeviceBackend *backend, DeviceType device)
Constructs a zero-initialized tensor with an explicit device tag.
DeviceType device() const
Which device this tensor's buffer conceptually resides on.
Definition tensor.hpp:122
Tensor & fill(float value)
Sets every element to value. Safe no-op on a zero-element tensor.
Tensor(Tensor &&other) noexcept
Tensor(Shape shape, DeviceBackend *backend, const std::vector< float > &values)
As above, tagged with backend->device().
Tensor & operator=(Tensor &&other) noexcept
Tensor & operator=(const Tensor &other)
int64_t rank() const
Number of dimensions – shape().rank().
Definition tensor.hpp:119
DeviceBackend * backend() const
The backend that owns this tensor's buffer. Not owned by the Tensor.
Definition tensor.hpp:130
bool requires_grad() const
Whether a module's backward() should accumulate a gradient for this tensor when it is a parameter,...
Definition tensor.hpp:141
Tensor & accumulate(const Tensor &other)
In-place elementwise accumulation: this[i] += other[i] for every element.
Tensor(Shape shape, DeviceBackend *backend, std::initializer_list< float > values, DeviceType device)
Constructs a tensor from explicit values.
int64_t numel() const
Total element count – shape().numel().
Definition tensor.hpp:116
const float & operator[](int64_t flat_index) const
Const overload of operator[].
Definition tensor.hpp:201
const Shape & shape() const
This tensor's shape.
Definition tensor.hpp:113
float read_element(int64_t flat_index) const
Reads one element to the host through the owning backend, on any device.
void set_requires_grad(bool requires_grad)
Sets requires_grad(); see there.
Definition tensor.hpp:144
Tensor(Shape shape, DeviceBackend *backend, std::initializer_list< float > values)
As above, tagged with backend->device().
float * data()
Raw buffer access (mutable). nullptr iff numel() == 0.
Definition tensor.hpp:170
float & at(std::initializer_list< int64_t > index)
Element access by multi-dimensional index (row-major).
float & operator[](int64_t flat_index)
Flat (rank-agnostic) element access by linear offset into the row-major buffer.
Definition tensor.hpp:195
Tensor & to(DeviceType target)
Same-device no-op form of to().
void write_element(int64_t flat_index, float value)
Writes one element from the host through the owning backend, on any device.
Tensor(Shape shape, DeviceBackend *backend, const std::vector< float > &values, DeviceType device)
Constructs a tensor from explicit values, runtime-sized source.
Tensor(Shape shape, DeviceBackend *backend)
Constructs a zero-initialized tensor on backend's own device (backend->device()).
std::vector< float > to_host_vector() const
Copies the whole buffer to a host vector through the owning backend, on any device.
const float * data() const
Raw buffer access. nullptr iff numel() == 0.
Definition tensor.hpp:167
Abstract interface isolating vendor-specific memory/compute operations from Tensor/ComputationGraph.
PULSATRIX_REQUIRE_HOST – always-on guard for code paths that dereference Tensor::data() on the host.
Definition acquisition_functions.hpp:16
DeviceType
Which physical device a Tensor's buffer resides on.
Definition device_backend.hpp:17
void require_device(const Tensor &t, DeviceType expected, const char *where)
Throws unless t lives on expected – the check every module and loss runs on the tensors handed to it ...
Definition tensor.hpp:285
Tensor dimension arithmetic – rank, element count, per-dimension access.