A linear probe: LinearModule(activation_dim, 1) + BCEWithLogitsLoss, trained on (activation, binary concept label) pairs. High post-training accuracy means the concept is linearly decodable from those activations; chance-level accuracy means it is not (at least not linearly).
More...
#include <linear_probe.hpp>
A linear probe: LinearModule(activation_dim, 1) + BCEWithLogitsLoss, trained on (activation, binary concept label) pairs. High post-training accuracy means the concept is linearly decodable from those activations; chance-level accuracy means it is not (at least not linearly).
The activations come from an ActivationSnapshot (this campaign's Phase 1 output) in production use, but nothing here depends on that: a probe consumes a plain (N, activation_dim) batch, so it is equally usable against hand-built synthetic data – which is exactly what the positive/negative-control verification in tests/linear_probe_test.cpp needs.
- Note
- Mirrors XorNetwork/MnistConvNet's established shape rather than inventing a training abstraction: constructor takes a backend,
train_step() drives one loss/backward/optimizer-step cycle, and the epoch loop is the caller's. The probe deliberately does not own its optimizer – same division of responsibility XorNetwork::train_step already uses, which is what lets a caller choose SGD vs. Adam and keep one optimizer's state across a whole training run.
-
train_step is a template on the optimizer type rather than taking an Optimizer&: this codebase has no Optimizer base class – SGDOptimizer and AdamOptimizer are unrelated types sharing only the step(Module&) / zero_grad(Module&) shape (a compile-time, duck-typed contract). Introducing an abstract base to give this one method a runtime-polymorphic parameter would be a change to two closed missions' classes for no behavioral gain; the template keeps the probe usable with either optimizer and leaves both untouched.
-
Small seeded random weight init, not zero init. Unlike XorNetwork there is no symmetry to break here – a single linear layer does receive non-zero gradients from all-zero weights (
grad_W = X^T (sigmoid(0) - y)), so zero init would train fine. The reason is accuracy(): with all-zero weight and bias every logit is exactly 0, so every example's predicted probability is exactly 0.5, landing the entire batch precisely on the decision threshold. Small random init removes that degenerate tie at step 0 while keeping the probe fully reproducible.
◆ LinearProbe() [1/2]
| pulsatrix::LinearProbe::LinearProbe |
( |
int64_t |
activation_dim, |
|
|
DeviceBackend * |
backend |
|
) |
| |
|
inline |
Constructs a probe over activations of a given dimension.
- Parameters
-
| activation_dim | Width of the activation vectors this probe reads. Must be > 0. |
| backend | Backend to allocate/compute through. Not owned; must outlive this probe. |
| seed | RNG seed for weight initialization – same seed gives the same probe. |
- Exceptions
-
| std::invalid_argument | if activation_dim <= 0. External boundary: a probe's dimension is caller-supplied (typically read off a snapshot tensor's shape), and nothing downstream rejects it – LinearModule(0, 1, ...) builds a well-formed zero-element weight and fails only later, confusingly. |
Seeded from the global seed stream (next_seed(), FND-7).
◆ LinearProbe() [2/2]
| pulsatrix::LinearProbe::LinearProbe |
( |
int64_t |
activation_dim, |
|
|
DeviceBackend * |
backend, |
|
|
unsigned |
seed |
|
) |
| |
|
inline |
◆ accuracy()
| float pulsatrix::LinearProbe::accuracy |
( |
const Tensor & |
activation_batch, |
|
|
const Tensor & |
label_batch |
|
) |
| const |
|
inline |
Fraction of the batch the probe classifies correctly.
- Parameters
-
| activation_batch | Shape (N, activation_dim), N > 0. |
| label_batch | Shape (N, 1), values in {0, 1}, same N. |
- Returns
- Correct predictions / N, in [0, 1].
- Exceptions
-
| std::invalid_argument | on any malformed batch – see validate_batch(). |
- Note
- Threshold convention: predicted class is 1 iff
sigmoid(logit) >= 0.5, i.e. iff logit >= 0. Evaluated on the logit directly – the sigmoid is monotonic, so the comparison is exactly equivalent and avoids an unnecessary exp() plus the rounding question of whether sigmoid(0) lands on 0.5f exactly. A label is read as positive iff it is >= 0.5, the same threshold, so soft/smoothed targets score sensibly rather than counting as neither class.
-
const although LinearModule::forward() is not: the probe's logical state is its learned parameters, which a forward-only scoring pass cannot change. The classifier is mutable purely so forward()'s internal input/pre-bias caches (backward()'s working state, not the probe's observable value) can be written.
-
Raw host loop over
Tensor::data() – the same convention every loss class in this codebase already uses internally, since Tensor has no elementwise comparison or reduction primitive. Adding one is out of this mission's scope.
◆ activation_dim()
| int64_t pulsatrix::LinearProbe::activation_dim |
( |
| ) |
const |
|
inline |
Width of the activation vectors this probe reads.
◆ classifier() [1/2]
The underlying linear classifier – inspection (learned weights are the whole point of a probe: their direction is the decoded concept vector) and test-time weight injection.
◆ classifier() [2/2]
| const LinearModule & pulsatrix::LinearProbe::classifier |
( |
| ) |
const |
|
inline |
◆ train_step()
template<typename OptimizerT >
| float pulsatrix::LinearProbe::train_step |
( |
const Tensor & |
activation_batch, |
|
|
const Tensor & |
label_batch, |
|
|
OptimizerT & |
optimizer |
|
) |
| |
|
inline |
Runs one training step: forward, BCE-with-logits loss, backward, one optimizer update of the probe's weight/bias.
- Template Parameters
-
- Parameters
-
| activation_batch | Shape (N, activation_dim), N > 0. |
| label_batch | Shape (N, 1), values in {0, 1}, same N. |
| optimizer | Optimizer to update this probe's parameters with. Not owned. |
- Returns
- The loss for this batch, measured before the update (mirroring XorNetwork::train_step's return contract).
- Exceptions
-
| std::invalid_argument | on any malformed batch – see validate_batch(). |
- Note
- Zeroes gradients before accumulating, for the exact reason XorNetwork:: train_step documents: this codebase's modules accumulate parameter gradients across backward() calls until something resets them, so without this every step would silently sum onto every previous step's gradient.
The documentation for this class was generated from the following file: