pulsatrix
Loading...
Searching...
No Matches
cross_entropy_loss.hpp
Go to the documentation of this file.
1
5#pragma once
6
7#include <cstdint>
8
10#include "pulsatrix/tensor.hpp"
11
12namespace pulsatrix {
13
27public:
32 explicit CrossEntropyLoss(DeviceBackend* backend);
33
46 [[nodiscard]] float forward(const Tensor& logits, int64_t target_class);
47
53 [[nodiscard]] Tensor backward() const;
54
55private:
56 DeviceBackend* backend_;
57 Tensor softmax_probs_;
58 int64_t target_class_ = 0;
59};
60
61} // namespace pulsatrix
loss = -log(softmax(logits)[target_class]), combined for numerical stability (subtract the max logit ...
Definition cross_entropy_loss.hpp:26
CrossEntropyLoss(DeviceBackend *backend)
Constructs a cross-entropy loss.
Tensor backward() const
Computes the gradient w.r.t. the logits: softmax(logits) - one_hot(target_class).
float forward(const Tensor &logits, int64_t target_class)
Computes the loss value and caches softmax probabilities/target for backward().
Vendor-agnostic compute/memory backend. CPUBackend, CUDABackend (Phase 1.5), and HIPBackend (Phase 1....
Definition device_backend.hpp:219
N-dimensional tensor. Owns its data buffer exclusively; a DeviceBackend* is injected (not owned) – th...
Definition tensor.hpp:29
Abstract interface isolating vendor-specific memory/compute operations from Tensor/ComputationGraph.
Definition acquisition_functions.hpp:16
N-dimensional tensor – owns a buffer via DeviceBackend*, RAII (Rule of Five).