loss = -log(softmax(logits)[target_class]), combined for numerical stability (subtract the max logit before exponentiating) rather than computing softmax and log separately.
More...
#include <cross_entropy_loss.hpp>
loss = -log(softmax(logits)[target_class]), combined for numerical stability (subtract the max logit before exponentiating) rather than computing softmax and log separately.
- Note
- Not a Module subclass, same as MSELoss – losses are the seed point relevance/ gradient propagation starts from, not something propagate_relevance is defined for.
-
Takes an integer class index, not a one-hot Tensor – the standard classification- loss convention (mirrors torch.nn.CrossEntropyLoss's (logits, target) signature), and avoids inventing a one-hot Tensor construction step every caller would otherwise need.
◆ CrossEntropyLoss()
| pulsatrix::CrossEntropyLoss::CrossEntropyLoss |
( |
DeviceBackend * |
backend | ) |
|
|
explicit |
Constructs a cross-entropy loss.
- Parameters
-
| backend | Backend to compute through. Not owned; must outlive this loss. |
◆ backward()
| Tensor pulsatrix::CrossEntropyLoss::backward |
( |
| ) |
const |
Computes the gradient w.r.t. the logits: softmax(logits) - one_hot(target_class).
- Returns
- Gradient tensor, same shape as the logits passed to forward().
- Note
- Must be called after forward() – uses the cached softmax probabilities.
◆ forward()
| float pulsatrix::CrossEntropyLoss::forward |
( |
const Tensor & |
logits, |
|
|
int64_t |
target_class |
|
) |
| |
Computes the loss value and caches softmax probabilities/target for backward().
- Parameters
-
| logits | Raw (pre-softmax) model output, shape (num_classes,). |
| target_class | Ground-truth class index, 0-based. Must be in [0, logits.numel()). |
- Returns
- The scalar cross-entropy loss.
- Note
- Device-generic (GPU-native-kernels Mission 1): only the loss scalar and the target logit cross to the host. target_class range is an PULSATRIX_ASSERT – an internal invariant for this loss's current (non-Python- bound) call sites, not yet a Python-reachable external boundary.
The documentation for this class was generated from the following file: