78 throw std::invalid_argument(
"LinearProbe: activation_dim must be positive");
80 std::mt19937 rng(seed);
82 for (
float& w : weights) {
83 w = uniform_symmetric(rng, 0.01f);
86 classifier_.
set_bias(std::vector<float>{uniform_symmetric(rng, 0.01f)});
106 template <
typename OptimizerT>
108 validate_batch(activation_batch, label_batch,
"LinearProbe::train_step");
110 optimizer.zero_grad(classifier_);
113 const float loss_value = loss_.
forward(logits, label_batch);
116 (void)classifier_.
backward(grad_logits);
118 optimizer.step(classifier_);
143 validate_batch(activation_batch, label_batch,
"LinearProbe::accuracy");
146 const int64_t n = logits.
numel();
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) {
155 return static_cast<float>(correct) /
static_cast<float>(n);
173 static float uniform_symmetric(std::mt19937& rng,
float scale) {
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;
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)");
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)");
205 if (activation_batch.shape().dim(1) != activation_dim_) {
206 throw std::invalid_argument(where +
": activation_batch width must equal activation_dim");
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");
211 if (activation_batch.shape().dim(0) <= 0) {
212 throw std::invalid_argument(where +
": batch must not be empty");
216 int64_t activation_dim_;
217 mutable LinearModule classifier_;
218 BCEWithLogitsLoss loss_;
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 ...