pulsatrix
Loading...
Searching...
No Matches
sgd_optimizer.hpp
Go to the documentation of this file.
1
5#pragma once
6
8
9namespace pulsatrix {
10
13public:
18 explicit SGDOptimizer(float learning_rate) : learning_rate_(learning_rate) {}
19
25 void step(Module& module);
26
31 void zero_grad(Module& module);
32
33private:
34 float learning_rate_;
35};
36
37} // namespace pulsatrix
Base class for every layer type (LinearModule, Conv2DModule, activations, ...).
Definition module.hpp:58
param -= learning_rate * grad, per parameter, for every parameter a Module exposes.
Definition sgd_optimizer.hpp:12
SGDOptimizer(float learning_rate)
Constructs an SGD optimizer.
Definition sgd_optimizer.hpp:18
void step(Module &module)
Applies one SGD update to every parameter the module exposes.
void zero_grad(Module &module)
Resets every parameter's gradient to zero.
Abstract base every layer subclasses – NVI forward(), pure-virtual LRP contract.
Definition acquisition_functions.hpp:16