pulsatrix
Loading...
Searching...
No Matches
linear_probe.hpp
Go to the documentation of this file.
1
6#pragma once
7
8#include <cmath>
9#include <cstdint>
10#include <random>
11#include <stdexcept>
12#include <string>
13#include <vector>
14
18
19namespace pulsatrix {
20
55public:
68 : LinearProbe(activation_dim, backend, static_cast<unsigned>(next_seed())) {}
69
70 LinearProbe(int64_t activation_dim, DeviceBackend* backend, unsigned seed)
71 : activation_dim_(activation_dim),
72 // Clamped only so a rejected dimension can't reach LinearModule/Shape and throw
73 // *their* message before the check below throws this class's own, clearer one --
74 // member initialization necessarily runs before the constructor body.
75 classifier_(activation_dim > 0 ? activation_dim : 1, 1, backend),
76 loss_(backend) {
77 if (activation_dim <= 0) {
78 throw std::invalid_argument("LinearProbe: activation_dim must be positive");
79 }
80 std::mt19937 rng(seed);
81 std::vector<float> weights(static_cast<size_t>(activation_dim));
82 for (float& w : weights) {
83 w = uniform_symmetric(rng, 0.01f);
84 }
85 classifier_.set_weight(weights);
86 classifier_.set_bias(std::vector<float>{uniform_symmetric(rng, 0.01f)});
87 }
88
106 template <typename OptimizerT>
107 float train_step(const Tensor& activation_batch, const Tensor& label_batch, OptimizerT& optimizer) {
108 validate_batch(activation_batch, label_batch, "LinearProbe::train_step");
109
110 optimizer.zero_grad(classifier_);
111
112 Tensor logits = classifier_.forward(activation_batch);
113 const float loss_value = loss_.forward(logits, label_batch);
114
115 Tensor grad_logits = loss_.backward();
116 (void)classifier_.backward(grad_logits); // the activations have no upstream to receive this
117
118 optimizer.step(classifier_);
119 return loss_value;
120 }
121
142 [[nodiscard]] float accuracy(const Tensor& activation_batch, const Tensor& label_batch) const {
143 validate_batch(activation_batch, label_batch, "LinearProbe::accuracy");
144
145 const Tensor logits = classifier_.forward(activation_batch);
146 const int64_t n = logits.numel();
147 int64_t correct = 0;
148 for (int64_t i = 0; i < n; ++i) {
149 const bool predicted_positive = logits.data()[i] >= 0.0f;
150 const bool labeled_positive = label_batch.data()[i] >= 0.5f;
151 if (predicted_positive == labeled_positive) {
152 ++correct;
153 }
154 }
155 return static_cast<float>(correct) / static_cast<float>(n);
156 }
157
159 [[nodiscard]] int64_t activation_dim() const { return activation_dim_; }
160
166 [[nodiscard]] LinearModule& classifier() { return classifier_; }
167
169 [[nodiscard]] const LinearModule& classifier() const { return classifier_; }
170
171private:
173 static float uniform_symmetric(std::mt19937& rng, float scale) {
174 // std::uniform_real_distribution's mapping is implementation-defined, so identical
175 // seeds would not give identical weights across standard libraries. This mapping is
176 // fixed here, making a seeded probe reproducible everywhere.
177 const float unit = static_cast<float>(rng() - std::mt19937::min()) /
178 static_cast<float>(std::mt19937::max() - std::mt19937::min() + 1ull);
179 return (unit * 2.0f - 1.0f) * scale;
180 }
181
197 void validate_batch(const Tensor& activation_batch, const Tensor& label_batch, const char* method) const {
198 const std::string where(method);
199 if (activation_batch.rank() != 2) {
200 throw std::invalid_argument(where + ": activation_batch must be rank-2 (N, activation_dim)");
201 }
202 if (label_batch.rank() != 2 || label_batch.shape().dim(1) != 1) {
203 throw std::invalid_argument(where + ": label_batch must be rank-2 (N, 1)");
204 }
205 if (activation_batch.shape().dim(1) != activation_dim_) {
206 throw std::invalid_argument(where + ": activation_batch width must equal activation_dim");
207 }
208 if (activation_batch.shape().dim(0) != label_batch.shape().dim(0)) {
209 throw std::invalid_argument(where + ": activation_batch and label_batch must have the same batch size");
210 }
211 if (activation_batch.shape().dim(0) <= 0) {
212 throw std::invalid_argument(where + ": batch must not be empty");
213 }
214 }
215
216 int64_t activation_dim_;
217 mutable LinearModule classifier_;
218 BCEWithLogitsLoss loss_;
219};
220
221} // namespace pulsatrix
Binary cross-entropy on raw logits – combined sigmoid + BCE, numerically stable.
Tensor backward() const
Gradient w.r.t. the logits: grad[i] = (sigmoid(x[i]) - y[i]) / numel.
float forward(const Tensor &logits, const Tensor &target)
Computes the loss value and caches logits/target for backward().
Vendor-agnostic compute/memory backend. CPUBackend, CUDABackend (Phase 1.5), and HIPBackend (Phase 1....
Definition device_backend.hpp:219
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
Tensor backward(const Tensor &grad_output) override
Computes the gradient w.r.t. this module's input, and accumulates the weight/bias gradients internall...
void set_bias(std::initializer_list< float > values)
Overwrites the bias buffer – test/initialization use only.
void set_weight(std::initializer_list< float > values)
Overwrites the weight buffer – test/initialization use only.
A linear probe: LinearModule(activation_dim, 1) + BCEWithLogitsLoss, trained on (activation,...
Definition linear_probe.hpp:54
LinearProbe(int64_t activation_dim, DeviceBackend *backend)
Constructs a probe over activations of a given dimension.
Definition linear_probe.hpp:67
int64_t activation_dim() const
Width of the activation vectors this probe reads.
Definition linear_probe.hpp:159
float train_step(const Tensor &activation_batch, const Tensor &label_batch, OptimizerT &optimizer)
Runs one training step: forward, BCE-with-logits loss, backward, one optimizer update of the probe's ...
Definition linear_probe.hpp:107
const LinearModule & classifier() const
Const overload of classifier().
Definition linear_probe.hpp:169
LinearProbe(int64_t activation_dim, DeviceBackend *backend, unsigned seed)
Definition linear_probe.hpp:70
LinearModule & classifier()
The underlying linear classifier – inspection (learned weights are the whole point of a probe: their ...
Definition linear_probe.hpp:166
float accuracy(const Tensor &activation_batch, const Tensor &label_batch) const
Fraction of the batch the probe classifies correctly.
Definition linear_probe.hpp:142
Tensor forward(const Tensor &input)
Runs this module's forward computation.
Definition module.hpp:73
N-dimensional tensor. Owns its data buffer exclusively; a DeviceBackend* is injected (not owned) – th...
Definition tensor.hpp:29
int64_t numel() const
Total element count – shape().numel().
Definition tensor.hpp:116
const float * data() const
Raw buffer access. nullptr iff numel() == 0.
Definition tensor.hpp:167
One global seed for everything that isn't given its own, and a deterministic mode that forbids nondet...
Dense/fully-connected layer – the reference Module implementation.
Definition acquisition_functions.hpp:16
uint64_t next_seed()
The next seed in the global stream: a distinct, well-mixed 64-bit value per call, reproducible for a ...