pulsatrix
Loading...
Searching...
No Matches
mnist_classifier_example.hpp
Go to the documentation of this file.
1
7#pragma once
8
9#include <vector>
10
18
19namespace pulsatrix {
20
35public:
41 explicit MnistConvNet(DeviceBackend* backend, unsigned seed = 42);
42
48 [[nodiscard]] Tensor forward(const Tensor& image);
49
55 [[nodiscard]] int64_t predict(const Tensor& image);
56
67 float train_step(const Tensor& image, int64_t target_class, AdamOptimizer& optimizer, MetricsSink& sink,
68 int step);
69
71 [[nodiscard]] const Tensor& classifier_weight() const { return classifier_.weight(); }
72
80 [[nodiscard]] std::vector<Module*> modules() { return {&conv_, &relu_, &flatten_, &classifier_}; }
81
82private:
83 DeviceBackend* backend_;
84 Conv2DModule conv_;
85 ReluModule relu_;
86 FlattenModule flatten_;
87 LinearModule classifier_;
88 CrossEntropyLoss loss_;
89};
90
91} // namespace pulsatrix
Adam optimizer – operates uniformly across any Module's parameters().
Adam (Kingma & Ba, 2015): per-parameter moving averages of gradient (m) and squared gradient (v),...
Definition adam_optimizer.hpp:20
2D convolution, batched (input/output are rank-4: N x channels x H x W) – migrated from the original ...
Definition conv2d_module.hpp:31
loss = -log(softmax(logits)[target_class]), combined for numerical stability (subtract the max logit ...
Definition cross_entropy_loss.hpp:26
Vendor-agnostic compute/memory backend. CPUBackend, CUDABackend (Phase 1.5), and HIPBackend (Phase 1....
Definition device_backend.hpp:219
y = reshape(x, {N, x.numel()/N}), N = x.shape().dim(0). No parameters, no gradient math beyond reshap...
Definition flatten_module.hpp:28
y = x @ W + b, batched (x is (N, in_features), y is (N, out_features)) – migrated from the original u...
Definition linear_module.hpp:37
const Tensor & weight() const
Definition linear_module.hpp:91
Interface the training loop logs scalars/histograms through. Concrete writers (TensorBoard event form...
Definition metrics_sink.hpp:22
Conv2D(1,8,5,5) -> ReLU -> Flatten -> Linear(4608,10), trained via CrossEntropyLoss + Adam,...
Definition mnist_classifier_example.hpp:34
int64_t predict(const Tensor &image)
forward() plus argmax – the predicted class index.
const Tensor & classifier_weight() const
Test/inspection accessor.
Definition mnist_classifier_example.hpp:71
std::vector< Module * > modules()
The network's layers in forward order (Conv2D, ReLU, Flatten, Linear), for building an ExplainerConte...
Definition mnist_classifier_example.hpp:80
Tensor forward(const Tensor &image)
Runs the network forward.
float train_step(const Tensor &image, int64_t target_class, AdamOptimizer &optimizer, MetricsSink &sink, int step)
Runs one training step: forward, cross-entropy loss, backward through every layer,...
MnistConvNet(DeviceBackend *backend, unsigned seed=42)
Constructs the network with randomly initialized weights.
y = max(x, 0), elementwise. No parameters, no parameter gradients.
Definition relu_module.hpp:15
N-dimensional tensor. Owns its data buffer exclusively; a DeviceBackend* is injected (not owned) – th...
Definition tensor.hpp:29
2D convolution – implemented via im2col + DeviceBackend::gemm (no new backend primitive).
Softmax + negative log-likelihood classification loss.
Reshape-only Module – flattens every non-batch dim of a (N, ...) input to (N, flattened_features),...
Dense/fully-connected layer – the reference Module implementation.
Keeps monitoring/visualization tools out of the training core – same OCP/DIP pattern as DeviceBackend...
Definition acquisition_functions.hpp:16
ReLU activation – the second Module subclass, following LinearModule's pattern.