pulsatrix
Loading...
Searching...
No Matches
adam_optimizer.hpp
Go to the documentation of this file.
1
5#pragma once
6
7#include <unordered_map>
8
10
11namespace pulsatrix {
12
21public:
30 explicit AdamOptimizer(float learning_rate, DeviceBackend* backend, float beta1 = 0.9f, float beta2 = 0.999f,
31 float eps = 1e-8f);
32
41 void step(Module& module);
42
47 void zero_grad(Module& module);
48
50 [[nodiscard]] float learning_rate() const { return learning_rate_; }
51
61 void set_learning_rate(float learning_rate) { learning_rate_ = learning_rate; }
62
63private:
64 struct AdamState {
65 Tensor m;
66 Tensor v;
67 int64_t t = 0;
68 };
69
70 float learning_rate_;
71 DeviceBackend* backend_;
72 float beta1_;
73 float beta2_;
74 float eps_;
75 std::unordered_map<const Tensor*, AdamState> state_;
76};
77
78} // namespace pulsatrix
Adam (Kingma & Ba, 2015): per-parameter moving averages of gradient (m) and squared gradient (v),...
Definition adam_optimizer.hpp:20
void step(Module &module)
Applies one Adam update to every parameter the module exposes.
void zero_grad(Module &module)
Resets every parameter's gradient to zero. Does not reset Adam's moment state.
void set_learning_rate(float learning_rate)
Overwrites the step size used by every subsequent step() call – necessary infrastructure for any mid-...
Definition adam_optimizer.hpp:61
AdamOptimizer(float learning_rate, DeviceBackend *backend, float beta1=0.9f, float beta2=0.999f, float eps=1e-8f)
Constructs an Adam optimizer.
float learning_rate() const
Current step size.
Definition adam_optimizer.hpp:50
Vendor-agnostic compute/memory backend. CPUBackend, CUDABackend (Phase 1.5), and HIPBackend (Phase 1....
Definition device_backend.hpp:219
Base class for every layer type (LinearModule, Conv2DModule, activations, ...).
Definition module.hpp:58
N-dimensional tensor. Owns its data buffer exclusively; a DeviceBackend* is injected (not owned) – th...
Definition tensor.hpp:29
Abstract base every layer subclasses – NVI forward(), pure-virtual LRP contract.
Definition acquisition_functions.hpp:16