Softmax + Categorical Crossentropy
I fixed the activation+loss function, you can't select it directly, it automaticly uses it if it can. I also fixed one-hot generator. Still haven't tested the activation + loss nor any other backward function. I'll do that when I get to the optimizers which is next.
This commit is contained in:
@@ -44,6 +44,11 @@ namespace panic{
|
||||
namespace neural_network{
|
||||
|
||||
|
||||
enum struct loss_type {
|
||||
unknown,
|
||||
categorical_crossentropy
|
||||
};
|
||||
|
||||
|
||||
/**
|
||||
* @brief Base loss for the rest of the neural network library to use
|
||||
@@ -58,7 +63,7 @@ namespace panic{
|
||||
struct loss{
|
||||
|
||||
/**
|
||||
* @brief Emphty vector to store sample losses
|
||||
* @brief Emphty vector to store sample losses for each sample in the batch.
|
||||
*
|
||||
*/
|
||||
panic::tensor::real_vector sample_losses;
|
||||
@@ -66,19 +71,15 @@ struct loss{
|
||||
/**
|
||||
* @brief Mean loss over the entire batch.
|
||||
*/
|
||||
panic::types::real_t data_loss;
|
||||
panic::types::real_t data_loss = 0;
|
||||
|
||||
/**
|
||||
* @brief Matrix for backwards pass
|
||||
* @brief Gradient with respect to the loss input.
|
||||
*
|
||||
* This will be used later during the backward pass.
|
||||
*/
|
||||
panic::tensor::real_matrix dinputs;
|
||||
|
||||
/**
|
||||
* @brief Matrix for output of loss function
|
||||
*/
|
||||
panic::tensor::real_matrix outputs;
|
||||
|
||||
|
||||
/**
|
||||
* @brief Default de-constructor
|
||||
*
|
||||
@@ -86,6 +87,16 @@ struct loss{
|
||||
virtual ~loss() = default;
|
||||
|
||||
|
||||
/**
|
||||
* @brief Virtual loss_type function for derivative losses
|
||||
*
|
||||
* @Note This returns the type of loss it is
|
||||
* unknown be default
|
||||
*/
|
||||
virtual loss_type get_type() const {
|
||||
return loss_type::unknown;
|
||||
}
|
||||
|
||||
/**
|
||||
* @brief Virtual forward function for derivative loss functions
|
||||
*
|
||||
@@ -146,7 +157,7 @@ struct loss{
|
||||
* @param y_true Vector of true label of data.
|
||||
*
|
||||
*/
|
||||
virtual bool calculate(
|
||||
bool calculate(
|
||||
const panic::tensor::real_matrix& y_pred,
|
||||
const panic::tensor::uint_vector& y_true);
|
||||
|
||||
@@ -157,7 +168,7 @@ struct loss{
|
||||
* @param y_true Matrix of true label of data.
|
||||
*
|
||||
*/
|
||||
virtual bool calculate(
|
||||
bool calculate(
|
||||
const panic::tensor::real_matrix& y_pred,
|
||||
const panic::tensor::real_matrix& y_true);
|
||||
|
||||
|
||||
Reference in New Issue
Block a user